{ "cells": [ { "cell_type": "code", "execution_count": null, "id": "32782932", "metadata": {}, "outputs": [], "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" ] }, { "cell_type": "code", "execution_count": null, "id": "9ff8a62d", "metadata": {}, "outputs": [], "source": [ "import os, sys, gc, cv2, numpy as np, pandas as pd, torch\n", "import torch.nn as nn\n", "import torch.nn.functional as F\n", "import torchvision.transforms as T\n", "import timm\n", "from pathlib import Path\n", "from tqdm.auto import tqdm\n", "from scipy.signal import savgol_filter\n", "\n", "# Paths\n", "BASELINE_PATH = '/kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet'\n", "WEIGHTS_PATH = '/kaggle/input/ecg-v10-best/pytorch/pytorch/4'\n", "COMPETITION_PATH = '/kaggle/input/physionet-ecg-image-digitization'\n", "\n", "sys.path.insert(0, BASELINE_PATH)\n", "\n", "# Import baseline stage 0/1/2\n", "import stage0_common as s0c\n", "import stage1_common as s1c\n", "import stage2_common as s2c\n", "from stage0_model import Net as Stage0Net\n", "from stage1_model import Net as Stage1Net\n", "from stage2_model import MyCoordUnetDecoder, encode_with_resnet\n", "\n", "device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n", "print(f\"Device: {device}\")" ] }, { "cell_type": "markdown", "id": "b838e186", "metadata": {}, "source": [ "## Constants" ] }, { "cell_type": "code", "execution_count": null, "id": "384a3691", "metadata": {}, "outputs": [], "source": [ "# V10.1 Constants\n", "TARGET_HEIGHT, TARGET_WIDTH = 1696, 4352\n", "ZERO_MV = np.array([703.5, 987.5, 1271.5, 1531.5])\n", "MV_TO_PIXEL = 78.5\n", "T0, T1 = 235, 4161\n", "X0, X1 = 0, 2176\n", "Y0, Y1 = 0, 1696\n", "OUTPUT_WIDTH = T1 - T0\n", "SOFT_ARGMAX_TEMP = 100.0\n", "\n", "LEAD_NAMES = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', 'V1', 'V2', 'V3', 'V4', 'V5', 'V6']\n", "ROW_LAYOUT = [\n", " ['I', 'aVR', 'V1', 'V4'],\n", " ['II', 'aVL', 'V2', 'V5'],\n", " ['III', 'aVF', 'V3', 'V6'],\n", "]\n", "LEADS_ORDER = [\"I\",\"II\",\"III\",\"aVR\",\"aVL\",\"aVF\",\"V1\",\"V2\",\"V3\",\"V4\",\"V5\",\"V6\"]\n", "\n", "# ECG amplitude limits (mV)\n", "ECG_MV_MIN, ECG_MV_MAX = -7.0, 7.0" ] }, { "cell_type": "markdown", "id": "70fb386f", "metadata": {}, "source": [ "## Model 1: V10.1 (EfficientNet-B4)" ] }, { "cell_type": "code", "execution_count": null, "id": "dbb9493d", "metadata": {}, "outputs": [], "source": [ "class CoordDecoderBlock(nn.Module):\n", " def __init__(self, in_ch, skip_ch, out_ch, scale=2):\n", " super().__init__()\n", " self.scale = scale\n", " self.conv = nn.Sequential(\n", " nn.Conv2d(in_ch + skip_ch + 2, out_ch, 3, padding=1, bias=False),\n", " nn.BatchNorm2d(out_ch),\n", " nn.ReLU(inplace=True),\n", " nn.Conv2d(out_ch, out_ch, 3, padding=1, bias=False),\n", " nn.BatchNorm2d(out_ch),\n", " nn.ReLU(inplace=True),\n", " )\n", "\n", " def forward(self, x, skip=None):\n", " x = F.interpolate(x, scale_factor=self.scale, mode='nearest')\n", " if skip is not None:\n", " x = torch.cat([x, skip], dim=1)\n", " b, c, h, w = x.shape\n", " cy, cx = torch.meshgrid(\n", " torch.linspace(-1, 1, h, device=x.device, dtype=x.dtype),\n", " torch.linspace(-1, 1, w, device=x.device, dtype=x.dtype),\n", " indexing='ij'\n", " )\n", " coord = torch.stack([cx, cy]).unsqueeze(0).expand(b, -1, -1, -1)\n", " x = torch.cat([x, coord], dim=1)\n", " return self.conv(x)\n", "\n", "\n", "class ECGNetV10(nn.Module):\n", " \"\"\"V10 with EfficientNet-B4 encoder.\"\"\"\n", " def __init__(self, encoder='efficientnet_b4', decoder_dims=[256, 128, 64, 32, 16]):\n", " super().__init__()\n", " self.encoder = timm.create_model(\n", " 'efficientnet_b4.ra2_in1k', pretrained=False, \n", " features_only=True, out_indices=(0, 1, 2, 3, 4)\n", " )\n", " enc_dims = [24, 32, 56, 160, 448]\n", " self.enc_dims = enc_dims\n", " \n", " self.dec_blocks = nn.ModuleList()\n", " in_ch = enc_dims[-1]\n", " skip_chs = enc_dims[:-1][::-1] + [0]\n", " while len(decoder_dims) < len(skip_chs):\n", " decoder_dims.append(decoder_dims[-1])\n", " decoder_dims = decoder_dims[:len(skip_chs)]\n", " \n", " for skip_ch, out_ch in zip(skip_chs, decoder_dims):\n", " self.dec_blocks.append(CoordDecoderBlock(in_ch, skip_ch, out_ch))\n", " in_ch = out_ch\n", " \n", " self.seg_head = nn.Conv2d(decoder_dims[-1], 4, 1)\n", " self.reg_head = nn.Sequential(\n", " nn.Conv2d(decoder_dims[-1], 64, 3, padding=1),\n", " nn.ReLU(inplace=True),\n", " nn.AdaptiveAvgPool2d((1, None)),\n", " )\n", " self.reg_out = nn.Sequential(\n", " nn.Conv1d(64, 32, 3, padding=1),\n", " nn.ReLU(inplace=True),\n", " nn.Conv1d(32, 4, 1),\n", " nn.Sigmoid()\n", " )\n", " \n", " def forward(self, x):\n", " input_size = x.shape[2:]\n", " enc = self.encoder(x)\n", " d = enc[-1]\n", " skips = enc[:-1][::-1] + [None]\n", " for block, skip in zip(self.dec_blocks, skips):\n", " d = block(d, skip)\n", " if d.shape[2:] != input_size:\n", " d = F.interpolate(d, size=input_size, mode='bilinear', align_corners=False)\n", " seg_logits = self.seg_head(d)\n", " reg_feat = self.reg_head(d).squeeze(2)\n", " reg_coords = self.reg_out(reg_feat)\n", " return seg_logits, reg_coords" ] }, { "cell_type": "markdown", "id": "23288d67", "metadata": {}, "source": [ "## Model 2: Net3 Pipeline (Public Solution)" ] }, { "cell_type": "code", "execution_count": null, "id": "bf56a039", "metadata": {}, "outputs": [], "source": [ "def 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", "\n", "def 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", "\n", "def 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", "\n", "def 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", "\n", "def 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", "\n", "def 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", "\n", "def 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", "\n", "def 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", "\n", "def 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)" ] }, { "cell_type": "code", "execution_count": null, "id": "57d5c7ef", "metadata": {}, "outputs": [], "source": [ "class Net3(nn.Module):\n", " def __init__(self, pretrained=True):\n", " super().__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',\n", " pretrained=pretrained,\n", " in_chans=3,\n", " num_classes=0,\n", " 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", " return self.pixel(last)\n", "\n", "\n", "class PhysioPipeline:\n", " def __init__(self, device=\"cuda:0\"):\n", " self.device = device\n", " self.stage0_net = self.stage1_net = self.stage2_net = None\n", " self.x0, self.x1 = 0, 2176\n", " self.y0, self.y1 = 0, 1696\n", " self.zero_mv = [703.5, 987.5, 1271.5, 1531.5]\n", " self.mv_to_pixel = 78.8\n", " self.t0, self.t1 = 235, 4161\n", " self.resize = T.Resize((1696, 4352), interpolation=T.InterpolationMode.BILINEAR)\n", "\n", " def load_models(self, stage0_w, stage1_w, stage2_w):\n", " self.stage0_net = s0c.load_net(Stage0Net(pretrained=False), stage0_w).to(self.device).eval()\n", " self.stage1_net = s1c.load_net(Stage1Net(pretrained=False), stage1_w).to(self.device).eval()\n", " self.stage2_net = Net3(pretrained=False).to(self.device).eval()\n", " st = torch.load(stage2_w, map_location=\"cpu\")\n", " if isinstance(st, dict) and \"state_dict\" in st: st = st[\"state_dict\"]\n", " self.stage2_net.load_state_dict(st, strict=True)\n", "\n", " def run_stage0(self, img_bgr):\n", " img_rgb = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB)\n", " img_for_model = change_color(img_rgb)\n", " batch = s0c.image_to_batch(img_for_model)\n", " with torch.no_grad(), torch.amp.autocast(self.device.split(\":\")[0], dtype=torch.float32):\n", " output = self.stage0_net(batch)\n", " rotated, keypoint = s0c.output_to_predict(img_rgb, batch, output)\n", " normalised, _, _ = s0c.normalise_by_homography(rotated, keypoint)\n", " return normalised\n", "\n", " def run_stage1(self, stage0_img_rgb):\n", " image = stage0_img_rgb\n", " batch = {'image': torch.from_numpy(np.ascontiguousarray(image.transpose(2, 0, 1))).unsqueeze(0)}\n", " with torch.no_grad(), torch.amp.autocast(self.device.split(\":\")[0], dtype=torch.float32):\n", " output = self.stage1_net(batch)\n", " gridpoint_xy, _ = s1c.output_to_predict(image, batch, output)\n", " return s1c.rectify_image(image, gridpoint_xy)\n", "\n", " def run_stage2_raw(self, stage1_img_rgb, length):\n", " \"\"\"Returns raw 4-row mV signal without savgol smoothing.\"\"\"\n", " img = stage1_img_rgb[self.y0:self.y1, self.x0:self.x1] / 255.0\n", " batch = self.resize(torch.from_numpy(np.ascontiguousarray(img.transpose(2, 0, 1))).unsqueeze(0)).float().to(self.device)\n", " with torch.no_grad(), torch.amp.autocast(self.device.split(\":\")[0], dtype=torch.float32):\n", " output = self.stage2_net(batch)\n", " pixel = torch.sigmoid(output).float().cpu().numpy()[0]\n", " series_in_pixel = s2c.pixel_to_series(pixel[..., self.t0:self.t1], self.zero_mv, length)\n", " series = (np.array(self.zero_mv).reshape(4, 1) - series_in_pixel) / self.mv_to_pixel\n", " return series # No smoothing - we'll smooth after ensemble" ] }, { "cell_type": "code", "execution_count": null, "id": "392653e8", "metadata": {}, "outputs": [], "source": [ "# Classifier for source detection\n", "CLS_MODEL_NAME = \"efficientnet_b2\"\n", "CLS_NUM_CLASSES = 12\n", "CLS_RESOLUTION = 256\n", "CLS_CKPT_PATH = \"/kaggle/input/physionet-image-multi-class-train/efficientnet_b2_full_train.pth\"\n", "\n", "def build_classifier(device=\"cuda\"):\n", " m = timm.create_model(CLS_MODEL_NAME, pretrained=False, num_classes=CLS_NUM_CLASSES)\n", " st = torch.load(CLS_CKPT_PATH, map_location=\"cpu\")\n", " if isinstance(st, dict) and \"state_dict\" in st: st = st[\"state_dict\"]\n", " st2 = {k.replace(\"module.\",\"\"): v for k, v in st.items()}\n", " m.load_state_dict(st2, strict=False)\n", " return m.to(device).eval()\n", "\n", "def cls_preprocess_bgr(img_bgr):\n", " img = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB)\n", " img = cv2.resize(img, (CLS_RESOLUTION, CLS_RESOLUTION), interpolation=cv2.INTER_AREA).astype(np.float32)/255.0\n", " mean = np.array([0.485,0.456,0.406], np.float32)\n", " 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", "\n", "@torch.no_grad()\n", "def 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}\"\n", "\n", "def select_stage1_with_source(pipeline, img_raw_bgr, pred_source_suffix, selector_margin=1.02):\n", " img_pp = preprocess_by_source(img_raw_bgr.copy(), pred_source_suffix)\n", " s1_raw = pipeline.run_stage1(pipeline.run_stage0(img_raw_bgr))\n", " q_raw = stage1_quality(s1_raw)\n", " s1_pp = pipeline.run_stage1(pipeline.run_stage0(img_pp))\n", " q_pp = stage1_quality(s1_pp)\n", " return s1_pp if q_pp > q_raw * selector_margin else s1_raw" ] }, { "cell_type": "markdown", "id": "9239905e", "metadata": {}, "source": [ "## V10.1 Inference Functions" ] }, { "cell_type": "code", "execution_count": null, "id": "3979d700", "metadata": {}, "outputs": [], "source": [ "def soft_argmax(heatmap, temperature=SOFT_ARGMAX_TEMP):\n", " \"\"\"Extract sub-pixel coordinates using soft-argmax.\"\"\"\n", " B, C, H, W = heatmap.shape\n", " y_coords = torch.arange(H, device=heatmap.device, dtype=heatmap.dtype).view(1, 1, H, 1)\n", " weights = F.softmax(heatmap * temperature, dim=2)\n", " return (weights * y_coords).sum(dim=2)\n", "\n", "\n", "def interpolate_nan(signal_1d):\n", " \"\"\"Interpolate NaN values from valid neighbors.\"\"\"\n", " valid_mask = np.isfinite(signal_1d)\n", " if valid_mask.all():\n", " return signal_1d\n", " if not valid_mask.any():\n", " return np.zeros_like(signal_1d)\n", " x = np.arange(len(signal_1d))\n", " signal_1d[~valid_mask] = np.interp(x[~valid_mask], x[valid_mask], signal_1d[valid_mask])\n", " return signal_1d\n", "\n", "\n", "@torch.no_grad()\n", "def process_v10_stage2(model, image_bgr):\n", " \"\"\"Stage 2: Signal extraction using V10 model. Returns 4-row mV signal.\"\"\"\n", " h, w = image_bgr.shape[:2]\n", " crop_h = min(h, Y1)\n", " crop_w = min(w, X1)\n", " image_cropped = image_bgr[:crop_h, :crop_w]\n", " image_resized = cv2.resize(image_cropped, (TARGET_WIDTH, TARGET_HEIGHT), interpolation=cv2.INTER_LINEAR)\n", " \n", " image_tensor = torch.from_numpy(image_resized.astype(np.float32) / 255.0).permute(2, 0, 1).unsqueeze(0)\n", " image_tensor = image_tensor.to(device)\n", " \n", " with torch.amp.autocast('cuda', dtype=torch.float32):\n", " seg_logits, reg_coords = model(image_tensor)\n", " \n", " seg_logits = torch.nan_to_num(seg_logits, nan=0.0, posinf=0.0, neginf=0.0)\n", " seg_probs = torch.sigmoid(seg_logits.float())\n", " signal_full = soft_argmax(seg_probs).cpu().numpy()[0] # [4, 4352]\n", " \n", " # Handle NaN\n", " for row_idx in range(4):\n", " if not np.isfinite(signal_full[row_idx]).all():\n", " signal_full[row_idx] = interpolate_nan(signal_full[row_idx].copy())\n", " remaining_nan = ~np.isfinite(signal_full[row_idx])\n", " if remaining_nan.any():\n", " signal_full[row_idx, remaining_nan] = ZERO_MV[row_idx]\n", " \n", " # Extract signal region and convert to mV\n", " signal_pixel = signal_full[:, T0:T1] # [4, OUTPUT_WIDTH]\n", " signal_mv = np.zeros_like(signal_pixel)\n", " for row_idx in range(4):\n", " signal_mv[row_idx] = (ZERO_MV[row_idx] - signal_pixel[row_idx]) / MV_TO_PIXEL\n", " signal_mv = np.clip(signal_mv, ECG_MV_MIN, ECG_MV_MAX)\n", " \n", " return signal_mv # [4, 3926]" ] }, { "cell_type": "markdown", "id": "0724ab36", "metadata": {}, "source": [ "## Load All Models" ] }, { "cell_type": "code", "execution_count": null, "id": "17778392", "metadata": {}, "outputs": [], "source": [ "# Load Stage 0/1 (shared)\n", "print(\"Loading Stage 0...\")\n", "stage0_net = s0c.load_net(Stage0Net(pretrained=False), f'{BASELINE_PATH}/weight/stage0-last.checkpoint.pth')\n", "stage0_net.to(device).eval()\n", "\n", "print(\"Loading Stage 1...\")\n", "stage1_net = s1c.load_net(Stage1Net(pretrained=False), f'{BASELINE_PATH}/weight/stage1-last.checkpoint.pth')\n", "stage1_net.to(device).eval()\n", "\n", "# Load V10.1 model\n", "print(\"Loading V10.1 model...\")\n", "v10_model = ECGNetV10(encoder='efficientnet_b4')\n", "checkpoint = torch.load(f'{WEIGHTS_PATH}/ecg_v10_best.pth', map_location='cpu', weights_only=False)\n", "state_dict = checkpoint['model']\n", "if list(state_dict.keys())[0].startswith('module.'):\n", " state_dict = {k[7:]: v for k, v in state_dict.items()}\n", "v10_model.load_state_dict(state_dict)\n", "v10_model.to(device).eval()\n", "print(f\"V10.1 loaded - epoch {checkpoint.get('epoch', '?')}, SNR: {checkpoint.get('holdout_snr_soft', 0):.2f} dB\")\n", "\n", "# Load Net3 pipeline\n", "print(\"Loading Net3 pipeline...\")\n", "pipeline = PhysioPipeline(device=\"cuda:0\" if device.type == \"cuda\" else \"cpu\")\n", "pipeline.load_models(\n", " stage0_w=f\"{BASELINE_PATH}/weight/stage0-last.checkpoint.pth\",\n", " stage1_w=f\"{BASELINE_PATH}/weight/stage1-last.checkpoint.pth\",\n", " stage2_w=\"/kaggle/input/physio-seg-public/pytorch/net3_009_4200/1/iter_0004200.pt\",\n", ")\n", "\n", "# Load classifier\n", "print(\"Loading source classifier...\")\n", "cls_model = build_classifier(device=str(device))\n", "\n", "print(\"All models loaded!\")" ] }, { "cell_type": "markdown", "id": "8566363c", "metadata": {}, "source": [ "## Ensemble Functions" ] }, { "cell_type": "code", "execution_count": null, "id": "587e9985", "metadata": {}, "outputs": [], "source": [ "def resample_signal(signal, target_length):\n", " \"\"\"Resample signal to target length.\"\"\"\n", " if len(signal) == target_length:\n", " return signal\n", " x_old = np.linspace(0, 1, len(signal))\n", " x_new = np.linspace(0, 1, target_length)\n", " return np.interp(x_new, x_old, signal)\n", "\n", "\n", "def series_to_leads_v10(series_mv):\n", " \"\"\"Convert 4-row series to 12-lead dictionary (V10 format).\"\"\"\n", " leads = {}\n", " segment_width = series_mv.shape[1] // 4\n", " \n", " for row_idx in range(3):\n", " for seg_idx, lead_name in enumerate(ROW_LAYOUT[row_idx]):\n", " start = seg_idx * segment_width\n", " end = (seg_idx + 1) * segment_width\n", " leads[lead_name] = series_mv[row_idx, start:end]\n", " \n", " leads['II'] = series_mv[3] # Full rhythm strip\n", " return leads\n", "\n", "\n", "def series_to_leads_net3(series_4row):\n", " \"\"\"Convert 4-row series to 12-lead dictionary (Net3 format).\"\"\"\n", " series_4row = np.asarray(series_4row)\n", " if series_4row.ndim == 3: \n", " series_4row = series_4row[0]\n", " if series_4row.shape[0] != 4 and series_4row.shape[1] == 4: \n", " 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)\n", " return d\n", "\n", "\n", "def ensemble_leads(leads_v10, leads_net3, weight_v10=0.5):\n", " \"\"\"Ensemble two lead dictionaries with weighted average.\"\"\"\n", " weight_net3 = 1.0 - weight_v10\n", " ensemble = {}\n", " \n", " for lead in LEADS_ORDER:\n", " if lead == 'II':\n", " sig_v10 = leads_v10.get('II', np.zeros(1))\n", " sig_net3 = leads_net3.get('II', np.zeros(1))\n", " else:\n", " sig_v10 = leads_v10.get(lead, np.zeros(1))\n", " sig_net3 = leads_net3.get(lead, np.zeros(1))\n", " \n", " # Resample to same length (use longer one)\n", " target_len = max(len(sig_v10), len(sig_net3))\n", " if target_len == 0:\n", " target_len = 1\n", " \n", " sig_v10 = resample_signal(sig_v10, target_len) if len(sig_v10) > 0 else np.zeros(target_len)\n", " sig_net3 = resample_signal(sig_net3, target_len) if len(sig_net3) > 0 else np.zeros(target_len)\n", " \n", " # Weighted average\n", " ensemble[lead] = weight_v10 * sig_v10 + weight_net3 * sig_net3\n", " \n", " return ensemble\n", "\n", "\n", "def ensemble_leads_adaptive(leads_v10, leads_net3, base_weight=0.5, shrink_factor=2.0):\n", " \"\"\"\n", " Adaptive ensemble based on model disagreement.\n", " When V10 and Net3 agree → use standard weighted average\n", " When they disagree → shrink prediction toward 0 (conservative)\n", " \n", " Returns:\n", " ensemble: Dict of ensembled signals\n", " confidence: Dict of per-point confidence scores (0-1)\n", " \"\"\"\n", " weight_net3 = 1.0 - base_weight\n", " ensemble = {}\n", " confidence = {}\n", " \n", " for lead in LEADS_ORDER:\n", " if lead == 'II':\n", " sig_v10 = leads_v10.get('II', np.zeros(1))\n", " sig_net3 = leads_net3.get('II', np.zeros(1))\n", " else:\n", " sig_v10 = leads_v10.get(lead, np.zeros(1))\n", " sig_net3 = leads_net3.get(lead, np.zeros(1))\n", " \n", " # Resample to same length\n", " target_len = max(len(sig_v10), len(sig_net3))\n", " if target_len == 0:\n", " target_len = 1\n", " \n", " sig_v10 = resample_signal(sig_v10, target_len) if len(sig_v10) > 0 else np.zeros(target_len)\n", " sig_net3 = resample_signal(sig_net3, target_len) if len(sig_net3) > 0 else np.zeros(target_len)\n", " \n", " # Compute weighted average (base prediction)\n", " mean_pred = base_weight * sig_v10 + weight_net3 * sig_net3\n", " \n", " # Compute per-point disagreement\n", " disagreement = np.abs(sig_v10 - sig_net3)\n", " \n", " # Normalize by signal scale\n", " combined_std = np.std(np.concatenate([sig_v10, sig_net3]))\n", " if combined_std > 0.01:\n", " disagreement = disagreement / combined_std\n", " \n", " # Compute confidence: high when models agree, low when they disagree\n", " conf = 1.0 / (1.0 + shrink_factor * disagreement)\n", " \n", " # Apply confidence: shrink toward 0 when uncertain\n", " ensemble[lead] = mean_pred * conf\n", " confidence[lead] = conf\n", " \n", " return ensemble, confidence\n", "\n", "\n", "def apply_smoothing(leads_dict, window=7, polyorder=2):\n", " \"\"\"Apply Savitzky-Golay smoothing to all leads.\"\"\"\n", " smoothed = {}\n", " for lead, signal in leads_dict.items():\n", " if len(signal) >= window:\n", " smoothed[lead] = savgol_filter(signal, window_length=window, polyorder=polyorder)\n", " else:\n", " smoothed[lead] = signal\n", " return smoothed" ] }, { "cell_type": "markdown", "id": "e2229731", "metadata": {}, "source": [ "## Shared Stage 0/1 Processing" ] }, { "cell_type": "code", "execution_count": null, "id": "4172932f", "metadata": {}, "outputs": [], "source": [ "@torch.no_grad()\n", "def process_stage0(image_rgb):\n", " \"\"\"Stage 0: Orientation correction.\"\"\"\n", " batch = s0c.image_to_batch(image_rgb)\n", " with torch.amp.autocast('cuda', dtype=torch.float32):\n", " output = stage0_net(batch)\n", " rotated, keypoint = s0c.output_to_predict(image_rgb, batch, output)\n", " normalized, _, _ = s0c.normalise_by_homography(rotated, keypoint)\n", " return normalized\n", "\n", "\n", "@torch.no_grad()\n", "def process_stage1(image_rgb):\n", " \"\"\"Stage 1: Grid rectification.\"\"\"\n", " batch = {'image': torch.from_numpy(np.ascontiguousarray(image_rgb.transpose(2, 0, 1))).unsqueeze(0)}\n", " with torch.amp.autocast('cuda', dtype=torch.float32):\n", " output = stage1_net(batch)\n", " gridpoint_xy, _ = s1c.output_to_predict(image_rgb, batch, output)\n", " rectified = s1c.rectify_image(image_rgb, gridpoint_xy)\n", " return rectified" ] }, { "cell_type": "markdown", "id": "e0c1b424", "metadata": {}, "source": [ "## Generate Ensemble Submission" ] }, { "cell_type": "code", "execution_count": null, "id": "7eab55d7", "metadata": {}, "outputs": [], "source": [ "# Ensemble configuration\n", "WEIGHT_V10 = 0.5 # Base weight for V10.1 (50/50 average)\n", "USE_ADAPTIVE = True # Use adaptive ensemble (shrink-on-disagreement)\n", "SHRINK_FACTOR = 2.0 # How strongly to shrink toward 0 when models disagree\n", "\n", "# Load test metadata\n", "test_df = pd.read_csv(f'{COMPETITION_PATH}/test.csv')\n", "test_dir = Path(f'{COMPETITION_PATH}/test')\n", "image_ids = test_df['id'].unique()\n", "\n", "print(f\"Processing {len(image_ids)} images...\")\n", "print(f\"Base ensemble weights: V10.1={WEIGHT_V10:.0%}, Net3={1-WEIGHT_V10:.0%}\")\n", "if USE_ADAPTIVE:\n", " print(f\"ADAPTIVE MODE: Shrink factor = {SHRINK_FACTOR} (higher = more conservative on disagreement)\")" ] }, { "cell_type": "markdown", "id": "5faff850", "metadata": {}, "source": [ "## Adaptive Ensemble Visualization\n", "\n", "This cell demonstrates how the adaptive ensemble works on a sample image." ] }, { "cell_type": "code", "execution_count": null, "id": "3d23ebe7", "metadata": {}, "outputs": [], "source": [ "# Demo: Compare fixed vs adaptive ensemble on a sample\n", "def demo_adaptive_ensemble(leads_v10, leads_net3, lead='II'):\n", " \"\"\"Visualize how adaptive ensemble differs from fixed ensemble.\"\"\"\n", " import matplotlib.pyplot as plt\n", " \n", " # Get signals\n", " sig_v10 = leads_v10.get(lead, np.zeros(100))\n", " sig_net3 = leads_net3.get(lead, np.zeros(100))\n", " \n", " # Resample to same length\n", " target_len = max(len(sig_v10), len(sig_net3))\n", " sig_v10 = resample_signal(sig_v10, target_len)\n", " sig_net3 = resample_signal(sig_net3, target_len)\n", " \n", " # Fixed ensemble (50/50)\n", " fixed = 0.5 * sig_v10 + 0.5 * sig_net3\n", " \n", " # Adaptive ensemble\n", " disagreement = np.abs(sig_v10 - sig_net3)\n", " combined_std = np.std(np.concatenate([sig_v10, sig_net3]))\n", " norm_disagreement = disagreement / max(combined_std, 0.01)\n", " confidence = 1.0 / (1.0 + SHRINK_FACTOR * norm_disagreement)\n", " adaptive = fixed * confidence\n", " \n", " # Plot\n", " fig, axes = plt.subplots(4, 1, figsize=(14, 10), sharex=True)\n", " \n", " x = np.arange(len(sig_v10))\n", " \n", " # Models\n", " axes[0].plot(x, sig_v10, 'b-', alpha=0.7, label='V10.1')\n", " axes[0].plot(x, sig_net3, 'r-', alpha=0.7, label='Net3')\n", " axes[0].set_ylabel('mV')\n", " axes[0].set_title(f'Lead {lead}: Model Predictions')\n", " axes[0].legend()\n", " axes[0].grid(True, alpha=0.3)\n", " \n", " # Disagreement\n", " axes[1].fill_between(x, 0, norm_disagreement, alpha=0.5, color='orange')\n", " axes[1].set_ylabel('Disagreement')\n", " axes[1].set_title('Normalized Disagreement |V10 - Net3| / std')\n", " axes[1].grid(True, alpha=0.3)\n", " \n", " # Confidence\n", " axes[2].fill_between(x, 0, confidence, alpha=0.5, color='green')\n", " axes[2].set_ylabel('Confidence')\n", " axes[2].set_title(f'Confidence = 1/(1 + {SHRINK_FACTOR}×disagreement)')\n", " axes[2].set_ylim(0, 1.1)\n", " axes[2].axhline(0.5, color='red', linestyle='--', alpha=0.5, label='50% threshold')\n", " axes[2].grid(True, alpha=0.3)\n", " axes[2].legend()\n", " \n", " # Compare ensembles\n", " axes[3].plot(x, fixed, 'b-', alpha=0.7, linewidth=2, label='Fixed Ensemble')\n", " axes[3].plot(x, adaptive, 'g-', alpha=0.7, linewidth=2, label='Adaptive Ensemble')\n", " axes[3].fill_between(x, fixed, adaptive, where=np.abs(fixed-adaptive)>0.01, \n", " alpha=0.3, color='yellow', label='Shrinkage')\n", " axes[3].set_xlabel('Sample')\n", " axes[3].set_ylabel('mV')\n", " axes[3].set_title('Fixed vs Adaptive Ensemble (yellow = shrinkage toward 0)')\n", " axes[3].legend()\n", " axes[3].grid(True, alpha=0.3)\n", " \n", " plt.tight_layout()\n", " plt.show()\n", " \n", " # Stats\n", " shrinkage = np.abs(fixed - adaptive)\n", " print(f\"\\n=== Adaptive Ensemble Stats for Lead {lead} ===\")\n", " print(f\"Mean confidence: {np.mean(confidence):.3f}\")\n", " print(f\"Min confidence: {np.min(confidence):.3f}\")\n", " print(f\"Mean shrinkage: {np.mean(shrinkage):.4f} mV\")\n", " print(f\"Max shrinkage: {np.max(shrinkage):.4f} mV\")\n", " \n", " return confidence\n", "\n", "# Skip this demo for now - will run after main inference if needed\n", "# demo_adaptive_ensemble(leads_v10, leads_net3, lead='II')" ] }, { "cell_type": "code", "execution_count": null, "id": "7b6df4af", "metadata": {}, "outputs": [], "source": [ "all_rows = []\n", "all_confidences = [] # Track confidence for analysis\n", "\n", "for img_id in tqdm(image_ids):\n", " img_path = test_dir / f\"{img_id}.png\"\n", " \n", " if not img_path.exists():\n", " print(f\"Missing: {img_path}\")\n", " continue\n", " \n", " img_df = test_df[test_df['id'] == img_id]\n", " fs = int(img_df['fs'].iloc[0])\n", " sig_len = int(img_df.loc[img_df['lead'] == 'II', 'number_of_rows'].iloc[0])\n", " \n", " try:\n", " # Load image\n", " img_bgr = cv2.imread(str(img_path))\n", " img_rgb = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB)\n", " \n", " # Predict source for preprocessing\n", " pred_src = predict_source_suffix(cls_model, img_bgr, device=str(device))\n", " \n", " # ===== V10.1 Pipeline =====\n", " try:\n", " normalized = process_stage0(img_rgb)\n", " except:\n", " normalized = img_rgb\n", " try:\n", " rectified_rgb = process_stage1(normalized)\n", " except:\n", " rectified_rgb = normalized\n", " rectified_bgr = cv2.cvtColor(rectified_rgb, cv2.COLOR_RGB2BGR)\n", " series_v10 = process_v10_stage2(v10_model, rectified_bgr)\n", " leads_v10 = series_to_leads_v10(series_v10)\n", " \n", " # ===== Net3 Pipeline =====\n", " s1_net3 = select_stage1_with_source(pipeline, img_bgr, pred_src, selector_margin=1.02)\n", " series_net3 = pipeline.run_stage2_raw(s1_net3, length=sig_len)\n", " leads_net3 = series_to_leads_net3(series_net3)\n", " \n", " # ===== Ensemble =====\n", " if USE_ADAPTIVE:\n", " # Adaptive ensemble: shrink toward 0 when models disagree\n", " leads_ensemble, confidence = ensemble_leads_adaptive(\n", " leads_v10, leads_net3, \n", " base_weight=WEIGHT_V10, \n", " shrink_factor=SHRINK_FACTOR\n", " )\n", " # Track mean confidence per image for analysis\n", " mean_conf = np.mean([np.mean(c) for c in confidence.values()])\n", " all_confidences.append({'id': img_id, 'mean_confidence': mean_conf})\n", " else:\n", " # Standard fixed-weight ensemble\n", " leads_ensemble = ensemble_leads(leads_v10, leads_net3, weight_v10=WEIGHT_V10)\n", " \n", " leads_ensemble = apply_smoothing(leads_ensemble, window=7, polyorder=2)\n", " \n", " except Exception as e:\n", " print(f\"Error {img_id}: {e}\")\n", " # Fallback to zeros\n", " for _, row in img_df.iterrows():\n", " for i in range(row['number_of_rows']):\n", " all_rows.append({'id': f\"{img_id}_{i}_{row['lead']}\", 'value': 0.0})\n", " continue\n", " \n", " # Create submission rows\n", " for _, row in img_df.iterrows():\n", " lead = row['lead']\n", " num_samples = row['number_of_rows']\n", " \n", " if lead in leads_ensemble:\n", " signal = resample_signal(leads_ensemble[lead], num_samples)\n", " for i, val in enumerate(signal):\n", " all_rows.append({'id': f\"{img_id}_{i}_{lead}\", 'value': float(val)})\n", " else:\n", " for i in range(num_samples):\n", " all_rows.append({'id': f\"{img_id}_{i}_{lead}\", 'value': 0.0})\n", " \n", " gc.collect()\n", "\n", "print(f\"\\nTotal rows: {len(all_rows)}\")\n", "\n", "# Report confidence stats if adaptive mode\n", "if USE_ADAPTIVE and all_confidences:\n", " conf_df = pd.DataFrame(all_confidences)\n", " print(f\"\\n=== Adaptive Ensemble Confidence Stats ===\")\n", " print(f\"Mean confidence: {conf_df['mean_confidence'].mean():.3f}\")\n", " print(f\"Min confidence: {conf_df['mean_confidence'].min():.3f}\")\n", " print(f\"Max confidence: {conf_df['mean_confidence'].max():.3f}\")\n", " print(f\"\\nLow confidence images (< 0.5):\")\n", " low_conf = conf_df[conf_df['mean_confidence'] < 0.5].sort_values('mean_confidence')\n", " print(low_conf.head(10) if len(low_conf) > 0 else \" None\")" ] }, { "cell_type": "code", "execution_count": null, "id": "730aa224", "metadata": {}, "outputs": [], "source": [ "# Create submission\n", "submission_df = pd.DataFrame(all_rows)\n", "\n", "# Check for NaN/Inf values\n", "nan_count = submission_df['value'].isna().sum()\n", "inf_count = (~np.isfinite(submission_df['value'])).sum() - nan_count\n", "print(f\"NaN values: {nan_count}, Inf values: {inf_count}\")\n", "\n", "if nan_count > 0 or inf_count > 0:\n", " print(\"WARNING: Found NaN/Inf values! Replacing with 0.0\")\n", " submission_df['value'] = submission_df['value'].replace([np.inf, -np.inf], 0.0)\n", " submission_df['value'] = submission_df['value'].fillna(0.0)\n", "\n", "# Verify\n", "assert submission_df['value'].isna().sum() == 0, \"Still have NaN values!\"\n", "assert np.isfinite(submission_df['value']).all(), \"Still have Inf values!\"\n", "\n", "submission_df.to_csv('/kaggle/working/submission.csv', index=False)\n", "\n", "print(f\"\\nSubmission saved!\")\n", "print(f\"Shape: {submission_df.shape}\")\n", "print(submission_df.head(10))\n", "print(f\"\\nValue stats:\")\n", "print(submission_df['value'].describe())" ] } ], "metadata": { "language_info": { "name": "python" } }, "nbformat": 4, "nbformat_minor": 5 }