#!/usr/bin/env python3
"""
FireViewer + RFDETR-Face-Detection TensorRT Export Script
Downloads models from HuggingFace and exports to ONNX + TensorRT Engine + Triton config.

Usage:
  python3 export_fireviewer_rfdetr_face.py --all
  # or specific models:
  python3 export_fireviewer_rfdetr_face.py --model fireviewer,rfd-etr-face

This script handles:
1. Model download from HuggingFace (if not present)
2. ONNX export with opset 19 + dynamic batching
3. TensorRT engine generation (FP16, dynamic batch shapes)
4. Triton config.pbtxt generation (comment-free for protobuf compatibility)
5. Repository structure at /opt/contextra/models/triton/<model_name>/1/
"""

import os
import sys
import subprocess
import gc
import logging
import argparse
from pathlib import Path

logging.basicConfig(level=logging.INFO, format="%(levelname)s: %(message)s")
logger = logging.getLogger("rfdetr_export")

# Triton model repository root (matches docker-compose volumes mapping)
TRITON_BASE_DIR = "/opt/contextra/models"
TRITON_REPO_ROOT = "triton"

# Dynamic batch configuration
DYNAMIC_BATCH = {
    "min_batch": 1,
    "opt_batch": 4,
    "max_batch": 8,
}

# Model configurations for FireViewer and RFDETR-Face-Detection
MODEL_CONFIGS = {
    "fireviewer/rf-detr-small-ground-elite-fire-smoke-v1": {
        "type": "detection",
        "resolution": 512,  # DINOv3 백본: patch_size=16 * num_windows=2 = 32로 나누어 떨어져야 함 (504는 불가)
        "max_proposals": 300,
        "num_classes": 91,
        "class_name": "RFDETRSmall",  # May need adjustment based on actual model class name
    },
    "fireviewer/rf-detr-large-ground-fire-smoke-v2": {
        "type": "detection",
        "resolution": 704,  # Large ground model resolution for detection
        "max_proposals": 300,
        "num_classes": 91,
        "class_name": "RFDETRLarge",  # May need adjustment based on actual model class name
    },
    # RFDETR-Face-Detection model
    "Herojayjay/RFDETR-Face-Detection": {
        "type": "detection",
        "resolution": 504,  # Divisible by 56 (patch_size=14 windowed attention)
        "max_proposals": 300,
        "num_classes": 91,
        "class_name": "RFDETRFaceDetection",  # May need adjustment based on actual model class name
    },
    # Additional user-requested model (Awiros person+head, DINOv3 backbone)
    "Awiros/person_and_head_detection": {
        "type": "detection",
        "resolution": 512,  # DINOv3: patch_size=16 * num_windows=2 = 32로 나누어 떨어져야 함 (모델 카드 Table 4: RF-DETR-S @ 512)
        "max_proposals": 300,
        "num_classes": 2,  # person + head (CrowdHuman fine-tune)
        "class_name": "RFDETRSmall",
        # HF 저장소의 fine-tuned 체크포인트 (safetensors -> .pth 자동 변환 후 pretrain_weights로 로드)
        "hf_checkpoint": "ckpts/rf-detrs/checkpoint_best_total.safetensors",
        "class_names": ["person", "head"],
    },
}


def _q(s):
    """Helper to produce a quoted string for pbtxt."""
    return '"' + s + '"'


def remove_comments(pbtxt_content):
    """Remove comments from Triton config.pbtxt for protobuf parser compatibility."""
    lines = []
    for line in pbtxt_content.split("\n"):
        # '#'로 시작하는 라인은 전체 제거
        if line.strip().startswith("#"):
            continue
        # inline 주석 제거 (첫 번째 '#'까지만 자름)
        if "#" in line:
            line = line.split("#")[0]
        lines.append(line)
    return "\n".join(lines)


