{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":97984,"databundleVersionId":14096757,"sourceType":"competition"},{"sourceId":14659613,"sourceType":"datasetVersion","datasetId":9365077},{"sourceId":14701999,"sourceType":"datasetVersion","datasetId":9364487}],"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\nimport numpy as np # linear algebra\nimport 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,"execution":{"iopub.status.busy":"2026-02-02T06:13:12.344408Z","iopub.execute_input":"2026-02-02T06:13:12.345070Z","iopub.status.idle":"2026-02-02T06:13:12.348891Z","shell.execute_reply.started":"2026-02-02T06:13:12.345037Z","shell.execute_reply":"2026-02-02T06:13:12.348102Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 파일명이 'connected_components_3d-3.26.1-cp312-cp312-...' 일 경우\n!pip install /kaggle/input/pipinstall/connected_components_3d-3.26.1-cp312-cp312-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T06:13:12.356865Z","iopub.execute_input":"2026-02-02T06:13:12.357120Z","iopub.status.idle":"2026-02-02T06:13:15.606647Z","shell.execute_reply.started":"2026-02-02T06:13:12.357097Z","shell.execute_reply":"2026-02-02T06:13:15.605599Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, sys, gc, re\nfrom pathlib import Path\n\nimport cv2\nimport numpy as np\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.models as models\n\nfrom IPython.display import display\n\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"expandable_segments:True\"\n\nROOT = Path(\"/kaggle/input/physionet-ecg-image-digitization\")\nHENGCK_BASE_DIR = Path(\"/kaggle/input/shsall/hengck23-submit-physionet\")\nCKPT_ROOT = Path(\"/kaggle/input/shsall/checkpoints\")\n\nOUT_PARQUET = Path(\"/kaggle/working/submission.parquet\")\nOUT_CSV = Path(\"/kaggle/working/submission.csv\")\n\nH_IN = 1700\nW_IN = 2200\nPATCH_SIZE = 100\nTARGET_WIDTH = 2200\n\nLEADS = [\"I\",\"II\",\"III\",\"aVR\",\"aVL\",\"aVF\",\"V1\",\"V2\",\"V3\",\"V4\",\"V5\",\"V6\"]\n\nUSE_CUDA = torch.cuda.is_available()\nDEVICE = torch.device(\"cuda:0\" if USE_CUDA else \"cpu\")\nAMP_DEVICE = \"cuda\" if USE_CUDA else \"cpu\"\nprint(\"USE_CUDA\", USE_CUDA)\nprint(\"DEVICE\", DEVICE)\n\ndef clear_gpu_memory():\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n        torch.cuda.reset_peak_memory_stats()\n\ndef assert_exists(p: Path, name: str):\n    if not p.exists():\n        raise FileNotFoundError(f\"{name} not found: {p}\")\n\ndef check_sample_submission_format(sample_df: pd.DataFrame):\n    if \"id\" not in sample_df.columns:\n        raise ValueError(\"sample_submission has no id column\")\n    s = sample_df[\"id\"].astype(str)\n    pat = re.compile(r\"^(\\d+)_([0-9]+)_([A-Za-z0-9]+)$\")\n    m = s.str.extract(pat)\n    bad = m.isna().any(axis=1)\n    if bad.any():\n        ex = sample_df.loc[bad, \"id\"].head(20).tolist()\n        raise ValueError(f\"sample_submission id parse failed. examples: {ex}\")\n\ndef interp_1d(x: np.ndarray, out_len: int) -> np.ndarray:\n    x = np.asarray(x, dtype=np.float32).reshape(-1)\n    n = int(x.shape[0])\n    if out_len <= 0:\n        return np.zeros((0,), dtype=np.float32)\n    if n == out_len:\n        return x.astype(np.float32)\n    if n < 2:\n        v = float(x[0]) if n == 1 else 0.0\n        return np.full((out_len,), v, dtype=np.float32)\n    src = np.linspace(0.0, 1.0, n, dtype=np.float32)\n    dst = np.linspace(0.0, 1.0, out_len, dtype=np.float32)\n    return np.interp(dst, src, x).astype(np.float32)\n\ndef build_coord_maps(H: int, W: int):\n    y_coords, x_coords = torch.meshgrid(\n        torch.linspace(0.0, 1.0, H),\n        torch.linspace(0.0, 1.0, W),\n        indexing=\"ij\",\n    )\n    return x_coords, y_coords\n\ndef build_local_fft_energy_map(img_t: torch.Tensor, patch_size: int):\n    H, W = img_t.shape\n    pad_h = (patch_size - H % patch_size) % patch_size\n    pad_w = (patch_size - W % patch_size) % patch_size\n\n    img_padded = F.pad(\n        img_t.unsqueeze(0).unsqueeze(0),\n        (0, pad_w, 0, pad_h),\n        mode=\"reflect\",\n    )\n\n    patches = img_padded.unfold(2, patch_size, patch_size).unfold(3, patch_size, patch_size)\n    patches_flat = patches.contiguous().view(-1, patch_size, patch_size)\n\n    f_coeffs = torch.fft.fft2(patches_flat)\n    mag = torch.abs(f_coeffs)\n    patch_energy = mag.mean(dim=(-1, -2)).view(1, 1, patches.shape[2], patches.shape[3])\n\n    grid_map = F.interpolate(\n        patch_energy,\n        size=(H + pad_h, W + pad_w),\n        mode=\"bilinear\",\n        align_corners=False,\n    )\n    grid_map = grid_map[0, 0, :H, :W]\n\n    denom = (grid_map.max() - grid_map.min()).clamp(min=1e-8)\n    grid_map = (grid_map - grid_map.min()) / denom\n    return grid_map\n\ndef stage1_rgb_to_input4ch(stage1_img_rgb: np.ndarray, x_coords: torch.Tensor, y_coords: torch.Tensor, H: int, W: int, patch_size: int):\n    if stage1_img_rgb is None:\n        raise ValueError(\"stage1_img_rgb is None\")\n\n    rgb = stage1_img_rgb\n    if rgb.dtype != np.uint8:\n        rgb = np.clip(rgb, 0, 255).astype(np.uint8)\n\n    gray = cv2.cvtColor(rgb, cv2.COLOR_RGB2GRAY)\n    gray = cv2.resize(gray, (W, H), interpolation=cv2.INTER_AREA)\n\n    img_t = torch.from_numpy(gray).float() / 255.0\n    grid_map = build_local_fft_energy_map(img_t, patch_size=patch_size)\n\n    input_4ch = torch.stack([img_t, grid_map, x_coords, y_coords], dim=0)\n    return input_4ch\n\nclass InhibitoryAttentionModule(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.gate_conv = nn.Sequential(\n            nn.Conv2d(1, 1, kernel_size=3, padding=1),\n            nn.Sigmoid()\n        )\n\n    def forward(self, x, fft_map):\n        mask = F.interpolate(fft_map, size=x.shape[2:], mode=\"bilinear\", align_corners=False)\n        gate = self.gate_conv(mask)\n        return x * (1.0 - gate) + x * 0.05\n\nclass ECG4RowDigitalizer(nn.Module):\n    def __init__(self, backbone_name=\"resnet34\", target_width=2200, use_dual_pool=False):\n        super().__init__()\n        self.target_width = target_width\n        self.disable_inhibitor = False\n        self.use_dual_pool = bool(use_dual_pool)\n\n        if backbone_name == \"resnet50\":\n            base_model = models.resnet50(weights=None)\n            self.feature_dim = 2048\n        else:\n            base_model = models.resnet34(weights=None)\n            self.feature_dim = 512\n\n        for layer in [base_model.layer3, base_model.layer4]:\n            for m in layer.modules():\n                if isinstance(m, nn.Conv2d):\n                    m.stride = (1, 1)\n\n        self.stem = nn.Conv2d(4, 64, kernel_size=7, stride=2, padding=3, bias=False)\n        self.backbone = nn.Sequential(\n            base_model.bn1, base_model.relu, base_model.maxpool,\n            base_model.layer1, base_model.layer2, base_model.layer3, base_model.layer4\n        )\n\n        self.inhibitor = InhibitoryAttentionModule()\n\n        self.pool_avg = nn.AdaptiveAvgPool2d((1, None))\n        if self.use_dual_pool:\n            self.pool_max = nn.AdaptiveMaxPool2d((1, None))\n            rnn_in = self.feature_dim * 2\n        else:\n            self.pool_max = None\n            rnn_in = self.feature_dim\n\n        self.rnn = nn.LSTM(\n            input_size=rnn_in,\n            hidden_size=256,\n            num_layers=2,\n            batch_first=True,\n            bidirectional=True\n        )\n\n        self.head_row1 = nn.Linear(512, 1)\n        self.head_row2 = nn.Linear(512, 1)\n        self.head_row3 = nn.Linear(512, 1)\n        self.head_row4 = nn.Linear(512, 1)\n\n    def forward(self, x):\n        fft_map = x[:, 1:2, :, :]\n        x = self.stem(x)\n        feat = self.backbone(x)\n\n        if not self.disable_inhibitor:\n            feat = self.inhibitor(feat, fft_map)\n\n        seq_avg = self.pool_avg(feat)\n        if self.use_dual_pool:\n            seq_max = self.pool_max(feat)\n            seq = torch.cat([seq_avg, seq_max], dim=1)\n        else:\n            seq = seq_avg\n\n        seq = seq.squeeze(2).permute(0, 2, 1)\n        rnn_out, _ = self.rnn(seq)\n\n        r1 = self.head_row1(rnn_out)\n        r2 = self.head_row2(rnn_out)\n        r3 = self.head_row3(rnn_out)\n        r4 = self.head_row4(rnn_out)\n\n        out = torch.cat([r1, r2, r3, r4], dim=-1).permute(0, 2, 1)\n        out = torch.tanh(out)\n        out = F.interpolate(out, size=self.target_width, mode=\"linear\")\n        return out\n\ndef extract_model_state_dict(ckpt_obj):\n    if isinstance(ckpt_obj, dict) and \"model_state_dict\" in ckpt_obj:\n        return ckpt_obj[\"model_state_dict\"], \"pack\"\n    if isinstance(ckpt_obj, dict):\n        return ckpt_obj, \"state\"\n    raise TypeError(\"unsupported ckpt type\")\n\ndef find_ecg_weight(prefer_name: str):\n    cand = list(CKPT_ROOT.rglob(prefer_name))\n    if len(cand) > 0:\n        return str(cand[0])\n    all_pth = list(CKPT_ROOT.rglob(\"*.pth\"))\n    if len(all_pth) == 0:\n        raise FileNotFoundError(\"no pth under checkpoints\")\n    all_pth.sort(key=lambda p: p.stat().st_mtime, reverse=True)\n    return str(all_pth[0])\n\ndef infer_use_dual_pool_from_ckpt(state_dict: dict) -> bool:\n    key = \"rnn.weight_ih_l0\"\n    if key not in state_dict:\n        raise KeyError(\"missing rnn.weight_ih_l0 in ckpt state_dict\")\n    in_dim = int(state_dict[key].shape[1])\n    if in_dim == 512:\n        return False\n    if in_dim == 1024:\n        return True\n    raise ValueError(f\"unexpected rnn input dim in ckpt: {in_dim}\")\n\ndef load_model_matching_ckpt(weight_path: str):\n    ckpt = torch.load(weight_path, map_location=\"cpu\")\n    state_dict, kind = extract_model_state_dict(ckpt)\n    use_dual = infer_use_dual_pool_from_ckpt(state_dict)\n    print(\"CKPT_KIND\", kind, \"use_dual_pool\", use_dual)\n\n    model = ECG4RowDigitalizer(backbone_name=\"resnet34\", target_width=TARGET_WIDTH, use_dual_pool=use_dual).to(DEVICE)\n    model.load_state_dict(state_dict, strict=True)\n    model.eval()\n    return model, use_dual\n\ndef split_into_4cols(x: np.ndarray):\n    W = int(x.shape[0])\n    q = W // 4\n    c1 = x[0:q]\n    c2 = x[q:2*q]\n    c3 = x[2*q:3*q]\n    c4 = x[3*q:W]\n    return c1, c2, c3, c4\n\ndef map_4row_to_12lead_from_layout(pred4: np.ndarray) -> dict:\n    p = np.asarray(pred4, dtype=np.float32)\n    if p.ndim != 2 or p.shape[0] != 4:\n        raise ValueError(f\"pred4 must be shape (4,W). got {p.shape}\")\n\n    row1 = p[0]\n    row2 = p[1]\n    row3 = p[2]\n    row4 = p[3]\n\n    r1_c1, r1_c2, r1_c3, r1_c4 = split_into_4cols(row1)\n    r2_c1, r2_c2, r2_c3, r2_c4 = split_into_4cols(row2)\n    r3_c1, r3_c2, r3_c3, r3_c4 = split_into_4cols(row3)\n\n    out = {}\n    out[\"I\"] = r1_c1\n    out[\"aVR\"] = r1_c2\n    out[\"V1\"] = r1_c3\n    out[\"V4\"] = r1_c4\n\n    out[\"II\"] = r2_c1\n    out[\"aVL\"] = r2_c2\n    out[\"V2\"] = r2_c3\n    out[\"V5\"] = r2_c4\n\n    out[\"III\"] = r3_c1\n    out[\"aVF\"] = r3_c2\n    out[\"V3\"] = r3_c3\n    out[\"V6\"] = r3_c4\n\n    out[\"II_long\"] = row4.astype(np.float32)\n    return out\n\ndef setup_hengck_imports_and_models():\n    assert_exists(HENGCK_BASE_DIR, \"HENGCK_BASE_DIR\")\n    sys.path.append(str(HENGCK_BASE_DIR))\n\n    global s0c, s1c, Stage0Net, Stage1Net, stage0_net, stage1_net, stage0_w, stage1_w\n\n    import stage0_common as s0c\n    import stage1_common as s1c\n    from stage0_model import Net as Stage0Net\n    from stage1_model import Net as Stage1Net\n\n    stage0_w = HENGCK_BASE_DIR / \"weight\" / \"stage0-last.checkpoint.pth\"\n    stage1_w = HENGCK_BASE_DIR / \"weight\" / \"stage1-last.checkpoint.pth\"\n    assert_exists(stage0_w, \"stage0 weight\")\n    assert_exists(stage1_w, \"stage1 weight\")\n\n    stage0_net = s0c.load_net(Stage0Net(pretrained=False), str(stage0_w)).to(DEVICE).eval()\n    stage1_net = s1c.load_net(Stage1Net(pretrained=False), str(stage1_w)).to(DEVICE).eval()\n\n@torch.no_grad()\ndef run_stage0(img_bgr):\n    img_rgb = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB)\n    batch = s0c.image_to_batch(img_rgb)\n    with torch.amp.autocast(AMP_DEVICE, enabled=USE_CUDA, dtype=torch.float32):\n        out = stage0_net(batch)\n    rotated, keypoint = s0c.output_to_predict(img_rgb, batch, out)\n    normalised, _, _ = s0c.normalise_by_homography(rotated, keypoint)\n    return normalised\n\n@torch.no_grad()\ndef run_stage1(stage0_img_rgb):\n    image = stage0_img_rgb\n    batch = {\"image\": torch.from_numpy(np.ascontiguousarray(image.transpose(2, 0, 1))).unsqueeze(0)}\n    with torch.amp.autocast(AMP_DEVICE, enabled=USE_CUDA, dtype=torch.float32):\n        out = stage1_net(batch)\n    gridpoint_xy, _ = s1c.output_to_predict(image, batch, out)\n    rectified = s1c.rectify_image(image, gridpoint_xy)\n    return rectified\n\ndef warp_raw_to_stage1_rgb(raw_png_path: str):\n    img_raw = cv2.imread(raw_png_path, cv2.IMREAD_COLOR)\n    if img_raw is None:\n        raise FileNotFoundError(raw_png_path)\n    s0 = run_stage0(img_raw)\n    s1 = run_stage1(s0)\n    return img_raw, s0, s1\n\n@torch.no_grad()\ndef predict_one_4row_raw(model: nn.Module, img_path: str, x_coords: torch.Tensor, y_coords: torch.Tensor):\n    _, _, s1_rgb = warp_raw_to_stage1_rgb(img_path)\n\n    input_4ch = stage1_rgb_to_input4ch(\n        s1_rgb, x_coords=x_coords, y_coords=y_coords, H=H_IN, W=W_IN, patch_size=PATCH_SIZE\n    )\n\n    x = input_4ch.unsqueeze(0).to(DEVICE, non_blocking=True)\n\n    with torch.amp.autocast(AMP_DEVICE, enabled=USE_CUDA, dtype=torch.float32):\n        pred = model(x).detach().float().cpu()[0]\n\n    pred4 = pred.numpy().astype(np.float32)\n    return pred4\n\ndef build_submission_minimal():\n    assert_exists(ROOT, \"ROOT\")\n    assert_exists(ROOT / \"test.csv\", \"test.csv\")\n    assert_exists(ROOT / \"sample_submission.parquet\", \"sample_submission.parquet\")\n    assert_exists(CKPT_ROOT, \"CKPT_ROOT\")\n\n    df_test = pd.read_csv(ROOT / \"test.csv\")\n    df_test[\"id\"] = df_test[\"id\"].astype(str)\n\n    sample = pd.read_parquet(ROOT / \"sample_submission.parquet\")[[\"id\"]]\n    check_sample_submission_format(sample)\n\n    x_coords, y_coords = build_coord_maps(H_IN, W_IN)\n\n    w = find_ecg_weight(\"ecg_high_res_ep50.pth\")\n    print(\"ECG_WEIGHT\", w)\n    model, use_dual = load_model_matching_ckpt(w)\n    print(\"MODEL use_dual_pool =\", use_dual, \"rnn.input_size =\", model.rnn.input_size)\n\n    rows = []\n    unique_ids = df_test[\"id\"].unique().tolist()\n\n    for idx, sample_id in enumerate(unique_ids):\n        img_path = ROOT / \"test\" / f\"{sample_id}.png\"\n        if not img_path.exists():\n            raise FileNotFoundError(str(img_path))\n\n        pred4 = predict_one_4row_raw(model, str(img_path), x_coords, y_coords)\n\n        lead_dict = map_4row_to_12lead_from_layout(pred4)\n\n        g = df_test[df_test[\"id\"] == sample_id]\n        for _, r in g.iterrows():\n            lead = str(r[\"lead\"])\n            n = int(r[\"number_of_rows\"])\n\n            if lead not in lead_dict:\n                raise KeyError(f\"missing lead after mapping: {lead}\")\n\n            sig = np.asarray(lead_dict[lead], dtype=np.float32).reshape(-1)\n            sig_n = interp_1d(sig, n)\n\n            for i, v in enumerate(sig_n):\n                rows.append({\"id\": f\"{sample_id}_{i}_{lead}\", \"value\": float(v)})\n\n        if idx % 50 == 0:\n            gc.collect()\n            print(\"progress\", idx, \"of\", len(unique_ids), \"rows\", len(rows))\n\n    df_sub = pd.DataFrame(rows)\n\n    df_sub = df_sub.set_index(\"id\").reindex(sample[\"id\"]).reset_index()\n\n    if df_sub[\"value\"].isna().any():\n        bad = df_sub[df_sub[\"value\"].isna()].head(20)\n        raise ValueError(f\"value NaN exists. first rows:\\n{bad}\")\n\n    if not (df_sub[\"id\"].values == sample[\"id\"].values).all():\n        raise ValueError(\"id order mismatch with sample_submission\")\n\n    df_sub.to_parquet(OUT_PARQUET, index=False)\n    df_sub.to_csv(OUT_CSV, index=False)\n\n    # print(\"OK\", str(OUT_PARQUET), df_sub.shape)\n    # print(\"OK\", str(OUT_CSV), df_sub.shape)\n\n    display(df_sub.head(20))\n    display(df_sub.describe(include=\"all\"))\n\n    return df_sub\n\nassert_exists(ROOT, \"ROOT\")\nassert_exists(HENGCK_BASE_DIR, \"HENGCK_BASE_DIR\")\nsetup_hengck_imports_and_models()\ndf_sub = build_submission_minimal()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T06:13:15.608665Z","iopub.execute_input":"2026-02-02T06:13:15.608939Z","iopub.status.idle":"2026-02-02T06:13:26.067088Z","shell.execute_reply.started":"2026-02-02T06:13:15.608911Z","shell.execute_reply":"2026-02-02T06:13:26.066555Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import matplotlib.pyplot as plt\n\n# def visualize_one_sample(model, sample_id: str):\n#     x_coords, y_coords = build_coord_maps(H_IN, W_IN)\n#     img_path = ROOT / \"test\" / f\"{sample_id}.png\"\n#     if not img_path.exists():\n#         raise FileNotFoundError(str(img_path))\n\n#     img_raw_bgr, s0_rgb, s1_rgb = warp_raw_to_stage1_rgb(str(img_path))\n#     pred4 = predict_one_4row_raw(model, str(img_path), x_coords, y_coords)\n#     lead_dict = map_4row_to_12lead_from_layout(pred4)\n\n#     fig = plt.figure(figsize=(16, 10))\n#     ax1 = fig.add_subplot(2, 2, 1)\n#     ax1.imshow(cv2.cvtColor(img_raw_bgr, cv2.COLOR_BGR2RGB))\n#     ax1.set_title(\"raw\")\n#     ax1.axis(\"off\")\n\n#     ax2 = fig.add_subplot(2, 2, 2)\n#     ax2.imshow(s0_rgb)\n#     ax2.set_title(\"stage0\")\n#     ax2.axis(\"off\")\n\n#     ax3 = fig.add_subplot(2, 2, 3)\n#     ax3.imshow(s1_rgb)\n#     ax3.set_title(\"stage1 rectified\")\n#     ax3.axis(\"off\")\n\n#     ax4 = fig.add_subplot(2, 2, 4)\n#     for r in range(4):\n#         ax4.plot(pred4[r])\n#     ax4.set_title(\"pred 4row\")\n#     ax4.set_xlabel(\"time index\")\n#     ax4.set_ylabel(\"value\")\n#     plt.tight_layout()\n#     plt.show()\n\n#     fig = plt.figure(figsize=(18, 12))\n#     lead_order = [\"I\",\"II\",\"III\",\"aVR\",\"aVL\",\"aVF\",\"V1\",\"V2\",\"V3\",\"V4\",\"V5\",\"V6\"]\n#     for i, lead in enumerate(lead_order, 1):\n#         ax = fig.add_subplot(4, 3, i)\n#         ax.plot(lead_dict[lead])\n#         ax.set_title(lead)\n#         ax.set_xlabel(\"time index\")\n#         ax.set_ylabel(\"value\")\n#     plt.tight_layout()\n#     plt.show()\n\n#     return {\n#         \"sample_id\": sample_id,\n#         \"raw_shape\": img_raw_bgr.shape,\n#         \"stage0_shape\": s0_rgb.shape,\n#         \"stage1_shape\": s1_rgb.shape,\n#         \"pred4_shape\": pred4.shape,\n#         \"lead_lengths\": {k: int(np.asarray(v).shape[0]) for k, v in lead_dict.items()},\n#     }\n\n# w = find_ecg_weight(\"ecg_high_res_ep50.pth\")\n# model, use_dual = load_model_matching_ckpt(w)\n\n# df_test = pd.read_csv(ROOT / \"test.csv\")\n# sample_id = str(df_test[\"id\"].iloc[0])\n\n# info = visualize_one_sample(model, sample_id)\n# display(pd.DataFrame([info]))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T06:13:26.068039Z","iopub.execute_input":"2026-02-02T06:13:26.068345Z","iopub.status.idle":"2026-02-02T06:13:26.072854Z","shell.execute_reply.started":"2026-02-02T06:13:26.068313Z","shell.execute_reply":"2026-02-02T06:13:26.072092Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import numpy as np\n# import pandas as pd\n# import torch\n# import torch.nn.functional as F\n# import matplotlib.pyplot as plt\n\n# torch.set_grad_enabled(False)\n\n# def _hi_ratio(x, split=0.15):\n#     x = np.asarray(x, dtype=np.float32).reshape(-1)\n#     if x.size < 8:\n#         return np.nan\n#     x = x - x.mean()\n#     spec = (np.abs(np.fft.rfft(x)) ** 2).astype(np.float64)\n#     freqs = np.fft.rfftfreq(x.size, d=1.0)\n#     fmax = freqs.max() if freqs.size > 0 else 0.0\n#     if fmax <= 0:\n#         return np.nan\n#     cut = split * fmax\n#     hi = spec[freqs > cut].sum()\n#     tot = spec.sum() + 1e-12\n#     return float(hi / tot)\n\n# def _der_energy(x):\n#     x = np.asarray(x, dtype=np.float32).reshape(-1)\n#     if x.size < 3:\n#         return np.nan\n#     d = np.diff(x)\n#     return float(np.mean(d * d))\n\n# @torch.no_grad()\n# def get_out_pre(model, x):\n#     fft_map = x[:, 1:2, :, :]\n#     z = model.stem(x)\n#     feat = model.backbone(z)\n#     if not model.disable_inhibitor:\n#         feat = model.inhibitor(feat, fft_map)\n\n#     seq_avg = model.pool_avg(feat)\n#     if model.use_dual_pool:\n#         seq_max = model.pool_max(feat)\n#         seq = torch.cat([seq_avg, seq_max], dim=1)\n#     else:\n#         seq = seq_avg\n\n#     seq = seq.squeeze(2).permute(0, 2, 1)\n#     rnn_out, _ = model.rnn(seq)\n\n#     r1 = model.head_row1(rnn_out)\n#     r2 = model.head_row2(rnn_out)\n#     r3 = model.head_row3(rnn_out)\n#     r4 = model.head_row4(rnn_out)\n\n#     out_pre = torch.cat([r1, r2, r3, r4], dim=-1).permute(0, 2, 1)\n#     return out_pre\n\n# def _downsample_like(pre_275, post_2200):\n#     pre_len = int(pre_275.shape[-1])\n#     post_len = int(post_2200.shape[-1])\n#     if post_len == pre_len:\n#         return post_2200\n#     idx = np.linspace(0, post_len - 1, pre_len).astype(np.int64)\n#     return post_2200[..., idx]\n\n# @torch.no_grad()\n# def compare_interp_modes(sample_id):\n#     x_coords, y_coords = build_coord_maps(H_IN, W_IN)\n\n#     w = find_ecg_weight(\"ecg_high_res_ep50.pth\")\n#     model, use_dual = load_model_matching_ckpt(w)\n\n#     img_path = ROOT / \"test\" / f\"{sample_id}.png\"\n#     img_raw_bgr, s0_rgb, s1_rgb = warp_raw_to_stage1_rgb(str(img_path))\n\n#     input_4ch, gray_u8, grid_map = stage1_rgb_to_input4ch(\n#         s1_rgb, x_coords=x_coords, y_coords=y_coords, H=H_IN, W=W_IN, patch_size=PATCH_SIZE\n#     )\n\n#     x = input_4ch.unsqueeze(0).to(DEVICE, non_blocking=True)\n\n#     with torch.amp.autocast(AMP_DEVICE, enabled=USE_CUDA, dtype=torch.float32):\n#         out_pre = get_out_pre(model, x)\n#         out_pre_tanh = torch.tanh(out_pre)\n\n#         post_linear = F.interpolate(out_pre_tanh, size=TARGET_WIDTH, mode=\"linear\", align_corners=False)\n#         post_nearest = F.interpolate(out_pre_tanh, size=TARGET_WIDTH, mode=\"nearest\")\n\n#     pre = out_pre_tanh.detach().float().cpu().numpy()[0]\n#     lin = post_linear.detach().float().cpu().numpy()[0]\n#     nea = post_nearest.detach().float().cpu().numpy()[0]\n\n#     lin_ds = _downsample_like(pre, lin)\n#     nea_ds = _downsample_like(pre, nea)\n\n#     rows = []\n#     for r in range(4):\n#         rows.append({\n#             \"row\": r + 1,\n#             \"pre_len\": int(pre[r].shape[0]),\n#             \"lin_len\": int(lin[r].shape[0]),\n#             \"nea_len\": int(nea[r].shape[0]),\n#             \"pre_hi\": _hi_ratio(pre[r]),\n#             \"lin_hi_full\": _hi_ratio(lin[r]),\n#             \"nea_hi_full\": _hi_ratio(nea[r]),\n#             \"lin_hi_ds\": _hi_ratio(lin_ds[r]),\n#             \"nea_hi_ds\": _hi_ratio(nea_ds[r]),\n#             \"pre_der\": _der_energy(pre[r]),\n#             \"lin_der_full\": _der_energy(lin[r]),\n#             \"nea_der_full\": _der_energy(nea[r]),\n#             \"lin_der_ds\": _der_energy(lin_ds[r]),\n#             \"nea_der_ds\": _der_energy(nea_ds[r]),\n#             \"pre_std\": float(np.std(pre[r])),\n#             \"lin_std_full\": float(np.std(lin[r])),\n#             \"nea_std_full\": float(np.std(nea[r])),\n#         })\n\n#     df = pd.DataFrame(rows)\n#     print(\"use_dual_pool\", use_dual, \"rnn_input_size\", int(model.rnn.input_size))\n#     print(df)\n\n#     fig = plt.figure(figsize=(18, 9))\n#     fig.suptitle(f\"interp mode compare sample_id {sample_id}\")\n\n#     for r in range(4):\n#         ax = plt.subplot(4, 1, r + 1)\n#         ax.plot(pre[r], linewidth=0.9, label=\"pre 275 tanh\")\n#         ax.plot(lin_ds[r], linewidth=0.9, label=\"linear up then down\")\n#         ax.plot(nea_ds[r], linewidth=0.9, label=\"nearest up then down\")\n#         ax.legend(loc=\"upper right\", fontsize=8)\n\n#     plt.tight_layout()\n#     plt.show()\n\n#     return df\n\n# df_test = pd.read_csv(ROOT / \"test.csv\")\n# sid = str(df_test[\"id\"].astype(str).iloc[0])\n# df_cmp = compare_interp_modes(sid)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-02T06:13:26.074404Z","iopub.execute_input":"2026-02-02T06:13:26.074637Z","iopub.status.idle":"2026-02-02T06:13:26.087898Z","shell.execute_reply.started":"2026-02-02T06:13:26.074616Z","shell.execute_reply":"2026-02-02T06:13:26.087203Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}}]}