{"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":"gpu","dataSources":[{"sourceId":97984,"databundleVersionId":14096757,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":13746387,"sourceType":"datasetVersion","datasetId":8747012},{"sourceId":13816899,"sourceType":"datasetVersion","datasetId":8620533},{"sourceId":14539615,"sourceType":"datasetVersion","datasetId":9286546},{"sourceId":271051632,"sourceType":"kernelVersion"},{"sourceId":272137252,"sourceType":"kernelVersion"},{"sourceId":677607,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":513841,"modelId":528480},{"sourceId":698706,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":527653,"modelId":541703}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Submission notebook","metadata":{}},{"cell_type":"code","source":"!pip uninstall -y tensorflow\n!uv pip install --no-deps --system --no-index --find-links='/kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet/setup' 'connected-components-3d'","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-18T17:05:10.827606Z","iopub.execute_input":"2026-01-18T17:05:10.828150Z","iopub.status.idle":"2026-01-18T17:05:50.453056Z","shell.execute_reply.started":"2026-01-18T17:05:10.828124Z","shell.execute_reply":"2026-01-18T17:05:50.452366Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nimport os\nimport warnings\nimport traceback\nfrom pathlib import Path\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.transforms as T\nfrom tqdm.auto import tqdm\nimport timm\nimport kagglehub\nfrom collections import defaultdict\nfrom scipy.signal import savgol_filter\n# --- 1. CONFIGURATION & IMPORTS ---\nwarnings.filterwarnings(\"ignore\", category=UserWarning, module=\"pydantic\")\n\n# Setup Paths (User specific paths)\nsys.path.append('/kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet')\n\n# Import custom modules from the appended path\n# We assume these exist based on your previous code\nfrom stage0_model import Net as Stage0Net\nfrom stage0_common import *\nfrom stage1_model import Net as Stage1Net\nfrom stage1_common import output_to_predict as stage1_output_to_predict, rectify_image\nfrom stage2_model import MyCoordUnetDecoder, encode_with_resnet\n\n# Configuration\nseed = 1\nCUDA0 = \"cuda:0\"\nFLOAT_TYPE = torch.float32\n\n# Deterministic setup\ndeterministic = kagglehub.package_import('wasupandceacar/deterministic').deterministic\ndeterministic.init_all(seed, disable_list=['cuda_block'])\n\n# Dataset Logic\nif_submit = os.getenv('KAGGLE_IS_COMPETITION_RERUN')\nif if_submit:\n    test_meta = Path(\"/kaggle/input/physionet-ecg-image-digitization/test.csv\")\n    test_dir = Path(\"/kaggle/input/physionet-ecg-image-digitization/test\")\nelse:\n    test_meta = Path(\"/kaggle/input/physio-test-fake-dataset/test_fake/test.csv\")\n    test_dir = Path(\"/kaggle/input/physio-test-fake-dataset/test_fake\")\n\nvalid_df = pd.read_csv(test_meta)\nvalid_df['id'] = valid_df['id'].astype(str)\n# We group by ID so we can process one image and generate rows for all its leads\nvalid_groups = valid_df.groupby('id')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T17:05:50.454487Z","iopub.execute_input":"2026-01-18T17:05:50.454751Z","iopub.status.idle":"2026-01-18T17:06:09.095284Z","shell.execute_reply.started":"2026-01-18T17:05:50.454726Z","shell.execute_reply":"2026-01-18T17:06:09.094741Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 2. STAGE 2 HELPER FUNCTIONS & MODEL ---\n# These were previously in stage2.py\n\ndef fill_linear(y, missing_mask):\n    x = np.arange(len(y))\n    y_filled = y.copy()\n    y_filled[missing_mask] = np.interp(\n        x[missing_mask],\n        x[~missing_mask],\n        y[~missing_mask]\n    )\n    return y_filled\n\ndef subpixel_centroid_1d(col, i_max, radius=2):\n    H = col.shape[0]\n    lo = max(i_max - radius, 0)\n    hi = min(i_max + radius + 1, H)\n    weights = col[lo:hi].astype(np.float32)\n    if weights.sum() <= 1e-6:\n        return float(i_max)\n    idxs = np.arange(lo, hi, dtype=np.float32)\n    return float((idxs * weights).sum() / weights.sum())\n\n# Constants for Stage 2\nS2_x0, S2_x1 = 0, 2176\nS2_y0, S2_y1 = 0, 1696\nS2_zero_mv = [703.5, 987.6, 1271.5, 1531.5]\nS2_mv_to_pixel = 78.5\nS2_t0, S2_t1 = 235, 4161\n\ndef pixel_to_series_here(pixel, zero_mv, length):\n    _, H, W = pixel.shape\n    assert H == 1696\n\n    series = []\n    for j in range(4):\n        p = pixel[j]\n        amax = p.argmax(0)\n        s = np.zeros(W, dtype=np.float32)\n        for x in range(W):\n            s[x] = subpixel_centroid_1d(p[:, x], amax[x], radius=2)\n        miss = ((p > 0.8).sum(0) == 0)\n        if miss.any():\n            s = fill_linear(s, miss)\n        series.append(s)\n\n    series = np.stack(series).astype(np.float32)\n    series = (np.array(zero_mv).reshape(4, 1) - series) / S2_mv_to_pixel\n\n    if length is not None and length != W:\n        series = torch.from_numpy(series).unsqueeze(1)\n        series = F.interpolate(series, size=length, mode='linear', align_corners=False)\n        series = series.squeeze(1).cpu().numpy()\n    return series\n\nclass Net3(nn.Module):\n    def __init__(self, pretrained=True):\n        super(Net3, self).__init__()\n        encoder_dim = [64, 128, 256, 512]\n        decoder_dim = [128, 64, 32, 16]\n        self.encoder = timm.create_model(\n            model_name='resnet34.a3_in1k', pretrained=pretrained, in_chans=3, num_classes=0, global_pool=''\n        )\n        self.decoder = MyCoordUnetDecoder(\n            in_channel=encoder_dim[-1],\n            skip_channel=encoder_dim[:-1][::-1] + [0],\n            out_channel=decoder_dim,\n            scale=[2, 2, 2, 2]\n        )\n        self.pixel = nn.Conv2d(decoder_dim[-1], 4, 1)\n\n    def forward(self, image):\n        encode = encode_with_resnet(self.encoder, image)\n        last, _ = self.decoder(feature=encode[-1], skip=encode[:-1][::-1] + [None])\n        pixel = self.pixel(last)\n        return pixel\n\ndef series_dict(series):\n    series_by_lead = dict()\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        split = np.array_split(series[l], 4)\n        for (k, s) in zip(lead_names, split):\n            series_by_lead[k] = s\n    series_by_lead['II'] = series[3]\n    return series_by_lead\n\n# --- 3. LOAD MODELS ---\n# We load all models once to avoid overhead in the loop\n\nprint(\"Loading models...\")\n\n# Load Stage 0\nstage0_net = Stage0Net(pretrained=False)\nstage0_net = load_net(stage0_net, '/kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet/weight/stage0-last.checkpoint.pth')\nstage0_net.to(CUDA0)\nstage0_net.eval()\n\n# Load Stage 1\nstage1_net = Stage1Net(pretrained=False)\nstage1_net = load_net(stage1_net, '/kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet/weight/stage1-last.checkpoint.pth')\nstage1_net.to(CUDA0)\nstage1_net.eval()\n\n# Load Stage 2\nstage2_net = Net3(pretrained=False).to(CUDA0)\nstage2_model_path = \"/kaggle/input/wasup-finetuned-on-l1-loss-with-true-ser/pytorch/default/12/checkpoint_epoch_001_step_004744.pt\"\nstage2_net.load_state_dict(torch.load(stage2_model_path)[\"model_state_dict\"])\nstage2_net.eval()\n\n# Stage 2 Transform\nresize_s2 = T.Resize((1696, 4352), interpolation=T.InterpolationMode.BILINEAR)\n\nprint(\"Models loaded successfully.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T17:06:09.096055Z","iopub.execute_input":"2026-01-18T17:06:09.096276Z","iopub.status.idle":"2026-01-18T17:06:12.762393Z","shell.execute_reply.started":"2026-01-18T17:06:09.096259Z","shell.execute_reply":"2026-01-18T17:06:12.761534Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import timm\nLEADS_ORDER = [\"I\",\"II\",\"III\",\"aVR\",\"aVL\",\"aVF\",\"V1\",\"V2\",\"V3\",\"V4\",\"V5\",\"V6\"]\n\ndef change_color(image_rgb):\n    hsv = cv2.cvtColor(image_rgb, cv2.COLOR_RGB2HSV)\n    h, s, v = cv2.split(hsv)\n    v_denoised = cv2.fastNlMeansDenoising(v, h=5.46)\n    std = np.std(v_denoised)\n    clip_limit = max(1.0, min(3.5, 2.0 + std / 25))\n    clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=(8, 8))\n    v_enhanced = clahe.apply(v_denoised)\n    hsv_enhanced = cv2.merge([h, s, v_enhanced])\n    return cv2.cvtColor(hsv_enhanced, cv2.COLOR_HSV2RGB)\n\ndef series_dict(series_4row):\n    series_4row = np.asarray(series_4row)\n    if series_4row.ndim == 3: series_4row = series_4row[0]\n    if series_4row.shape[0] != 4 and series_4row.shape[1] == 4: series_4row = series_4row.T\n\n    d = {}\n    names = [\n        ['I','aVR','V1','V4'],\n        ['II_short','aVL','V2','V5'],   \n        ['III','aVF','V3','V6'],\n    ]\n    for r in range(3):\n        for lead, arr in zip(names[r], np.array_split(series_4row[r], 4)):\n            d[lead] = np.asarray(arr, dtype=np.float32)\n\n    d['II'] = np.asarray(series_4row[3], dtype=np.float32)  # full 10s II\n    return d\n\n# ✅ FIX: Einthoven correction on SHORT leads only\ndef dw(d, alpha=0.33):\n    if all(k in d for k in ['I','II_short','III']):\n        L1, L2s, L3 = d['I'], d['II_short'], d['III']\n        e = L2s - (L1 + L3)\n        d['I']        = L1 + alpha*e\n        d['III']      = L3 + alpha*e\n        d['II_short'] = L2s - alpha*e\n    return d\n\ndef clahe_luminance_bgr(img_bgr, clip=2.0, tile=8):\n    lab = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2LAB)\n    l, a, b = cv2.split(lab)\n    clahe = cv2.createCLAHE(clipLimit=float(clip), tileGridSize=(int(tile), int(tile)))\n    l2 = clahe.apply(l)\n    return cv2.cvtColor(cv2.merge([l2, a, b]), cv2.COLOR_LAB2BGR)\n\ndef grayworld_white_balance(img_bgr):\n    img = img_bgr.astype(np.float32)\n    b, g, r = cv2.split(img)\n    mb, mg, mr = b.mean(), g.mean(), r.mean()\n    m = (mb + mg + mr) / 3.0\n    b *= (m / (mb + 1e-6)); g *= (m / (mg + 1e-6)); r *= (m / (mr + 1e-6))\n    return np.clip(cv2.merge([b, g, r]), 0, 255).astype(np.uint8)\n\ndef denoise_median(img_bgr, k=3):\n    k = int(k); k = k if k % 2 == 1 else k + 1\n    return cv2.medianBlur(img_bgr, k)\n\ndef denoise_bilateral(img_bgr, d=7, sigmaColor=50, sigmaSpace=50):\n    return cv2.bilateralFilter(img_bgr, d=int(d), sigmaColor=float(sigmaColor), sigmaSpace=float(sigmaSpace))\n\ndef illumination_strength(img_bgr, sigma=35):\n    gray = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2GRAY).astype(np.float32) / 255.0\n    blur = cv2.GaussianBlur(gray, (0, 0), sigma)\n    return float(np.std(blur))\n\ndef bg_correct_lab_l(img_bgr, k=81):\n    lab = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2LAB)\n    l, a, b = cv2.split(lab)\n    k = int(k); k = k if k % 2 == 1 else k + 1\n    kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (k, k))\n    bg = cv2.morphologyEx(l, cv2.MORPH_OPEN, kernel)\n    l_corr = cv2.subtract(l, bg)\n    l_corr = cv2.normalize(l_corr, None, 0, 255, cv2.NORM_MINMAX).astype(np.uint8)\n    return cv2.cvtColor(cv2.merge([l_corr, a, b]), cv2.COLOR_LAB2BGR)\n\ndef preprocess_by_source(img_bgr, source):\n    s = str(source)\n    if s == \"0001\": return img_bgr\n    if s == \"0003\": return clahe_luminance_bgr(grayworld_white_balance(img_bgr), clip=1.2, tile=8)\n    if s == \"0004\": return img_bgr\n    if s == \"0006\":\n        x = denoise_bilateral(img_bgr, d=5, sigmaColor=25, sigmaSpace=25)\n        return clahe_luminance_bgr(x, clip=1.2, tile=8)\n    if s == \"0005\":\n        x = img_bgr\n        if illumination_strength(x, sigma=35) > 0.14: x = bg_correct_lab_l(x, k=81)\n        if cv2.cvtColor(x, cv2.COLOR_BGR2GRAY).std() < 30: x = clahe_luminance_bgr(x, clip=1.1, tile=8)\n        return x\n    if s == \"0009\":\n        x = img_bgr\n        if illumination_strength(x, sigma=35) > 0.14: x = bg_correct_lab_l(x, k=101)\n        return denoise_median(x, k=3)\n    if s == \"0010\":\n        x = img_bgr\n        if illumination_strength(x, sigma=35) > 0.14: x = bg_correct_lab_l(x, k=81)\n        if cv2.cvtColor(x, cv2.COLOR_BGR2GRAY).std() < 30: x = clahe_luminance_bgr(x, clip=1.15, tile=8)\n        return x\n    if s == \"0011\": return clahe_luminance_bgr(grayworld_white_balance(img_bgr), clip=1.2, tile=8)\n    if s == \"0012\": return img_bgr\n    return img_bgr\n\ndef stage1_quality(s1_rgb):\n    g = cv2.cvtColor(s1_rgb.astype(np.uint8), cv2.COLOR_RGB2GRAY)\n    e = cv2.Canny(g, 50, 150)\n    density = e.mean() / 255.0\n    gx = cv2.Sobel(g, cv2.CV_32F, 1, 0, ksize=3)\n    gy = cv2.Sobel(g, cv2.CV_32F, 0, 1, ksize=3)\n    ax = float(np.mean(np.abs(gx))); ay = float(np.mean(np.abs(gy)))\n    anis = max(ax, ay) / (min(ax, ay) + 1e-6)\n    return float(density * 0.7 + np.tanh(anis - 1.0) * 0.3)\n\nCLS_MODEL_NAME=\"efficientnet_b2\"\nCLS_NUM_CLASSES=12\nCLS_RESOLUTION=256\nCLS_CKPT_PATH=\"/kaggle/input/physionet-image-multi-class-train/efficientnet_b2_fold_0_best.pth\"\ndevice = \"cuda\"\n\ndef build_classifier(device=\"cuda\"):\n    model =  timm.create_model(CLS_MODEL_NAME, pretrained=False)\n    model.reset_classifier(CLS_NUM_CLASSES)\n    checkpoint = torch.load(CLS_CKPT_PATH, weights_only=False)\n    model.load_state_dict(checkpoint['model'])\n    return model.to(device).eval()\n\ndef cls_preprocess_bgr(img_bgr):\n    img = cv2.resize(img_bgr, (CLS_RESOLUTION, CLS_RESOLUTION), interpolation=cv2.INTER_LINEAR).astype(np.float32)/255.0\n    mean = np.array([0.485,0.456,0.406], np.float32); std = np.array([0.229,0.224,0.225], np.float32)\n    img = (img - mean) / std\n    return torch.from_numpy(img).permute(2,0,1).unsqueeze(0)\n\ncls_model = build_classifier(device=device)\n\n@torch.no_grad()\ndef predict_source_suffix(model, img_bgr, device=\"cuda\"):\n    x = cls_preprocess_bgr(img_bgr).to(device)\n    p = F.softmax(model(x), dim=1)[0]\n    cls = int(torch.argmax(p).item())\n    return f\"{cls+1:04d}\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T17:06:12.763791Z","iopub.execute_input":"2026-01-18T17:06:12.764059Z","iopub.status.idle":"2026-01-18T17:06:13.449683Z","shell.execute_reply.started":"2026-01-18T17:06:12.764040Z","shell.execute_reply":"2026-01-18T17:06:13.448687Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def make_all_0001_sigs(all_results):\n    \"\"\"\n    Scans all results for '0001' (clean) predictions and builds a searchable library.\n    Library structure: { signature_tuple: [ (sample_id, leads_dict), ... ] }\n    \"\"\"\n    library = defaultdict(list)\n    \n    for sample_id, predictions in all_results.items():\n        # predictions is dict like {'0001': series_by_lead}\n        # We only care if the key is '0001'\n        if '0001' in predictions:\n            leads = predictions['0001']\n            \n            # Create a signature based on lead names and lengths\n            # Signature format: (('I', 500), ('II', 500), ('V1', 500)...)\n            sorted_keys = sorted(leads.keys())\n            signature = tuple((k, len(leads[k])) for k in sorted_keys)\n            \n            library[signature].append((sample_id, leads))\n            \n    print(f\"Created 0001 library with {sum(len(v) for v in library.values())} unique clean signals.\")\n    return library\n\ndef potentially_replace(series_by_lead, all_sigs_library, threshold):\n    \"\"\"\n    Checks if 'series_by_lead' is close enough to any signal in 'all_sigs_library'.\n    Returns the REPLACEMENT leads if a match is found, otherwise returns ORIGINAL.\n    \"\"\"\n    # 1. Generate signature for the candidate\n    sorted_keys = sorted(series_by_lead.keys())\n    signature = tuple((k, len(series_by_lead[k])) for k in sorted_keys)\n    \n    # 2. Get compatible clean signals (Exact length match only)\n    candidates = all_sigs_library.get(signature)\n    \n    # If no clean images share this structure, we can't replace anything.\n    if not candidates:\n        return series_by_lead\n    \n    best_dist = float('inf')\n    best_match_leads = None\n    \n    # 3. Find nearest neighbor\n    for _, clean_leads in candidates:\n        current_dist = 0.0\n        # Calculate Euclidean distance across all leads\n        for lead_name in series_by_lead:\n            # We know lengths match because of the signature check\n            diff = series_by_lead[lead_name] - clean_leads[lead_name]\n            current_dist += np.linalg.norm(diff)\n            \n        if current_dist < best_dist:\n            best_dist = current_dist\n            best_match_leads = clean_leads\n\n    # 4. Check Threshold\n    if best_dist < threshold:\n        # print(f\"Match found! Replaced with distance {best_dist:.2f}\") # Optional debug\n        return best_match_leads\n    \n    return series_by_lead\n\ndef load_clean_library_npz(path):\n    data = np.load(path, allow_pickle=True)\n\n    clean_library = defaultdict(list)\n\n    meta = data[\"meta\"]\n    for entry_id, sig, name, lead_keys in meta:\n        leads = {\n            k: data[f\"{entry_id}:{k}\"]\n            for k in lead_keys\n        }\n        clean_library[tuple(sig)].append((name, leads))\n\n    return clean_library\n\ndef concat_clean_libraries(lib_a, lib_b):\n    merged = defaultdict(list)\n\n    # Copy lib_a\n    for sig, items in lib_a.items():\n        merged[sig].extend(items)\n\n    # Append lib_b\n    for sig, items in lib_b.items():\n        merged[sig].extend(items)\n\n    return merged\n    \nclean_lib = load_clean_library_npz(\"/kaggle/input/clean-lib-2/clean_lib.npz\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T17:06:32.210148Z","iopub.execute_input":"2026-01-18T17:06:32.210427Z","iopub.status.idle":"2026-01-18T17:06:37.619814Z","shell.execute_reply.started":"2026-01-18T17:06:32.210391Z","shell.execute_reply":"2026-01-18T17:06:37.619018Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 4. PIPELINE EXECUTION ---\n\nsubmit_df_list = []\n\nall_results = {}\nmeta_lookup = {}\n\nprint(\"Starting inference loop...\")\n\nfor sample_id, df_meta in tqdm(valid_groups):\n\n    meta_lookup[str(sample_id)] = df_meta\n    \n    # 1. READ IMAGE\n    input_path = test_dir / f'{sample_id}.png'\n\n    use_preprocess = True\n    if not use_preprocess:\n        # Use standard cv2 reading (BGR) then convert to RGB\n        original_image = cv2.imread(str(input_path))\n        pred_src = predict_source_suffix(cls_model, original_image, device=device)\n        if original_image is None:\n            print(f\"Error reading {input_path}\")\n            continue\n        original_image = cv2.cvtColor(original_image, cv2.COLOR_BGR2RGB)\n    else:\n        image_bgr = cv2.imread(input_path, cv2.IMREAD_COLOR)\n        original_image = cv2.imread(str(input_path))\n        pred_src = predict_source_suffix(cls_model, original_image, device=device)\n        original_image = cv2.cvtColor(preprocess_by_source(image_bgr.copy(), pred_src), cv2.COLOR_BGR2RGB)\n    \n    # This variable tracks the image state through the pipeline\n    current_image = original_image\n    \n    # We use a flag to catch failures in early stages\n    pipeline_failed = False\n    \n    # --- STAGE 0: Homography ---\n    try:\n        batch = image_to_batch(current_image)\n        with torch.no_grad(), torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n            output = stage0_net(batch)\n        rotated, keypoint = output_to_predict(current_image, batch, output)\n        normalised, _, _ = normalise_by_homography(rotated, keypoint)\n        current_image = normalised # Output of Stage 0 is input to Stage 1\n    except Exception:\n        traceback.print_exc() # Uncomment for debug\n        pipeline_failed = True\n\n    # --- STAGE 1: Rectification ---\n    if not pipeline_failed:\n        try:\n            # Prepare batch for Stage 1. \n            # Note: The original code saved normalised as BGR then read as RGB. \n            # Since current_image is RGB (from Stage 0), we use it directly.\n            # Usually stage1 expect (C, H, W)\n            img_tensor = torch.from_numpy(np.ascontiguousarray(current_image.transpose(2, 0, 1))).unsqueeze(0)\n            batch = {'image': img_tensor} \n            \n            with torch.no_grad(), torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n                output = stage1_net(batch)\n            gridpoint_xy, _ = stage1_output_to_predict(current_image, batch, output)\n            rectified = rectify_image(current_image, gridpoint_xy)\n            current_image = rectified # Output of Stage 1 is input to Stage 2\n        except Exception:\n            traceback.print_exc()\n            pipeline_failed = True\n\n    # --- STAGE 2: Signal Extraction ---\n    # We need to determine the length for the 'II' lead to resize properly/interpolate later\n    length = df_meta[df_meta['lead']=='II'].iloc[0].number_of_rows\n    \n    series_result = None\n\n    if not pipeline_failed:\n        try:\n            # Crop and Normalize\n            # Original code: image = image[y0:y1, x0:x1] / 255\n            # current_image is RGB here\n            crop = current_image[S2_y0:S2_y1, S2_x0:S2_x1] / 255.0\n            \n            # Prepare Tensor\n            crop_tensor = torch.from_numpy(np.ascontiguousarray(crop).transpose(2, 0, 1)).unsqueeze(0).float().to(CUDA0)\n            batch = resize_s2(crop_tensor)\n\n            with torch.no_grad(), torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n                output = stage2_net(batch)\n            \n            pixel = torch.sigmoid(output).float().data.cpu().numpy()[0]\n            series_result = pixel_to_series_here(pixel[..., S2_t0:S2_t1], S2_zero_mv, length)\n            for i in range(4):\n                series_result[i] = savgol_filter(series_result[i], window_length=9, polyorder=2)\n            \n        except Exception:\n            traceback.print_exc()\n            pipeline_failed = True\n    \n    # --- FALLBACK ---\n    if pipeline_failed or series_result is None:\n        # If any stage crashed, we return zeros for this ID\n        series_result = np.zeros((4, length))\n\n    # --- FORMAT SUBMISSION ---\n    series_by_lead = series_dict(series_result)\n\n    if str(sample_id) not in all_results:\n        all_results[str(sample_id)] = {}\n        all_results[str(sample_id)][str(pred_src)] = series_by_lead\n\n# --- PHASE 2: MATCHING & REPLACEMENT ---\nprint(\"Building Clean Library and attempting replacements...\")\n\n# 1. Build the library of clean '0001' signals\nclean_library = make_all_0001_sigs(all_results)\n\nclean_library = concat_clean_libraries(clean_library, clean_lib)\n# 2. Define your Threshold (You need to tune this using the plots from previous steps!)\n# Example: If your distances are usually 100-200 for matches and 5000+ for non-matches\nMATCH_THRESHOLD = 50.0\n\nfor sample_id, predictions in tqdm(all_results.items()):\n    \n    # Retrieve the metadata we saved earlier\n    df_meta = meta_lookup[str(sample_id)]\n    \n    # Get the predicted class and the signal\n    pred_class = list(predictions.keys())[0]\n    series_by_lead = predictions[pred_class]\n    \n    # APPLY THE TRICK: \n    # If the image was NOT predicted as '0001', try to find a '0001' match to replace it\n    if pred_class != \"0001\":\n        series_by_lead = potentially_replace(series_by_lead, clean_library, MATCH_THRESHOLD)\n    \n    # Now write to dataframe (Moved from inside the loop)\n    for _, row in df_meta.iterrows():\n        s = series_by_lead[row.lead]\n        \n        # Final interpolation if length mismatch\n        if len(s) != row.number_of_rows:\n            x_old = np.linspace(0.0, 1.0, len(s))\n            x_new = np.linspace(0.0, 1.0, row.number_of_rows)\n            s = np.interp(x_new, x_old, s)\n            \n        row_id = [f'{sample_id}_{t}_{row.lead}' for t in range(row.number_of_rows)]\n        \n        submit_df_list.append(pd.DataFrame({\n            'id': row_id,\n            'value': s,\n        }))\n\n# --- 5. SAVE CSV ---\nprint(\"Generating submission.csv...\")\nif submit_df_list:\n    submit_df = pd.concat(submit_df_list, axis=0, ignore_index=True, sort=False, copy=False)\n    submit_df.to_csv('submission.csv', index=False)\n    print(\"Done!\")\nelse:\n    print(\"Error: No data generated.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T17:06:37.621052Z","iopub.execute_input":"2026-01-18T17:06:37.621280Z","iopub.status.idle":"2026-01-18T17:07:15.878258Z","shell.execute_reply.started":"2026-01-18T17:06:37.621261Z","shell.execute_reply":"2026-01-18T17:07:15.877565Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub = pd.read_csv('submission.csv')\nprint(len(sub))\nsub.head(30)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-18T17:07:15.879244Z","iopub.execute_input":"2026-01-18T17:07:15.879468Z","iopub.status.idle":"2026-01-18T17:07:16.091367Z","shell.execute_reply.started":"2026-01-18T17:07:15.879449Z","shell.execute_reply":"2026-01-18T17:07:16.090760Z"}},"outputs":[],"execution_count":null}]}