def generate_detection_config(model_name, input_size, max_proposals, num_classes):
    """Detection 모델 config.pbtxt 생성."""
    lines = []
    lines.append("name: " + _q(model_name))
    lines.append("platform: " + _q("tensorrt_plan"))
    lines.append("")
    lines.append("max_batch_size: {}".format(DYNAMIC_BATCH["max_batch"]))
    lines.append("")
    lines.append("input [")
    lines.append('  {"')
    lines.append('    name: ' + _q("input"))
    lines.append('    data_type: TYPE_FP32')
    lines.append('    dims: [ 3, {}, {} ]'.format(input_size, input_size))
    lines.append('  }')
    lines.append("]")
    lines.append("")
    lines.append("output [")
    lines.append('  {"')
    lines.append('    name: ' + _q("dets"))
    lines.append('    data_type: TYPE_FP32')
    lines.append('    dims: [ {}, -1 ]'.format(max_proposals))
    lines.append('  },')
    lines.append('  {"')
    lines.append('    name: ' + _q("labels"))
    lines.append('    data_type: TYPE_FP32')
    lines.append('    dims: [ {}, -1 ]'.format(max_proposals, num_classes))
    lines.append('  }')
    lines.append("]")
    lines.append("")
    return "\n".join(lines)


def generate_segmentation_config(model_name, input_size, max_proposals, num_classes):
    """Segmentation 모델 config.pbtxt 생성 (masks 출력 추가)."""
    lines = []
    lines.append("name: " + _q(model_name))
    lines.append("platform: " + _q("tensorrt_plan"))
    lines.append("")
    lines.append("max_batch_size: {}".format(DYNAMIC_BATCH["max_batch"]))
    lines.append("")
    lines.append("input [")
    lines.append('  {"')
    lines.append('    name: ' + _q("input"))
    lines.append('    data_type: TYPE_FP32')
    lines.append('    dims: [ 3, {}, {} ]'.format(input_size, input_size))
    lines.append('  }')
    lines.append("]")
    lines.append("")
    lines.append("output [")
    lines.append('  {"')
    lines.append('    name: ' + _q("dets"))
    lines.append('    data_type: TYPE_FP32')
    lines.append('    dims: [ {}, -1 ]'.format(max_proposals))
    lines.append('  },')
    lines.append('  {"')
    lines.append('    name: ' + _q("labels"))
    lines.append('    data_type: TYPE_FP32')
    lines.append('    dims: [ {}, -1 ]'.format(max_proposals, num_classes))
    lines.append('  },')
    lines.append('  {"')
    lines.append('    name: ' + _q("masks"))
    lines.append('    data_type: TYPE_FP32')
    # mask resolution = input_size / 4 (typical for segmentation heads)
    mask_res = input_size // 4
    lines.append('    dims: [ {}, {}, {} ]'.format(max_proposals, mask_res, mask_res))
    lines.append("  }")
    lines.append("]")
    lines.append("")
    return "\n".join(lines)


def download_model(repo_id, cache_dir="/home/tamasic/remote-contextra/models"):
    """HuggingFace 모델 다운로드."""
    from huggingface_hub import snapshot_download

    logger.info("Downloading {} ...".format(repo_id))
    try:
        local_dir = snapshot_download(
            repo_id=repo_id,
            cache_dir=cache_dir,
            resume_download=True,
        )
        logger.info("Downloaded to: {}".format(local_dir))
        return local_dir
    except Exception as e:
        logger.error("Failed to download {}: {}".format(repo_id, e))
        return None


