{"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 - cc3d compatibility shim using scipy\n# This replaces cc3d.connected_components and cc3d.statistics with scipy equivalents\n\nimport numpy as np\nfrom scipy import ndimage\n\nclass cc3d_compat:\n    \"\"\"cc3d compatibility layer using scipy.ndimage\"\"\"\n    \n    @staticmethod\n    def connected_components(binary_mask):\n        \"\"\"Replace cc3d.connected_components with scipy.ndimage.label\"\"\"\n        labeled, num_features = ndimage.label(binary_mask)\n        return labeled\n    \n    @staticmethod\n    def statistics(labeled_array):\n        \"\"\"Replace cc3d.statistics - compute centroids and voxel_counts\n        \n        Note: cc3d includes background (label 0) at index 0 of output arrays.\n        The calling code uses [1:] to skip background.\n        \"\"\"\n        unique_labels = np.unique(labeled_array)\n        \n        # Initialize with background at index 0\n        centroids = [[0.0, 0.0]]  # Background centroid (not used)\n        voxel_counts = [0]  # Background count (not used)\n        \n        # Process each non-zero label\n        for label in unique_labels:\n            if label == 0:\n                continue  # Skip background\n                \n            mask = labeled_array == label\n            voxel_count = np.sum(mask)\n            \n            # Compute centroid (y, x format like cc3d)\n            coords = np.argwhere(mask)\n            if len(coords) > 0:\n                centroid = coords.mean(axis=0)\n            else:\n                centroid = np.array([0.0, 0.0])\n            \n            centroids.append(centroid.tolist())\n            voxel_counts.append(voxel_count)\n        \n        return {\n            'centroids': np.array(centroids),\n            'voxel_counts': np.array(voxel_counts)\n        }\n\n# Inject into sys.modules so \"import cc3d\" works\nimport sys\nsys.modules['cc3d'] = cc3d_compat\n\nprint(\"Starting offline inference (cc3d shimmed with scipy)...\")"}, {"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: Use baseline model (ResNet34 + custom UNet decoder from timm)\nfrom stage2_model import Net as Stage2Net\nfrom stage2_common import load_net, pixel_to_series\n\n# Load baseline Stage 2 model\nstage2_net = Stage2Net(pretrained=False)\nstage2_net = load_net(stage2_net, f'{model_base}/weight/stage2-00005810.checkpoint.pth')\nstage2_net.to(CUDA0).eval()\n\nprint(f\"Stage 2 model loaded (baseline ResNet34 UNet)\")"}, {"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\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 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\n\nprint(\"Signal processing helpers ready\")"}, {"cell_type": "code", "execution_count": null, "metadata": {}, "outputs": [], "source": "# Stage 2: Process with baseline model (CPU optimized - small resolution)\nstage1_dir = Path(global_dict[\"stage1_dir\"])\nstage2_dir = Path(global_dict[\"stage2_dir\"])\nstage2_dir.mkdir(exist_ok=True, parents=True)\n\n# Use resolution divisible by 16 for UNet (16x downsampling)\n# 432 = 16 * 27, 1104 = 16 * 69\nTARGET_H = 432\nTARGET_W = 1104\n\n# Original expected size (for pixel_to_series which asserts H==1696)\nORIG_H = 1696\nORIG_W = 4352\n\nprint(f\"Stage 2: Processing at {TARGET_W}x{TARGET_H}, upscale to {ORIG_W}x{ORIG_H}...\")\n\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    orig_h, orig_w = image.shape[:2]\n    \n    # Resize to target size for fast inference\n    image_resized = cv2.resize(image, (TARGET_W, TARGET_H), interpolation=cv2.INTER_LINEAR)\n    print(f\"Processing {sample_id}: {orig_w}x{orig_h} -> {TARGET_W}x{TARGET_H}\")\n    \n    # Prepare input batch\n    batch = {'image': torch.from_numpy(image_resized.transpose(2, 0, 1)).unsqueeze(0).float().to(CUDA0)}\n    \n    try:\n        with torch.no_grad():\n            output = stage2_net(batch)\n        \n        # Get prediction mask (ensure float32 for numpy)\n        pixel_mask = output['pixel'].float().cpu().numpy()[0]  # (4, H, W)\n        _, mask_h, mask_w = pixel_mask.shape\n        print(f\"Mask shape: {mask_w}x{mask_h}\")\n        \n        # Resize mask back to original expected size (1696x4352)\n        # pixel_to_series expects H=1696\n        pixel_mask_full = np.zeros((4, ORIG_H, ORIG_W), dtype=np.float32)\n        for c in range(4):\n            pixel_mask_full[c] = cv2.resize(pixel_mask[c], (ORIG_W, ORIG_H), interpolation=cv2.INTER_LINEAR)\n        print(f\"Upscaled mask: {ORIG_W}x{ORIG_H}\")\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        # Use original parameters (no scaling needed - mask is now full size)\n        series_in_pixel = pixel_to_series(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        print(f\"Saved {sample_id}, range: [{series.min():.3f}, {series.max():.3f}]\")\n        \n    except Exception as e:\n        print(f\"Error {sample_id}: {e}\")\n        import traceback\n        traceback.print_exc()\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\ndel stage2_net\ngc.collect()\ntorch.cuda.empty_cache()\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}