{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":97984,"databundleVersionId":14096757,"sourceType":"competition"},{"sourceId":13731160,"sourceType":"datasetVersion","datasetId":8733970},{"sourceId":13746387,"sourceType":"datasetVersion","datasetId":8747012},{"sourceId":14216669,"sourceType":"datasetVersion","datasetId":8989372},{"sourceId":14352323,"sourceType":"datasetVersion","datasetId":9150821},{"sourceId":14355772,"sourceType":"datasetVersion","datasetId":9166780},{"sourceId":14511751,"sourceType":"datasetVersion","datasetId":9268705},{"sourceId":14517996,"sourceType":"datasetVersion","datasetId":9272362},{"sourceId":14524678,"sourceType":"datasetVersion","datasetId":9276665},{"sourceId":14549997,"sourceType":"datasetVersion","datasetId":9204491},{"sourceId":14566159,"sourceType":"datasetVersion","datasetId":8810161},{"sourceId":14566236,"sourceType":"datasetVersion","datasetId":9138572},{"sourceId":14575133,"sourceType":"datasetVersion","datasetId":8810164}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Liu parts","metadata":{}},{"cell_type":"code","source":"\"\"\"\n================================================================================\nECG INFERENCE - 16dB SOLUTION + TTA (flip + intensity) + Physics Correction\n+ VALIDATION (Train-set scoring: official-like SNR(dB) with alignment)\n================================================================================\nProperly using hengck23's pretrained models with fixed imports\nAdded: Einthoven's Law (II = I + III) Post-processing (Expanded to correct II)\nAdded: Validation scoring on train images using provided scoring functions\n\n[MOD]\n① Validation Outputに、GT(GrandTrue) / Pred / 差分 を含むCSVを出力\n   - per-image: OUT_DIR/validation_output/detail/{image_id}-{type_id}.csv\n   - all-in-one: OUT_DIR/validation_output/validation_detail_all.csv\n② トータルSNR(dB)（全validationサンプルで signal/noise power 合算）を計算して表示\n③ アイントホーフェンの法則を拡張し、Lead II も補正対象に追加 (II = I + III)\n\n[NEW]\n★ Stage2 を ResNet34(Net3互換) に差し替え + ckpt[\"cfg\"] 反映\n   - timm resnet34.a3_in1k + MyCoordUnetDecoder(scale=[2,2,2,2], dec=[128,64,32,16])\n   - 特に X_SCALE / T0 / T1 / PAD_MULTIPLE / ZERO_MV / MV_TO_PIXEL を推論側へ反映\n\n[NEW-MOD]\n★ pixel_to_series を「局所 Sub-pixel Soft-argmax」に置き換え（改良版）\n   - 全YでのSoft-argmaxは禁止（上端/罫線に吸着しやすい）\n   - argmaxでトレースを確定→その近傍(±R)だけでsoft化してsub-pixel化\n   - sigmoid確率pではなく logit(p) を使う（power極大依存を緩和）\n   - 欠損列は NaN→時間方向に補間 + 端は線形外挿（0埋め禁止）\n\n[NEW-MOD-2]\n★ Stage1 の gridpoint_xy に「安全な」後処理を追加（stage1_commonは未改造）\n   - griddata の外挿で破綻しやすい問題を回避\n   - affine + residual を作り、residual を masked/normalized gaussian で平滑化\n   - クリップするのは residual ではなく「平滑化で動かした量 delta」\n   - 単調性/範囲チェックに通らない場合は自動で元の gridpoint_xy に戻す\n\n[NEW-MOD-3]\n★ Stage0 change_color を常時ONで適用（Stage0モデル入力のみ）\n   - Stage0の幾何復元は元画像基準（s0_outputは image_rgb を使う）\n\n[NEW-MOD-3 + TTA requested]\n★ Stage0 change_color ON/OFF を追加TTAとして扱う（2パターン）\n   ① Stage0 change_color ON/OFF の分岐それぞれで Stage2 まで回し、series を平均/中央値でブレンド（Stage2側TTA）\n   ② Stage0 change_color ON/OFF の分岐それぞれで Stage1 まで回し、gridpoint をブレンドして rectified を作り、Stage2 は1回だけ（Stage1だけTTA）\n\n[NEW: STAGE2 ENSEMBLE / BLEND]\n★ Stage2 ckpt 複数をBlendする\n   - 3方式を実装（精度が高い可能性が高い順）\n     1) pixel_logit_mean  : Pixel-level Logit Averaging（推奨・Default）\n     2) series_median     : Signal-level Median（堅牢性重視）\n     3) series_mean       : Signal-level Mean（平滑化）\n   - X_SCALEが異なっても、pixel系は共通幅へリサイズしてから平均\n   - series系は共通ROI(base)を各モデルscaleへ写像して揃えてからseries化→median/mean\n\n[POST ORDER MOD]\n★ 要望: POST_II_HEAD_MODE=\"blend\" を先に適用 → その後 POST_EINTHOVEN_MODE=\"project\" で II重み高で補正\n\n[NEW-MOD-4  (MINIMAL ADDITION)]\n★ Stage1 Quality Gate (sparse metrics) + ZERO-FILL\n   - Stage1 の補間前 gridpoint_xy (sparse) を stage1_common.interpolate_mapping の一時差し替えで捕捉\n   - hole_rate / mono_x / mono_y は sparse で計算\n   - n_points は stage1_out[\"gridpoint\"] の connected components 数で計算\n   - gate 条件を満たさない画像は以降の処理を打ち切り、submission を全ゼロにする\n     gate: hole_rate < 0.10 AND mono_x > 0.95 AND mono_y > 0.95 AND 2000 < n_points < 3000\n================================================================================\n\"\"\"\n\nimport os\nimport sys\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nfrom tabulate import tabulate\nfrom tqdm import tqdm\nimport warnings\nimport contextlib\nimport timm  # ResNet/ConvNeXt 用\nfrom dataclasses import dataclass\nfrom concurrent.futures import ThreadPoolExecutor\n\nfrom scipy.ndimage import gaussian_filter  # Stage1 post-process\nfrom scipy.signal import savgol_filter  # [ADDED] final Savgol\n\nwarnings.filterwarnings(\"ignore\")\n\nimport threading\n_S1_SPARSE_LOCK = threading.Lock()\n\n# Install cc3d silently\ntry:\n    import cc3d\nexcept Exception:\n    import subprocess\n\n    subprocess.run(\n        [\n            \"pip\",\n            \"install\",\n            \"connected-components-3d\",\n            \"--no-index\",\n            \"--find-links=file:///kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet/setup\",\n            \"-q\",\n        ],\n        capture_output=True,\n    )\n    import cc3d\n\n\n# ============================================================\n# PATHS / CONFIG\n# ============================================================\nprint(\"=\" * 80)\nprint(\"ECG INFERENCE + VALIDATION (SNR dB) [Stage2 = ResNet(Net3) + Local Subpixel Soft-argmax (logit)]\")\nprint(\"=\" * 80)\n\nbase_path = \"/kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet\"\nsys.path.insert(0, base_path)\n\nKAGGLE_DIR = \"/kaggle/input/physionet-ecg-image-digitization\"\nWEIGHT_DIR = f\"{base_path}/weight\"\nOUT_DIR = \"/kaggle/working/output-submit\"\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nFLOAT_TYPE = torch.float16  # keep as original\n\nos.makedirs(f\"{OUT_DIR}/normalised\", exist_ok=True)\nos.makedirs(f\"{OUT_DIR}/rectified\", exist_ok=True)\nos.makedirs(f\"{OUT_DIR}/digitalised\", exist_ok=True)\n\nprint(f\"\\n🔧 Device: {DEVICE}\")\nprint(f\"📁 Weights: {WEIGHT_DIR}\")\nprint(f\"📁 OUT_DIR: {OUT_DIR}\")\n\n# ============================================================\n# CUDA empty_cache policy (reduced frequency; default OFF)\n#   - Set CUDA_EMPTY_CACHE_ENABLE=True to enable periodic empty_cache.\n#   - CUDA_EMPTY_CACHE_EVERY controls cadence (e.g., 50 -> every 50 images).\n#   - CUDA_EMPTY_CACHE_ON_ERROR triggers a one-off empty_cache on exceptions.\n# ============================================================\nCUDA_EMPTY_CACHE_ENABLE = False\nCUDA_EMPTY_CACHE_EVERY = 50\nCUDA_EMPTY_CACHE_ON_ERROR = True\n\n\ndef maybe_cuda_empty_cache(step: int | None = None, force: bool = False):\n    \"\"\"Reduced-frequency CUDA cache clearing.\n\n    NOTE:\n      - No recursion.\n      - force=True triggers immediate empty_cache.\n    \"\"\"\n    if DEVICE != \"cuda\" or (not torch.cuda.is_available()):\n        return\n    if force:\n        torch.cuda.empty_cache()\n        return\n    if not CUDA_EMPTY_CACHE_ENABLE:\n        return\n    every = int(CUDA_EMPTY_CACHE_EVERY)\n    if every <= 0 or step is None:\n        return\n    if (int(step) % every) == 0:\n        torch.cuda.empty_cache()\n\n\n# ============================================================\n# [NEW] 2GPU branch-parallel controls (MINIMAL ADDITION)\n# ============================================================\nN_GPU = torch.cuda.device_count() if torch.cuda.is_available() else 0\nUSE_TWO_GPU_BRANCH = (DEVICE == \"cuda\") and (N_GPU >= 2)\nDEVICE0 = \"cuda:0\" if DEVICE == \"cuda\" else DEVICE\nDEVICE1 = \"cuda:1\" if USE_TWO_GPU_BRANCH else DEVICE0\n\nif USE_TWO_GPU_BRANCH:\n    print(\"🚀 2GPU detected: Stage0 color TTA ON-branch -> cuda:0, OFF-branch -> cuda:1 (parallel)\")\nelse:\n    print(\"ℹ️  <2GPU: Stage0 color TTA branches run sequentially on single device (original behavior)\")\n\n_executor = ThreadPoolExecutor(max_workers=2) if USE_TWO_GPU_BRANCH else None\n\n\ndef _set_cuda_device(device_str: str):\n    if device_str.startswith(\"cuda:\"):\n        torch.cuda.set_device(int(device_str.split(\":\")[1]))\n    elif device_str == \"cuda\":\n        torch.cuda.set_device(0)\n\n\ndef _to_device_batch(batch: dict, device_str: str):\n    \"\"\"Force batch tensors to the specified device.\"\"\"\n    dev = torch.device(device_str) if device_str.startswith(\"cuda\") else torch.device(\"cpu\")\n    out = {}\n    for k, v in batch.items():\n        if torch.is_tensor(v):\n            out[k] = v.to(dev, non_blocking=True)\n        else:\n            out[k] = v\n    return out\n\n\n# ============================================================\n# [ADDED] POSTPROCESS CFG (switchable)\n# ============================================================\nPOST_EINTHOVEN_ENABLE = True\nPOST_EINTHOVEN_MODE = \"project\"   # \"baseline\" or \"project\" or \"off\"\nPOST_EINTHOVEN_WEIGHT = 0.5       # used when mode=\"baseline\"\nPOST_EINTHOVEN_STRENGTH = 1.0     # used when mode=\"project\" (0.5~0.8 recommended, 1.0 strong)\nPOST_EINTHOVEN_PROJECT_II_FRAC = 0.6  # 0~1, larger => II changes more, I/III change less\n\nPOST_II_HEAD_MODE = \"blend\"       # \"overwrite\" or \"blend\" or \"off\"\nPOST_II_BLEND_ALPHA = 0.5\nPOST_II_BLEND_ADAPTIVE = False    # keep False for minimal predictable behavior\n\nPOST_SAVGOL_ENABLE = True\nPOST_SAVGOL_POLYORDER = 2\nPOST_SAVGOL_LEN_THRESHOLD = 9000\nPOST_SAVGOL_WIN_SHORT = 5\nPOST_SAVGOL_WIN_LONG = 9\n\n\ndef savgol_final_1d(x: np.ndarray) -> np.ndarray:\n    \"\"\"[ADDED] final-stage smoothing (switchable)\"\"\"\n    x = x.astype(np.float32, copy=False)\n    if not POST_SAVGOL_ENABLE:\n        return x\n    n = len(x)\n    if n < 7:\n        return x\n    w = POST_SAVGOL_WIN_SHORT if n < POST_SAVGOL_LEN_THRESHOLD else POST_SAVGOL_WIN_LONG\n    w = int(w)\n    if w >= n:\n        w = n - 1\n    if w % 2 == 0:\n        w -= 1\n    if w < 5 or w <= int(POST_SAVGOL_POLYORDER):\n        return x\n    return savgol_filter(x, window_length=w, polyorder=int(POST_SAVGOL_POLYORDER)).astype(np.float32)\n\n\ndef blend_ii_head(long_ii: np.ndarray, short_ii: np.ndarray) -> np.ndarray:\n    \"\"\"[ADDED] II_short と long II 先頭の更新（overwrite/blend/off）\"\"\"\n    long_ii = long_ii.astype(np.float32, copy=True)\n    short_ii = short_ii.astype(np.float32, copy=False)\n    n = min(len(long_ii), len(short_ii))\n    if n < 5:\n        return long_ii\n\n    mode = str(POST_II_HEAD_MODE)\n    if mode == \"off\":\n        return long_ii\n\n    if mode == \"overwrite\":\n        long_ii[:n] = short_ii[:n]\n        return long_ii\n\n    if mode == \"blend\":\n        alpha = float(np.clip(POST_II_BLEND_ALPHA, 0.0, 1.0))\n        if POST_II_BLEND_ADAPTIVE:\n            a = long_ii[:n]\n            b = short_ii[:n]\n            rmse = float(np.sqrt(((a - b) ** 2).mean()))\n            scale = float(np.std(a) + 1e-6)\n            nr = rmse / scale\n            gate = float(np.clip(1.0 - nr, 0.0, 1.0))\n            alpha *= gate\n        long_ii[:n] = (1.0 - alpha) * long_ii[:n] + alpha * short_ii[:n]\n        return long_ii\n\n    return long_ii\n\n\ndef apply_einthoven_switchable(series_by_lead: dict):\n    \"\"\"Apply Einthoven correction on I/II/III (II = I + III).\n\n    - baseline: only adjusts II toward (I+III)\n    - project : distributes correction among I/II/III with II fraction control\n\n    This function is intentionally minimal and only touches keys present.\n    \"\"\"\n    if not POST_EINTHOVEN_ENABLE:\n        return\n    mode = str(POST_EINTHOVEN_MODE).lower()\n    if mode in (\"off\", \"none\", \"disabled\"):\n        return\n    if (\"I\" not in series_by_lead) or (\"II\" not in series_by_lead) or (\"III\" not in series_by_lead):\n        return\n\n    I = np.asarray(series_by_lead[\"I\"], dtype=np.float32)\n    II = np.asarray(series_by_lead[\"II\"], dtype=np.float32)\n    III = np.asarray(series_by_lead[\"III\"], dtype=np.float32)\n\n    n = int(min(I.size, II.size, III.size))\n    if n < 8:\n        return\n\n    I0 = I[:n].copy()\n    II0 = II[:n].copy()\n    III0 = III[:n].copy()\n\n    r = (I0 + III0 - II0)  # residual (want 0)\n\n    if mode == \"baseline\":\n        w = float(np.clip(POST_EINTHOVEN_WEIGHT, 0.0, 1.0))\n        II[:n] = ((1.0 - w) * II0 + w * (I0 + III0)).astype(np.float32)\n\n    elif mode == \"project\":\n        strength = float(np.clip(POST_EINTHOVEN_STRENGTH, 0.0, 2.0))\n        ii_frac = float(np.clip(POST_EINTHOVEN_PROJECT_II_FRAC, 0.0, 1.0))\n\n        dII = strength * ii_frac * r\n        dI = -strength * (1.0 - ii_frac) * 0.5 * r\n        dIII = -strength * (1.0 - ii_frac) * 0.5 * r\n\n        I[:n] = (I0 + dI).astype(np.float32)\n        II[:n] = (II0 + dII).astype(np.float32)\n        III[:n] = (III0 + dIII).astype(np.float32)\n\n    else:\n        return\n\n    series_by_lead[\"I\"] = I\n    series_by_lead[\"II\"] = II\n    series_by_lead[\"III\"] = III\n\n\n# ============================================================\n# Stage0 change_color (BASELINE ALWAYS ON)\n# ============================================================\n\nSTAGE0_COLOR_TTA = True\n# 1: Stage0 change_color ON/OFF branches run to Stage2 and blend series\n# 2: Stage0 change_color ON/OFF branches run to Stage1 and blend gridpoint (Stage2 once)\nSTAGE0_COLOR_TTA_PATTERN = 1\nSTAGE0_COLOR_TTA_REDUCE = \"mean\"  # \"mean\" or \"median\"\nSTAGE0_COLOR_TTA_VERBOSE = False\n\nprint(\"🎨 Stage0 change_color: BASELINE ALWAYS ON (Stage0 model input only)\")\nif STAGE0_COLOR_TTA:\n    print(f\"🧪 Stage0 change_color ON/OFF TTA: ENABLED  (pattern={STAGE0_COLOR_TTA_PATTERN}, reduce={STAGE0_COLOR_TTA_REDUCE})\")\nelse:\n    print(\"🧪 Stage0 change_color ON/OFF TTA: DISABLED (baseline=ON only)\")\n\n\ndef change_color(image_rgb: np.ndarray) -> np.ndarray:\n    \"\"\"Stage0 用のコントラスト強調\"\"\"\n    hsv = cv2.cvtColor(image_rgb, cv2.COLOR_RGB2HSV)\n    h, s, v = cv2.split(hsv)\n\n    v_denoised = cv2.fastNlMeansDenoising(v, h=6)\n\n    std = float(np.std(v_denoised))\n    clip_limit = max(1.0, min(3.5, 2.0 + std / 25.0))\n    clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=(8, 8))\n    v_enhanced = clahe.apply(v_denoised)\n\n    hsv_enhanced = cv2.merge([h, s, v_enhanced])\n    out = cv2.cvtColor(hsv_enhanced, cv2.COLOR_HSV2RGB)\n    return out\n\n\n# =========================\n# Local Subpixel parameters\n# =========================\nSUBPIX_POWER = 40.0\nSUBPIX_THRESH = 0.05\nSUBPIX_RADIUS = 8\nSUBPIX_LOGIT_EPS = 1e-6\nprint(f\"🎯 Subpixel params: power={SUBPIX_POWER}, thresh={SUBPIX_THRESH}, radius={SUBPIX_RADIUS} (logit eps={SUBPIX_LOGIT_EPS})\")\n\n# ------------------------------------------------------------\n# Stage2 training-cfg defaults (will be overwritten per-ckpt from ckpt[\"cfg\"])\n# ------------------------------------------------------------\nST2_X_SCALE = 1\nST2_T0 = 118\nST2_T1 = 2080\nST2_PAD_MULTIPLE = 32\nST2_ZERO_MV = [703.5, 987.5, 1271.5, 1531.5]\nST2_MV_TO_PIXEL = 78.0\nST2_ARCH = \"resnet34.a3_in1k\"\nST2_USE_IMAGENET_NORM = False\n\n# ============================================================\n# Stage2 ENSEMBLE CKPTS + BLEND MODE\n# ============================================================\nSTAGE2_CKPTS = [\n    # \"/kaggle/input/physionetexp8002/stage2_best.pth\",\n    # \"/kaggle/input/physionetexp8003/exp8001stage2_best.pth\",\n    \"/kaggle/input/physionetexp8003/exp8003stage2_best.pth\",\n    # \"/kaggle/input/physionetexp8008/exp8004stage2_best.pth\",\n    \"/kaggle/input/physionetexp8008/exp8007stage2_epoch05.pth\",\n    \"/kaggle/input/physionetexp8008/exp8008stage2_best.pth\",\n]\n\n# 3方式\nSTAGE2_BLEND_MODE = \"pixel_logit_mean\"\n\n# ============================================================\n# [NEW] Stage1 Gate thresholds (as requested)\n# ============================================================\nSTAGE1_GATE_ENABLE = True\nSTAGE1_GATE_HOLE_MAX = 0.10\nSTAGE1_GATE_MONO_MIN = 0.95\nSTAGE1_GATE_NPOINTS_MIN = 2000\nSTAGE1_GATE_NPOINTS_MAX = 3000\nSTAGE1_GATE_POINT_THR = 0.5  # connected components threshold on stage1_out[\"gridpoint\"]\n\nprint(f\"🧱 Stage1 Gate: enable={STAGE1_GATE_ENABLE} hole<{STAGE1_GATE_HOLE_MAX} mono>{STAGE1_GATE_MONO_MIN} n_points({STAGE1_GATE_NPOINTS_MIN},{STAGE1_GATE_NPOINTS_MAX}) thr={STAGE1_GATE_POINT_THR}\")\n\n\n# safe autocast helper (works on CPU too)\ndef amp_autocast(device: str, dtype=torch.float16, enabled: bool = True):\n    if device.startswith(\"cuda\"):\n        return torch.amp.autocast(\"cuda\", dtype=dtype, enabled=enabled)\n    return contextlib.nullcontext()\n\n\n# ============================================================\n# Local Sub-pixel Soft-argmax (Improved)\n# ============================================================\ndef _interp_with_linear_extrap(cols_f: np.ndarray, y: np.ndarray, fill_value: float | None = None) -> np.ndarray:\n    \"\"\"NaN を線形補間し、端は線形外挿\"\"\"\n    ok = np.isfinite(y)\n    if ok.sum() >= 2:\n        x_ok = cols_f[ok]\n        y_ok = y[ok].astype(np.float32)\n\n        y_out = np.interp(cols_f, x_ok, y_ok).astype(np.float32)\n\n        # left extrap\n        if x_ok[0] > cols_f[0]:\n            x0, x1 = x_ok[0], x_ok[1]\n            y0, y1 = y_ok[0], y_ok[1]\n            slope = (y1 - y0) / max(x1 - x0, 1e-6)\n            m = cols_f < x0\n            y_out[m] = y0 + slope * (cols_f[m] - x0)\n\n        # right extrap\n        if x_ok[-1] < cols_f[-1]:\n            x0, x1 = x_ok[-2], x_ok[-1]\n            y0, y1 = y_ok[-2], y_ok[-1]\n            slope = (y1 - y0) / max(x1 - x0, 1e-6)\n            m = cols_f > x1\n            y_out[m] = y1 + slope * (cols_f[m] - x1)\n\n        return y_out\n\n    if ok.sum() == 1:\n        return np.full_like(y, y[ok][0], dtype=np.float32)\n\n    if fill_value is None:\n        return y.astype(np.float32)\n\n    return np.full_like(y, fill_value, dtype=np.float32)\n\n\n\ndef pixel_to_series_subpixel_local(\n    pixel: np.ndarray,\n    zero_mv: list,\n    length: int | None,\n    power: float = 40.0,\n    threshold: float = 0.02,\n    radius: int = 8,\n    eps: float = 1e-6,\n) -> np.ndarray:\n    \"\"\"pixel: (4,H,W) float32 prob map (sigmoid output)\"\"\"\n    C, H, W = pixel.shape\n    series = []\n\n    cols = np.arange(W, dtype=np.int32)\n    cols_f = cols.astype(np.float32)\n\n    radius = max(int(radius), 0)\n    offs = np.arange(-radius, radius + 1, dtype=np.int32)[:, None]  # (K,1)\n\n    for j in range(C):\n        p = pixel[j].astype(np.float32)  # (H,W)\n        amax = np.argmax(p, axis=0).astype(np.int32)  # (W,)\n\n        p_peak = p[amax, cols]\n        miss = p_peak < float(threshold)\n\n        if radius == 0:\n            y = amax.astype(np.float32)\n            y[miss] = np.nan\n        else:\n            idx = amax[None, :] + offs  # (K,W)\n            idx = np.clip(idx, 0, H - 1)\n\n            pwin = p[idx, cols[None, :]]  # (K,W)\n\n            # use logit(p)\n            pclip = np.clip(pwin, float(eps), 1.0 - float(eps))\n            logit = np.log(pclip) - np.log1p(-pclip)  # (K,W)\n\n            logit_max = np.max(logit, axis=0, keepdims=True)\n            logits = (float(power) * (logit - logit_max)).astype(np.float64)  # stable exp\n            w = np.exp(logits)  # float64\n\n            wsum = np.sum(w, axis=0) + 1e-12\n            y = (w * idx.astype(np.float64)).sum(axis=0) / wsum\n            y = y.astype(np.float32)\n            y[miss] = np.nan\n\n        # fill NaNs: interp + linear extrap at edges\n        y = _interp_with_linear_extrap(cols_f, y, fill_value=float(zero_mv[j]))\n        series.append(y)\n\n    series = np.stack(series).astype(np.float32)  # (4,W)\n\n    if length is not None and length != W:\n        t = torch.from_numpy(series).unsqueeze(1)  # (4,1,W)\n        t = F.interpolate(t, size=length, mode=\"linear\", align_corners=False)\n        series = t.squeeze(1).cpu().numpy().astype(np.float32)\n\n    return series\n\n\n# ============================================================\n# Stage1 post-process (safe)\n# ============================================================\nUSE_STAGE1_POSTPROC = False\nSTAGE1_SIGMA = 0.8\nSTAGE1_MAX_DELTA = 0.25\nSTAGE1_RMS_THRESH = 0.9\nSTAGE1_DEBUG = False\n\nSTAGE1_MIN_MONO_FRAC = 0.85  # monotonic check (x/y increasing)\nSTAGE1_MIN_STEP = 0.20  # minimum positive step in pixel-space\nSTAGE1_BOUND_MARGIN = 2.0  # allow small margin outside image\n\n\ndef _grid_is_reasonable(g: np.ndarray, image_hw: tuple[int, int]) -> bool:\n    \"\"\"単調性（折り返し）と画像範囲をざっくり検査。\"\"\"\n    Himg, Wimg = image_hw\n    gx = g[..., 0]\n    gy = g[..., 1]\n    finite = np.isfinite(gx) & np.isfinite(gy)\n    if finite.sum() < 4:\n        return False\n\n    in_range = (\n        (gx > -STAGE1_BOUND_MARGIN)\n        & (gx < (Wimg - 1) + STAGE1_BOUND_MARGIN)\n        & (gy > -STAGE1_BOUND_MARGIN)\n        & (gy < (Himg - 1) + STAGE1_BOUND_MARGIN)\n    )\n    if (in_range[finite].mean() < 0.98):\n        return False\n\n    dx = np.diff(gx, axis=1)\n    dy = np.diff(gy, axis=0)\n\n    finite_dx = np.isfinite(dx)\n    finite_dy = np.isfinite(dy)\n\n    if finite_dx.sum() > 0:\n        frac_pos_x = (dx[finite_dx] > STAGE1_MIN_STEP).mean()\n        if frac_pos_x < STAGE1_MIN_MONO_FRAC:\n            return False\n    if finite_dy.sum() > 0:\n        frac_pos_y = (dy[finite_dy] > STAGE1_MIN_STEP).mean()\n        if frac_pos_y < STAGE1_MIN_MONO_FRAC:\n            return False\n\n    return True\n\n\n\ndef stage1_affine_residual_smooth_safe(\n    grid_xy: np.ndarray,\n    image_hw: tuple[int, int],\n    sigma: float = 0.8,\n    max_delta: float = 0.25,\n    rms_thresh: float = 0.9,\n):\n    \"\"\"Stage1 grid smoothing with safety checks.\"\"\"\n    g0 = grid_xy.astype(np.float32)\n    valid0 = (g0[..., 0] > 0) & (g0[..., 1] > 0)\n    holes = int((~valid0).sum())\n\n    Hg, Wg = g0.shape[:2]\n    if int(valid0.sum()) < 4:\n        return g0, 999.0, False\n\n    I, J = np.meshgrid(np.arange(Wg, dtype=np.float32), np.arange(Hg, dtype=np.float32))\n    ii = I[valid0].reshape(-1)\n    jj = J[valid0].reshape(-1)\n    x = g0[..., 0][valid0].reshape(-1)\n    y = g0[..., 1][valid0].reshape(-1)\n\n    A = np.stack([ii, jj, np.ones_like(ii)], axis=1).astype(np.float32)\n    px, *_ = np.linalg.lstsq(A, x.astype(np.float32), rcond=None)\n    py, *_ = np.linalg.lstsq(A, y.astype(np.float32), rcond=None)\n    px = px.astype(np.float32)\n    py = py.astype(np.float32)\n\n    Aall = np.stack([I.reshape(-1), J.reshape(-1), np.ones(Hg * Wg, np.float32)], axis=1)\n    x_aff = (Aall @ px).reshape(Hg, Wg)\n    y_aff = (Aall @ py).reshape(Hg, Wg)\n    aff = np.stack([x_aff, y_aff], axis=2).astype(np.float32)\n\n    r = np.zeros_like(aff, dtype=np.float32)\n    r[valid0] = (g0[valid0] - aff[valid0]).astype(np.float32)\n\n    r_valid = r[valid0]\n    rms = float(np.sqrt(np.mean((r_valid[:, 0] ** 2) + (r_valid[:, 1] ** 2))))\n\n    if holes == 0 and rms <= float(rms_thresh):\n        return g0, rms, False\n\n    mask = valid0.astype(np.float32)\n    if float(sigma) > 0:\n        m_s = gaussian_filter(mask, sigma=float(sigma), mode=\"nearest\")\n        r_s = np.zeros_like(r, dtype=np.float32)\n        for c in range(2):\n            num = gaussian_filter(r[..., c] * mask, sigma=float(sigma), mode=\"nearest\")\n            r_s[..., c] = num / (m_s + 1e-6)\n    else:\n        r_s = r\n\n    delta = (r_s - r).astype(np.float32)\n    delta = np.clip(delta, -float(max_delta), +float(max_delta))\n\n    out = (aff + (r + delta)).astype(np.float32)\n\n    if not _grid_is_reasonable(out, image_hw):\n        return g0, rms, False\n\n    return out, rms, True\n\n\n# ============================================================\n# DATA INSPECTION\n# ============================================================\ntrain_df = pd.read_csv(f\"{KAGGLE_DIR}/train.csv\")\ntest_df = pd.read_csv(f\"{KAGGLE_DIR}/test.csv\")\ntrain_df[\"id\"] = train_df[\"id\"].astype(str)\ntest_df[\"id\"] = test_df[\"id\"].astype(str)\nvalid_test_ids = test_df[\"id\"].unique().tolist()\n\nprint(\"\\n📊 DATASET OVERVIEW\")\nstats = [\n    [\"Training Samples\", len(train_df)],\n    [\"Test Images\", len(valid_test_ids)],\n    [\"Total Test Rows\", len(test_df)],\n    [\"Sampling Frequency (test)\", f\"{test_df['fs'].iloc[0]} Hz\"],\n]\nprint(tabulate(stats, headers=[\"Metric\", \"Value\"], tablefmt=\"fancy_grid\"))\n\n\n# ============================================================\n# TTA OPTIONS (Stage2 internal TTA)\n# ============================================================\nUSE_TTA = False\nTTA_USE_FLIP = False\nTTA_USE_BC = False\nTTA_USE_GAMMA = False\n\n\n# ============================================================\n# MODEL IMPORTS\n# ============================================================\nfrom stage0_model import Net as Stage0Net\nfrom stage0_common import (\n    load_net as s0_load,\n    image_to_batch,\n    output_to_predict as s0_output,\n    normalise_by_homography,\n)\nfrom stage1_model import Net as Stage1Net\n\n# stage1_common is imported twice:\n#   - \"as s1c\" for sparse-capture hook\n#   - specific symbols for minimal changes elsewhere\nimport stage1_common as s1c\nfrom stage1_common import (\n    load_net as s1_load,\n    output_to_predict as s1_output,  # kept for compatibility (not used directly)\n    rectify_image,\n)\n\nimport stage2_model as s2m\nfrom stage2_common import filter_series_by_limits  # pixel_to_series は使わない\n\n\n# ============================================================\n# [NEW] stage1 output_to_predict wrapper (capture sparse gridpoint_xy)\n# ============================================================\ndef output_to_predict_with_sparse(image, batch, output):\n    \"\"\"Capture sparse gridpoint_xy by temporarily monkey-patching stage1_common.interpolate_mapping.\"\"\"\n    captured = {}\n    orig_interp_attr = s1c.interpolate_mapping\n    orig_interp_global = s1c.output_to_predict.__globals__.get(\"interpolate_mapping\", orig_interp_attr)\n\n    def _interp_capture(gridpoint_xy):\n        try:\n            captured[\"gridpoint_xy_sparse\"] = np.asarray(gridpoint_xy).copy()\n        except Exception:\n            captured[\"gridpoint_xy_sparse\"] = None\n        return orig_interp_attr(gridpoint_xy)\n\n    with _S1_SPARSE_LOCK:\n        s1c.interpolate_mapping = _interp_capture\n        s1c.output_to_predict.__globals__[\"interpolate_mapping\"] = _interp_capture\n        try:\n            gridpoint_xy, more = s1c.output_to_predict(image, batch, output)\n        finally:\n            s1c.interpolate_mapping = orig_interp_attr\n            s1c.output_to_predict.__globals__[\"interpolate_mapping\"] = orig_interp_global\n\n    gridpoint_xy_sparse = captured.get(\"gridpoint_xy_sparse\", None)\n    return gridpoint_xy, more, gridpoint_xy_sparse\n\n\n\ndef stage1_sparse_metrics(stage1_out: dict, grid_xy_sparse: np.ndarray | None) -> dict:\n    \"\"\"Compute hole_rate / mono_x / mono_y using sparse grid; n_points from connected components.\"\"\"\n    metrics = dict(hole_rate=np.nan, mono_x=np.nan, mono_y=np.nan, n_points=np.nan)\n    metrics[\"sparse_captured\"] = grid_xy_sparse is not None\n\n    # n_points\n    try:\n        gp = stage1_out[\"gridpoint\"][0, 0].detach().float().cpu().numpy()\n        cc = cc3d.connected_components(gp > float(STAGE1_GATE_POINT_THR))\n        stats = cc3d.statistics(cc)\n        n_points = int(len(stats.get(\"centroids\", [])) - 1)  # exclude background\n        metrics[\"n_points\"] = n_points\n    except Exception:\n        pass\n\n    if grid_xy_sparse is None:\n        return metrics\n\n    try:\n        x = grid_xy_sparse[..., 0]\n        y = grid_xy_sparse[..., 1]\n        hole = (x == 0) & (y == 0)\n        valid = ~hole\n        metrics[\"hole_rate\"] = float(hole.mean())\n\n        mono_x = 0.0\n        mono_y = 0.0\n        if valid.sum() > 10:\n            dx = np.diff(x, axis=1)\n            dy = np.diff(y, axis=0)\n\n            vdx = valid[:, 1:] & valid[:, :-1]\n            vdy = valid[1:, :] & valid[:-1, :]\n\n            mono_x = float((dx[vdx] > 0).mean()) if vdx.any() else 0.0\n            mono_y = float((dy[vdy] > 0).mean()) if vdy.any() else 0.0\n\n        metrics[\"mono_x\"] = float(mono_x)\n        metrics[\"mono_y\"] = float(mono_y)\n    except Exception:\n        pass\n\n    return metrics\n\n\n\ndef stage1_gate(metrics: dict) -> tuple[bool, list[str]]:\n    \"\"\"Gate: hole_rate < 0.10 AND mono_x > 0.95 AND mono_y > 0.95 AND 2000 < n_points < 3000\"\"\"\n    if not STAGE1_GATE_ENABLE:\n        return True, []\n\n    # If sparse mapping could not be captured, skip gate to avoid false negatives.\n    if not metrics.get(\"sparse_captured\", True):\n        return True, []\n\n    reasons = []\n    ok = True\n\n    hole_rate = metrics.get(\"hole_rate\", np.nan)\n    mono_x = metrics.get(\"mono_x\", np.nan)\n    mono_y = metrics.get(\"mono_y\", np.nan)\n    n_points = metrics.get(\"n_points\", np.nan)\n\n    if not np.isfinite(hole_rate) or hole_rate >= float(STAGE1_GATE_HOLE_MAX):\n        ok = False\n        reasons.append(f\"hole_rate={hole_rate:.4f} (max<{STAGE1_GATE_HOLE_MAX})\")\n\n    if not np.isfinite(mono_x) or mono_x <= float(STAGE1_GATE_MONO_MIN):\n        ok = False\n        reasons.append(f\"mono_x={mono_x:.4f} (min>{STAGE1_GATE_MONO_MIN})\")\n\n    if not np.isfinite(mono_y) or mono_y <= float(STAGE1_GATE_MONO_MIN):\n        ok = False\n        reasons.append(f\"mono_y={mono_y:.4f} (min>{STAGE1_GATE_MONO_MIN})\")\n\n    if not np.isfinite(n_points) or (n_points <= int(STAGE1_GATE_NPOINTS_MIN)) or (n_points >= int(STAGE1_GATE_NPOINTS_MAX)):\n        ok = False\n        reasons.append(f\"n_points={n_points} (need {STAGE1_GATE_NPOINTS_MIN}<n<{STAGE1_GATE_NPOINTS_MAX})\")\n\n    return ok, reasons\n\n\n# ============================================================\n# Stage2 ResNet Net (Inference)\n# ============================================================\nclass Stage2ResNet(nn.Module):\n    def __init__(\n        self,\n        pretrained: bool = False,\n        arch: str = \"resnet34.a3_in1k\",\n        decoder_dim=(128, 64, 32, 16),\n        use_imagenet_norm: bool = False,\n    ):\n        super().__init__()\n        self.register_buffer(\"D\", torch.tensor(0))\n        self.use_imagenet_norm = bool(use_imagenet_norm)\n        if self.use_imagenet_norm:\n            self.register_buffer(\"mean\", torch.tensor([0.485, 0.456, 0.406]).reshape(1, 3, 1, 1))\n            self.register_buffer(\"std\",  torch.tensor([0.229, 0.224, 0.225]).reshape(1, 3, 1, 1))\n\n        encoder_dim = [64, 128, 256, 512]  # resnet34\n\n        self.encoder = timm.create_model(\n            model_name=arch,\n            pretrained=pretrained,\n            in_chans=3,\n            num_classes=0,\n            global_pool=\"\",\n        )\n\n        self.decoder = s2m.MyCoordUnetDecoder(\n            in_channel=encoder_dim[-1],\n            skip_channel=encoder_dim[:-1][::-1] + [0],\n            out_channel=list(decoder_dim),\n            scale=[2, 2, 2, 2],\n        )\n\n        self.pixel = nn.Conv2d(decoder_dim[-1], 4, 1)\n\n    def forward(self, batch):\n        device = self.D.device\n        image = batch[\"image\"].to(device)\n        x = image.float() / 255.0\n        if self.use_imagenet_norm:\n            x = (x - self.mean) / self.std\n\n        encode = s2m.encode_with_resnet(self.encoder, x)\n        last, _ = self.decoder(feature=encode[-1], skip=encode[:-1][::-1] + [None])\n\n        logits = self.pixel(last)\n        pixel = torch.sigmoid(logits)\n        return {\"pixel\": pixel}\n\n\n\ndef load_stage2_ckpt_and_cfg(model: nn.Module, ckpt_path: str, strict: bool = True):\n    \"\"\"Load checkpoint and return cfg dict if present.\"\"\"\n    if not os.path.exists(ckpt_path):\n        raise FileNotFoundError(f\"Stage2 checkpoint not found: {ckpt_path}\")\n\n    ckpt = torch.load(ckpt_path, map_location=\"cpu\")\n\n    ckpt_cfg = {}\n    if isinstance(ckpt, dict) and \"cfg\" in ckpt and isinstance(ckpt[\"cfg\"], dict):\n        ckpt_cfg = ckpt[\"cfg\"]\n\n    state_dict = ckpt[\"state_dict\"] if isinstance(ckpt, dict) and \"state_dict\" in ckpt else ckpt\n\n    if isinstance(state_dict, dict) and any(k.startswith(\"module.\") for k in state_dict.keys()):\n        state_dict = {k[len(\"module.\"):]: v for k, v in state_dict.items()}\n\n    try:\n        msg = model.load_state_dict(state_dict, strict=strict)\n        if not strict:\n            print(f\"[Stage2ResNet] missing_keys={msg.missing_keys}, unexpected_keys={msg.unexpected_keys}\")\n    except RuntimeError as e:\n        print(f\"[Stage2ResNet] strict load failed: {e}\")\n        print(\"[Stage2ResNet] retry with strict=False ...\")\n        msg = model.load_state_dict(state_dict, strict=False)\n        print(f\"[Stage2ResNet] missing_keys={msg.missing_keys}, unexpected_keys={msg.unexpected_keys}\")\n\n    print(f\"[Stage2ResNet] loaded checkpoint from {ckpt_path}\")\n    return ckpt_cfg\n\n\n\ndef pad_to_multiple_2d(x: torch.Tensor, multiple: int, value: float):\n    B, C, H, W = x.shape\n    pad_h = (multiple - (H % multiple)) % multiple\n    pad_w = (multiple - (W % multiple)) % multiple\n    if pad_h == 0 and pad_w == 0:\n        return x, (0, 0)\n    x = F.pad(x, (0, pad_w, 0, pad_h), mode=\"constant\", value=value)\n    return x, (pad_h, pad_w)\n\n\n# ============================================================\n# LOAD MODELS (MINIMAL CHANGE: load 2nd copy if 2GPU)\n# ============================================================\nprint(\"\\n\" + \"=\" * 80)\nprint(\"LOADING MODELS\")\nprint(\"=\" * 80)\n\nprint(\"\\n🔧 Loading Stage 0...\")\n\nstage0_net = Stage0Net(pretrained=False)\nstage0_net = s0_load(stage0_net, f\"{WEIGHT_DIR}/stage0-last.checkpoint.pth\")\nstage0_net.to(DEVICE0).eval()\n\nstage0_net_1 = None\nif USE_TWO_GPU_BRANCH:\n    stage0_net_1 = Stage0Net(pretrained=False)\n    stage0_net_1 = s0_load(stage0_net_1, f\"{WEIGHT_DIR}/stage0-last.checkpoint.pth\")\n    stage0_net_1.to(DEVICE1).eval()\n\nprint(\"✅ Stage 0 loaded\")\n\nprint(\"\\n🔧 Loading Stage 1...\")\n\nstage1_net = Stage1Net(pretrained=False)\nstage1_net = s1_load(stage1_net, f\"{WEIGHT_DIR}/stage1-last.checkpoint.pth\")\nstage1_net.to(DEVICE0).eval()\n\nstage1_net_1 = None\nif USE_TWO_GPU_BRANCH:\n    stage1_net_1 = Stage1Net(pretrained=False)\n    stage1_net_1 = s1_load(stage1_net_1, f\"{WEIGHT_DIR}/stage1-last.checkpoint.pth\")\n    stage1_net_1.to(DEVICE1).eval()\n\nprint(\"✅ Stage 1 loaded\")\n\n\n# ============================================================\n# Stage2 ENSEMBLE LOAD\n# ============================================================\n@dataclass\nclass Stage2Wrapper:\n    ckpt_path: str\n    model: nn.Module\n    arch: str\n    use_imagenet_norm: bool\n    pad_multiple: int\n    x_scale: int\n    t0: int\n    t1: int\n    zero_mv: list\n    mv_to_pixel: float\n\n\nprint(\"\\n🔧 Loading Stage 2 ENSEMBLE (multiple ckpts)...\")\n\n\ndef _read_stage2_cfg_from_ckpt(ckpt_path: str) -> dict:\n    tmp = torch.load(ckpt_path, map_location=\"cpu\")\n    if isinstance(tmp, dict) and \"cfg\" in tmp and isinstance(tmp[\"cfg\"], dict):\n        return tmp[\"cfg\"]\n    return {}\n\n\n\ndef _load_stage2_ensemble_to(device_str: str) -> list[Stage2Wrapper]:\n    _set_cuda_device(device_str)\n    wrappers: list[Stage2Wrapper] = []\n    for ckpt_path in STAGE2_CKPTS:\n        if not os.path.exists(ckpt_path):\n            raise FileNotFoundError(f\"Stage2 checkpoint not found: {ckpt_path}\")\n\n        cfg = _read_stage2_cfg_from_ckpt(ckpt_path)\n\n        arch = str(cfg.get(\"ARCH\", ST2_ARCH))\n        use_imagenet_norm = bool(cfg.get(\"USE_IMAGENET_NORM\", ST2_USE_IMAGENET_NORM))\n        pad_multiple = int(cfg.get(\"PAD_MULTIPLE\", ST2_PAD_MULTIPLE))\n        x_scale = int(cfg.get(\"X_SCALE\", ST2_X_SCALE))\n        t0 = int(cfg.get(\"T0\", ST2_T0))\n        t1 = int(cfg.get(\"T1\", ST2_T1))\n        zero_mv = cfg.get(\"ZERO_MV\", ST2_ZERO_MV)\n        mv_to_pixel = float(cfg.get(\"MV_TO_PIXEL\", ST2_MV_TO_PIXEL))\n\n        m = Stage2ResNet(\n            pretrained=False,\n            arch=arch,\n            use_imagenet_norm=use_imagenet_norm,\n        )\n        _ = load_stage2_ckpt_and_cfg(m, ckpt_path, strict=True)\n        m.to(device_str).eval()\n\n        wrappers.append(\n            Stage2Wrapper(\n                ckpt_path=ckpt_path,\n                model=m,\n                arch=arch,\n                use_imagenet_norm=use_imagenet_norm,\n                pad_multiple=pad_multiple,\n                x_scale=x_scale,\n                t0=t0,\n                t1=t1,\n                zero_mv=list(zero_mv),\n                mv_to_pixel=mv_to_pixel,\n            )\n        )\n    return wrappers\n\n\nstage2_wrappers = _load_stage2_ensemble_to(DEVICE0)\nstage2_wrappers_1 = _load_stage2_ensemble_to(DEVICE1) if USE_TWO_GPU_BRANCH else None\n\nprint(\"✅ Stage 2 ensemble loaded:\", len(stage2_wrappers))\nfor i, w in enumerate(stage2_wrappers):\n    print(f\"  [{i}] x_scale={w.x_scale} t0={w.t0} t1={w.t1} pad={w.pad_multiple} arch={w.arch} norm={w.use_imagenet_norm}\")\n    print(f\"      zero_mv={w.zero_mv} mv_to_pixel={w.mv_to_pixel} ckpt={w.ckpt_path}\")\n\n\n# ============================================================\n# OFFICIAL-LIKE SCORING\n# ============================================================\nfrom typing import Tuple, Dict\nimport scipy.optimize\nimport scipy.signal\n\nLEADS = [\"I\", \"II\", \"III\", \"aVR\", \"aVL\", \"aVF\", \"V1\", \"V2\", \"V3\", \"V4\", \"V5\", \"V6\"]\nMAX_TIME_SHIFT = 0.2\nPERFECT_SCORE = 1e6\n\n\nclass ParticipantVisibleError(Exception):\n    pass\n\n\n\ndef compute_power(label: np.ndarray, prediction: np.ndarray) -> Tuple[float, float]:\n    if label.ndim != 1 or prediction.ndim != 1:\n        raise ParticipantVisibleError(\"Inputs must be 1-dimensional arrays.\")\n    finite_mask = np.isfinite(prediction)\n    if not np.any(finite_mask):\n        raise ParticipantVisibleError(\"The 'prediction' array contains no finite values (all NaN or inf).\")\n    prediction = prediction.copy()\n    prediction[~np.isfinite(prediction)] = 0\n    noise = label - prediction\n    p_signal = np.sum(label**2)\n    p_noise = np.sum(noise**2)\n    return p_signal, p_noise\n\n\n\ndef compute_snr(signal: float, noise: float) -> float:\n    if noise == 0:\n        snr = PERFECT_SCORE\n    elif signal == 0:\n        snr = 0\n    else:\n        snr = min((signal / noise), PERFECT_SCORE)\n    return snr\n\n\n\ndef align_signals(label: np.ndarray, pred: np.ndarray, max_shift: float = float(\"inf\")) -> np.ndarray:\n    if np.any(~np.isfinite(label)):\n        raise ParticipantVisibleError(\"values in label should all be finite\")\n    if np.sum(np.isfinite(pred)) == 0:\n        raise ParticipantVisibleError(\"prediction can not all be infinite\")\n\n    label_arr = np.asarray(label, dtype=np.float64)\n    pred_arr = np.asarray(pred, dtype=np.float64)\n\n    label_mean = np.mean(label_arr)\n    pred_mean = np.mean(pred_arr)\n\n    label_arr_centered = label_arr - label_mean\n    pred_arr_centered = pred_arr - pred_mean\n\n    correlation = scipy.signal.correlate(label_arr_centered, pred_arr_centered, mode=\"full\")\n    n_label = np.size(label_arr)\n    n_pred = np.size(pred_arr)\n\n    lags = scipy.signal.correlation_lags(n_label, n_pred, mode=\"full\")\n    valid_lags_mask = (lags >= -max_shift) & (lags <= max_shift)\n\n    max_correlation = np.nanmax(correlation[valid_lags_mask])\n    all_max_indices = np.flatnonzero(correlation == max_correlation)\n    best_idx = min(all_max_indices, key=lambda i: abs(lags[i]))\n    time_shift = lags[best_idx]\n\n    start_padding_len = max(time_shift, 0)\n    pred_slice_start = max(-time_shift, 0)\n    pred_slice_end = min(n_label - time_shift, n_pred)\n    end_padding_len = max(n_label - n_pred - time_shift, 0)\n\n    aligned_pred = np.concatenate(\n        (\n            np.full(start_padding_len, np.nan),\n            pred_arr[pred_slice_start:pred_slice_end],\n            np.full(end_padding_len, np.nan),\n        )\n    )\n\n    def objective_func(v_shift):\n        return np.nansum((label_arr - (aligned_pred - v_shift)) ** 2)\n\n    if np.any(np.isfinite(label_arr) & np.isfinite(aligned_pred)):\n        results = scipy.optimize.minimize_scalar(objective_func, method=\"Brent\")\n        vertical_shift = results.x\n        aligned_pred -= vertical_shift\n\n    return aligned_pred\n\n\n\ndef _calculate_image_score(group: pd.DataFrame) -> float:\n    unique_fs_values = group[\"fs\"].unique()\n    if len(unique_fs_values) != 1:\n        raise ParticipantVisibleError(\"Sampling frequency should be consistent across each ecg\")\n    sampling_frequency = unique_fs_values[0]\n    if sampling_frequency != int(len(group[group[\"lead\"] == \"II\"]) / 10):\n        raise ParticipantVisibleError(\"The sequence_length should be sampling frequency * 10s\")\n\n    sum_signal = 0.0\n    sum_noise = 0.0\n\n    for lead in LEADS:\n        sub = group[group[\"lead\"] == lead]\n        label = sub[\"value_true\"].values\n        pred = sub[\"value_pred\"].values\n\n        aligned_pred = align_signals(label, pred, int(sampling_frequency * MAX_TIME_SHIFT))\n        p_signal, p_noise = compute_power(label, aligned_pred)\n        sum_signal += p_signal\n        sum_noise += p_noise\n\n    return compute_snr(sum_signal, sum_noise)\n\n\n\ndef score_df(solution: pd.DataFrame, submission: pd.DataFrame, row_id_column_name: str = \"id\") -> float:\n    for df in [solution, submission]:\n        if row_id_column_name not in df.columns:\n            raise ParticipantVisibleError(f\"'{row_id_column_name}' column not found in DataFrame.\")\n        if df[\"value\"].isna().any():\n            raise ParticipantVisibleError(\"NaN exists in solution/submission\")\n        if not np.isfinite(df[\"value\"]).all():\n            raise ParticipantVisibleError(\"Infinity exists in solution/submission\")\n\n    submission = submission[[\"id\", \"value\"]]\n    merged_df = pd.merge(solution, submission, on=row_id_column_name, suffixes=(\"_true\", \"_pred\"))\n    merged_df[\"image_id\"] = merged_df[row_id_column_name].str.split(\"_\").str[0]\n    merged_df[\"row_id\"] = merged_df[row_id_column_name].str.split(\"_\").str[1].astype(\"int64\")\n    merged_df[\"lead\"] = merged_df[row_id_column_name].str.split(\"_\").str[2]\n    merged_df.sort_values(by=[\"image_id\", \"row_id\", \"lead\"], inplace=True)\n\n    image_scores = merged_df.groupby(\"image_id\").apply(_calculate_image_score, include_groups=False)\n    return max(float(10 * np.log10(image_scores.mean())), -PERFECT_SCORE)\n\n\n# ============================================================\n# VALIDATION HELPERS\n# ============================================================\nTYPE_IDS = [\"0001\", \"0003\", \"0004\", \"0005\", \"0006\", \"0009\", \"0010\", \"0011\", \"0012\"]\nLEAD_NAMES_BY_ROW = [\n    [\"I\", \"aVR\", \"V1\", \"V4\"],\n    [\"II_short\", \"aVL\", \"V2\", \"V5\"],\n    [\"III\", \"aVF\", \"V3\", \"V6\"],\n]\n\n\ndef interpolate_signal_to_length(signal: np.ndarray, target_length: int) -> np.ndarray:\n    if len(signal) == target_length:\n        return signal\n    x_old = np.linspace(0.0, 1.0, len(signal), endpoint=False)\n    x_new = np.linspace(0.0, 1.0, target_length, endpoint=False)\n    return np.interp(x_new, x_old, signal)\n\n\n\ndef extract_leads_from_series(series: np.ndarray) -> dict:\n    L = series.shape[1]\n    pred = {}\n    for row_idx in range(3):\n        split = np.array_split(series[row_idx], 4)\n        for name, sig in zip(LEAD_NAMES_BY_ROW[row_idx], split):\n            pred[name] = sig\n    pred[\"II\"] = series[3]\n    return pred\n\n\n\ndef test_from_train_df(data_dir=KAGGLE_DIR) -> pd.DataFrame:\n    valid_df = pd.read_csv(f\"{data_dir}/train.csv\")\n    valid_df[\"id\"] = valid_df[\"id\"].astype(str)\n    fake_test_df = []\n    for _, d in valid_df.iterrows():\n        image_id = d[\"id\"]\n        truth_df = pd.read_csv(f\"{data_dir}/train/{image_id}/{image_id}.csv\")\n        non_nan_count = truth_df.count()\n        this_df = pd.DataFrame(\n            {\n                \"id\": image_id,\n                \"lead\": non_nan_count.index,\n                \"fs\": d[\"fs\"],\n                \"number_of_rows\": non_nan_count.values,\n            }\n        )\n        fake_test_df.append(this_df)\n    return pd.concat(fake_test_df)\n\n\n\ndef validation_subset(image_ids, type_ids, data_dir=KAGGLE_DIR):\n    train_df2 = pd.read_csv(f\"{data_dir}/train.csv\")\n    train_df2[\"id\"] = train_df2[\"id\"].astype(str)\n    valid_image_ids = set(train_df2[\"id\"].values)\n\n    meta_df = test_from_train_df(data_dir)\n    filtered_image_ids = [img_id for img_id in image_ids if img_id in valid_image_ids]\n\n    for image_id in filtered_image_ids:\n        lead_records = meta_df[meta_df[\"id\"] == image_id].copy()\n        for type_id in type_ids:\n            filename = f\"{data_dir}/train/{image_id}/{image_id}-{type_id}.png\"\n            yield filename, lead_records\n\n\n\ndef load_ground_truth(image_id: str, data_dir=KAGGLE_DIR) -> pd.DataFrame:\n    return pd.read_csv(f\"{data_dir}/train/{image_id}/{image_id}.csv\")\n\n\n\ndef build_submission_rows(predicted_leads: dict, image_id: str, lead_records: pd.DataFrame) -> list:\n    rows = []\n    for _, record in lead_records.iterrows():\n        lead_name = record[\"lead\"]\n        n = int(record[\"number_of_rows\"])\n        if lead_name not in predicted_leads:\n            continue\n        pred_sig = interpolate_signal_to_length(np.asarray(predicted_leads[lead_name]), n).astype(np.float32)\n        pred_sig = savgol_final_1d(pred_sig)\n        for t in range(n):\n            rows.append({\"id\": f\"{image_id}_{t}_{lead_name}\", \"value\": float(pred_sig[t])})\n    return rows\n\n\n\ndef build_solution_rows(ground_truth: pd.DataFrame, image_id: str, lead_records: pd.DataFrame) -> list:\n    rows = []\n    for _, record in lead_records.iterrows():\n        lead_name = record[\"lead\"]\n        n = int(record[\"number_of_rows\"])\n        fs = record[\"fs\"]\n        if lead_name not in ground_truth.columns:\n            continue\n        gt = ground_truth[lead_name].dropna().values[:n]\n        for t in range(len(gt)):\n            rows.append({\"id\": f\"{image_id}_{t}_{lead_name}\", \"fs\": fs, \"value\": float(gt[t])})\n    return rows\n\n\n\ndef make_validation_detail_and_powers(\n    solution_df: pd.DataFrame,\n    submission_df: pd.DataFrame,\n    image_id: str,\n    type_id: str,\n) -> Tuple[pd.DataFrame, Dict[str, float], float, float]:\n    merged = pd.merge(\n        solution_df[[\"id\", \"fs\", \"value\"]],\n        submission_df[[\"id\", \"value\"]],\n        on=\"id\",\n        suffixes=(\"_true\", \"_pred\"),\n    )\n\n    parts = merged[\"id\"].str.split(\"_\", expand=True)\n    merged[\"image_id\"] = parts[0]\n    merged[\"t\"] = parts[1].astype(\"int32\")\n    merged[\"lead\"] = parts[2]\n    merged[\"type_id\"] = type_id\n\n    merged[\"GrandTrue\"] = merged[\"value_true\"]\n    merged[\"Pred\"] = merged[\"value_pred\"]\n    merged[\"Diff\"] = merged[\"Pred\"] - merged[\"GrandTrue\"]\n\n    fs = float(merged[\"fs\"].iloc[0])\n    merged[\"PredAligned\"] = np.nan\n    merged[\"DiffAligned\"] = np.nan\n\n    per_lead_snr_db: Dict[str, float] = {}\n    sum_signal = 0.0\n    sum_noise = 0.0\n\n    for lead in LEADS:\n        m = merged[\"lead\"] == lead\n        if not m.any():\n            continue\n\n        label = merged.loc[m, \"GrandTrue\"].to_numpy()\n        pred = merged.loc[m, \"Pred\"].to_numpy()\n\n        aligned = align_signals(label, pred, int(fs * MAX_TIME_SHIFT))\n        merged.loc[m, \"PredAligned\"] = aligned\n        merged.loc[m, \"DiffAligned\"] = aligned - label\n\n        p_signal, p_noise = compute_power(label, aligned)\n        sum_signal += p_signal\n        sum_noise += p_noise\n\n        per_lead_snr_db[lead] = 10.0 * np.log10(compute_snr(p_signal, p_noise))\n\n    merged.sort_values([\"image_id\", \"type_id\", \"lead\", \"t\"], inplace=True)\n\n    detail_df = merged[\n        [\n            \"image_id\",\n            \"type_id\",\n            \"lead\",\n            \"t\",\n            \"fs\",\n            \"GrandTrue\",\n            \"Pred\",\n            \"Diff\",\n            \"PredAligned\",\n            \"DiffAligned\",\n            \"id\",\n        ]\n    ].copy()\n\n    return detail_df, per_lead_snr_db, sum_signal, sum_noise\n\n\n# ============================================================\n# CORE PIPELINE  (MINIMAL: add device_str + net selection + gate)\n# ============================================================\ndef _select_nets_and_wrappers(device_str: str):\n    if device_str == DEVICE1 and USE_TWO_GPU_BRANCH:\n        return stage0_net_1, stage1_net_1, stage2_wrappers_1\n    return stage0_net, stage1_net, stage2_wrappers\n\n\n\ndef run_stage0(image_rgb: np.ndarray, apply_change_color: bool = True, device_str: str = DEVICE0):\n    \"\"\"Stage0: model input optionally change_color; geometry uses original image.\"\"\"\n    _set_cuda_device(device_str)\n    s0_net, _, _ = _select_nets_and_wrappers(device_str)\n\n    image_for_model = change_color(image_rgb) if apply_change_color else image_rgb\n    batch = image_to_batch(image_for_model)\n    batch = _to_device_batch(batch, device_str)\n\n    with amp_autocast(device_str, dtype=FLOAT_TYPE, enabled=False):\n        with torch.no_grad():\n            output = s0_net(batch)\n            rotated, keypoint = s0_output(image_rgb, batch, output)  # <- always original\n            normalised, keypoint, homo = normalise_by_homography(rotated, keypoint)\n    return normalised, homo\n\n\n\ndef run_stage1(normalised_rgb: np.ndarray, device_str: str = DEVICE0):\n    \"\"\"Stage1 with sparse gate metrics.\"\"\"\n    _set_cuda_device(device_str)\n    _, s1_net, _ = _select_nets_and_wrappers(device_str)\n\n    batch = {\n        \"image\": torch.from_numpy(np.ascontiguousarray(normalised_rgb.transpose(2, 0, 1)))\n        .unsqueeze(0)\n        .to(device_str)\n    }\n    with amp_autocast(device_str, dtype=FLOAT_TYPE, enabled=False):\n        with torch.no_grad():\n            stage1_out = s1_net(batch)\n\n            gridpoint_xy, more, gridpoint_xy_sparse = output_to_predict_with_sparse(normalised_rgb, batch, stage1_out)\n\n            if USE_STAGE1_POSTPROC:\n                Himg, Wimg = normalised_rgb.shape[:2]\n                gridpoint_xy2, rms, applied = stage1_affine_residual_smooth_safe(\n                    gridpoint_xy,\n                    image_hw=(Himg, Wimg),\n                    sigma=STAGE1_SIGMA,\n                    max_delta=STAGE1_MAX_DELTA,\n                    rms_thresh=STAGE1_RMS_THRESH,\n                )\n                if STAGE1_DEBUG:\n                    print(f\"[stage1_post] holes={(gridpoint_xy[...,0]<=0).sum()} rms={rms:.3f} applied={applied}\")\n                gridpoint_xy = gridpoint_xy2\n\n            rectified = rectify_image(normalised_rgb, gridpoint_xy)\n\n    metrics = stage1_sparse_metrics(stage1_out, gridpoint_xy_sparse)\n    ok, reasons = stage1_gate(metrics)\n\n    gate_dbg = dict(**metrics)\n    gate_dbg[\"gate_ok\"] = bool(ok)\n    gate_dbg[\"gate_reasons\"] = \"; \".join(reasons) if reasons else \"\"\n\n    return rectified, gridpoint_xy, gate_dbg\n\n\n\ndef blend_gridpoints_safe(g1: np.ndarray, g2: np.ndarray, image_hw: tuple[int, int], reducer: str = \"median\") -> np.ndarray:\n    \"\"\"Blend two interpolated gridpoints (pattern=2).\"\"\"\n    if g1 is None:\n        return g2\n    if g2 is None:\n        return g1\n    a = np.stack([g1.astype(np.float32), g2.astype(np.float32)], axis=0)\n    if reducer == \"mean\":\n        g = a.mean(axis=0)\n    else:\n        g = np.median(a, axis=0)\n    Himg, Wimg = image_hw\n    g[..., 0] = np.clip(g[..., 0], -2.0, (Wimg - 1) + 2.0)\n    g[..., 1] = np.clip(g[..., 1], -2.0, (Himg - 1) + 2.0)\n    return g.astype(np.float32)\n\n\n@torch.no_grad()\ndef _stage2_forward_pixelprob_model(crop_rgb: np.ndarray, w: Stage2Wrapper, device_str: str):\n    _set_cuda_device(device_str)\n    Hc, Wc = crop_rgb.shape[:2]\n    img = torch.from_numpy(np.ascontiguousarray(crop_rgb.transpose(2, 0, 1))).unsqueeze(0).to(device_str)\n    img, _ = pad_to_multiple_2d(img, w.pad_multiple, value=255)\n\n    out = w.model({\"image\": img})\n    pix = out[\"pixel\"].float()[:, :, :Hc, :Wc]\n    return pix[0].detach().cpu().numpy().astype(np.float32)\n\n\n\ndef _resize_pixel_to_width(pixel: np.ndarray, target_w: int) -> np.ndarray:\n    if pixel.shape[-1] == target_w:\n        return pixel\n    t = torch.from_numpy(pixel).unsqueeze(0)\n    t = F.interpolate(t, size=(pixel.shape[1], target_w), mode=\"bilinear\", align_corners=False)\n    return t.squeeze(0).numpy().astype(np.float32)\n\n\n@torch.no_grad()\ndef _stage2_forward_pixelprob_model_torch(crop_rgb: np.ndarray, w: Stage2Wrapper, device_str: str) -> torch.Tensor:\n    \"\"\"GPU-side forward. Returns (4,H,W) float32 tensor on device.\"\"\"\n    _set_cuda_device(device_str)\n    Hc, Wc = crop_rgb.shape[:2]\n    img = torch.from_numpy(np.ascontiguousarray(crop_rgb.transpose(2, 0, 1))).unsqueeze(0).to(device_str, non_blocking=True)\n    img, _ = pad_to_multiple_2d(img, w.pad_multiple, value=255)\n    with amp_autocast(device_str, dtype=FLOAT_TYPE, enabled=True):\n        out = w.model({\"image\": img})\n    pix = out[\"pixel\"].float()[:, :, :Hc, :Wc]\n    return pix[0]\n\n\n\ndef _resize_pixel_to_width_torch(pixel: torch.Tensor, target_w: int) -> torch.Tensor:\n    \"\"\"pixel: (4,H,W) tensor on device -> (4,H,target_w)\"\"\"\n    if int(pixel.shape[-1]) == int(target_w):\n        return pixel\n    t = pixel.unsqueeze(0)\n    t = F.interpolate(t, size=(int(pixel.shape[1]), int(target_w)), mode=\"bilinear\", align_corners=False)\n    return t.squeeze(0)\n\n\n\ndef _get_common_roi_base(wrappers: list[Stage2Wrapper]) -> tuple[int, int]:\n    t0s = [w.t0 for w in wrappers]\n    t1s = [w.t1 for w in wrappers]\n    t0 = max(t0s)\n    t1 = min(t1s)\n    if t1 <= t0:\n        t0 = int(np.median(t0s))\n        t1 = int(np.median(t1s))\n    if t1 <= t0:\n        t0 = min(t0s)\n        t1 = max(t1s)\n    return int(t0), int(t1)\n\n\n\ndef _robust_median_list(values: list):\n    a = np.asarray(values, dtype=np.float32)\n    return np.median(a, axis=0).tolist()\n\n\n\ndef run_stage2_single(\n    rectified_rgb: np.ndarray, length: int, w: Stage2Wrapper, roi_base=None, device_str: str = DEVICE0\n) -> np.ndarray:\n    _set_cuda_device(device_str)\n    x0, x1 = 0, 2176\n    y0, y1 = 0, 1696\n\n    crop = rectified_rgb[y0:y1, x0:x1]\n\n    if w.x_scale != 1:\n        crop = cv2.resize(crop, (crop.shape[1] * w.x_scale, crop.shape[0]), interpolation=cv2.INTER_LINEAR)\n\n    if roi_base is None:\n        t0s = int(w.t0 * w.x_scale)\n        t1s = int(w.t1 * w.x_scale)\n    else:\n        t0b, t1b = roi_base\n        t0s = int(t0b * w.x_scale)\n        t1s = int(t1b * w.x_scale)\n\n    with amp_autocast(device_str, dtype=FLOAT_TYPE, enabled=True):\n        if USE_TTA:\n            tta_pixels = []\n            tta_pixels.append(_stage2_forward_pixelprob_model(crop, w, device_str))\n\n            if TTA_USE_BC:\n                img_bc = cv2.convertScaleAbs(crop, alpha=1.1, beta=10)\n                tta_pixels.append(_stage2_forward_pixelprob_model(img_bc, w, device_str))\n\n            if TTA_USE_GAMMA:\n                gamma = 0.8\n                lut = np.array([((i / 255.0) ** gamma) * 255 for i in range(256)], dtype=np.uint8)\n                img_g = cv2.LUT(crop, lut)\n                tta_pixels.append(_stage2_forward_pixelprob_model(img_g, w, device_str))\n\n            if TTA_USE_FLIP:\n                img_f = cv2.flip(crop, 1)\n                pix_f = _stage2_forward_pixelprob_model(img_f, w, device_str)\n                pix_f = pix_f[..., ::-1]\n                tta_pixels.append(pix_f)\n\n            pixel = np.mean(tta_pixels, axis=0).astype(np.float32)\n        else:\n            pixel = _stage2_forward_pixelprob_model(crop, w, device_str)\n\n    pixel_roi = pixel[..., t0s:t1s].astype(np.float32)\n\n    series_in_pixel = pixel_to_series_subpixel_local(\n        pixel_roi,\n        zero_mv=w.zero_mv,\n        length=length,\n        power=SUBPIX_POWER,\n        threshold=SUBPIX_THRESH,\n        radius=SUBPIX_RADIUS,\n        eps=SUBPIX_LOGIT_EPS,\n    )\n    series = (np.array(w.zero_mv, dtype=np.float32).reshape(4, 1) - series_in_pixel) / float(w.mv_to_pixel)\n    return series.astype(np.float32)\n\n\n\ndef run_stage2_blend(rectified_rgb: np.ndarray, length: int, mode: str = \"pixel_logit_mean\", device_str: str = DEVICE0) -> np.ndarray:\n    _set_cuda_device(device_str)\n    _, _, wrappers = _select_nets_and_wrappers(device_str)\n    assert len(wrappers) >= 1\n\n    t0_base, t1_base = _get_common_roi_base(wrappers)\n    roi_base = (t0_base, t1_base)\n\n    if mode in [\"series_mean\", \"series_median\"]:\n        series_list = [run_stage2_single(rectified_rgb, length, w, roi_base=roi_base, device_str=device_str) for w in wrappers]\n        stack = np.stack(series_list, axis=0).astype(np.float32)\n        if mode == \"series_mean\":\n            return stack.mean(axis=0).astype(np.float32)\n        else:\n            return np.median(stack, axis=0).astype(np.float32)\n\n    if mode != \"pixel_logit_mean\":\n        raise ValueError(f\"Unknown Stage2 blend mode: {mode}\")\n\n    x0, x1 = 0, 2176\n    y0, y1 = 0, 1696\n    base_crop = rectified_rgb[y0:y1, x0:x1]\n\n    ref_scale = max([w.x_scale for w in wrappers])\n    ref_w = (x1 - x0) * ref_scale\n\n    t0_ref = int(t0_base * ref_scale)\n    t1_ref = int(t1_base * ref_scale)\n\n    zero_mv_ref = _robust_median_list([w.zero_mv for w in wrappers])\n    mv_to_pixel_ref = float(np.median(np.asarray([w.mv_to_pixel for w in wrappers], dtype=np.float32)))\n\n    # ============================================================\n    # GPU-complete pixel_logit_mean:\n    #   - per-model pixel stays on GPU\n    #   - resize/logit/sum/mean/sigmoid on GPU\n    #   - move to CPU only once after ROI is selected\n    # ============================================================\n    logit_sum_t: torch.Tensor | None = None\n    n_models = 0\n    eps = 1e-6\n\n    for w in wrappers:\n        crop = base_crop\n        if w.x_scale != 1:\n            crop = cv2.resize(crop, (crop.shape[1] * w.x_scale, crop.shape[0]), interpolation=cv2.INTER_LINEAR)\n\n        pix_t = _stage2_forward_pixelprob_model_torch(crop, w, device_str)  # (4,H,W) on GPU\n        pix_t = _resize_pixel_to_width_torch(pix_t, ref_w)                  # (4,H,ref_w) on GPU\n\n        p_t = pix_t.clamp(float(eps), 1.0 - float(eps))\n        logit_t = torch.log(p_t) - torch.log1p(-p_t)\n\n        if logit_sum_t is None:\n            logit_sum_t = logit_t\n        else:\n            logit_sum_t = logit_sum_t + logit_t\n        n_models += 1\n\n    if logit_sum_t is None or n_models <= 0:\n        raise RuntimeError(\"Stage2 pixel_logit_mean: no models accumulated\")\n\n    logit_mean_t = logit_sum_t / float(n_models)\n    pixel_blend_t = torch.sigmoid(logit_mean_t)\n\n    # single CPU transfer (ROI only)\n    pixel_roi = pixel_blend_t[..., t0_ref:t1_ref].detach().cpu().numpy().astype(np.float32)\n\n    series_in_pixel = pixel_to_series_subpixel_local(\n        pixel_roi,\n        zero_mv=zero_mv_ref,\n        length=length,\n        power=SUBPIX_POWER,\n        threshold=SUBPIX_THRESH,\n        radius=SUBPIX_RADIUS,\n        eps=SUBPIX_LOGIT_EPS,\n    )\n    series = (np.array(zero_mv_ref, dtype=np.float32).reshape(4, 1) - series_in_pixel) / float(mv_to_pixel_ref)\n    return series.astype(np.float32)\n\n\n# ============================================================\n# [FIX] Define infer_one_ecg_image (was missing)\n# ============================================================\ndef _blend_series(series_list: list[np.ndarray], reducer: str = \"mean\") -> np.ndarray:\n    stack = np.stack(series_list, axis=0).astype(np.float32)\n    if str(reducer).lower() == \"median\":\n        return np.median(stack, axis=0).astype(np.float32)\n    return stack.mean(axis=0).astype(np.float32)\n\n\n\ndef infer_one_ecg_image(image_rgb: np.ndarray, length_ii: int):\n    \"\"\"End-to-end inference for one ECG image.\n\n    Returns:\n      series(4,L), normalised_rgb, rectified_rgb, gate_dbg\n\n    Behavior:\n      - Stage0 baseline uses change_color ON for model input\n      - Optional Stage0 color TTA (ON/OFF) pattern 1 or 2\n      - If Stage1 gate fails, returns zero series (and still returns images for debug)\n    \"\"\"\n\n    L = int(length_ii)\n\n    def _branch_full(apply_cc: bool, device_str: str):\n        normalised, _ = run_stage0(image_rgb, apply_change_color=apply_cc, device_str=device_str)\n        rectified, grid, gate_dbg = run_stage1(normalised, device_str=device_str)\n        ok = bool(gate_dbg.get(\"gate_ok\", True))\n        if not ok:\n            return np.zeros((4, L), np.float32), normalised, rectified, grid, gate_dbg, False\n        series = run_stage2_blend(rectified, L, mode=STAGE2_BLEND_MODE, device_str=device_str)\n        return series, normalised, rectified, grid, gate_dbg, True\n\n    def _branch_stage1(apply_cc: bool, device_str: str):\n        normalised, _ = run_stage0(image_rgb, apply_change_color=apply_cc, device_str=device_str)\n        rectified, grid, gate_dbg = run_stage1(normalised, device_str=device_str)\n        ok = bool(gate_dbg.get(\"gate_ok\", True))\n        return normalised, rectified, grid, gate_dbg, ok\n\n    # No TTA\n    if not STAGE0_COLOR_TTA:\n        series, normalised, rectified, _, gate_dbg, ok = _branch_full(True, DEVICE0)\n        if not ok:\n            return np.zeros((4, L), np.float32), normalised, rectified, gate_dbg\n        return series, normalised, rectified, gate_dbg\n\n    pattern = int(STAGE0_COLOR_TTA_PATTERN)\n    reducer = str(STAGE0_COLOR_TTA_REDUCE).lower()\n\n    # Pattern 1: run to Stage2 twice and blend series\n    if pattern == 1:\n        if USE_TWO_GPU_BRANCH:\n            f_on = _executor.submit(_branch_full, True, DEVICE0)\n            f_off = _executor.submit(_branch_full, False, DEVICE1)\n            on = f_on.result()\n            off = f_off.result()\n        else:\n            on = _branch_full(True, DEVICE0)\n            off = _branch_full(False, DEVICE0)\n\n        series_list = []\n        # Choose baseline artifacts by default\n        series_on, norm_on, rect_on, _, gate_on, ok_on = on[0], on[1], on[2], on[3], on[4], on[5]\n        series_off, norm_off, rect_off, _, gate_off, ok_off = off[0], off[1], off[2], off[3], off[4], off[5]\n\n        normalised, rectified, gate_dbg = norm_on, rect_on, gate_on\n\n        if ok_on:\n            series_list.append(series_on)\n        if ok_off:\n            series_list.append(series_off)\n            if not ok_on:\n                normalised, rectified, gate_dbg = norm_off, rect_off, gate_off\n\n        if len(series_list) == 0:\n            return np.zeros((4, L), np.float32), normalised, rectified, gate_dbg\n\n        series = _blend_series(series_list, reducer=reducer)\n        gate_dbg = dict(gate_dbg)\n        gate_dbg[\"tta_color\"] = True\n        gate_dbg[\"tta_pattern\"] = 1\n        gate_dbg[\"tta_branches_ok\"] = int(len(series_list))\n        return series, normalised, rectified, gate_dbg\n\n    # Pattern 2: run to Stage1 twice, blend grid, Stage2 once\n    else:\n        if USE_TWO_GPU_BRANCH:\n            f_on = _executor.submit(_branch_stage1, True, DEVICE0)\n            f_off = _executor.submit(_branch_stage1, False, DEVICE1)\n            on = f_on.result()\n            off = f_off.result()\n        else:\n            on = _branch_stage1(True, DEVICE0)\n            off = _branch_stage1(False, DEVICE0)\n\n        norm_on, rect_on, grid_on, gate_on, ok_on = on\n        norm_off, rect_off, grid_off, gate_off, ok_off = off\n\n        normalised, rectified, gate_dbg = norm_on, rect_on, gate_on\n\n        if ok_on and ok_off:\n            Himg, Wimg = normalised.shape[:2]\n            grid_blend = blend_gridpoints_safe(grid_on, grid_off, (Himg, Wimg), reducer=reducer)\n            rectified = rectify_image(normalised, grid_blend)\n            series = run_stage2_blend(rectified, L, mode=STAGE2_BLEND_MODE, device_str=DEVICE0)\n            gate_dbg = dict(gate_dbg)\n            gate_dbg[\"tta_color\"] = True\n            gate_dbg[\"tta_pattern\"] = 2\n            gate_dbg[\"tta_branches_ok\"] = 2\n            return series, normalised, rectified, gate_dbg\n\n        if ok_on:\n            series = run_stage2_blend(rect_on, L, mode=STAGE2_BLEND_MODE, device_str=DEVICE0)\n            gate_dbg = dict(gate_on)\n            gate_dbg[\"tta_color\"] = True\n            gate_dbg[\"tta_pattern\"] = 2\n            gate_dbg[\"tta_branches_ok\"] = 1\n            return series, norm_on, rect_on, gate_dbg\n\n        if ok_off:\n            series = run_stage2_blend(rect_off, L, mode=STAGE2_BLEND_MODE, device_str=DEVICE0)\n            gate_dbg = dict(gate_off)\n            gate_dbg[\"tta_color\"] = True\n            gate_dbg[\"tta_pattern\"] = 2\n            gate_dbg[\"tta_branches_ok\"] = 1\n            return series, norm_off, rect_off, gate_dbg\n\n        # both failed\n        return np.zeros((4, L), np.float32), normalised, rectified, gate_dbg\n\n\n# ============================================================\n# INFERENCE (TEST)\n# ============================================================\nprint(\"\\n\" + \"=\" * 80)\nprint(\"INFERENCE ON TEST\")\nprint(\"=\" * 80)\n\nFAIL_ID = []\nGATE_LOG_TEST = []\nprint(\"\\n🔄 Stage 0/1/2 full pipeline over TEST...\")\n\nfor step, sample_id in enumerate(tqdm(valid_test_ids, desc=\"Test Inference\"), start=1):\n    try:\n        image = cv2.imread(f\"{KAGGLE_DIR}/test/{sample_id}.png\", cv2.IMREAD_COLOR_RGB)\n        d = test_df[(test_df[\"id\"] == sample_id) & (test_df[\"lead\"] == \"II\")].iloc[0]\n        length = int(d.number_of_rows)\n\n        series, normalised, rectified, gate_dbg = infer_one_ecg_image(image, length)\n\n        # log per image\n        gate_dbg = gate_dbg or {}\n        GATE_LOG_TEST.append(\n            dict(\n                id=str(sample_id),\n                gate_ok=bool(gate_dbg.get(\"gate_ok\", True)),\n                hole_rate=float(gate_dbg.get(\"hole_rate\", np.nan)),\n                mono_x=float(gate_dbg.get(\"mono_x\", np.nan)),\n                mono_y=float(gate_dbg.get(\"mono_y\", np.nan)),\n                n_points=float(gate_dbg.get(\"n_points\", np.nan)),\n                gate_reasons=str(gate_dbg.get(\"gate_reasons\", \"\")),\n            )\n        )\n\n        # print per-image to log\n        print(\n            f\"[Gate] id={sample_id} ok={gate_dbg.get('gate_ok', True)} \"\n            f\"hole={gate_dbg.get('hole_rate', np.nan):.4f} \"\n            f\"mono=({gate_dbg.get('mono_x', np.nan):.4f},{gate_dbg.get('mono_y', np.nan):.4f}) \"\n            f\"n_points={gate_dbg.get('n_points', np.nan)} \"\n            f\"{('reason=' + gate_dbg.get('gate_reasons','')) if not gate_dbg.get('gate_ok', True) else ''}\"\n        )\n\n        cv2.imwrite(f\"{OUT_DIR}/normalised/{sample_id}.norm.png\", cv2.cvtColor(normalised, cv2.COLOR_RGB2BGR))\n        cv2.imwrite(f\"{OUT_DIR}/rectified/{sample_id}.rect.png\", cv2.cvtColor(rectified, cv2.COLOR_RGB2BGR))\n        np.save(f\"{OUT_DIR}/digitalised/{sample_id}.series.npy\", series)\n\n    except Exception as e:\n        print(f\"\\n⚠️  {sample_id}: {e}\")\n        FAIL_ID.append(sample_id)\n\n        # ensure zero series saved so submission is zero-filled\n        try:\n            d = test_df[(test_df[\"id\"] == sample_id) & (test_df[\"lead\"] == \"II\")].iloc[0]\n            length = int(d.number_of_rows)\n            np.save(f\"{OUT_DIR}/digitalised/{sample_id}.series.npy\", np.zeros((4, length), np.float32))\n        except Exception:\n            pass\n        if CUDA_EMPTY_CACHE_ON_ERROR:\n            maybe_cuda_empty_cache(force=True)\n\n    maybe_cuda_empty_cache(step=step)\n\nprint(f\"✅ Test inference done: {len(valid_test_ids) - len(FAIL_ID)}/{len(valid_test_ids)} success\")\n\n\n# save test gate log\ntry:\n    gate_csv_path = os.path.join(OUT_DIR, \"stage1_gate_test.csv\")\n    pd.DataFrame(GATE_LOG_TEST).to_csv(gate_csv_path, index=False)\n    print(f\"📝 Saved test gate log => {gate_csv_path} (n={len(GATE_LOG_TEST)})\")\nexcept Exception as e:\n    print(f\"Failed to save gate log: {e}\")\n\n\n# ============================================================\n# SUBMISSION (TEST) + EINTHOVEN\n# ============================================================\nprint(\"\\n\" + \"=\" * 80)\nprint(\"CREATING SUBMISSION (TEST) + EINTHOVEN\")\nprint(\"=\" * 80)\n\nsubmission_data = []\ngb = test_df.groupby(\"id\")\n\nfor sample_id, df in tqdm(gb, desc=\"Building submission\"):\n    try:\n        series = np.load(f\"{OUT_DIR}/digitalised/{sample_id}.series.npy\")\n        _4_, L = series.shape\n\n        series_by_lead = {}\n        for l in range(3):\n            lead_names = [\n                [\"I\", \"aVR\", \"V1\", \"V4\"],\n                [\"II\", \"aVL\", \"V2\", \"V5\"],\n                [\"III\", \"aVF\", \"V3\", \"V6\"],\n            ][l]\n\n            idx = [int(round(1 * L / 4)), int(round(2 * L / 4)), int(round(3 * L / 4))]\n            split = np.split(series[l], idx)\n            for k, s in zip(lead_names, split):\n                series_by_lead[k] = s\n\n        long_ii = series[3].copy().astype(np.float32)\n        short_ii = series_by_lead.get(\"II\", long_ii[: min(len(long_ii), 10)]).astype(np.float32)\n        long_ii = blend_ii_head(long_ii, short_ii)  # [FIRST]\n        series_by_lead[\"II\"] = long_ii\n\n        apply_einthoven_switchable(series_by_lead)  # [SECOND]\n\n    except Exception:\n        series_by_lead = {}\n        for _, d in df.iterrows():\n            series_by_lead[d.lead] = np.zeros(int(d.number_of_rows), dtype=np.float32)\n\n    for _, d in df.iterrows():\n        s = series_by_lead[d.lead]\n        target_len = int(d.number_of_rows)\n\n        if len(s) != target_len:\n            x_old = np.linspace(0.0, 1.0, len(s), endpoint=False)\n            x_new = np.linspace(0.0, 1.0, target_len, endpoint=False)\n            s = np.interp(x_new, x_old, s)\n\n        s = s.astype(np.float32)\n        s = savgol_final_1d(s)\n\n        for t in range(target_len):\n            submission_data.append({\"id\": f\"{sample_id}_{t}_{d.lead}\", \"value\": float(s[t])})\n\nsubmission = pd.DataFrame(submission_data)\nsubmission.to_csv(\"liu_submission.csv\", index=False)\nprint(f\"\\n✅ Submission created: {len(submission):,} rows\")\n\nsubmission[\"lead\"] = submission[\"id\"].str.split(\"_\").str[2]\nprint(\"\\n📊 PER-LEAD STATISTICS (TEST submission)\")\nlead_stats = []\nfor lead in LEADS:\n    data = submission[submission[\"lead\"] == lead][\"value\"]\n    lead_stats.append(\n        [\n            lead,\n            len(data),\n            f\"{data.mean():.6f}\",\n            f\"{data.std():.6f}\",\n            f\"{data.min():.6f}\",\n            f\"{data.max():.6f}\",\n            (data != 0).sum(),\n        ]\n    )\n\nprint(tabulate(lead_stats, headers=[\"Lead\", \"Count\", \"Mean\", \"Std\", \"Min\", \"Max\", \"Non-Zero\"], tablefmt=\"fancy_grid\"))\nprint(f\"\\n✅ Failed IDs: {len(FAIL_ID)}\")\nif FAIL_ID:\n    print(f\"   {FAIL_ID}\")\n\n\n# ============================================================\n# VALIDATION (TRAIN)\n# ============================================================\nprint(\"\\n\" + \"=\" * 80)\nprint(\"VALIDATION ON TRAIN (official-like scoring)\")\nprint(\"=\" * 80)\n\nif os.getenv(\"KAGGLE_IS_COMPETITION_RERUN\"):\n    print(\"Skipping validation during competition rerun (KAGGLE_IS_COMPETITION_RERUN set).\")\nelse:\n    VALID_N_IMAGES = 0\n    VALID_TYPE_IDS = TYPE_IDS\n    VALID_PLOT_FIRST = 3\n\n    train_ids = train_df[\"id\"].unique().tolist()[900 : 900 + VALID_N_IMAGES]\n    type_snrs = {tid: [] for tid in VALID_TYPE_IDS}\n    lead_snrs = {lead: [] for lead in LEADS}\n    image_stats = []\n\n    VAL_OUT_DIR = os.path.join(OUT_DIR, \"validation_output\")\n    VAL_PLOT_DIR = os.path.join(VAL_OUT_DIR, \"plots\")\n    VAL_DETAIL_DIR = os.path.join(VAL_OUT_DIR, \"detail\")\n    VAL_DETAIL_ALL = os.path.join(VAL_OUT_DIR, \"validation_detail_all.csv\")\n\n    os.makedirs(VAL_OUT_DIR, exist_ok=True)\n    os.makedirs(VAL_PLOT_DIR, exist_ok=True)\n    os.makedirs(VAL_DETAIL_DIR, exist_ok=True)\n\n    if os.path.exists(VAL_DETAIL_ALL):\n        os.remove(VAL_DETAIL_ALL)\n\n    total_signal = 0.0\n    total_noise = 0.0\n\n    for n, (filename, lead_records) in enumerate(validation_subset(train_ids, VALID_TYPE_IDS, KAGGLE_DIR)):\n        if not os.path.exists(filename):\n            continue\n\n        image_name = os.path.basename(filename)\n        parts = image_name.replace(\".png\", \"\").split(\"-\")\n        image_id = parts[0]\n        type_id = parts[1] if len(parts) > 1 else \"NA\"\n\n        print(f\"\\n[{n+1}] {image_name}\")\n\n        try:\n            image = cv2.imread(filename, cv2.IMREAD_COLOR_RGB)\n            if image is None:\n                raise RuntimeError(\"cv2.imread failed\")\n\n            length_ii = int(lead_records[lead_records[\"lead\"] == \"II\"].iloc[0][\"number_of_rows\"])\n\n            series, normalised, rectified, gate_dbg = infer_one_ecg_image(image, length_ii)\n\n            # log per image\n            gate_dbg = gate_dbg or {}\n            print(\n                f\"[Gate] id={image_id}-{type_id} ok={gate_dbg.get('gate_ok', True)} \"\n                f\"hole={gate_dbg.get('hole_rate', np.nan):.4f} \"\n                f\"mono=({gate_dbg.get('mono_x', np.nan):.4f},{gate_dbg.get('mono_y', np.nan):.4f}) \"\n                f\"n_points={gate_dbg.get('n_points', np.nan)} \"\n                f\"{('reason=' + gate_dbg.get('gate_reasons','')) if not gate_dbg.get('gate_ok', True) else ''}\"\n            )\n\n            predicted_leads = extract_leads_from_series(series)\n\n            L = series.shape[1]\n            idx = [int(round(1 * L / 4)), int(round(2 * L / 4)), int(round(3 * L / 4))]\n            row1_split = np.split(series[1], idx)\n            short_ii = row1_split[0].astype(np.float32)\n\n            short_I = np.split(series[0], idx)[0].astype(np.float32)\n            short_III = np.split(series[2], idx)[0].astype(np.float32)\n\n            long_ii = predicted_leads[\"II\"].copy().astype(np.float32)\n            long_ii = blend_ii_head(long_ii, short_ii)  # [FIRST]\n\n            tmp_for_corr = {\"I\": short_I, \"II\": long_ii, \"III\": short_III}\n            apply_einthoven_switchable(tmp_for_corr)  # [SECOND]\n\n            predicted_leads[\"I\"] = tmp_for_corr[\"I\"]\n            predicted_leads[\"III\"] = tmp_for_corr[\"III\"]\n            predicted_leads[\"II\"] = tmp_for_corr[\"II\"]\n\n            gt = load_ground_truth(image_id, KAGGLE_DIR)\n\n            submission_rows = build_submission_rows(predicted_leads, image_id, lead_records)\n            solution_rows = build_solution_rows(gt, image_id, lead_records)\n            submission_df = pd.DataFrame(submission_rows)\n            solution_df = pd.DataFrame(solution_rows)\n\n            detail_df, per_lead, sum_sig, sum_noise = make_validation_detail_and_powers(\n                solution_df, submission_df, image_id, type_id\n            )\n\n            overall_snr_db = 10.0 * np.log10(compute_snr(sum_sig, sum_noise))\n            print(f\"Overall SNR: {overall_snr_db:.2f} dB  (Blend={STAGE2_BLEND_MODE})\")\n\n            if type_id in type_snrs:\n                type_snrs[type_id].append(overall_snr_db)\n            for ld, db in per_lead.items():\n                lead_snrs[ld].append(db)\n\n            total_signal += sum_sig\n            total_noise += sum_noise\n\n            detail_path = os.path.join(VAL_DETAIL_DIR, f\"{image_id}-{type_id}.csv\")\n            detail_df.to_csv(detail_path, index=False)\n\n            write_header = not os.path.exists(VAL_DETAIL_ALL)\n            detail_df.to_csv(VAL_DETAIL_ALL, mode=\"a\", header=write_header, index=False)\n\n            row = {\n                \"image_id\": image_id,\n                \"type_id\": type_id,\n                \"overall_db\": overall_snr_db,\n                \"gate_ok\": bool(gate_dbg.get(\"gate_ok\", True)),\n                \"hole_rate\": float(gate_dbg.get(\"hole_rate\", np.nan)),\n                \"mono_x\": float(gate_dbg.get(\"mono_x\", np.nan)),\n                \"mono_y\": float(gate_dbg.get(\"mono_y\", np.nan)),\n                \"n_points\": float(gate_dbg.get(\"n_points\", np.nan)),\n                \"gate_reasons\": str(gate_dbg.get(\"gate_reasons\", \"\")),\n            }\n            for lead_name in LEADS:\n                row[lead_name] = per_lead.get(lead_name, np.nan)\n            image_stats.append(row)\n\n            if n < VALID_PLOT_FIRST:\n                plt.figure(figsize=(16, 4))\n                plt.imshow(rectified)\n                plt.title(f\"{image_id}-{type_id} rectified\")\n                plt.axis(\"off\")\n                plt.show()\n\n        except Exception as e:\n            print(f\"Failed on {filename}: {e}\")\n            continue\n\n    stats_df = pd.DataFrame(image_stats)\n    stats_csv_path = os.path.join(VAL_OUT_DIR, \"stats.csv\")\n    stats_df.to_csv(stats_csv_path, index=False)\n    print(f\"\\nSaved validation stats => {stats_csv_path} (n={len(stats_df)})\")\n    print(f\"Saved validation detail(all) => {VAL_DETAIL_ALL}\")\n\n    print(\"\\n\" + \"=\" * 80)\n    print(\"VALIDATION SUMMARY\")\n    print(\"=\" * 80)\n\n    all_db = []\n    for tid, arr in type_snrs.items():\n        all_db.extend(arr)\n\n    if len(all_db) > 0:\n        overall_avg_linear = np.mean([10 ** (db / 10) for db in all_db])\n        overall_avg_db = 10 * np.log10(overall_avg_linear)\n        print(f\"Overall Average SNR: {overall_avg_db:.2f} dB (n={len(all_db)})\")\n    else:\n        print(\"No validation results collected.\")\n\n    print(\"\\nAverage SNR by Type:\")\n    for tid in sorted(type_snrs.keys()):\n        if len(type_snrs[tid]) == 0:\n            continue\n        avg_linear = np.mean([10 ** (db / 10) for db in type_snrs[tid]])\n        avg_db = 10 * np.log10(avg_linear)\n        print(f\"  Type {tid}: {avg_db:.2f} dB (n={len(type_snrs[tid])})\")\n\n    print(\"\\nAverage SNR by Lead:\")\n    for lead in LEADS:\n        if len(lead_snrs[lead]) == 0:\n            continue\n        avg_linear = np.mean([10 ** (db / 10) for db in lead_snrs[lead]])\n        avg_db = 10 * np.log10(avg_linear)\n        print(f\"  {lead:>3}: {avg_db:.2f} dB (n={len(lead_snrs[lead])})\")\n\n    if total_signal > 0 and total_noise >= 0:\n        total_snr_db = 10.0 * np.log10(compute_snr(total_signal, total_noise))\n        print(f\"\\nTotal SNR (sum power over all validation samples): {total_snr_db:.2f} dB\")\n    else:\n        print(\"\\nTotal SNR: cannot compute (no collected signals).\")\n\n    if len(all_db) > 0:\n        plt.figure(figsize=(10, 6))\n        data = [type_snrs[tid] for tid in sorted(type_snrs.keys()) if len(type_snrs[tid]) > 0]\n        labels = [tid for tid in sorted(type_snrs.keys()) if len(type_snrs[tid]) > 0]\n        plt.boxplot(data, labels=labels)\n        plt.title(\"SNR Distribution by Image Type\")\n        plt.xlabel(\"type_id\")\n        plt.ylabel(\"SNR (dB)\")\n        plt.grid(True, axis=\"y\", alpha=0.3)\n        plt.tight_layout()\n        plt.savefig(os.path.join(VAL_PLOT_DIR, \"summary-snr-by-type.png\"), dpi=110)\n        plt.show()\n        plt.close()\n\n        plt.figure(figsize=(14, 6))\n        data = [lead_snrs[ld] for ld in LEADS if len(lead_snrs[ld]) > 0]\n        labels = [ld for ld in LEADS if len(lead_snrs[ld]) > 0]\n        plt.boxplot(data, labels=labels)\n        plt.title(\"SNR Distribution by Lead\")\n        plt.xlabel(\"lead\")\n        plt.ylabel(\"SNR (dB)\")\n        plt.grid(True, axis=\"y\", alpha=0.3)\n        plt.tight_layout()\n        plt.savefig(os.path.join(VAL_PLOT_DIR, \"summary-snr-by-lead.png\"), dpi=110)\n        plt.show()\n        plt.close()\n\n    print(\"\\n✅ Validation complete!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:29:27.331448Z","iopub.execute_input":"2026-01-22T07:29:27.331733Z","iopub.status.idle":"2026-01-22T07:30:29.644493Z","shell.execute_reply.started":"2026-01-22T07:29:27.331710Z","shell.execute_reply":"2026-01-22T07:30:29.643730Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import gc\ngc.collect()\ntorch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:30:29.645756Z","iopub.execute_input":"2026-01-22T07:30:29.645973Z","iopub.status.idle":"2026-01-22T07:30:29.898537Z","shell.execute_reply.started":"2026-01-22T07:30:29.645955Z","shell.execute_reply":"2026-01-22T07:30:29.897970Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%reset -f","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:30:29.899170Z","iopub.execute_input":"2026-01-22T07:30:29.899425Z","iopub.status.idle":"2026-01-22T07:30:30.332539Z","shell.execute_reply.started":"2026-01-22T07:30:29.899395Z","shell.execute_reply":"2026-01-22T07:30:30.331749Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# James Parts","metadata":{}},{"cell_type":"markdown","source":"# Stage 1 - Perspective correction","metadata":{}},{"cell_type":"code","source":"CANONICAL_WIDTH = 2200\nCANONICAL_HEIGHT = 1700","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:30:30.333704Z","iopub.execute_input":"2026-01-22T07:30:30.334234Z","iopub.status.idle":"2026-01-22T07:30:30.337335Z","shell.execute_reply.started":"2026-01-22T07:30:30.334208Z","shell.execute_reply":"2026-01-22T07:30:30.336823Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json\nfrom pathlib import Path\nfrom typing import List, Tuple, Dict\n\nimport cv2\nimport numpy as np\nimport torch\nimport torch.nn as nn\nfrom torch.amp import autocast\nfrom torchvision.models.segmentation import deeplabv3_resnet50\nfrom torchvision.models import ResNet50_Weights\nfrom tqdm import tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:30:30.339059Z","iopub.execute_input":"2026-01-22T07:30:30.339303Z","iopub.status.idle":"2026-01-22T07:30:30.351988Z","shell.execute_reply.started":"2026-01-22T07:30:30.339286Z","shell.execute_reply":"2026-01-22T07:30:30.351346Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_canonical_keypoints() -> Dict[str, Dict[str, Tuple[float, float]]]:\n    POINTS_JSON = \"\"\"\n    {\n      \"I\": {\n        \"text_center\": [123, 752],\n        \"bar_line_intersect\": [118, 716]\n      },\n      \"aVR\": {\n        \"text_center\": [641, 753],\n        \"bar_line_intersect\": [610, 716]\n      },\n      \"V1\": {\n        \"text_center\": [1122, 753],\n        \"bar_line_intersect\": [1102, 716]\n      },\n      \"V4\": {\n        \"text_center\": [1614, 752],\n        \"bar_line_intersect\": [1594, 716],\n        \"end\": [2088, 716]\n      },\n      \"II_subset\": {\n        \"text_center\": [127, 1037],\n        \"bar_line_intersect\": [119, 992]\n      },\n      \"aVL\": {\n        \"text_center\": [640, 1037],\n        \"bar_line_intersect\": [610, 992]\n      },\n      \"V2\": {\n        \"text_center\": [1122, 1036],\n        \"bar_line_intersect\": [1102, 992]\n      },\n      \"V5\": {\n        \"text_center\": [1615, 1037],\n        \"bar_line_intersect\": [1594, 992],\n        \"end\": [2087, 992]\n      },\n      \"III\": {\n        \"text_center\": [131, 1319],\n        \"bar_line_intersect\": [118, 1268]\n      },\n      \"aVF\": {\n        \"text_center\": [641, 1319],\n        \"bar_line_intersect\": [610, 1268]\n      },\n      \"V3\": {\n        \"text_center\": [1122, 1319],\n        \"bar_line_intersect\": [1102, 1268]\n      },\n      \"V6\": {\n        \"text_center\": [1615, 1319],\n        \"bar_line_intersect\": [1594, 1268],\n        \"end\": [2088, 1268]\n      },\n      \"II\": {\n        \"text_center\": [127, 1588],\n        \"start\": [118, 1544],\n        \"end\": [2087, 1544]\n      }\n    }\n    \"\"\"\n    raw_points = json.loads(POINTS_JSON)\n    canonical_points: Dict[str, Dict[str, Tuple[float, float]]] = {}\n    for lead_name, point_mapping in raw_points.items():\n        canonical_points[lead_name] = {}\n        for point_type, coordinate_pair in point_mapping.items():\n            x_coordinate, y_coordinate = coordinate_pair\n            canonical_points[lead_name][point_type] = (float(x_coordinate), float(y_coordinate))\n    return canonical_points\n\n\ndef flatten_canonical_keypoints(\n    canonical_points: Dict[str, Dict[str, Tuple[float, float]]]\n) -> List[Tuple[str, float, float]]:\n    \"\"\"\n    Flatten the nested dict {lead_name: {point_type: (x, y)}} into an ordered\n    list of (keypoint_name, x, y).\n\n    keypoint_name is \"lead_name.point_type\".\n    \"\"\"\n    flattened: List[Tuple[str, float, float]] = []\n    for lead_name in sorted(canonical_points.keys()):\n        point_mapping = canonical_points[lead_name]\n        for point_type in sorted(point_mapping.keys()):\n            x_coordinate, y_coordinate = point_mapping[point_type]\n            keypoint_name = f\"{lead_name}.{point_type}\"\n            flattened.append((keypoint_name, x_coordinate, y_coordinate))\n    return flattened\n\ndef extract_predicted_keypoints_from_heatmaps(\n    logits: torch.Tensor,\n) -> torch.Tensor:\n    \"\"\"\n    Args:\n        logits: (B, K, H, W) raw scores from the model.\n\n    Returns:\n        predicted_keypoints: (B, K, 2) tensor with (x, y) coordinates\n        in pixel space (0..W-1, 0..H-1).\n    \"\"\"\n    batch_size, num_keypoints, height, width = logits.shape\n\n    flattened = logits.view(batch_size, num_keypoints, -1)  # (B, K, H*W)\n    max_indices = torch.argmax(flattened, dim = -1)  # (B, K)\n\n    y_coordinates = max_indices // width\n    x_coordinates = max_indices % width\n\n    predicted_keypoints = torch.stack(\n        [x_coordinates, y_coordinates],\n        dim = -1,\n    )  # (B, K, 2)\n\n    return predicted_keypoints.to(torch.float32)\n\n\ndef compute_homography_from_keypoints(\n    detected_points: np.ndarray,\n    canonical_points: np.ndarray,\n) -> Tuple[np.ndarray, np.ndarray]:\n    \"\"\"\n    detected_points: (K, 2) in input image coordinates\n    canonical_points: (K, 2) in canonical coordinates\n\n    Returns:\n        H: (3, 3) homography mapping detected_points -> canonical_points\n        inliers: mask of inliers returned by RANSAC\n    \"\"\"\n    if detected_points.shape[0] < 4:\n        raise RuntimeError(\"Need at least 4 keypoints to estimate homography.\")\n\n    H, inliers = cv2.findHomography(\n        detected_points.astype(np.float32),\n        canonical_points.astype(np.float32),\n        method = cv2.RANSAC,\n        ransacReprojThreshold = 7.0,\n    )\n    if H is None:\n        raise RuntimeError(\"cv2.findHomography failed to estimate a valid homography.\")\n\n    return H, inliers\n\n\ndef rectify_image_with_homography(\n    image_bgr: np.ndarray,\n    homography_matrix: np.ndarray,\n    output_size: Tuple[int, int],\n) -> np.ndarray:\n    \"\"\"\n    Apply homography to the input image to produce a rectified canonical image.\n\n    image_bgr: input image (H, W, 3) BGR\n    homography_matrix: (3, 3) mapping input coords -> canonical coords\n    output_size: (canonical_width, canonical_height)\n    \"\"\"\n    canonical_width, canonical_height = output_size\n\n    rectified_bgr = cv2.warpPerspective(\n        image_bgr,\n        homography_matrix,\n        dsize = (canonical_width, canonical_height),\n        flags = cv2.INTER_LINEAR,\n        borderMode = cv2.BORDER_CONSTANT,\n        borderValue = (255, 255, 255),\n    )\n    return rectified_bgr","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:30:30.352788Z","iopub.execute_input":"2026-01-22T07:30:30.353015Z","iopub.status.idle":"2026-01-22T07:30:30.368360Z","shell.execute_reply.started":"2026-01-22T07:30:30.353000Z","shell.execute_reply":"2026-01-22T07:30:30.367687Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport timm\n\nLEAD_ORDER = [\"I\", \"aVR\", \"V1\", \"V4\", \"II_subset\", \"aVL\", \"V2\", \"V5\", \"III\", \"aVF\", \"V3\", \"V6\", \"II\"]\n\nclass DifferentiableSignalExtractor(nn.Module):\n    def __init__(self, num_classes, height, mask_scale_y, top_k=1, \n                 reference_height=CANONICAL_HEIGHT, lead_order=LEAD_ORDER):\n        super().__init__()\n        self.num_classes = num_classes\n        self.height = height\n        self.top_k = top_k\n        \n        # --- SCALING FIX ---\n        # If the model input height differs from the reference template height,\n        # we must scale the LEAD_CONFIG coordinates.\n        # e.g. If Input is 800 and Ref is 1600, scale is 0.5. BaseY 708 becomes 354.\n        # Note: 'height' passed here is input_size[0] * mask_scale_y\n        effective_model_height = height / mask_scale_y\n        vertical_scale = effective_model_height / reference_height\n        \n        print(f\"DEBUG: SignalExtractor Vertical Scale: {vertical_scale:.4f} (Ref: {reference_height} -> In: {effective_model_height})\")\n\n        # Grid of Y coordinates\n        y_coords = torch.arange(height, dtype=torch.float32)\n        \n        voltage_grids = []\n        for lead_name in lead_order:\n            if lead_name not in LEAD_CONFIG:\n                base_y = height / 2.0\n            else:\n                # 1. Scale Baseline to Model Resolution\n                base_y_ref = LEAD_CONFIG[lead_name][\"base_y\"]\n                base_y = (base_y_ref * vertical_scale) * mask_scale_y\n            \n            # 2. Scale Pixels-Per-MV to Model Resolution\n            # If image shrinks, pixels get smaller, so we need FEWER pixels per mV.\n            scaled_px_per_mv = (PX_PER_MV * vertical_scale) * mask_scale_y\n            \n            volts = (base_y - y_coords) / scaled_px_per_mv\n            voltage_grids.append(volts)\n            \n        self.voltage_grids = torch.stack(voltage_grids) \n        self.register_buffer(\"voltage_lookup\", self.voltage_grids.view(1, num_classes, height, 1))\n\n    def forward(self, logits):\n        \"\"\"\n        logits: (B, C, H, W)\n        \"\"\"\n        B, C, H, W = logits.shape\n        \n        # 1. Top-K Selection\n        topk_logits, topk_indices = torch.topk(logits, k=self.top_k, dim=2)\n        \n        # 2. Local Softmax\n        topk_probs = torch.softmax(topk_logits, dim=2) # (B, C, K, W)\n        \n        # 3. Gather Voltage Values\n        batch_voltage_grid = self.voltage_lookup.expand(B, -1, -1, W)\n        gathered_volts = torch.gather(batch_voltage_grid, 2, topk_indices)\n        \n        # 4. Weighted Average\n        predicted_signals = torch.sum(topk_probs * gathered_volts, dim=2)\n        \n        return predicted_signals\n\n# ==========================================\n# MODEL: ViT Segmenter\n# ==========================================\nclass SimpleViTSeg(nn.Module):\n    def __init__(self, model_name, num_classes=13, img_size=(1694, 2198), patch_size=14, \n                 pretrained=True, mask_scale_x=1, mask_scale_y=1, use_signal_extractor=True):\n        super().__init__()\n        self.num_classes = num_classes\n        self.patch_size = patch_size\n        \n        self.encoder = timm.create_model(\n            model_name, pretrained=pretrained, in_chans=3, img_size=img_size,\n            num_classes=0, global_pool='', dynamic_img_size=True\n        )\n        embed_dim = self.encoder.num_features\n        \n        self.upsampler = self._build_upsampler(embed_dim, mask_scale_x, mask_scale_y)\n        \n        self.output_dim_per_patch = (patch_size * patch_size) * num_classes\n        self.head = nn.Linear(embed_dim, self.output_dim_per_patch)\n\n        target_h = int(img_size[0] * mask_scale_y)\n        \n        self.use_signal_extractor = use_signal_extractor\n        if use_signal_extractor:\n            self.signal_extractor = DifferentiableSignalExtractor(\n                num_classes, target_h, mask_scale_y, \n                top_k=10,\n                reference_height=CANONICAL_HEIGHT\n            )\n\n    def _build_upsampler(self, dim, scale_x, scale_y):\n        layers = []\n        steps_x = int(math.log2(scale_x)) if scale_x > 1 else 0\n        steps_y = int(math.log2(scale_y)) if scale_y > 1 else 0\n        total_steps = max(steps_x, steps_y)\n        \n        for i in range(total_steps):\n            s_x = 2 if i < steps_x else 1\n            s_y = 2 if i < steps_y else 1\n            if s_x == 1 and s_y == 1: continue\n            layers.append(nn.ConvTranspose2d(dim, dim, (s_y, s_x), (s_y, s_x), 0))\n            layers.append(nn.GELU()) \n        return nn.Sequential(*layers)\n\n    def forward(self, x):\n        features = self.encoder.forward_features(x)\n        B, H, W = x.shape[0], x.shape[2], x.shape[3]\n        H_grid, W_grid = H // self.patch_size, W // self.patch_size\n        \n        features = features[:, -(H_grid * W_grid):, :]\n        features = features.permute(0, 2, 1).view(B, -1, H_grid, W_grid)\n        \n        if len(self.upsampler) > 0:\n            features = self.upsampler(features)\n            \n        H_grid_new, W_grid_new = features.shape[2], features.shape[3]\n        features = features.flatten(2).transpose(1, 2)\n        patches = self.head(features)\n        \n        P, C = self.patch_size, self.num_classes\n        patches = patches.view(B, H_grid_new, W_grid_new, P, P, C)\n        reconstructed = patches.permute(0, 5, 1, 3, 2, 4).contiguous()\n        masks = reconstructed.view(B, C, H_grid_new * P, W_grid_new * P)\n        \n        if self.use_signal_extractor:\n            signals = self.signal_extractor(masks)\n            return masks, signals\n        else:\n            return masks\n\ndef build_keypoint_detector(num_keypoints: int, model_path: Path, device: torch.device) -> nn.Module:\n    # model = deeplabv3_resnet50(num_classes = num_keypoints, weights_backbone = None)\n    model = SimpleViTSeg(\n        model_name='vit_base_patch14_reg4_dinov2.lvd142m',\n        num_classes=num_keypoints,\n        img_size=(1036, 1036),\n        patch_size=14,\n        pretrained=False,\n        mask_scale_x=1,\n        mask_scale_y=1,\n        use_signal_extractor=False,\n    )\n    \n    state_dict = torch.load(model_path, map_location = \"cpu\")\n    cleaned_state_dict = {\n        key.removeprefix(\"_orig_mod.\"): value\n        for key, value in state_dict.items()\n    }\n    \n    model.load_state_dict(cleaned_state_dict)\n    model.to(device)\n    model.eval()\n    model=torch.compile(model)\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:30:30.369061Z","iopub.execute_input":"2026-01-22T07:30:30.369218Z","iopub.status.idle":"2026-01-22T07:30:30.385848Z","shell.execute_reply.started":"2026-01-22T07:30:30.369206Z","shell.execute_reply":"2026-01-22T07:30:30.385023Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_rectified_image(image_path: str, model: nn.Module) -> np.ndarray:\n    image_bgr = cv2.imread(str(image_path), cv2.IMREAD_COLOR)\n\n    original_height, original_width = image_bgr.shape[:2]\n    image_rgb = cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB)\n\n    # input_width, input_height = 1024, 1024\n    input_width, input_height = 1036, 1036\n    resized_rgb = cv2.resize(\n        image_rgb,\n        (input_width, input_height),\n        interpolation = cv2.INTER_LINEAR,\n    )\n\n    image_tensor = torch.from_numpy(resized_rgb.astype(np.float32) / 255.0)\n    device = next(model.parameters()).device\n    image_tensor = image_tensor.permute(2, 0, 1).unsqueeze(0).to(device, non_blocking = True)\n\n    with torch.no_grad():\n        with autocast(\"cuda\", dtype = torch.float16):\n            output_dict = model(image_tensor)\n            # logits = output_dict[\"out\"]  # (1, K, H, W)\n            logits = output_dict\n\n    predicted_keypoints_resized = extract_predicted_keypoints_from_heatmaps(logits)[0]  # (K, 2)\n    predicted_keypoints_resized = predicted_keypoints_resized.cpu().numpy()\n\n    scale_x = float(original_width) / float(input_width)\n    scale_y = float(original_height) / float(input_height)\n\n    detected_points_original = predicted_keypoints_resized.copy()\n    detected_points_original[:, 0] *= scale_x\n    detected_points_original[:, 1] *= scale_y\n\n    homography_matrix, inliers = compute_homography_from_keypoints(\n        detected_points = detected_points_original,\n        canonical_points = canonical_coords,\n    )\n\n    # print('Stage 1 homography:', homography_matrix)\n\n    rectified_bgr: np.ndarray = rectify_image_with_homography(\n        image_bgr = image_bgr,\n        homography_matrix = homography_matrix,\n        output_size = (CANONICAL_WIDTH, CANONICAL_HEIGHT),\n    )\n    return rectified_bgr","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:30:30.386466Z","iopub.execute_input":"2026-01-22T07:30:30.386742Z","iopub.status.idle":"2026-01-22T07:30:30.404145Z","shell.execute_reply.started":"2026-01-22T07:30:30.386726Z","shell.execute_reply":"2026-01-22T07:30:30.403356Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"canonical_points_dict = load_canonical_keypoints()\nflattened_canonical = flatten_canonical_keypoints(canonical_points_dict)\nkeypoint_names: List[str] = [name for name, _, _ in flattened_canonical]\ncanonical_coords = np.array(\n    [(x_coordinate, y_coordinate) for _, x_coordinate, y_coordinate in flattened_canonical],\n    dtype = np.float32,\n)\n\nkeypoint_detector_0 = build_keypoint_detector(\n    num_keypoints=canonical_coords.shape[0], \n    model_path='/kaggle/input/ecg-stage-1-keypoint-detectors/fold_1_24e_5sigma_albu.pth',\n    device=torch.device('cuda:0')\n)\n\nkeypoint_detector_1 = build_keypoint_detector(\n    num_keypoints=canonical_coords.shape[0], \n    model_path='/kaggle/input/ecg-stage-1-keypoint-detectors/fold_1_48e_5sigma_heavyAlbu.pth',\n    device=torch.device('cuda:1')\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:30:30.404959Z","iopub.execute_input":"2026-01-22T07:30:30.405180Z","iopub.status.idle":"2026-01-22T07:30:39.639106Z","shell.execute_reply.started":"2026-01-22T07:30:30.405160Z","shell.execute_reply":"2026-01-22T07:30:39.638422Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Test stage 1","metadata":{}},{"cell_type":"code","source":"import cv2\nimport matplotlib.pyplot as plt\n\n# SAMPLE_IMAGE_ID = '4275164612'\nSAMPLE_IMAGE_ID = '3944898701'\nimage_bgr = get_rectified_image(\n    # image_path=f'/kaggle/input/physionet-ecg-image-digitization/train/{SAMPLE_IMAGE_ID}/{SAMPLE_IMAGE_ID}-0005.png',\n    image_path=f'/kaggle/input/physionet-ecg-image-digitization/train/{SAMPLE_IMAGE_ID}/{SAMPLE_IMAGE_ID}-0006.png',\n    model=keypoint_detector_0\n)\n\n# Convert BGR to RGB for correct colors in matplotlib\nimage_rgb = cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB)\n\nplt.imshow(image_rgb)\nplt.axis('off')  # optional: hide axes\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:30:39.639952Z","iopub.execute_input":"2026-01-22T07:30:39.640395Z","iopub.status.idle":"2026-01-22T07:31:01.533958Z","shell.execute_reply.started":"2026-01-22T07:30:39.640374Z","shell.execute_reply":"2026-01-22T07:31:01.533153Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_bgr = get_rectified_image(\n    image_path=f'/kaggle/input/physionet-ecg-image-digitization/train/{SAMPLE_IMAGE_ID}/{SAMPLE_IMAGE_ID}-0006.png',\n    model=keypoint_detector_1\n)\n\n# Convert BGR to RGB for correct colors in matplotlib\nimage_rgb = cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB)\n\nplt.imshow(image_rgb)\nplt.axis('off')  # optional: hide axes\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:31:01.534678Z","iopub.execute_input":"2026-01-22T07:31:01.534932Z","iopub.status.idle":"2026-01-22T07:31:13.497392Z","shell.execute_reply.started":"2026-01-22T07:31:01.534911Z","shell.execute_reply":"2026-01-22T07:31:13.496776Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Stage 3 - lead segmentation","metadata":{}},{"cell_type":"code","source":"import torch\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport cv2\nfrom PIL import Image\nfrom torchvision.transforms import v2\nimport timm\nimport pandas as pd\nimport math","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:31:13.498284Z","iopub.execute_input":"2026-01-22T07:31:13.498557Z","iopub.status.idle":"2026-01-22T07:31:13.590069Z","shell.execute_reply.started":"2026-01-22T07:31:13.498537Z","shell.execute_reply":"2026-01-22T07:31:13.589341Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n# Must match training config\nMODEL_NAME = \"vit_base_patch14_reg4_dinov2.lvd142m\"\nINPUT_SIZE = (1694, 4396) # (H, W)\nPATCH_SIZE = 14\nNUM_CLASSES = 13","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:31:13.590859Z","iopub.execute_input":"2026-01-22T07:31:13.591584Z","iopub.status.idle":"2026-01-22T07:31:13.595067Z","shell.execute_reply.started":"2026-01-22T07:31:13.591560Z","shell.execute_reply":"2026-01-22T07:31:13.594307Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Constants ---\nPX_PER_MV = 78.0\n\n# Final verified config from HostSegmentationDataset\nLEAD_CONFIG = {\n  # Row 1:\n  \"I\":         {\"start\": [117.5, 708],  \"base_y\": 708},\n  \"aVR\":       {\"start\": [610, 708],  \"base_y\": 708},\n  \"V1\":        {\"start\": [1102, 708], \"base_y\": 708},\n  \"V4\":        {\"start\": [1594, 708], \"base_y\": 708},\n\n  # Row 2:\n  \"II_subset\": {\"start\": [117.5, 991],  \"base_y\": 991}, \n  \"aVL\":       {\"start\": [610, 991],  \"base_y\": 991},\n  \"V2\":        {\"start\": [1102, 991], \"base_y\": 991},\n  \"V5\":        {\"start\": [1594, 991], \"base_y\": 991},\n\n  # Row 3:\n  \"III\":       {\"start\": [117.5, 1274.5], \"base_y\": 1274.5},\n  \"aVF\":       {\"start\": [610, 1274.5], \"base_y\": 1274.5},\n  \"V3\":        {\"start\": [1102, 1274.5], \"base_y\": 1274.5},\n  \"V6\":        {\"start\": [1594, 1274.5], \"base_y\": 1274.5},\n\n  # Row 4 (Rhythm):\n  \"II\":        {\"start\": [117, 1534], \"base_y\": 1534, \"end\": [2086, 1534]}\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:31:13.598090Z","iopub.execute_input":"2026-01-22T07:31:13.598330Z","iopub.status.idle":"2026-01-22T07:31:13.610963Z","shell.execute_reply.started":"2026-01-22T07:31:13.598314Z","shell.execute_reply":"2026-01-22T07:31:13.610363Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_segmentation_model(checkpoint_path, device):\n    print(f\"Creating model: {MODEL_NAME}...\")\n    # Instantiate the model with the exact same args as training\n    model = SimpleViTSeg(\n        model_name=MODEL_NAME,\n        num_classes=NUM_CLASSES,\n        img_size=INPUT_SIZE,\n        patch_size=PATCH_SIZE,\n        pretrained=False,\n        mask_scale_x=2\n    )\n    \n    # Load weights\n    print(f\"Loading weights from {checkpoint_path}...\")\n    state_dict = torch.load(checkpoint_path, map_location=DEVICE)\n\n    cleaned_state_dict = {\n        key.removeprefix(\"_orig_mod.\"): value\n        for key, value in state_dict.items()\n    }\n    \n    model.load_state_dict(cleaned_state_dict, strict=True)\n    model.to(device)\n    model.eval()\n    model=torch.compile(model)\n    print(\"Model loaded successfully.\")\n    return model\n\n# Initialize\nSTANDARD_CHECKPOINT_PATH_0 = \"/kaggle/input/ecg-lead-segmentation-models/vit_ep12_snr12.4_super_res_Sx2_bs2_lr1en4_12e_finetune.pth\"\nSTANDARD_CHECKPOINT_PATH_1 = \"/kaggle/input/ecg-lead-segmentation-models/vit_ep9_snr16.4_super_res_Sx2_bs1_lr7en5_12e_3hrHeavyAug_finetune.pth\"\nFLIPPED_CHECKPOINT_PATH_0 = \"/kaggle/input/ecg-lead-segmentation-models/vit_ep12_snr11.9_super_res_Sx2_bs2_lr1en4_12e_finetune_flipped_mse02.pth\"\nFLIPPED_CHECKPOINT_PATH_1 = \"/kaggle/input/ecg-lead-segmentation-models/vit_ep12_snr13.1_super_res_Sx2_bs2_lr1en4_12e_finetune_flipped.pth\"\n\nstandard_segmentation_model_0 = load_segmentation_model(STANDARD_CHECKPOINT_PATH_0, 'cuda:0')\nstandard_segmentation_model_1 = load_segmentation_model(STANDARD_CHECKPOINT_PATH_1, 'cuda:0')\nflipped_segmentation_model_0 = load_segmentation_model(FLIPPED_CHECKPOINT_PATH_0, 'cuda:1')\nflipped_segmentation_model_1 = load_segmentation_model(FLIPPED_CHECKPOINT_PATH_1, 'cuda:1')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:31:13.611582Z","iopub.execute_input":"2026-01-22T07:31:13.611806Z","iopub.status.idle":"2026-01-22T07:31:34.550195Z","shell.execute_reply.started":"2026-01-22T07:31:13.611786Z","shell.execute_reply":"2026-01-22T07:31:34.549529Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn.functional as F\n\nclass MultiKSignalExtractor(nn.Module):\n    \"\"\"\n    Differentiable signal extractor that computes voltage traces and confidence\n    metrics for multiple top-k values along the vertical (H) dimension.\n\n    Input:\n        logits: Tensor of shape (B, C, H, W)\n            B = batch size\n            C = num_classes / leads\n            H = height (vertical axis over which argmax/top-k is applied)\n            W = width (time axis)\n\n    Output:\n        metrics: Tensor of shape (B, C, F, W)\n\n        The feature axis F is a concatenation over all requested k values.\n        For k_values = [1, 3, 5, 10, 20], the feature layout is:\n\n        F index | k   | metric\n        --------+-----+----------------------\n          0     |  1  | expected_voltage\n          1     |  1  | top1_probability (sigmoid of max logit)\n\n          2     |  3  | expected_voltage\n          3     |  3  | min_voltage (over top-3 support)\n          4     |  3  | max_voltage (over top-3 support)\n          5     |  3  | entropy_nats (over top-3 probs)\n\n          6     |  5  | expected_voltage\n          7     |  5  | min_voltage\n          8     |  5  | max_voltage\n          9     |  5  | entropy_nats\n\n          10    | 10  | expected_voltage\n          11    | 10  | min_voltage\n          12    | 10  | max_voltage\n          13    | 10  | entropy_nats\n\n          14    | 20  | expected_voltage\n          15    | 20  | min_voltage\n          16    | 20  | max_voltage\n          17    | 20  | entropy_nats\n\n        General rule:\n          - For k == 1:\n              [ expected_voltage, top1_probability ]\n          - For k > 1:\n              [ expected_voltage, min_voltage, max_voltage, entropy_nats ]\n\n        You can inspect `self.feature_specs` for a programmatic mapping:\n            self.feature_specs[i] = (k_value, metric_name)\n    \"\"\"\n\n    def __init__(\n        self,\n        num_classes=NUM_CLASSES,\n        height=INPUT_SIZE[0],\n        mask_scale_y=1,\n        top_k_values=(1, 3, 5, 10, 20),\n        reference_height=CANONICAL_HEIGHT,\n        lead_order=LEAD_ORDER,\n        min_max_prob_threshold=1e-3\n    ):\n        super().__init__()\n        self.num_classes = num_classes\n        self.height = height\n        self.mask_scale_y = mask_scale_y\n        self.min_max_prob_threshold = float(min_max_prob_threshold)\n\n        # Normalize and check k values\n        unique_sorted_k_values = sorted(set(int(k) for k in top_k_values))\n        if 1 not in unique_sorted_k_values:\n            raise ValueError(\"k=1 must be present in top_k_values.\")\n        self.top_k_values = unique_sorted_k_values\n        self.max_k = max(self.top_k_values)\n\n        # --- SCALING FIX (as in original) ---\n        effective_model_height = height / mask_scale_y\n        vertical_scale = effective_model_height / reference_height\n\n        print(\n            f\"DEBUG: SignalExtractor Vertical Scale: \"\n            f\"{vertical_scale:.4f} (Ref: {reference_height} -> In: {effective_model_height})\"\n        )\n\n        # Grid of Y coordinates\n        y_coords = torch.arange(height, dtype=torch.float32)\n\n        voltage_grids = []\n        for lead_name in lead_order:\n            if lead_name not in LEAD_CONFIG:\n                base_y = height / 2.0\n            else:\n                base_y_ref = LEAD_CONFIG[lead_name][\"base_y\"]\n                base_y = (base_y_ref * vertical_scale) * mask_scale_y\n\n            scaled_px_per_mv = (PX_PER_MV * vertical_scale) * mask_scale_y\n            volts = (base_y - y_coords) / scaled_px_per_mv\n            voltage_grids.append(volts)\n\n        self.voltage_grids = torch.stack(voltage_grids)\n        self.register_buffer(\"voltage_lookup\", self.voltage_grids.view(1, num_classes, height, 1))\n\n        # Build a feature specification list for the F dimension\n        feature_specs = []\n        for k in self.top_k_values:\n            if k == 1:\n                feature_specs.append((k, \"expected_voltage\"))\n                feature_specs.append((k, \"top1_probability\"))\n            else:\n                feature_specs.append((k, \"expected_voltage\"))\n                feature_specs.append((k, \"min_voltage\"))\n                feature_specs.append((k, \"max_voltage\"))\n                feature_specs.append((k, \"entropy_nats\"))\n        self.feature_specs = feature_specs  # length == F\n\n    def forward(self, logits):\n        \"\"\"\n        logits: (B, C, H, W)\n\n        Returns:\n            metrics: (B, C, F, W)\n\n        F dimension layout is described in the class docstring and\n        in self.feature_specs.\n        \"\"\"\n        B, C, H, W = logits.shape\n        if H != self.height:\n            raise ValueError(f\"Logits height H={H} does not match configured height={self.height}.\")\n\n        # Compute top-max_k once and reuse for all smaller k\n        topk_logits_all, topk_indices_all = torch.topk(logits, k=self.max_k, dim=2)  # (B, C, max_k, W)\n\n        # Voltage grid broadcasted to batch\n        batch_voltage_grid = self.voltage_lookup.expand(B, -1, -1, W)  # (B, C, H, W)\n\n        feature_tensors = []\n\n        for k in self.top_k_values:\n            # Slice the first k elements from the global top-max_k\n            current_logits = topk_logits_all[:, :, :k, :]      # (B, C, k, W)\n            current_indices = topk_indices_all[:, :, :k, :]    # (B, C, k, W)\n\n            current_probs = F.softmax(current_logits, dim=2)   # (B, C, k, W)\n            current_volts = torch.gather(batch_voltage_grid, 2, current_indices)  # (B, C, k, W)\n\n            # Expected value (soft argmax over top-k support)\n            expected_voltage = torch.sum(current_probs * current_volts, dim=2)  # (B, C, W)\n\n            if k == 1:\n                # For k=1:\n                #   - expected_voltage is just the voltage at the single top position\n                #   - top1_probability is sigmoid(max_logit)\n                top1_logits = current_logits[:, :, 0, :]  # (B, C, W)\n                top1_probability = torch.sigmoid(top1_logits)\n\n                metrics_for_k = torch.stack(\n                    [\n                        expected_voltage,   # (B, C, W)\n                        top1_probability,   # (B, C, W)\n                    ],\n                    dim=2,  # -> (B, C, 2, W)\n                )\n            else:\n                # Unmasked min/max (fallback if threshold masks everything)\n                unmasked_min_voltage, _ = torch.min(current_volts, dim=2)  # (B, C, W)\n                unmasked_max_voltage, _ = torch.max(current_volts, dim=2)  # (B, C, W)\n\n                # Masked min/max: only consider voltages where prob > threshold\n                prob_mask = current_probs > self.min_max_prob_threshold  # (B, C, k, W)\n                any_valid = prob_mask.any(dim=2)  # (B, C, W)\n\n                positive_infinity = torch.tensor(float(\"inf\"), device=current_volts.device, dtype=current_volts.dtype)\n                negative_infinity = torch.tensor(float(\"-inf\"), device=current_volts.device, dtype=current_volts.dtype)\n\n                volts_for_min = current_volts.masked_fill(~prob_mask, positive_infinity)\n                volts_for_max = current_volts.masked_fill(~prob_mask, negative_infinity)\n\n                masked_min_voltage, _ = torch.min(volts_for_min, dim=2)  # (B, C, W) (may be +inf)\n                masked_max_voltage, _ = torch.max(volts_for_max, dim=2)  # (B, C, W) (may be -inf)\n\n                min_voltage = torch.where(any_valid, masked_min_voltage, unmasked_min_voltage)\n                max_voltage = torch.where(any_valid, masked_max_voltage, unmasked_max_voltage)\n\n                # Entropy in nats (natural log)\n                safe_probs = current_probs.clamp_min(1e-12)\n                entropy_nats = -torch.sum(safe_probs * torch.log(safe_probs), dim=2)  # (B, C, W)\n\n                metrics_for_k = torch.stack(\n                    [\n                        expected_voltage,  # (B, C, W)\n                        min_voltage,       # (B, C, W)\n                        max_voltage,       # (B, C, W)\n                        entropy_nats,      # (B, C, W)\n                    ],\n                    dim=2,  # -> (B, C, 4, W)\n                )\n\n            feature_tensors.append(metrics_for_k)  # list of (B, C, m_k, W)\n\n        # Concatenate along the feature axis to get (B, C, F, W)\n        metrics = torch.cat(feature_tensors, dim=2)\n\n        return metrics\n\nstandard_segmentation_model_0.signal_extractor = MultiKSignalExtractor().to('cuda:0')\nstandard_segmentation_model_1.signal_extractor = MultiKSignalExtractor().to('cuda:0')\nflipped_segmentation_model_0.signal_extractor = MultiKSignalExtractor().to('cuda:1')\nflipped_segmentation_model_1.signal_extractor = MultiKSignalExtractor().to('cuda:1')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:31:34.551093Z","iopub.execute_input":"2026-01-22T07:31:34.551374Z","iopub.status.idle":"2026-01-22T07:31:34.584181Z","shell.execute_reply.started":"2026-01-22T07:31:34.551351Z","shell.execute_reply":"2026-01-22T07:31:34.583455Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport cv2\nimport torch\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom torchvision.transforms import v2\n\n# Define a distinct color palette (RGB) for visualization\n\n# We skip pure Green because that is used for Ground Truth\n\nLEAD_COLORS = [\n    (255, 0, 0),      # Red\n    (0, 255, 255),    # Cyan\n    (255, 0, 255),    # Magenta\n    (255, 255, 0),    # Yellow\n    (255, 165, 0),    # Orange\n    (147, 112, 219),  # Medium Purple\n    (255, 192, 203),  # Pink\n    (30, 144, 255),   # Dodger Blue\n    (173, 255, 47),   # Green Yellow\n    (255, 99, 71),    # Tomato\n    (240, 230, 140),  # Khaki\n    (255, 255, 255)   # White\n]\n\ndef get_host_ground_truth_heatmap(signal_csv_path, train_csv_path, target_shape, sigma=5.0, sharp=False):\n    \"\"\"\n    Generates GT heatmap from CSV signals using the pixel-perfect projection logic.\n    \"\"\"\n    h, w = target_shape\n    \n    # 1. Load Metadata (FS)\n    # Note: In a loop, you'd cache train.csv, but for single-shot vis, reading it here is fine.\n    df_meta = pd.read_csv(train_csv_path)\n    record_id = Path(signal_csv_path).stem\n    \n    # Robust ID matching\n    if df_meta['id'].dtype == np.int64:\n        try:\n            row = df_meta[df_meta['id'] == int(record_id)]\n        except ValueError:\n             row = df_meta[df_meta['id'] == record_id]\n    else:\n        row = df_meta[df_meta['id'] == str(record_id)]\n        \n    if row.empty:\n        print(f\"⚠️ ID {record_id} not found in {train_csv_path}\")\n        return np.zeros((h, w), dtype=np.float32)\n\n    fs = float(row.iloc[0]['fs'])\n    \n    # 2. Load Signal\n    df_sig = pd.read_csv(signal_csv_path)\n    \n    # 3. Calculate Scale\n    rhythm_start_x = LEAD_CONFIG[\"II\"][\"start\"][0]\n    rhythm_end_x = LEAD_CONFIG[\"II\"][\"end\"][0]\n    rhythm_width_px = rhythm_end_x - rhythm_start_x\n    px_per_sec = rhythm_width_px / 10.0\n    \n    lead_masks = []\n    \n    for lead_key, config in LEAD_CONFIG.items():\n        start_xy = config[\"start\"]\n        base_y = config[\"base_y\"]\n        \n        # Extract Signal\n        signal_data = None\n        if lead_key == \"II_subset\":\n            ref_len = len(df_sig[\"I\"].dropna())\n            full_ii = df_sig[\"II\"].values\n            signal_data = full_ii[:ref_len]\n        elif lead_key == \"II\":\n            signal_data = df_sig[\"II\"].values\n        else:\n            if lead_key in df_sig.columns:\n                signal_data = df_sig[lead_key].dropna().values\n                \n        if signal_data is None or len(signal_data) == 0:\n            continue\n            \n        # Project to Pixels\n        num_samples = len(signal_data)\n        time_steps = np.arange(num_samples) / fs\n        x_coords = start_xy[0] + (time_steps * px_per_sec)\n        y_coords = base_y - (signal_data * PX_PER_MV)\n        \n        # Rounding (Matches Dataset Class)\n        pts = np.stack([np.rint(x_coords), np.rint(y_coords)], axis=1).astype(np.int32)\n        \n        # Clip\n        valid_mask = (pts[:,0] >= 0) & (pts[:,0] < w) & (pts[:,1] >= 0) & (pts[:,1] < h)\n        pts = pts[valid_mask]\n        \n        if len(pts) < 2:\n            continue\n            \n        pts_array = pts.reshape((-1, 1, 2))\n        temp_layer = np.zeros((h, w), dtype=np.float32)\n        \n        if sharp:\n            # Draw Sharp (Thickness 2 for visibility in matplotlib plots)\n            cv2.polylines(temp_layer, [pts_array], isClosed=False, color=1.0, thickness=2, lineType=cv2.LINE_AA)\n            lead_masks.append(temp_layer)\n        else:\n            # Draw Soft (Blur)\n            cv2.polylines(temp_layer, [pts_array], isClosed=False, color=1.0, thickness=1, lineType=cv2.LINE_AA)\n            k_size = int(6 * sigma) | 1\n            blurred = cv2.GaussianBlur(temp_layer, (k_size, k_size), sigmaX=sigma, sigmaY=sigma)\n            if blurred.max() > 0:\n                blurred /= blurred.max()\n            lead_masks.append(blurred)\n            \n    if not lead_masks:\n        return np.zeros((h, w), dtype=np.float32)\n        \n    combined_gt = np.stack(lead_masks, axis=0).max(axis=0)\n    return combined_gt\n\n# --- Helper: Calculate Lead Boundaries ---\ndef get_lead_pixel_ranges(config):\n    \"\"\"\n    Parses LEAD_CONFIG to determine the [start_x, end_x) for each lead.\n    \n    Logic: \n    1. The end of one lead is the start of the next lead on the same row.\n    2. The LAST lead in a row is clipped to the end of the Rhythm II strip.\n    \"\"\"\n    ranges = {}\n    \n    # 1. Determine Global Max X from Rhythm Strip (II)\n    # This fixes the \"too far to the right\" issue\n    if \"II\" in config and \"end\" in config[\"II\"]:\n        global_end_x = int(config[\"II\"][\"end\"][0])\n    else:\n        # Fallback just in case, though your config has it\n        global_end_x = 2086 \n\n    # 2. Group keys by base_y to handle rows\n    rows = {}\n    for key, params in config.items():\n        y = params[\"base_y\"]\n        if y not in rows:\n            rows[y] = []\n        rows[y].append((key, params))\n    \n    for y, items in rows.items():\n        # Sort by start x coordinate\n        items.sort(key=lambda x: x[1][\"start\"][0])\n        \n        for i in range(len(items)):\n            key, params = items[i]\n            start_x = int(params[\"start\"][0])\n            \n            # Determine End X\n            if \"end\" in params:\n                # Explicit end defined (e.g., Rhythm strip II itself)\n                end_x = int(params[\"end\"][0])\n            elif i < len(items) - 1:\n                # End is the start of the next lead in the row\n                end_x = int(items[i+1][1][\"start\"][0])\n            else:\n                # Last lead in row extends ONLY to the Global Max X\n                end_x = global_end_x\n            \n            ranges[key] = (start_x, end_x)\n            \n    return ranges\n\ndef visualize_host_digitization(model, rgb_image: np.ndarray, signal_csv_path, train_csv_path, flip_image):\n    # 1. Format conversion\n    original_pil = Image.fromarray(rgb_image)\n    orig_w, orig_h = original_pil.size\n    \n    # 2. Generate Ground Truth (Sharp Lines) using Host Logic\n    gt_heatmap_sharp = get_host_ground_truth_heatmap(\n        signal_csv_path, \n        train_csv_path, \n        (orig_h, orig_w), \n        sharp=True\n    )\n    \n    # 3. Generate Model Prediction\n    transform = v2.Compose([\n        v2.RandomHorizontalFlip(p=int(flip_image)),\n        v2.Resize(INPUT_SIZE, antialias=True),\n        v2.ToImage(), \n        v2.ToDtype(torch.float32, scale=True),\n    ])\n    model_device = next(model.parameters()).device\n    input_tensor = transform(original_pil).unsqueeze(0).to(model_device)\n    \n    with torch.no_grad():\n        with torch.amp.autocast(DEVICE, dtype=torch.float16):\n            logits, _ = model(input_tensor)\n            preds = torch.sigmoid(logits)\n\n    if flip_image:\n        preds = torch.flip(preds, [3])\n    \n    # Remove batch dim -> [num_classes, H, W]\n    preds_squeeze = preds[0]\n    \n    # Map channels to Lead Keys\n    lead_keys = list(LEAD_CONFIG.keys())\n    \n    # Pre-calculate valid X-ranges for every lead\n    # UPDATED: Now strictly clips to the rhythm strip end\n    lead_ranges = get_lead_pixel_ranges(LEAD_CONFIG)\n\n    # 4. Visualization Setup\n    canvas = np.zeros((orig_h, orig_w, 3), dtype=np.uint8)\n    \n    # --- LAYER 1: Ground Truth (Faint Green) ---\n    gt_mask = gt_heatmap_sharp > 0.5\n    canvas[gt_mask] = [0, 30, 0] \n\n    total_points = 0\n\n    # --- LAYER 2: Model Predictions (Spatially Constrained Argmax) ---\n    for lead_idx, lead_key in enumerate(lead_keys):\n        \n        if lead_idx >= preds_squeeze.shape[0]:\n            break\n\n        # Get specific lead heatmap\n        lead_heatmap_small = preds_squeeze[lead_idx].float().cpu().numpy()\n        \n        # Resize to Original Resolution\n        lead_heatmap_full = cv2.resize(lead_heatmap_small, (orig_w, orig_h), interpolation=cv2.INTER_LINEAR)\n        \n        # Get valid X boundaries for this specific lead\n        x_start, x_end = lead_ranges.get(lead_key, (0, 0))\n        \n        # Sanity clip bounds to image size\n        x_start = max(0, x_start)\n        x_end = min(orig_w, x_end)\n        \n        if x_end <= x_start:\n            continue\n\n        # --- STRICT COLUMN EXTRACTION ---\n        # 1. Slice the heatmap only in the valid X range\n        roi = lead_heatmap_full[:, x_start:x_end]\n        \n        # 2. Argmax down the columns (find best Y for every X)\n        y_indices = np.argmax(roi, axis=0)\n        \n        # 3. Create X indices corresponding to the slice\n        x_indices = np.arange(x_start, x_end)\n        \n        # 4. Zip into points\n        digitized_points = list(zip(x_indices, y_indices))\n        \n        total_points += len(digitized_points)\n        \n        # Select Color\n        color = LEAD_COLORS[lead_idx % len(LEAD_COLORS)]\n        \n        # Draw Points\n        for x, y in digitized_points:\n            cv2.circle(canvas, (x, y), 0, color, -1)\n\n    # 5. Plot\n    plt.figure(figsize=(24, 10))\n    plt.title(f\"Digitization: Spatially Constrained Argmax\\n(Clipped to Rhythm Strip Length) - Points: {total_points}\")\n    plt.imshow(canvas)\n    plt.axis(\"off\")\n    plt.tight_layout()\n    plt.show()\n    \n    print(f\"Extracted {total_points} total trace points.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:31:34.585094Z","iopub.execute_input":"2026-01-22T07:31:34.585304Z","iopub.status.idle":"2026-01-22T07:31:34.991360Z","shell.execute_reply.started":"2026-01-22T07:31:34.585289Z","shell.execute_reply":"2026-01-22T07:31:34.990744Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_host_digitization(\n    standard_segmentation_model_0, \n    image_rgb, \n    f'/kaggle/input/physionet-ecg-image-digitization/train/{SAMPLE_IMAGE_ID}/{SAMPLE_IMAGE_ID}.csv', \n    f'/kaggle/input/physionet-ecg-image-digitization/train.csv',\n    flip_image=False\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:31:34.992150Z","iopub.execute_input":"2026-01-22T07:31:34.992452Z","iopub.status.idle":"2026-01-22T07:32:11.826887Z","shell.execute_reply.started":"2026-01-22T07:31:34.992434Z","shell.execute_reply":"2026-01-22T07:32:11.826234Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_host_digitization(\n    standard_segmentation_model_1, \n    image_rgb, \n    f'/kaggle/input/physionet-ecg-image-digitization/train/{SAMPLE_IMAGE_ID}/{SAMPLE_IMAGE_ID}.csv', \n    f'/kaggle/input/physionet-ecg-image-digitization/train.csv',\n    flip_image=False\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:32:11.827594Z","iopub.execute_input":"2026-01-22T07:32:11.827834Z","iopub.status.idle":"2026-01-22T07:32:18.434479Z","shell.execute_reply.started":"2026-01-22T07:32:11.827809Z","shell.execute_reply":"2026-01-22T07:32:18.433873Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_host_digitization(\n    flipped_segmentation_model_0, \n    image_rgb, \n    f'/kaggle/input/physionet-ecg-image-digitization/train/{SAMPLE_IMAGE_ID}/{SAMPLE_IMAGE_ID}.csv', \n    f'/kaggle/input/physionet-ecg-image-digitization/train.csv',\n    flip_image=True\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:32:18.435235Z","iopub.execute_input":"2026-01-22T07:32:18.435528Z","iopub.status.idle":"2026-01-22T07:32:41.549752Z","shell.execute_reply.started":"2026-01-22T07:32:18.435485Z","shell.execute_reply":"2026-01-22T07:32:41.549128Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_host_digitization(\n    flipped_segmentation_model_1, \n    image_rgb, \n    f'/kaggle/input/physionet-ecg-image-digitization/train/{SAMPLE_IMAGE_ID}/{SAMPLE_IMAGE_ID}.csv', \n    f'/kaggle/input/physionet-ecg-image-digitization/train.csv',\n    flip_image=True\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:32:41.550438Z","iopub.execute_input":"2026-01-22T07:32:41.550711Z","iopub.status.idle":"2026-01-22T07:32:47.831544Z","shell.execute_reply.started":"2026-01-22T07:32:41.550691Z","shell.execute_reply":"2026-01-22T07:32:47.830783Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nfrom scipy.ndimage import grey_opening, grey_closing\n\ndef sharpen_signal(raw_signal, kernel_size=6, amount=0.3):\n    # 1. White Top-Hat: Extracts positive peaks (R-waves)\n    #    Formula: Image - Opening(Image)\n    opening = grey_opening(raw_signal, size=kernel_size)\n    white_tophat = raw_signal - opening\n\n    # 2. Black Top-Hat: Extracts negative peaks (Q/S-waves)\n    #    Formula: Closing(Image) - Image\n    closing = grey_closing(raw_signal, size=kernel_size)\n    black_tophat = closing - raw_signal\n\n    # 3. Combine: Add positive peaks, subtract negative peaks\n    #    This stretches the signal vertically only at local extrema.\n    sharpened_signal = raw_signal + (amount * (white_tophat - black_tophat))\n    return sharpened_signal\n\ndef align_einthoven_triad(lead_I, lead_II, lead_III, max_shift=10):\n    \"\"\"\n    Aligns I and III to the coordinate system of Lead II (Anchor).\n    Returns: aligned_I, aligned_II, aligned_III, shift_I, shift_III\n    \"\"\"\n    # Center signals\n    s1, s2, s3 = lead_I - np.mean(lead_I), lead_II - np.mean(lead_II), lead_III - np.mean(lead_III)\n    \n    # Compare only on the shortest common length\n    n_compare = min(len(s1), len(s2), len(s3))\n    strict_slice_2 = s2[:n_compare]\n    \n    # Grid search for shifts\n    shifts = np.arange(-max_shift, max_shift + 1)\n    best_loss, best_tau_1, best_tau_3 = float('inf'), 0, 0\n    \n    s1_pad = np.pad(s1, max_shift, mode='constant', constant_values=np.nan)\n    s3_pad = np.pad(s3, max_shift, mode='constant', constant_values=np.nan)\n    \n    for t1 in shifts:\n        slice_1 = s1_pad[max_shift + t1 : max_shift + t1 + n_compare]\n        for t3 in shifts:\n            slice_3 = s3_pad[max_shift + t3 : max_shift + t3 + n_compare]\n            residual = (slice_1 + slice_3) - strict_slice_2\n            valid = np.isfinite(residual)\n            if np.sum(valid) > n_compare * 0.5:\n                loss = np.sum(residual[valid]**2)\n                if loss < best_loss:\n                    best_loss, best_tau_1, best_tau_3 = loss, t1, t3\n\n    # Project to Anchor (Lead II) timeline\n    def project(arr, shift, length):\n        out = np.full(length, np.nan)\n        idx = np.arange(length) + shift\n        valid = (idx >= 0) & (idx < len(arr))\n        out[valid] = arr[idx[valid]]\n        return out\n        \n    return (project(lead_I, best_tau_1, len(lead_II)), \n            lead_II.copy(), \n            project(lead_III, best_tau_3, len(lead_II)), \n            best_tau_1, best_tau_3)\n\ndef align_augmented_triad(lead_1, lead_2, lead_3, max_shift=10):\n    \"\"\"Aligns aVR, aVL, aVF (Sum=0). Anchor is lead_2 (aVL).\"\"\"\n    s1, s2, s3 = lead_1 - np.mean(lead_1), lead_2 - np.mean(lead_2), lead_3 - np.mean(lead_3)\n    n_compare = min(len(s1), len(s2), len(s3))\n    strict_slice_2 = s2[:n_compare]\n    \n    shifts = np.arange(-max_shift, max_shift + 1)\n    best_loss, best_tau_1, best_tau_3 = float('inf'), 0, 0\n    s1_pad = np.pad(s1, max_shift, mode='constant', constant_values=np.nan)\n    s3_pad = np.pad(s3, max_shift, mode='constant', constant_values=np.nan)\n    \n    for t1 in shifts:\n        slice_1 = s1_pad[max_shift + t1 : max_shift + t1 + n_compare]\n        for t3 in shifts:\n            slice_3 = s3_pad[max_shift + t3 : max_shift + t3 + n_compare]\n            residual = slice_1 + strict_slice_2 + slice_3 # Sum = 0\n            valid = np.isfinite(residual)\n            if np.sum(valid) > n_compare * 0.5:\n                loss = np.sum(residual[valid]**2)\n                if loss < best_loss:\n                    best_loss, best_tau_1, best_tau_3 = loss, t1, t3\n\n    def project(arr, shift, length):\n        out = np.full(length, np.nan)\n        idx = np.arange(length) + shift\n        valid = (idx >= 0) & (idx < len(arr))\n        out[valid] = arr[idx[valid]]\n        return out\n\n    return (project(lead_1, best_tau_1, len(lead_2)), \n            lead_2.copy(), \n            project(lead_3, best_tau_3, len(lead_2)), \n            best_tau_1, best_tau_3)\n\ndef process_einthoven_integration(results: dict, alpha: float = 0.5) -> None:\n    # Requires I, II, III\n    if not all(k in results for k in ['I', 'II', 'III']): return\n    \n    raw_I = results['I']['voltage']\n    raw_II = results['II']['voltage']\n    raw_III = results['III']['voltage']\n\n    # 1. Align Triad\n    aI, aII, aIII, t1, t3 = align_einthoven_triad(raw_I, raw_II, raw_III)\n\n    # 2. DC Correct\n    resid = (aI + aIII) - aII\n    dc = np.nanmean(resid)\n    aI -= dc/3.0; aIII -= dc/3.0; aII += dc/3.0\n\n    # 3. Compute Expectations\n    exp_I, exp_III, exp_II = aII - aIII, aII - aI, aI + aIII\n\n    # 4. Write Back Helper\n    def write_back(target, aligned, shift):\n        k = np.arange(len(target))\n        src_idx = k - shift\n        valid = (src_idx >= 0) & (src_idx < len(aligned))\n        valid_data = np.isfinite(aligned[src_idx[valid]])\n        target[k[valid][valid_data]] = aligned[src_idx[valid][valid_data]]\n\n    # Blend II (Direct)\n    mask = np.isfinite(aII) & np.isfinite(exp_II)\n    results['II']['voltage'][mask] = (aII[mask] + alpha * exp_II[mask]) / (1 + alpha)\n\n    # Blend I (Reverse Shift)\n    mask = np.isfinite(aI) & np.isfinite(exp_I)\n    aI[mask] = (aI[mask] + alpha * exp_I[mask]) / (1 + alpha)\n    write_back(results['I']['voltage'], aI, t1)\n\n    # Blend III (Reverse Shift)\n    mask = np.isfinite(aIII) & np.isfinite(exp_III)\n    aIII[mask] = (aIII[mask] + alpha * exp_III[mask]) / (1 + alpha)\n    write_back(results['III']['voltage'], aIII, t3)\n\n    # 5. Merge II_subset if it exists\n    if 'II_subset' in results:\n        raw_sub = results['II_subset']['voltage']\n        # We process II_subset against the *updated* II prefix\n        sub_len = min(len(raw_sub), len(results['II']['voltage']))\n        target_prefix = results['II']['voltage'][:sub_len]\n        \n        # Simple Pairwise Correlation Alignment (subset -> II)\n        # Replaces calling 'align_signals' to avoid heavy scipy.optimize dependency\n        shifts = np.arange(-10, 11)\n        best_s, best_err = 0, float('inf')\n        t_nomean = target_prefix - np.mean(target_prefix)\n        s_nomean = raw_sub - np.mean(raw_sub)\n        \n        for s in shifts:\n            # We shift 'raw_sub' to match 'target_prefix'\n            # Index map: target[i] ~ raw_sub[i+s]\n            idx = np.arange(sub_len) + s\n            valid = (idx >= 0) & (idx < len(s_nomean))\n            if np.sum(valid) > sub_len * 0.5:\n                diff = t_nomean[valid] - s_nomean[idx[valid]]\n                err = np.sum(diff**2)\n                if err < best_err: best_err, best_s = err, s\n        \n        # Create aligned subset array\n        aligned_sub = np.full(sub_len, np.nan)\n        idx = np.arange(sub_len) + best_s\n        valid = (idx >= 0) & (idx < len(raw_sub))\n        aligned_sub[valid] = raw_sub[idx[valid]]\n        \n        # Fill gaps with current II values\n        inv = ~np.isfinite(aligned_sub)\n        aligned_sub[inv] = target_prefix[inv]\n        \n        # Weighted Blend into II (using your 0.8 / 1.8 logic)\n        results['II']['voltage'][:sub_len] += aligned_sub * 0.8\n        results['II']['voltage'][:sub_len] /= 1.8\n\ndef process_augmented_integration(results: dict, alpha: float = 0.5) -> None:\n    if not all(k in results for k in ['aVR', 'aVL', 'aVF']): return\n    \n    raw_VR = results['aVR']['voltage']\n    raw_VL = results['aVL']['voltage']\n    raw_VF = results['aVF']['voltage']\n\n    # 1. Align Triad (Anchor aVL)\n    aVR, aVL, aVF, tVR, tVF = align_augmented_triad(raw_VR, raw_VL, raw_VF)\n\n    # 2. DC Correct (Sum=0)\n    resid = aVR + aVL + aVF\n    dc = np.nanmean(resid)\n    aVR -= dc/3.0; aVL -= dc/3.0; aVF -= dc/3.0\n\n    # 3. Expectations\n    exp_VR = -(aVL + aVF)\n    exp_VL = -(aVR + aVF)\n    exp_VF = -(aVR + aVL)\n    \n    def write_back(target, aligned, shift):\n        k = np.arange(len(target))\n        src_idx = k - shift\n        valid = (src_idx >= 0) & (src_idx < len(aligned))\n        valid_data = np.isfinite(aligned[src_idx[valid]])\n        target[k[valid][valid_data]] = aligned[src_idx[valid][valid_data]]\n\n    # Blend aVL (Anchor)\n    mask = np.isfinite(aVL) & np.isfinite(exp_VL)\n    results['aVL']['voltage'][mask] = (aVL[mask] + alpha * exp_VL[mask]) / (1 + alpha)\n\n    # Blend aVR (Shifted)\n    mask = np.isfinite(aVR) & np.isfinite(exp_VR)\n    aVR[mask] = (aVR[mask] + alpha * exp_VR[mask]) / (1 + alpha)\n    write_back(results['aVR']['voltage'], aVR, tVR)\n\n    # Blend aVF (Shifted)\n    mask = np.isfinite(aVF) & np.isfinite(exp_VF)\n    aVF[mask] = (aVF[mask] + alpha * exp_VF[mask]) / (1 + alpha)\n    write_back(results['aVF']['voltage'], aVF, tVF)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:32:47.832423Z","iopub.execute_input":"2026-01-22T07:32:47.832651Z","iopub.status.idle":"2026-01-22T07:32:47.858439Z","shell.execute_reply.started":"2026-01-22T07:32:47.832634Z","shell.execute_reply":"2026-01-22T07:32:47.857835Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport cv2\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport torch\nfrom torchvision.transforms import v2\nfrom PIL import Image\n\n# Re-using the logic from the previous step to determine X-bounds\ndef get_lead_pixel_ranges(config):\n    ranges = {}\n    \n    # 1. Determine Global Max X from Rhythm Strip (II)\n    if \"II\" in config and \"end\" in config[\"II\"]:\n        global_end_x = int(config[\"II\"][\"end\"][0])\n    else:\n        global_end_x = 2086 \n\n    # 2. Group keys by base_y\n    rows = {}\n    for key, params in config.items():\n        y = params[\"base_y\"]\n        if y not in rows:\n            rows[y] = []\n        rows[y].append((key, params))\n    \n    for y, items in rows.items():\n        items.sort(key=lambda x: x[1][\"start\"][0])\n        for i in range(len(items)):\n            key, params = items[i]\n            start_x = int(params[\"start\"][0])\n            \n            if \"end\" in params:\n                end_x = int(params[\"end\"][0])\n            elif i < len(items) - 1:\n                end_x = int(items[i+1][1][\"start\"][0])\n            else:\n                end_x = global_end_x\n            \n            ranges[key] = (start_x, end_x)\n    return ranges\n\n\n\ndef extract_signals_directly(\n    predicted_signals: torch.Tensor, \n    lead_config: dict = LEAD_CONFIG, \n    lead_order: list = LEAD_ORDER,\n    canonical_width: int = CANONICAL_WIDTH,\n    safety_margin_px: int = 2\n) -> dict:\n    \"\"\"\n    Extracts physical signals directly from vector outputs.\n    \n    Includes:\n    1. Geometric Time Mapping (Fixes Phase Shift)\n    2. Safety Cropping (Fixes Boundary Spikes/Dips)\n    3. Tail Extrapolation (Fixes Zero-Drop at end)\n    \"\"\"\n    \n    # 1. Setup\n    if isinstance(predicted_signals, torch.Tensor):\n        signals_np = predicted_signals.detach().cpu().float().numpy()\n    else:\n        signals_np = predicted_signals\n        \n    num_channels, tensor_width = signals_np.shape\n    \n    scale_x = tensor_width / float(canonical_width)\n    \n    # Rhythm II setup for time scaling\n    rhythm_start = lead_config[\"II\"][\"start\"][0]\n    rhythm_end = lead_config[\"II\"][\"end\"][0]\n    total_rhythm_pixels = rhythm_end - rhythm_start\n    px_per_sec_canonical = total_rhythm_pixels / 10.0\n    \n    lead_ranges = get_lead_pixel_ranges(lead_config)\n    \n    extracted_results = {}\n    \n    for ch_idx, lead_name in enumerate(lead_order):\n        if ch_idx >= num_channels:\n            break\n            \n        full_signal = signals_np[ch_idx]\n        can_start, can_end = lead_ranges.get(lead_name, (0, 0))\n        \n        # 2. Geometric Mapping\n        idx_start = int(np.floor(can_start * scale_x))\n        idx_end = int(np.ceil(can_end * scale_x))\n        \n        # # 3. Apply Safety Margin (The Fix for Spikes)\n        # # We deliberately stop reading N pixels before the geometric boundary\n        # # to avoid the artifacts where leads bleed into each other or padding.\n        # idx_end = idx_end - safety_margin_px\n        \n        # Sanity checks\n        idx_start = max(0, idx_start)\n        idx_end = min(tensor_width, idx_end)\n        \n        # Ensure we have enough data to be meaningful\n        if idx_end <= idx_start:\n            continue\n            \n        # 4. Extract Voltage\n        voltage_segment = full_signal[idx_start:idx_end]\n        voltage_segment[-safety_margin_px:] = voltage_segment[-(safety_margin_px + 1)]\n        \n        # 5. Generate Time Vector\n        tensor_indices = np.arange(idx_start, idx_end)\n        \n        # Map Tensor Index -> Canonical Pixel -> Relative to Lead Start -> Seconds\n        canonical_locs = tensor_indices / scale_x\n        relative_locs = canonical_locs - can_start\n        time_segment = relative_locs / px_per_sec_canonical\n        \n        # 6. Clean Extrapolation (The Fix for Length Mismatch)\n        # Instead of letting resample_to_target zero-fill the missing ms,\n        # we append a point at the EXACT expected end time using the \n        # last known CLEAN voltage.\n        expected_duration = (can_end - can_start) / px_per_sec_canonical\n        \n        if len(time_segment) > 0:\n            # Check if we have a gap to fill (we almost always will due to the safety crop)\n            if time_segment[-1] < expected_duration:\n                time_segment = np.append(time_segment, expected_duration)\n                # Hold the last clean value flat\n                voltage_segment = np.append(voltage_segment, voltage_segment[-1])\n\n        # 7. Sharpen voltage segment.\n        voltage_segment = sharpen_signal(voltage_segment)\n        \n        extracted_results[lead_name] = {\n            \"time\": time_segment,\n            \"voltage\": voltage_segment\n        }\n\n    # 8. Exploit redundancy between leads.\n    # 8a. Apply Einthoven's law (I, II, III) + II_subset merge\n    process_einthoven_integration(extracted_results, alpha=0.5)\n\n    # # 8b. Apply augmented lead projection (aVR, aVL, aVF)\n    # process_augmented_integration(extracted_results, alpha=0.5)\n        \n    return extracted_results","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:32:47.859143Z","iopub.execute_input":"2026-01-22T07:32:47.859369Z","iopub.status.idle":"2026-01-22T07:32:47.877550Z","shell.execute_reply.started":"2026-01-22T07:32:47.859345Z","shell.execute_reply":"2026-01-22T07:32:47.876875Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport cv2\nimport torch\nfrom pathlib import Path\nfrom torchvision.transforms import v2\nfrom PIL import Image\nfrom scipy.interpolate import interp1d\n\ndef calculate_snr(gt_signal, pred_signal):\n    \"\"\"\n    Calculates SNR in dB.\n    Args:\n        gt_signal: Ground truth numpy array\n        pred_signal: Predicted numpy array (must be same length)\n    \"\"\"\n    # Calculate Signal Power\n    signal_power = np.sum(gt_signal ** 2)\n    \n    # Calculate Noise (Error) Power\n    noise_power = np.sum((gt_signal - pred_signal) ** 2)\n    \n    # Avoid division by zero (perfect reconstruction)\n    if noise_power < 1e-6:\n        return float('inf')\n    \n    # Avoid log of zero if signal is perfectly flat 0 (rare)\n    if signal_power < 1e-6:\n        return 0.0\n\n    snr = 10 * np.log10(signal_power / noise_power)\n    return snr\n\ndef interpolate_low_confidence_regions(batch_tensor: torch.Tensor, threshold: float = 0.4) -> torch.Tensor:\n    \"\"\"\n    Interpolates channel 10 (voltage) IN-PLACE for each lead where the confidence \n    interval (ch 12 - ch 11) exceeds the given threshold.\n    \n    Args:\n        batch_tensor: torch.Tensor of shape (batch_size, lead_count, channel_count, width)\n                      Example: (1, 13, 18, 8792)\n        threshold: float, maximum acceptable spread between upper and lower bounds.\n        \n    Returns:\n        torch.Tensor: The modified batch_tensor (returned for convenience, though modified in-place).\n    \"\"\"\n    batch_size, lead_count, channel_count, width = batch_tensor.shape\n    \n    # Pre-allocate x-indices for interpolation\n    x_indices = np.arange(width)\n\n    for b in range(batch_size):\n        for l in range(lead_count):\n            # 1. Access the specific lead's confidence bounds\n            # Shape is (Width,) e.g., (8792,)\n            # Channel 11: Lower Bound\n            # Channel 12: Upper Bound\n            lower_bound = batch_tensor[b, l, 11, :]\n            upper_bound = batch_tensor[b, l, 12, :]\n            \n            confidence_spread = upper_bound - lower_bound\n            \n            # 2. Identify reliable data points\n            valid_mask = confidence_spread <= threshold\n            \n            # 3. Skip if no work is needed (all good or all bad)\n            if valid_mask.all() or (~valid_mask).all():\n                continue\n\n            # 4. Move data to CPU/NumPy for interpolation\n            mask_np = valid_mask.cpu().numpy()\n            \n            # Extract signal from Channel 10\n            signal_np = batch_tensor[b, l, 10, :].cpu().numpy()\n            \n            valid_x = x_indices[mask_np]\n            valid_y = signal_np[mask_np]\n            target_x = x_indices[~mask_np]\n            \n            # 5. Interpolate\n            imputed_values = np.interp(target_x, valid_x, valid_y)\n            \n            # 6. Assign back In-Place\n            batch_tensor[b, l, 10, ~valid_mask] = torch.from_numpy(imputed_values).to(\n                device=batch_tensor.device, \n                dtype=batch_tensor.dtype\n            )\n\n    return batch_tensor\n\ndef visualize_signal_comparison(model, rgb_image: np.ndarray, signal_csv_path, train_csv_path, flip_image):\n    \"\"\"\n    Runs model inference, extracts signals, and plots them overlaying the Ground Truth CSV.\n    Displays SNR (dB) in titles.\n    \"\"\"\n    # 1. Load Metadata for Sampling Rate\n    df_meta = pd.read_csv(train_csv_path)\n    record_id = Path(signal_csv_path).stem\n    \n    if df_meta['id'].dtype == np.int64:\n        try:\n            row = df_meta[df_meta['id'] == int(record_id)]\n        except ValueError:\n            row = df_meta[df_meta['id'] == record_id]\n    else:\n        row = df_meta[df_meta['id'] == str(record_id)]\n\n    if row.empty:\n        fs = 500.0\n    else:\n        fs = float(row.iloc[0]['fs'])\n\n    # 2. Load Ground Truth\n    try:\n        df_sig = pd.read_csv(signal_csv_path)\n    except Exception as e:\n        print(f\"Error loading CSV: {e}\")\n        return\n\n    # 3. Inference\n    original_pil = Image.fromarray(rgb_image)\n    orig_w, orig_h = original_pil.size\n    \n    transform = v2.Compose([\n        v2.RandomHorizontalFlip(p=int(flip_image)),\n        v2.Resize(INPUT_SIZE, antialias=True),\n        v2.ToImage(), \n        v2.ToDtype(torch.float32, scale=True),\n    ])\n    model_device = next(model.parameters()).device\n    input_tensor = transform(original_pil).unsqueeze(0).to(model_device)\n    \n    with torch.no_grad():\n        with torch.amp.autocast(DEVICE, dtype=torch.float16):\n            _, predicted_signals = model(input_tensor)\n\n    if flip_image:\n        predicted_signals = torch.flip(predicted_signals, [3])\n\n    predicted_signals = interpolate_low_confidence_regions(predicted_signals)\n    predicted_signals = predicted_signals[:,:,10]\n    \n    # 4. Extract Signals\n    # Note: INPUT_SIZE is (1694, 2198), so width is index 1\n    pred_data = extract_signals_directly(predicted_signals[0])\n    \n    # 5. Plotting\n    fig = plt.figure(figsize=(24, 12))\n    gs = fig.add_gridspec(5, 4) \n    \n    layout_map = {\n        \"I\": (0, 0), \"aVR\": (0, 1), \"V1\": (0, 2), \"V4\": (0, 3),\n        \"II_subset\": (1, 0), \"aVL\": (1, 1), \"V2\": (1, 2), \"V5\": (1, 3),\n        \"III\": (2, 0), \"aVF\": (2, 1), \"V3\": (2, 2), \"V6\": (2, 3),\n    }\n    \n    snr_values = []\n    \n    # --- Plot 12 Standard Leads ---\n    for lead_name, (row, col) in layout_map.items():\n        ax = fig.add_subplot(gs[row, col])\n        csv_col = \"II\" if lead_name == \"II_subset\" else lead_name\n        \n        gt_interp_target = None\n        gt_time_target = None\n        \n        # A. Plot Ground Truth\n        if csv_col in df_sig.columns:\n            raw_data = df_sig[csv_col].dropna().values\n            if len(raw_data) > 0:\n                gt_time = np.arange(len(raw_data)) / fs\n                \n                # Crop II_subset to ~2.5s for grid visualization\n                if lead_name == \"II_subset\":\n                    crop_idx = int(2.5 * fs)\n                    if len(raw_data) > crop_idx:\n                        raw_data = raw_data[:crop_idx]\n                        gt_time = gt_time[:crop_idx]\n                \n                ax.plot(gt_time, raw_data, color='black', alpha=0.5, linewidth=1.5, label='GT')\n                gt_interp_target = raw_data\n                gt_time_target = gt_time\n\n        # B. Plot Prediction\n        if lead_name in pred_data:\n            p_time = pred_data[lead_name][\"time\"]\n            p_volt = pred_data[lead_name][\"voltage\"]\n            \n            snr_str = \"nan\"\n            if gt_interp_target is not None and len(p_volt) > 1:\n                # Interpolate Prediction to match GT timestamps\n                p_interp = np.interp(gt_time_target, p_time, p_volt)\n                # interpolation_fn = interp1d(p_time, p_volt, kind='cubic')\n                # p_interp = interpolation_fn(gt_time_target)\n                \n                # Calculate SNR\n                snr = calculate_snr(gt_interp_target, p_interp)\n                \n                if snr != float('inf'):\n                    snr_values.append(snr)\n                    snr_str = f\"{snr:.1f} dB\"\n                else:\n                    snr_str = \"Inf dB\"\n\n            ax.plot(p_time, p_volt, color='red', linewidth=1, label='Pred')\n            ax.set_title(f\"{lead_name} (SNR: {snr_str})\")\n        else:\n            ax.set_title(f\"{lead_name} (No Pred)\")\n        \n        ax.grid(True, linestyle=':', alpha=0.6)\n        if row != 2: ax.set_xticklabels([])\n\n    # --- Plot Rhythm Strip (II) ---\n    ax_rhythm = fig.add_subplot(gs[3:, :]) \n    \n    if \"II\" in df_sig.columns:\n        rhythm_data = df_sig[\"II\"].dropna().values\n        rhythm_time = np.arange(len(rhythm_data)) / fs\n        ax_rhythm.plot(rhythm_time, rhythm_data, color='black', alpha=0.5, linewidth=1.5, label='Ground Truth')\n        gt_interp_target = rhythm_data\n        gt_time_target = rhythm_time\n        \n    if \"II\" in pred_data:\n        p_time = pred_data[\"II\"][\"time\"]\n        p_volt = pred_data[\"II\"][\"voltage\"]\n        \n        snr_str = \"nan\"\n        if gt_interp_target is not None and len(p_volt) > 1:\n            p_interp = np.interp(gt_time_target, p_time, p_volt)\n            snr = calculate_snr(gt_interp_target, p_interp)\n            if snr != float('inf'):\n                snr_values.append(snr)\n                snr_str = f\"{snr:.1f} dB\"\n            else:\n                snr_str = \"Inf dB\"\n\n        ax_rhythm.plot(p_time, p_volt, color='red', linewidth=1, label='Prediction')\n        ax_rhythm.set_title(f\"Rhythm II (SNR: {snr_str})\")\n    \n    ax_rhythm.legend(loc='upper right')\n    ax_rhythm.grid(True, linestyle=':', alpha=0.6)\n    ax_rhythm.set_xlabel(\"Time (s)\")\n    \n    avg_snr = np.mean(snr_values) if snr_values else 0\n    plt.suptitle(f\"Digitization vs Ground Truth (ID: {record_id} | Avg SNR: {avg_snr:.1f} dB)\", fontsize=16)\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:32:47.878258Z","iopub.execute_input":"2026-01-22T07:32:47.878585Z","iopub.status.idle":"2026-01-22T07:32:47.900489Z","shell.execute_reply.started":"2026-01-22T07:32:47.878563Z","shell.execute_reply":"2026-01-22T07:32:47.899984Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_signal_comparison(\n    standard_segmentation_model_0, \n    image_rgb, \n    f'/kaggle/input/physionet-ecg-image-digitization/train/{SAMPLE_IMAGE_ID}/{SAMPLE_IMAGE_ID}.csv', \n    '/kaggle/input/physionet-ecg-image-digitization/train.csv',\n    flip_image=False\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:32:47.901226Z","iopub.execute_input":"2026-01-22T07:32:47.901647Z","iopub.status.idle":"2026-01-22T07:32:55.413411Z","shell.execute_reply.started":"2026-01-22T07:32:47.901623Z","shell.execute_reply":"2026-01-22T07:32:55.412755Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_signal_comparison(\n    standard_segmentation_model_1, \n    image_rgb, \n    f'/kaggle/input/physionet-ecg-image-digitization/train/{SAMPLE_IMAGE_ID}/{SAMPLE_IMAGE_ID}.csv', \n    '/kaggle/input/physionet-ecg-image-digitization/train.csv',\n    flip_image=False\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:32:55.414554Z","iopub.execute_input":"2026-01-22T07:32:55.414818Z","iopub.status.idle":"2026-01-22T07:33:02.314685Z","shell.execute_reply.started":"2026-01-22T07:32:55.414790Z","shell.execute_reply":"2026-01-22T07:33:02.313859Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_signal_comparison(\n    flipped_segmentation_model_0, \n    image_rgb, \n    f'/kaggle/input/physionet-ecg-image-digitization/train/{SAMPLE_IMAGE_ID}/{SAMPLE_IMAGE_ID}.csv', \n    '/kaggle/input/physionet-ecg-image-digitization/train.csv',\n    flip_image=True\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:33:02.315366Z","iopub.execute_input":"2026-01-22T07:33:02.315595Z","iopub.status.idle":"2026-01-22T07:33:09.097844Z","shell.execute_reply.started":"2026-01-22T07:33:02.315579Z","shell.execute_reply":"2026-01-22T07:33:09.097081Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_signal_comparison(\n    flipped_segmentation_model_1, \n    image_rgb, \n    f'/kaggle/input/physionet-ecg-image-digitization/train/{SAMPLE_IMAGE_ID}/{SAMPLE_IMAGE_ID}.csv', \n    '/kaggle/input/physionet-ecg-image-digitization/train.csv',\n    flip_image=True\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:33:09.098570Z","iopub.execute_input":"2026-01-22T07:33:09.098782Z","iopub.status.idle":"2026-01-22T07:33:15.953280Z","shell.execute_reply.started":"2026-01-22T07:33:09.098765Z","shell.execute_reply":"2026-01-22T07:33:15.952563Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SAMPLE_IMAGE_ID","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:33:15.954014Z","iopub.execute_input":"2026-01-22T07:33:15.954272Z","iopub.status.idle":"2026-01-22T07:33:15.959129Z","shell.execute_reply.started":"2026-01-22T07:33:15.954253Z","shell.execute_reply":"2026-01-22T07:33:15.958396Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"from concurrent.futures import ThreadPoolExecutor\nfrom typing import Dict, Optional, List, Tuple\nimport warnings\n\nimport cv2\nimport numpy as np\nfrom PIL import Image\n\ndef resample_to_target(time_src, volt_src, target_fs, num_samples):\n    \"\"\"\n    Resamples the physical signal (time_src, volt_src) to the specific \n    integer number of samples required by the submission format.\n    \"\"\"\n    if len(volt_src) < 2:\n        return np.zeros(num_samples)\n        \n    # Create the target time grid (0 to duration)\n    duration = num_samples / target_fs\n    target_time = np.linspace(0, duration, num_samples, endpoint=False)\n    \n    # Linear interpolation\n    # We use numpy interp for robustness\n    resampled_volt = np.interp(target_time, time_src, volt_src, left=0, right=0)\n    \n    return resampled_volt\n\ndef resample_to_target_v2(time_src, volt_src, target_fs, num_samples) -> np.ndarray:\n    predicted_voltages = volt_src[1:]\n    interpolated_predicted_voltages = np.interp(\n        np.linspace(0, len(predicted_voltages) - 1, num=num_samples),\n        np.arange(len(predicted_voltages)),\n        predicted_voltages,\n    )\n    return interpolated_predicted_voltages\n\ndef align_signals(label: np.ndarray, pred: np.ndarray, max_shift: float = float('inf')) -> np.ndarray:\n    # Initialize the reference and digitized signals.\n    label_arr = np.asarray(label, dtype=np.float64)\n    pred_arr = np.asarray(pred, dtype=np.float64)\n\n    label_mean = np.mean(label_arr)\n    pred_mean = np.mean(pred_arr)\n\n    label_arr_centered = label_arr - label_mean\n    pred_arr_centered = pred_arr - pred_mean\n\n    # Compute the correlation between the reference and digitized signals and locate the maximum correlation.\n    correlation = scipy.signal.correlate(label_arr_centered, pred_arr_centered, mode='full')\n\n    n_label = np.size(label_arr)\n    n_pred = np.size(pred_arr)\n\n    lags = scipy.signal.correlation_lags(n_label, n_pred, mode='full')\n    valid_lags_mask = (lags >= -max_shift) & (lags <= max_shift)\n\n    max_correlation = np.nanmax(correlation[valid_lags_mask])\n    all_max_indices = np.flatnonzero(correlation == max_correlation)\n    best_idx = min(all_max_indices, key=lambda i: abs(lags[i]))\n    time_shift = lags[best_idx]\n    start_padding_len = max(time_shift, 0)\n    pred_slice_start = max(-time_shift, 0)\n    pred_slice_end = min(n_label - time_shift, n_pred)\n    end_padding_len = max(n_label - n_pred - time_shift, 0)\n    aligned_pred = np.concatenate((np.full(start_padding_len, np.nan), pred_arr[pred_slice_start:pred_slice_end], np.full(end_padding_len, np.nan)))\n\n    def objective_func(v_shift):\n        return np.nansum((label_arr - (aligned_pred - v_shift)) ** 2)\n\n    if np.any(np.isfinite(label_arr) & np.isfinite(aligned_pred)):\n        results = scipy.optimize.minimize_scalar(objective_func, method='Brent')\n        vertical_shift = results.x\n        aligned_pred -= vertical_shift\n    return aligned_pred\n\ndef get_lead_heatmap(model, image_pil, flip_input, device_id):\n    \"\"\"\n    Worker function to run inside a thread.\n    Handles preprocessing, inference, and raw signal extraction for one model.\n    \"\"\"\n    # 1. Preprocess (CPU side)\n    # If flip_input is True, flip the image before resizing\n    if flip_input:\n        img_input = image_pil.transpose(Image.FLIP_LEFT_RIGHT)\n    else:\n        img_input = image_pil\n\n    transform = v2.Compose([\n        v2.Resize(INPUT_SIZE, antialias=True),\n        v2.ToImage(), \n        v2.ToDtype(torch.float32, scale=True),\n    ])\n    \n    input_tensor = transform(img_input).unsqueeze(0).to(device_id)\n\n    # 2. Inference (GPU side)\n    # We use a new CUDA stream for this thread to ensure true parallelism on the GPU\n    stream = torch.cuda.Stream(device=device_id)\n    with torch.cuda.stream(stream):\n        with torch.no_grad():\n            with torch.amp.autocast(\"cuda\", dtype=torch.float16):\n                heatmap, predicted_signals = model(input_tensor)\n                \n                # If we flipped the input, we must flip the output width back\n                if flip_input:\n                    heatmap = torch.flip(heatmap, [3])\n                #     predicted_signals = torch.flip(predicted_signals, [3])\n                \n                # predicted_signals = interpolate_low_confidence_regions(predicted_signals)\n                \n                # # Select the specific channel (e.g., 10)\n                # raw_signal_tensor = predicted_signals[:, :, 10]\n\n    # Synchronize to ensure inference is done before moving to CPU\n    stream.synchronize()\n    \n    # # 3. Extract Signals (CPU side)\n    # # Move to CPU and extract raw voltage/time data\n    # # This is often CPU-heavy, so doing it in the thread is beneficial\n    # raw_signal_cpu = raw_signal_tensor[0].cpu()\n    # extracted_data = extract_signals_directly(raw_signal_cpu)\n\n    # heatmap = heatmap.cpu()\n    \n    return heatmap\n\nshared_signal_extractor = MultiKSignalExtractor().to('cuda:0')\n\ndef process_single_image(\n    standard_model_0, \n    standard_model_1, \n    flipped_model_0, \n    flipped_model_1, \n    keypoint_model_0, \n    keypoint_model_1,\n    image_path, \n    target_fs, \n    df_rows_for_id\n):\n    # RECTIFY IMAGES.\n    # We define the task to run inside the threads\n    def rectify_task(model):\n        # get_rectified_image typically handles moving data to the model's device\n        img_bgr = get_rectified_image(image_path=image_path, model=model)\n        img_rgb = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB)\n        return Image.fromarray(img_rgb)\n\n    # Launch rectification for both keypoint detector versions\n    with ThreadPoolExecutor(max_workers=2) as kp_executor:\n        future_kp0 = kp_executor.submit(rectify_task, keypoint_model_0)\n        future_kp1 = kp_executor.submit(rectify_task, keypoint_model_1)\n        \n        # Resolve images\n        image_pil_0 = future_kp0.result()\n        image_pil_1 = future_kp1.result()\n\n    # COMPUTE HEATMAPS.\n    # To avoid race conditions on the same model instance, we process images sequentially.\n    # We still run Standard and Flipped models in parallel for each image.\n    \n    all_heatmaps = []\n    \n    # Helper to run the pair of models on a single image\n    def run_models_on_image(img_pil, standard_model, flipped_model):\n        with ThreadPoolExecutor(max_workers=2) as seg_executor:\n            dev_std = next(standard_model.parameters()).device\n            dev_flip = next(flipped_model.parameters()).device\n            \n            standard_heatmap_future = seg_executor.submit(\n                get_lead_heatmap, \n                standard_model, img_pil, False, dev_std\n            )\n            flipped_heatmap_future = seg_executor.submit(\n                get_lead_heatmap, \n                flipped_model, img_pil, True, dev_flip\n            )\n            return standard_heatmap_future.result(), flipped_heatmap_future.result()\n\n    # Process Image 0 (from Keypoint Model 0)\n    image0_standard_heatmap, image0_flip_heatmap = run_models_on_image(image_pil_0, standard_model_0, flipped_model_0)\n    \n    # Process Image 1 (from Keypoint Model 1)\n    image1_standard_heatmap, image1_flip_heatmap = run_models_on_image(image_pil_1, standard_model_1, flipped_model_1)\n\n    # AVERAGE HEATMAPS.\n    mean_heatmap_logits_0 = torch.stack([\n        image0_standard_heatmap,\n        image1_standard_heatmap, \n    ], dim = 0).mean(dim = 0)\n    \n    mean_heatmap_logits_1 = torch.stack([\n        image0_flip_heatmap,\n        image1_flip_heatmap, \n    ], dim = 0).to('cuda:0').mean(dim = 0)\n\n    mean_heatmap_logits = torch.stack([\n        mean_heatmap_logits_0,\n        mean_heatmap_logits_1.to('cuda:0'),\n    ], dim = 0).mean(dim = 0)\n    \n    # mean_heatmap_logits = torch.stack([\n    #     image0_standard_heatmap,\n    #     image0_flip_heatmap.to('cuda:0'),\n    #     image1_standard_heatmap, \n    #     image1_flip_heatmap.to('cuda:0')\n    # ], dim = 0).to('cuda:0').mean(dim = 0)\n\n    # EXTRACT & POST-PROCESS SIGNALS.\n    predicted_signals = shared_signal_extractor(mean_heatmap_logits)\n\n    predicted_signals = interpolate_low_confidence_regions(predicted_signals).cpu()\n    predicted_signals = predicted_signals[0, :, 10]\n    predicted_signals = extract_signals_directly(predicted_signals)\n    \n    # RESAMPLE SIGNALS.\n    final_signals = {}\n    \n    for _, row in df_rows_for_id.iterrows():\n        lead_name = row.lead\n        num_rows = row.number_of_rows\n        \n        resampled_signal = resample_to_target(\n            predicted_signals[lead_name][\"time\"],\n            predicted_signals[lead_name][\"voltage\"],\n            target_fs, \n            num_rows\n        )\n        final_signals[lead_name] = resampled_signal\n\n    return final_signals","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:33:15.959941Z","iopub.execute_input":"2026-01-22T07:33:15.960292Z","iopub.status.idle":"2026-01-22T07:33:15.983327Z","shell.execute_reply.started":"2026-01-22T07:33:15.960266Z","shell.execute_reply":"2026-01-22T07:33:15.982804Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import scipy.signal\nimport scipy.optimize\n\ndef visualize_inference_output(final_signals, signal_csv_path, train_csv_path):\n    \"\"\"\n    Visualizes the dictionary output of process_single_image overlaid with the Ground Truth.\n    \n    Args:\n        final_signals (dict): Output from process_single_image {lead_name: np.array}.\n        signal_csv_path (str): Path to the ground truth CSV for this record.\n        train_csv_path (str): Path to the metadata CSV (to get fs).\n    \"\"\"\n    # 1. Load Metadata for Sampling Rate\n    record_id = Path(signal_csv_path).stem\n    df_meta = pd.read_csv(train_csv_path)\n    \n    # Handle ID types (int vs str)\n    if df_meta['id'].dtype == np.int64:\n        try:\n            row = df_meta[df_meta['id'] == int(record_id)]\n        except ValueError:\n            row = df_meta[df_meta['id'] == record_id]\n    else:\n        row = df_meta[df_meta['id'] == str(record_id)]\n\n    fs = float(row.iloc[0]['fs']) if not row.empty else 500.0\n\n    # 2. Load Ground Truth\n    try:\n        df_sig = pd.read_csv(signal_csv_path)\n    except Exception as e:\n        print(f\"Error loading GT CSV: {e}\")\n        return\n\n    # 3. Setup Plotting Grid (Standard 12-lead layout + Rhythm Strip)\n    fig = plt.figure(figsize=(24, 12))\n    gs = fig.add_gridspec(5, 4) \n    \n    # Standard 12-Lead Layout (3 rows x 4 columns)\n    # Note: (1,0) is usually Lead II in the grid view (short segment)\n    layout_map = {\n        \"I\": (0, 0), \"aVR\": (0, 1), \"V1\": (0, 2), \"V4\": (0, 3),\n        \"II\": (1, 0), \"aVL\": (1, 1), \"V2\": (1, 2), \"V5\": (1, 3), # \"II\" here acts as the subset\n        \"III\": (2, 0), \"aVF\": (2, 1), \"V3\": (2, 2), \"V6\": (2, 3),\n    }\n\n    snr_values = []\n    total_signal_power, total_noise_power = 0.0, 0.0\n\n    # --- Plot 12 Standard Leads (Grid) ---\n    for lead_name, (row, col) in layout_map.items():\n        ax = fig.add_subplot(gs[row, col])\n        \n        # Ground Truth Extraction\n        gt_data = None\n        gt_time = None\n        \n        # Determine CSV column name (CSV usually has \"II\" for the long strip)\n        csv_col = lead_name\n        \n        if csv_col in df_sig.columns:\n            raw_data = df_sig[csv_col].dropna().values\n            if len(raw_data) > 0:\n                time_axis = np.arange(len(raw_data)) / fs\n                \n                # If this is the \"Grid View\" for Lead II (position 1,0), crop to ~2.5s\n                if lead_name == \"II\" and row == 1:\n                    crop_idx = int(2.5 * fs)\n                    if len(raw_data) > crop_idx:\n                        raw_data = raw_data[:crop_idx]\n                        time_axis = time_axis[:crop_idx]\n                \n                ax.plot(time_axis, raw_data, color='black', alpha=0.5, linewidth=1.5, label='GT')\n                gt_data = raw_data\n                gt_time = time_axis\n\n        # Prediction Extraction\n        pred_data = None\n        if lead_name in final_signals:\n            pred_data = final_signals[lead_name]\n            \n            # Create time axis for prediction based on length and fs\n            pred_time = np.arange(len(pred_data)) / fs\n            \n            # If this is the \"Grid View\" for Lead II, crop the long prediction\n            if lead_name == \"II\" and row == 1:\n                crop_idx = int(2.5 * fs)\n                if len(pred_data) > crop_idx:\n                    pred_data = pred_data[:crop_idx]\n                    pred_time = pred_time[:crop_idx]\n\n            ax.plot(pred_time, pred_data, color='red', linewidth=1, label='Pred')\n\n        # SNR Calculation\n        snr_str = \"nan\"\n        if gt_data is not None and pred_data is not None:\n            # Truncate to the shorter length for comparison\n            min_len = min(len(gt_data), len(pred_data))\n            if min_len > 1:\n\n                aligned_pred_data = align_signals(gt_data, pred_data)\n                aligned_pred_data[~np.isfinite(aligned_pred_data)] = 0\n                \n                # Inline SNR calculation for portability\n                noise = gt_data[:min_len] - aligned_pred_data[:min_len]\n                signal_power = np.sum(gt_data[:min_len] ** 2)\n                noise_power = np.sum(noise ** 2)\n\n                if lead_name != 'II':\n                    total_signal_power += signal_power\n                    total_noise_power += noise_power\n                \n                if noise_power > 1e-9:\n                    snr = 10 * np.log10(signal_power / noise_power)\n                else:\n                    snr = float('inf')\n\n                if snr != float('inf'):\n                    snr_values.append(snr)\n                    snr_str = f\"{snr:.2f} dB\"\n                else:\n                    snr_str = \"Inf dB\"\n\n        ax.set_title(f\"{lead_name} (SNR: {snr_str})\")\n        ax.grid(True, linestyle=':', alpha=0.6)\n        if row != 2: ax.set_xticklabels([])\n\n    # --- Plot Rhythm Strip (Full Lead II) ---\n    ax_rhythm = fig.add_subplot(gs[3:, :])\n    \n    if \"II\" in df_sig.columns:\n        rhythm_gt = df_sig[\"II\"].dropna().values\n        rhythm_time = np.arange(len(rhythm_gt)) / fs\n        ax_rhythm.plot(rhythm_time, rhythm_gt, color='black', alpha=0.5, linewidth=1.5, label='Ground Truth')\n        \n        if \"II\" in final_signals:\n            rhythm_pred = final_signals[\"II\"]\n            # Ensure prediction matches GT length roughly for display\n            r_pred_time = np.arange(len(rhythm_pred)) / fs\n            ax_rhythm.plot(r_pred_time, rhythm_pred, color='red', linewidth=1, label='Prediction')\n            \n            # Calculate SNR for full strip\n            min_len = min(len(rhythm_gt), len(rhythm_pred))\n            if min_len > 1:\n                noise = rhythm_gt[:min_len] - rhythm_pred[:min_len]\n                signal_power = np.sum(rhythm_gt[:min_len] ** 2)\n                noise_power = np.sum(noise ** 2)\n                total_signal_power += signal_power\n                total_noise_power += noise_power\n                if noise_power > 1e-9:\n                    snr = 10 * np.log10(signal_power / noise_power)\n                    snr_values.append(snr)\n                    snr_str = f\"{snr:.2f} dB\"\n                else:\n                    snr_str = \"Inf dB\"\n                ax_rhythm.set_title(f\"Rhythm II (SNR: {snr_str})\")\n\n    ax_rhythm.legend(loc='upper right')\n    ax_rhythm.grid(True, linestyle=':', alpha=0.6)\n    ax_rhythm.set_xlabel(\"Time (s)\")\n\n    avg_snr = np.mean(snr_values) if snr_values else 0\n    overall_snr = 10 * np.log10(total_signal_power / total_noise_power)\n    plt.suptitle(f\"Inference Result vs GT (ID: {record_id} | Avg SNR: {overall_snr:.2f} dB)\", fontsize=16)\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2026-01-22T07:33:15.984232Z","iopub.execute_input":"2026-01-22T07:33:15.984409Z","iopub.status.idle":"2026-01-22T07:33:16.004090Z","shell.execute_reply.started":"2026-01-22T07:33:15.984394Z","shell.execute_reply":"2026-01-22T07:33:16.003355Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_inference_output(\n    final_signals=process_single_image(\n        standard_model_0=standard_segmentation_model_0,\n        standard_model_1=standard_segmentation_model_1,\n        flipped_model_0=flipped_segmentation_model_0,\n        flipped_model_1=flipped_segmentation_model_1,\n        keypoint_model_0=keypoint_detector_0, \n        keypoint_model_1=keypoint_detector_1,\n        image_path=f'/kaggle/input/physionet-ecg-image-digitization/train/{SAMPLE_IMAGE_ID}/{SAMPLE_IMAGE_ID}-0005.png', \n        target_fs=512, \n        df_rows_for_id=pd.DataFrame([\n            {'lead': 'I', 'number_of_rows': 1280},\n            {'lead': 'II', 'number_of_rows': 5120},\n            {'lead': 'III', 'number_of_rows': 1280},\n            {'lead': 'aVR', 'number_of_rows': 1280},\n            {'lead': 'aVL', 'number_of_rows': 1280},\n            {'lead': 'aVF', 'number_of_rows': 1280},\n            {'lead': 'V1', 'number_of_rows': 1280},\n            {'lead': 'V2', 'number_of_rows': 1280},\n            {'lead': 'V3', 'number_of_rows': 1280},\n            {'lead': 'V4', 'number_of_rows': 1280},\n            {'lead': 'V5', 'number_of_rows': 1280},\n            {'lead': 'V6', 'number_of_rows': 1280},\n        ])\n    ), \n    signal_csv_path=f'/kaggle/input/physionet-ecg-image-digitization/train/{SAMPLE_IMAGE_ID}/{SAMPLE_IMAGE_ID}.csv', \n    train_csv_path='/kaggle/input/physionet-ecg-image-digitization/train.csv'\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:33:16.004729Z","iopub.execute_input":"2026-01-22T07:33:16.004937Z","iopub.status.idle":"2026-01-22T07:33:29.236672Z","shell.execute_reply.started":"2026-01-22T07:33:16.004923Z","shell.execute_reply":"2026-01-22T07:33:29.235876Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nimport torch\nimport cv2\nfrom torchvision.transforms import v2\nfrom PIL import Image\nfrom tqdm import tqdm\n\n# --- Configuration ---\nOUTPUT_PATH = '/kaggle/working/james_submission.csv'\nTEST_PATH = '/kaggle/input/physionet-ecg-image-digitization/test.csv'\nIMAGES_DIR = '/kaggle/input/physionet-ecg-image-digitization/test'\n\n# --- Main Submission Loop ---\n\ntest_df = pd.read_csv(TEST_PATH)\n\n# Prepare Output File\nif os.path.exists(OUTPUT_PATH):\n    os.remove(OUTPUT_PATH)\n\n# Write Header\npd.DataFrame(columns=[\"id\", \"value\"]).to_csv(OUTPUT_PATH, index=False)\n\ncached_id = None\ncached_predictions = {}\n\n# Filter rows to process by ID groups to minimize file I/O\n# We group by ID so we process each image exactly once\ngrouped = test_df.groupby('id', sort=False)\n\nprint(f\"Starting inference on {len(grouped)} images...\")\n\nfor current_id, group_rows in tqdm(grouped):\n    \n    # 1. Get FS from the first row of the group (assuming const per file)\n    fs = group_rows.iloc[0].fs\n    \n    # 2. Process the Image (Heavy Lifting)\n    # Pass the group_rows so we know exactly how many samples each lead needs\n    image_path = f\"{IMAGES_DIR}/{current_id}.png\"\n    predictions_map = process_single_image(\n        standard_model_0=standard_segmentation_model_0,\n        standard_model_1=standard_segmentation_model_1,\n        flipped_model_0=flipped_segmentation_model_0,\n        flipped_model_1=flipped_segmentation_model_1,\n        keypoint_model_0=keypoint_detector_0, \n        keypoint_model_1=keypoint_detector_1,\n        image_path=image_path, \n        target_fs=fs, \n        df_rows_for_id=group_rows\n    )\n    \n    # 3. Write to CSV\n    chunks = []\n    \n    for _, row in group_rows.iterrows():\n        lead_name = row.lead\n        \n        # Get data (or zeros if missing)\n        if lead_name in predictions_map:\n            lead_data = predictions_map[lead_name]\n        else:\n            lead_data = np.zeros(row.number_of_rows)\n            \n        # Safety Length Check\n        if len(lead_data) != row.number_of_rows:\n            # Force resize if rounding errors occurred\n            lead_data = np.interp(\n                np.linspace(0, 1, row.number_of_rows),\n                np.linspace(0, 1, len(lead_data)),\n                lead_data\n            )\n\n        # Create ID strings\n        # Format: {file_id}_{sample_index}_{lead_name}\n        # Vectorized string creation is faster than list comprehension loop\n        indices = np.arange(row.number_of_rows)\n        id_strs = [f\"{current_id}_{i}_{lead_name}\" for i in indices]\n        \n        # Append to chunks\n        chunk_df = pd.DataFrame({\"id\": id_strs, \"value\": lead_data})\n        chunks.append(chunk_df)\n    \n    # Flush this image's data to disk\n    if chunks:\n        pd.concat(chunks).to_csv(OUTPUT_PATH, mode='a', index=False, header=False)\n\nprint(\"Submission generated successfully.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:33:29.237642Z","iopub.execute_input":"2026-01-22T07:33:29.237836Z","iopub.status.idle":"2026-01-22T07:33:52.401584Z","shell.execute_reply.started":"2026-01-22T07:33:29.237821Z","shell.execute_reply":"2026-01-22T07:33:52.400783Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Free memory","metadata":{}},{"cell_type":"code","source":"# del keypoint_detector_0\n# del keypoint_detector_1\n# del standard_segmentation_model_0\n# del standard_segmentation_model_1\n# del flipped_segmentation_model_0\n# del flipped_segmentation_model_1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:33:52.402376Z","iopub.execute_input":"2026-01-22T07:33:52.402642Z","iopub.status.idle":"2026-01-22T07:33:52.406045Z","shell.execute_reply.started":"2026-01-22T07:33:52.402625Z","shell.execute_reply":"2026-01-22T07:33:52.405271Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import gc\ngc.collect()\nimport torch\ntorch.cuda.empty_cache()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:33:52.406844Z","iopub.execute_input":"2026-01-22T07:33:52.407150Z","iopub.status.idle":"2026-01-22T07:33:53.264914Z","shell.execute_reply.started":"2026-01-22T07:33:52.407129Z","shell.execute_reply":"2026-01-22T07:33:53.264347Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%reset -f","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:33:53.265721Z","iopub.execute_input":"2026-01-22T07:33:53.265958Z","iopub.status.idle":"2026-01-22T07:33:54.390634Z","shell.execute_reply.started":"2026-01-22T07:33:53.265930Z","shell.execute_reply":"2026-01-22T07:33:54.390033Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Imanishi parts","metadata":{}},{"cell_type":"code","source":"!pip install --find-links /kaggle/input/phisio-wheels --no-index connected-components-3d segmentation-models","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:33:54.391854Z","iopub.execute_input":"2026-01-22T07:33:54.392070Z","iopub.status.idle":"2026-01-22T07:33:57.892969Z","shell.execute_reply.started":"2026-01-22T07:33:54.392053Z","shell.execute_reply":"2026-01-22T07:33:57.892231Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile run.py\n\nimport os\nimport pandas as pd\nimport numpy as np\nimport cv2\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nimport random\nimport tensorflow as tf\nfrom tensorflow.keras.models import Model\nfrom sklearn.model_selection import KFold\nimport shutil\nimport segmentation_models as sm\nimport torch\nimport torch.nn.functional as F\nfrom scipy.signal import savgol_filter\nimport argparse\n\nimport sys\nsys.path.append('/kaggle/input/hengck23-demo-submit-physionet')\nsys.path.append('/kaggle/input/physio-src')\n\nfrom model_v2 import build_model\nfrom utils_v2 import run_my_stage0, rectify_image_v2, clustering_y_v2, mask2series3, einth_correction1\nfrom stage1_model import Net as Net1\nimport stage1_common\n\n\nparser = argparse.ArgumentParser()\nparser.add_argument('--half', type=int)\nparser.add_argument('--debug', action='store_true')\nargs = parser.parse_args()\n\n\ngpus = tf.config.list_physical_devices('GPU')\nfor gpu in gpus:\n    tf.config.experimental.set_memory_growth(gpu, True)\n\n\n\nclass Config:\n    TRAIN_DIR = '/kaggle/input/physionet-ecg-image-digitization/train'\n    TEST_DIR = '/kaggle/input/physionet-ecg-image-digitization/test'\n    SEGMENTS = ['0001', '0003', '0004', '0005', '0006', '0009', '0010', '0011', '0012']\n    \n    DATA_DIR = '/kaggle/input/physionet-ecg-image-digitization'\n    SIGNAL_GROUPS = [\n        ['I', 'aVR', 'V1', 'V4'],\n        ['II_quarter', 'aVL', 'V2', 'V5'],\n        ['III', 'aVF', 'V3', 'V6'],\n    ]\n    TMP_DIR = '/kaggle_tmp'\n    \n    HEIGHT1, WIDTH1 = 1152, 1440   # aligned の座標系（gridpoint_xy が参照している系）\n\n    GRID_SIZE = 40\n    X_SCALE = 2.0\n    \n    HEIGHT2, WIDTH2 = 43*GRID_SIZE, int(56*GRID_SIZE*X_SCALE)\n    TOP_CUT = 2\n    \n    ZERO_YS = [17.8, 25, 32.2, 38.75]\n    START_X = 3.\n    X_RANGE = 50\n    \n    RECT_SIZE = (HEIGHT2 - TOP_CUT * GRID_SIZE, WIDTH2)\n    STAGE2_TOP_CROP = RECT_SIZE[0]%32\n    STAGE2_RIGHT_CROP = RECT_SIZE[1]%32\n    STAGE2_INSIZE = (RECT_SIZE[0] - STAGE2_TOP_CROP, RECT_SIZE[1] - STAGE2_RIGHT_CROP)\n    \n    \n    ##### Stage0 #####\n    STAGE0_RESIZE = 1024\n    N_KPTS = 9\n    STAGE0_BB_NAME = 'effv2b1'\n    STAGE0_WEIGHT_PATH = '/kaggle/input/physio-work/keypoint_effv2b1_1024_fold0_run2_dice0.9638.h5'\n    STAGE0_AREA_MIN = 50\n    \n    ##### Stage2 #####\n    STAGE2_BB_NAME = 'effv2b1'\n    STAGE2_WEIGHT_PATH = [\n        '/kaggle/input/physio-work/effv2b1_1632_1280_coordloss2_2xdata_25ep_fold0_run4_mse2.19.h5',\n        '/kaggle/input/physio-work/effv2b1_1632_1280_coordloss2_3xdata_15ep_foldfull_run1.h5',\n    ]\n    STAGE2_TTA = True\n    STAGE2_MIN_MASK_SCORE = 0.2\n    STAGE2_MIN_SERIES_SCORE = 0.02\n    \n    USE_SEG0001_PREPROC = True\n    USE_V2_ENSEMBLE = True\n    USE_SAVGOL = True\n    USE_EINTHOVEN = True\n    USE_Y_CLUSTERING = True\n    USE_2HEAD_ENSEMBLE = False\n    \n\nclass EnsembleModel2Head(Model):\n    def __init__(self, models):\n        super().__init__()\n        self.models = models\n        \n    def call(self, inputs, training=None):\n        preds = []\n        for model in self.models:\n            preds.append(model(inputs))\n\n        ens_preds = {}\n        for key in preds[0].keys():\n            stack = []\n            for p in preds:\n                stack.append(p[key])\n            stack = tf.stack(stack, axis=0)\n            ens_preds[key] = tf.reduce_mean(stack, axis=0)\n        return ens_preds\n        \n\ncfg = Config()\nos.makedirs(cfg.TMP_DIR, exist_ok=True)\n\ndummy = sm.Unet(\n    'resnet18',\n    input_shape=(224, 224, 3),\n    classes=3,\n    activation=None,\n    encoder_weights=None,\n)\ndel dummy\n\n\ndef run_stage1_gridout(normalised):\n    batch = {\n        'image': torch.from_numpy(np.ascontiguousarray(normalised.transpose(2, 0, 1))).unsqueeze(0),\n    }\n    \n    with torch.amp.autocast('cuda', dtype=torch.float16):\n        with torch.no_grad():\n            output = net1(batch)\n    \n    gridpoint_xy, more = stage1_common.output_to_predict(normalised, batch, output)\n    # rectified = stage1_common.rectify_image(normalised, gridpoint_xy)\n    return gridpoint_xy\n\n\n\nif args.debug:\n    data_df = pd.read_csv('/kaggle/input/physionet-ecg-image-digitization/train.csv')\n    data_df['fold'] = 0\n    \n    kf = KFold(n_splits=5, random_state=12, shuffle=True)\n    for i, (_, test_index) in enumerate(kf.split(data_df)):\n        data_df.loc[test_index, 'fold'] = i\n    \n    fold = 0\n    data_df = data_df[data_df['fold']==fold].sample(50, random_state=0)\n    \n    fake_test_df = []\n    for i, d in tqdm(data_df.iterrows(), total=len(data_df)):\n        image_id = d['id']\n    \n        truth_df = pd.read_csv(f'{cfg.DATA_DIR}/train/{image_id}/{image_id}.csv')\n        non_nan_count = truth_df.count()\n        \n        this_df = pd.DataFrame({\n            'id':image_id ,\n            'lead':non_nan_count.index,\n            'fs': d['fs'],\n            'number_of_rows':non_nan_count.values \n        })\n        fake_test_df.append(this_df)\n    test_df = pd.concat(fake_test_df)\nelse:\n    test_df = pd.read_csv('/kaggle/input/physionet-ecg-image-digitization/test.csv')\n\n\nhalf_size = len(test_df) // 2\nif args.half == 0:\n    test_df = test_df.iloc[:half_size]\nelif args.half == 1:\n    test_df = test_df.iloc[half_size:]\n\n\n\n##### Stage0-1 #####\nmodel0 = build_model(cfg, stage=0)\nmodel0.load_weights(cfg.STAGE0_WEIGHT_PATH)\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nnet1 = Net1(pretrained=False)\nnet1 = stage1_common.load_net(net1, '/kaggle/input/hengck23-demo-submit-physionet/weight/stage1-last.checkpoint.pth')\nnet1.to(DEVICE)\nprint('model load done')\n\n# pytorchに先にGPUメモリを確保させる\nwith torch.amp.autocast('cuda', dtype=torch.float16):\n    with torch.no_grad():\n        _ = net1({'image': torch.zeros([2, 3, cfg.HEIGHT1, cfg.WIDTH1])})\n\n\ngrid_save_dir = f'{cfg.TMP_DIR}/grid'\nhomo_save_dir = f'{cfg.TMP_DIR}/homo'\nos.makedirs(grid_save_dir, exist_ok=True)\nos.makedirs(homo_save_dir, exist_ok=True)\nrotk_records = []\n\nfor i, sample_id in enumerate(tqdm(test_df['id'].unique())):\n    if args.debug:\n        seg = cfg.SEGMENTS[int(str(sample_id)[-2:]) % len(cfg.SEGMENTS)]\n        png_path = f'{cfg.TRAIN_DIR}/{sample_id}/{sample_id}-{seg}.png'\n    else:\n        png_path = f'{cfg.TEST_DIR}/{sample_id}.png'\n    sig_len = test_df[test_df['id']==sample_id]['fs'].iat[0] * 10\n    \n    image, rotated, normalised, homo, rotk = run_my_stage0(png_path, model0, cfg)\n    gridpoint_xy = run_stage1_gridout(normalised)\n\n    np.save(f'{grid_save_dir}/{sample_id}_grid', gridpoint_xy.astype(np.float32))\n    np.save(f'{homo_save_dir}/{sample_id}_homo', homo.astype(np.float32))\n    rotk_records.append({'id': sample_id, 'rotk': int(rotk)})\nrotk_df = pd.DataFrame(rotk_records)\ndel model0, net1\ntorch.cuda.empty_cache()\n\n\n##### Stage2 #####\n\nif len(cfg.STAGE2_WEIGHT_PATH) > 1:\n    models = []\n    for path in cfg.STAGE2_WEIGHT_PATH:\n        model = build_model(cfg, stage=2)\n        model.load_weights(path)\n        models.append(model)\n    model2 = EnsembleModel2Head(models)\nelse:\n    model2 = build_model(cfg, stage=2)\n    model2.load_weights(cfg.STAGE2_WEIGHT_PATH[0])\n\nseries_save_dir = f'{cfg.TMP_DIR}/series'\nos.makedirs(series_save_dir, exist_ok=True)\n\nsub_df = []\nfor sample_id in tqdm(test_df['id'].unique()):\n    if args.debug:\n        seg = cfg.SEGMENTS[int(str(sample_id)[-2:]) % len(cfg.SEGMENTS)]\n        png_path = f'{cfg.TRAIN_DIR}/{sample_id}/{sample_id}-{seg}.png'\n    else:\n        png_path = f'{cfg.TEST_DIR}/{sample_id}.png'\n    sig_len = test_df[test_df['id']==sample_id]['fs'].iat[0] * 10\n\n    image = cv2.imread(png_path, cv2.IMREAD_COLOR_RGB)\n    gridpoint_xy = np.load(f'{grid_save_dir}/{sample_id}_grid.npy')\n    homo = np.load(f'{homo_save_dir}/{sample_id}_homo.npy')\n    rotk = rotk_df[rotk_df['id']==sample_id]['rotk'].iat[0]\n    \n    rotated = np.ascontiguousarray(np.rot90(image, rotk, axes=(0, 1)))\n    rectified = rectify_image_v2(rotated, homo, gridpoint_xy, cfg)\n    rectified = rectified[cfg.STAGE2_TOP_CROP:, :cfg.STAGE2_INSIZE[1], :]\n    \n    if cfg.STAGE2_TTA:\n        pred = model2(np.stack([rectified, rectified[:, ::-1, :]], axis=0))\n        mask_pred = pred['mask'].numpy()\n        series_pred = pred['series'].numpy()\n        mask_pred = (mask_pred[0, :, :, :] + mask_pred[1, :, ::-1, :]) / 2\n        series_pred = (series_pred[0, :, :, :] + series_pred[1, :, ::-1, :]) / 2\n    else:\n        pred = model2(rectified[None, :, :, :])\n        mask_pred = pred['mask'].numpy()[0]\n        series_pred = pred['series'].numpy()[0]\n        \n    series_pred = series_pred * (series_pred > cfg.STAGE2_MIN_SERIES_SCORE)\n    mask_pred = mask_pred * (mask_pred > cfg.STAGE2_MIN_MASK_SCORE)\n\n    # y-clustering and apply y-mask\n    y_mask_mask = clustering_y_v2(mask_pred)\n    \n    series_pred = series_pred * y_mask_mask[:, None, :]\n    mask_pred = mask_pred * y_mask_mask[:, None, :]\n\n    # x-mask (replace with np.nan)\n    null_pos_series = (series_pred.max(axis=0)==0)\n    null_pos_mask = (mask_pred.max(axis=0)==0)\n\n    # two x-mask OR merge\n    null_pos_series += null_pos_mask\n\n    # argmax and coordinate => mv\n    series_pred_mv = mask2series3(series_pred, sig_len, null_pos_series, cfg)\n    \n    # mask and series ensemble\n    if cfg.USE_2HEAD_ENSEMBLE:\n        mask_pred_mv = mask2series3(mask_pred, sig_len, null_pos_mask, cfg)\n        series_pred_mv = (series_pred_mv + mask_pred_mv) / 2\n\n    if cfg.USE_V2_ENSEMBLE:\n        nrow = test_df[(test_df['id']==sample_id) & (test_df['lead']=='I')]['number_of_rows'].iat[0]\n        data2a = series_pred_mv[:nrow, 1]\n        data2b = series_pred_mv[:nrow, 3]\n        series_pred_mv[:nrow, 3] = (data2a + data2b) / 2\n\n    if cfg.USE_SAVGOL:\n        if sig_len < 9000:\n            window_length = 5\n        else:\n            window_length = 9\n            \n        for j in range(4):\n            series_pred_mv[:, j] = savgol_filter(series_pred_mv[:, j], window_length=window_length, polyorder=2)\n\n    if cfg.USE_EINTHOVEN:\n        series = einth_correction1(series_pred_mv, sample_id, test_df)\n    \n    np.save(f'{series_save_dir}/{sample_id}', series)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:33:57.894187Z","iopub.execute_input":"2026-01-22T07:33:57.894511Z","iopub.status.idle":"2026-01-22T07:33:57.906392Z","shell.execute_reply.started":"2026-01-22T07:33:57.894460Z","shell.execute_reply":"2026-01-22T07:33:57.905766Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nos.environ['TF_USE_LEGACY_KERAS'] = '1'\nos.environ['SM_FRAMEWORK'] = 'tf.keras'\n\nimport sys\nsys.path.append('/kaggle/input/hengck23-demo-submit-physionet')\nsys.path.append('/kaggle/input/physio-src')\n\nimport subprocess\nimport shutil\nimport pandas as pd\nimport numpy as np\nfrom sklearn.model_selection import KFold\nfrom tqdm import tqdm\n\n# from metric import score","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:33:57.907187Z","iopub.execute_input":"2026-01-22T07:33:57.907442Z","iopub.status.idle":"2026-01-22T07:33:58.272820Z","shell.execute_reply.started":"2026-01-22T07:33:57.907426Z","shell.execute_reply":"2026-01-22T07:33:58.272045Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DEBUG = False\n\n\nenv1 = os.environ.copy()\nenv2 = os.environ.copy()\nenv1['CUDA_VISIBLE_DEVICES'] = '0'\nenv2['CUDA_VISIBLE_DEVICES'] = '1'\n\ncmd1 = f'python run.py --half 0'\nif DEBUG:\n    cmd1 += ' --debug'\nproc1 = subprocess.Popen(cmd1.split(' '), env=env1)\n\ncmd2 = f'python run.py --half 1'\nif DEBUG:\n    cmd2 += ' --debug'\nproc2 = subprocess.Popen(cmd2.split(' '), env=env2)\n\n_ = proc1.communicate()\n_ = proc2.communicate()\n\n\n# estimate time for 1000 images: 72 minutes","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:33:58.273688Z","iopub.execute_input":"2026-01-22T07:33:58.273960Z","iopub.status.idle":"2026-01-22T07:35:08.154666Z","shell.execute_reply.started":"2026-01-22T07:33:58.273937Z","shell.execute_reply":"2026-01-22T07:35:08.154031Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def make_pred_df(sample_id, series, test_df):\n    result_dict = {}\n    for plot_idx in range(3):\n        ofset = 0\n        for signal in SIGNAL_GROUPS[plot_idx]:\n            if signal == 'II_quarter':\n                nrow = test_df[(test_df['id']==sample_id) & (test_df['lead']=='I')]['number_of_rows'].iat[0]\n                ofset += nrow\n            else:\n                nrow = test_df[(test_df['id']==sample_id) & (test_df['lead']==signal)]['number_of_rows'].iat[0]\n                result_dict[signal] = series[plot_idx][ofset:ofset + nrow]\n                ofset += nrow\n        result_dict['II'] = series[3]\n    \n    row_ids = []\n    concat_values = []\n    for signal, values in result_dict.items():\n        concat_values.append(values)\n        row_ids += [f'{sample_id}_{i}_{signal}' for i in range(len(values))]\n    concat_values = np.concatenate(concat_values)\n    \n    pred_df = pd.DataFrame({'id': row_ids, 'value': concat_values})\n    return pred_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:35:08.160407Z","iopub.execute_input":"2026-01-22T07:35:08.160701Z","iopub.status.idle":"2026-01-22T07:35:08.166834Z","shell.execute_reply.started":"2026-01-22T07:35:08.160682Z","shell.execute_reply":"2026-01-22T07:35:08.166100Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_DIR = '/kaggle/input/physionet-ecg-image-digitization'\nTMP_DIR = '/kaggle_tmp'\nSIGNAL_GROUPS = [\n    ['I', 'aVR', 'V1', 'V4'],\n    ['II_quarter', 'aVL', 'V2', 'V5'],\n    ['III', 'aVF', 'V3', 'V6'],\n]\n\nif DEBUG:\n    data_df = pd.read_csv('/kaggle/input/physionet-ecg-image-digitization/train.csv')\n    data_df['fold'] = 0\n    \n    kf = KFold(n_splits=5, random_state=12, shuffle=True)\n    for i, (_, test_index) in enumerate(kf.split(data_df)):\n        data_df.loc[test_index, 'fold'] = i\n    \n    fold = 0\n    data_df = data_df[data_df['fold']==fold].sample(50, random_state=0)\n    \n    fake_test_df = []\n    gt_df = []\n    for i, d in tqdm(data_df.iterrows(), total=len(data_df)):\n        image_id = d['id']\n    \n        truth_df = pd.read_csv(f'{DATA_DIR}/train/{image_id}/{image_id}.csv')\n        non_nan_count = truth_df.count()\n    \n        #lead\tfs\tnumber_of_rows \n        this_df = pd.DataFrame({\n            'id':image_id ,\n            'lead':non_nan_count.index,\n            'fs': d['fs'],\n            'number_of_rows':non_nan_count.values \n        })\n        fake_test_df.append(this_df)\n    \n        values_with_nan = truth_df.values.T.ravel()\n        values = values_with_nan[~np.isnan(values_with_nan)]\n        row_ids = [\n            f\"{image_id}_{i}_{signal}\"\n            for signal in truth_df.columns\n            for i in range(non_nan_count[signal])\n        ]\n        this_gt_df = pd.DataFrame({'id': row_ids, 'value': values})\n        this_gt_df['fs'] = d['fs']\n        gt_df.append(this_gt_df)\n    test_df = pd.concat(fake_test_df)\n    gt_df = pd.concat(gt_df)\n    \n    gt_df['lead'] = gt_df['id'].apply(lambda x: int(x.split('_')[0]))\n    \n    gt_df['sample_id'] = gt_df['id'].apply(lambda x: int(x.split('_')[0]))\n    gt_df = gt_df.reset_index(drop=True)\nelse:\n    test_df = pd.read_csv('/kaggle/input/physionet-ecg-image-digitization/test.csv')\n\n\nsub_df = []\nfor i, sample_id in enumerate(test_df['id'].unique()):\n    sig_len = test_df[test_df['id']==sample_id]['fs'].iat[0] * 10\n    series = np.load(f'{TMP_DIR}/series/{sample_id}.npy')\n\n    pred_df = make_pred_df(sample_id, series.T, test_df)\n    sub_df.append(pred_df)\n\n    if DEBUG:\n        cv_score = score(gt_df[gt_df['sample_id']==sample_id], pred_df, row_id_column_name='id')\n        print(f'{sample_id}, score: {cv_score:.4f}')\n\nsub_df = pd.concat(sub_df)\n\nif DEBUG:\n    cv_score = score(gt_df, sub_df, row_id_column_name='id')\n    print(f'cv_score: {cv_score:.4f}')\n\n# cv_score: 23.7419\n# val200 cv_score: 23.9266\n\n# Ensemble 2 model\n# cv_score: 24.0984","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:35:08.167725Z","iopub.execute_input":"2026-01-22T07:35:08.168006Z","iopub.status.idle":"2026-01-22T07:35:08.242844Z","shell.execute_reply.started":"2026-01-22T07:35:08.167991Z","shell.execute_reply":"2026-01-22T07:35:08.242273Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub_df.to_csv('imanishi_submission.csv',index=False)\n\nsub_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:35:08.243583Z","iopub.execute_input":"2026-01-22T07:35:08.243832Z","iopub.status.idle":"2026-01-22T07:35:08.397022Z","shell.execute_reply.started":"2026-01-22T07:35:08.243808Z","shell.execute_reply":"2026-01-22T07:35:08.396427Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Ensemble","metadata":{}},{"cell_type":"code","source":"!ls -lh","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:35:08.397683Z","iopub.execute_input":"2026-01-22T07:35:08.398003Z","iopub.status.idle":"2026-01-22T07:35:08.600325Z","shell.execute_reply.started":"2026-01-22T07:35:08.397983Z","shell.execute_reply":"2026-01-22T07:35:08.599076Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!head -n 10 liu_submission.csv","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:35:08.601814Z","iopub.execute_input":"2026-01-22T07:35:08.602178Z","iopub.status.idle":"2026-01-22T07:35:08.797623Z","shell.execute_reply.started":"2026-01-22T07:35:08.602135Z","shell.execute_reply":"2026-01-22T07:35:08.796770Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!head -n 10 james_submission.csv","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:35:08.798991Z","iopub.execute_input":"2026-01-22T07:35:08.799346Z","iopub.status.idle":"2026-01-22T07:35:08.992160Z","shell.execute_reply.started":"2026-01-22T07:35:08.799308Z","shell.execute_reply":"2026-01-22T07:35:08.991217Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!head -n 10 imanishi_submission.csv","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:35:08.993466Z","iopub.execute_input":"2026-01-22T07:35:08.993953Z","iopub.status.idle":"2026-01-22T07:35:09.188784Z","shell.execute_reply.started":"2026-01-22T07:35:08.993914Z","shell.execute_reply":"2026-01-22T07:35:09.187984Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\n\ndef align_signals(label: np.ndarray, pred: np.ndarray, max_shift: float = float('inf')) -> np.ndarray:\n    # Initialize the reference and digitized signals.\n    label_arr = np.asarray(label, dtype=np.float64)\n    pred_arr = np.asarray(pred, dtype=np.float64)\n\n    label_mean = np.mean(label_arr)\n    pred_mean = np.mean(pred_arr)\n\n    label_arr_centered = label_arr - label_mean\n    pred_arr_centered = pred_arr - pred_mean\n\n    # Compute the correlation between the reference and digitized signals and locate the maximum correlation.\n    correlation = scipy.signal.correlate(label_arr_centered, pred_arr_centered, mode='full')\n\n    n_label = np.size(label_arr)\n    n_pred = np.size(pred_arr)\n\n    lags = scipy.signal.correlation_lags(n_label, n_pred, mode='full')\n    valid_lags_mask = (lags >= -max_shift) & (lags <= max_shift)\n\n    max_correlation = np.nanmax(correlation[valid_lags_mask])\n    all_max_indices = np.flatnonzero(correlation == max_correlation)\n    best_idx = min(all_max_indices, key=lambda i: abs(lags[i]))\n    time_shift = lags[best_idx]\n    start_padding_len = max(time_shift, 0)\n    pred_slice_start = max(-time_shift, 0)\n    pred_slice_end = min(n_label - time_shift, n_pred)\n    end_padding_len = max(n_label - n_pred - time_shift, 0)\n    aligned_pred = np.concatenate((np.full(start_padding_len, np.nan), pred_arr[pred_slice_start:pred_slice_end], np.full(end_padding_len, np.nan)))\n\n    def objective_func(v_shift):\n        return np.nansum((label_arr - (aligned_pred - v_shift)) ** 2)\n\n    if np.any(np.isfinite(label_arr) & np.isfinite(aligned_pred)):\n        results = scipy.optimize.minimize_scalar(objective_func, method='Brent')\n        vertical_shift = results.x\n        aligned_pred -= vertical_shift\n    return aligned_pred","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:35:09.189968Z","iopub.execute_input":"2026-01-22T07:35:09.190260Z","iopub.status.idle":"2026-01-22T07:35:09.201438Z","shell.execute_reply.started":"2026-01-22T07:35:09.190225Z","shell.execute_reply":"2026-01-22T07:35:09.200708Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import re\nfrom dataclasses import dataclass\nfrom typing import Dict, Iterable, Tuple\n\nimport numpy as np\nimport pandas as pd\n\nimport scipy.optimize\nimport scipy.signal\n\n\ndef _parse_submission_id(submission_id: str) -> tuple[str, int, str, str]:\n    \"\"\"\n    Parse an id like:\n      1053922973_0_I\n      <record_prefix>_<timestep>_<lead>\n\n    Returns:\n      signal_key: \"<record_prefix>_<lead>\"   (e.g., \"1053922973_I\")\n      timestep: int                          (e.g., 0)\n      record_prefix: str                     (e.g., \"1053922973\")\n      lead: str                              (e.g., \"I\")\n    \"\"\"\n    pieces = submission_id.split(\"_\")\n    if len(pieces) < 3:\n        raise ValueError(f\"Unexpected id format: {submission_id!r}\")\n\n    lead = pieces[-1]\n    timestep_str = pieces[-2]\n    record_prefix = \"_\".join(pieces[:-2])\n\n    try:\n        timestep = int(timestep_str)\n    except ValueError as exc:\n        raise ValueError(f\"Unexpected timestep in id: {submission_id!r}\") from exc\n\n    signal_key = f\"{record_prefix}_{lead}\"\n    return signal_key, timestep, record_prefix, lead\n\ndef blend_submissions_by_aligned_signal(\n    submission_1_df: pd.DataFrame,\n    submission_2_df: pd.DataFrame,\n    weight_submission_2: float = 0.5,\n    max_shift: float = float(\"inf\"),\n) -> pd.DataFrame:\n    \"\"\"\n    Takes two submission dataframes with columns: [\"id\", \"value\"].\n\n    For each unique signal (record_prefix + lead):\n      1) Align submission_2 signal to submission_1 signal via cross-correlation + vertical shift.\n      2) Where aligned submission_2 overlaps (finite), output a weighted average:\n            out = (1-w) * sub1 + w * aligned_sub2\n         Where it does not overlap (NaN), output raw submission_1.\n\n    Output preserves the original row order of submission_1_df.\n    \"\"\"\n    if not (0.0 <= weight_submission_2 <= 1.0):\n        raise ValueError(f\"weight_submission_2 must be in [0, 1], got {weight_submission_2}\")\n\n    required_columns = {\"id\", \"value\"}\n    if not required_columns.issubset(submission_1_df.columns):\n        raise ValueError(f\"submission_1_df must contain {required_columns}, got {set(submission_1_df.columns)}\")\n    if not required_columns.issubset(submission_2_df.columns):\n        raise ValueError(f\"submission_2_df must contain {required_columns}, got {set(submission_2_df.columns)}\")\n\n    submission_1_work_df = submission_1_df[[\"id\", \"value\"]].copy()\n    submission_2_work_df = submission_2_df[[\"id\", \"value\"]].copy()\n\n    # Parse ids for submission 1 (also defines the output indexing/order).\n    parsed_1 = submission_1_work_df[\"id\"].map(_parse_submission_id)\n    submission_1_work_df[\"signal_key\"] = parsed_1.map(lambda t: t[0])\n    submission_1_work_df[\"timestep\"] = parsed_1.map(lambda t: t[1])\n\n    # Parse ids for submission 2 (used for lookup).\n    parsed_2 = submission_2_work_df[\"id\"].map(_parse_submission_id)\n    submission_2_work_df[\"signal_key\"] = parsed_2.map(lambda t: t[0])\n    submission_2_work_df[\"timestep\"] = parsed_2.map(lambda t: t[1])\n\n    # Build per-signal arrays for submission 2 for fast access.\n    # (Assumes at most one value per (signal_key, timestep).)\n    submission_2_grouped: Dict[str, pd.DataFrame] = {\n        signal_key: group_df for signal_key, group_df in submission_2_work_df.groupby(\"signal_key\", sort=False)\n    }\n\n    blended_values = submission_1_work_df[\"value\"].to_numpy(dtype=np.float64, copy=True)\n\n    # Iterate signals as defined by submission 1.\n    for signal_key, sub1_group_df in submission_1_work_df.groupby(\"signal_key\", sort=False):\n        sub1_group_sorted_df = sub1_group_df.sort_values(\"timestep\", kind=\"mergesort\")\n        sub1_row_indices = sub1_group_sorted_df.index.to_numpy()\n\n        sub1_values = submission_1_work_df.loc[sub1_row_indices, \"value\"].to_numpy(dtype=np.float64, copy=False)\n\n        sub2_group_df = submission_2_grouped.get(signal_key)\n        if sub2_group_df is None:\n            # No matching signal in submission 2: keep submission 1 as-is.\n            continue\n\n        sub2_group_sorted_df = sub2_group_df.sort_values(\"timestep\", kind=\"mergesort\")\n        sub2_values = sub2_group_sorted_df[\"value\"].to_numpy(dtype=np.float64, copy=False)\n\n        aligned_sub2_values = align_signals(sub1_values, sub2_values, max_shift=max_shift)\n\n        overlap_mask = np.isfinite(aligned_sub2_values)\n        if not np.any(overlap_mask):\n            continue\n\n        w = float(weight_submission_2)\n        blended_signal_values = sub1_values.copy()\n        blended_signal_values[overlap_mask] = (\n            (1.0 - w) * sub1_values[overlap_mask] + w * aligned_sub2_values[overlap_mask]\n        )\n\n        blended_values[sub1_row_indices] = blended_signal_values\n\n    out_df = submission_1_df.copy()\n    out_df[\"value\"] = blended_values\n    return out_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:35:09.202283Z","iopub.execute_input":"2026-01-22T07:35:09.202616Z","iopub.status.idle":"2026-01-22T07:35:09.221364Z","shell.execute_reply.started":"2026-01-22T07:35:09.202592Z","shell.execute_reply":"2026-01-22T07:35:09.220723Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"liu_df = pd.read_csv(\"liu_submission.csv\")\njames_df = pd.read_csv(\"james_submission.csv\")\nimanishi_df = pd.read_csv(\"imanishi_submission.csv\")\n\nblended_df = blend_submissions_by_aligned_signal(\n    submission_1_df=liu_df,\n    submission_2_df=james_df,\n    weight_submission_2=0.58182\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:35:09.222151Z","iopub.execute_input":"2026-01-22T07:35:09.222723Z","iopub.status.idle":"2026-01-22T07:35:09.711827Z","shell.execute_reply.started":"2026-01-22T07:35:09.222706Z","shell.execute_reply":"2026-01-22T07:35:09.711026Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"blended_df = blend_submissions_by_aligned_signal(\n    submission_1_df=blended_df,\n    submission_2_df=imanishi_df,\n    weight_submission_2=0.45\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:35:09.712809Z","iopub.execute_input":"2026-01-22T07:35:09.713612Z","iopub.status.idle":"2026-01-22T07:35:10.027895Z","shell.execute_reply.started":"2026-01-22T07:35:09.713584Z","shell.execute_reply":"2026-01-22T07:35:10.027306Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"blended_df.to_csv('submission.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:35:10.028619Z","iopub.execute_input":"2026-01-22T07:35:10.028815Z","iopub.status.idle":"2026-01-22T07:35:10.208400Z","shell.execute_reply.started":"2026-01-22T07:35:10.028793Z","shell.execute_reply":"2026-01-22T07:35:10.207417Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Sanity check","metadata":{}},{"cell_type":"code","source":"!echo \"Line counts\"\n!cat liu_submission.csv | wc -l\n!cat james_submission.csv | wc -l\n!cat imanishi_submission.csv | wc -l\n!cat submission.csv | wc -l\n\n!echo\n!echo \"File prefixes\"\n!head -n 5 liu_submission.csv\n!head -n 5 james_submission.csv\n!head -n 5 imanishi_submission.csv\n!head -n 5 submission.csv","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:35:10.209365Z","iopub.execute_input":"2026-01-22T07:35:10.209599Z","iopub.status.idle":"2026-01-22T07:35:12.265762Z","shell.execute_reply.started":"2026-01-22T07:35:10.209582Z","shell.execute_reply":"2026-01-22T07:35:12.265019Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pathlib import Path\nfrom typing import Iterable, List, Optional\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\ndef overlay_first_ecg_from_csvs(\n    signal_csv_paths: List[str],\n    *,\n    labels: Optional[List[str]] = None,\n    figure_width: float = 24.0,\n    figure_height: float = 12.0,\n    line_width: float = 1.1,\n    alpha: float = 0.9,\n) -> None:\n    \"\"\"\n    Overlay multiple ECG signal CSVs (same schema) on a standard 12-lead layout.\n\n    Behavior:\n      - Only plots the first ECG ID prefix found in the first CSV.\n      - Plots all leads for that ECG.\n      - Uses a fixed time axis: non-II leads are 2.5 s, II is 10 s.\n      - Assumes data are clean: no NaNs and no row count mismatches across CSVs.\n\n    Expected CSV format:\n      - Either:\n          A) One record per CSV: columns are lead names (I, II, III, aVR, aVL, aVF, V1..V6)\n             and rows are timesteps (with II potentially longer).\n        or:\n          B) Multiple records in one CSV with an 'id' column like \"<record>_<timestep>_<lead>\" and 'value' column.\n             In this case we extract the first record prefix and pivot into lead columns.\n\n    Notes:\n      - If your files are of type (A), this will just read columns and plot directly.\n      - If your files are of type (B), this will pivot each file into wide format for the first record prefix.\n    \"\"\"\n    if len(signal_csv_paths) == 0:\n        raise ValueError(\"signal_csv_paths must be non-empty\")\n\n    if labels is not None and len(labels) != len(signal_csv_paths):\n        raise ValueError(\"If provided, labels must have same length as signal_csv_paths\")\n\n    default_labels = [Path(path).name for path in signal_csv_paths]\n    labels_to_use = labels if labels is not None else default_labels\n\n    # Standard 12-lead layout (grid view) + rhythm strip (II long).\n    layout_map = {\n        \"I\": (0, 0), \"aVR\": (0, 1), \"V1\": (0, 2), \"V4\": (0, 3),\n        \"II\": (1, 0), \"aVL\": (1, 1), \"V2\": (1, 2), \"V5\": (1, 3),\n        \"III\": (2, 0), \"aVF\": (2, 1), \"V3\": (2, 2), \"V6\": (2, 3),\n    }\n    lead_order = [\"I\", \"aVR\", \"V1\", \"V4\", \"II\", \"aVL\", \"V2\", \"V5\", \"III\", \"aVF\", \"V3\", \"V6\"]\n\n    def read_first_record_wide(csv_path: str) -> pd.DataFrame:\n        df = pd.read_csv(csv_path)\n\n        # Wide format: lead columns exist directly.\n        if \"id\" not in df.columns and \"value\" not in df.columns:\n            return df\n\n        # Long format: pivot first record prefix into lead columns.\n        if not {\"id\", \"value\"}.issubset(df.columns):\n            raise ValueError(f\"{csv_path} has unexpected columns: {list(df.columns)}\")\n\n        first_id = str(df[\"id\"].iloc[0])\n        id_pieces = first_id.split(\"_\")\n        if len(id_pieces) < 3:\n            raise ValueError(f\"{csv_path}: cannot parse first id: {first_id!r}\")\n\n        record_prefix = \"_\".join(id_pieces[:-2])\n\n        def parse_id(submission_id: str) -> tuple[str, int, str]:\n            parts = str(submission_id).split(\"_\")\n            lead_name = parts[-1]\n            timestep = int(parts[-2])\n            prefix = \"_\".join(parts[:-2])\n            return prefix, timestep, lead_name\n\n        parsed = df[\"id\"].map(parse_id)\n        df_work = df.copy()\n        df_work[\"record_prefix\"] = parsed.map(lambda t: t[0])\n        df_work[\"timestep\"] = parsed.map(lambda t: t[1])\n        df_work[\"lead\"] = parsed.map(lambda t: t[2])\n\n        df_work = df_work[df_work[\"record_prefix\"] == record_prefix]\n        wide_df = (\n            df_work.pivot(index=\"timestep\", columns=\"lead\", values=\"value\")\n            .sort_index()\n            .reset_index(drop=True)\n        )\n        wide_df.columns.name = None\n        return wide_df\n\n    wide_dfs: List[pd.DataFrame] = [read_first_record_wide(path) for path in signal_csv_paths]\n\n    # Determine which leads are available (use first dataframe as the reference).\n    available_leads = [lead for lead in lead_order if lead in wide_dfs[0].columns]\n    if len(available_leads) == 0:\n        raise ValueError(\n            f\"No recognized lead columns found in first CSV. \"\n            f\"Expected some of {lead_order}. Got columns: {list(wide_dfs[0].columns)}\"\n        )\n\n    fig = plt.figure(figsize=(figure_width, figure_height))\n    gs = fig.add_gridspec(5, 4)\n\n    # Helper to create fixed time axes without fs.\n    def make_time_axis(num_samples: int, duration_seconds: float) -> np.ndarray:\n        if num_samples <= 1:\n            return np.array([0.0], dtype=np.float64)\n        return np.linspace(0.0, duration_seconds, num_samples, endpoint=False, dtype=np.float64)\n\n    # Plot 12-lead grid view.\n    for lead_name, (grid_row, grid_col) in layout_map.items():\n        if lead_name not in available_leads:\n            continue\n\n        ax = fig.add_subplot(gs[grid_row, grid_col])\n\n        duration_seconds = 2.5\n        for df_idx, wide_df in enumerate(wide_dfs):\n            if lead_name not in wide_df.columns:\n                continue\n\n            y = wide_df[lead_name].to_numpy(dtype=np.float64, copy=False)\n\n            # Grid view \"II\" is a short segment (2.5s), even though the rhythm strip below is 10s.\n            if lead_name == \"II\":\n                if y.size > 0:\n                    short_len = max(1, int(round(y.size * (2.5 / 10.0))))\n                    y = y[:short_len]\n\n            t = make_time_axis(y.size, duration_seconds)\n            ax.plot(t, y, linewidth=line_width, alpha=alpha, label=labels_to_use[df_idx])\n\n        ax.set_title(lead_name)\n        ax.grid(True, linestyle=\":\", alpha=0.6)\n        if grid_row != 2:\n            ax.set_xticklabels([])\n\n        # if lead_name == \"I\":\n        #     ax.legend(loc=\"upper right\", fontsize=9, framealpha=0.8)\n\n    # Rhythm strip (full Lead II, 10s) across bottom 2 rows.\n    ax_rhythm = fig.add_subplot(gs[3:, :])\n    if \"II\" in available_leads:\n        for df_idx, wide_df in enumerate(wide_dfs):\n            if \"II\" not in wide_df.columns:\n                continue\n            y = wide_df[\"II\"].to_numpy(dtype=np.float64, copy=False)\n            t = make_time_axis(y.size, 10.0)\n            ax_rhythm.plot(t, y, linewidth=line_width, alpha=alpha, label=labels_to_use[df_idx])\n\n        ax_rhythm.set_title(\"Rhythm II (10s)\")\n        ax_rhythm.grid(True, linestyle=\":\", alpha=0.6)\n        ax_rhythm.set_xlabel(\"Time (s)\")\n        ax_rhythm.legend(loc=\"upper right\", fontsize=10, framealpha=0.8)\n\n    title_record_hint = Path(signal_csv_paths[0]).stem\n    plt.suptitle(f\"Overlayed ECG signals\", fontsize=16)\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:35:12.266961Z","iopub.execute_input":"2026-01-22T07:35:12.267202Z","iopub.status.idle":"2026-01-22T07:35:12.287297Z","shell.execute_reply.started":"2026-01-22T07:35:12.267177Z","shell.execute_reply":"2026-01-22T07:35:12.286536Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"overlay_first_ecg_from_csvs([\n    'liu_submission.csv',\n    'james_submission.csv',\n    'imanishi_submission.csv',\n    'submission.csv',\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T07:35:12.288040Z","iopub.execute_input":"2026-01-22T07:35:12.288238Z","iopub.status.idle":"2026-01-22T07:35:14.733290Z","shell.execute_reply.started":"2026-01-22T07:35:12.288216Z","shell.execute_reply":"2026-01-22T07:35:14.732384Z"}},"outputs":[],"execution_count":null}]}