def prepare_finetuned_weights(model_name, config, base_dir):
    """HF fine-tuned 체크포인트 다운로드 + rfdetr .pth 포맷 변환.

    Awiros 공식 convert_release_checkpoint.py와 동일한 포맷:
      {"model": state_dict, "args": Namespace(class_names=[...])}
    반환: 로컬 .pth 경로 (RFDETRSmall(pretrain_weights=...)로 로드 가능)
    """
    from huggingface_hub import hf_hub_download
    from argparse import Namespace
    import torch

    ckpt_filename = config["hf_checkpoint"]
    class_names = config.get("class_names", ["person", "head"])

    cache_dir = os.path.join(base_dir, "hf_cache", model_name.replace("/", "_"))
    out_pth = os.path.join(
        base_dir, "converted", model_name.replace("/", "_") + ".pth"
    )

    if os.path.exists(out_pth):
        logger.info("Converted checkpoint already exists: {}".format(out_pth))
        return out_pth

    # 1. HF에서 체크포인트 다운로드 (safetensors)
    ckpt_path = hf_hub_download(
        repo_id=model_name,
        filename=ckpt_filename,
        local_dir=cache_dir,
    )
    logger.info("Downloaded checkpoint: {}".format(ckpt_path))

    # 2. 이미 .pth면 그대로 사용
    if not str(ckpt_path).endswith(".safetensors"):
        return ckpt_path

    # 3. safetensors -> rfdetr .pth 포맷 변환
    from safetensors.torch import load_file

    state_dict = load_file(str(ckpt_path))
    payload = {"model": state_dict, "args": Namespace(class_names=class_names)}

    os.makedirs(os.path.dirname(out_pth), exist_ok=True)
    torch.save(payload, out_pth)
    logger.info(
        "Converted to rfdetr checkpoint format: {} (classes={})".format(
            out_pth, class_names
        )
    )
    return out_pth


def export_model(model_name, config, base_dir):
    """단일 모델 export (ONNX -> TensorRT -> Triton config)."""
    mtype = config["type"]
    resolution = config["resolution"]
    max_proposals = config["max_proposals"]
    num_classes = config["num_classes"]

    logger.info("=" * 60)
    logger.info(
        "Exporting {} ({} model, resolution={})".format(model_name, mtype, resolution)
    )
    logger.info("=" * 60)

    # Triton model repository 구조 생성
    triton_repo_dir = os.path.join(base_dir, TRITON_REPO_ROOT, model_name.replace("/", "_"))
    version_dir = os.path.join(triton_repo_dir, "1")
    os.makedirs(version_dir, exist_ok=True)

    try:
        # 1. rfdetr 라이브러리에서 모델 로드 및 ONNX export
        logger.info(
            "Loading RF-DETR model: {} (resolution={})".format(model_name, resolution)
        )

        # rfdetr import - models are already cached locally
        sys.path.insert(0, "/home/tamasic/remote-contextra/models")
        import rfdetr

        # 모델 클래스 동적 로드
        model_class_name = config.get("class_name", model_name.split("/")[-1])
        try:
            ModelClass = getattr(rfdetr, model_class_name)
        except AttributeError:
            # Try as-is or with RFDETR prefix
            try:
                ModelClass = getattr(rfdetr, "RFDETR" + model_class_name.replace("RFDETR", ""))
            except AttributeError:
                logger.warning(
                    "Could not find class {} in rfdetr, trying alternative".format(
                        model_class_name
                    )
                )
                # Last resort: try to import the model directly from path
                ModelClass = rfdetr.RFDETRBase  # fallback

        # fine-tuned 가중치 로드 (HF 체크포인트가 지정된 경우)
        if "hf_checkpoint" in config:
            weights_path = prepare_finetuned_weights(model_name, config, base_dir)
            model = ModelClass(
                resolution=resolution,
                pretrain_weights=str(weights_path),
            )
            logger.info("Loaded fine-tuned weights from {}".format(weights_path))
        else:
            model = ModelClass(resolution=resolution)
        logger.info("Model loaded: {} (num_classes={})".format(
            type(model).__name__, num_classes))

        # ONNX export
        onnx_file_path = model.export(
            output_dir=base_dir,
            format="onnx",
            opset_version=19,
            dynamic_batch=True,
            shape=(resolution, resolution),
            verbose=False,
        )
        onnx_path = str(onnx_file_path)
        logger.info("ONNX exported: {}".format(onnx_path))

        # 2. TensorRT 엔진 생성 (trtexec 사용)
        final_engine_path = os.path.join(version_dir, "model.plan")

        min_batch = DYNAMIC_BATCH["min_batch"]
        opt_batch = DYNAMIC_BATCH["opt_batch"]
        max_batch = DYNAMIC_BATCH["max_batch"]

        trtexec_cmd = [
            "/usr/bin/trtexec",
            "--onnx={}".format(onnx_path),
            "--saveEngine={}".format(final_engine_path),
            "--fp16",
            "--verbose",
            "--minShapes=input:{}x3x{}x{}".format(min_batch, resolution, resolution),
            "--optShapes=input:{}x3x{}x{}".format(opt_batch, resolution, resolution),
            "--maxShapes=input:{}x3x{}x{}".format(max_batch, resolution, resolution),
        ]

        # 모델별 메모리 설정
        if mtype == "segmentation":
            trtexec_cmd.append("--memPoolSize=workspace:2048")
        elif mtype == "detection":
            trtexec_cmd.append("--memPoolSize=workspace:1024")
        else:
            trtexec_cmd.append("--memPoolSize=workspace:1536")

        logger.info("Running trtexec... (this may take a while)")
        result = subprocess.run(
            trtexec_cmd,
            check=True,
            capture_output=True,
            text=True,
            timeout=900,  # 15분 타임아웃
        )
        logger.info("TensorRT engine created: {}".format(final_engine_path))

        # 3. config.pbtxt 생성 (주석 제거)
        if mtype == "segmentation":
            config_content = generate_segmentation_config(
                model_name, resolution, max_proposals, num_classes
            )
        else:
            config_content = generate_detection_config(
                model_name, resolution, max_proposals, num_classes
            )

        # 주석 제거 후 저장
        clean_config = remove_comments(config_content)
        config_path = os.path.join(triton_repo_dir, "config.pbtxt")
        with open(config_path, "w") as f:
            f.write(clean_config)
        logger.info("Config written to: {}".format(config_path))

        # 4. 완료 플래그 생성
        init_flag = os.path.join(base_dir, model_name.replace("/", "_"), ".initialized")
        os.makedirs(os.path.dirname(init_flag), exist_ok=True)
        with open(init_flag, "w") as f:
            f.write(
                "Exported at {} | TRT v10.16.00 | Triton 2.67.0\n".format(
                    __import__("time").time()
                )
            )

        logger.info("OK Successfully exported {}".format(model_name))
        return True

    except Exception as e:
        logger.error("X Failed to export {}: {}".format(model_name, e), exc_info=True)
        return False


