# model.py — DETECT/SEG/POSE 공통 스키마 (detect 전용 repo에서도 동일 사용 가능)
import json
import numpy as np
import triton_python_backend_utils as pb_utils
from ultralytics import YOLO
import torch
from pathlib import Path

class TritonPythonModel:
    def initialize(self, args):
        self.device = "cuda" if torch.cuda.is_available() else "cpu"
        repo_dir = Path(args['model_repository'])
        model_dir = repo_dir / args['model_version']

        # task 추론
        name = repo_dir.name.lower()
        if "pose" in name:
            self.task_type = "pose"
        elif "seg" in name or "segment" in name:
            self.task_type = "seg"
        elif "cls" in name or "classify" in name:
            self.task_type = "classify"
        else:
            self.task_type = "detect"

        # YOLO 엔진 로드 (헤드 강제)
        self.model = YOLO(model_dir / "model.engine", task=self.task_type)

        # 클래스 맵
        labels_path = model_dir / "labels.json"
        if labels_path.exists():
            with open(labels_path, "r", encoding="utf-8") as f:
                self.class_names = {int(k): v for k, v in json.load(f).items()}
        else:
            self.class_names = {}

        print(f"[INIT] repo='{repo_dir.name}' v={args['model_version']} task='{self.task_type}' device='{self.device}'")

    def _to_numpy(self, x):
        if isinstance(x, torch.Tensor):
            return x.detach().cpu().numpy()
        return np.array(x)

    def _kp_to_x_y_conf(self, kobj, H, W, max_k=17):
        """
        keypoints를 [N, K, 3] 픽셀 좌표 + conf로 변환.
        conf가 없으면 1.0으로 채움.
        """
        if kobj is None:
            return []
        # 우선순위: data([N,K,3]) → xy([N,K,2])+conf → xyn([N,K,2]) 정규화
        # 1) data
        try:
            arr = self._to_numpy(getattr(kobj, "data", None))
            if arr is not None and arr.ndim == 3 and arr.shape[2] >= 2:
                N, K = arr.shape[0], min(max_k, arr.shape[1])
                out = []
                for i in range(N):
                    one = []
                    for j in range(K):
                        x = float(arr[i, j, 0]); y = float(arr[i, j, 1])
                        c = float(arr[i, j, 2]) if arr.shape[2] >= 3 else 1.0
                        one.append([x, y, c])
                    out.append(one)
                return out
        except Exception:
            pass
        # 2) xy + conf
        try:
            xy = self._to_numpy(getattr(kobj, "xy", None))
            cf = self._to_numpy(getattr(kobj, "conf", None))
            if xy is not None and xy.ndim == 3 and xy.shape[2] == 2:
                N, K = xy.shape[0], min(max_k, xy.shape[1])
                out = []
                for i in range(N):
                    one = []
                    for j in range(K):
                        x = float(xy[i, j, 0]); y = float(xy[i, j, 1])
                        c = float(cf[i, j]) if cf is not None and cf.ndim == 2 else 1.0
                        one.append([x, y, c])
                    out.append(one)
                return out
        except Exception:
            pass
        # 3) xyn 정규화 → 픽셀로 환산
        try:
            xyn = self._to_numpy(getattr(kobj, "xyn", None))
            if xyn is not None and xyn.ndim == 3 and xyn.shape[2] == 2:
                N, K = xyn.shape[0], min(max_k, xyn.shape[1])
                out = []
                for i in range(N):
                    one = []
                    for j in range(K):
                        x = float(xyn[i, j, 0] * W); y = float(xyn[i, j, 1] * H)
                        one.append([x, y, 1.0])
                    out.append(one)
                return out
        except Exception:
            pass
        return []

    def execute(self, requests):
        responses = []
        for request in requests:
            try:
                raw = pb_utils.get_input_tensor_by_name(request, "RAW_IMAGE").as_numpy()
                shapes = pb_utils.get_input_tensor_by_name(request, "ORIGINAL_IMAGE_SHAPE").as_numpy()  # [N,2] = [H,W]
                imgs = list(raw) if raw.dtype == object else list(raw)

                results_batch = self.model.predict(imgs, device=self.device, verbose=False)
                batch_out = []

                for i, res in enumerate(results_batch):
                    H, W = int(shapes[i][0]), int(shapes[i][1])
                    boxes = getattr(res, "boxes", None)
                    has_boxes = boxes is not None and len(boxes) > 0

                    b_abs, b_rel, labels, scores = [], [], [], []
                    if has_boxes:
                        xyxy = self._to_numpy(getattr(boxes, "xyxy", None))
                        if xyxy is not None:
                            b_abs = xyxy.astype(float).tolist()
                            scale = np.array([W, H, W, H], dtype=float)
                            b_rel = (xyxy / scale).astype(float).tolist()
                        cls_ids = self._to_numpy(getattr(boxes, "cls", None))
                        if cls_ids is not None:
                            labels = [self.class_names.get(int(c), str(int(c))) for c in cls_ids.astype(int)]
                        confs = self._to_numpy(getattr(boxes, "conf", None))
                        if confs is not None:
                            scores = confs.astype(float).tolist()

                    item = {
                        "task": self.task_type,
                        "img_size": [H, W],
                        "coord_type": "pixel",
                        "bboxes": b_abs,
                        "relative_bboxes": b_rel,
                        "labels": labels,
                        "scores": scores
                    }

                    if self.task_type == "pose":
                        kpts = getattr(res, "keypoints", None)
                        item["keypoints"] = self._kp_to_x_y_conf(kpts, H, W, max_k=17)

                    if self.task_type == "seg":
                        masks = getattr(res, "masks", None)
                        polys = []
                        if masks is not None and hasattr(masks, "xy"):
                            try:
                                polys = [poly.tolist() for poly in masks.xy]
                            except Exception:
                                polys = []
                        item["masks"] = polys

                    batch_out.append(item)

                out_json = json.dumps(batch_out, ensure_ascii=False)
                out_tensor = pb_utils.Tensor("FINAL_RESULT", np.array([out_json], dtype=object))
                responses.append(pb_utils.InferenceResponse(output_tensors=[out_tensor]))

            except Exception as e:
                responses.append(pb_utils.InferenceResponse(error=pb_utils.TritonError(f"Model execution failed: {e}")))
        return responses

    def finalize(self):
        print("Model unloaded.")
