{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":97984,"databundleVersionId":14096757,"sourceType":"competition"},{"sourceId":14582340,"sourceType":"datasetVersion","datasetId":9135743},{"sourceId":14583494,"sourceType":"datasetVersion","datasetId":9299902}],"dockerImageVersionId":31260,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# # This Python 3 environment comes with many helpful analytics libraries installed\n# # It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# # For example, here's several helpful packages to load\n\n# import numpy as np # linear algebra\n# import pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# # Input data files are available in the read-only \"../input/\" directory\n# # For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\n# import os\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# # You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# # You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!df -h","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T16:29:56.919177Z","iopub.execute_input":"2026-01-22T16:29:56.919414Z","iopub.status.idle":"2026-01-22T16:29:57.077584Z","shell.execute_reply.started":"2026-01-22T16:29:56.919388Z","shell.execute_reply":"2026-01-22T16:29:57.076625Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!mkdir -p /kaggle/temp/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T16:29:57.622863Z","iopub.execute_input":"2026-01-22T16:29:57.623200Z","iopub.status.idle":"2026-01-22T16:29:57.747372Z","shell.execute_reply.started":"2026-01-22T16:29:57.623163Z","shell.execute_reply":"2026-01-22T16:29:57.746359Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!uv pip install --no-index --no-reinstall --no-deps \\\n    --find-links=/kaggle/input/yagm-dependencies/dependencies/ \\\n    setproctitle hydra-core lightning albumentations==1.4.16 huggingface-hub==1.2.1 transformers==5.0.0.rc0\n\n!uv pip install --no-index --no-reinstall --no-deps \\\n    --find-links=/kaggle/input/yagm-dependencies monai","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T16:29:57.847341Z","iopub.execute_input":"2026-01-22T16:29:57.847761Z","iopub.status.idle":"2026-01-22T16:29:58.374489Z","shell.execute_reply.started":"2026-01-22T16:29:57.847718Z","shell.execute_reply":"2026-01-22T16:29:58.373517Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!rm -rf /kaggle/working/ecg\n!cp -r /kaggle/input/ecg-checkpoints/ecg /kaggle/working/ecg","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T16:29:58.376280Z","iopub.execute_input":"2026-01-22T16:29:58.376577Z","iopub.status.idle":"2026-01-22T16:30:03.994358Z","shell.execute_reply.started":"2026-01-22T16:29:58.376535Z","shell.execute_reply":"2026-01-22T16:30:03.993391Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%cd /kaggle/working/ecg/src/ecg/inference/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T16:30:03.995741Z","iopub.execute_input":"2026-01-22T16:30:03.996085Z","iopub.status.idle":"2026-01-22T16:30:04.001857Z","shell.execute_reply.started":"2026-01-22T16:30:03.996042Z","shell.execute_reply":"2026-01-22T16:30:04.001175Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T16:30:04.003405Z","iopub.execute_input":"2026-01-22T16:30:04.003679Z","iopub.status.idle":"2026-01-22T16:30:04.126147Z","shell.execute_reply.started":"2026-01-22T16:30:04.003658Z","shell.execute_reply":"2026-01-22T16:30:04.125513Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T16:30:04.127231Z","iopub.execute_input":"2026-01-22T16:30:04.127552Z","iopub.status.idle":"2026-01-22T16:30:04.133818Z","shell.execute_reply.started":"2026-01-22T16:30:04.127509Z","shell.execute_reply":"2026-01-22T16:30:04.133173Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport sys\n\nos.environ[\"NO_ALBUMENTATIONS_UPDATE\"] = \"1\"\nsys.path = [\n    \"/kaggle/working/ecg/yagm/src/\",\n    \"/kaggle/working/ecg/src/\",\n    \"/kaggle/working/ecg/\",\n    \"/kaggle/working/ecg/third_party/segmentation_models_pytorch_3d\",\n    \"/kaggle/working/ecg/third_party/slowfast\",\n    \"/kaggle/working/ecg/third_party/timm_3d\",\n] + sys.path\nprint(\"SYS PATH:\", sys.path, sep=\"\\n\")\n\n\nimport csv\nimport gc\nimport json\nimport logging\nimport math\nimport pickle\nimport random\nimport shutil\nimport time\n\nimport albumentations as A\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport polars as pl\nimport pyarrow as pa\nimport pyarrow.parquet as pq\nimport torch\nfrom omegaconf import OmegaConf\nfrom torch.nn import functional as F\nfrom torch.utils.data import DataLoader, Dataset\nfrom tqdm import tqdm\nfrom yagm.transforms import albumentations_custom as AC\nfrom yagm.transforms.keypoints.decode import decode_heatmap_2d_batched\nfrom yagm.utils import lightning as l_utils\nfrom yagm.utils.geometric import batch_perspective_transform_2d\n# Setting up custom logging\nfrom yagm.utils.logging import init_logging, setup_logging\n\nfrom ecg.utils import constants\nfrom ecg.utils.codecs import ColumnGaussianHeatmapCodec\nfrom ecg.utils.crop import crop_by_global_homo, crop_by_unwrap_flow\nfrom ecg.utils.interpolate import interp_1d_cv2\nfrom ecg.utils.keypoint_detection_postprocess import register_grid_keypoints\nfrom ecg.utils.misc import get_scale_xy_from_homo_mat\nfrom ecg.utils.viz import viz_img_keypoints\n\n### SETUP LOGGER FOR IPYTHON ###\n# ref: https://github.com/ipython/ipykernel/issues/111\n# Create logger\nlogger = logging.getLogger()\nlogger.setLevel(logging.INFO)\n# Create STDERR handler\nhandler = logging.StreamHandler(sys.stderr)\n# Create formatter and add it to the handler\nformatter = logging.Formatter(\"[%(levelname)s] %(message)s\")\nhandler.setFormatter(formatter)\n# Set STDERR handler as the only handler\nlogger.handlers = [handler]\n\n\n# ===================== GLOBAL CONFIG =====================\nMODE = \"KAGGLE_TEST\"\n\nDEVICE = torch.device(\"cuda:0\")\n\nDO_ROTATION_INFERENCE = True\nDO_DETECTION = True\nDO_REMOVE_TMP_DETECTION = False\nDO_REGISTER = True\nDO_VIZ_DETECTION = True\nDO_SIGNAL_HEATMAP_INFERENCE = True\nDO_EVALUATE = False\n\n\nif MODE == \"LOCAL_TEST\":\n    TEST_ON_TOPK = None\n    LOADER_NUM_WORKERS = 16\n    ASSETS_DIR = \"./assets\"\n    TEST_IMG_DIR = \"/home/dangnh36/datasets/.comp/ecg/raw/test/\"\n    TEST_CSV_PATH = \"/home/dangnh36/datasets/.comp/ecg/raw/test.csv\"\n    TMP_DIR = \"./tmp/\"\n    TEMPLATE_ROOT_DIR = \"/home/dangnh36/datasets/.comp/ecg/processed/templates/fs100/\"\n    SUBMISSION_CSV_PATH = os.path.join(TMP_DIR, \"submission.csv\")\nelif MODE == \"LOCAL_VAL\":\n    TEST_ON_TOPK = None\n    LOADER_NUM_WORKERS = 16\n    ASSETS_DIR = \"./assets\"\n    TEST_IMG_DIR = \"/home/dangnh36/datasets/.comp/ecg/raw/train/\"\n    # pseudo_test_fold0_1200images.csv | pseudo_test_fold0.csv | pseudo_test.csv\n    TEST_CSV_PATH = (\n        \"/home/dangnh36/datasets/.comp/ecg/processed/pseudo_test_fold0_1200images.csv\"\n    )\n    TMP_DIR = \"./tmp/\"\n    TEMPLATE_ROOT_DIR = \"/home/dangnh36/datasets/.comp/ecg/processed/templates/fs100/\"\n    SUBMISSION_CSV_PATH = os.path.join(TMP_DIR, \"submission.csv\")\nelif MODE == \"KAGGLE_TEST\":\n    SUBMISSION_CSV_PATH = \"/kaggle/working/submission.csv\"\n    LOADER_NUM_WORKERS = 4\n    TEST_ON_TOPK = None\n    ASSETS_DIR = \"/kaggle/input/ecg-checkpoints/\"\n    TEST_IMG_DIR = \"/kaggle/input/physionet-ecg-image-digitization/test/\"\n    # pseudo_test_fold0_1200images.csv | pseudo_test_fold0.csv | pseudo_test.csv\n    TEST_CSV_PATH = \"/kaggle/input/physionet-ecg-image-digitization/test.csv\"\n    TMP_DIR = \"/kaggle/temp/\"\n    TEMPLATE_ROOT_DIR = (\n        \"/kaggle/input/ecg-checkpoints/need_to_keep/need_to_keep/templates/fs100/\"\n    )\nelse:\n    raise ValueError\n\n\n# cv2.setNumThreads(0)\n# cv2.ocl.setUseOpenCL(False)\n\nCV2_INTERPOLATION_METHODS = {\n    \"linear\": cv2.INTER_LINEAR,\n    \"area\": cv2.INTER_AREA,\n    \"bilinear\": cv2.INTER_LINEAR,\n    \"trilinear\": cv2.INTER_LINEAR,\n    \"nearest\": cv2.INTER_NEAREST,\n    \"cubic\": cv2.INTER_CUBIC,\n    \"lanczos\": cv2.INTER_LANCZOS4,\n}\n\nREF_MAIN_KPT_XYS = constants.REF_KPT_XYS[2365:]\nassert REF_MAIN_KPT_XYS.shape == (57, 2)\n\n\ndef create_submission_file(\n    predictions_dict, test_csv_path, output_path=\"submission.csv\"\n):\n    \"\"\"\n    Streams submission to CSV using the exact 'number_of_rows'\n    specified in test.csv, avoiding manual 'fs * duration' calculation.\n    \"\"\"\n    print(f\"Reading metadata from {test_csv_path}...\")\n    test_df = pd.read_csv(test_csv_path)\n\n    # Validation: Ensure test.csv has the expected columns\n    required_cols = {\"id\", \"lead\", \"number_of_rows\"}\n    if not required_cols.issubset(test_df.columns):\n        print(\n            f\"⚠️ WARNING: test.csv is missing columns: {required_cols - set(test_df.columns)}\"\n        )\n        print(\n            \"Falling back to manual calculation logic might be necessary if this fails.\"\n        )\n\n    print(f\"Streaming submission to {output_path}...\")\n\n    # 1. Open CSV Writer\n    # newline='' is required to prevent blank lines on Windows\n    with open(output_path, 'w', newline='') as f:\n        writer = csv.writer(f)\n        \n        # 2. Write Header\n        writer.writerow(['id', 'value'])\n\n        # 3. Iterate through test.csv rows directly\n        for _, row in tqdm(\n            test_df.iterrows(), total=len(test_df), desc=\"Processing test.csv rows\"\n        ):\n\n            base_id = str(row[\"id\"])\n            lead_name = row[\"lead\"]\n            target_len = int(row[\"number_of_rows\"])  # STRICT: Use this exact value\n\n            # 4. Retrieve Prediction\n            if base_id in predictions_dict:\n                img_preds = predictions_dict[base_id]\n            elif int(base_id) in predictions_dict:  # fallback if keys are ints\n                img_preds = predictions_dict[int(base_id)]\n            else:\n                img_preds = {}  # Missing ID\n\n            # Get specific lead signal\n            if lead_name in img_preds:\n                signal = img_preds[lead_name]\n            else:\n                signal = np.zeros(target_len, dtype=np.float32)\n\n            # 5. Enforce Exact Length (Resample if needed)\n            current_len = len(signal)\n\n            if current_len != target_len:\n                x_old = np.linspace(0, 1, current_len)\n                x_new = np.linspace(0, 1, target_len)\n                signal = np.interp(x_new, x_old, signal)\n\n            # 6. Generate IDs\n            # Format: {base_id}_{row_id}_{lead}\n            # row_id is 0-based index\n            ids = [f\"{base_id}_{i}_{lead_name}\" for i in range(target_len)]\n\n            # 7. Write Chunk\n            # zip() creates a lazy iterator, very memory efficient\n            writer.writerows(zip(ids, signal))\n\n    print(f\"Done! CSV Submission created at {output_path}.\")\n\n\ndef csv_to_nested_dict(csv_path):\n    \"\"\"\n    Reads the submission.csv and reconstructs the original nested dictionary:\n    { 'base_id': { 'lead': np.array([value1, value2, ...]) } }\n\n    Uses Polars for high-performance string parsing and aggregation.\n    \"\"\"\n    print(f\"Reading {csv_path}...\")\n\n    # 1. Lazy Scan (Memory Efficient)\n    # Switch from scan_parquet to scan_csv\n    lf = pl.scan_csv(csv_path)\n\n    # Regex Explanation:\n    #   ^(.*)     -> Group 1: Capture Base ID (greedy, handles slashes/underscores)\n    #   _(\\d+)    -> Group 2: Capture Row ID (digits only, preceded by _)\n    #   _([^_]+)$ -> Group 3: Capture Lead (everything after last _)\n\n    processed_lf = lf.with_columns(\n        [\n            pl.col(\"id\").str.extract(r\"^(.*)_(\\d+)_([^_]+)$\", 1).alias(\"base_id\"),\n            pl.col(\"id\")\n            .str.extract(r\"^(.*)_(\\d+)_([^_]+)$\", 2)\n            .cast(pl.Int32)\n            .alias(\"row_idx\"),\n            pl.col(\"id\").str.extract(r\"^(.*)_(\\d+)_([^_]+)$\", 3).alias(\"lead\"),\n        ]\n    )\n\n    # 2. Group by (Base ID, Lead) and Aggregate Values\n    # We MUST sort by 'row_idx' inside the aggregation to ensure signal order is correct.\n    aggregated_df = (\n        processed_lf.group_by([\"base_id\", \"lead\"])\n        .agg(pl.col(\"value\").sort_by(\"row_idx\").alias(\"signal\"))\n        .collect()\n    )  # Materialize results here\n\n    print(f\"Reconstructing dictionary from {len(aggregated_df)} groupings...\")\n\n    # 3. Convert to Dictionary\n    # Iterating over the aggregated groups is fast (approx 25k iterations for 1200 images * 12 leads)\n    reconstructed_dict = {}\n\n    for row in tqdm(aggregated_df.iter_rows(named=True), total=len(aggregated_df)):\n        base_id = row[\"base_id\"]\n        lead = row[\"lead\"]\n\n        # Polars returns lists, convert back to numpy float32\n        signal_array = np.array(row[\"signal\"], dtype=np.float32)\n\n        if base_id not in reconstructed_dict:\n            reconstructed_dict[base_id] = {}\n\n        reconstructed_dict[base_id][lead] = signal_array\n\n    return reconstructed_dict\n\n\ndef load_sample_signal(img_dir, img_id):\n    if \"/\" in img_id:\n        ori_sample_id = img_id.split(\"/\")[0]\n        signal_csv_path = os.path.join(img_dir, ori_sample_id, f\"{ori_sample_id}.csv\")\n    else:\n        raise NotImplementedError\n    signal_df = (\n        pl.scan_csv(signal_csv_path).select(pl.col(\"*\").cast(pl.Float64)).collect()\n    )\n    sig_len = len(signal_df)\n    signal = signal_df.to_numpy()\n    null_mask = ~np.isnan(signal)\n    ret = {}\n\n    for lead_idx, lead_name in enumerate(signal_df.columns):\n        assert lead_name in constants.LEAD_TO_SIG_LEN\n        start_frac, end_frac = constants.LEAD_TO_SIG_LEN[lead_name]\n        valid_indices = null_mask[:, lead_idx].nonzero()[0]\n        first = int(valid_indices[0])\n        last = int(valid_indices[-1])\n        expect_start = int(start_frac * sig_len)\n        expect_end = int(end_frac * sig_len - 1)\n        assert first == expect_start and last == expect_end\n\n        lead_sig = signal[first : last + 1, lead_idx]\n        assert not np.isnan(lead_sig).any()\n        if lead_name == \"II\":\n            assert len(lead_sig) == sig_len\n        else:\n            if sig_len % 4 == 0:\n                assert len(lead_sig) == sig_len // 4\n            else:\n                # int: 810, ceil: 972\n                assert len(lead_sig) == int(sig_len / 4) or len(lead_sig) == math.ceil(\n                    sig_len / 4\n                )\n        ret[lead_name] = lead_sig\n    return ret\n\n\ndef clear_dir_content(folder_path):\n    if os.path.isdir(folder_path):\n        for item in os.listdir(folder_path):\n            item_path = os.path.join(folder_path, item)\n            if os.path.isdir(item_path):\n                shutil.rmtree(item_path)\n            else:\n                os.remove(item_path)\n\n\ndef is_homography_near_identity(H: np.ndarray, threshold: float = 1e-2) -> bool:\n    \"\"\"\n    Checks if a 3x3 Homography matrix is near identity using Frobenius norm.\n    \n    Args:\n        H: 3x3 numpy array\n        threshold: The tolerance for 'nearness'. \n                   1e-3 is usually very close (sub-pixel shift implications).\n                   1e-2 allows for slight shifts/rotations.\n    \"\"\"\n    if H.shape != (3, 3):\n        raise ValueError(\"Input must be a 3x3 matrix\")\n\n    # 1. Handle scale ambiguity\n    # H[3,3] should be close to 1 for identity. If it's 0, it's definitely not identity.\n    if np.isclose(H[2, 2], 0):\n        return False\n    \n    # Normalize so H[2,2] == 1.0\n    H_norm = H / H[2, 2]\n    \n    # 2. Calculate distance from Identity\n    I = np.eye(3)\n    diff = H_norm - I\n    \n    # 3. Frobenius norm (square root of sum of squared elements)\n    error = np.linalg.norm(diff)\n    \n    return error < threshold\n\n\nclass ECGImageRotationDataset(Dataset):\n    def __init__(self, cfg, img_ids, img_dir):\n        self.cfg = cfg\n        self.img_ids = img_ids\n        self.img_dir = img_dir\n        print(\"NUMBER OF IMAGES:\", len(self.img_ids))\n\n        # BUILD TRANSFORM FUNC\n        print(f\"USING IMAGE SIZE={cfg.data.img_size} INTERPOLATION={cfg.data.interp}\")\n        self.transform = self.build_transform()\n        print(\"Transform:\\n\", self.transform)\n\n    def build_transform(self):\n        acfg = self.cfg.data.aug\n        # BUILD TRANSFORM\n        if acfg.keep_ratio:\n            resize_op = AC.LongestMaxHW(\n                self.cfg.data.img_size,\n                CV2_INTERPOLATION_METHODS[self.cfg.data.interp],\n                p=1.0,\n            )\n        else:\n            resize_op = A.Resize(\n                *self.cfg.data.img_size,\n                interpolation=CV2_INTERPOLATION_METHODS[self.cfg.data.interp],\n                p=1.0,\n            )\n        transform = A.Compose(\n            [\n                resize_op,\n                A.PadIfNeeded(\n                    self.cfg.data.img_size[0],\n                    self.cfg.data.img_size[1],\n                    position=\"center\",  # important to protect near-boundary consistency\n                    border_mode=cv2.BORDER_CONSTANT,\n                    value=(128, 128, 128),\n                    p=1.0,\n                ),\n            ],\n            p=1.0,\n        )\n        return transform\n\n    def __len__(self):\n        return len(self.img_ids)\n\n    def __getitem__(self, idx):\n        img_id = self.img_ids[idx]\n        ### LOAD RGB IMAGE\n        img_path = os.path.join(self.img_dir, f\"{img_id}.png\")\n        img = cv2.imread(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        img = self.transform(image=img)[\"image\"]\n        return {\n            \"idx\": torch.tensor(idx, dtype=torch.long),\n            \"image\": torch.from_numpy(img.transpose(2, 0, 1)),  # CHW\n        }\n\n\nclass ECGKeypointDetectionDataset(Dataset):\n    def __init__(self, cfg, img_ids, img_dir, rotation_map=None):\n        self.cfg = cfg\n        self.img_ids = img_ids\n        self.img_dir = img_dir\n        print(\"NUMBER OF IMAGES:\", len(self.img_ids))\n        self.rotation_map = (\n            rotation_map\n            if rotation_map is not None\n            else {img_id: 0 for img_id in self.img_ids}\n        )\n\n        # BUILD TRANSFORM FUNC\n        print(f\"USING IMAGE SIZE={cfg.data.img_size} INTERPOLATION={cfg.data.interp}\")\n        self.transform = self.build_transform()\n        print(\"Transform:\\n\", self.transform)\n\n    def build_transform(self):\n        acfg = self.cfg.data.aug\n        # BUILD TRANSFORM\n        keypoint_params = A.KeypointParams(\n            format=\"xys\",\n            remove_invisible=False,\n            check_each_transform=False,\n        )\n        compose_kwargs = dict(keypoint_params=keypoint_params)\n\n        if acfg.keep_ratio:\n            resize_op = AC.LongestMaxHW(\n                self.cfg.data.img_size,\n                CV2_INTERPOLATION_METHODS[self.cfg.data.interp],\n                p=1.0,\n            )\n        else:\n            resize_op = A.Resize(\n                *self.cfg.data.img_size,\n                interpolation=CV2_INTERPOLATION_METHODS[self.cfg.data.interp],\n                p=1.0,\n            )\n        transform = A.Compose(\n            [\n                resize_op,\n                A.PadIfNeeded(\n                    self.cfg.data.img_size[0],\n                    self.cfg.data.img_size[1],\n                    position=\"center\",  # important to protect near-boundary consistency\n                    border_mode=cv2.BORDER_CONSTANT,\n                    value=(128, 128, 128),\n                    p=1.0,\n                ),\n            ],\n            p=1.0,\n            **compose_kwargs,\n        )\n        return transform\n\n    def __len__(self):\n        return len(self.img_ids)\n\n    def __getitem__(self, idx):\n        img_id = self.img_ids[idx]\n\n        ### LOAD RGB IMAGE\n        img_path = os.path.join(self.img_dir, f\"{img_id}.png\")\n        img = cv2.imread(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        # predicted rotation index, from 0 to 3\n        # 0: 0 degree, 1: 90 degree, 2: 180 degree, 3: 270 degree\n        rot_code = self.rotation_map[img_id]\n        # after using this model to predict rotation index\n        # standardize the orientation of 0 degree\n        std_rotation_code = [\n            None,\n            cv2.ROTATE_90_COUNTERCLOCKWISE,\n            cv2.ROTATE_180,\n            cv2.ROTATE_90_CLOCKWISE,\n        ][rot_code]\n        if std_rotation_code is not None:\n            img = cv2.rotate(img, std_rotation_code)\n        else:\n            # 0 to 0, unchanged\n            pass\n\n        # 4 corners of original image\n        ori_h, ori_w = img.shape[:2]\n        src_corners = np.array(\n            [[0, 0], [ori_w, 0], [ori_w, ori_h], [0, ori_h]], dtype=np.float32\n        )  # xy\n\n        # albumentation keypoint: (x,y,scale), we use original scale as 1.0\n        albu_keypoints = np.ones((4, 3), dtype=np.float32)\n        albu_keypoints[:, :2] = src_corners\n\n        tdata = self.transform(image=img, keypoints=albu_keypoints)\n        img = tdata[\"image\"]\n        albu_keypoints = tdata[\"keypoints\"]\n        assert albu_keypoints.shape == (4, 3)\n        # calculate the transformation matrix M mapping transformed coordinates to original coordinates\n        dst_corners = albu_keypoints[:, :2]  # xy\n        M = cv2.getPerspectiveTransform(\n            src=dst_corners[None], dst=src_corners[None, ...]\n        )  # (3,3)\n\n        return {\n            \"idx\": torch.tensor(idx, dtype=torch.long),\n            \"image\": torch.from_numpy(img.transpose(2, 0, 1)),  # CHW\n            \"ori_wh\": torch.tensor([ori_w, ori_h], dtype=torch.float32),\n            \"M\": torch.from_numpy(M),\n        }\n\n\nclass ECGSignalHeatmapDataset(Dataset):\n    def __init__(self, cfg, img_ids, img_dir, template_root_dir, rotation_map=None, pairs = None):\n        self.cfg = cfg\n        self.img_ids = img_ids\n        self.img_dir = img_dir\n        print(\"NUMBER OF IMAGES:\", len(self.img_ids))\n        self.rotation_map = (\n            rotation_map\n            if rotation_map is not None\n            else {img_id: 0 for img_id in self.img_ids}\n        )\n\n        # infer some configs\n        self.GT_H, self.GT_W, self.GT_L = self.cfg.data.heatmap_hwl\n        self.IMG_H = round(self.GT_H * self.cfg.data.heatmap_stride[0])\n        self.IMG_W = round(self.GT_W * self.cfg.data.heatmap_stride[1])\n        self.img_size = [self.IMG_H, self.IMG_W]\n        # signal cover exactly 500 pixels width in groundtruth mask (width=512 pixels)\n        # 503.93700787401576 = 492.12598425196853 / (500 / 512)\n        self.CROP_W = 492.12598425196853 / (self.GT_L / self.GT_W)\n        if cfg.data.keep_aspect_ratio:\n            _crop_h = self.CROP_W * self.IMG_H / self.IMG_W\n            if cfg.data.crop_h is None:\n                self.CROP_H = _crop_h\n            else:\n                assert abs(cfg.data.crop_h - _crop_h) < 1e-6\n                self.CROP_H = cfg.data.crop_h\n        else:\n            assert cfg.data.crop_h is not None\n            self.CROP_H = cfg.data.crop_h\n        self.MAX_ABS_SIGNAL = self.CROP_H / 2 / constants.REF_Y_UNIT_PIXELS\n        logger.info(\n            \"GT_H=%d GT_W=%d GT_L=%d IMG_H=%d IMG_W=%d CROP_H=%f CROP_W=%f\",\n            self.GT_H,\n            self.GT_W,\n            self.GT_L,\n            self.IMG_H,\n            self.IMG_W,\n            self.CROP_H,\n            self.CROP_W,\n        )\n        logger.info(\n            \"Current codec allow to encode/decode signal in range %s\",\n            [-self.MAX_ABS_SIGNAL, self.MAX_ABS_SIGNAL],\n        )\n\n        self.crop_names = [\n            \"I\",\n            \"II_short\",\n            \"III\",\n            \"aVR\",\n            \"aVL\",\n            \"aVF\",\n            \"V1\",\n            \"V2\",\n            \"V3\",\n            \"V4\",\n            \"V5\",\n            \"V6\",\n            \"II_long_0\",\n            \"II_long_1\",\n            \"II_long_2\",\n            \"II_long_3\",\n        ]\n        self.samples = []\n        if pairs is None:\n            for img_idx, img_id in enumerate(img_ids):\n                fs = EXPECTED_GT_FS[img_id]\n                for crop_name in self.crop_names:\n                    if crop_name in [\"II_short\", \"II_long_0\", \"II_long_2\"]:\n                        len10 = EXPECTED_GT_LEN[img_id][\"II\"]\n                        # 10250 // 4 = 2562\n                        gt_len = len10 // 4\n                    elif crop_name in [\"II_long_1\", \"II_long_3\"]:\n                        len10 = EXPECTED_GT_LEN[img_id][\"II\"]\n                        # 10250 // 4 + 1 = 2563\n                        gt_len = len10 // 2 - len10 // 4\n                    else:\n                        gt_len = EXPECTED_GT_LEN[img_id][crop_name]\n                    self.samples.append([img_idx, img_id, crop_name, gt_len, fs])\n        else:\n            for img_id, crop_name in pairs:\n                fs = EXPECTED_GT_FS[img_id]\n                if crop_name in [\"II_short\", \"II_long_0\", \"II_long_2\"]:\n                    len10 = EXPECTED_GT_LEN[img_id][\"II\"]\n                    # 10250 // 4 = 2562\n                    gt_len = len10 // 4\n                elif crop_name in [\"II_long_1\", \"II_long_3\"]:\n                    len10 = EXPECTED_GT_LEN[img_id][\"II\"]\n                    # 10250 // 4 + 1 = 2563\n                    gt_len = len10 // 2 - len10 // 4\n                else:\n                    gt_len = EXPECTED_GT_LEN[img_id][crop_name]\n                self.samples.append([img_ids.index(img_id), img_id, crop_name, gt_len, fs])\n\n        print(\"NUMBER OF SAMPLES:\", len(self.samples))\n\n        # LOAD TEMPLATE\n        if cfg.data.use_template is not None:\n            template_dir = os.path.join(\n                template_root_dir,\n                f\"{self.IMG_H}x{self.IMG_W}\",\n            )\n            _template_names = os.listdir(template_dir)\n            self.cached_template_imgs = {}\n            for _template_name in os.listdir(template_dir):\n                template_img = cv2.imread(os.path.join(template_dir, _template_name))\n                if cfg.data.use_template == \"rgb\":\n                    template_img = cv2.cvtColor(template_img, cv2.COLOR_BGR2RGB)\n                elif cfg.data.use_template == \"gray\":\n                    template_img = cv2.cvtColor(template_img, cv2.COLOR_BGR2GRAY)[\n                        ..., None\n                    ]\n                elif cfg.data.use_template == \"red\":\n                    # take R from BGR -> (H, W, 1)\n                    template_img = template_img[:, :, 2:3]\n                else:\n                    raise ValueError\n                assert template_img.ndim == 3 and template_img.shape[:2] == (\n                    self.IMG_H,\n                    self.IMG_W,\n                )\n                self.cached_template_imgs[_template_name.replace(\".png\", \"\")] = (\n                    template_img\n                )\n            logger.info(\n                \"Loaded %d template images: %s\",\n                len(self.cached_template_imgs),\n                list(self.cached_template_imgs.keys()),\n            )\n\n        # load numpy array of XY keypoints\n        self.detected_kpt_xys = np.load(DETECTION_NPY_SAVE_PATH)\n        assert self.detected_kpt_xys.shape == (len(img_ids), constants.NUM_KPTS, 2)\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        img_idx, img_id, crop_name, gt_len, fs = self.samples[idx]\n        detected_keypoints = self.detected_kpt_xys[img_idx]\n        assert detected_keypoints.shape == (constants.NUM_KPTS, 2)\n\n        img_path = os.path.join(self.img_dir, f\"{img_id}.png\")\n        img = cv2.imread(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        # predicted rotation index, from 0 to 3\n        # 0: 0 degree, 1: 90 degree, 2: 180 degree, 3: 270 degree\n        rot_code = self.rotation_map[img_id]\n        # after using this model to predict rotation index\n        # standardize the orientation of 0 degree\n        std_rotation_code = [\n            None,\n            cv2.ROTATE_90_COUNTERCLOCKWISE,\n            cv2.ROTATE_180,\n            cv2.ROTATE_90_CLOCKWISE,\n        ][rot_code]\n        if std_rotation_code is not None:\n            img = cv2.rotate(img, std_rotation_code)\n        else:\n            # 0 to 0, unchanged\n            pass\n\n        # CROP BY WARP\n        if self.cfg.data.warp.method == \"global_homo\":\n            crop = crop_by_global_homo(\n                crop_name,\n                img,\n                detected_keypoints,\n                signal_length_secs=self.cfg.data.signal_length_secs,\n                warp_h=self.IMG_H,\n                warp_w=self.IMG_W,\n                crop_h=self.CROP_H,\n                crop_w=self.CROP_W,\n                filter_nearby=self.cfg.data.warp.global_homo.filter_nearby,\n                interpolation_mode=CV2_INTERPOLATION_METHODS[self.cfg.data.warp.interp],\n            )\n        elif self.cfg.data.warp.method == \"unwrap_flow\":\n            # -------------------------------------------\n            # dirty hack to quickly obtain TEMPLATE CROPS\n            # DON'T COMMENT OUT\n            # img = cv2.imread('/home/dangnh36/datasets/.comp/ecg/processed/REFERENCE_TEMPLATE_FS100.png')\n            # img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n            # augmented_detected_keypoints = constants.REF_KPT_XYS.copy()\n            # -------------------------------------------\n\n            try:\n                crop = crop_by_unwrap_flow(\n                    crop_name,\n                    img,\n                    detected_keypoints=detected_keypoints,\n                    signal_length_secs=self.cfg.data.signal_length_secs,\n                    warp_h=self.IMG_H,\n                    warp_w=self.IMG_W,\n                    crop_h=self.CROP_H,\n                    crop_w=self.CROP_W,\n                    filter_inside=self.cfg.data.warp.unwrap_flow.filter_inside,\n                    unwrap_method=self.cfg.data.warp.unwrap_flow.method,\n                    interpolation=self.cfg.data.warp.interp,\n                    warp_backend=self.cfg.data.warp.backend,\n                    warp_device=self.cfg.data.warp.device,\n                    flow_smooth_sigma=self.cfg.data.warp.unwrap_flow.smooth_sigma,\n                    _shift_0dot5_in_cv2_remap=self.cfg.data.warp.unwrap_flow._shift_0dot5_in_cv2_remap,\n                    **self.cfg.data.warp.unwrap_flow.kwargs,\n                )\n            except:\n                print(\"EXCEPTION OCCUR, CAN NOT CROP BY WARP FLOW\")\n                print('FALLBACK TO GLOBAL HOMO')\n                try:\n                    crop = crop_by_global_homo(\n                        crop_name,\n                        img,\n                        detected_keypoints,\n                        signal_length_secs=self.cfg.data.signal_length_secs,\n                        warp_h=self.IMG_H,\n                        warp_w=self.IMG_W,\n                        crop_h=self.CROP_H,\n                        crop_w=self.CROP_W,\n                        filter_nearby=self.cfg.data.warp.global_homo.filter_nearby,\n                        interpolation_mode=CV2_INTERPOLATION_METHODS[self.cfg.data.warp.interp],\n                    )\n                except:\n                    # SAFETY ONLY :) YOLO\n                    crop = img[:self.img_H, :self.img_W]\n\n\n            # -----------------------------------------------\n            # DON'T COMMENT OUT\n            # template_crop_save_path = f'/home/dangnh36/datasets/.comp/ecg/processed/templates/fs100/{self.IMG_H}x{self.IMG_W}/{crop_name}.png'\n            # os.makedirs(os.path.dirname(template_crop_save_path), exist_ok=True)\n            # cv2.imwrite(template_crop_save_path, cv2.cvtColor(crop, cv2.COLOR_RGB2BGR))\n            # -----------------------------------------------\n        else:\n            raise ValueError\n\n        # concat with empty template if needed\n        if self.cfg.data.use_template is not None:\n            template = self.cached_template_imgs[crop_name]\n            # usually (H, W, 6) or (H, W, 4)\n            crop = np.concatenate([crop, template], axis=-1)\n\n        return {\n            \"idx\": torch.tensor(idx, dtype=torch.long),\n            \"image\": torch.from_numpy(crop.transpose(2, 0, 1)),  # CHW\n        }\n\n\n# =========== LOAD METADATA ============\n# Load CSV using Polars for speed\nTEST_DF = pl.read_csv(TEST_CSV_PATH)\nif TEST_ON_TOPK is not None and MODE != \"KAGGLE_TEST\":\n    TEST_DF = TEST_DF[: TEST_ON_TOPK * 12]\n# F**K, TEST_DF['id'].unique().to_list() is not deterministic\nALL_IMG_IDS = sorted(list(set(TEST_DF[\"id\"].cast(pl.String).to_list())))\nprint(\"NUMBER OF TEST IDS:\", len(ALL_IMG_IDS))\nprint(\"FIRST 5 IMAGE IDS:\", ALL_IMG_IDS[:5])\n\nEXPECTED_GT_LEN = {}\nEXPECTED_GT_FS = {}\nfor row in TEST_DF.iter_rows(named=True):\n    EXPECTED_GT_LEN.setdefault(str(row[\"id\"]), {})[row[\"lead\"]] = row[\"number_of_rows\"]\n    EXPECTED_GT_FS[str(row[\"id\"])] = row[\"fs\"]\n\n\n# ======================= FACE ROTATION ========================\n\n# IMAGE_ROTATION_CONFIG_PATH = os.path.join(ASSETS_DIR, \"ROTATION_EFFL3_512.yaml\")\n# IMAGE_ROTATION_CHECKPOINT_PATH = os.path.join(\n#     ASSETS_DIR,\n#     \"ROTATION_EFFL3_512_ep=6_step=10003_val_accuracy=1.ckpt\",\n# )\n\n\nIMAGE_ROTATION_CONFIG_PATH = os.path.join(ASSETS_DIR, \"ROTATE_EFFVIT_B2_512_config.yaml\")\nIMAGE_ROTATION_CHECKPOINT_PATH = os.path.join(\n    ASSETS_DIR,\n    \"ROTATE_EFFVIT_B2_512_ep11_step12005_val_accuracy1.000000_val_angle_mae1.103764.ckpt\",\n)\n\n\nIMAGE_ROTATION_CONFIG = OmegaConf.load(IMAGE_ROTATION_CONFIG_PATH)\n\n# !!! overwrite some configs !!!\nIMAGE_ROTATION_CONFIG.model.encoder.pretrained = False\n\nROTATION_TMP_DIR = os.path.join(TMP_DIR, \"rotation\")\nos.makedirs(ROTATION_TMP_DIR, exist_ok=True)\nROTATION_MAP_PATH = os.path.join(ROTATION_TMP_DIR, \"rotation_map.json\")\n\n\ndef image_rotation_inference():\n    global IMAGE_ROTATION_CONFIG\n    image_rotation_dataset = ECGImageRotationDataset(\n        IMAGE_ROTATION_CONFIG, img_ids=ALL_IMG_IDS, img_dir=TEST_IMG_DIR\n    )\n    image_rotation_loader = DataLoader(\n        image_rotation_dataset,\n        batch_size=4,\n        shuffle=False,\n        num_workers=LOADER_NUM_WORKERS,\n        drop_last=False,\n        pin_memory=False,\n    )\n    # load model\n    image_rotation_task = l_utils.build_task(IMAGE_ROTATION_CONFIG)\n\n    l_utils.load_lightning_state_dict(\n        model=image_rotation_task,\n        ckpt_path=IMAGE_ROTATION_CHECKPOINT_PATH,\n        cfg=IMAGE_ROTATION_CONFIG,\n    )\n    logger.info(\"Loaded Pytorch state dict from %s\", IMAGE_ROTATION_CHECKPOINT_PATH)\n    model = image_rotation_task.model\n    del image_rotation_task\n    gc.collect()\n    model.to(DEVICE).eval()\n    # print(model)\n\n    # PERFORM INFERENCE PER-FOLD\n    all_pred_cls = []\n    with torch.inference_mode(), torch.autocast(\n        device_type=\"cuda\", dtype=torch.float16\n    ):\n        for batch_idx, batch in tqdm(\n            enumerate(image_rotation_loader),\n            total=len(image_rotation_loader),\n        ):\n            # print(batch_idx, [(k, v.shape) for k, v in batch.items()])\n            batch_imgs = batch[\"image\"].to(DEVICE)\n            pred_cls_logit, _pred_reg_logit = model(batch_imgs)\n            # (N, 4) -> (N,)\n            pred_cls = torch.argmax(pred_cls_logit, dim=1)\n            all_pred_cls.append(pred_cls.cpu().numpy())\n\n    all_pred_cls = np.concatenate(all_pred_cls, axis=0).tolist()\n    assert len(all_pred_cls) == len(ALL_IMG_IDS)\n\n    rotation_map = {}\n    for img_id, rotation_idx in zip(ALL_IMG_IDS, all_pred_cls):\n        rotation_map[img_id] = rotation_idx\n        if rotation_idx != 0:\n            print(\"WRONG ROTATION:\", img_id, \"-->\", rotation_idx)\n\n    with open(ROTATION_MAP_PATH, \"w\") as f:\n        json.dump(rotation_map, f)\n\n    del model\n    gc.collect()\n    torch.cuda.empty_cache()\n    return rotation_map\n\n\nif DO_ROTATION_INFERENCE:\n    ROTATION_MAP = image_rotation_inference()\nelse:\n    with open(ROTATION_MAP_PATH, \"r\") as f:\n        ROTATION_MAP = json.load(f)\n\n\n# ========== KEYPOINTS DETECTION STAGE ==========\nDETECTION_TMP_DIR = os.path.join(TMP_DIR, \"detection\")\nos.makedirs(DETECTION_TMP_DIR, exist_ok=True)\n\nKEYPOINT_DETECTION_CONFIG_PATH = os.path.join(\n    ASSETS_DIR, \"KEYPOINT_DETECTION_ROUND2_DSNT_config.yaml\"\n)\nKEYPOINT_DETECTION_CONFIG = OmegaConf.load(KEYPOINT_DETECTION_CONFIG_PATH)\nDETECTION_NPY_SAVE_PATH = os.path.join(DETECTION_TMP_DIR, \"all_predicted_keypoints.npy\")\n\n\n# !!! overwrite some configs !!!\nKEYPOINT_DETECTION_CONFIG.model.encoder.pretrained = False\nKEYPOINT_DETECTION_CONFIG.model.change_stem_stride = None\n\n# print(\"KEYPOINT DETECTION CONFIG:\", KEYPOINT_DETECTION_CONFIG, sep=\"\\n\")\n\n\ndef keypoint_detection_inference():\n    global KEYPOINT_DETECTION_CONFIG, ROTATION_MAP\n    keypoint_detection_dataset = ECGKeypointDetectionDataset(\n        KEYPOINT_DETECTION_CONFIG,\n        img_ids=ALL_IMG_IDS,\n        img_dir=TEST_IMG_DIR,\n        rotation_map=ROTATION_MAP,\n    )\n    keypoint_detection_loader = DataLoader(\n        keypoint_detection_dataset,\n        batch_size=1,\n        shuffle=False,\n        num_workers=LOADER_NUM_WORKERS,\n        drop_last=False,\n        pin_memory=False,\n    )\n    keypoint_detection_input_wh = (\n        torch.tensor(list(KEYPOINT_DETECTION_CONFIG.data.img_size))\n        .reshape(1, 1, 2)\n        .contiguous()\n        .to(DEVICE)\n    )\n\n    FOLD_CKPT_NAMES = [\n        \"KEYPOINT_DETECTION_ROUND2_DSNT_FOLD0_ep0_step2000_val_best__AVG_ACC0.999509.ckpt\",\n        \"KEYPOINT_DETECTION_ROUND2_DSNT_FOLD1_ep1_step4000_val_best__AVG_ACC0.999206.ckpt\",\n        \"KEYPOINT_DETECTION_ROUND2_DSNT_FOLD2_ep0_step2000_val_best__AVG_ACC0.999356.ckpt\",\n        \"KEYPOINT_DETECTION_ROUND2_DSNT_FOLD3_ep0_step2000_val_best__AVG_ACC0.999238.ckpt\",\n        \"KEYPOINT_DETECTION_ROUND2_DSNT_FOLD4_ep0_step3000_val_best__AVG_ACC0.999585.ckpt\",\n    ]\n    ACTIVE_FOLD_IDS = [0, 1, 2, 3, 4]\n\n    # LOAD MODELS\n    for fold_idx, fold_id in enumerate(ACTIVE_FOLD_IDS):\n        fold_ckpt_name = FOLD_CKPT_NAMES[fold_idx]\n        fold_ckpt_path = os.path.join(ASSETS_DIR, fold_ckpt_name)\n        keypoint_detection_task = l_utils.build_task(KEYPOINT_DETECTION_CONFIG)\n        l_utils.load_lightning_state_dict(\n            model=keypoint_detection_task,\n            ckpt_path=fold_ckpt_path,\n            cfg=KEYPOINT_DETECTION_CONFIG,\n        )\n        logger.info(\"Loaded Pytorch state dict from %s\", fold_ckpt_path)\n        fold_model = keypoint_detection_task.model\n        fold_model.to(DEVICE).eval()\n        # print(fold_model)\n\n        # FOLD_DETECTION_TMP_DIR = os.path.join(DETECTION_TMP_DIR, f'fold_{fold_idx}')\n        # os.makedirs(FOLD_DETECTION_TMP_DIR, exist_ok = True)\n\n        # PERFORM INFERENCE PER-FOLD\n        count = -1\n        with torch.inference_mode(), torch.autocast(\n            device_type=\"cuda\", dtype=torch.float16\n        ):\n            for batch_idx, batch in tqdm(\n                enumerate(keypoint_detection_loader),\n                total=len(keypoint_detection_loader),\n            ):\n                # print(batch_idx, [(k, v.shape) for k, v in batch.items()])\n                batch_imgs = batch[\"image\"].to(DEVICE)\n                B = batch_imgs.shape[0]\n                # (N, 58, 2048, 2048), (N, 57, 2)\n                heatmap_all_pred, dsnt_main_kpt_pred, _reg_main_kpt_pred = fold_model(\n                    batch_imgs\n                )\n                assert _reg_main_kpt_pred is None\n                # (N, 2048, 2048)\n                heatmap_grid_pred = (\n                    F.sigmoid(heatmap_all_pred[:, 0]).half().cpu().numpy()\n                )\n                # from [0, 1] to [0, W] and [0, H]\n                dsnt_main_kpt_xy_pixels_pred = (\n                    dsnt_main_kpt_pred * keypoint_detection_input_wh\n                )\n                dsnt_main_kpt_xy_pixels_pred = (\n                    dsnt_main_kpt_xy_pixels_pred.cpu().numpy()\n                )\n                batch_M = batch[\"M\"].numpy()\n\n                # from model input (2048x2048) image space, transform back to original image space (original resolution)\n                dsnt_main_kpt_xy_pixels_pred = batch_perspective_transform_2d(\n                    dsnt_main_kpt_xy_pixels_pred,  # (N, L, 2)\n                    batch_M,  # (N, 3, 3)\n                )\n\n                # print('OUTPUT:', heatmap_all_pred.shape, heatmap_grid_pred.shape, heatmap_grid_pred.dtype,\n                #     dsnt_main_kpt_xy_pixels_pred.shape, dsnt_main_kpt_xy_pixels_pred.dtype)\n\n                # gather the main keypoint's confident score\n                # left for future work :)\n\n                # N, C, H, W = heatmap.shape\n                # indices = dsnt_main_kpt_xy_pixels_pred\n                # x = indices[:, :, 0].long()\n                # y = indices[:, :, 1].long()\n                # x = x.clamp(0, W - 1)\n                # y = y.clamp(0, H - 1)\n                # # Create auxiliary indices for Batch and Channel dimensions\n                # # batch_idx: Shape (N, 1) -> Broadcasts to (N, C)\n                # batch_idx = torch.arange(N, device=heatmap.device).unsqueeze(1)\n                # # channel_idx: Shape (1, C) -> Broadcasts to (N, C)\n                # channel_idx = torch.arange(C, device=heatmap.device).unsqueeze(0)\n                # # Use Advanced Indexing\n                # conf_values = heatmap[batch_idx, channel_idx, y, x]\n\n                # SAVE GRID HEATMAP TO NPY FILE\n                batch_idxs = batch[\"idx\"].tolist()\n                for batch_element_idx in range(B):\n                    count += 1\n                    img_idx = batch_idxs[batch_element_idx]\n                    assert img_idx == count\n                    img_id = ALL_IMG_IDS[img_idx]\n                    if \"VAL\" in MODE:\n                        save_heatmap_npz_path = os.path.join(\n                            DETECTION_TMP_DIR, f'{img_id.replace(\"/\", \"--\")}.npz'\n                        )\n                    else:\n                        save_heatmap_npz_path = os.path.join(\n                            DETECTION_TMP_DIR, f\"{img_id}.npz\"\n                        )\n                    save_heatmap_grid = heatmap_grid_pred[batch_element_idx]\n                    save_dsnt_main_kpt_xy_pixels = dsnt_main_kpt_xy_pixels_pred[\n                        batch_element_idx\n                    ]\n                    M = batch_M[batch_element_idx]\n                    ori_wh = batch[\"ori_wh\"][batch_element_idx].numpy()\n\n                    if fold_idx == 0:\n                        # (2048, 2048)\n                        # save_heatmap_grid = save_heatmap_grid\n                        # (1, 57, 2)\n                        save_dsnt_main_xy = save_dsnt_main_kpt_xy_pixels[None]\n                    else:\n                        npzfile = np.load(save_heatmap_npz_path)\n                        prev_avg_heatmap_grid = npzfile[\"heatmap_grid\"]\n                        prev_dsnt_main_xy = npzfile[\"dsnt_main_xy\"]\n                        prev_M = npzfile[\"M\"]\n                        prev_ori_wh = npzfile[\"ori_wh\"]\n                        assert np.array_equal(prev_M, M) and np.array_equal(\n                            prev_ori_wh, ori_wh\n                        )\n                        # (2048, 2048)\n                        save_heatmap_grid = (\n                            prev_avg_heatmap_grid * fold_idx + save_heatmap_grid\n                        ) / (fold_idx + 1)\n                        # (NUM_FOLDS, 57, 2)\n                        save_dsnt_main_xy = np.concatenate(\n                            [prev_dsnt_main_xy, save_dsnt_main_kpt_xy_pixels[None]],\n                            axis=0,\n                        )\n\n                    print(\n                        f\"FOLD_IDX={fold_idx} FOLD_ID={fold_id} HEATMAP={save_heatmap_grid.shape} DSNT={save_dsnt_main_xy.shape}\"\n                    )\n                    np.savez(\n                        save_heatmap_npz_path,\n                        heatmap_grid=save_heatmap_grid,\n                        dsnt_main_xy=save_dsnt_main_xy,\n                        M=M,\n                        ori_wh=ori_wh,\n                    )\n\n        del fold_model\n        gc.collect()\n        torch.cuda.empty_cache()\n\n    del keypoint_detection_dataset, keypoint_detection_loader\n    del keypoint_detection_input_wh\n    gc.collect()\n    torch.cuda.empty_cache()\n    print(\"DONE \")\n\n\nif DO_DETECTION:\n    keypoint_detection_inference()\n\n\ndef keypoint_detection_postprocess():\n    global KEYPOINT_DETECTION_CONFIG, DETECTION_NPY_SAVE_PATH\n    # (1, 2)\n    heatmap_stride_wh = np.array(\n        [\n            [\n                KEYPOINT_DETECTION_CONFIG.data.heatmap_stride[1],\n                KEYPOINT_DETECTION_CONFIG.data.heatmap_stride[0],\n            ]\n        ]\n    )\n    print(\"HEATMAP STRIDE WH:\", heatmap_stride_wh)\n\n    all_final_kpt = []\n    for img_id in tqdm(ALL_IMG_IDS):\n        if \"VAL\" in MODE:\n            save_npz_path = os.path.join(\n                DETECTION_TMP_DIR, f'{img_id.replace(\"/\", \"--\")}.npz'\n            )\n        else:\n            save_npz_path = os.path.join(DETECTION_TMP_DIR, f\"{img_id}.npz\")\n        npzfile = np.load(save_npz_path)\n        heatmap_grid = npzfile[\"heatmap_grid\"]\n        dsnt_main_xy = npzfile[\"dsnt_main_xy\"]\n        M = npzfile[\"M\"]\n        ori_w, ori_h = npzfile[\"ori_wh\"]\n\n        heatmap_h, heatmap_w = heatmap_grid.shape\n\n        # estimate Homography transformation matrix from DSNT main keypoints\n        # H: from current image TO reference image (to_ref_H)\n        assert dsnt_main_xy.shape[1:] == (57, 2)\n        # (57, 2) -> (NUM_FOLDS, 57, 2)\n        ref_main_kpt_xys = REF_MAIN_KPT_XYS[None].repeat(dsnt_main_xy.shape[0], axis=0)\n        # ransacReprojThreshold: Even though MAGSAC is \"threshold-free\", OpenCV still requires\n        # this parameter as an upper bound for internal optimizations. 3.0 - 5.0 is standard.\n        for ransacReprojThreshold in [3, 5, 7, 9]:\n            H, mask = cv2.findHomography(\n                dsnt_main_xy.reshape(-1, 2),\n                ref_main_kpt_xys.reshape(-1, 2),\n                cv2.USAC_MAGSAC,  # MAGSAC++ algorithm\n                ransacReprojThreshold=ransacReprojThreshold,\n                maxIters=100_000,\n                confidence=0.9999,\n            )\n            if H is not None and mask.sum() / mask.size > 0.1:\n                print(\n                    f\"findHomography mask sum at threshold {ransacReprojThreshold}: {mask.sum()}/{mask.size}\"\n                )\n                break\n        else:\n            print(\n                \"FIND HOMOGRAPHY FAILED, FALLBACK TO PERSPECTIVE TRANSFORM:\",\n                ransacReprojThreshold,\n                mask.sum(),\n            )\n            # ref_kpt_names.index('dc_0_0'), ref_kpt_names.index('s_0_2_t'), ref_kpt_names.index('s_2_2_b'), ref_kpt_names.index('dc_3_3')\n            # [2391, 2369, 2386, 2406] - 2365 = [26, 4, 21, 41]\n            src_xy = np.median(dsnt_main_xy[:, [26, 4, 21, 41]], axis=0).astype(\n                \"float32\"\n            )  # (4, 2)\n            dst_xy = REF_MAIN_KPT_XYS[[26, 4, 21, 41]].astype(\"float32\")  # (4, 2)\n            H = cv2.getPerspectiveTransform(\n                src_xy,\n                dst_xy,\n            )\n\n        try:\n            is_identity = is_homography_near_identity(H, threshold = 1e-2)\n        except:\n            is_identity = False\n\n        # type 0001 image\n        if is_identity and ori_w == 2200 and ori_h == 1700:\n            print(\"\\n\\n\\n\\n\\n !!! FOUND IS IDENTITY !!! \\n\\n\\n\\n\\n\")\n            all_final_kpt.append(constants.REF_KPT_XYS.copy())\n            continue\n            \n        \n        to_ref_H = H\n        # estimate relative scale\n        from_ref_H = np.linalg.pinv(to_ref_H)\n        # roughtly estimate the scale from current original image (000x) relative to the reference image 0001\n        ori_scale_x, ori_scale_y = get_scale_xy_from_homo_mat(from_ref_H)\n\n        # radius: ~20 on type1 1700x2200 image\n        # image is isotropically resize + padding to `cfg.data.img_size`\n        nms_thres_x = 20 * ori_scale_x * min(heatmap_w / ori_w, heatmap_h / ori_h)\n        nms_thres_y = 20 * ori_scale_y * min(heatmap_w / ori_w, heatmap_h / ori_h)\n        nms_thres = ((nms_thres_x**2 + nms_thres_y**2) / 2) ** 0.5\n        nms_thres = max(5, nms_thres)\n        # print('NMS THRES:', nms_thres_x, nms_thres_y, nms_thres)\n\n        # decode -> grid points\n        # YX order, heatmap pixel indices space\n        _t0 = time.time()\n        raw_heatmap_grid_kpt_pred = decode_heatmap_2d_batched(\n            # step_output[\"heatmap_all_pred\"][:, 0:1],\n            torch.from_numpy(heatmap_grid[None, None]).to(DEVICE),  # (1, C, H, W)\n            pool_ksize=[3, 3],\n            nms_radius_thres=nms_thres,\n            blur_operator=None,\n            conf_thres=0.05,\n            max_dets=10_000,\n            timeout=10,\n        )\n        _t1 = time.time()\n        logger.debug(\"Decode grid points take %.2f sec\", _t1 - _t0)\n        # un-fixed number of grid keypoints detected for each image\n        # so we need a for loop over each image/element of the batch\n        heatmap_grid_kpt_pred = []\n        # we don't handle batch here, just single heatmap at a time\n        assert len(raw_heatmap_grid_kpt_pred) == 1\n        img_grid_kpt = raw_heatmap_grid_kpt_pred[0]\n        assert len(img_grid_kpt) == 1  # 1-channel only\n        img_grid_kpt = img_grid_kpt[0].numpy()\n        if len(img_grid_kpt) == 0:\n            # we met empty array of shape (0, 3) here\n            decoded_img_grid_kpt = img_grid_kpt\n        else:\n            # (L, 3)\n            assert img_grid_kpt.shape[1] == 3\n            # yx -> xy -> contigous coordinate\n            xy = img_grid_kpt[:, [1, 0]] + 0.5\n            # back to input image space\n            xy = xy * heatmap_stride_wh\n            # back to original image space\n            xy = batch_perspective_transform_2d(\n                xy[None],  # (1, L, 2)\n                M[None],  # (1, 3, 3)\n            )[0]\n            # (x, y, conf) format\n            decoded_img_grid_kpt = np.concatenate([xy, img_grid_kpt[:, 2:3]], axis=1)\n\n        # print('AFTER NMS:', decoded_img_grid_kpt.shape)\n\n        # @TODO - ADD TRY EXCEPT HERE\n        try:\n            final_grid_kpt = register_grid_keypoints(\n                decoded_img_grid_kpt,\n                to_ref_H,\n                viz=False,\n                img=None,\n                ref_img=None,\n                verbose=False,\n            )\n        except:\n            print('UNEXPECTED EXCEPTION IN REGISTERING GRID KEYPOINTS :( ')\n            final_grid_kpt = np.zeros((2365, 2), dtype = 'float32')\n\n        final_main_kpt = np.median(dsnt_main_xy, axis=0)\n        final_kpt = np.concatenate([final_grid_kpt, final_main_kpt], axis=0)\n        assert final_kpt.shape == (2422, 2)\n        all_final_kpt.append(final_kpt)\n\n        # remove temporary npz files\n        # if DO_REMOVE_TMP_DETECTION:\n        #     os.remove(save_npz_path)\n\n    all_final_kpt = np.stack(all_final_kpt, axis=0)\n    assert all_final_kpt.shape == (len(ALL_IMG_IDS), 2422, 2)\n\n    # save\n    np.save(DETECTION_NPY_SAVE_PATH, all_final_kpt)\n\n    print(\"DONE DETECTION POSTPROCESSING!!!\\n\\n\\n\")\n\n    gc.collect()\n    torch.cuda.empty_cache()\n\n\nif DO_REGISTER:\n    keypoint_detection_postprocess()\n\n\nif DO_VIZ_DETECTION:\n    all_detections = np.load(DETECTION_NPY_SAVE_PATH)\n    viz_img_idxs = random.choices(\n        list(range(len(ALL_IMG_IDS))), k=min(2, len(ALL_IMG_IDS))\n    )\n    for viz_img_idx in viz_img_idxs:\n        img_id = ALL_IMG_IDS[viz_img_idx]\n        img_path = os.path.join(TEST_IMG_DIR, f\"{img_id}.png\")\n        img = cv2.imread(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        # predicted rotation index, from 0 to 3\n        # 0: 0 degree, 1: 90 degree, 2: 180 degree, 3: 270 degree\n        rot_code = ROTATION_MAP[img_id]\n        # after using this model to predict rotation index\n        # standardize the orientation of 0 degree\n        std_rotation_code = [\n            None,\n            cv2.ROTATE_90_COUNTERCLOCKWISE,\n            cv2.ROTATE_180,\n            cv2.ROTATE_90_CLOCKWISE,\n        ][rot_code]\n        if std_rotation_code is not None:\n            img = cv2.rotate(img, std_rotation_code)\n        else:\n            # 0 to 0, unchanged\n            pass\n\n        detected_kpt_xys = all_detections[viz_img_idx]\n        if \"KAGGLE\" in MODE:\n            save_path = None\n        else:\n            save_path = os.path.join(TMP_DIR, \"viz\", f\"{img_id.replace('/', '--')}.png\")\n        viz_img_keypoints(\n            img,\n            kpt_xys=detected_kpt_xys,\n            kpt_classes=[0] * 2365 + list(range(57)),\n            save_path=save_path,\n            figsize=(12, 12),\n            point_size=10,\n            alpha=0.9,\n            show_legend=False,\n        )\n        print(\"Done visualization, saved to\", save_path)\n\n\n# ======================== SIGNAL HEATMAP INFERENCE ======================\n\nSIGNAL_RESULT_DIR = os.path.join(TMP_DIR, \"signal\")\nos.makedirs(SIGNAL_RESULT_DIR, exist_ok=True)\n\n\ndef signal_heatmap_inference(cfg, ckpt_path, save_name, pairs = None):\n    global ROTATION_MAP\n    signal_heatmap_dataset = ECGSignalHeatmapDataset(\n        cfg,\n        img_ids=ALL_IMG_IDS,\n        img_dir=TEST_IMG_DIR,\n        template_root_dir=TEMPLATE_ROOT_DIR,\n        rotation_map=ROTATION_MAP,\n        pairs = pairs\n    )\n    signal_heatmap_loader = DataLoader(\n        signal_heatmap_dataset,\n        batch_size=1,\n        shuffle=False,\n        num_workers=LOADER_NUM_WORKERS,\n        drop_last=False,\n        pin_memory=False,\n    )\n\n    # load model\n    signal_heatmap_task = l_utils.build_task(cfg)\n\n    l_utils.load_lightning_state_dict(\n        model=signal_heatmap_task,\n        ckpt_path=ckpt_path,\n        cfg=cfg,\n    )\n    logger.info(\"Loaded Pytorch state dict from %s\", ckpt_path)\n    model = signal_heatmap_task.model\n    model.to(DEVICE).eval()\n    # print(model)\n\n    # BUILD CODEC\n    ds = signal_heatmap_dataset\n    heatmap_codec = ColumnGaussianHeatmapCodec(\n        sigma=cfg.data.sigma,\n        crop_aspect_ratio=ds.CROP_H / ds.CROP_W,\n        H=ds.GT_H,\n        W=ds.GT_W,\n        L=ds.GT_L,\n        cutoff_thres=0.011108996538242308,\n        adaptive_sigma_scale=cfg.data.adaptive_sigma_scale,\n        subpixel=False,\n        dtype=torch.float32,\n        ref_signal_pixels=constants.REF_SIGNAL_PIXELS\n        * cfg.data.signal_length_secs\n        / 2.5,\n    )\n\n    test_samples = ds.samples\n\n    # PERFORM INFERENCE\n\n    cur_idx = -1\n    all_heatmap_preds = {}\n    with torch.inference_mode(), torch.autocast(\n        device_type=\"cuda\", dtype=torch.float16\n    ):\n        for batch_idx, batch in tqdm(\n            enumerate(signal_heatmap_loader),\n            total=len(signal_heatmap_loader),\n        ):\n            # print(batch_idx, [(k, v.shape) for k, v in batch.items()])\n            batch_imgs = batch[\"image\"].to(DEVICE)\n            prob_heatmap_pred, _offset_heatmap_pred, _regress_signal_pred = model(\n                batch_imgs\n            )\n            assert _offset_heatmap_pred is None\n            # usually nn.Softmax(dim=2)\n            prob_heatmap_pred = model.heatmap_act(prob_heatmap_pred)\n            # we don't use regress_signal_pred, so don't decode it\n            heatmap_signal_pred, _ = heatmap_codec.batch_decode(\n                prob_heatmap_pred,\n                _offset_heatmap_pred,\n                None,  # _regress_signal_pred\n            )\n            # .float() to deal with bf16\n            heatmap_signal_pred = heatmap_signal_pred.cpu().float().numpy()\n            B = heatmap_signal_pred.shape[0]\n            assert B == 1, \"Just to assert, >1 is still okay\"\n\n            for i in range(B):\n                cur_idx += 1\n                img_idx, img_id, crop_name, gt_len, fs = test_samples[cur_idx]\n\n                # interpolate to original length\n                ori_heatmap_signal_pred = interp_1d_cv2(\n                    heatmap_signal_pred[i], gt_len, method=\"adaptive\"\n                )\n                all_heatmap_preds.setdefault(img_id, {})[\n                    crop_name\n                ] = ori_heatmap_signal_pred\n\n    # save predictions result as pickle\n    pkl_save_path = os.path.join(SIGNAL_RESULT_DIR, f\"{save_name}.pkl\")\n    os.makedirs(os.path.dirname(pkl_save_path), exist_ok=True)\n    with open(pkl_save_path, \"wb\") as f:\n        pickle.dump(all_heatmap_preds, f)\n    print(\"SAVED PREDICTIONS TO\", pkl_save_path, \"\\n\\n\\n\\n\\n\")\n\n\nif DO_SIGNAL_HEATMAP_INFERENCE:\n    # FOLD 0 COAT 512\n    # cfg_path = \"./assets/EXP1_COAT512_FOLD0_ep3_step50000_val_SNR23.380785_config.yaml\"\n    # ckpt_path = \"./assets/EXP1_COAT512_FOLD0_ep3_step50000_val_SNR23.380785.ckpt\"\n    # save_name = \"COAT512_FOLD0\"\n    # cfg = OmegaConf.load(cfg_path)\n    # cfg.data.warp.unwrap_flow.kwargs.K = 16\n    # cfg.data.use_template = None\n    # cfg.data.warp.unwrap_flow._shift_0dot5_in_cv2_remap = True\n    # print(\"MODEL 1 CONFIG:\", cfg, sep=\"\\n\")\n    # signal_heatmap_inference(cfg, ckpt_path, save_name)\n\n    \n    # FINAL COAT + VGG LR 2e-4\n    cfg_path = \"/kaggle/input/ecg-checkpoints/FINAL_01-19__21-15-59.746642_ALLDATA_FINAL_COAT512_VGG1024_GT1000_LR2en4_SEED980611_config.yaml\"\n    ckpt_path = \"/kaggle/input/ecg-checkpoints/FINAL_01-19__21-15-59.746642_ALLDATA_FINAL_COAT512_VGG1024_GT1000_LR2en4_SEED980611_ep3_step70007_val_SNR28.006985.ckpt\"\n    save_name = \"ALLDATA_FINAL_COAT512_VGG1024_GT1000_LR2en4\"\n    cfg = OmegaConf.load(cfg_path)\n    cfg.model.encoder.pretrained = False\n    cfg.model.fine_encoder.pretrained=False\n    print(\"MODEL 1 CONFIG:\", cfg, sep=\"\\n\")\n    signal_heatmap_inference(cfg, ckpt_path, save_name)\n\n\n    # FINAL CONVNEXT + VGG LR 2e-4\n    cfg_path = \"/kaggle/input/ecg-checkpoints/FINAL_01-19__22-22-04.345114_ALLDATA_FINAL_CONVNEXT512_VGG1024_GT1000_LR1en4_SEED981022_config.yaml\"\n    ckpt_path = \"/kaggle/input/ecg-checkpoints/FINAL_01-19__22-22-04.345114_ALLDATA_FINAL_CONVNEXT512_VGG1024_GT1000_LR1en4_SEED981022_ep3_step60011_val_SNR27.484089.ckpt\"\n    save_name = \"ALLDATA_FINAL_CONVNEXT512_VGG1024_GT1000_LR1en4\"\n    cfg = OmegaConf.load(cfg_path)\n    cfg.model.encoder.pretrained = False\n    cfg.model.fine_encoder.pretrained=False\n    print(\"MODEL 1 CONFIG:\", cfg, sep=\"\\n\")\n    signal_heatmap_inference(cfg, ckpt_path, save_name)\n\n    \n\n\n    # # FINAL CONVNEXT + VGG LR 2e-4\n    # cfg_path = \"/kaggle/input/ecg-checkpoints/FINAL_01-20__13-29-11.248285_ALLDATA_FINAL_COAT512_GT500_LR3en4_SEED1022_config.yaml\"\n    # ckpt_path = \"/kaggle/input/ecg-checkpoints/FINAL_01-20__13-29-11.248285_ALLDATA_FINAL_COAT512_GT500_LR3en4_SEED1022_ep3_step64994_val_SNR26.050978.ckpt\"\n    # save_name = \"ALLDATA_FINAL_COAT512_GT500_LR3en4\"\n    # cfg = OmegaConf.load(cfg_path)\n    # cfg.model.encoder.pretrained = False\n    # # cfg.model.fine_encoder.pretrained=False\n    # print(\"MODEL 1 CONFIG:\", cfg, sep=\"\\n\")\n    # signal_heatmap_inference(cfg, ckpt_path, save_name)\n\n\n\n    # # FINAL CONVNEXT SMALL 1024\n    # cfg_path = \"/kaggle/input/ecg-checkpoints/FINAL_01-20__14-14-56.667323_ALLDATA_FINAL_CONVNEXTSMALL1024_GT1000_LR5en5_SEED1998_config.yaml\"\n    # ckpt_path = \"/kaggle/input/ecg-checkpoints/FINAL_01-20__14-14-56.667323_ALLDATA_FINAL_CONVNEXTSMALL1024_GT1000_LR5en5_SEED1998_ep3_step64994_val_SNR26.081009.ckpt\"\n    # save_name = \"ALLDATA_FINAL_CONVNEXTSMALL1024_GT1000_LR5en5\"\n    # cfg = OmegaConf.load(cfg_path)\n    # cfg.model.encoder.pretrained = False\n    # # cfg.model.fine_encoder.pretrained=False\n    # print(\"MODEL 1 CONFIG:\", cfg, sep=\"\\n\")\n    # signal_heatmap_inference(cfg, ckpt_path, save_name)\n\n\n# ================== SIGNAL POST PROCESSING =============\n# post processing\n# merge multiple II_long crop into one\n# ensemble (e.g, average) II_short vs II_long\n\nwith open(\n    os.path.join(SIGNAL_RESULT_DIR, \"ALLDATA_FINAL_COAT512_VGG1024_GT1000_LR2en4.pkl\"),\n    \"rb\",\n) as f:\n    model1_preds = pickle.load(f)\n\n\nwith open(\n    os.path.join(SIGNAL_RESULT_DIR, \"ALLDATA_FINAL_CONVNEXT512_VGG1024_GT1000_LR1en4.pkl\"),\n    \"rb\",\n) as f:\n    model2_preds = pickle.load(f)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T16:30:04.135265Z","iopub.execute_input":"2026-01-22T16:30:04.135472Z","iopub.status.idle":"2026-01-22T16:32:35.083013Z","shell.execute_reply.started":"2026-01-22T16:30:04.135453Z","shell.execute_reply":"2026-01-22T16:32:35.082234Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# [\"I\", \"II\", \"III\", \"aVR\", \"aVL\", \"aVF\", \"V1\", \"V2\", \"V3\", \"V4\", \"V5\", \"V6\"]\n\nall_ensemble_preds = {}\n\nfor img_id in ALL_IMG_IDS:\n    ensemble_preds = {}\n    all_ensemble_preds[img_id] = ensemble_preds\n    for lead in [\"I\", \"III\", \"aVR\", \"aVL\", \"aVF\", \"V1\", \"V2\", \"V3\", \"V4\", \"V5\", \"V6\"]:\n        ensemble_preds[lead] = 0.7 * model1_preds[img_id][lead] + 0.3 * model2_preds[img_id][lead]\n\n    ii_long1 = np.concatenate(\n        [\n            model1_preds[img_id][k]\n            for k in (\"II_long_0\", \"II_long_1\", \"II_long_2\", \"II_long_3\")\n        ],\n        axis=0,\n    )\n\n    ii_long2 = np.concatenate(\n        [\n            model2_preds[img_id][k]\n            for k in (\"II_long_0\", \"II_long_1\", \"II_long_2\", \"II_long_3\")\n        ],\n        axis=0,\n    )\n\n    ensemble_ii_long = ii_long1 * 0.7 + ii_long2 * 0.3\n\n    ii_short1 = model1_preds[img_id]['II_short']\n    ii_short2 = model2_preds[img_id]['II_short']\n    ensemble_ii_short = ii_short1 * 0.7 + ii_short2 * 0.3\n    ensemble_ii_long[:len(ensemble_ii_short)] = (ensemble_ii_long[:len(ensemble_ii_short)] + ensemble_ii_short) / 2\n    ensemble_preds['II'] = ensemble_ii_long\n    \n    len10 = EXPECTED_GT_LEN[img_id][\"II\"]\n    if len(ensemble_preds[\"II\"]) != len10:\n        ensemble_preds[\"II\"] = interp_1d_cv2(ensemble_preds[\"II\"], len10)\n    assert len(ensemble_preds) == 12\n    \nall_ensemble_preds ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T16:32:35.084872Z","iopub.execute_input":"2026-01-22T16:32:35.085188Z","iopub.status.idle":"2026-01-22T16:32:35.101231Z","shell.execute_reply.started":"2026-01-22T16:32:35.085113Z","shell.execute_reply":"2026-01-22T16:32:35.100483Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# TALL MODEL INFER\nneed_tall = []\nfor img_id, img_preds in all_ensemble_preds.items():\n    for lead in [\"I\", \"II\", \"III\", \"aVR\", \"aVL\", \"aVF\", \"V1\", \"V2\", \"V3\", \"V4\", \"V5\", \"V6\"]:\n        signal_pred = img_preds[lead]\n        if np.any(np.abs(signal_pred) > 3.195):\n            print(np.max(np.abs(signal_pred)))\n            if lead == 'II':\n                for _sub_lead in ['II_long_0', 'II_long_1', 'II_long_2', 'II_long_3']:\n                    need_tall.append([img_id, _sub_lead])\n            else:\n                need_tall.append([img_id, lead])\n\nprint(len(need_tall))\nneed_tall","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T16:32:35.102058Z","iopub.execute_input":"2026-01-22T16:32:35.102321Z","iopub.status.idle":"2026-01-22T16:32:35.111297Z","shell.execute_reply.started":"2026-01-22T16:32:35.102295Z","shell.execute_reply":"2026-01-22T16:32:35.110624Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# TALL MODEL\ncfg_path = \"/kaggle/input/ecg-checkpoints/CONVNEXT_TALL_1024x512_GT500_config.yaml\"\nckpt_path = \"/kaggle/input/ecg-checkpoints/CONVNEXT_TALL_1024x512_GT500_ep3_step49996_val_SNR21.749077.ckpt\"\nsave_name = \"CONVNEXT_TALL_1024x512_GT500\"\ncfg = OmegaConf.load(cfg_path)\ncfg.model.encoder.pretrained = False\n# cfg.model.fine_encoder.pretrained=False\nprint(\"MODEL 1 CONFIG:\", cfg, sep=\"\\n\")\nsignal_heatmap_inference(cfg, ckpt_path, save_name, pairs = need_tall)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T16:32:52.285838Z","iopub.execute_input":"2026-01-22T16:32:52.286356Z","iopub.status.idle":"2026-01-22T16:33:00.648137Z","shell.execute_reply.started":"2026-01-22T16:32:52.286315Z","shell.execute_reply":"2026-01-22T16:33:00.647284Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with open(\n    os.path.join(SIGNAL_RESULT_DIR, \"CONVNEXT_TALL_1024x512_GT500.pkl\"),\n    \"rb\",\n) as f:\n    tall_preds = pickle.load(f)\n\nfor img_id, v in tall_preds.items():\n    print(img_id, '-->', v.keys())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T16:33:39.353422Z","iopub.execute_input":"2026-01-22T16:33:39.354211Z","iopub.status.idle":"2026-01-22T16:33:39.359319Z","shell.execute_reply.started":"2026-01-22T16:33:39.354168Z","shell.execute_reply":"2026-01-22T16:33:39.358551Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for img_id, img_tall_preds in tall_preds.items():\n    if 'II_long_0' in img_tall_preds:\n        img_tall_preds['II'] = np.concatenate(\n            [\n                img_tall_preds.pop(k)\n                for k in (\"II_long_0\", \"II_long_1\", \"II_long_2\", \"II_long_3\")\n            ],\n            axis=0,\n        )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T16:40:45.007147Z","iopub.execute_input":"2026-01-22T16:40:45.007758Z","iopub.status.idle":"2026-01-22T16:40:45.012003Z","shell.execute_reply.started":"2026-01-22T16:40:45.007718Z","shell.execute_reply":"2026-01-22T16:40:45.011192Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for img_id, img_tall_preds in tall_preds.items():\n    for lead, tall_signal in img_tall_preds.items():\n        ensemble_signal = all_ensemble_preds[img_id][lead]\n        need_to_fix_mask = np.abs(ensemble_signal) > 3.195\n        src = ensemble_signal[need_to_fix_mask]\n        dst = tall_signal[need_to_fix_mask]\n        # print('\\n-------', src, dst, sep='\\n')\n        replace = src.copy()\n        for i, (src_e, dst_e) in enumerate(zip(src, dst)):\n            if (src_e > 0 and dst_e > src_e) or (src_e < 0 and dst_e < src_e):\n                replace[i] = dst_e\n        ensemble_signal[need_to_fix_mask] = replace\n        # print('\\n-------', src, dst, replace, sep='\\n')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T16:45:13.006462Z","iopub.execute_input":"2026-01-22T16:45:13.007055Z","iopub.status.idle":"2026-01-22T16:45:13.015738Z","shell.execute_reply.started":"2026-01-22T16:45:13.007019Z","shell.execute_reply":"2026-01-22T16:45:13.014948Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_preds = all_ensemble_preds\n\n# for img_id, img_preds in final_preds.items():\n#     print(\n#         img_id,\n#         \"-->\",\n#         {lead_name: len(signal) for lead_name, signal in img_preds.items()},\n#     )\n\n# print(final_preds)\n\n\nif MODE != \"KAGGLE_TEST\" and DO_EVALUATE:\n    # grab the GT\n    all_gts = {}\n    for img_id in ALL_IMG_IDS:\n        all_gts[img_id] = load_sample_signal(TEST_IMG_DIR, img_id)\n\n    from ecg.utils.comp_metrics_fast import compute_metrics\n\n    global_snr, per_sample_snrs = compute_metrics(\n        final_preds,\n        all_gts,\n        EXPECTED_GT_FS,\n        do_align=True,\n    )\n\n    print(\"\\n\\n\\n================ EVALUATION ==================\")\n    # print('PER SAMPLE SNR:', per_sample_snrs, sep = '\\n')\n    print(\"GLOBAL SNR:\", global_snr)\n\n\n# ==================== FINAL STEP: SUBMISSION PARQUET ==================\n\nif MODE == 'KAGGLE_TEST':\n    clear_dir_content('/kaggle/working/')\n\ncreate_submission_file(\n    predictions_dict=final_preds,\n    test_csv_path=TEST_CSV_PATH,\n    output_path=SUBMISSION_CSV_PATH,\n)\nprint(\"DONE ! SUBMISSION PARQUET SAVED !\")\n\n\nif 1:\n    print(\"=========== TEST RECOVERING ========\")\n    submission_df = pd.read_csv(SUBMISSION_CSV_PATH)\n    print(\"SUBMISSION DF:\", submission_df.shape)\n    # print(submission_df)\n    display(submission_df)\n\nif 0:\n    recovered_preds = csv_to_nested_dict(SUBMISSION_CSV_PATH)\n    global_snr, per_sample_snrs = compute_metrics(\n        recovered_preds,\n        all_gts,\n        EXPECTED_GT_FS,\n        do_align=True,\n    )\n    print(\"GLOBAL SNR:\", global_snr)\n\n\nprint(\"ALL DONE!\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T16:47:59.473215Z","iopub.execute_input":"2026-01-22T16:47:59.473940Z","iopub.status.idle":"2026-01-22T16:47:59.749185Z","shell.execute_reply.started":"2026-01-22T16:47:59.473907Z","shell.execute_reply":"2026-01-22T16:47:59.748526Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}