{"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,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":13746387,"sourceType":"datasetVersion","datasetId":8747012},{"sourceId":13816899,"sourceType":"datasetVersion","datasetId":8620533},{"sourceId":14537986,"sourceType":"datasetVersion","datasetId":9285392},{"sourceId":14541538,"sourceType":"datasetVersion","datasetId":9287793},{"sourceId":272137252,"sourceType":"kernelVersion"},{"sourceId":292594549,"sourceType":"kernelVersion"},{"sourceId":292669675,"sourceType":"kernelVersion"},{"sourceId":292730384,"sourceType":"kernelVersion"},{"sourceId":292893503,"sourceType":"kernelVersion"},{"sourceId":677607,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":513841,"modelId":528480},{"sourceId":724677,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":551493,"modelId":564084}],"dockerImageVersionId":31154,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# End-to-end summary (final 98th place solution)\n\n**Pipeline overview (training → inference):**\n1. **Train source-suffix classifier** on training images to predict the acquisition source (0001–0012).\n2. **Inference on test images**: optional source‑aware preprocessing → Stage 0 (alignment) → Stage 1 (grid rectification) → Stage 2 (Net3 signal extraction) → submission CSV.\n\n**Key results:** Public 18.44 / Private 18.33 (98th place).\n\n**Why it works (high level):**\n- Stage 0/1 explicitly correct camera geometry and grid alignment, which is the dominant error source in this competition.\n- Stage 2 (Net3) focuses on waveform extraction once the geometry is consistent.\n- Source‑aware preprocessing reduces grid/lead distortion on out‑of‑distribution sources, improving SNR and lead consistency.\n\n**Why `eca_nfnet_l0` works well here:**\n- **Strong texture + edge sensitivity** for ECG grids and lead patterns, with stable feature scales (NFNet design).\n- **ECA attention** improves channel selection without heavy overhead, which helps distinguish subtle source‑style differences.\n- Performs well at **small resolution (256)** with high class imbalance (source classes), making it efficient and accurate for the classifier.\n\n**Data flow (simplified):**\n`test.png → source classifier → preprocess_by_source → stage0 → stage1 → stage2 → series → submission.csv`\n\n\ncredits:\n\nhengck23’s excellent solution (demo submission) - 16.1 baseline: https://www.kaggle.com/code/hengck23/demo-submission\n\nwasupandceacar’s Net3 pipeline - 17.75 baseline: https://www.kaggle.com/code/wasupandceacar/physio-v2-3-public\n\nSaner Turhaner ’s Visual QA - 18.17 baseline: https://www.kaggle.com/code/sanpier/visual-qa-for-all-stages-of-ecg-digitization\n\nOrginal base inference code - 18.21 https://www.kaggle.com/code/tonylica/physionet-ecg-streamlined-inference","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":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile physionet_ensemble_pipeline.py\nimport argparse\nimport gc\nimport os\nimport re\nimport shutil\nimport subprocess\nimport sys\nimport tempfile\nfrom dataclasses import dataclass\nfrom pathlib import Path\nfrom typing import Dict, List, Optional, Tuple\n\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.transforms as T\nfrom scipy.signal import savgol_filter\n\ntry:\n    import timm\nexcept Exception as exc:\n    raise RuntimeError(\"timm is required. Install timm before running.\") from exc\n\ntry:\n    import albumentations as A\n    from albumentations.pytorch import ToTensorV2\nexcept Exception:\n    A = None\n    ToTensorV2 = None\n\nLEADS_ORDER = [\"I\", \"II\", \"III\", \"aVR\", \"aVL\", \"aVF\", \"V1\", \"V2\", \"V3\", \"V4\", \"V5\", \"V6\"]\nSOURCE_SUFFIXES = [f\"{i:04d}\" for i in range(1, 13)]\n\nDEFAULT_STAGE0_WEIGHTS = (\n    \"/kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet/\"\n    \"weight/stage0-last.checkpoint.pth\"\n)\nDEFAULT_STAGE1_WEIGHTS = (\n    \"/kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet/\"\n    \"weight/stage1-last.checkpoint.pth\"\n)\nDEFAULT_STAGE2_WEIGHTS = (\n    \"/kaggle/input/physio-seg-public/pytorch/net3_009_4200/1/iter_0004200.pt\"\n)\nDEFAULT_HENGCK_ROOT = \"/kaggle/input/hengck23-submit-physionet/hengck23-submit-physionet\"\n\n\nY_SHIFT_RATIO = {\n    \"I\": 12.6 / 21.59,\n    \"II\": 9 / 21.59,\n    \"III\": 5.4 / 21.59,\n    \"aVR\": 12.6 / 21.59,\n    \"aVL\": 9 / 21.59,\n    \"aVF\": 5.4 / 21.59,\n    \"V1\": 12.59 / 21.59,\n    \"V2\": 9 / 21.59,\n    \"V3\": 5.4 / 21.59,\n    \"V4\": 12.59 / 21.59,\n    \"V5\": 9 / 21.59,\n    \"V6\": 5.4 / 21.59,\n    \"full\": 2.1 / 21.59,\n}\n\nLEAD_LABEL_MAPPING = {\n    \"I\": 1,\n    \"II\": 2,\n    \"III\": 3,\n    \"aVR\": 4,\n    \"aVL\": 5,\n    \"aVF\": 6,\n    \"V1\": 7,\n    \"V2\": 8,\n    \"V3\": 9,\n    \"V4\": 10,\n    \"V5\": 11,\n    \"V6\": 12,\n}\n\n\ndef change_color(image_rgb: np.ndarray) -> np.ndarray:\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\ndef series_dict(series_4row: np.ndarray) -> Dict[str, np.ndarray]:\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: Dict[str, np.ndarray] = {}\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\ndef dw(series: Dict[str, np.ndarray], alpha: float = 0.33) -> Dict[str, np.ndarray]:\n    if all(k in series for k in [\"I\", \"II_short\", \"III\"]):\n        l1, l2s, l3 = series[\"I\"], series[\"II_short\"], series[\"III\"]\n        error = l2s - (l1 + l3)\n        series[\"I\"] = l1 + alpha * error\n        series[\"III\"] = l3 + alpha * error\n        series[\"II_short\"] = l2s - alpha * error\n    return series\n\n\ndef clahe_luminance_bgr(img_bgr: np.ndarray, clip: float = 2.0, tile: int = 8) -> np.ndarray:\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\ndef grayworld_white_balance(img_bgr: np.ndarray) -> np.ndarray:\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)\n    g *= m / (mg + 1e-6)\n    r *= m / (mr + 1e-6)\n    return np.clip(cv2.merge([b, g, r]), 0, 255).astype(np.uint8)\n\n\ndef denoise_median(img_bgr: np.ndarray, k: int = 3) -> np.ndarray:\n    k = int(k)\n    k = k if k % 2 == 1 else k + 1\n    return cv2.medianBlur(img_bgr, k)\n\n\ndef denoise_bilateral(\n    img_bgr: np.ndarray, d: int = 7, sigma_color: float = 50, sigma_space: float = 50\n) -> np.ndarray:\n    return cv2.bilateralFilter(\n        img_bgr,\n        d=int(d),\n        sigmaColor=float(sigma_color),\n        sigmaSpace=float(sigma_space),\n    )\n\n\ndef illumination_strength(img_bgr: np.ndarray, sigma: float = 35) -> float:\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\ndef bg_correct_lab_l(img_bgr: np.ndarray, k: int = 81) -> np.ndarray:\n    lab = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2LAB)\n    l, a, b = cv2.split(lab)\n    k = int(k)\n    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\ndef preprocess_by_source(img_bgr: np.ndarray, source: str) -> np.ndarray:\n    s = str(source)\n    if s == \"0001\":\n        return img_bgr\n    if s == \"0003\":\n        return clahe_luminance_bgr(grayworld_white_balance(img_bgr), clip=1.2, tile=8)\n    if s == \"0004\":\n        return img_bgr\n    if s == \"0006\":\n        x = denoise_bilateral(img_bgr, d=5, sigma_color=25, sigma_space=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:\n            x = bg_correct_lab_l(x, k=81)\n        if cv2.cvtColor(x, cv2.COLOR_BGR2GRAY).std() < 30:\n            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:\n            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:\n            x = bg_correct_lab_l(x, k=81)\n        if cv2.cvtColor(x, cv2.COLOR_BGR2GRAY).std() < 30:\n            x = clahe_luminance_bgr(x, clip=1.15, tile=8)\n        return x\n    if s == \"0011\":\n        return clahe_luminance_bgr(grayworld_white_balance(img_bgr), clip=1.2, tile=8)\n    if s == \"0012\":\n        return img_bgr\n    return img_bgr\n\n\ndef stage1_quality(s1_rgb: np.ndarray) -> float:\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)))\n    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\n\ndef load_hengck_modules(hengck_root: str):\n    sys.path.append(hengck_root)\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    return s0c, s1c, s2c, Stage0Net, Stage1Net, MyCoordUnetDecoder, encode_with_resnet\n\n\nclass Net3(nn.Module):\n    def __init__(self, encode_with_resnet, MyCoordUnetDecoder, pretrained: bool = 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        self._encode_with_resnet = encode_with_resnet\n\n    def forward(self, image: torch.Tensor) -> torch.Tensor:\n        encode = self._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\nclass PhysioPipeline:\n    def __init__(self, device: str = \"cuda:0\"):\n        self.device = device\n        self.stage0_net = None\n        self.stage1_net = None\n        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        self._s0c = None\n        self._s1c = None\n        self._s2c = None\n\n    def load_models(\n        self,\n        s0c,\n        s1c,\n        s2c,\n        stage0_w: str,\n        stage1_w: str,\n        stage2_w: str,\n        Stage0Net,\n        Stage1Net,\n        MyCoordUnetDecoder,\n        encode_with_resnet,\n    ) -> None:\n        self._s0c = s0c\n        self._s1c = s1c\n        self._s2c = s2c\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(\n            encode_with_resnet=encode_with_resnet,\n            MyCoordUnetDecoder=MyCoordUnetDecoder,\n            pretrained=False,\n        ).to(self.device).eval()\n        st = torch.load(stage2_w, map_location=\"cpu\")\n        if isinstance(st, dict) and \"state_dict\" in st:\n            st = st[\"state_dict\"]\n        self.stage2_net.load_state_dict(st, strict=True)\n\n    def run_stage0(self, img_bgr: np.ndarray) -> np.ndarray:\n        img_rgb = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB)\n        img_for_model = change_color(img_rgb)\n        batch = self._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 = self._s0c.output_to_predict(img_rgb, batch, output)\n        normalised, _, _ = self._s0c.normalise_by_homography(rotated, keypoint)\n        return normalised\n\n    def run_stage1(self, stage0_img_rgb: np.ndarray) -> np.ndarray:\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, _ = self._s1c.output_to_predict(image, batch, output)\n        return self._s1c.rectify_image(image, gridpoint_xy)\n\n    def run_stage2(self, stage1_img_rgb: np.ndarray, length: int) -> np.ndarray:\n        img = stage1_img_rgb[self.y0 : self.y1, self.x0 : self.x1] / 255.0\n        batch = (\n            self.resize(torch.from_numpy(np.ascontiguousarray(img.transpose(2, 0, 1))).unsqueeze(0))\n            .float()\n            .to(self.device)\n        )\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 = self._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        for i in range(4):\n            series[i] = savgol_filter(series[i], window_length=7, polyorder=2)\n        return series\n\n\ndef cls_preprocess_bgr(img_bgr: np.ndarray, resolution: int) -> torch.Tensor:\n    img = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB)\n    img = cv2.resize(img, (resolution, 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\n@torch.no_grad()\ndef predict_source_suffix(model: nn.Module, img_bgr: np.ndarray, resolution: int, device: str) -> str:\n    x = cls_preprocess_bgr(img_bgr, resolution=resolution).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\ndef select_stage1_with_source(\n    pipeline: PhysioPipeline,\n    img_raw_bgr: np.ndarray,\n    pred_source_suffix: str,\n    selector_margin: float = 1.02,\n) -> np.ndarray:\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\n\n\n@dataclass\nclass TrainConfig:\n    model_name: str = \"efficientnet_b2\"\n    resolution: int = 256\n    batch_size: int = 32\n    epochs: int = 6\n    lr: float = 3e-4\n    weight_decay: float = 1e-2\n    num_workers: int = 4\n    val_split: float = 0.1\n    seed: int = 42\n\n\ndef seed_everything(seed: int) -> None:\n    import random\n\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\n\ndef build_cls_transforms(resolution: int) -> Tuple[A.Compose, A.Compose]:\n    if A is None:\n        raise RuntimeError(\"albumentations is required for classifier training.\")\n\n    train_t = A.Compose(\n        [\n            A.Resize(height=resolution, width=resolution, p=1),\n            A.HorizontalFlip(p=0.5),\n            A.RandomBrightnessContrast(brightness_limit=0.15, contrast_limit=0.15, p=0.4),\n            A.GaussianBlur(blur_limit=(3, 7), sigma_limit=(0.1, 1.5), p=0.2),\n            A.Normalize(p=1),\n            ToTensorV2(p=1),\n        ]\n    )\n    valid_t = A.Compose(\n        [\n            A.Resize(height=resolution, width=resolution, p=1),\n            A.Normalize(p=1),\n            ToTensorV2(p=1),\n        ]\n    )\n    return train_t, valid_t\n\n\nclass SourceDataset(torch.utils.data.Dataset):\n    def __init__(self, items: List[Tuple[str, int]], transforms: A.Compose):\n        self.items = items\n        self.transforms = transforms\n\n    def __len__(self) -> int:\n        return len(self.items)\n\n    def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]:\n        image_path, label = self.items[idx]\n        image = cv2.imread(image_path, cv2.IMREAD_COLOR)\n        if image is None:\n            raise FileNotFoundError(image_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        image = self.transforms(image=image)[\"image\"]\n        return {\"image\": image, \"label\": torch.tensor(label, dtype=torch.long)}\n\n\ndef gather_source_items(train_root: Path) -> List[Tuple[str, int]]:\n    items: List[Tuple[str, int]] = []\n    pattern = re.compile(r\"-(\\d{4})\\.png$\")\n    for path in train_root.glob(\"**/*.png\"):\n        match = pattern.search(path.name)\n        if not match:\n            continue\n        suffix = match.group(1)\n        if suffix not in SOURCE_SUFFIXES:\n            continue\n        label = int(suffix) - 1\n        items.append((str(path), label))\n    if not items:\n        raise RuntimeError(f\"No training images found under {train_root}\")\n    return items\n\n\ndef train_classifier(data_root: str, output_path: str, cfg: TrainConfig) -> None:\n    seed_everything(cfg.seed)\n    train_root = Path(data_root) / \"train\"\n    items = gather_source_items(train_root)\n    rng = np.random.default_rng(cfg.seed)\n    rng.shuffle(items)\n\n    split_idx = int(len(items) * (1 - cfg.val_split))\n    train_items = items[:split_idx]\n    val_items = items[split_idx:]\n\n    train_t, valid_t = build_cls_transforms(cfg.resolution)\n    train_ds = SourceDataset(train_items, train_t)\n    val_ds = SourceDataset(val_items, valid_t)\n\n    train_loader = torch.utils.data.DataLoader(\n        train_ds,\n        batch_size=cfg.batch_size,\n        shuffle=True,\n        num_workers=cfg.num_workers,\n        pin_memory=True,\n    )\n    val_loader = torch.utils.data.DataLoader(\n        val_ds,\n        batch_size=cfg.batch_size,\n        shuffle=False,\n        num_workers=cfg.num_workers,\n        pin_memory=True,\n    )\n\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    model = timm.create_model(cfg.model_name, pretrained=True, num_classes=len(SOURCE_SUFFIXES))\n    model.to(device)\n\n    optimizer = torch.optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay)\n    criterion = nn.CrossEntropyLoss()\n    scaler = torch.cuda.amp.GradScaler(enabled=torch.cuda.is_available())\n\n    best_acc = 0.0\n    for epoch in range(cfg.epochs):\n        model.train()\n        running_loss = 0.0\n        for batch in train_loader:\n            optimizer.zero_grad()\n            images = batch[\"image\"].to(device)\n            labels = batch[\"label\"].to(device)\n            with torch.cuda.amp.autocast(enabled=torch.cuda.is_available()):\n                logits = model(images)\n                loss = criterion(logits, labels)\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n            running_loss += loss.item() * images.size(0)\n\n        model.eval()\n        correct = 0\n        total = 0\n        with torch.no_grad():\n            for batch in val_loader:\n                images = batch[\"image\"].to(device)\n                labels = batch[\"label\"].to(device)\n                logits = model(images)\n                preds = torch.argmax(logits, dim=1)\n                correct += (preds == labels).sum().item()\n                total += labels.size(0)\n\n        val_acc = correct / max(total, 1)\n        epoch_loss = running_loss / max(len(train_loader.dataset), 1)\n        print(f\"Epoch {epoch + 1}/{cfg.epochs} - loss={epoch_loss:.4f} val_acc={val_acc:.4f}\")\n\n        if val_acc > best_acc:\n            best_acc = val_acc\n            torch.save(model.state_dict(), output_path)\n            print(f\"Saved best classifier to {output_path}\")\n\n\ndef load_classifier(weights_path: Optional[str], device: str, model_name: str, num_classes: int) -> Optional[nn.Module]:\n    if not weights_path:\n        return None\n    if not os.path.exists(weights_path):\n        raise FileNotFoundError(weights_path)\n    model = timm.create_model(model_name, pretrained=False, num_classes=num_classes)\n    state = torch.load(weights_path, map_location=\"cpu\")\n    if isinstance(state, dict) and \"state_dict\" in state:\n        state = state[\"state_dict\"]\n    state = {k.replace(\"module.\", \"\"): v for k, v in state.items()}\n    model.load_state_dict(state, strict=False)\n    model.to(device).eval()\n    return model\n\n\ndef get_lines(np_image: np.ndarray, threshold: int = 1200, rho_resolution: int = 1) -> Optional[np.ndarray]:\n    image = cv2.cvtColor(np_image, cv2.COLOR_RGB2BGR)\n    gray_image = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)\n    edges = cv2.Canny(gray_image, 50, 150, apertureSize=3)\n    return cv2.HoughLines(edges, rho_resolution, np.pi / 180, threshold, None, 0, 0)\n\n\ndef is_within_x_degrees_of_horizontal(theta: float, degree_window: float) -> bool:\n    theta_degrees = theta * 180 / np.pi\n    deviation_from_horizontal = abs(90 - theta_degrees)\n    return deviation_from_horizontal < degree_window\n\n\ndef filter_lines(\n    lines: Optional[np.ndarray], degree_window: float = 30, parallelism_count: int = 3, parallelism_window: float = 2\n) -> Optional[np.ndarray]:\n    parallelism_radian = np.deg2rad(parallelism_window)\n    filtered_lines = []\n    if lines is not None:\n        for line in lines:\n            for rho, theta in line:\n                if is_within_x_degrees_of_horizontal(theta, degree_window):\n                    filtered_lines.append((rho, theta))\n\n    parallel_lines = []\n    if filtered_lines:\n        for rho, theta in filtered_lines:\n            count = 0\n            for comp_rho, comp_theta in filtered_lines:\n                if (\n                    abs(theta - comp_theta) < parallelism_radian\n                    or abs((theta - comp_theta) - np.pi) < parallelism_radian\n                ):\n                    count += 1\n            if count >= parallelism_count:\n                parallel_lines.append((rho, theta))\n\n    if not parallel_lines:\n        return None\n    return np.array(parallel_lines)[:, np.newaxis, :]\n\n\ndef get_median_degrees(lines: np.ndarray) -> float:\n    lines = lines[:, 0, :]\n    line_angles = [-(90 - line[1] * 180 / np.pi) for line in lines]\n    return round(np.median(line_angles), 4)\n\n\ndef get_rotation_angle(np_image: np.ndarray) -> float:\n    lines = get_lines(np_image, threshold=1200)\n    filtered_lines = filter_lines(lines, degree_window=30, parallelism_count=3, parallelism_window=2)\n    if filtered_lines is None:\n        return 0.0\n    return get_median_degrees(filtered_lines)\n\n\n\ndef cut_to_mask(mask: torch.Tensor, return_y1: bool = False) -> Tuple[torch.Tensor, Optional[int], Optional[int]]:\n    coords = torch.where(mask[0] >= 1)\n    y_min, y_max = coords[0].min().item(), coords[0].max().item()\n    x_min, x_max = coords[1].min().item(), coords[1].max().item()\n    if return_y1:\n        return mask[:, y_min : y_max + 1, x_min : x_max + 1], y_min, x_min\n    return mask[:, y_min : y_max + 1, x_min : x_max + 1], None, None\n\n\ndef cut_binary(mask_to_use: torch.Tensor) -> Tuple[Dict[str, torch.Tensor], Dict[str, Dict[str, int]]]:\n    signal_masks: Dict[str, torch.Tensor] = {}\n    signal_positions: Dict[str, Dict[str, int]] = {}\n    for lead_name, lead_value in LEAD_LABEL_MAPPING.items():\n        binary_mask = torch.where(mask_to_use == lead_value, 1, 0)\n        if binary_mask.sum() > 0:\n            cropped_mask, y1, x1 = cut_to_mask(binary_mask, True)\n            signal_masks[lead_name] = cropped_mask\n            signal_positions[lead_name] = {\"y1\": y1, \"x1\": x1}\n        else:\n            signal_masks[lead_name] = None\n            signal_positions[lead_name] = None\n    return signal_masks, signal_positions\n\n\ndef vectorise_by_fs(\n    image_height: int,\n    mask: torch.Tensor,\n    signal_cropped_y: int,\n    sec_per_pixel: float,\n    mV_per_pixel: float,\n    y_shift_ratio: Dict[str, float],\n    lead: str,\n    fs: int,\n) -> torch.Tensor:\n    total_seconds_from_mask = round(sec_per_pixel * mask.shape[2], 1)\n    if total_seconds_from_mask > 5:\n        total_seconds = 10.0\n        y_shift_ratio_ = y_shift_ratio[\"full\"]\n    else:\n        total_seconds = 2.5\n        y_shift_ratio_ = y_shift_ratio[lead]\n    values_needed = int(total_seconds * fs)\n\n    non_zero_mean = torch.tensor(\n        [\n            torch.mean(torch.nonzero(mask[0, :, i]).type(torch.float32))\n            for i in range(mask.shape[2])\n        ]\n    )\n    signal_cropped_shifted = (1 - y_shift_ratio_) * image_height - signal_cropped_y\n    predicted_signal = (signal_cropped_shifted - non_zero_mean) * mV_per_pixel\n\n    n = predicted_signal.shape[0]\n    data_reshaped = predicted_signal.view(1, 1, n)\n    resampled_data = F.interpolate(\n        data_reshaped, size=values_needed, mode=\"linear\", align_corners=False\n    )\n    return resampled_data.view(-1)\n\n\n\ndef make_submission_from_pred(base_id: str, fs: int, sig_len: int, d_series: Dict[str, np.ndarray]) -> pd.DataFrame:\n    base_id = str(base_id)\n    fs = int(fs)\n    sig_len = int(sig_len)\n    n_short = int(np.floor(fs * 2.5))\n\n    def take_segment(y: np.ndarray, n: int) -> np.ndarray:\n        y = np.asarray(y, dtype=np.float64)\n        if len(y) >= n:\n            return y[:n]\n        if len(y) == 0:\n            return np.zeros(n, np.float64)\n        return np.concatenate([y, np.full(n - len(y), y[-1], np.float64)])\n\n    rows = []\n    for lead in LEADS_ORDER:\n        y = np.asarray(d_series[lead], dtype=np.float64)\n        seg = take_segment(y, sig_len if lead == \"II\" else n_short)\n        rows.append(\n            pd.DataFrame(\n                {\n                    \"id\": [f\"{base_id}_{i}_{lead}\" for i in range(len(seg))],\n                    \"value\": seg.astype(np.float32),\n                }\n            )\n        )\n    return pd.concat(rows, ignore_index=True)\n\n\ndef resample_to_length(y: np.ndarray, n: int) -> np.ndarray:\n    if y is None:\n        return None\n    if len(y) == n:\n        return y\n    x_old = np.linspace(0, 1, len(y))\n    x_new = np.linspace(0, 1, n)\n    return np.interp(x_new, x_old, y).astype(np.float32)\n\n\ndef blend_signals(\n    stage_series: Dict[str, np.ndarray],\n    digit_series: Dict[str, np.ndarray],\n    fs: int,\n    sig_len: int,\n    alpha: float,\n) -> Dict[str, np.ndarray]:\n    n_short = int(np.floor(fs * 2.5))\n    out: Dict[str, np.ndarray] = {}\n    for lead in LEADS_ORDER:\n        target_len = sig_len if lead == \"II\" else n_short\n        s_stage = resample_to_length(stage_series.get(lead), target_len) if stage_series else None\n        s_digit = resample_to_length(digit_series.get(lead), target_len) if digit_series else None\n\n        if s_digit is None and s_stage is None:\n            out[lead] = np.zeros(target_len, dtype=np.float32)\n        elif s_digit is None:\n            out[lead] = s_stage\n        elif s_stage is None:\n            out[lead] = s_digit\n        else:\n            out[lead] = (alpha * s_stage + (1 - alpha) * s_digit).astype(np.float32)\n\n    # create II_short for Einthoven correction\n    out[\"II_short\"] = out[\"II\"][:n_short]\n    out = dw(out)\n    out.pop(\"II_short\", None)\n    return out\n\n\ndef run_inference(\n    data_root: str,\n    output_path: str,\n    hengck_root: str,\n    stage0_w: str,\n    stage1_w: str,\n    stage2_w: str,\n    cls_weights: Optional[str],\n    cls_model_name: str,\n    cls_resolution: int,\n    selector_margin: float,\n    use_stage: bool,\n    use_digitiser: bool,\n    blend_alpha: float,\n) -> None:\n    if not use_stage and not use_digitiser:\n        raise ValueError(\"At least one of use_stage/use_digitiser must be true.\")\n\n    s0c = s1c = s2c = Stage0Net = Stage1Net = MyCoordUnetDecoder = encode_with_resnet = None\n    pipeline = None\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    if use_stage:\n        s0c, s1c, s2c, Stage0Net, Stage1Net, MyCoordUnetDecoder, encode_with_resnet = load_hengck_modules(\n            hengck_root\n        )\n        pipeline = PhysioPipeline(device=\"cuda:0\" if device == \"cuda\" else \"cpu\")\n        pipeline.load_models(\n            s0c=s0c,\n            s1c=s1c,\n            s2c=s2c,\n            stage0_w=stage0_w,\n            stage1_w=stage1_w,\n            stage2_w=stage2_w,\n            Stage0Net=Stage0Net,\n            Stage1Net=Stage1Net,\n            MyCoordUnetDecoder=MyCoordUnetDecoder,\n            encode_with_resnet=encode_with_resnet,\n        )\n\n    cls_model = load_classifier(\n        cls_weights, device=device, model_name=cls_model_name, num_classes=len(SOURCE_SUFFIXES)\n    )\n\n    df_test = pd.read_csv(Path(data_root) / \"test.csv\")\n    df_test[\"id\"] = df_test[\"id\"].astype(str)\n    sample_submission = pd.read_parquet(Path(data_root) / \"sample_submission.parquet\")[[\"id\"]]\n\n    res = []\n    for sample_id, group in df_test.groupby(\"id\", sort=True):\n        img_path = Path(data_root) / \"test\" / f\"{sample_id}.png\"\n        img_raw = cv2.imread(str(img_path), cv2.IMREAD_COLOR)\n        if img_raw is None:\n            raise FileNotFoundError(img_path)\n\n        fs = int(group.fs.iloc[0])\n        sig_len = int(group.loc[group.lead == \"II\", \"number_of_rows\"].iloc[0])\n\n        stage_series = None\n        if use_stage:\n            if cls_model is None:\n                s1 = pipeline.run_stage1(pipeline.run_stage0(img_raw))\n            else:\n                pred_src = predict_source_suffix(cls_model, img_raw, cls_resolution, device=device)\n                s1 = select_stage1_with_source(pipeline, img_raw, pred_src, selector_margin=selector_margin)\n            series_4row = pipeline.run_stage2(s1, length=sig_len)\n            stage_series = dw(series_dict(series_4row))\n            if \"II_short\" in stage_series:\n                stage_series.pop(\"II_short\")\n\n        final_series = stage_series\n\n        res.append(make_submission_from_pred(sample_id, fs, sig_len, final_series))\n        gc.collect()\n\n    df_submission = pd.concat(res, ignore_index=True)\n    df_submission = df_submission.set_index(\"id\").reindex(sample_submission[\"id\"]).reset_index()\n    if df_submission[\"value\"].isna().any():\n        raise RuntimeError(\"Submission has NaNs.\")\n    if not (df_submission[\"id\"].values == sample_submission[\"id\"].values).all():\n        raise RuntimeError(\"Submission IDs do not match sample_submission.\")\n\n    df_submission.to_csv(output_path, index=False)\n    print(f\"Saved submission to {output_path} with shape {df_submission.shape}\")\n\n\ndef parse_args() -> argparse.Namespace:\n    parser = argparse.ArgumentParser(description=\"PhysioNet ECG ensemble pipeline\")\n    subparsers = parser.add_subparsers(dest=\"mode\", required=True)\n\n    train_parser = subparsers.add_parser(\"train_classifier\", help=\"Train image-source classifier\")\n    train_parser.add_argument(\"--data-root\", required=True, help=\"Path to competition data root\")\n    train_parser.add_argument(\"--output\", required=True, help=\"Where to save classifier weights (.pth)\")\n    train_parser.add_argument(\"--epochs\", type=int, default=6)\n    train_parser.add_argument(\"--batch-size\", type=int, default=32)\n    train_parser.add_argument(\"--resolution\", type=int, default=256)\n    train_parser.add_argument(\"--model\", type=str, default=\"eca_nfnet_l0\")\n    train_parser.add_argument(\"--lr\", type=float, default=3e-4)\n    train_parser.add_argument(\"--weight-decay\", type=float, default=1e-2)\n    train_parser.add_argument(\"--val-split\", type=float, default=0.1)\n    train_parser.add_argument(\"--num-workers\", type=int, default=4)\n    train_parser.add_argument(\"--seed\", type=int, default=42)\n\n    infer_parser = subparsers.add_parser(\"predict\", help=\"Generate submission from test images\")\n    infer_parser.add_argument(\"--data-root\", required=True, help=\"Path to competition data root\")\n    infer_parser.add_argument(\"--output\", default=\"submission.csv\")\n    infer_parser.add_argument(\"--hengck-root\", default=DEFAULT_HENGCK_ROOT)\n    infer_parser.add_argument(\"--stage0-weights\", default=DEFAULT_STAGE0_WEIGHTS)\n    infer_parser.add_argument(\"--stage1-weights\", default=DEFAULT_STAGE1_WEIGHTS)\n    infer_parser.add_argument(\"--stage2-weights\", default=DEFAULT_STAGE2_WEIGHTS)\n    infer_parser.add_argument(\"--cls-weights\", default=None)\n    infer_parser.add_argument(\"--cls-model\", default=\"eca_nfnet_l0\")\n    infer_parser.add_argument(\"--cls-resolution\", type=int, default=256)\n    infer_parser.add_argument(\"--selector-margin\", type=float, default=1.02)\n    infer_parser.add_argument(\"--use-stage\", action=\"store_true\", default=True)\n    infer_parser.add_argument(\"--no-stage\", dest=\"use_stage\", action=\"store_false\")\n    infer_parser.add_argument(\"--use-digitiser\", action=\"store_true\", default=True)\n    infer_parser.add_argument(\"--no-digitiser\", dest=\"use_digitiser\", action=\"store_false\")\n    infer_parser.add_argument(\"--blend-alpha\", type=float, default=0.6)\n\n    return parser.parse_args()\n\n\ndef main() -> None:\n    args = parse_args()\n    if args.mode == \"train_classifier\":\n        cfg = TrainConfig(\n            model_name=args.model,\n            resolution=args.resolution,\n            batch_size=args.batch_size,\n            epochs=args.epochs,\n            lr=args.lr,\n            weight_decay=args.weight_decay,\n            num_workers=args.num_workers,\n            val_split=args.val_split,\n            seed=args.seed,\n        )\n        train_classifier(args.data_root, args.output, cfg)\n    elif args.mode == \"predict\":\n        run_inference(\n            data_root=args.data_root,\n            output_path=args.output,\n            hengck_root=args.hengck_root,\n            stage0_w=args.stage0_weights,\n            stage1_w=args.stage1_weights,\n            stage2_w=args.stage2_weights,\n            cls_weights=args.cls_weights,\n            cls_model_name=args.cls_model,\n            cls_resolution=args.cls_resolution,\n            selector_margin=args.selector_margin,\n            use_stage=args.use_stage,\n            use_digitiser=args.use_digitiser,\n            blend_alpha=args.blend_alpha,\n        )\n    else:\n        raise RuntimeError(f\"Unknown mode {args.mode}\")\n\n\nif __name__ == \"__main__\":\n    main()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python physionet_ensemble_pipeline.py train_classifier \\\n  --data-root /kaggle/input/physionet-ecg-image-digitization \\\n  --output /kaggle/working/source_classifier.pth","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python physionet_ensemble_pipeline.py predict \\\n  --data-root /kaggle/input/physionet-ecg-image-digitization \\\n  --output submission.csv --selector-margin 0.85 \\\n  --cls-weights /kaggle/working/source_classifier.pth \\\n  --no-digitiser","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}