{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":97984,"databundleVersionId":14096757,"sourceType":"competition"},{"sourceId":14209975,"sourceType":"datasetVersion","datasetId":8894366}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"- Grid 17 Keypoints Detection (Unet segmentation x k-means clustering based) -> Signal Extraction (Unet segmentation x y-coord extraction head)  \n- Used training samples ~ 450 images (grid 17 keypoints manually annotated) for keypoint detection model, all images (GT + Pseudo keypoints) for signal extraction model.  ","metadata":{}},{"cell_type":"code","source":"!pip install -q --no-deps /kaggle/input/physionet2025-private-dataset/timm-1.0.22-py3-none-any.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:15:43.076812Z","iopub.execute_input":"2025-12-19T07:15:43.077073Z","iopub.status.idle":"2025-12-19T07:15:49.764330Z","shell.execute_reply.started":"2025-12-19T07:15:43.077026Z","shell.execute_reply":"2025-12-19T07:15:49.763120Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json\nimport os\nfrom collections import OrderedDict\nimport random\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom scipy.interpolate import interp1d\nfrom scipy.spatial import Delaunay\nfrom scipy.interpolate import Rbf\nfrom PIL import Image\nimport cv2\nimport torch\nfrom torch.utils.data import Dataset\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.linalg as LA\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom fastprogress import progress_bar as pb\nimport easyocr","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:16:09.028513Z","iopub.execute_input":"2025-12-19T07:16:09.029115Z","iopub.status.idle":"2025-12-19T07:16:59.952998Z","shell.execute_reply.started":"2025-12-19T07:16:09.029079Z","shell.execute_reply":"2025-12-19T07:16:59.952043Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Config","metadata":{}},{"cell_type":"code","source":"BASE_IMAGE_DIR = '/kaggle/input/physionet-ecg-image-digitization'\nANNO_FILEPATH = '/kaggle/input/physionet2025-private-dataset/leads_keypoints_annotation_224/leads_keypoints_annotation_224/annotations/person_keypoints_default.json'\n\nLEADS_CROP_IDX_PAIRS = {\n    'I': (0, 1),\n    'II': (5, 6),\n    'III': (10, 11),\n    'aVR': (1, 2),\n    'aVL': (6, 7),\n    'aVF': (11, 12),\n    'V1': (2, 3),\n    'V2': (7, 8),\n    'V3': (12, 13),\n    'V4': (3, 4),\n    'V5': (8, 9),\n    'V6': (13, 14),\n    'II_full': (15, 16),\n}\n\nKPT_CONNECTIONS = [\n    (0,1), (1,2), (2,3), (3,4),        # row 0\n    (5,6), (6,7), (7,8), (8,9),        # row 1\n    (10,11),(11,12),(12,13),(13,14),   # row 2\n    (15,16),                           # row 3\n\n    (0,5),(5,10),(10,15),              # col 0\n    (1,6),(6,11),                      # col 1\n    (2,7),(7,12),                      # col 2\n    (3,8),(8,13),                      # col 3\n    (4,9),(9,14),(14,16),              # col 4\n]\nN_CONNS = len(KPT_CONNECTIONS)\n\nKPTDET_IMAGE_SIZE = (256*3, 256*5)  # (H, W)\nSIG_SEG_IMAGE_SIZE = (256, int(256*2.0))  # (H, W)\n\nN_KPTS = 17#len(KPT_ID_TO_LEAD_MAPPING)\n# keypoints definitions\n# kpt id : (rel_x, rel_y)\n# 0 : (0, 0)\n# 1 : (1, 0)\n# 2 : (2, 0)\n# 3 : (3, 0)\n# 4 : (4, 0)\n# 5 : (0, 1)\n# 6 : (1, 1)\n# 7 : (2, 1)\n# 8 : (3, 1)\n# 9 : (4, 1)\n# 10 : (0, 2)\n# 11 : (1, 2)\n# 12 : (2, 2)\n# 13 : (3, 2)\n# 14 : (4, 2)\n# 15 : (0, 3)\n# 16 : (4, 3)\n\ndevice = 'cuda'\n\nUSE_AMP = True\nAMP_DTYPE = torch.float16\n\nENABLE_AUTO_ROTATION_ADJUST = True\nENABLE_RIGHT_GRID_POINTS_ADJUST = False\nENABLE_GRID_POINTS_ADJUST_BY_REF_AFFINE = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:16:59.954306Z","iopub.execute_input":"2025-12-19T07:16:59.954828Z","iopub.status.idle":"2025-12-19T07:16:59.964199Z","shell.execute_reply.started":"2025-12-19T07:16:59.954805Z","shell.execute_reply":"2025-12-19T07:16:59.963338Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Preparation","metadata":{}},{"cell_type":"code","source":"def load_and_format_keypoint_anno(anno_filepath):\n    anno_dict = json.load(open(anno_filepath, 'r'))\n    df_images = pd.DataFrame(anno_dict['images'])\n    df_images = df_images.rename(columns={'id': 'image_id'})\n    df_images = df_images.drop(columns=['license', 'flickr_url', 'coco_url', 'date_captured'])\n    df_images['ts_label_file_name'] = df_images['file_name'].map(lambda x: x.split('-')[0] + '.csv')\n    df_images['ecg_id'] = df_images['file_name'].map(lambda x: x.split('/')[-2])\n    df_anno = pd.DataFrame(anno_dict['annotations'])\n    df_anno = df_anno.drop(columns=['category_id', 'segmentation', 'iscrowd', 'attributes'])\n    # [x1, y1, visible1, x2, y2, visible2, x3, y3, visible3, ...] -> [[x1, y1], [x2, y2], [x3, y3],...]\n    df_anno['keypoints'] = df_anno['keypoints'].map(lambda flat_kpts: [[x, y] for x, y in zip(flat_kpts[::3], flat_kpts[1::3])])\n    df_anno_image = df_anno.merge(df_images, on='image_id', how='left')\n    return df_anno_image","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:16:59.965316Z","iopub.execute_input":"2025-12-19T07:16:59.965897Z","iopub.status.idle":"2025-12-19T07:16:59.996488Z","shell.execute_reply.started":"2025-12-19T07:16:59.965865Z","shell.execute_reply":"2025-12-19T07:16:59.995467Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_anno_image = load_and_format_keypoint_anno(ANNO_FILEPATH)\nsr_anno_ref = df_anno_image.iloc[9]\n\ndisplay(df_anno_image)\ndisplay(sr_anno_ref)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:16:59.998241Z","iopub.execute_input":"2025-12-19T07:16:59.998624Z","iopub.status.idle":"2025-12-19T07:17:00.204779Z","shell.execute_reply.started":"2025-12-19T07:16:59.998592Z","shell.execute_reply":"2025-12-19T07:17:00.204170Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_meta_test = pd.read_csv(os.path.join(BASE_IMAGE_DIR, 'test.csv'))\ndf_meta_train = pd.read_csv(os.path.join(BASE_IMAGE_DIR, 'train.csv'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:17:00.205475Z","iopub.execute_input":"2025-12-19T07:17:00.205679Z","iopub.status.idle":"2025-12-19T07:17:00.227555Z","shell.execute_reply.started":"2025-12-19T07:17:00.205662Z","shell.execute_reply":"2025-12-19T07:17:00.226809Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Pre Processing","metadata":{}},{"cell_type":"code","source":"def add_vertical_extension_points(src_pts, dst_pts, ratio=0.2):\n    \"\"\"Add vertically extended points along two specific lines.\n\n    Extends the following lines by a certain ratio in both upward and downward\n    directions:\n\n    - Left line: index 0 → index 15\n    - Right line: index 4 → index 16\n\n    The extension amount is determined by `ratio`. For example, if `ratio = 0.2`,\n    the function extends each line by 20% of its original length both upward and\n    downward.\n\n    Args:\n        src_pts (array-like): Source points array of shape (N, 2).\n        dst_pts (array-like): Destination points array of shape (N, 2).\n        ratio (float, optional): Extension ratio relative to line length.\n            Defaults to 0.2.\n\n    Returns:\n        tuple:\n            - numpy.ndarray: Augmented source points (original + 4 extended points).\n            - numpy.ndarray: Augmented destination points (original + 4 extended points).\n    \"\"\"\n\n    # Left line (index 0 → index 15)\n    p0_src = src_pts[0]\n    p15_src = src_pts[15]\n    p0_dst = dst_pts[0]\n    p15_dst = dst_pts[15]\n\n    # Right line (index 4 → index 16)\n    p4_src = src_pts[4]\n    p16_src = src_pts[16]\n    p4_dst = dst_pts[4]\n    p16_dst = dst_pts[16]\n\n    # Vector (downward direction)\n    v_src_left  = p15_src - p0_src\n    v_dst_left  = p15_dst - p0_dst\n    v_src_right = p16_src - p4_src\n    v_dst_right = p16_dst - p4_dst\n\n    # Upward = -v, downward = +v\n    # Create additional points (4 points)\n    src_add = []\n    dst_add = []\n\n    # Left upward extension\n    src_add.append(p0_src - v_src_left * ratio)\n    dst_add.append(p0_dst - v_dst_left * ratio)\n\n    # Right upward extension\n    src_add.append(p4_src - v_src_right * ratio)\n    dst_add.append(p4_dst - v_dst_right * ratio)\n\n    # Left downward extension\n    src_add.append(p15_src + v_src_left * ratio)\n    dst_add.append(p15_dst + v_dst_left * ratio)\n\n    # Right downward extension\n    src_add.append(p16_src + v_src_right * ratio)\n    dst_add.append(p16_dst + v_dst_right * ratio)\n\n    # Combine original points with added points\n    src_all = np.vstack([src_pts, np.array(src_add)])\n    dst_all = np.vstack([dst_pts, np.array(dst_add)])\n\n    return src_all, dst_all\n\n\ndef crop_np(image_np, x1, y1, x2, y2):\n    \"\"\"Safely crop a NumPy image (H, W, C).\n\n    The cropping coordinates are automatically clipped to prevent out-of-range\n    indexing. If the resulting region is empty, an empty array with zero size is\n    returned.\n\n    Args:\n        image_np (numpy.ndarray): Image array of shape (H, W, C).\n        x1 (int): Left coordinate.\n        y1 (int): Top coordinate.\n        x2 (int): Right coordinate.\n        y2 (int): Bottom coordinate.\n\n    Returns:\n        numpy.ndarray: Cropped image array.\n    \"\"\"\n    H, W = image_np.shape[:2]\n\n    # Clip coordinates to valid range\n    x1c = max(0, min(W, x1))\n    x2c = max(0, min(W, x2))\n    y1c = max(0, min(H, y1))\n    y2c = max(0, min(H, y2))\n\n    # Return empty array if the region is invalid\n    if x1c >= x2c or y1c >= y2c:\n        return image_np[:0, :0].copy()\n\n    return image_np[y1c:y2c, x1c:x2c].copy()\n\n\ndef crop_leads_region_from_warped_image(\n    image,\n    dst_pts_after_crop,\n    leads_crop_idx_pairs=LEADS_CROP_IDX_PAIRS,\n    height=130,\n):\n    \"\"\"Crop lead regions from a warped image.\n\n    Given warped image coordinates (`dst_pts_after_crop`), this function extracts\n    rectangular regions corresponding to ECG leads or other indexed line pairs.\n    For each lead, it crops a vertical slice defined by the line connecting two\n    points and extended vertically by `height`.\n\n    Args:\n        image (PIL.Image.Image or numpy.ndarray):\n            Input image (H, W, C).\n        dst_pts_after_crop (array-like):\n            Array of points of shape (N, 2) after warping/cropping.\n        leads_crop_idx_pairs (dict, optional):\n            Mapping of lead names to (start_idx, end_idx) pairs.\n        height (int, optional):\n            Vertical half-height of the crop region. Defaults to 130.\n\n    Returns:\n        OrderedDict:\n            Dictionary mapping each lead name to its cropped image region.\n    \"\"\"\n\n    lead_image_dict = OrderedDict()\n\n    is_pil = isinstance(image, Image.Image)\n    is_np = isinstance(image, np.ndarray)\n\n    if not (is_pil or is_np):\n        raise TypeError(\"image must be PIL.Image.Image or numpy.ndarray\")\n\n    for lead_name, (sid, eid) in leads_crop_idx_pairs.items():\n\n        x1 = int(dst_pts_after_crop[sid][0])\n        y1 = int(dst_pts_after_crop[sid][1] - height)\n        x2 = int(dst_pts_after_crop[eid][0])\n        y2 = int(dst_pts_after_crop[eid][1] + height)\n\n        if is_pil:\n            lead_image_dict[lead_name] = image.crop((x1, y1, x2, y2))\n        else:\n            lead_image_dict[lead_name] = crop_np(image, x1, y1, x2, y2)\n\n    return lead_image_dict","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:17:03.841753Z","iopub.execute_input":"2025-12-19T07:17:03.842599Z","iopub.status.idle":"2025-12-19T07:17:03.854666Z","shell.execute_reply.started":"2025-12-19T07:17:03.842570Z","shell.execute_reply":"2025-12-19T07:17:03.853634Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"easyocr_reader = easyocr.Reader(['en'], gpu=True)\nOCR_KEYWORDS = ['name', 'date', 'age', 'height', 'weight', 'sex', 'avr', 'avl', 'avf']\n\n\ndef detect_rotation_ocr_easyocr(\n    img: np.ndarray,\n    reader=easyocr_reader,\n    keywords=OCR_KEYWORDS,\n):\n    \"\"\"Detect the correct rotation angle of an image using EasyOCR keyword matching.\n\n    The function tries multiple rotations (0°, 90°, 180°, 270°), performs OCR on\n    each rotated image, and scores them based on keyword matches. The rotation\n    angle with the highest score is selected.\n\n    Args:\n        img (numpy.ndarray): Input image array (H, W, C or H, W).\n        reader: EasyOCR reader instance. Defaults to `easyocr_reader`.\n        keywords (list[str]): List of keywords to evaluate OCR results.\n            Defaults to `OCR_KEYWORDS`.\n\n    Returns:\n        int: The best rotation angle in degrees (0, 90, 180, or 270).\n    \"\"\"\n    best_angle = 0\n    best_score = -1\n\n    for angle in [0, 90, 180, 270]:\n        # Rotate image (rot90 rotates in 90-degree steps)\n        rot = np.rot90(img, k=angle // 90)\n\n        # Run EasyOCR\n        results = reader.readtext(rot, detail=1)\n\n        # Calculate keyword-matching score\n        score = 0.0\n        for bbox, text, conf in results:\n            t = text.lower()\n            for kw in keywords:\n                # Partial match\n                if kw in t:\n                    score += float(conf)\n\n        # Update best score/angle\n        if score > best_score:\n            best_score = score\n            best_angle = angle\n\n    return best_angle\n\n\ndef adjust_rotation(img: np.ndarray, best_angle: int):\n    \"\"\"Rotate an image according to the detected best angle.\n\n    Args:\n        img (numpy.ndarray): Input image array.\n        best_angle (int): Angle in degrees (0, 90, 180, 270).\n\n    Returns:\n        numpy.ndarray: Rotated image.\n    \"\"\"\n    if best_angle != 0:\n        img_rot_adjusted = np.rot90(img, k=best_angle // 90)\n    else:\n        img_rot_adjusted = img\n    return img_rot_adjusted\n\n\ndef auto_detect_and_adjust_rotation(\n    img: np.ndarray,\n    reader=easyocr_reader,\n    keywords=OCR_KEYWORDS,\n):\n    \"\"\"Automatically detect and correct the rotation of an image using OCR keywords.\n\n    This function first determines the best rotation angle using OCR-based\n    detection and then rotates the image accordingly.\n\n    Args:\n        img (numpy.ndarray): Input image array.\n        reader: EasyOCR reader instance. Defaults to `easyocr_reader`.\n        keywords (list[str]): List of keywords for scoring OCR results.\n            Defaults to `OCR_KEYWORDS`.\n\n    Returns:\n        tuple:\n            - numpy.ndarray: The rotation-adjusted image.\n            - int: The rotation angle applied.\n    \"\"\"\n    best_angle = detect_rotation_ocr_easyocr(\n        img=img,\n        reader=reader,\n        keywords=keywords,\n    )\n    return adjust_rotation(img, best_angle), best_angle","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:17:04.001644Z","iopub.execute_input":"2025-12-19T07:17:04.001978Z","iopub.status.idle":"2025-12-19T07:17:08.373500Z","shell.execute_reply.started":"2025-12-19T07:17:04.001954Z","shell.execute_reply":"2025-12-19T07:17:08.372805Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_ecg_test_transforms_for_kptdet(img_size=KPTDET_IMAGE_SIZE):\n    \"\"\"Create test-time augmentation transforms for keypoint detection.\n\n    Resizes the image to the specified size `(H, W)`. If `img_size` is `None`,\n    no resizing is applied.\n\n    Args:\n        img_size (tuple[int, int] or None): Target image size as (H, W).\n            If None, resizing is skipped.\n\n    Returns:\n        albumentations.Compose: A composed set of test-time transforms.\n    \"\"\"\n    transforms = [\n        A.Resize(*img_size) if img_size is not None else A.NoOp(),\n    ]\n\n    return A.Compose(transforms)\n\n\ndef get_ecg_test_transforms_for_sigseg(img_size=SIG_SEG_IMAGE_SIZE):\n    \"\"\"Create test-time augmentation transforms for signal segmentation.\n\n    Resizes the image to the specified size `(H, W)`. If `img_size` is `None`,\n    no resizing is applied.\n\n    Args:\n        img_size (tuple[int, int] or None): Target image size as (H, W).\n            If None, resizing is skipped.\n\n    Returns:\n        albumentations.Compose: A composed set of test-time transforms.\n    \"\"\"\n    transforms = [\n        A.Resize(*img_size) if img_size is not None else A.NoOp(),\n    ]\n\n    return A.Compose(transforms)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:17:08.374646Z","iopub.execute_input":"2025-12-19T07:17:08.374901Z","iopub.status.idle":"2025-12-19T07:17:08.380400Z","shell.execute_reply.started":"2025-12-19T07:17:08.374872Z","shell.execute_reply":"2025-12-19T07:17:08.379675Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def prepare_image_for_keypoint_model(\n    image_filepath: str,\n    transforms,\n):\n    \"\"\"Load and preprocess an image for the keypoint detection model.\n\n    The function loads an image, applies optional auto-rotation correction,\n    performs optional center cropping for vertically long images, applies\n    augmentation transforms, and converts the image into a normalized tensor.\n\n    Args:\n        image_filepath (str): Path to the input image file.\n        transforms: Albumentations transform pipeline. If None, no transforms are applied.\n\n    Returns:\n        tuple:\n            - torch.Tensor: Image tensor of shape (C, H, W), normalized to [0, 1].\n            - PIL.Image.Image: Original (post-rotation, pre-transform) PIL image.\n    \"\"\"\n    image_np = cv2.imread(image_filepath)\n    image_np = cv2.cvtColor(image_np, cv2.COLOR_BGR2RGB)\n\n    # ### debug ###\n    # random_rot_angle = random.choice([0, 90, 180, 270])\n    # image_np = np.rot90(image_np, k=random_rot_angle//90)\n    # print(f'random_rot_angle : {random_rot_angle}')\n    # #############\n\n    # Apply automatic rotation adjustment if enabled\n    if ENABLE_AUTO_ROTATION_ADJUST:\n        image_np, best_angle = auto_detect_and_adjust_rotation(image_np)\n        if best_angle != 0:\n            print(f'[DEBUG] image rotation adjusted : angle -> {best_angle}')\n\n    H_org, W_org = image_np.shape[:2]\n    discard_height = int(0.1 * H_org)\n\n    # If the image is vertically too long, crop the vertical center region\n    if (H_org / W_org) > 1.2:\n        print('[DEBUG] cropping vertical center..')\n        image_np = image_np[discard_height:-discard_height]\n\n    org_pil_image = Image.fromarray(image_np)\n\n    # Apply augmentation transforms\n    if transforms:\n        augmented = transforms(image=image_np)\n        image_np = augmented['image']\n\n    # Convert to tensor: (H, W, C) → (C, H, W)\n    image_tensor = torch.from_numpy(image_np.transpose(2, 0, 1)).float() / 255.0\n\n    return image_tensor, org_pil_image\n\n\ndef prepare_image_for_signal_seg_model(\n    warp_cropped_image: \"PIL.Image\",\n    dst_pts_after_crop: torch.Tensor,\n    transforms,\n    resize_image_size=SIG_SEG_IMAGE_SIZE,\n):\n    \"\"\"Prepare multi-lead input data for the signal segmentation model.\n\n    The function crops 12-lead (and II_full) regions from a warped image,\n    resizes each lead sub-image, optionally applies augmentations, and stacks\n    them into a tensor suitable for model inference.\n\n    Args:\n        warp_cropped_image (PIL.Image): Warped and cropped input ECG image.\n        dst_pts_after_crop (torch.Tensor): Array of post-warped keypoints (N, 2).\n        transforms: Albumentations transform pipeline. If None, no transforms are applied.\n        resize_image_size (tuple[int, int], optional): Target size (H, W) for each lead image.\n            Defaults to `SIG_SEG_IMAGE_SIZE`.\n\n    Returns:\n        torch.Tensor: Tensor of shape (n_leads, 3, H, W), normalized to [0, 1].\n    \"\"\"\n    resize_h, resize_w = resize_image_size\n    leads_crop_image_dict = crop_leads_region_from_warped_image(\n        warp_cropped_image, dst_pts_after_crop\n    )\n\n    SORTED_LEAD_NAMES = [\n        'I', 'aVR', 'V1', 'V4',\n        'II', 'aVL', 'V2', 'V5',\n        'III', 'aVF', 'V3', 'V6',\n        'II_full',\n    ]\n\n    lead_image_list = []\n\n    # Construct the lead image list in the required order\n    for k in SORTED_LEAD_NAMES:\n        if k != 'II_full':\n            lead_image_list.append(\n                np.array(leads_crop_image_dict[k].resize((resize_w, resize_h)))\n            )\n        else:\n            img = leads_crop_image_dict[k]\n            w, h = img.size\n            # Split II_full into 4 equal horizontal segments\n            lead_image_list += [\n                np.array(\n                    img.crop((i * w // 4, 0, (i + 1) * w // 4, h)).resize((resize_w, resize_h))\n                )\n                for i in range(4)\n            ]\n\n    # Apply transforms to each lead image\n    if transforms:\n        aug_image_np_list = []\n        for image_np in lead_image_list:\n            augmented = transforms(image=image_np)\n            image_np = augmented['image']\n            aug_image_np_list.append(image_np)\n        lead_image_list = aug_image_np_list\n\n    # Stack into array: (n_leads, H, W, 3)\n    image_np = np.stack(lead_image_list, axis=0)\n\n    # Convert to tensor: (n_leads, H, W, 3) → (n_leads, 3, H, W)\n    image_tensor = torch.from_numpy(image_np).float().permute(0, 3, 1, 2) / 255.0\n\n    return image_tensor","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:17:08.381409Z","iopub.execute_input":"2025-12-19T07:17:08.381698Z","iopub.status.idle":"2025-12-19T07:17:08.405897Z","shell.execute_reply.started":"2025-12-19T07:17:08.381680Z","shell.execute_reply":"2025-12-19T07:17:08.405193Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Post Processing","metadata":{}},{"cell_type":"code","source":"from scipy.spatial.distance import cdist\nfrom sklearn.decomposition import PCA\nfrom scipy.optimize import linear_sum_assignment\nfrom sklearn.cluster import KMeans\n\n\ndef extract_kpts_by_kmeans(pred_probs: torch.Tensor, n_kpts: int = N_KPTS):\n    \"\"\"Extract keypoints from heatmaps using weighted KMeans clustering.\n\n    The function collapses the channel dimension into a single heatmap,\n    filters low-response pixels, pads if necessary, and applies KMeans to\n    generate `n_kpts` representative keypoint coordinates.\n\n    Args:\n        pred_probs (torch.Tensor): Heatmap tensor of shape (C, H, W) after\n            sigmoid or softmax.\n        n_kpts (int, optional): Number of keypoints to extract. Defaults to\n            `N_KPTS`.\n\n    Returns:\n        np.ndarray: Array of shape (n_kpts, 2) containing keypoint centers\n        in (x, y) format.\n    \"\"\"\n    C, H, W = pred_probs.shape\n\n    # Combine channel responses into a single heatmap\n    heat = pred_probs.sum(dim=0).cpu().numpy()  # (H, W)\n\n    # Prepare flattened pixel coordinates and weights\n    ys, xs = np.indices((H, W))\n    xs = xs.reshape(-1)\n    ys = ys.reshape(-1)\n    w = heat.reshape(-1)\n\n    # Filter out low-weight pixels\n    mask = w > (0.1 * w.max())  # threshold (tunable)\n    xs = xs[mask]\n    ys = ys[mask]\n    w = w[mask]\n\n    # ------------------------------------------------------------------\n    # Padding if points are fewer than n_kpts\n    # ------------------------------------------------------------------\n    num_pts = len(xs)\n    if num_pts < n_kpts:\n        pad_len = n_kpts - num_pts\n\n        xs = np.concatenate([xs, np.zeros(pad_len, dtype=xs.dtype)])\n        ys = np.concatenate([ys, np.zeros(pad_len, dtype=ys.dtype)])\n\n        # Pad weights with small values\n        pad_w = np.full(pad_len, 0.1 * w.max() if len(w) > 0 else 0.1)\n        w = np.concatenate([w, pad_w])\n\n    # Weighted KMeans\n    km = KMeans(n_clusters=n_kpts, n_init='auto')\n    km.fit(np.stack([xs, ys], axis=1), sample_weight=w)\n\n    # Cluster centers (x, y)\n    centers = km.cluster_centers_\n    return centers\n\n\ndef remove_duplicate_kpts(centers: np.ndarray, min_dist: float = 10.0):\n    \"\"\"Remove duplicate keypoints based on Euclidean distance.\n\n    Keypoints closer than `min_dist` are considered duplicates; only one\n    point from each duplicate group is kept.\n\n    Args:\n        centers (np.ndarray): Array of keypoints of shape (N, 2).\n        min_dist (float): Minimum allowed distance between keypoints. Points\n            closer than this are treated as duplicates.\n\n    Returns:\n        np.ndarray: Filtered keypoints (M, 2), where M ≤ N.\n    \"\"\"\n    if len(centers) == 0:\n        return centers\n\n    centers = centers.copy()\n    keep = np.ones(len(centers), dtype=bool)\n\n    # Compute pairwise distance matrix\n    D = cdist(centers, centers)  # (N, N)\n\n    # Only examine upper-triangular portion (ignore self-distances)\n    N = len(centers)\n    for i in range(N):\n        if not keep[i]:\n            continue\n        for j in range(i + 1, N):\n            if not keep[j]:\n                continue\n            if D[i, j] < min_dist:\n                keep[j] = False\n\n    return centers[keep]\n\n\ndef interpolate_missing_keypoints_by_affine(pred_ref_sorted):\n    \"\"\"Interpolate missing keypoints using an affine transformation.\n\n    Given a predefined 17-point reference grid, this function fits an affine\n    transformation from the valid (non-NaN) keypoints and uses it to fill\n    in the missing ones.\n\n    Args:\n        pred_ref_sorted (np.ndarray): Array of shape (17, 2). Known points are\n            given directly; missing points contain NaN entries.\n\n    Returns:\n        np.ndarray: Filled keypoints of shape (17, 2), dtype float32.\n    \"\"\"\n    # Define 17-point grid\n    GRID_IDX = [\n        (0,0),(1,0),(2,0),(3,0),(4,0),\n        (0,1),(1,1),(2,1),(3,1),(4,1),\n        (0,2),(1,2),(2,2),(3,2),(4,2),\n        (0,3),(4,3)\n    ]\n    pred = pred_ref_sorted.astype(float)\n    ref = np.array(GRID_IDX).astype(float)\n\n    valid = ~np.isnan(pred[:, 0])\n    missing = ~valid\n\n    if valid.sum() < 3:\n        raise ValueError(\"At least 3 points are required to determine an affine transform.\")\n\n    ref_valid = ref[valid]\n    pred_valid = pred[valid]\n\n    # --- Least-squares affine estimation ---\n    ref_aug = np.hstack([ref_valid, np.ones((len(ref_valid), 1))])  # (M, 3)\n    X, *_ = np.linalg.lstsq(ref_aug, pred_valid, rcond=None)        # (3, 2)\n\n    # --- Fill only NaN positions ---\n    if missing.any():\n        ref_missing = ref[missing]\n        ref_missing_aug = np.hstack([ref_missing, np.ones((len(ref_missing), 1))])\n        pred_filled = ref_missing_aug @ X\n        pred[missing] = pred_filled\n\n    return pred.astype(np.float32)\n\n\ndef match_keypoints_2d_with_rot180(pred_centers, ref_centers, eps=1e-6):\n    \"\"\"Match 2D keypoints considering possible 180-degree rotation.\n\n    This function uses PCA-based normalization, scaling, translation alignment,\n    and Hungarian matching to determine the best correspondence between\n    predicted keypoints and reference keypoints. Both the original and\n    180-degree rotated versions of the predicted points are evaluated.\n\n    Args:\n        pred_centers (np.ndarray): Predicted keypoints of shape (N, 2).\n        ref_centers (np.ndarray): Reference keypoints of shape (N, 2).\n        eps (float): Small constant to avoid division by zero.\n\n    Returns:\n        tuple:\n            - np.ndarray: Matching index pairs of shape (N, 2).\n            - np.ndarray: Transformed predicted keypoints after PCA and scaling.\n            - np.ndarray: Predicted keypoints reordered to match reference order,\n                with unmatched positions filled with NaN.\n            - float: Total matching cost.\n    \"\"\"\n    pred = np.asarray(pred_centers, dtype=np.float32)\n    ref = np.asarray(ref_centers, dtype=np.float32)\n\n    N_pred = len(pred)\n    N_ref = len(ref)\n\n    # --- PCA alignment of reference (align major axis to x-axis) ---\n    pca_ref = PCA(n_components=2).fit(ref)\n    ref_trans = pca_ref.transform(ref)\n    if np.ptp(ref_trans[:, 0]) < 0:\n        ref_trans[:, 0] *= -1\n\n    # Normalize by the longest axis\n    ref_trans /= (np.ptp(ref_trans, axis=0).max() + eps)\n\n    # Reference bounding box\n    ref_bbox = ref_trans.max(axis=0) - ref_trans.min(axis=0)\n\n    # --- Candidates: original and 180° rotated (x,y flipped) ---\n    pred_variants = [pred.copy(), pred.copy() * np.array([-1, -1])]\n\n    best_cost = np.inf\n    best_pred_trans = None\n    best_matches = None\n\n    for pred_var in pred_variants:\n        # PCA alignment\n        pca_pred = PCA(n_components=2).fit(pred_var)\n        pred_trans = pca_pred.transform(pred_var)\n        if np.ptp(pred_trans[:, 0]) < 0:\n            pred_trans[:, 0] *= -1\n\n        # Scale predicted keypoints to match reference bounding box\n        pred_bbox = pred_trans.max(axis=0) - pred_trans.min(axis=0)\n        scale_factor = ref_bbox / (pred_bbox + eps)\n        scale_uniform = scale_factor.min()\n        pred_trans *= scale_uniform\n\n        # Center alignment before Hungarian matching\n        pred_centroid = pred_trans.mean(axis=0)\n        ref_centroid = ref_trans.mean(axis=0)\n        translation = ref_centroid - pred_centroid\n        pred_trans += translation\n\n        # Compute distance matrix and optimal assignment\n        D = np.linalg.norm(pred_trans[:, None, :] - ref_trans[None, :, :], axis=-1)\n        pred_ids, ref_ids = linear_sum_assignment(D)\n        cost = D[pred_ids, ref_ids].sum()\n\n        if cost < best_cost:\n            best_cost = cost\n            best_pred_trans = pred_trans[pred_ids]\n            best_matches = np.stack([pred_ids, ref_ids], axis=1)\n\n    # --- Reorder to reference indexing; fill unmatched entries with NaN ---\n    pred_ref_sorted = np.full((N_ref, 2), np.nan, dtype=np.float32)\n    for pid, rid in best_matches:\n        pred_ref_sorted[rid] = pred[pid]\n\n    return best_matches, best_pred_trans, pred_ref_sorted, best_cost","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:17:08.408128Z","iopub.execute_input":"2025-12-19T07:17:08.408476Z","iopub.status.idle":"2025-12-19T07:17:08.996875Z","shell.execute_reply.started":"2025-12-19T07:17:08.408452Z","shell.execute_reply":"2025-12-19T07:17:08.996223Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def extract_keypoints_from_kpt_model_pred(\n    pred_logits_mask: torch.Tensor,\n    org_pil_image: \"PIL.Image\",\n    n_kpts=N_KPTS,\n    ref_keypoints=sr_anno_ref.keypoints,\n):\n    \"\"\"Extract and align keypoints from the keypoint detection model output.\n\n    The function performs the following steps:\n    1. Converts model logits to probabilities.\n    2. Extracts keypoints using weighted KMeans clustering.\n    3. Removes duplicate keypoints.\n    4. Matches extracted keypoints with reference keypoints, considering\n       potential 180° rotation.\n    5. Fills missing keypoints using affine interpolation.\n    6. Rescales keypoints to match the original image resolution.\n    7. Applies horizontal flip correction if required.\n\n    Args:\n        pred_logits_mask (torch.Tensor): Raw logits of shape (C, H, W).\n        org_pil_image (PIL.Image): Original image before preprocessing.\n        n_kpts (int, optional): Expected number of keypoints. Defaults to `N_KPTS`.\n        ref_keypoints (array-like): Reference keypoint coordinates.\n\n    Returns:\n        tuple:\n            - np.ndarray: Sorted keypoints (n_kpts, 2) in original image scale.\n            - np.ndarray: Unsorted extracted keypoints prior to matching.\n    \"\"\"\n    C, H, W = pred_logits_mask.shape\n    W_org, H_org = org_pil_image.size\n\n    pred_probs_mask = pred_logits_mask.sigmoid()\n\n    # Extract keypoints from heatmap using KMeans\n    pred_kpts = extract_kpts_by_kmeans(pred_probs_mask, n_kpts=n_kpts)\n\n    # Remove duplicates\n    pred_kpts = remove_duplicate_kpts(pred_kpts)\n\n    # Match unsorted predicted keypoints to sorted reference keypoints\n    _matches, _, pred_kpts_ref_sorted, _ = match_keypoints_2d_with_rot180(\n        pred_kpts,\n        ref_centers=torch.tensor(sr_anno_ref.keypoints)\n    )\n\n    # Fill NaN keypoints via affine interpolation\n    pred_kpts_ref_sorted = interpolate_missing_keypoints_by_affine(pred_kpts_ref_sorted)\n    assert pred_kpts_ref_sorted.shape == (n_kpts, 2)\n    assert not np.isnan(pred_kpts_ref_sorted).any()\n\n    # Rescale to original image size\n    pred_kpts_ref_sorted[:, 0] = ((W_org / W) * pred_kpts_ref_sorted[:, 0]).astype(int)\n    pred_kpts_ref_sorted[:, 1] = ((H_org / H) * pred_kpts_ref_sorted[:, 1]).astype(int)\n    pred_kpts_ref_sorted = pred_kpts_ref_sorted.astype(int)\n\n    # Horizontal flip correction based on keypoint ordering\n    X_POINTS_FLIP_IDX_MAPPING = np.array([\n        (4, 0),\n        (3, 1),\n        (2, 2),\n        (1, 3),\n        (0, 4),\n        (9, 5),\n        (8, 6),\n        (7, 7),\n        (6, 8),\n        (5, 9),\n        (14, 10),\n        (13, 11),\n        (12, 12),\n        (11, 13),\n        (10, 14),\n        (16, 15),\n        (15, 16),\n    ])\n\n    if pred_kpts_ref_sorted[0, 0] > pred_kpts_ref_sorted[4, 0]:\n        # Apply horizontal flip reordering\n        pred_kpts_ref_sorted = pred_kpts_ref_sorted[X_POINTS_FLIP_IDX_MAPPING[:, 0]]\n    else:\n        pass\n\n    return pred_kpts_ref_sorted, pred_kpts\n\n\ndef apply_warp_affine_from_keypoints(\n    org_pil_image,\n    pred_kpts_ref_sorted: np.ndarray,\n    ref_keypoints=sr_anno_ref.keypoints,\n    W_ref=sr_anno_ref.width,\n    H_ref=sr_anno_ref.height,\n    extend_point_ratio=0.2,\n):\n    \"\"\"Apply Delaunay-based piecewise affine warping using matched keypoints.\n\n    This function warps an image to a reference coordinate system using\n    pointwise affine transforms computed from Delaunay triangulation.\n\n    Steps:\n    1. Add 4 vertically extended keypoints to stabilize edges.\n    2. Perform Delaunay triangulation on the destination points.\n    3. For each triangle, compute an affine transform and warp the region.\n    4. Composite warped triangles into the output canvas.\n    5. Crop the warped image using the bounding box of extended keypoints.\n    6. Adjust destination keypoints to cropped coordinate space.\n\n    Args:\n        org_pil_image (PIL.Image): Original image.\n        pred_kpts_ref_sorted (np.ndarray): Sorted predicted keypoints (N, 2).\n        ref_keypoints (array-like): Reference keypoints corresponding to output grid.\n        W_ref (int): Width of the reference image.\n        H_ref (int): Height of the reference image.\n        extend_point_ratio (float): Ratio used to add vertical extension points.\n\n    Returns:\n        tuple:\n            - PIL.Image: Cropped warped image.\n            - np.ndarray: Extended destination keypoints (N + 4, 2).\n            - np.ndarray: Destination keypoints after crop (N, 2).\n    \"\"\"\n    target = np.array(org_pil_image)\n\n    warped = np.zeros((H_ref, W_ref, 3), dtype=np.uint8)\n\n    src_pts = pred_kpts_ref_sorted.astype(np.float32)\n    dst_pts = np.array(ref_keypoints, dtype=np.float32)\n\n    # Add 4 vertical extension points\n    src_pts_ext, dst_pts_ext = add_vertical_extension_points(src_pts, dst_pts, extend_point_ratio)\n\n    # Delaunay triangulation\n    tri = Delaunay(dst_pts_ext)\n\n    for simplex in tri.simplices:\n        src_tri = src_pts_ext[simplex]\n        dst_tri = dst_pts_ext[simplex]\n\n        M = cv2.getAffineTransform(src_tri, dst_tri)\n\n        xmin = max(int(dst_tri[:, 0].min()), 0)\n        xmax = min(int(dst_tri[:, 0].max()) + 1, W_ref)\n        ymin = max(int(dst_tri[:, 1].min()), 0)\n        ymax = min(int(dst_tri[:, 1].max()) + 1, H_ref)\n\n        if xmax <= xmin or ymax <= ymin:\n            continue\n\n        warped_patch = cv2.warpAffine(target, M, (W_ref, H_ref))\n\n        mask = np.zeros((H_ref, W_ref), dtype=np.uint8)\n        cv2.fillConvexPoly(mask, dst_tri.astype(np.int32), 255)\n        mask3 = np.stack([mask] * 3, axis=-1)\n\n        warped = np.where(mask3 > 0, warped_patch, warped)\n\n    after = warped.copy()\n\n    # Crop warped output using extended keypoints\n    x_min = int(dst_pts_ext[:, 0].min())\n    x_max = int(dst_pts_ext[:, 0].max())\n    y_min = int(dst_pts_ext[:, 1].min())\n    y_max = int(dst_pts_ext[:, 1].max())\n    after_cropped = after[y_min:y_max, x_min:x_max]\n\n    # Adjust destination keypoints to cropped coordinates\n    dst_pts_ext_after_crop = dst_pts_ext.copy()\n    dst_pts_ext_after_crop[:, 0] -= x_min\n    dst_pts_ext_after_crop[:, 1] -= y_min\n    dst_pts_after_crop = dst_pts_ext_after_crop[:-4].astype(int)  # exclude extension points\n\n    return Image.fromarray(after_cropped), dst_pts_ext, dst_pts_after_crop","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:17:08.997748Z","iopub.execute_input":"2025-12-19T07:17:08.998027Z","iopub.status.idle":"2025-12-19T07:17:09.014212Z","shell.execute_reply.started":"2025-12-19T07:17:08.998001Z","shell.execute_reply":"2025-12-19T07:17:09.013332Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from scipy.spatial import procrustes\n\ndef shape_similarity(true_kpts, pred_kpts):\n    \"\"\"Compute shape similarity between two keypoint sets using Procrustes analysis.\n\n    This function performs Procrustes alignment between the ground-truth\n    keypoints and predicted keypoints, and returns a similarity score.\n\n    The `scipy.spatial.procrustes` function normalizes both shapes, aligns them,\n    and returns a disparity value:\n        - disparity = 0.0 → perfect match\n        - disparity > 1.0 → large difference\n\n    A similarity score is computed as `1 - disparity`, clipped implicitly by\n    nature of Procrustes output to roughly the 0–1 range.\n\n    Args:\n        true_kpts (array-like): Ground-truth keypoints of shape (N, 2).\n        pred_kpts (array-like): Predicted keypoints of shape (N, 2).\n\n    Returns:\n        float: Similarity score, where 1.0 indicates identical shapes and\n               values near 0 indicate low similarity.\n    \"\"\"\n    mtx1, mtx2, disparity = procrustes(true_kpts, pred_kpts)\n\n    # If disparity = 0 → perfect alignment; if disparity ≥ 1 → very different\n    return 1 - disparity  # similarity score (approximately 0–1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:17:09.015251Z","iopub.execute_input":"2025-12-19T07:17:09.015517Z","iopub.status.idle":"2025-12-19T07:17:09.046371Z","shell.execute_reply.started":"2025-12-19T07:17:09.015498Z","shell.execute_reply":"2025-12-19T07:17:09.045512Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def fit_and_extrapolate(points, t_extr):\n    \"\"\"Fit a 2D line to points and extrapolate to a given parameter.\n\n    Fits a linear function separately for x and y coordinates over a sequence\n    of points, then evaluates the line at a specified extrapolation position.\n\n    Args:\n        points (array-like): Array of shape (N, 2) representing points to fit.\n        t_extr (float): Parameter value at which to extrapolate the fitted line\n                        (e.g., 4/3 * (N-1)).\n\n    Returns:\n        np.ndarray: Extrapolated point of shape (2,).\n    \"\"\"\n    pts = np.asarray(points)\n    N = pts.shape[0]\n\n    # t = 0,1,2,...,N-1\n    t = np.arange(N)\n\n    # Fit x(t) = ax * t + bx, y(t) = ay * t + by\n    A = np.vstack([t, np.ones_like(t)]).T\n    ax, bx = np.linalg.lstsq(A, pts[:,0], rcond=None)[0]\n    ay, by = np.linalg.lstsq(A, pts[:,1], rcond=None)[0]\n\n    # Evaluate at t = t_extr\n    x_new = ax * t_extr + bx\n    y_new = ay * t_extr + by\n    return np.array([x_new, y_new])\n\n\ndef adjust_right_edge_keypoints(kpts):\n    \"\"\"Adjust right-edge keypoints by affine-invariant linear extrapolation.\n\n    Certain right-edge keypoints (indices 4, 9, 14, 16) are replaced by\n    extrapolated values computed from neighboring points along rows and columns.\n\n    Args:\n        kpts (np.ndarray): Array of shape (17, 2) representing keypoints.\n\n    Returns:\n        np.ndarray: Keypoints array of shape (17, 2) with adjusted right-edge points.\n    \"\"\"\n    kpts = kpts.copy()\n\n    # --- Row 0: ids [0,1,2,3] -> extrapolate for id=4 ---\n    row0 = kpts[[0,1,2,3]]\n    kpts[4] = fit_and_extrapolate(row0, t_extr=4)\n\n    # --- Row 1: ids [5,6,7,8] -> extrapolate for id=9 ---\n    row1 = kpts[[5,6,7,8]]\n    kpts[9] = fit_and_extrapolate(row1, t_extr=4)\n\n    # --- Row 2: ids [10,11,12,13] -> extrapolate for id=14 ---\n    row2 = kpts[[10,11,12,13]]\n    kpts[14] = fit_and_extrapolate(row2, t_extr=4)\n\n    # --- Column: ids [4,9,14] -> extrapolate for id=16 ---\n    col = kpts[[4,9,14]]\n    kpts[16] = fit_and_extrapolate(col, t_extr=3)\n\n    return kpts","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:17:09.047363Z","iopub.execute_input":"2025-12-19T07:17:09.047636Z","iopub.status.idle":"2025-12-19T07:17:09.065664Z","shell.execute_reply.started":"2025-12-19T07:17:09.047609Z","shell.execute_reply":"2025-12-19T07:17:09.064997Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def map_ref_to_pred_affine_only(\n    pred_kpts, ref_kpts,\n    ransac_thresh=0.05,\n    debug=False,\n    apply_only_outliers=True,\n):\n    \"\"\"\n    Affine-only correction: map reference keypoints to prediction coordinates while \n    preserving the shape of ref_kpts.\n\n    The function normalizes coordinates internally to estimate an affine transformation,\n    so the estimation is scale-invariant.\n\n    Args:\n        pred_kpts (np.ndarray): Array of shape (N,2) representing predicted keypoints.\n        ref_kpts (np.ndarray): Array of shape (N,2) representing reference keypoints.\n        ransac_thresh (float): RANSAC residual threshold (normalized coordinate space).\n        debug (bool): If True, scatter plot of the mapped keypoints is shown.\n        apply_only_outliers (bool): If True, inliers remain unchanged and only \n                                     outliers are corrected using affine.\n\n    Returns:\n        final_out (np.ndarray): (N,2) mapped keypoints in pred coordinate system.\n        inliers (np.ndarray): (N,) boolean array indicating inlier keypoints.\n        M_full (np.ndarray): (2,3) affine matrix in pred coordinate system.\n    \"\"\"\n    # ----------------------------\n    # Step 0: Normalize coordinates\n    # ----------------------------\n    def normalize(pts):\n        min_xy = pts.min(axis=0)\n        max_xy = pts.max(axis=0)\n        scale = max_xy - min_xy\n        scale[scale == 0] = 1.0\n        pts_norm = (pts - min_xy) / scale\n        return pts_norm, min_xy, scale\n\n    pred = np.asarray(pred_kpts, dtype=np.float32)\n    ref  = np.asarray(ref_kpts, dtype=np.float32)\n\n    pred_norm, pred_min, pred_scale = normalize(pred)\n    ref_norm,  ref_min,  ref_scale  = normalize(ref)\n\n    # ----------------------------\n    # Step 1: Robust affine estimation (normalized space)\n    # ----------------------------\n    M_norm, inliers_mask = cv2.estimateAffine2D(\n        ref_norm, pred_norm,\n        method=cv2.RANSAC,\n        ransacReprojThreshold=ransac_thresh\n    )\n\n    if M_norm is None:\n        M_norm = np.eye(2,3, dtype=np.float32)\n        inliers_mask = np.ones(len(ref), dtype=bool)\n\n    inliers = inliers_mask.ravel().astype(bool)\n\n    # ----------------------------\n    # Step 2: Map back to pred coordinates\n    # ----------------------------\n    ref_scaled = (ref - ref_min) / ref_scale\n    ref_scaled_homo = np.hstack([ref_scaled, np.ones((len(ref_scaled), 1))])\n\n    mapped_ref_norm = (M_norm @ ref_scaled_homo.T).T\n    mapped_ref = mapped_ref_norm * pred_scale + pred_min\n\n    # ----------------------------\n    # Step 3: Convert M_norm → M_full in pred coordinates\n    # ----------------------------\n    M_full = np.zeros((2,3), dtype=np.float32)\n    M_full[:,0] = M_norm[:,0] * (pred_scale / ref_scale)\n    M_full[:,1] = M_norm[:,1] * (pred_scale / ref_scale)\n    M_full[:,2] = (\n        M_norm[:,0] * (-ref_min[0] / ref_scale[0]) +\n        M_norm[:,1] * (-ref_min[1] / ref_scale[1]) +\n        M_norm[:,2]\n    ) * pred_scale + pred_min\n\n    # ----------------------------\n    # Step 4: Optionally preserve inliers\n    # ----------------------------\n    if apply_only_outliers:\n        final_out = mapped_ref.copy()\n        final_out[inliers] = pred[inliers]\n    else:\n        final_out = mapped_ref\n\n    # ----------------------------\n    # Step 5: Debug visualization\n    # ----------------------------\n    if debug:\n        plt.figure(figsize=(6,6))\n        plt.scatter(final_out[:,0], final_out[:,1], c='blue', alpha=0.4, label='mapped_ref')\n        plt.scatter(final_out[inliers,0], final_out[inliers,1],\n                    facecolors='none', edgecolors='blue', s=100, label='inliers')\n        plt.scatter(pred[:,0], pred[:,1], c='orange', label='pred')\n        plt.legend()\n        plt.gca().invert_yaxis()\n        plt.axis('equal')\n        plt.show()\n\n    return final_out, inliers, M_full\n\n\ndef map_ref_to_pred_affine_tps(pred_kpts, ref_kpts,\n                               ransac_thresh=0.05,\n                               debug=False):\n    \"\"\"\n    TPS correction: apply thin-plate spline (TPS) transform to map reference keypoints\n    to predicted keypoints, only correcting outliers.\n\n    Procedure:\n        1) Robustly obtain inliers via RANSAC affine (scale-invariant)\n        2) Learn TPS(ref→pred) from inliers\n        3) Apply TPS to outlier reference keypoints to get corrected pred positions\n\n    Args:\n        pred_kpts (np.ndarray): (N,2) predicted keypoints.\n        ref_kpts (np.ndarray): (N,2) reference keypoints.\n        ransac_thresh (float): RANSAC threshold in normalized space.\n        debug (bool): If True, shows scatter plot.\n\n    Returns:\n        adjusted (np.ndarray): (N,2) adjusted keypoints after TPS correction.\n        inliers (np.ndarray): (N,) boolean array indicating inliers.\n    \"\"\"\n    # ----------------------------\n    # Step 0: Normalize coordinates\n    # ----------------------------\n    def normalize(pts):\n        mn = pts.min(axis=0)\n        mx = pts.max(axis=0)\n        scale = mx - mn\n        scale[scale == 0] = 1.0\n        pts_norm = (pts - mn) / scale\n        return pts_norm, mn, scale\n\n    pred = np.asarray(pred_kpts, dtype=np.float32)\n    ref  = np.asarray(ref_kpts, dtype=np.float32)\n\n    pred_norm, pred_min, pred_scale = normalize(pred)\n    ref_norm,  ref_min,  ref_scale  = normalize(ref)\n\n    # ----------------------------\n    # Step 1: Robust affine estimation\n    # ----------------------------\n    M_norm, inliers_mask = cv2.estimateAffine2D(\n        ref_norm, pred_norm,\n        method=cv2.RANSAC,\n        ransacReprojThreshold=ransac_thresh\n    )\n\n    if M_norm is None:\n        M_norm = np.eye(2,3, dtype=np.float32)\n        inliers_mask = np.ones(len(pred), bool)\n\n    inliers = inliers_mask.ravel().astype(bool)\n\n    # ----------------------------\n    # Step 2: Learn TPS from inliers\n    # ----------------------------\n    ref_in  = ref[inliers]\n    pred_in = pred[inliers]\n\n    # Thin Plate Spline via Rbf\n    rbf_x = Rbf(ref_in[:,0], ref_in[:,1], pred_in[:,0], function='thin_plate')\n    rbf_y = Rbf(ref_in[:,0], ref_in[:,1], pred_in[:,1], function='thin_plate')\n\n    # ----------------------------\n    # Step 3: Apply TPS to outliers\n    # ----------------------------\n    adjusted = pred.copy()\n    outliers = ~inliers\n\n    if np.any(outliers):\n        ox = rbf_x(ref[outliers,0], ref[outliers,1])\n        oy = rbf_y(ref[outliers,0], ref[outliers,1])\n        adjusted[outliers,0] = ox\n        adjusted[outliers,1] = oy\n\n    # ----------------------------\n    # Step 4: Debug visualization\n    # ----------------------------\n    if debug:\n        plt.figure(figsize=(6,6))\n        plt.scatter(pred[:,0], pred[:,1], c='orange', label='pred')\n        plt.scatter(adjusted[:,0], adjusted[:,1], c='blue', alpha=0.5, label='adjusted')\n        plt.scatter(adjusted[outliers,0], adjusted[outliers,1],\n                    facecolors='none', edgecolors='red', s=120,\n                    label='corrected outliers')\n        plt.legend()\n        plt.axis('equal')\n        plt.gca().invert_yaxis()\n        plt.show()\n\n    return adjusted, inliers","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:17:12.986033Z","iopub.execute_input":"2025-12-19T07:17:12.986840Z","iopub.status.idle":"2025-12-19T07:17:13.006228Z","shell.execute_reply.started":"2025-12-19T07:17:12.986811Z","shell.execute_reply":"2025-12-19T07:17:13.005424Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def format_pred_signal_to_submit_df(pred_sig_y_coord, ecg_id, df_meta):\n    \"\"\"\n    Convert predicted ECG signal tensor into a submission-ready DataFrame.\n\n    Args:\n        pred_sig_y_coord (torch.Tensor): Predicted signal values of shape (N, W).\n        ecg_id (int or str): ECG record ID.\n        df_meta (pd.DataFrame): Metadata DataFrame containing either:\n            - 'lead' and 'number_of_rows' columns for test data\n            - 'fs' column for training data\n\n    Returns:\n        pd.DataFrame: Submission DataFrame with columns:\n            - 'id': formatted as \"{ecg_id}_{row_id}_{lead_name}\"\n            - 'value': resampled signal values\n        The index is set to 'id' and sorted.\n    \"\"\"\n    part1_row_1_3 = pred_sig_y_coord[:4*3, :]\n    part2_row_4 = pred_sig_y_coord[-4:, :]\n    \n    PART1_LEAD_NAMES = [\n        'I', 'aVR', 'V1', 'V4',\n        'II', 'aVL', 'V2', 'V5',\n        'III', 'aVF', 'V3', 'V6',\n    ]\n    \n    signal_dict = {}\n    \n    for lead_name, sig in zip(PART1_LEAD_NAMES, part1_row_1_3):\n        signal_dict[lead_name] = sig.cpu().numpy()\n    \n    # Overwrite II lead with full concatenated signal\n    ii_full_sig = torch.concat([part2_row_4[i] for i in range(4)]).cpu().numpy()\n    signal_dict['II'] = ii_full_sig\n\n    df_list = []\n\n    if 'lead' in df_meta.columns.tolist():\n        # Test data: metadata contains 'lead' and 'number_of_rows'\n        df_meta_tgt_ecg = df_meta[(df_meta['id'] == int(ecg_id))]\n        assert len(df_meta_tgt_ecg) > 0\n        for i, sr_row in df_meta_tgt_ecg.iterrows():\n            lead_name = sr_row.lead\n            target_n_points = sr_row.number_of_rows\n            sig = signal_dict[lead_name]\n            sig_resampled = interp1d(np.linspace(0, 1, len(sig)), sig, kind='linear')(\n                np.linspace(0, 1, target_n_points)\n            )\n            df = pd.DataFrame({\n                'id': [f'{ecg_id}_{row_id}_{lead_name}' for row_id in np.arange(target_n_points)],\n                'value': sig_resampled\n            })\n            df_list.append(df)\n    else:\n        # Training data: metadata contains 'fs'\n        for lead_name, sig in signal_dict.items():\n            df_meta_tgt = df_meta[(df_meta['id'] == int(ecg_id))]\n            assert len(df_meta_tgt) == 1, len(df_meta_tgt)\n            fs = df_meta_tgt.iloc[0].fs\n            if lead_name == 'II':\n                target_n_points = int(fs * 10)\n            else:\n                target_n_points = int(fs * 2.5)\n            sig_resampled = interp1d(np.linspace(0, 1, len(sig)), sig, kind='linear')(\n                np.linspace(0, 1, target_n_points)\n            )\n            df = pd.DataFrame({\n                'id': [f'{ecg_id}_{row_id}_{lead_name}' for row_id in np.arange(target_n_points)],\n                'value': sig_resampled\n            })\n            df_list.append(df)\n\n    df = pd.concat(df_list).set_index('id').sort_index()\n    return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:17:13.146845Z","iopub.execute_input":"2025-12-19T07:17:13.147648Z","iopub.status.idle":"2025-12-19T07:17:13.157879Z","shell.execute_reply.started":"2025-12-19T07:17:13.147623Z","shell.execute_reply":"2025-12-19T07:17:13.156923Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Utility","metadata":{}},{"cell_type":"code","source":"def plot_pred_vs_gt_grid(pred_y, gt_y=None, n_rows=4, n_cols=4, figsize=(12, 8)):\n    \"\"\"\n    Plot predicted vs. ground truth ECG signals in a grid.\n\n    Args:\n        pred_y (np.ndarray or torch.Tensor): Predicted signals, shape (N_LEADS, W)\n        gt_y (np.ndarray or torch.Tensor, optional): Ground truth signals, same shape as pred_y\n        n_rows (int): Number of rows in the grid\n        n_cols (int): Number of columns in the grid\n        figsize (tuple): Figure size\n    \"\"\"\n    if isinstance(pred_y, np.ndarray):\n        pred = pred_y\n        gt = gt_y\n    else:\n        pred = pred_y.cpu().numpy()\n        if gt_y is not None:\n            gt = gt_y.cpu().numpy()\n        else:\n            gt = None\n\n    N_LEADS, W = pred.shape\n    \n    # Create figure with no spacing between subplots\n    fig, axes = plt.subplots(\n        n_rows, n_cols, \n        figsize=figsize,\n        gridspec_kw={'hspace': 0, 'wspace': 0}\n    )\n    axes = axes.flatten()\n\n    for i in range(N_LEADS):\n        row = i // n_cols\n        col = i % n_cols\n        \n        # Plot signals\n        axes[i].plot(pred[i], label='Pred', color='red', linewidth=1.5)\n        if gt is not None:\n            axes[i].plot(gt[i], label='GT', color='blue', linewidth=1.5, alpha=0.7)\n        \n        # Add legend with lead label\n        axes[i].legend(title=f'Lead {i}', loc='upper right', fontsize=8, framealpha=0.9)\n        \n        axes[i].grid(True, alpha=0.3)\n        axes[i].set_ylim(-2.5, 2.5)\n        \n        # Y-axis: only show for first column\n        if col == 0:\n            axes[i].set_ylabel('Amplitude', fontsize=9)\n        else:\n            axes[i].set_yticklabels([])\n        \n        # X-axis: only show for last row\n        if row == n_rows - 1:\n            axes[i].set_xlabel('Time (samples)', fontsize=9)\n        else:\n            axes[i].set_xticklabels([])\n\n    # Hide any extra subplots\n    for j in range(N_LEADS, n_rows * n_cols):\n        axes[j].axis('off')\n\n    plt.tight_layout(pad=0.5)\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:17:21.787534Z","iopub.execute_input":"2025-12-19T07:17:21.788233Z","iopub.status.idle":"2025-12-19T07:17:21.797315Z","shell.execute_reply.started":"2025-12-19T07:17:21.788204Z","shell.execute_reply":"2025-12-19T07:17:21.796277Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"def extract_y_coord_from_mask(\n    mask_probs: torch.Tensor,\n    threshold: float = 0.3,\n    sharpening: float = None,  # None → do not apply sharpening\n):\n    \"\"\"\n    Extracts the y-coordinate from mask probabilities using soft-argmax along the height axis.\n\n    Args:\n        mask_probs (torch.Tensor): Tensor of shape (B, N_LEADS, 1, H, W).\n        threshold (float, optional): Values in mask_probs below this threshold are set to 0.\n        sharpening (float, optional): >1 sharpens the distribution. None means no sharpening.\n\n    Returns:\n        torch.Tensor: y-coordinates of shape (B, N_LEADS, W), centered and inverted.\n    \n    Note:\n        mask_probs will be detached during computation.\n    \"\"\"\n    # Squeeze channel dimension if exists\n    if mask_probs.dim() == 5:\n        mask_probs = mask_probs.squeeze(2)  # (B, N_LEADS, H, W)\n    B, N_LEADS, H, W = mask_probs.shape\n\n    # --- optional thresholding ---\n    if threshold is not None:\n        mask_probs = mask_probs.clone()\n        mask_probs[mask_probs < threshold] = 0.0\n\n    # --- optional sharpening ---\n    if sharpening is not None and sharpening > 1:\n        p = (mask_probs + 1e-8) ** sharpening\n    else:\n        p = mask_probs\n\n    # --- normalize over height axis ---\n    col_sum = p.sum(dim=2, keepdim=True)  # (B, N_LEADS, 1, W)\n    p_norm = p / (col_sum + 1e-6)\n\n    # --- weighted mean (soft-argmax along height) ---\n    ys = torch.arange(H, device=mask_probs.device).view(1, 1, H, 1)\n    y_coord = (p_norm * ys).sum(dim=2)  # (B, N_LEADS, W)\n\n    # --- mask columns that were all zero ---\n    all_zero_cols = (col_sum.squeeze(2) == 0)  # (B, N_LEADS, W)\n    y_coord = y_coord.masked_fill(all_zero_cols, float('nan'))\n\n    # --- linear interpolation to fill NaNs ---\n    y_coord_np = y_coord.detach().cpu().numpy()\n    for b in range(B):\n        for l in range(N_LEADS):\n            y_series = y_coord_np[b, l]\n            if np.isnan(y_series).all():\n                # If all NaNs, set to center line\n                y_series[:] = H // 2\n                y_coord_np[b, l] = y_series\n                continue\n            nans = np.isnan(y_series)\n            not_nans = ~nans\n            # Fill NaNs at edges\n            if nans[0]:\n                first_valid = np.flatnonzero(not_nans)[0]\n                y_series[:first_valid] = y_series[first_valid]\n            if nans[-1]:\n                last_valid = np.flatnonzero(not_nans)[-1]\n                y_series[last_valid+1:] = y_series[last_valid]\n            # Linear interpolation\n            y_series[nans] = np.interp(\n                np.flatnonzero(nans),\n                np.flatnonzero(not_nans),\n                y_series[not_nans]\n            )\n            y_coord_np[b, l] = y_series\n\n    y_coord = torch.from_numpy(y_coord_np).to(mask_probs.device)\n\n    # Optional transformation (original spec)\n    y_coord -= H // 2\n    return -y_coord\n\n\ndef extract_y_coord_argmax_continuity_with_fillna(\n    mask_probs: torch.Tensor,\n    threshold: float = 0.3,\n    bonus_sigma: float = SIG_SEG_IMAGE_SIZE[0] / 10,  # 12.8\n    bonus_weight: float = 1.0,\n    bonus_decay: str = 'quadratic',  # 'exp', 'linear', 'quadratic'\n    bonus_min: float = 0.001,  # minimum bonus\n    sharpening: float = None,\n    fill_strategy: str = 'forward_backward',  # 'forward_backward' or 'interpolate'\n    max_extrap_dist: float = SIG_SEG_IMAGE_SIZE[0] / 36,  # 10.67\n):\n    \"\"\"\n    Extracts y-coordinate using column-wise argmax with continuity bonus and fills empty columns in advance.\n\n    Args:\n        mask_probs (torch.Tensor): Tensor of shape (B, N, H, W) or (B, N, 1, H, W).\n        threshold (float, optional): Values below this threshold are set to 0.\n        bonus_sigma (float, optional): Sigma for continuity bonus calculation.\n        bonus_weight (float, optional): Weight of continuity bonus.\n        bonus_decay (str, optional): Decay type for bonus ('exp', 'linear', 'quadratic').\n        bonus_min (float, optional): Minimum bonus value.\n        sharpening (float, optional): >1 sharpens the distribution. None means no sharpening.\n        fill_strategy (str, optional): Strategy to fill empty columns.\n            - 'forward_backward': forward fill → backward fill\n            - 'interpolate': linear interpolation\n        max_extrap_dist (float, optional): Maximum extrapolation distance.\n\n    Returns:\n        torch.Tensor: y-coordinates of shape (B, N, W), centered and inverted.\n    \"\"\"\n    # ========== SHAPE FIX ==========\n    if mask_probs.dim() == 5:\n        mask_probs = mask_probs.squeeze(2)  # (B,N,H,W)\n\n    B, N, H, W = mask_probs.shape\n    device = mask_probs.device\n\n    # ========== SHARPENING (optional) ==========\n    p = mask_probs.clone()\n    if sharpening is not None and sharpening > 1:\n        p = (p + 1e-8) ** sharpening\n\n    # ========== THRESHOLD ==========\n    if threshold is not None:\n        p[p < threshold] = 0.0\n\n    # ========== FILL EMPTY COLUMNS ==========\n    if fill_strategy == 'forward_backward':\n        p = _fill_columns_forward_backward(p)\n    elif fill_strategy == 'interpolate':\n        p = _fill_columns_interpolate(p, threshold)\n    else:\n        raise ValueError(f\"Unknown fill_strategy: {fill_strategy}\")\n\n    # ========== SETUP ==========\n    ys = torch.arange(H, device=device, dtype=torch.float32).view(1, 1, H, 1)  # (1,1,H,1)\n    y_argmax = torch.zeros(B, N, W, device=device, dtype=torch.float32)\n\n    # ========== COLUMN-WISE ARGMAX WITH CONTINUITY BONUS ==========\n    for x in range(W):\n        col = p[:, :, :, x]  # (B, N, H)\n\n        if x == 0:\n            # --- First column: pure argmax without bonus ---\n            weighted = col\n\n        elif x == 1:\n            # --- Second column: higher bonus for positions close to previous column ---\n            y_prev = y_argmax[:, :, 0].unsqueeze(-1)  # (B, N, 1)\n            dist = torch.abs(ys.squeeze(-1) - y_prev)  # (B, N, H)\n\n            # Select decay function\n            if bonus_decay == 'exp':\n                bonus = torch.clamp(torch.exp(-(dist ** 2) / (2 * bonus_sigma ** 2)), min=bonus_min)\n            elif bonus_decay == 'linear':\n                bonus = torch.clamp(1.0 - dist / bonus_sigma, min=bonus_min)\n            elif bonus_decay == 'quadratic':\n                bonus = torch.clamp(1.0 - (dist / bonus_sigma) ** 2, min=bonus_min)\n            else:\n                raise ValueError(f\"Unknown bonus_decay: {bonus_decay}\")\n\n            # Apply bonus\n            weighted = col * (1.0 + bonus_weight * (bonus - 1.0))\n\n        else:\n            # --- Third column onwards: higher bonus for positions close to extrapolated line from previous two columns ---\n            y1 = y_argmax[:, :, x - 1].unsqueeze(-1)  # (B, N, 1)\n            y0 = y_argmax[:, :, x - 2].unsqueeze(-1)  # (B, N, 1)\n\n            # Limit extrapolation distance\n            delta = y1 - y0\n            if max_extrap_dist is not None:\n                delta = torch.clamp(delta, -max_extrap_dist, max_extrap_dist)\n\n            y_extrap = y1 + delta  # extrapolated prediction\n\n            # Clip to image range\n            y_extrap = torch.clamp(y_extrap, 0, H - 1)\n\n            dist = torch.abs(ys.squeeze(-1) - y_extrap)  # (B, N, H)\n\n            # Select decay function\n            if bonus_decay == 'exp':\n                bonus = torch.clamp(torch.exp(-(dist ** 2) / (2 * bonus_sigma ** 2)), min=bonus_min)\n            elif bonus_decay == 'linear':\n                bonus = torch.clamp(1.0 - dist / bonus_sigma, min=bonus_min)\n            elif bonus_decay == 'quadratic':\n                bonus = torch.clamp(1.0 - (dist / bonus_sigma) ** 2, min=bonus_min)\n            else:\n                raise ValueError(f\"Unknown bonus_decay: {bonus_decay}\")\n\n            # Apply bonus\n            weighted = col * (1.0 + bonus_weight * (bonus - 1.0))\n\n        # Check if the column is entirely zero\n        col_sum = weighted.sum(dim=2)  # (B, N)\n        is_empty = (col_sum == 0)  # (B, N)\n        is_std_zero = (weighted.std(dim=2) < 1e-7)\n\n        # Argmax\n        y_current = weighted.argmax(dim=2).float()\n\n        if is_std_zero.any():\n            y_current[is_std_zero] = H // 2\n\n        # Handle empty columns just in case\n        if is_empty.any():\n            if x == 0:\n                y_current = torch.where(is_empty, torch.tensor(H // 2, dtype=torch.float32, device=device), y_current)\n            elif x == 1:\n                y_current = torch.where(is_empty, y_argmax[:, :, 0], y_current)\n            else:\n                y1 = y_argmax[:, :, x - 1]\n                y0 = y_argmax[:, :, x - 2]\n                y_extrap_fill = torch.clamp(y1 + (y1 - y0), 0, H - 1)\n                y_current = torch.where(is_empty, y_extrap_fill, y_current)\n\n        y_argmax[:, :, x] = y_current\n\n    # ========== COORDINATE TRANSFORM ==========\n    y_coord = -(y_argmax - (H // 2))\n\n    return y_coord\n\n\ndef _fill_columns_forward_backward(p: torch.Tensor) -> torch.Tensor:\n    \"\"\"\n    Fill empty columns using forward fill followed by backward fill.\n\n    Args:\n        p (torch.Tensor): Tensor of shape (B, N, H, W)\n\n    Returns:\n        torch.Tensor: Filled tensor of shape (B, N, H, W)\n    \"\"\"\n    B, N, H, W = p.shape\n    p_filled = p.clone()\n\n    for b in range(B):\n        for n in range(N):\n            # Check if each column has any non-zero element\n            col_valid = (p[b, n].sum(dim=0) > 0)  # (W,)\n\n            # Forward fill (left to right)\n            last_valid_col = None\n            for x in range(W):\n                if col_valid[x]:\n                    last_valid_col = x\n                elif last_valid_col is not None:\n                    p_filled[b, n, :, x] = p_filled[b, n, :, last_valid_col]\n\n            # Backward fill (right to left)\n            next_valid_col = None\n            for x in reversed(range(W)):\n                if col_valid[x]:\n                    next_valid_col = x\n                elif next_valid_col is not None:\n                    p_filled[b, n, :, x] = p_filled[b, n, :, next_valid_col]\n\n    return p_filled\n\n\ndef _fill_columns_interpolate(p: torch.Tensor, threshold: float) -> torch.Tensor:\n    \"\"\"\n    Fill empty columns using linear interpolation via NumPy.\n\n    Args:\n        p (torch.Tensor): Tensor of shape (B, N, H, W)\n        threshold (float): Threshold value for interpolation scaling\n\n    Returns:\n        torch.Tensor: Filled tensor of shape (B, N, H, W)\n    \"\"\"\n    B, N, H, W = p.shape\n    device = p.device\n\n    # Get representative y-coordinate per column using argmax\n    p_np = p.detach().cpu().numpy()\n    p_filled = p.clone()\n\n    for b in range(B):\n        for n in range(N):\n            col_sum = p_np[b, n].sum(axis=0)  # (W,)\n            valid_cols = col_sum > 0\n\n            if not valid_cols.any():\n                # If all columns are empty, place a small Gaussian at the center\n                center = H // 2\n                y_grid = np.arange(H).reshape(-1, 1)\n                gaussian = np.exp(-((y_grid - center) ** 2) / (2 * 10 ** 2))\n                p_filled[b, n] = torch.from_numpy(gaussian).to(device)\n                continue\n\n            # Get argmax of valid columns\n            y_positions = np.zeros(W)\n            for x in range(W):\n                if valid_cols[x]:\n                    y_positions[x] = p_np[b, n, :, x].argmax()\n\n            # Linear interpolation\n            invalid_cols = ~valid_cols\n            if invalid_cols.any():\n                valid_x = np.where(valid_cols)[0]\n                invalid_x = np.where(invalid_cols)[0]\n\n                # Handle edges\n                if invalid_cols[0]:\n                    first_valid = valid_x[0]\n                    y_positions[:first_valid] = y_positions[first_valid]\n                if invalid_cols[-1]:\n                    last_valid = valid_x[-1]\n                    y_positions[last_valid + 1:] = y_positions[last_valid]\n\n                # Interpolate intermediate positions\n                invalid_x = invalid_x[(invalid_x > valid_x[0]) & (invalid_x < valid_x[-1])]\n                if len(invalid_x) > 0:\n                    y_positions[invalid_x] = np.interp(invalid_x, valid_x, y_positions[valid_x])\n\n                # Place Gaussian at interpolated y positions\n                for x in invalid_x:\n                    y_center = y_positions[x]\n                    y_grid = np.arange(H)\n                    gaussian = np.exp(-((y_grid - y_center) ** 2) / (2 * 5 ** 2))\n                    p_filled[b, n, :, x] = torch.from_numpy(gaussian).to(device) * threshold * 2\n\n    return p_filled","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:17:22.279649Z","iopub.execute_input":"2025-12-19T07:17:22.279964Z","iopub.status.idle":"2025-12-19T07:17:22.310998Z","shell.execute_reply.started":"2025-12-19T07:17:22.279940Z","shell.execute_reply":"2025-12-19T07:17:22.310312Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def patch_first_conv(model, new_in_channels, default_in_channels=3, pretrained=True):\n    \"\"\"Modify the first convolution layer to accept a different number of input channels.\n\n    Handles cases where:\n        - new_in_channels == 1 or 2 → reuse the original weights\n        - new_in_channels > 3 → initialize weights randomly using Kaiming normal\n\n    Args:\n        model (nn.Module): The model containing the convolution to patch.\n        new_in_channels (int): Number of input channels for the new convolution.\n        default_in_channels (int, optional): Number of channels in the original convolution. Defaults to 3.\n        pretrained (bool, optional): Whether to reuse pretrained weights. Defaults to True.\n    \"\"\"\n\n    # Find first convolution layer with default_in_channels\n    for module in model.modules():\n        if isinstance(module, nn.Conv2d) and module.in_channels == default_in_channels:\n            break\n\n    weight = module.weight.detach()\n    module.in_channels = new_in_channels\n\n    if not pretrained:\n        module.weight = nn.parameter.Parameter(\n            torch.Tensor(\n                module.out_channels,\n                new_in_channels // module.groups,\n                *module.kernel_size,\n            )\n        )\n        module.reset_parameters()\n\n    elif new_in_channels == 1:\n        # Sum weights across input channels for grayscale conversion\n        new_weight = weight.sum(1, keepdim=True)\n        module.weight = nn.parameter.Parameter(new_weight)\n\n    else:\n        # Repeat existing weights for new channels\n        new_weight = torch.Tensor(\n            module.out_channels, new_in_channels // module.groups, *module.kernel_size\n        )\n\n        for i in range(new_in_channels):\n            new_weight[:, i] = weight[:, i % default_in_channels]\n\n        new_weight = new_weight * (default_in_channels / new_in_channels)\n        module.weight = nn.parameter.Parameter(new_weight)\n\n\ndef replace_strides_with_dilation(module, dilation_rate):\n    \"\"\"Replace strides with dilation in all Conv2d modules.\n\n    Args:\n        module (nn.Module): Module containing Conv2d layers to patch.\n        dilation_rate (int): Dilation rate to set.\n    \"\"\"\n    for mod in module.modules():\n        if isinstance(mod, nn.Conv2d):\n            mod.stride = (1, 1)\n            mod.dilation = (dilation_rate, dilation_rate)\n            kh, kw = mod.kernel_size\n            mod.padding = ((kh // 2) * dilation_rate, (kh // 2) * dilation_rate)\n\n            # Special handling for EfficientNet static padding\n            if hasattr(mod, \"static_padding\"):\n                mod.static_padding = nn.Identity()\n\n\ntry:\n    from inplace_abn import InPlaceABN\nexcept ImportError:\n    InPlaceABN = None\n\n\nclass Conv2dReLU(nn.Sequential):\n    \"\"\"Convolution followed by optional batch normalization and ReLU activation.\"\"\"\n\n    def __init__(\n        self,\n        in_channels,\n        out_channels,\n        kernel_size,\n        padding=0,\n        stride=1,\n        use_batchnorm=True,\n    ):\n        if use_batchnorm == \"inplace\" and InPlaceABN is None:\n            raise RuntimeError(\n                \"To use `use_batchnorm='inplace'`, inplace_abn must be installed. \"\n                \"See: https://github.com/mapillary/inplace_abn\"\n            )\n\n        conv = nn.Conv2d(\n            in_channels,\n            out_channels,\n            kernel_size,\n            stride=stride,\n            padding=padding,\n            bias=not (use_batchnorm),\n        )\n        relu = nn.ReLU(inplace=True)\n\n        if use_batchnorm == \"inplace\":\n            bn = InPlaceABN(out_channels, activation=\"leaky_relu\", activation_param=0.0)\n            relu = nn.Identity()\n        elif use_batchnorm and use_batchnorm != \"inplace\":\n            bn = nn.BatchNorm2d(out_channels)\n        else:\n            bn = nn.Identity()\n\n        super(Conv2dReLU, self).__init__(conv, bn, relu)\n\n\nclass SCSEModule(nn.Module):\n    \"\"\"Concurrent Spatial and Channel Squeeze & Excitation Module.\"\"\"\n\n    def __init__(self, in_channels, reduction=16):\n        super().__init__()\n        self.cSE = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Conv2d(in_channels, in_channels // reduction, 1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(in_channels // reduction, in_channels, 1),\n            nn.Sigmoid(),\n        )\n        self.sSE = nn.Sequential(nn.Conv2d(in_channels, 1, 1), nn.Sigmoid())\n\n    def forward(self, x):\n        \"\"\"Forward pass applying both cSE and sSE attention.\"\"\"\n        return x * self.cSE(x) + x * self.sSE(x)\n\n\nclass ArgMax(nn.Module):\n    \"\"\"Module that computes argmax along the specified dimension.\"\"\"\n\n    def __init__(self, dim=None):\n        super().__init__()\n        self.dim = dim\n\n    def forward(self, x):\n        return torch.argmax(x, dim=self.dim)\n\n\nclass Clamp(nn.Module):\n    \"\"\"Clamp tensor values between min and max.\"\"\"\n\n    def __init__(self, min=0, max=1):\n        super().__init__()\n        self.min, self.max = min, max\n\n    def forward(self, x):\n        return torch.clamp(x, self.min, self.max)\n\n\nclass Activation(nn.Module):\n    \"\"\"Flexible activation module supporting various activations.\"\"\"\n\n    def __init__(self, name, **params):\n        super().__init__()\n\n        if name is None or name == \"identity\":\n            self.activation = nn.Identity(**params)\n        elif name == \"sigmoid\":\n            self.activation = nn.Sigmoid()\n        elif name == \"softmax2d\":\n            self.activation = nn.Softmax(dim=1, **params)\n        elif name == \"softmax\":\n            self.activation = nn.Softmax(**params)\n        elif name == \"logsoftmax\":\n            self.activation = nn.LogSoftmax(**params)\n        elif name == \"tanh\":\n            self.activation = nn.Tanh()\n        elif name == \"argmax\":\n            self.activation = ArgMax(**params)\n        elif name == \"argmax2d\":\n            self.activation = ArgMax(dim=1, **params)\n        elif name == \"clamp\":\n            self.activation = Clamp(**params)\n        elif callable(name):\n            self.activation = name(**params)\n        else:\n            raise ValueError(\n                f\"Activation must be callable or one of sigmoid/softmax/logsoftmax/tanh/\"\n                f\"argmax/argmax2d/clamp/None; got {name}\"\n            )\n\n    def forward(self, x):\n        return self.activation(x)\n\n\nclass Attention(nn.Module):\n    \"\"\"Attention module wrapper supporting multiple types.\"\"\"\n\n    def __init__(self, name, **params):\n        super().__init__()\n\n        if name is None:\n            self.attention = nn.Identity(**params)\n        elif name == \"scse\":\n            self.attention = SCSEModule(**params)\n        else:\n            raise ValueError(f\"Attention {name} is not implemented\")\n\n    def forward(self, x):\n        return self.attention(x)\n\n\nclass DecoderBlock(nn.Module):\n    \"\"\"Decoder block with optional attention and Conv2dReLU layers.\"\"\"\n\n    def __init__(\n        self,\n        in_channels,\n        skip_channels,\n        out_channels,\n        use_batchnorm=True,\n        attention_type=None,\n    ):\n        super().__init__()\n        self.conv1 = Conv2dReLU(\n            in_channels + skip_channels,\n            out_channels,\n            kernel_size=3,\n            padding=1,\n            use_batchnorm=use_batchnorm,\n        )\n        self.attention1 = Attention(\n            attention_type, in_channels=in_channels + skip_channels\n        )\n        self.conv2 = Conv2dReLU(\n            out_channels,\n            out_channels,\n            kernel_size=3,\n            padding=1,\n            use_batchnorm=use_batchnorm,\n        )\n        self.attention2 = Attention(attention_type, in_channels=out_channels)\n\n    def forward(self, x, skip=None):\n        \"\"\"Forward pass with optional skip connection and attention.\"\"\"\n        x = F.interpolate(x, scale_factor=2, mode=\"nearest\")\n        if skip is not None:\n            x = torch.cat([x, skip], dim=1)\n            x = self.attention1(x)\n        x = self.conv1(x)\n        x = self.conv2(x)\n        x = self.attention2(x)\n        return x\n\n\nclass CenterBlock(nn.Sequential):\n    \"\"\"Center block of UNet composed of two Conv2dReLU layers.\"\"\"\n\n    def __init__(self, in_channels, out_channels, use_batchnorm=True):\n        conv1 = Conv2dReLU(\n            in_channels,\n            out_channels,\n            kernel_size=3,\n            padding=1,\n            use_batchnorm=use_batchnorm,\n        )\n        conv2 = Conv2dReLU(\n            out_channels,\n            out_channels,\n            kernel_size=3,\n            padding=1,\n            use_batchnorm=use_batchnorm,\n        )\n        super().__init__(conv1, conv2)\n\n\nclass UnetDecoder(nn.Module):\n    \"\"\"UNet decoder module composed of multiple DecoderBlocks.\"\"\"\n\n    def __init__(\n        self,\n        encoder_channels,\n        decoder_channels,\n        n_blocks=5,\n        use_batchnorm=True,\n        attention_type=None,\n        center=False,\n    ):\n        super().__init__()\n\n        if n_blocks != len(decoder_channels):\n            raise ValueError(\n                \"Model depth is {}, but `decoder_channels` is provided for {} blocks.\".format(\n                    n_blocks, len(decoder_channels)\n                )\n            )\n\n        # Remove first skip connection (same spatial resolution)\n        encoder_channels = encoder_channels[1:]\n        # Reverse channels to start from the encoder's head\n        encoder_channels = encoder_channels[::-1]\n\n        # Compute block input/output channels\n        head_channels = encoder_channels[0]\n        in_channels = [head_channels] + list(decoder_channels[:-1])\n        skip_channels = list(encoder_channels[1:]) + [0]\n        out_channels = decoder_channels\n\n        if center:\n            self.center = CenterBlock(\n                head_channels, head_channels, use_batchnorm=use_batchnorm\n            )\n        else:\n            self.center = nn.Identity()\n\n        # Create decoder blocks\n        kwargs = dict(use_batchnorm=use_batchnorm, attention_type=attention_type)\n        blocks = [\n            DecoderBlock(in_ch, skip_ch, out_ch, **kwargs)\n            for in_ch, skip_ch, out_ch in zip(in_channels, skip_channels, out_channels)\n        ]\n        self.blocks = nn.ModuleList(blocks)\n\n    def forward(self, *features):\n        \"\"\"Forward pass through decoder.\n\n        Args:\n            features: Tuple of encoder feature maps.\n\n        Returns:\n            Tensor: Decoded feature map.\n        \"\"\"\n        features = features[1:]  # remove first skip with same spatial resolution\n        features = features[::-1]  # reverse channels to start from head of encoder\n\n        head = features[0]\n        skips = features[1:]\n\n        x = self.center(head)\n        for i, decoder_block in enumerate(self.blocks):\n            skip = skips[i] if i < len(skips) else None\n            x = decoder_block(x, skip)\n\n        return x\n\n\nclass SegmentationHead(nn.Sequential):\n    \"\"\"Segmentation head with optional upsampling and activation.\"\"\"\n\n    def __init__(\n        self, in_channels, out_channels, kernel_size=3, activation=None, upsampling=1\n    ):\n        conv2d = nn.Conv2d(\n            in_channels, out_channels, kernel_size=kernel_size, padding=kernel_size // 2\n        )\n        upsampling_layer = (\n            nn.UpsamplingBilinear2d(scale_factor=upsampling)\n            if upsampling > 1\n            else nn.Identity()\n        )\n        activation_layer = Activation(activation)\n        super().__init__(conv2d, upsampling_layer, activation_layer)\n\n\nclass ClassificationHead(nn.Sequential):\n    \"\"\"Classification head with pooling, dropout, linear layer, and activation.\"\"\"\n\n    def __init__(\n        self, in_channels, classes, pooling=\"avg\", dropout=0.2, activation=None\n    ):\n        if pooling not in (\"max\", \"avg\"):\n            raise ValueError(\n                \"Pooling should be one of ('max', 'avg'), got {}.\".format(pooling)\n            )\n        pool = nn.AdaptiveAvgPool2d(1) if pooling == \"avg\" else nn.AdaptiveMaxPool2d(1)\n        flatten = nn.Flatten()\n        dropout_layer = nn.Dropout(p=dropout, inplace=True) if dropout else nn.Identity()\n        linear = nn.Linear(in_channels, classes, bias=True)\n        activation_layer = Activation(activation)\n        super().__init__(pool, flatten, dropout_layer, linear, activation_layer)\n\n\nclass DenoisingYCoordExtractionWeightedAvgHeadV3(nn.Module):\n    \"\"\"Lightweight Y-coordinate extraction head using Conv1D, GRU, and weighted average.\n\n    This module predicts denoised y-coordinates from a probability mask by applying\n    Conv1D feature extraction, a bidirectional GRU, and a weighted average over the\n    y-axis. Learnable scale and offset parameters allow fine adjustment of outputs.\n    \"\"\"\n\n    def __init__(\n        self,\n        hidden_dim: int = 256,\n        gru_hidden_dim: int = 256,\n        gru_num_layers: int = 2,\n        dropout: float = 0.1,\n        fixed_y_coord_scale: float = 2 * 0.01,\n        input_h_dim: int = 256,  # Input height of the image (dynamic)\n    ):\n        \"\"\"\n        Args:\n            hidden_dim (int): Hidden dimension for Conv1D layers.\n            gru_hidden_dim (int): Hidden dimension for GRU layers.\n            gru_num_layers (int): Number of GRU layers.\n            dropout (float): Dropout rate for Conv1D, GRU, and linear layers.\n            fixed_y_coord_scale (float): Scaling factor applied to predicted y-coordinate.\n            input_h_dim (int): Input height dimension H of the image.\n        \"\"\"\n        super().__init__()\n        self.fixed_y_coord_scale = fixed_y_coord_scale\n        self.input_h_dim = input_h_dim\n        # ---------------------------------------------------------------------\n        # Conv1D feature extraction along H dimension\n        # ---------------------------------------------------------------------\n        self.conv_layers = nn.Sequential(\n            nn.Conv1d(input_h_dim, hidden_dim, kernel_size=5, padding=2),\n            nn.BatchNorm1d(hidden_dim),\n            nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Conv1d(hidden_dim, hidden_dim, kernel_size=5, padding=2),\n            nn.BatchNorm1d(hidden_dim),\n            nn.GELU(),\n            nn.Dropout(dropout),\n        )\n        # ---------------------------------------------------------------------\n        # Bidirectional GRU for sequential modeling along width\n        # ---------------------------------------------------------------------\n        self.gru = nn.GRU(\n            input_size=hidden_dim,\n            hidden_size=gru_hidden_dim,\n            num_layers=gru_num_layers,\n            batch_first=True,\n            dropout=dropout if gru_num_layers > 1 else 0.0,\n            bidirectional=True,\n        )\n        # ---------------------------------------------------------------------\n        # Denoised attention / score head to produce height-wise weights\n        # ---------------------------------------------------------------------\n        self.denoised_score_head = nn.Sequential(\n            nn.Linear(gru_hidden_dim * 2, gru_hidden_dim),\n            nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(gru_hidden_dim, input_h_dim),\n            nn.Softmax(dim=-1),\n        )\n        # ---------------------------------------------------------------------\n        # Learnable scale and offset parameters for final y-coordinate\n        # ---------------------------------------------------------------------\n        self.y_scale_param = nn.Parameter(torch.tensor(1.0))\n        self.y_offset_param = nn.Parameter(torch.tensor(0.0))\n\n    def forward(self, prob_mask: torch.Tensor):\n        \"\"\"\n        Forward pass to extract y-coordinate from probability mask.\n\n        Args:\n            prob_mask (torch.Tensor): Input mask of shape (B, 1, H, W).\n\n        Returns:\n            torch.Tensor: Predicted y-coordinate of shape (B, W).\n        \"\"\"\n        B, C, H, W = prob_mask.shape\n        assert C == 1\n        # Remove channel dimension: (B, 1, H, W) -> (B, H, W)\n        x = prob_mask.squeeze(1)\n        # ---------------------------------------------------------------------\n        # Conv1D feature extraction\n        # ---------------------------------------------------------------------\n        x = self.conv_layers(x)  # (B, hidden_dim, W)\n        # Transpose for GRU: (B, W, hidden_dim)\n        x = x.transpose(2, 1)\n        # ---------------------------------------------------------------------\n        # GRU forward\n        # ---------------------------------------------------------------------\n        gru_out, _ = self.gru(x)  # (B, W, gru_hidden_dim * 2)\n        # ---------------------------------------------------------------------\n        # Denoised attention / score prediction over height\n        # ---------------------------------------------------------------------\n        denoised_score = self.denoised_score_head(gru_out)  # (B, W, H)\n        # ---------------------------------------------------------------------\n        # Compute y-coordinate via weighted average along height\n        # ---------------------------------------------------------------------\n        y_grid = torch.arange(H, device=x.device).float() / H  # normalized grid [0,1)\n        y_grid = -(y_grid - 0.5) * H  # map to [-H/2, H/2]\n        y_grid = y_grid.view(1, 1, H)  # (1, 1, H)\n        # Weighted sum along H dimension\n        pred_y_coord_normalized = (y_grid * denoised_score).sum(dim=-1)\n        # Apply learnable scale and offset\n        pred_y_coord = (\n            self.y_offset_param\n            + self.fixed_y_coord_scale * torch.sigmoid(self.y_scale_param) * pred_y_coord_normalized\n        )\n\n        return pred_y_coord","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:17:22.312661Z","iopub.execute_input":"2025-12-19T07:17:22.313128Z","iopub.status.idle":"2025-12-19T07:17:22.354351Z","shell.execute_reply.started":"2025-12-19T07:17:22.313094Z","shell.execute_reply":"2025-12-19T07:17:22.353580Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class KeypointSegModelV2(nn.Module):\n    \"\"\"\n    Keypoint and Part Affinity Field (PAF) segmentation model.\n\n    This model uses a U-Net style architecture with two separate decoders:\n    - One decoder for keypoint heatmaps\n    - One decoder for part affinity fields (connections between keypoints)\n\n    The model concatenates keypoint and PAF outputs to produce the final\n    multi-channel segmentation mask.\n    \"\"\"\n\n    def __init__(\n        self,\n        # encoder_out_channels,  # must contain input raw image channels (usually 3) at the beginning of the list\n        num_classes=N_KPTS + 2 * N_CONNS,\n        encoder_model_name=\"timm-efficientnet-b0\",\n        decoder_channels=(256, 128, 64, 32, 16),\n    ):\n        \"\"\"\n        Args:\n            num_classes (int): Total number of output channels (keypoints + PAFs).\n            encoder_model_name (str): Encoder backbone name from timm.\n            decoder_channels (tuple[int]): Channel configuration for U-Net decoders.\n        \"\"\"\n        super().__init__()\n        # ---------------------------------------------------------------------\n        # Encoder\n        # ---------------------------------------------------------------------\n        self._encoder = timm.create_model(\n            encoder_model_name, features_only=True, pretrained=False\n        )\n        # Include raw input image channels at the beginning\n        encoder_out_channels = [3] + self._encoder.feature_info.channels()\n        print(encoder_out_channels)\n        # ---------------------------------------------------------------------\n        # Ensure decoder_channels length matches encoder stages\n        # ---------------------------------------------------------------------\n        if len(decoder_channels) > len(encoder_out_channels) - 1:\n            print(\n                f\"len(decoder_channels) > len(encoder_out_channels) - 1, deleting tail decoder_channels\"\n            )\n            print(\"[WARNING] predicted mask will be upscaled by interpolation\")\n            decoder_channels = decoder_channels[: len(encoder_out_channels) - 1]\n        # ---------------------------------------------------------------------\n        # Decoders for keypoints and PAFs\n        # ---------------------------------------------------------------------\n        self._decoder_kpt = UnetDecoder(\n            encoder_channels=encoder_out_channels,\n            decoder_channels=decoder_channels,\n            n_blocks=len(decoder_channels),\n            use_batchnorm=True,\n            center=True if encoder_model_name.startswith(\"vgg\") else False,\n            attention_type=None,\n        )\n        self._decoder_paf = UnetDecoder(\n            encoder_channels=encoder_out_channels,\n            decoder_channels=decoder_channels,\n            n_blocks=len(decoder_channels),\n            use_batchnorm=True,\n            center=True if encoder_model_name.startswith(\"vgg\") else False,\n            attention_type=None,\n        )\n        # ---------------------------------------------------------------------\n        # Segmentation heads\n        # ---------------------------------------------------------------------\n        n_kpts = N_KPTS\n        n_pafs = 2 * N_CONNS\n        self._segmentation_head_kpt = SegmentationHead(\n            in_channels=decoder_channels[-1],\n            out_channels=n_kpts,\n            activation=None,\n            kernel_size=3,\n        )\n        self._segmentation_head_paf = SegmentationHead(\n            in_channels=decoder_channels[-1],\n            out_channels=n_pafs,\n            activation=None,\n            kernel_size=3,\n        )\n\n    def forward(self, x):\n        \"\"\"\n        Forward pass for keypoint and PAF segmentation.\n\n        Args:\n            x (torch.Tensor): Input tensor of shape (B, C, H, W).\n\n        Returns:\n            torch.Tensor: Concatenated keypoint and PAF masks of shape\n            (B, N_KPTS + 2*N_CONNS, H, W).\n        \"\"\"\n        bs, ch, h, w = x.shape\n        # ---------------------------------------------------------------------\n        # Encoder forward\n        # ---------------------------------------------------------------------\n        feat_list = self._encoder(x)\n        # Prepend raw input image to feature list\n        feat_list = [x] + feat_list\n        # ---------------------------------------------------------------------\n        # Decoder forward\n        # ---------------------------------------------------------------------\n        dec_out_kpt = self._decoder_kpt(*feat_list)\n        dec_out_paf = self._decoder_paf(*feat_list)\n        # ---------------------------------------------------------------------\n        # Segmentation heads forward\n        # ---------------------------------------------------------------------\n        pred_mask_kpt = self._segmentation_head_kpt(dec_out_kpt)\n        pred_mask_paf = self._segmentation_head_paf(dec_out_paf)\n        # Concatenate keypoints and PAFs along channel dimension\n        pred_mask = torch.cat([pred_mask_kpt, pred_mask_paf], dim=1)\n        # ---------------------------------------------------------------------\n        # Upsample if output size does not match input size\n        # ---------------------------------------------------------------------\n        if x.shape[-2:] != pred_mask.shape[-2:]:\n            pred_mask = torch.nn.functional.interpolate(pred_mask, size=x.shape[-2:], mode=\"bilinear\")\n\n        return pred_mask\n\n\n# class SignalExtractorModel(nn.Module):\n#     \"\"\"Signal segmentation model with optional y-coordinate extraction from masks.\"\"\"\n\n#     def __init__(\n#         self,\n#         num_classes=1,\n#         encoder_model_name=\"timm-efficientnet-b0\",\n#         decoder_channels=(256, 128, 64, 32, 16),\n#         fixed_y_coord_scale=2 * 0.01,\n#         detach_y_coord_path=True,\n#     ):\n#         \"\"\"\n#         Args:\n#             num_classes (int): Number of output channels for segmentation.\n#             encoder_model_name (str): Name of encoder model from timm.\n#             decoder_channels (tuple[int]): Decoder channel configuration.\n#             fixed_y_coord_scale (float): Scaling factor for predicted y-coordinates.\n#             detach_y_coord_path (bool): Whether to detach y-coordinate path during backprop.\n#         \"\"\"\n#         super().__init__()\n#         self._encoder = timm.create_model(encoder_model_name, features_only=True, pretrained=False)\n#         encoder_out_channels = [3] + self._encoder.feature_info.channels()\n#         self._decoder = UnetDecoder(\n#             encoder_channels=encoder_out_channels,\n#             decoder_channels=decoder_channels,\n#             n_blocks=len(decoder_channels),\n#             use_batchnorm=True,\n#             center=True if encoder_model_name.startswith(\"vgg\") else False,\n#             attention_type=None,\n#         )\n#         self._segmentation_head = SegmentationHead(\n#             in_channels=decoder_channels[-1],\n#             out_channels=num_classes,\n#             activation=None,\n#             kernel_size=3,\n#         )\n#         self.num_classes = num_classes\n#         self.y_scale_param = nn.Parameter(torch.tensor(1.0))\n#         self.y_offset_param = nn.Parameter(torch.tensor(0.0))\n#         self.fixed_y_coord_scale = fixed_y_coord_scale\n#         self.detach_y_coord_path = detach_y_coord_path\n\n#     def forward(self, x):\n#         \"\"\"Forward pass for signal segmentation and y-coordinate prediction.\n\n#         Args:\n#             x (torch.Tensor): Input tensor of shape (B, N_LEADS, C, H, W).\n\n#         Returns:\n#             tuple[torch.Tensor, torch.Tensor]: Predicted mask (B, N_LEADS, num_classes, H, W)\n#                 and predicted y-coordinates (B, N_LEADS, W).\n#         \"\"\"\n#         bs, n_leads, ch, h, w = x.shape\n#         x = x.reshape(bs * n_leads, ch, h, w)\n#         feat_list = self._encoder(x)\n#         feat_list = [x] + feat_list  # prepend raw input image\n#         dec_out = self._decoder(*feat_list)\n#         pred_mask = self._segmentation_head(dec_out)\n#         pred_mask = pred_mask.reshape(bs, n_leads, self.num_classes, h, w)\n\n#         # Extract y-coordinates from mask\n#         pred_y_coord = extract_y_coord_from_mask(pred_mask.sigmoid())\n#         if self.detach_y_coord_path:\n#             pred_y_coord = pred_y_coord.detach()\n#         pred_y_coord = self.y_offset_param + self.fixed_y_coord_scale * torch.sigmoid(self.y_scale_param) * pred_y_coord\n#         return pred_mask, pred_y_coord\n\n\nclass SignalSegModelV7(nn.Module):\n    \"\"\"\n    Advanced signal segmentation model with deformation-aware warping,\n    auxiliary segmentation, lead-wise positional encoding, and\n    context-augmented y-coordinate extraction.\n\n    This model extends a standard U-Net–based segmentation architecture by:\n    - Predicting a dense warp field to spatially align segmentation outputs\n    - Using auxiliary segmentation to stabilize y-coordinate estimation\n    - Incorporating lead-wise positional encoding\n    - Aggregating vertical context from neighboring leads for y-coordinate extraction\n    \"\"\"\n\n    def __init__(\n        self,\n        # encoder_out_channels,  # must contain input raw image channels (usually 3) at the beginning of the list\n        num_classes=1,\n        encoder_model_name=\"timm-efficientnet-b0\",\n        decoder_channels=None,\n        fixed_y_coord_scale=2 * 0.01,\n        detach_y_coord_path=True,\n        y_coord_extract_head_cls=DenoisingYCoordExtractionWeightedAvgHeadV3,\n        y_coord_extract_head_kwargs={},\n        max_warp_pixels=30.0,  # Maximum displacement in real pixel space\n        apply_lead_pe=True,\n    ):\n        \"\"\"\n        Args:\n            num_classes (int): Number of output channels for segmentation.\n            encoder_model_name (str): Encoder backbone name from timm.\n            decoder_channels (list[int] | None): Decoder channel configuration.\n                If None, channels are automatically inferred from encoder.\n            fixed_y_coord_scale (float): Fixed scaling factor for y-coordinate prediction.\n            detach_y_coord_path (bool): Whether to detach segmentation output\n                when feeding into y-coordinate extraction head.\n            y_coord_extract_head_cls (nn.Module): Class used for y-coordinate extraction.\n            y_coord_extract_head_kwargs (dict): Additional kwargs for y-coordinate head.\n            max_warp_pixels (float): Maximum allowable warp displacement in pixels.\n            apply_lead_pe (bool): Whether to apply lead-wise positional encoding.\n        \"\"\"\n        super().__init__()\n        # ---------------------------------------------------------------------\n        # Encoder\n        # ---------------------------------------------------------------------\n        self._encoder = timm.create_model(\n            encoder_model_name, features_only=True, pretrained=False\n        )\n        encoder_out_channels = self._encoder.feature_info.channels()  # Channels per encoder stage\n        encoder_reductions = [info[\"reduction\"] for info in self._encoder.feature_info.info]\n        # ---------------------------------------------------------------------\n        # Auto-generate decoder channels if not provided\n        # ---------------------------------------------------------------------\n        if decoder_channels is None:\n            decoder_channels = []\n            for ch in reversed(encoder_out_channels):\n                # Reduce channels by half, but keep a minimum of 16\n                decoder_channels.append(max(ch // 2, 16))\n            # Truncate to match the number of encoder stages\n            decoder_channels = decoder_channels[: len(encoder_out_channels)]\n        # ---------------------------------------------------------------------\n        # Compute final upsampling factor for segmentation heads\n        # ---------------------------------------------------------------------\n        print(len(decoder_channels), encoder_reductions)\n        final_scale = encoder_reductions[-1] / (2 ** len(decoder_channels))\n        print(f\"final_scale : {final_scale}\")\n\n        # ---------------------------------------------------------------------\n        # Decoders\n        #   - Main decoder for segmentation\n        #   - Warp decoder for deformation field prediction\n        #   - Auxiliary decoder for auxiliary segmentation\n        # ---------------------------------------------------------------------\n        self._decoder = UnetDecoder(\n            encoder_channels=[3] + encoder_out_channels,\n            decoder_channels=decoder_channels,\n            n_blocks=len(decoder_channels),\n            use_batchnorm=True,\n            center=True if encoder_model_name.startswith(\"vgg\") else False,\n            attention_type=None,\n        )\n        self._decoder_warp = UnetDecoder(\n            encoder_channels=[3] + encoder_out_channels,\n            decoder_channels=decoder_channels,\n            n_blocks=len(decoder_channels),\n            use_batchnorm=True,\n            center=True if encoder_model_name.startswith(\"vgg\") else False,\n            attention_type=None,\n        )\n        self._decoder_aux = UnetDecoder(\n            encoder_channels=[3] + encoder_out_channels,\n            decoder_channels=decoder_channels,\n            n_blocks=len(decoder_channels),\n            use_batchnorm=True,\n            center=True if encoder_model_name.startswith(\"vgg\") else False,\n            attention_type=None,\n        )\n        # ---------------------------------------------------------------------\n        # Segmentation heads\n        # ---------------------------------------------------------------------\n        self._segmentation_head = SegmentationHead(\n            in_channels=decoder_channels[-1],\n            out_channels=num_classes,\n            activation=None,\n            kernel_size=3,\n            upsampling=final_scale if final_scale > 1 else 1,\n        )\n        self._segmentation_warp_head = SegmentationHead(\n            in_channels=decoder_channels[-1],\n            out_channels=2,  # (dx, dy) warp field\n            activation=None,\n            kernel_size=3,\n            upsampling=final_scale if final_scale > 1 else 1,\n        )\n        self._segmentation_aux_head = SegmentationHead(\n            in_channels=decoder_channels[-1],\n            out_channels=num_classes,\n            activation=None,\n            kernel_size=3,\n            upsampling=final_scale if final_scale > 1 else 1,\n        )\n        # ---------------------------------------------------------------------\n        # Warp scaling (normalized grid space)\n        # ---------------------------------------------------------------------\n        warp_scale_h = max_warp_pixels / (SIG_SEG_IMAGE_SIZE[0] / 2.0)\n        warp_scale_w = max_warp_pixels / (SIG_SEG_IMAGE_SIZE[1] / 2.0)\n        self.register_buffer(\n            \"warp_scale\", torch.tensor([warp_scale_w, warp_scale_h])\n        )\n        # ---------------------------------------------------------------------\n        # Y-coordinate extraction head\n        # ---------------------------------------------------------------------\n        head_kwargs = y_coord_extract_head_kwargs.copy()\n        if \"input_h_dim\" not in head_kwargs:\n            # Input height becomes 3 * H due to vertical concatenation\n            head_kwargs[\"input_h_dim\"] = 3 * SIG_SEG_IMAGE_SIZE[0]\n        self._y_coord_extract_head = y_coord_extract_head_cls(**head_kwargs)\n        # ---------------------------------------------------------------------\n        # Learnable parameters and configuration flags\n        # ---------------------------------------------------------------------\n        self.num_classes = num_classes\n        self.y_scale_param = nn.Parameter(torch.tensor(1.0))\n        self.y_offset_param = nn.Parameter(torch.tensor(0.0))\n        self.fixed_y_coord_scale = fixed_y_coord_scale\n        self.detach_y_coord_path = detach_y_coord_path\n        # Blending coefficient between main and auxiliary segmentation\n        self.alpha = nn.Parameter(torch.tensor(0.7))\n        # ---------------------------------------------------------------------\n        # Positional encodings\n        # ---------------------------------------------------------------------\n        self.center_pe = nn.Parameter(torch.randn(*SIG_SEG_IMAGE_SIZE))\n\n        grid_y, grid_x = torch.meshgrid(\n            torch.linspace(-1, 1, SIG_SEG_IMAGE_SIZE[0]),\n            torch.linspace(-1, 1, SIG_SEG_IMAGE_SIZE[1]),\n            indexing=\"ij\",\n        )\n        base_grid = torch.stack((grid_x, grid_y), dim=-1)  # (H, W, 2)\n        self.register_buffer(\"base_grid\", base_grid)\n        # Lead-wise positional encoding\n        self.apply_lead_pe = apply_lead_pe\n        self.n_leads = 16\n        self.lead_pe = nn.Parameter(\n            torch.zeros(self.n_leads, 3, *SIG_SEG_IMAGE_SIZE)\n        )  # (n_leads, C, H, W)\n\n    def forward(self, x):\n        \"\"\"\n        Forward pass for segmentation, warp-based alignment,\n        and context-aware y-coordinate extraction.\n\n        Args:\n            x (torch.Tensor): Input tensor of shape (B, N_LEADS, C, H, W).\n\n        Returns:\n            tuple[torch.Tensor, torch.Tensor]:\n                - pred_mask: Segmentation mask of shape (B, N_LEADS, num_classes, H, W)\n                - pred_y_coord: Predicted y-coordinates of shape (B, N_LEADS, W)\n        \"\"\"\n        bs, n_leads, ch, h, w = x.shape\n        assert self.n_leads == n_leads\n        # Apply lead-wise positional encoding if enabled\n        if self.apply_lead_pe:\n            x = x + self.lead_pe.unsqueeze(0)\n        # Merge batch and lead dimensions\n        x = x.reshape(bs * n_leads, ch, h, w)\n        # ---------------------------------------------------------------------\n        # Encoder forward\n        # ---------------------------------------------------------------------\n        feat_list = self._encoder(x)\n        feat_list = [x] + feat_list  # Prepend raw input image\n        # ---------------------------------------------------------------------\n        # Decoder forward\n        # ---------------------------------------------------------------------\n        dec_out = self._decoder(*feat_list)\n        dec_warp_out = self._decoder_warp(*feat_list)\n        dec_aux_out = self._decoder_aux(*feat_list)\n        # ---------------------------------------------------------------------\n        # Segmentation predictions\n        # ---------------------------------------------------------------------\n        pred_mask = self._segmentation_head(dec_out)\n        pred_warp_field_raw = self._segmentation_warp_head(dec_warp_out)\n        pred_mask_aux = self._segmentation_aux_head(dec_aux_out)\n        # ---------------------------------------------------------------------\n        # Warp field normalization and application\n        # ---------------------------------------------------------------------\n        pred_warp_field = torch.tanh(pred_warp_field_raw) * self.warp_scale.view(1, 2, 1, 1)\n        warped_grid = self.base_grid.unsqueeze(0) + pred_warp_field.permute(0, 2, 3, 1)\n        pred_mask = F.grid_sample(pred_mask, warped_grid, align_corners=False)\n        pred_mask_aux = F.grid_sample(pred_mask_aux, warped_grid, align_corners=False)\n        # ---------------------------------------------------------------------\n        # Blend main and auxiliary segmentation for y-coordinate extraction\n        # ---------------------------------------------------------------------\n        y_coord_head_input = (\n            self.alpha\n            * (pred_mask.detach() if self.detach_y_coord_path else pred_mask)\n            + (1.0 - self.alpha) * pred_mask_aux\n        )\n        # ---------------------------------------------------------------------\n        # Incorporate vertical context from neighboring leads\n        # Concatenate (lead-1, current lead, lead+1) along height dimension\n        # ---------------------------------------------------------------------\n        y_coord_head_input = y_coord_head_input.reshape(bs, n_leads, self.num_classes, h, w)\n        pred_mask_aux = pred_mask_aux.reshape(bs, n_leads, self.num_classes, h, w)\n        top_bottom_zero_padding = torch.zeros_like(pred_mask_aux)[:, :4, :, :, :].to(pred_mask_aux.device)\n        pred_mask_aux_zero_padded = torch.cat(\n            [top_bottom_zero_padding, pred_mask_aux, top_bottom_zero_padding],\n            dim=1,\n        )  # (B, 4 + N_LEADS + 4, NC, H, W)\n        pred_mask_aux_top = pred_mask_aux_zero_padded[:, :n_leads]\n        pred_mask_aux_bottom = pred_mask_aux_zero_padded[:, -n_leads:]\n        y_coord_head_input = torch.cat(\n            [\n                pred_mask_aux_top,\n                self.center_pe.reshape((1, 1, 1, h, w)) + y_coord_head_input,\n                pred_mask_aux_bottom,\n            ],\n            dim=-2,\n        )  # (B, N_LEADS, NC, 3*H, W)\n        y_coord_head_input = y_coord_head_input.reshape(bs * n_leads, self.num_classes, 3 * h, w)\n        # ---------------------------------------------------------------------\n        # Y-coordinate prediction\n        # ---------------------------------------------------------------------\n        pred_y_coord = self._y_coord_extract_head(y_coord_head_input)\n        # Restore batch and lead dimensions\n        pred_mask = pred_mask.reshape(bs, n_leads, self.num_classes, h, w)\n        pred_y_coord = pred_y_coord.reshape(bs, n_leads, w)\n\n        return pred_mask, pred_y_coord","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:17:38.203005Z","iopub.execute_input":"2025-12-19T07:17:38.203727Z","iopub.status.idle":"2025-12-19T07:17:38.236916Z","shell.execute_reply.started":"2025-12-19T07:17:38.203691Z","shell.execute_reply":"2025-12-19T07:17:38.235988Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Prediction","metadata":{}},{"cell_type":"code","source":"@torch.inference_mode()\ndef predict(\n    ecg_id: str,\n    image_filepath: str,\n    keypoint_seg_model,\n    signal_seg_model,\n    kptdet_transforms,\n    sigseg_transforms,\n    sr_anno_ref: pd.Series,\n    df_meta: pd.DataFrame,\n):\n    try:\n        # raise Exception()\n        # 1. Keypoint detection\n        kptdet_input_image_tensor, org_pil_image = prepare_image_for_keypoint_model(\n            image_filepath=image_filepath,\n            transforms=kptdet_transforms,\n        )\n        with torch.inference_mode():\n            kptdet_pred_llogits_mask = keypoint_seg_model(kptdet_input_image_tensor.unsqueeze(0).to(device)).squeeze(0)  # (N_KPTS(17), H, W)\n            kptdet_pred_llogits_mask = kptdet_pred_llogits_mask[:N_KPTS, :, :]\n        pred_kpts_ref_sorted, _ = extract_keypoints_from_kpt_model_pred(\n            pred_logits_mask=kptdet_pred_llogits_mask,\n            org_pil_image=org_pil_image, \n            n_kpts=N_KPTS,\n            ref_keypoints=sr_anno_ref.keypoints,\n        )\n\n        if ENABLE_GRID_POINTS_ADJUST_BY_REF_AFFINE:\n            pred_kpts_ref_sorted, inliers, M = map_ref_to_pred_affine_only(pred_kpts_ref_sorted, np.array(sr_anno_ref.keypoints).astype(float), ransac_thresh=0.05)\n            # pred_kpts_ref_sorted, inliers = map_ref_to_pred_affine_tps(pred_kpts_ref_sorted, np.array(sr_anno_ref.keypoints).astype(float), ransac_thresh=0.05)\n            print(f'[DEBUG] sum(inliers) : {sum(inliers)}')\n        if ENABLE_RIGHT_GRID_POINTS_ADJUST:\n            pred_kpts_ref_sorted = adjust_right_edge_keypoints(pred_kpts_ref_sorted)\n\n        kpt_shape_sim = shape_similarity(np.array(sr_anno_ref.keypoints).astype(float), pred_kpts_ref_sorted)\n    \n        # 2. Signal Extraction\n        if kpt_shape_sim < 0.9:\n            # keypoint detection likely failed\n            pred_sig_y_coord = torch.zeros(16, SIG_SEG_IMAGE_SIZE[1])\n            warp_cropped_pil_image = org_pil_image\n            print('[warning] keypoint detection failed')\n        else:\n            warp_cropped_pil_image, dst_pts_ext, dst_pts_after_crop = apply_warp_affine_from_keypoints(\n                org_pil_image=org_pil_image,\n                pred_kpts_ref_sorted=pred_kpts_ref_sorted,\n                ref_keypoints=sr_anno_ref.keypoints,\n                W_ref=sr_anno_ref.width,\n                H_ref=sr_anno_ref.height,\n                extend_point_ratio=0.2,\n            )\n            sigseg_input_image_tensor = prepare_image_for_signal_seg_model(\n                warp_cropped_image=warp_cropped_pil_image,\n                dst_pts_after_crop=dst_pts_after_crop,\n                transforms=sigseg_transforms,\n            )\n            with torch.inference_mode():\n                pred_sigseg_logits_mask, pred_sig_y_coord = signal_seg_model(sigseg_input_image_tensor.unsqueeze(0).to(device))\n                pred_sigseg_logits_mask = pred_sigseg_logits_mask.squeeze(0)\n                pred_sig_y_coord = pred_sig_y_coord.squeeze(0)\n            assert pred_sig_y_coord.shape == (16, SIG_SEG_IMAGE_SIZE[1])\n    except Exception as e:\n        print(e)\n        pred_sig_y_coord = torch.zeros(16, SIG_SEG_IMAGE_SIZE[1])\n        warp_cropped_pil_image = None\n\n    # 3. format to submission format\n    df_sub = format_pred_signal_to_submit_df(pred_sig_y_coord, ecg_id, df_meta)\n    return df_sub, warp_cropped_pil_image, pred_sig_y_coord","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:17:38.350348Z","iopub.execute_input":"2025-12-19T07:17:38.350641Z","iopub.status.idle":"2025-12-19T07:17:38.360483Z","shell.execute_reply.started":"2025-12-19T07:17:38.350622Z","shell.execute_reply":"2025-12-19T07:17:38.359661Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- prepare model","metadata":{}},{"cell_type":"code","source":"keypoint_seg_model = KeypointSegModelV2(\n    encoder_model_name='convnext_small.dinov3_lvd1689m',\n    num_classes=N_KPTS+2*N_CONNS,  # 2*N_CONNS for PAF\n).eval().to(device)\nkeypoint_seg_model.load_state_dict(torch.load('/kaggle/input/physionet2025-private-dataset/exp0005_run1_50ep.pth', map_location=torch.device(device)))\n\n# signal_seg_model = SignalExtractorModel(\n#     encoder_model_name='tf_efficientnetv2_m.in21k_ft_in1k',\n#     num_classes=1\n# ).eval().to(device)\n# signal_seg_model.load_state_dict(torch.load('/kaggle/input/physionet2025-private-dataset/exp1002_run1_40ep.pth', map_location=torch.device(device)))\nsignal_seg_model = SignalSegModelV7(\n    encoder_model_name='convnext_small.dinov3_lvd1689m',\n    num_classes=1,\n    y_coord_extract_head_cls=DenoisingYCoordExtractionWeightedAvgHeadV3,\n    y_coord_extract_head_kwargs={},\n    apply_lead_pe=True,\n    max_warp_pixels=30,\n).eval().to(device)\nsignal_seg_model.load_state_dict(torch.load('/kaggle/input/physionet2025-private-dataset/exp1003_run6_50ep.pth', map_location=torch.device(device)))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:17:41.346151Z","iopub.execute_input":"2025-12-19T07:17:41.347000Z","iopub.status.idle":"2025-12-19T07:17:49.299274Z","shell.execute_reply.started":"2025-12-19T07:17:41.346975Z","shell.execute_reply":"2025-12-19T07:17:49.298466Z"},"scrolled":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- prepare transforms","metadata":{}},{"cell_type":"code","source":"kptdet_transforms = get_ecg_test_transforms_for_kptdet()\nsigseg_transforms = get_ecg_test_transforms_for_sigseg()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:17:49.300692Z","iopub.execute_input":"2025-12-19T07:17:49.300982Z","iopub.status.idle":"2025-12-19T07:17:49.307952Z","shell.execute_reply.started":"2025-12-19T07:17:49.300963Z","shell.execute_reply":"2025-12-19T07:17:49.307118Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"uniq_ids = df_meta_test['id'].unique().tolist()\n\ndebug = False if os.getenv('KAGGLE_IS_COMPETITION_RERUN') else True\n\ndf_sub_list = []\n\nfor i, ecg_id in enumerate(pb(uniq_ids[::-1])):\n    image_filepath = f\"/kaggle/input/physionet-ecg-image-digitization/test/{ecg_id}.png\"\n    df_sub, warp_cropped_pil_image, pred_sig_y_coord = predict(\n        ecg_id=ecg_id,\n        image_filepath=image_filepath,\n        keypoint_seg_model=keypoint_seg_model,\n        signal_seg_model=signal_seg_model,\n        kptdet_transforms=kptdet_transforms,\n        sigseg_transforms=sigseg_transforms,\n        sr_anno_ref=sr_anno_ref,\n        df_meta=df_meta_test,\n    )\n    if debug:\n        print(ecg_id)\n        if warp_cropped_pil_image is not None:\n            display(warp_cropped_pil_image.resize((700, 400)))\n        plot_pred_vs_gt_grid(pred_sig_y_coord)\n    assert not df_sub['value'].isna().any()\n    assert  np.isfinite(df_sub['value']).all()\n    df_sub_list.append(df_sub)\n\ndf_sub = pd.concat(df_sub_list).sort_index()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:17:49.308697Z","iopub.execute_input":"2025-12-19T07:17:49.308922Z","iopub.status.idle":"2025-12-19T07:18:10.793471Z","shell.execute_reply.started":"2025-12-19T07:17:49.308905Z","shell.execute_reply":"2025-12-19T07:18:10.792674Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_sub","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:18:10.794861Z","iopub.execute_input":"2025-12-19T07:18:10.795162Z","iopub.status.idle":"2025-12-19T07:18:10.805499Z","shell.execute_reply.started":"2025-12-19T07:18:10.795141Z","shell.execute_reply":"2025-12-19T07:18:10.804802Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Submission","metadata":{}},{"cell_type":"code","source":"df_sub.to_csv('submission.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:18:10.806490Z","iopub.execute_input":"2025-12-19T07:18:10.806849Z","iopub.status.idle":"2025-12-19T07:18:11.028235Z","shell.execute_reply.started":"2025-12-19T07:18:10.806827Z","shell.execute_reply":"2025-12-19T07:18:11.027564Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pd.read_csv('submission.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T07:18:11.029047Z","iopub.execute_input":"2025-12-19T07:18:11.029314Z","iopub.status.idle":"2025-12-19T07:18:11.096703Z","shell.execute_reply.started":"2025-12-19T07:18:11.029296Z","shell.execute_reply":"2025-12-19T07:18:11.095947Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Sample Evaluation","metadata":{}},{"cell_type":"code","source":"debug = False if os.getenv('KAGGLE_IS_COMPETITION_RERUN') else True\n\nif debug:\n\n    uniq_ids = df_meta_train['id'].unique().tolist()\n    \n    df_sub_list = []\n    \n    for i, ecg_id in enumerate(pb(uniq_ids[::-1][:30])):  # last 50 sample images\n        # image_filepath = f\"/kaggle/input/physionet-ecg-image-digitization/test/{ecg_id}.png\"\n        image_filepath = f\"/kaggle/input/physionet-ecg-image-digitization/train/{ecg_id}/{ecg_id}-{random.choice([1, 3, 4, 5, 6, 9, 10, 11, 12]):04d}.png\"\n        df_sub, warp_cropped_pil_image, pred_sig_y_coord = predict(\n            ecg_id=ecg_id,\n            image_filepath=image_filepath,\n            keypoint_seg_model=keypoint_seg_model,\n            signal_seg_model=signal_seg_model,\n            kptdet_transforms=kptdet_transforms,\n            sigseg_transforms=sigseg_transforms,\n            sr_anno_ref=sr_anno_ref,\n            df_meta=df_meta_train,\n        )\n        if i % 5 == 0:\n            print(image_filepath)\n            display(warp_cropped_pil_image.resize((700, 400)))\n            plot_pred_vs_gt_grid(pred_sig_y_coord)\n        df_sub_list.append(df_sub)\n    \n    df_sub = pd.concat(df_sub_list).sort_index()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T03:18:54.081225Z","iopub.execute_input":"2025-12-08T03:18:54.081957Z","iopub.status.idle":"2025-12-08T03:25:30.450913Z","shell.execute_reply.started":"2025-12-08T03:18:54.081929Z","shell.execute_reply":"2025-12-08T03:25:30.450197Z"},"scrolled":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}