def main():
    parser = argparse.ArgumentParser(
        description="FireViewer + RFDETR-Face-Detection TensorRT Export Pipeline"
    )
    parser.add_argument(
        "--all",
        action="store_true",
        help="Export all FireViewer and RFDETR-Face-Detection models",
    )
    parser.add_argument(
        "--model",
        type=str,
        help="Comma-separated list of model names (use exact repo IDs from MODEL_CONFIGS)",
    )
    parser.add_argument(
        "--base-dir",
        type=str,
        default=TRITON_BASE_DIR,
        help="Triton model repository root directory",
    )

    args = parser.parse_args()

    # 내보낼 모델 목록 결정
    if args.all:
        target_models = list(MODEL_CONFIGS.keys())
    elif args.model:
        target_models = [m.strip() for m in args.model.split(",")]
        # 유효성 검사
        invalid = [m for m in target_models if m not in MODEL_CONFIGS]
        if invalid:
            logger.error("Invalid model names: {}".format(invalid))
            logger.info("Available models: {}".format(list(MODEL_CONFIGS.keys())))
            sys.exit(1)
    else:
        # 기본: 전체 모델
        target_models = list(MODEL_CONFIGS.keys())
        logger.info(
            "No --all or --model specified. Defaulting to all models."
        )

    success_count = 0
    total_count = len(target_models)

    for model_name in target_models:
        config = MODEL_CONFIGS[model_name]
        # 모델이 로컬에 없으면 다운로드 시도
        if not os.path.exists(
            "/home/tamasic/remote-contextra/models"
            + "/models--"
            + model_name.replace("/", "--")
        ):
            logger.info("Model not local, attempting download...")
            download_model(model_name)

        if export_model(model_name, config, args.base_dir):
            success_count += 1
        gc.collect()  # 메모리 해제

    logger.info("=" * 60)
    logger.info(
        "Export complete: {}/{} models".format(success_count, total_count)
    )
    logger.info("=" * 60)

    if success_count < total_count:
        sys.exit(1)


if __name__ == "__main__":
    main()