"""
Contextra AI Analytics — RF-DETR TensorRT Export Script (v3)
RF-DETR v1.9+ -> ONNX -> TensorRT Engine -> Triton Inference Server

주요 개선:
- 모든 모델 일괄 export (12개 전체)
- 실제 디렉토리 구조에 맞는 동적 경로 생성
- config.pbtxt 주석 자동 제거 (Triton protobuf 파서 호환성)
- max_batch_size 기본 8 설정 (dynamic batching 지원)
- 모델별 해상도 및 출력 shape 정확 반영
- Segmentation/Detection/Pose 별도 처리 로직 분리

사용 방법:
  cd /workspace/contextra-analytics/models
  python3 export_all_trt_v3.py --all
  
  또는 특정 모델만:
  python3 export_all_trt_v3.py --model RFDETRSmall,RFDETRLarge
"""

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("trt_export_v3")

# Triton model repository 기본 경로 (docker-compose.yml의 volumes 기준)
DEFAULT_BASE_DIR = "/opt/contextra/models"
TRITON_REPO_ROOT = "triton"

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

# 모델별 설정: resolution, proposals, classes, type
MODEL_CONFIGS = {
    # Detection models (COCO 91 classes)
    "RFDETRNano": {
        "type": "detection",
        "resolution": 384,
        "max_proposals": 300,
        "num_classes": 91,
    },
    "RFDETRSmall": {
        "type": "detection",
        "resolution": 512,
        "max_proposals": 300,
        "num_classes": 91,
    },
    "RFDETRMedium": {
        "type": "detection",
        "resolution": 576,
        "max_proposals": 300,
        "num_classes": 91,
    },
    "RFDETRLarge": {
        "type": "detection",
        "resolution": 704,
        "max_proposals": 300,
        "num_classes": 91,
    },
    "RFDETRBase": {
        "type": "detection",
        "resolution": 504,  # divisible by 56 (patch_size=14)
        "max_proposals": 300,
        "num_classes": 91,
    },
    "RFDETRKeypointPreview": {
        "type": "pose",
        "resolution": 528,
        "max_proposals": 100,
        "num_classes": 2,  # person + background
        "keypoints_dim": 17 * 3,  # COCO 17 keypoints x (x,y,v)
    },
    # Segmentation models (COCO 91 classes)
    "RFDETRSegNano": {
        "type": "segmentation",
        "resolution": 384,
        "max_proposals": 100,
        "num_classes": 91,
    },
    "RFDETRSegSmall": {
        "type": "segmentation",
        "resolution": 504,  # divisible by 24 (patch_size=12)
        "max_proposals": 100,
        "num_classes": 91,
    },
    "RFDETRSegMedium": {
        "type": "segmentation",
        "resolution": 576,
        "max_proposals": 200,
        "num_classes": 91,
    },
    "RFDETRSegLarge": {
        "type": "segmentation",
        "resolution": 720,  # divisible by 24 (was 704)
        "max_proposals": 200,
        "num_classes": 91,
    },
    "RFDETRSegXLarge": {
        "type": "segmentation",
        "resolution": 888,  # divisible by 24 (was 896)
        "max_proposals": 300,
        "num_classes": 91,
    },
    "RFDETRSeg2XLarge": {
        "type": "segmentation",
        "resolution": 1008,  # divisible by 24 (was 1024)
        "max_proposals": 300,
        "num_classes": 91,
    },
}


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


def remove_comments(pbtxt_content):
    """
    Triton protobuf 파서 호환성을 위해 주석 제거
    (주석 라인과 trailing 주석 모두 제거)
    """
    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 generate_pose_config(model_name, input_size, max_proposals, num_classes, keypoints_dim):
    """Pose estimation 모델 config.pbtxt 생성 (keypoints 출력 추가)"""
    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("keypoints"))
    lines.append('    data_type: TYPE_FP32')
    lines.append('    dims: [ {}, -1 ]'.format(max_proposals, keypoints_dim))
    lines.append('  }')
    lines.append("]")
    lines.append("")
    return "\n".join(lines)


def export_model(model_name, config, base_dir):
    """단일 모델 export (ONNX -> TensorRT -> Triton config)

    Args:
        model_name: 모델 이름
        config: MODEL_CONFIGS에서 가져온 설정 dict
        base_dir: Triton model repository 루트 경로

    Returns:
        bool: 성공 여부
    """
    mtype = config["type"]
    resolution = config["resolution"]

    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)
    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)
        )

        # 동적 import로 필요한 클래스만 로드
        model_module = __import__("rfdetr", fromlist=[config.get("class", model_name)])
        ModelClass = getattr(
            model_module, config.get("class", model_name)
        )

        model = ModelClass(resolution=resolution)
        logger.info("Model loaded: {}".format(type(model).__name__))

        # 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 == "pose":
            trtexec_cmd.append("--memPoolSize=workspace:1536")
        else:
            trtexec_cmd.append("--memPoolSize=workspace:1024")

        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,
                config["max_proposals"],
                config["num_classes"]
            )
        elif mtype == "pose":
            config_content = generate_pose_config(
                model_name, resolution,
                config["max_proposals"],
                config["num_classes"],
                config.get("keypoints_dim", 0)
            )
        else:
            config_content = generate_detection_config(
                model_name, resolution,
                config["max_proposals"],
                config["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, ".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="RF-DETR TensorRT Export Pipeline v3"
    )
    parser.add_argument(
        "--all",
        action="store_true",
        help="Export all 12 models",
    )
    parser.add_argument(
        "--model",
        type=str,
        help="Comma-separated list of model names to export "
        "(e.g., RFDETRSmall,RFDETRLarge)",
    )
    parser.add_argument(
        "--base-dir",
        type=str,
        default=DEFAULT_BASE_DIR,
        help="Triton model repository root directory",
    )

    args = parser.parse_args()

    logger.info("=" * 60)
    logger.info("RF-DETR TensorRT Export Pipeline (v3)")
    logger.info("Base dir: {}".format(args.base_dir))
    logger.info("=" * 60)

    # 내보낼 모델 목록 결정
    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 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()