{"cells":[{"cell_type":"markdown","metadata":{},"source":"# PhysioNet ECG - Custom Model Inference\n\n**Pipeline**:\n- Stage 0 & 1: LB 17.95 baseline (rotation + grid rectification)\n- Stage 2: Custom trained U-Net model"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Offline notebook - no pip install needed\n# connected-components-3d is pre-installed in Kaggle environment\n# We define UNet model inline instead of using segmentation-models-pytorch\nprint(\"Starting offline inference...\")"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"import warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport os\nimport sys\nimport gc\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport cv2\nimport pandas as pd\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\nfrom shutil import copyfile\nfrom scipy.signal import savgol_filter\nimport torchvision.transforms as T\nimport timm\n\nCUDA0 = 'cuda:0' if torch.cuda.is_available() else 'cpu'\nFLOAT_TYPE = torch.float32\n\n# Paths\ntest_meta = Path(\"/kaggle/input/physionet-ecg-image-digitization/test.csv\")\ntest_dir = Path(\"/kaggle/input/physionet-ecg-image-digitization/test\")\n\n# LB 17.95 baseline models\nmodel_base = '/kaggle/input/hengck23-submit-physionet'\nif os.path.exists(os.path.join(model_base, 'hengck23-submit-physionet')):\n    model_base = os.path.join(model_base, 'hengck23-submit-physionet')\nsys.path.append(model_base)\n\n# Custom model path\nCUSTOM_MODEL_PATH = '/kaggle/input/ecg-segmentation-model-custom/kaggle_model/best_model.pth'\n\nvalid_df = pd.read_csv(test_meta)\nvalid_df['id'] = valid_df['id'].astype(str)\nvalid_id = valid_df['id'].unique().tolist()\n\nglobal_dict = {\n    \"stage0_dir\": \"/kaggle/working/stage0\",\n    \"stage1_dir\": \"/kaggle/working/stage1\",\n    \"stage2_dir\": \"/kaggle/working/stage2\",\n    \"model_base\": model_base,\n}\n\nprint(f'Device: {CUDA0}')\nprint(f'Model base: {model_base}')\nprint(f'Custom model: {CUSTOM_MODEL_PATH}')\nprint(f'Custom model exists: {os.path.exists(CUSTOM_MODEL_PATH)}')\nprint(f'Test images: {len(valid_id)}')"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Stage 0: Rotation correction (from LB 17.95)\nfrom stage0_model import Net as Stage0Net\nfrom stage0_common import *\n\ndef apply_grayscale_guidance(image_rgb):\n    gray = cv2.cvtColor(image_rgb, cv2.COLOR_RGB2GRAY)\n    denoised = cv2.fastNlMeansDenoising(gray, h=10)\n    clahe = cv2.createCLAHE(clipLimit=3.0, tileGridSize=(8,8))\n    contrast_enhanced = clahe.apply(denoised)\n    return cv2.cvtColor(contrast_enhanced, cv2.COLOR_GRAY2RGB)\n\nstage0_dir = Path(global_dict[\"stage0_dir\"])\nstage0_dir.mkdir(exist_ok=True, parents=True)\n\nstage0_net = Stage0Net(pretrained=False)\nstage0_net = load_net(stage0_net, f'{model_base}/weight/stage0-last.checkpoint.pth')\nstage0_net.to(CUDA0).eval()\n\nprint(\"Stage 0: Processing...\")\nfor sample_id in tqdm(valid_id):\n    path = test_dir / f'{sample_id}.png'\n    output_path = stage0_dir / f'{sample_id}.png'\n    \n    image_original = cv2.imread(str(path), cv2.IMREAD_COLOR)\n    if image_original is None:\n        continue\n    image_original = cv2.cvtColor(image_original, cv2.COLOR_BGR2RGB)\n    image_for_model = apply_grayscale_guidance(image_original)\n    \n    batch = image_to_batch(image_for_model)\n    batch = {k: v.to(CUDA0) if isinstance(v, torch.Tensor) else v for k, v in batch.items()}\n    \n    try:\n        with torch.no_grad(), torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n            output = stage0_net(batch)\n        rotated, keypoint = output_to_predict(image_original, batch, output)\n        normalised, _, _ = normalise_by_homography(rotated, keypoint)\n        cv2.imwrite(str(output_path), cv2.cvtColor(normalised, cv2.COLOR_RGB2BGR))\n    except:\n        copyfile(path, output_path)\n\ndel stage0_net\ngc.collect()\ntorch.cuda.empty_cache()\nprint(\"Stage 0 complete\")"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Stage 1: Grid rectification (from LB 17.95)\nfrom stage1_model import Net as Stage1Net\nfrom stage1_common import *\n\nstage1_dir = Path(global_dict[\"stage1_dir\"])\nstage1_dir.mkdir(exist_ok=True, parents=True)\n\nstage1_net = Stage1Net(pretrained=False)\nstage1_net = load_net(stage1_net, f'{model_base}/weight/stage1-last.checkpoint.pth')\nstage1_net.to(CUDA0).eval()\n\nprint(\"Stage 1: Processing...\")\nfor sample_id in tqdm(valid_id):\n    path = stage0_dir / f'{sample_id}.png'\n    output_path = stage1_dir / f'{sample_id}.png'\n    \n    image = cv2.imread(str(path), cv2.IMREAD_COLOR)\n    if image is None:\n        image = cv2.imread(str(test_dir / f'{sample_id}.png'), cv2.IMREAD_COLOR)\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    \n    batch = {'image': torch.from_numpy(image.transpose(2, 0, 1)).unsqueeze(0).to(CUDA0)}\n    \n    try:\n        with torch.no_grad(), torch.amp.autocast('cuda', dtype=FLOAT_TYPE):\n            output = stage1_net(batch)\n        gridpoint_xy, _ = output_to_predict(image, batch, output)\n        rectified = rectify_image(image, gridpoint_xy)\n        cv2.imwrite(str(output_path), cv2.cvtColor(rectified, cv2.COLOR_RGB2BGR))\n    except:\n        copyfile(path, output_path)\n\ndel stage1_net\ngc.collect()\ntorch.cuda.empty_cache()\nprint(\"Stage 1 complete\")"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Stage 2: Custom Model - Define UNet inline (no smp dependency)\n\nclass ConvBlock(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.conv = nn.Sequential(\n            nn.Conv2d(in_ch, out_ch, 3, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_ch, out_ch, 3, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True)\n        )\n    def forward(self, x):\n        return self.conv(x)\n\nclass DecoderBlock(nn.Module):\n    def __init__(self, in_ch, skip_ch, out_ch):\n        super().__init__()\n        self.up = nn.ConvTranspose2d(in_ch, in_ch, 2, stride=2)\n        self.conv = ConvBlock(in_ch + skip_ch, out_ch)\n    \n    def forward(self, x, skip=None):\n        x = self.up(x)\n        if skip is not None:\n            # Handle size mismatch\n            if x.shape[2:] != skip.shape[2:]:\n                x = F.interpolate(x, size=skip.shape[2:], mode='bilinear', align_corners=False)\n            x = torch.cat([x, skip], dim=1)\n        return self.conv(x)\n\nclass UNetResNet34(nn.Module):\n    \"\"\"UNet with ResNet34 encoder - compatible with segmentation_models_pytorch\"\"\"\n    def __init__(self, in_channels=3, classes=4):\n        super().__init__()\n        # Encoder (ResNet34)\n        resnet = timm.create_model('resnet34', pretrained=False, in_chans=in_channels)\n        \n        self.encoder0 = nn.Sequential(resnet.conv1, resnet.bn1, resnet.act1)\n        self.pool0 = resnet.maxpool\n        self.encoder1 = resnet.layer1  # 64\n        self.encoder2 = resnet.layer2  # 128\n        self.encoder3 = resnet.layer3  # 256\n        self.encoder4 = resnet.layer4  # 512\n        \n        # Decoder\n        self.decoder4 = DecoderBlock(512, 256, 256)\n        self.decoder3 = DecoderBlock(256, 128, 128)\n        self.decoder2 = DecoderBlock(128, 64, 64)\n        self.decoder1 = DecoderBlock(64, 64, 32)\n        self.decoder0 = DecoderBlock(32, 0, 16)\n        \n        self.final = nn.Conv2d(16, classes, 1)\n    \n    def forward(self, x):\n        # Encoder\n        e0 = self.encoder0(x)\n        e0_pool = self.pool0(e0)\n        e1 = self.encoder1(e0_pool)\n        e2 = self.encoder2(e1)\n        e3 = self.encoder3(e2)\n        e4 = self.encoder4(e3)\n        \n        # Decoder\n        d4 = self.decoder4(e4, e3)\n        d3 = self.decoder3(d4, e2)\n        d2 = self.decoder2(d3, e1)\n        d1 = self.decoder1(d2, e0)\n        d0 = self.decoder0(d1)\n        \n        # Upsample to input size\n        d0 = F.interpolate(d0, size=x.shape[2:], mode='bilinear', align_corners=False)\n        \n        return self.final(d0)\n\n# Try loading with smp format, fall back to custom\ntry:\n    import segmentation_models_pytorch as smp\n    custom_model = smp.Unet(\n        encoder_name='resnet34',\n        encoder_weights=None,\n        in_channels=3,\n        classes=4,\n        activation=None\n    )\n    custom_model.load_state_dict(torch.load(CUSTOM_MODEL_PATH, map_location=CUDA0))\n    print(\"Loaded model using segmentation_models_pytorch\")\nexcept:\n    print(\"smp not available, using inline UNet definition\")\n    custom_model = UNetResNet34(in_channels=3, classes=4)\n    # Try to load weights (may need adjustment for key names)\n    try:\n        state_dict = torch.load(CUSTOM_MODEL_PATH, map_location=CUDA0)\n        custom_model.load_state_dict(state_dict, strict=False)\n        print(\"Loaded weights (some keys may be missing)\")\n    except Exception as e:\n        print(f\"Warning: Could not load weights: {e}\")\n\ncustom_model.to(CUDA0).eval()\nprint(f\"Custom model ready on {CUDA0}\")"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Signal extraction parameters (from LB 17.95)\nmv_to_pixel = 78.5\nzero_mv = [703.5, 987.5, 1271.5, 1531.5]\nt0, t1 = 235, 4161\n\n# Resize transform\nresize = T.Resize((424, 1088), interpolation=T.InterpolationMode.BILINEAR)\nresize_full = T.Resize((1696, 4352), interpolation=T.InterpolationMode.BILINEAR)\n\nclass MedicalConstraintRefiner:\n    def __init__(self, alpha=0.33):\n        self.alpha = alpha\n\n    def apply_einthoven_law(self, series_dict):\n        if all(k in series_dict for k in ['I', 'II', 'III']):\n            L1, L2, L3 = series_dict['I'], series_dict['II'], series_dict['III']\n            error = L2 - (L1 + L3)\n            series_dict['I'] = L1 + (self.alpha * error)\n            series_dict['III'] = L3 + (self.alpha * error)\n            series_dict['II'] = L2 - (self.alpha * error)\n        return series_dict\n\nrefiner = MedicalConstraintRefiner()\n\ndef pixel_to_series_custom(pixel_mask, zero_mv_list, target_length):\n    \"\"\"Convert pixel mask to signal series\"\"\"\n    num_rows, h, w = pixel_mask.shape\n    series = np.zeros((num_rows, target_length))\n    \n    for row in range(num_rows):\n        row_mask = pixel_mask[row]\n        signal = []\n        \n        for x in range(w):\n            col = row_mask[:, x]\n            if col.max() > 0.3:  # threshold\n                # Weighted centroid for sub-pixel accuracy\n                positions = np.arange(len(col))\n                weights = col / (col.sum() + 1e-8)\n                y_pos = np.sum(positions * weights)\n            else:\n                y_pos = zero_mv_list[row] if row < len(zero_mv_list) else h // 2\n            signal.append(y_pos)\n        \n        # Resample to target length\n        signal = np.array(signal)\n        x_old = np.linspace(0, 1, len(signal))\n        x_new = np.linspace(0, 1, target_length)\n        series[row] = np.interp(x_new, x_old, signal)\n    \n    return series\n\ndef series_to_dict(series):\n    d = {}\n    names = [['I', 'aVR', 'V1', 'V4'], ['II', 'aVL', 'V2', 'V5'], ['III', 'aVF', 'V3', 'V6']]\n    for i in range(3):\n        splits = np.array_split(series[i], 4)\n        for name, data in zip(names[i], splits):\n            d[name] = data\n    d['II_Long'] = series[3]\n    return d\n\ndef dict_to_series(d, shape):\n    s = np.zeros(shape)\n    s[0] = np.concatenate([d['I'], d['aVR'], d['V1'], d['V4']])\n    s[1] = np.concatenate([d['II'], d['aVL'], d['V2'], d['V5']])\n    s[2] = np.concatenate([d['III'], d['aVF'], d['V3'], d['V6']])\n    s[3] = d['II_Long']\n    return s"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Stage 2: Process with custom model\nstage1_dir = Path(global_dict[\"stage1_dir\"])\nstage2_dir = Path(global_dict[\"stage2_dir\"])\nstage2_dir.mkdir(exist_ok=True, parents=True)\n\n# Normalization for model input\nmean = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1).to(CUDA0)\nstd = torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1).to(CUDA0)\n\nprint(\"Stage 2: Processing with custom model...\")\nfor sample_id in tqdm(valid_id):\n    path = stage1_dir / f'{sample_id}.png'\n    output_path = stage2_dir / f'{sample_id}.npy'\n    \n    image = cv2.imread(str(path), cv2.IMREAD_COLOR)\n    if image is None:\n        length = valid_df[(valid_df['id']==sample_id) & (valid_df['lead']=='II')].iloc[0].number_of_rows\n        np.save(output_path, np.zeros((4, length)))\n        continue\n    \n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    \n    # Prepare input (resize to model input size)\n    img_tensor = torch.from_numpy(image.transpose(2, 0, 1)).float().unsqueeze(0) / 255.0\n    img_tensor = resize(img_tensor).to(CUDA0)\n    img_tensor = (img_tensor - mean) / std\n    \n    try:\n        with torch.no_grad(), torch.amp.autocast('cuda'):\n            output = custom_model(img_tensor)\n        \n        # Get prediction\n        pixel_mask = torch.sigmoid(output).cpu().numpy()[0]  # (4, H, W)\n        \n        # Resize mask back to original size for accurate pixel-to-mV conversion\n        pixel_mask_full = np.zeros((4, 1696, 4352))\n        for i in range(4):\n            pixel_mask_full[i] = cv2.resize(pixel_mask[i], (4352, 1696))\n        \n        # Get target length\n        length = valid_df[(valid_df['id']==sample_id) & (valid_df['lead']=='II')].iloc[0].number_of_rows\n        \n        # Extract series from mask\n        series_in_pixel = pixel_to_series_custom(pixel_mask_full[..., t0:t1], zero_mv, length)\n        \n        # Convert to mV\n        series = (np.array(zero_mv).reshape(4, 1) - series_in_pixel) / mv_to_pixel\n        \n        # Apply smoothing\n        for i in range(series.shape[0]):\n            series[i] = savgol_filter(series[i], window_length=7, polyorder=2)\n        \n        # Apply medical constraints\n        s_dict = series_to_dict(series)\n        s_dict = refiner.apply_einthoven_law(s_dict)\n        series = dict_to_series(s_dict, series.shape)\n        \n        np.save(output_path, series)\n        \n    except Exception as e:\n        print(f\"Error {sample_id}: {e}\")\n        length = valid_df[(valid_df['id']==sample_id) & (valid_df['lead']=='II')].iloc[0].number_of_rows\n        np.save(output_path, np.zeros((4, length)))\n\nprint(\"Stage 2 complete\")"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Generate submission\ndef series_dict_final(series):\n    d = {}\n    for l in range(3):\n        names = [['I', 'aVR', 'V1', 'V4'], ['II', 'aVL', 'V2', 'V5'], ['III', 'aVF', 'V3', 'V6']][l]\n        splits = np.array_split(series[l], 4)\n        for k, s in zip(names, splits):\n            d[k] = s\n    d['II'] = series[3]  # Use rhythm strip for Lead II\n    return d\n\nstage2_dir = Path(global_dict[\"stage2_dir\"])\nsubmit_df = []\ngb = valid_df.groupby('id')\n\nprint(\"Generating submission...\")\nfor rec_idx, (sample_id, df) in enumerate(tqdm(gb)):\n    try:\n        series = np.load(stage2_dir / f'{sample_id}.npy')\n        series_by_lead = series_dict_final(series)\n        \n        for _, d in df.iterrows():\n            s = series_by_lead.get(d.lead, np.zeros(d.number_of_rows))\n            \n            if len(s) != d.number_of_rows:\n                x_old = np.linspace(0, 1, len(s))\n                x_new = np.linspace(0, 1, d.number_of_rows)\n                s = np.interp(x_new, x_old, s)\n            \n            row_id = [f'{sample_id}_{t}_{d.lead}' for t in range(d.number_of_rows)]\n            submit_df.append(pd.DataFrame({'id': row_id, 'value': s}))\n            \n    except Exception as e:\n        print(f\"Error {sample_id}: {e}\")\n    \n    if rec_idx % 100 == 0:\n        gc.collect()\n\nfinal_df = pd.concat(submit_df, axis=0, ignore_index=True)\nfinal_df.to_csv('submission.csv', index=False)\n\nprint(f\"\\nSubmission saved!\")\nprint(f\"Shape: {final_df.shape}\")\nprint(final_df.head())\nprint(f\"\\nValue stats:\")\nprint(final_df['value'].describe())"}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.0"}},"nbformat":4,"nbformat_minor":4}