{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":120286,"databundleVersionId":14384428,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import logging\nfrom typing import Dict, List, Optional, Union\n\nimport numpy as np\nimport torch\nfrom torchmetrics import Metric\n\nfrom data_types import (\n    ActionEnum,\n    ActionLabel,\n    ActionPoint,\n    ActionPrediction,\n    Prediction,\n    RawTarget,\n)\n\n# Average bone lengths in meters computed on our 3D pose dataset\n# over 10K poses from different stores, cameras etc.\nBONE_LENGTH_MEANS = {\n    (\"neck\", \"nose\"): 0.19354,\n    (\"left_shoulder\", \"left_elbow\"): 0.27096,\n    (\"left_elbow\", \"left_wrist\"): 0.21228,\n    (\"right_shoulder\", \"right_elbow\"): 0.27210,\n    (\"right_elbow\", \"right_wrist\"): 0.21316,\n    (\"left_hip\", \"left_knee\"): 0.39204,\n    (\"left_knee\", \"left_ankle\"): 0.39530,\n    (\"right_hip\", \"right_knee\"): 0.39266,\n    (\"right_knee\", \"right_ankle\"): 0.39322,\n    (\"left_shoulder\", \"right_shoulder\"): 0.35484,\n    (\"left_hip\", \"right_hip\"): 0.17150,\n    (\"neck\", \"left_shoulder\"): 0.18136,\n    (\"neck\", \"right_shoulder\"): 0.18081,\n    (\"left_shoulder\", \"left_hip\"): 0.51375,\n    (\"right_shoulder\", \"right_hip\"): 0.51226,\n}\n\nM_PX_FACTOR_AVG = 3.07\n\nlogger = logging.getLogger(__name__)\n\n\ndef compute_m_px_factor_from_bones(raw_target: RawTarget) -> Optional[Dict]:\n    \"\"\"Compute meters / pixel factors (one per view) based on precomputed\n    bones lenghts in meters used as reference\n    \n    Args:\n        raw_target: RawTarget with poses stored as dictionaries mapping joint names to coordinates\n    \"\"\"\n    m_px_factors = dict()\n    for rank_name, rank_poses in raw_target.poses.items():\n        m_px_factors[rank_name] = []\n        for pose in rank_poses:  # Iterate through frames\n            if pose is None:\n                continue\n            for (first_bone, second_bone), bone_length_meters in BONE_LENGTH_MEANS.items():\n                # Access joints from dictionary instead of Pose2D attributes\n                first_joint = pose.get(first_bone)\n                second_joint = pose.get(second_bone)\n                \n                if first_joint is not None and second_joint is not None:\n                    bone_length_pixels = np.sqrt(\n                        np.sum((first_joint - second_joint) ** 2)\n                    )\n                    if bone_length_pixels > 0:\n                        m_px_factors[rank_name].append(bone_length_meters / bone_length_pixels)\n    if len(m_px_factors) == 0:\n        return None\n    else:\n        output = {rank_name: np.nanmedian(factors) for rank_name, factors in m_px_factors.items()}\n        output = {\n            rank_name: M_PX_FACTOR_AVG if np.isnan(factor) else factor\n            for rank_name, factor in output.items()\n        }\n        return output\n\n\nclass GlobalMetric(Metric):\n    def __init__(\n        self,\n        confidence_thresholds: np.ndarray = np.linspace(1, 0, 101),\n        null_action_index: int = 3,\n        tiou_thresholds: np.ndarray = np.linspace(0.05, 0.5, 10),\n        spatial_thresholds: np.ndarray = np.linspace(1, 0.05, 20),\n    ):\n        super().__init__()\n        self.eps = 1e-4\n        self.null_action_index = null_action_index\n        self.confidence_thresholds = confidence_thresholds\n        self.tiou_thresholds = tiou_thresholds\n        self.spatial_thresholds = spatial_thresholds\n\n        self.add_state(\n            \"tp\",\n            default=torch.zeros(\n                (\n                    len(self.spatial_thresholds),\n                    len(self.tiou_thresholds),\n                    len(self.confidence_thresholds),\n                ),\n                dtype=torch.float,\n            ),\n            dist_reduce_fx=\"sum\",\n        )\n\n        self.add_state(\n            \"fp\",\n            default=torch.zeros(\n                (\n                    len(self.spatial_thresholds),\n                    len(self.tiou_thresholds),\n                    len(self.confidence_thresholds),\n                ),\n                dtype=torch.float,\n            ),\n            dist_reduce_fx=\"sum\",\n        )\n        self.add_state(\n            \"fn\",\n            default=torch.zeros(\n                (\n                    len(self.spatial_thresholds),\n                    len(self.tiou_thresholds),\n                    len(self.confidence_thresholds),\n                ),\n                dtype=torch.float,\n            ),\n            dist_reduce_fx=\"sum\",\n        )\n\n    def update(self, pred_batch: List[Prediction], raw_target_batch: List[RawTarget]):\n        assert len(pred_batch) == len(raw_target_batch)\n        for i, raw_target in enumerate(raw_target_batch):  # Iterate through batches\n            # Use precomputed m_px_factors if available, otherwise compute from bones\n            if raw_target.m_px_factors is not None:\n                m_px_factors = raw_target.m_px_factors\n            else:\n                m_px_factors = compute_m_px_factor_from_bones(raw_target)\n            \n            if m_px_factors is None or any(np.isnan(factor) for factor in m_px_factors.values()):\n                continue\n            pred = pred_batch[i]\n            pred_actions, target_actions, shelves = (\n                pred.actions,\n                raw_target.action_labels,\n                raw_target.shelves,\n            )\n            \n            # Helper function to get confidence score from probs\n            def get_confidence(action_pred: ActionPrediction) -> float:\n                if isinstance(action_pred.probs, (list, tuple)):\n                    # Multi-class probs: return max prob excluding null_action\n                    return max(p for j, p in enumerate(action_pred.probs) if j != self.null_action_index)\n                else:\n                    # Single score: return directly\n                    return action_pred.probs\n            \n            if self.null_action_index == -1 and len(pred_actions) > 0:\n                if isinstance(pred_actions[0].probs, (list, tuple)):\n                    self.null_action_index = len(pred_actions[0].probs) - 1\n\n            # Sort predictions by confidence score\n            pred_actions = sorted(\n                pred_actions,\n                key=get_confidence,\n                reverse=True,\n            )\n\n            # Iterate through spatial and temporal thresholds\n            for s_idx, spatial_thresh in enumerate(self.spatial_thresholds):\n                for t_idx, tiou_thresh in enumerate(self.tiou_thresholds):\n                    # For each confidence threshold\n                    for c_idx, conf_thresh in enumerate(self.confidence_thresholds):\n                        matched_targets = set()\n                        tp = fp = 0\n\n                        # Check each prediction against confidence threshold\n                        for pred_action in pred_actions:\n                            # Get confidence score\n                            max_prob = get_confidence(pred_action)\n                            if max_prob < conf_thresh:\n                                continue\n\n                            best_iou = 0\n                            best_target_idx = -1\n\n                            # Find best matching target\n                            for target_idx, target in enumerate(target_actions):\n                                if target_idx in matched_targets:\n                                    continue\n\n                                # Compute temporal IoU\n                                t_intersection = min(pred_action.end, target.end) - max(\n                                    pred_action.start, target.start\n                                )\n                                t_union = max(pred_action.end, target.end) - min(\n                                    pred_action.start, target.start\n                                )\n                                tiou = max(0, t_intersection / (t_union + self.eps))\n\n                                # Check spatial match (assuming both have same structure)\n                                spatial_match = True\n                                for cam, ranks in pred_action.spatial.items():\n                                    if cam not in target.spatial:\n                                        continue\n                                    valid_distances = []\n                                    for rank, pred_point in ranks.items():\n                                        if (\n                                            rank not in target.spatial[cam]\n                                            or rank not in m_px_factors\n                                        ):\n                                            continue\n                                        target_point = target.spatial[cam][rank]\n                                        # Convert pixel distance to meters using m_px_factor\n                                        dist_pixels = (\n                                            (pred_point.x - target_point.x) ** 2\n                                            + (pred_point.y - target_point.y) ** 2\n                                        ) ** 0.5\n                                        dist_meters = dist_pixels * m_px_factors[rank]\n                                        valid_distances.append(dist_meters)\n                                    if all(d > spatial_thresh for d in valid_distances):\n                                        spatial_match = False\n                                        break\n                                    if not spatial_match:\n                                        break\n\n                                if tiou >= tiou_thresh and spatial_match:\n                                    if tiou > best_iou:\n                                        best_iou = tiou\n                                        best_target_idx = target_idx\n\n                            # Update counts\n                            if best_target_idx >= 0:\n                                tp += 1\n                                matched_targets.add(best_target_idx)\n                            else:\n                                fp += 1\n\n                        # Compute false negatives\n                        fn = len(target_actions) - len(matched_targets)\n\n                        # Update metric states\n                        self.tp[s_idx, t_idx, c_idx] += tp\n                        self.fp[s_idx, t_idx, c_idx] += fp\n                        self.fn[s_idx, t_idx, c_idx] += fn\n\n    def compute(self) -> Dict[str, Union[torch.Tensor, float]]:\n        \"\"\"Compute AP metrics across spatial and temporal dimensions using vectorized operations.\n\n        Returns:\n            Dict containing:\n            - ap_spatial: AP for each spatial threshold (averaged over tIoU thresholds)\n            - ap_temporal: AP for each tIoU threshold (averaged over spatial thresholds)\n            - ap_global: Single AP value averaged over all thresholds\n        \"\"\"\n        # Calculate precision and recall for all thresholds at once\n        precisions = self.tp / (self.tp + self.fp + self.eps)\n        recalls = self.tp / (self.tp + self.fn + self.eps)\n\n        # For each spatial and temporal threshold combination, sort by recall\n        # and compute maximum precision for each recall level\n        for s_idx in range(len(self.spatial_thresholds)):\n            for t_idx in range(len(self.tiou_thresholds)):\n                # Compute maximum precision for each recall level (going backwards)\n                precisions[s_idx, t_idx] = torch.flip(\n                    torch.cummax(torch.flip(precisions[s_idx, t_idx], [0]), dim=0)[0], [0]\n                )\n\n        # Calculate AP using vectorized operations\n        # Compute recall differences\n        recall_diffs = recalls[..., 1:] - recalls[..., :-1]\n\n        # Compute areas under PR curves using precision values and recall differences\n        ap_values = torch.sum(\n            recall_diffs * precisions[..., :-1], dim=-1\n        )  # Sum over confidence thresholds\n\n        mAP_spatial_curve = ap_values[:, 0]\n        mAP_spatial = torch.mean(\n            ap_values[:, 0]\n        )  # Average over all spatial thresholds with min temporal threshold only\n\n        mAP_temporal_curve = ap_values[0]\n        mAP_temporal = torch.mean(\n            ap_values[0]\n        )  # Average over all temporal thresholds with min spatial threshold only\n\n        # Global mAP is mean of all values\n        mAP_global = torch.mean(ap_values)\n\n        return {\n            \"mAPspatialcurve\": (mAP_spatial_curve, self.spatial_thresholds),\n            \"mAPtemporalcurve\": (mAP_temporal_curve, self.tiou_thresholds),\n            \"mAPspatial\": mAP_spatial,\n            \"mAPtemporal\": mAP_temporal,\n            \"mAPglobal\": mAP_global,\n        }\n\n\nif __name__ == \"__main__\":\n    import numpy as np\n\n    from data_types import (\n        ActionEnum,\n        ActionLabel,\n        ActionPoint,\n        ActionPrediction,\n        Prediction,\n        RawTarget,\n    )\n\n    # Create dummy pose as dictionary mapping joint names to (row, col) coordinates\n    pose1 = {\n        \"left_shoulder\": np.array([100, 100]),\n        \"right_shoulder\": np.array([100, 140]),\n        \"left_hip\": np.array([160, 100]),\n        \"right_hip\": np.array([160, 140]),\n    }\n\n    # Create dummy predictions with correct ActionPrediction format\n    action_point = ActionPoint(x=0.5, y=0.5, visible=True)\n    spatial_dict = {\"action_cam\": {\"rank_0\": action_point}}\n\n    pred_actions = [\n        ActionPrediction(\n            probs=[0.1, 0.9, 0.0, 0.0, 0.0],  # High probability for \"put\" action\n            start=0.0,  # normalized time\n            end=0.33,  # normalized time\n            spatial=spatial_dict,\n        ),\n        ActionPrediction(\n            probs=[0.8, 0.1, 0.1, 0.0, 0.0],  # High probability for \"take\" action\n            start=0.5,  # normalized time\n            end=0.83,  # normalized time\n            spatial=spatial_dict,\n        ),\n    ]\n    prediction = Prediction(actions=pred_actions)\n\n    # Create dummy targets\n    target_actions = [\n        ActionLabel(label=ActionEnum.put, start=0.0, end=0.4, spatial=spatial_dict),\n        ActionLabel(label=ActionEnum.take, start=0.53, end=0.87, spatial=spatial_dict),\n    ]\n\n    # Create poses dictionary with multiple frames\n    poses_dict = {\"rank_0\": [pose1] * 30}  # 30 frames of the same pose\n\n    raw_target = RawTarget(poses=poses_dict, action_labels=target_actions, shelves={}, m_px_factors=None)\n\n    # Initialize and test the metric\n    metric = GlobalMetric()\n    metric.update([prediction], [raw_target])\n    metric_result = metric.compute()\n\n    print(\"Test completed successfully!\")\n    print(f\"TP shape: {metric.tp.shape}\")\n    print(f\"FP shape: {metric.fp.shape}\")\n    print(f\"FN shape: {metric.fn.shape}\")\n    print(f\"Metric result: {metric_result}\")\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null}]}