{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.12","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,"sourceType":"competition"}],"dockerImageVersionId":31260,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"c4392f97","cell_type":"markdown","source":"# PhysioNet ECG","metadata":{}},{"id":"32c90861","cell_type":"code","source":"import os\nimport random\nimport numpy as np\nimport pandas as pd\nimport cv2\nfrom scipy import signal as scipy_signal\nfrom tqdm import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\ntry:\n    import timm\n    import albumentations as A\n    from albumentations.pytorch import ToTensorV2\nexcept ImportError:\n    !pip install -q timm albumentations tqdm\n    import timm\n    import albumentations as A\n    from albumentations.pytorch import ToTensorV2\n\nprint(f\"PyTorch: {torch.__version__}, CUDA: {torch.cuda.is_available()}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T06:16:51.798309Z","iopub.execute_input":"2026-01-22T06:16:51.798607Z","iopub.status.idle":"2026-01-22T06:17:10.635746Z","shell.execute_reply.started":"2026-01-22T06:16:51.798580Z","shell.execute_reply":"2026-01-22T06:17:10.634991Z"}},"outputs":[],"execution_count":null},{"id":"b6a9ae2b","cell_type":"code","source":"class Config:\n    SEED = 42\n    IMAGE_SIZE = (512, 1024)  # Same as V3\n    BATCH_SIZE = 4\n    EPOCHS = 80\n    LR = 1e-4\n    WEIGHT_DECAY = 1e-5\n    DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    OUTPUT_LEN = 5000\n    \n    if os.path.exists(\"/kaggle/input\"):\n        BASE_DIR = \"/kaggle/input/physionet-ecg-image-digitization\"\n    else:\n        BASE_DIR = \"d:/physionet-ecg-image-digitization\"\n\n    TRAIN_CSV = os.path.join(BASE_DIR, \"train.csv\")\n    TEST_CSV = os.path.join(BASE_DIR, \"test.csv\")\n    TRAIN_IMG_DIR = os.path.join(BASE_DIR, \"train\")\n    TEST_IMG_DIR = os.path.join(BASE_DIR, \"test\")\n    \ndef seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n    \nseed_everything(Config.SEED)\nprint(f\"Device: {Config.DEVICE}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T06:17:10.637179Z","iopub.execute_input":"2026-01-22T06:17:10.637429Z","iopub.status.idle":"2026-01-22T06:17:10.654974Z","shell.execute_reply.started":"2026-01-22T06:17:10.637405Z","shell.execute_reply":"2026-01-22T06:17:10.654194Z"}},"outputs":[],"execution_count":null},{"id":"5b18e2e3","cell_type":"code","source":"def align_signals(pred, gt, max_shift_samples):\n    if len(pred) == 0 or len(gt) == 0:\n        return pred, gt, 0\n    correlation = scipy_signal.correlate(gt, pred, mode='full')\n    lags = scipy_signal.correlation_lags(len(gt), len(pred), mode='full')\n    \n    valid_mask = np.abs(lags) <= max_shift_samples\n    valid_lags = lags[valid_mask]\n    valid_corr = correlation[valid_mask]\n    \n    if len(valid_corr) == 0:\n        return pred[:len(gt)], gt, 0\n    \n    optimal_idx = np.argmax(valid_corr)\n    optimal_shift = valid_lags[optimal_idx]\n    \n    if optimal_shift > 0:\n        aligned_pred = np.pad(pred, (optimal_shift, 0), mode='edge')[:len(gt)]\n    elif optimal_shift < 0:\n        aligned_pred = np.pad(pred, (0, -optimal_shift), mode='edge')[-optimal_shift:len(gt)-optimal_shift]\n    else:\n        aligned_pred = pred[:len(gt)]\n    \n    min_len = min(len(aligned_pred), len(gt))\n    return aligned_pred[:min_len], gt[:min_len], optimal_shift\n\ndef compute_snr_single_record(pred_12leads, gt_12leads, fs=500):\n    max_shift_samples = int(0.2 * fs)\n    total_signal_power = 0.0\n    total_noise_power = 0.0\n    \n    for lead_idx in range(12):\n        pred = pred_12leads[lead_idx]\n        gt = gt_12leads[lead_idx]\n        \n        if np.std(gt) < 1e-6:\n            continue\n        \n        aligned_pred, aligned_gt, _ = align_signals(pred, gt, max_shift_samples)\n        offset = np.mean(aligned_gt) - np.mean(aligned_pred)\n        aligned_pred = aligned_pred + offset\n        \n        signal_power = np.sum(aligned_gt ** 2)\n        noise = aligned_gt - aligned_pred\n        noise_power = np.sum(noise ** 2)\n        \n        total_signal_power += signal_power\n        total_noise_power += noise_power\n    \n    if total_noise_power < 1e-10:\n        return 100.0\n    \n    snr = 10 * np.log10(total_signal_power / (total_noise_power + 1e-10))\n    return snr\n\ndef compute_competition_snr(all_preds, all_gts, fs_list):\n    snrs = []\n    for pred, gt, fs in zip(all_preds, all_gts, fs_list):\n        snr = compute_snr_single_record(pred, gt, fs=fs)\n        snrs.append(snr)\n    return np.mean(snrs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T06:17:10.655847Z","iopub.execute_input":"2026-01-22T06:17:10.656096Z","iopub.status.idle":"2026-01-22T06:17:10.666157Z","shell.execute_reply.started":"2026-01-22T06:17:10.656076Z","shell.execute_reply":"2026-01-22T06:17:10.665597Z"}},"outputs":[],"execution_count":null},{"id":"ee0a2bb2","cell_type":"code","source":"def get_transforms(cfg, mode=\"train\"):\n    if mode == \"train\":\n        return A.Compose([\n            A.Resize(cfg.IMAGE_SIZE[0], cfg.IMAGE_SIZE[1]),\n            A.ShiftScaleRotate(shift_limit=0.02, scale_limit=0.05, rotate_limit=2, p=0.3, border_mode=cv2.BORDER_CONSTANT),\n            A.GaussianBlur(blur_limit=(3, 5), p=0.2),\n            A.RandomBrightnessContrast(brightness_limit=0.1, contrast_limit=0.1, p=0.2),\n            A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n            ToTensorV2(),\n        ])\n    else:\n        return A.Compose([\n            A.Resize(cfg.IMAGE_SIZE[0], cfg.IMAGE_SIZE[1]),\n            A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n            ToTensorV2(),\n        ])\n\nclass ECGDataset(Dataset):\n    def __init__(self, df, img_dir, mode=\"train\", transform=None):\n        self.df = df.reset_index(drop=True)\n        self.img_dir = img_dir\n        self.mode = mode\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        ecg_id = str(row['id'])\n        fs = row['fs'] if 'fs' in row else 500\n        \n        signal = np.zeros((12, Config.OUTPUT_LEN), dtype=np.float32)\n        \n        if self.mode != \"test\":\n            sig_path = os.path.join(self.img_dir, ecg_id, f\"{ecg_id}.csv\")\n            if os.path.exists(sig_path):\n                try:\n                    sig_df = pd.read_csv(sig_path)\n                    leads = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', 'V1', 'V2', 'V3', 'V4', 'V5', 'V6']\n                    \n                    for i, lead in enumerate(leads):\n                        if lead in sig_df.columns:\n                            lead_data = sig_df[lead].values.astype(np.float32)\n                            lead_data = np.nan_to_num(lead_data, nan=0.0)\n                            if len(lead_data) > 0:\n                                signal[i] = np.interp(\n                                    np.linspace(0, 1, Config.OUTPUT_LEN),\n                                    np.linspace(0, 1, len(lead_data)),\n                                    lead_data\n                                )\n                except:\n                    pass\n\n        suffixes = [\"-0001.png\", \"-0002.png\", \"-0003.png\", \"-0004.png\", \"-0005.png\"]\n        found_img_path = None\n        subfolder = os.path.join(self.img_dir, ecg_id)\n        \n        if os.path.isdir(subfolder):\n            if self.mode == \"train\":\n                candidates = [s for s in suffixes if os.path.exists(os.path.join(subfolder, f\"{ecg_id}{s}\"))]\n                if candidates:\n                    found_img_path = os.path.join(subfolder, f\"{ecg_id}{random.choice(candidates)}\")\n            if found_img_path is None:\n                for s in suffixes:\n                    if os.path.exists(os.path.join(subfolder, f\"{ecg_id}{s}\")):\n                        found_img_path = os.path.join(subfolder, f\"{ecg_id}{s}\")\n                        break\n        \n        if found_img_path is None:\n            direct_path = os.path.join(self.img_dir, f\"{ecg_id}.png\")\n            if os.path.exists(direct_path):\n                found_img_path = direct_path\n        if found_img_path is None:\n            found_img_path = os.path.join(self.img_dir, ecg_id, f\"{ecg_id}-0001.png\")\n\n        image = cv2.imread(found_img_path)\n        if image is None:\n            image = np.zeros((*Config.IMAGE_SIZE, 3), dtype=np.uint8)\n        else:\n            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n            \n        if self.transform:\n            augmented = self.transform(image=image)\n            image = augmented['image']\n            \n        return image, torch.tensor(signal), torch.tensor([fs])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T06:17:10.667659Z","iopub.execute_input":"2026-01-22T06:17:10.667969Z","iopub.status.idle":"2026-01-22T06:17:10.686172Z","shell.execute_reply.started":"2026-01-22T06:17:10.667947Z","shell.execute_reply":"2026-01-22T06:17:10.685595Z"}},"outputs":[],"execution_count":null},{"id":"b89d5c9c","cell_type":"code","source":"class ECGModel(nn.Module):\n    def __init__(self, model_name='convnext_small', num_leads=12, output_len=5000, pretrained=False):\n        super().__init__()\n        \n        self.backbone = timm.create_model(model_name, pretrained=pretrained, features_only=True)\n        \n        dummy_input = torch.randn(1, 3, Config.IMAGE_SIZE[0], Config.IMAGE_SIZE[1])\n        with torch.no_grad():\n            feats = self.backbone(dummy_input)\n        \n        self.feat_channels = [f.shape[1] for f in feats[-3:]]\n        \n        self.lateral1 = nn.Conv2d(self.feat_channels[0], 256, 1)\n        self.lateral2 = nn.Conv2d(self.feat_channels[1], 256, 1)\n        self.lateral3 = nn.Conv2d(self.feat_channels[2], 256, 1)\n        \n        self.fuse = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1),\n            nn.BatchNorm2d(256),\n            nn.GELU(),\n        )\n        \n        self.decoder = nn.Sequential(\n            nn.ConvTranspose1d(256, 256, 4, stride=2, padding=1),\n            nn.BatchNorm1d(256),\n            nn.GELU(),\n            nn.ConvTranspose1d(256, 128, 4, stride=2, padding=1),\n            nn.BatchNorm1d(128),\n            nn.GELU(),\n            nn.ConvTranspose1d(128, 64, 4, stride=2, padding=1),\n            nn.BatchNorm1d(64),\n            nn.GELU(),\n            nn.ConvTranspose1d(64, 32, 4, stride=2, padding=1),\n            nn.BatchNorm1d(32),\n            nn.GELU(),\n            nn.Conv1d(32, num_leads, 1),\n        )\n        \n        self.output_len = output_len\n        \n    def forward(self, x):\n        feats = self.backbone(x)\n        f1, f2, f3 = feats[-3], feats[-2], feats[-1]\n        \n        p3 = self.lateral3(f3)\n        p2 = self.lateral2(f2) + F.interpolate(p3, size=f2.shape[2:], mode='bilinear', align_corners=False)\n        p1 = self.lateral1(f1) + F.interpolate(p2, size=f1.shape[2:], mode='bilinear', align_corners=False)\n        \n        fused = self.fuse(p1)\n        x = fused.mean(dim=2)\n        x = self.decoder(x)\n        x = F.interpolate(x, size=self.output_len, mode='linear', align_corners=False)\n        \n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T06:17:10.688023Z","iopub.execute_input":"2026-01-22T06:17:10.688649Z","iopub.status.idle":"2026-01-22T06:17:10.706470Z","shell.execute_reply.started":"2026-01-22T06:17:10.688615Z","shell.execute_reply":"2026-01-22T06:17:10.705951Z"}},"outputs":[],"execution_count":null},{"id":"5bd377e4","cell_type":"code","source":"def train_fn(model, loader, optimizer, criterion, device, scaler):\n    model.train()\n    total_loss = 0\n    \n    pbar = tqdm(loader, desc=\"Training\")\n    for images, signals, meta in pbar:\n        images = images.to(device)\n        signals = signals.to(device)\n        \n        optimizer.zero_grad()\n        \n        with torch.amp.autocast('cuda'):\n            preds = model(images)\n            loss = criterion(preds, signals)\n        \n        if torch.isnan(loss):\n            continue\n\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        scaler.step(optimizer)\n        scaler.update()\n        \n        total_loss += loss.item()\n        pbar.set_postfix({'loss': f'{loss.item():.4f}'})\n    \n    return total_loss / len(loader)\n\ndef eval_fn(model, loader, criterion, device):\n    model.eval()\n    total_loss = 0\n    all_preds = []\n    all_gts = []\n    all_fs = []\n    \n    with torch.no_grad():\n        for images, signals, meta in tqdm(loader, desc=\"Evaluating\"):\n            images = images.to(device)\n            signals = signals.to(device)\n            \n            with torch.amp.autocast('cuda'):\n                preds = model(images)\n                loss = criterion(preds, signals)\n            \n            total_loss += loss.item()\n            \n            preds_np = preds.float().cpu().numpy()\n            signals_np = signals.float().cpu().numpy()\n            fs_np = meta[:, 0].numpy()\n            \n            for pred, gt, fs in zip(preds_np, signals_np, fs_np):\n                all_preds.append(pred)\n                all_gts.append(gt)\n                all_fs.append(int(fs))\n    \n    avg_loss = total_loss / len(loader)\n    snr_db = compute_competition_snr(all_preds, all_gts, all_fs)\n    \n    return avg_loss, snr_db\n\ndef run_training():\n    if not os.path.exists(Config.TRAIN_CSV):\n        print(f\"Train CSV not found.\")\n        return\n        \n    df = pd.read_csv(Config.TRAIN_CSV)\n    print(f\"Total: {len(df)}\")\n    \n    train_df = df.sample(frac=0.9, random_state=Config.SEED)\n    valid_df = df.drop(train_df.index)\n    print(f\"Train: {len(train_df)}, Valid: {len(valid_df)}\")\n    \n    train_ds = ECGDataset(train_df, Config.TRAIN_IMG_DIR, mode=\"train\", transform=get_transforms(Config, \"train\"))\n    valid_ds = ECGDataset(valid_df, Config.TRAIN_IMG_DIR, mode=\"valid\", transform=get_transforms(Config, \"valid\"))\n    \n    train_loader = DataLoader(train_ds, batch_size=Config.BATCH_SIZE, shuffle=True, num_workers=0, pin_memory=True)\n    valid_loader = DataLoader(valid_ds, batch_size=Config.BATCH_SIZE, shuffle=False, num_workers=0, pin_memory=True)\n    \n    print(\"Initializing ConvNeXt-Small...\")\n    model = ECGModel().to(Config.DEVICE)\n    print(f\"Params: {sum(p.numel() for p in model.parameters()):,}\")\n    \n    optimizer = torch.optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=Config.WEIGHT_DECAY)\n    criterion = nn.SmoothL1Loss()\n    scaler = torch.amp.GradScaler('cuda')\n    \n    scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, T_0=10, T_mult=2, eta_min=1e-6)\n    \n    best_snr = -float('inf')\n    patience = 15\n    patience_counter = 0\n    \n    for epoch in range(Config.EPOCHS):\n        print(f\"\\n{'='*40}\")\n        print(f\"Epoch {epoch+1}/{Config.EPOCHS} | LR: {optimizer.param_groups[0]['lr']:.2e}\")\n        \n        train_loss = train_fn(model, train_loader, optimizer, criterion, Config.DEVICE, scaler)\n        valid_loss, valid_snr = eval_fn(model, valid_loader, criterion, Config.DEVICE)\n        scheduler.step()\n        \n        print(f\"Train: {train_loss:.4f} | Valid: {valid_loss:.4f} | SNR: {valid_snr:.2f} dB\")\n        \n        if valid_snr > best_snr:\n            best_snr = valid_snr\n            torch.save(model.state_dict(), \"best_model.pth\")\n            print(f\">>> Best SNR: {best_snr:.2f} dB\")\n            patience_counter = 0\n        else:\n            patience_counter += 1\n            if patience_counter >= patience:\n                print(\"Early stop\")\n                break\n    \n    print(f\"\\n*** Final Best: {best_snr:.2f} dB ***\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T06:17:10.707222Z","iopub.execute_input":"2026-01-22T06:17:10.707432Z","iopub.status.idle":"2026-01-22T06:17:10.732896Z","shell.execute_reply.started":"2026-01-22T06:17:10.707411Z","shell.execute_reply":"2026-01-22T06:17:10.732350Z"}},"outputs":[],"execution_count":null},{"id":"a2b274fd","cell_type":"code","source":"def inference():\n    model = ECGModel().to(Config.DEVICE)\n    \n    weights_path = \"best_model.pth\"\n    if not os.path.exists(weights_path) and os.path.exists(\"/kaggle/input\"):\n        for root, dirs, files in os.walk(\"/kaggle/input\"):\n            if \"best_model.pth\" in files:\n                weights_path = os.path.join(root, \"best_model.pth\")\n                break\n    \n    if os.path.exists(weights_path):\n        model.load_state_dict(torch.load(weights_path, map_location=Config.DEVICE, weights_only=True))\n        print(f\"Loaded: {weights_path}\")\n    \n    model.eval()\n    \n    if not os.path.exists(Config.TEST_CSV):\n        return\n\n    test_df = pd.read_csv(Config.TEST_CSV)\n    transforms = get_transforms(Config, \"test\")\n    unique_ids = test_df['id'].unique()\n    \n    submission_rows = []\n    leads_map = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', 'V1', 'V2', 'V3', 'V4', 'V5', 'V6']\n    \n    print(f\"Inference: {len(unique_ids)} records\")\n    \n    for ecg_id in tqdm(unique_ids):\n        ecg_id_str = str(ecg_id)\n        record_df = test_df[test_df['id'] == ecg_id]\n        fs = record_df['fs'].iloc[0]\n        \n        lead_lengths = {row['lead']: row['number_of_rows'] for _, row in record_df.iterrows()}\n        \n        img_path = os.path.join(Config.TEST_IMG_DIR, f\"{ecg_id_str}.png\")\n        if not os.path.exists(img_path):\n            img_path = os.path.join(Config.TEST_IMG_DIR, ecg_id_str, f\"{ecg_id_str}.png\")\n            \n        image = cv2.imread(img_path)\n        if image is None:\n            image = np.zeros((*Config.IMAGE_SIZE, 3), dtype=np.uint8)\n        else:\n            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n            \n        augmented = transforms(image=image)\n        img_tensor = augmented['image'].unsqueeze(0).to(Config.DEVICE)\n        \n        with torch.no_grad():\n            with torch.amp.autocast('cuda'):\n                pred = model(img_tensor)\n        \n        pred = pred.float().cpu().numpy()[0]\n        \n        for lead_idx, lead_name in enumerate(leads_map):\n            lead_pred = pred[lead_idx]\n            target_len = lead_lengths.get(lead_name, int(fs * 2.5))\n            \n            resampled = np.interp(\n                np.linspace(0, 1, target_len),\n                np.linspace(0, 1, len(lead_pred)),\n                lead_pred\n            )\n                \n            for i, value in enumerate(resampled):\n                submission_rows.append({\"id\": f\"{ecg_id}_{i}_{lead_name}\", \"value\": float(value)})\n                \n    submission = pd.DataFrame(submission_rows)\n    submission.to_csv(\"submission.csv\", index=False)\n    print(f\"Saved {len(submission)} rows\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T06:17:10.733791Z","iopub.execute_input":"2026-01-22T06:17:10.734078Z","iopub.status.idle":"2026-01-22T06:17:10.755658Z","shell.execute_reply.started":"2026-01-22T06:17:10.734048Z","shell.execute_reply":"2026-01-22T06:17:10.754941Z"}},"outputs":[],"execution_count":null},{"id":"fc6fa43b-97f9-4fcf-8734-b20642b0e7de","cell_type":"code","source":"run_training()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T06:17:10.756563Z","iopub.execute_input":"2026-01-22T06:17:10.757014Z","iopub.status.idle":"2026-01-22T11:59:07.931737Z","shell.execute_reply.started":"2026-01-22T06:17:10.756991Z","shell.execute_reply":"2026-01-22T11:59:07.931015Z"}},"outputs":[],"execution_count":null},{"id":"f2d85008-aa6a-41cd-97dd-fad96199a309","cell_type":"code","source":"inference()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-22T11:59:07.932690Z","iopub.execute_input":"2026-01-22T11:59:07.932988Z","iopub.status.idle":"2026-01-22T11:59:14.660962Z","shell.execute_reply.started":"2026-01-22T11:59:07.932965Z","shell.execute_reply":"2026-01-22T11:59:14.660350Z"}},"outputs":[],"execution_count":null}]}