{"cells":[{"cell_type":"markdown","metadata":{},"source":"# PhysioNet ECG - U-Net Training\nTrain segmentation model for better signal extraction"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"import os\nimport cv2\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nfrom PIL import Image\nfrom tqdm import tqdm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom scipy.interpolate import interp1d\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f'Device: {device}')\n\nDATA_DIR = Path('/kaggle/input/physionet-ecg-image-digitization')\ntrain_df = pd.read_csv(DATA_DIR / 'train.csv')\ntest_df = pd.read_csv(DATA_DIR / 'test.csv')\nprint(f'Train: {len(train_df)}, Test: {len(test_df[\"id\"].unique())}')"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# U-Net Model\nclass DoubleConv(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.conv = nn.Sequential(\n            nn.Conv2d(in_ch, out_ch, 3, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_ch, out_ch, 3, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True)\n        )\n    def forward(self, x):\n        return self.conv(x)\n\nclass UNet(nn.Module):\n    def __init__(self, in_ch=3, out_ch=1):\n        super().__init__()\n        self.inc = DoubleConv(in_ch, 64)\n        self.down1 = nn.Sequential(nn.MaxPool2d(2), DoubleConv(64, 128))\n        self.down2 = nn.Sequential(nn.MaxPool2d(2), DoubleConv(128, 256))\n        self.down3 = nn.Sequential(nn.MaxPool2d(2), DoubleConv(256, 512))\n        self.down4 = nn.Sequential(nn.MaxPool2d(2), DoubleConv(512, 512))\n        \n        self.up1 = nn.ConvTranspose2d(512, 512, 2, stride=2)\n        self.conv1 = DoubleConv(1024, 512)\n        self.up2 = nn.ConvTranspose2d(512, 256, 2, stride=2)\n        self.conv2 = DoubleConv(512, 256)\n        self.up3 = nn.ConvTranspose2d(256, 128, 2, stride=2)\n        self.conv3 = DoubleConv(256, 128)\n        self.up4 = nn.ConvTranspose2d(128, 64, 2, stride=2)\n        self.conv4 = DoubleConv(128, 64)\n        self.outc = nn.Conv2d(64, out_ch, 1)\n        \n    def forward(self, x):\n        x1 = self.inc(x)\n        x2 = self.down1(x1)\n        x3 = self.down2(x2)\n        x4 = self.down3(x3)\n        x5 = self.down4(x4)\n        \n        x = self.up1(x5)\n        x = torch.cat([x, x4], dim=1)\n        x = self.conv1(x)\n        x = self.up2(x)\n        x = torch.cat([x, x3], dim=1)\n        x = self.conv2(x)\n        x = self.up3(x)\n        x = torch.cat([x, x2], dim=1)\n        x = self.conv3(x)\n        x = self.up4(x)\n        x = torch.cat([x, x1], dim=1)\n        x = self.conv4(x)\n        return self.outc(x)\n\nmodel = UNet().to(device)\nprint(f'Parameters: {sum(p.numel() for p in model.parameters()):,}')"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Dataset\ndef create_mask(image):\n    gray = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)\n    _, binary = cv2.threshold(gray, 80, 255, cv2.THRESH_BINARY_INV)\n    for ksize in [30, 50]:\n        h_k = cv2.getStructuringElement(cv2.MORPH_RECT, (ksize, 1))\n        v_k = cv2.getStructuringElement(cv2.MORPH_RECT, (1, ksize))\n        grid = cv2.bitwise_or(\n            cv2.morphologyEx(binary, cv2.MORPH_OPEN, h_k),\n            cv2.morphologyEx(binary, cv2.MORPH_OPEN, v_k)\n        )\n        binary = cv2.bitwise_and(binary, cv2.bitwise_not(grid))\n    kernel = np.ones((2,2), np.uint8)\n    return cv2.morphologyEx(cv2.morphologyEx(binary, cv2.MORPH_CLOSE, kernel), cv2.MORPH_OPEN, kernel)\n\nclass ECGDataset(Dataset):\n    def __init__(self, data_dir, record_ids, size=512):\n        self.samples = []\n        self.size = size\n        for rid in record_ids:\n            d = data_dir / 'train' / str(rid)\n            if d.exists():\n                for p in d.glob('*.png'):\n                    self.samples.append(p)\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        img = cv2.imread(str(self.samples[idx]))\n        if img is None:\n            img = np.array(Image.open(self.samples[idx]))\n            if img.shape[2] == 4:\n                img = cv2.cvtColor(img, cv2.COLOR_RGBA2BGR)\n            else:\n                img = cv2.cvtColor(img, cv2.COLOR_RGB2BGR)\n        \n        mask = create_mask(img)\n        img = cv2.resize(img, (self.size, self.size))\n        mask = cv2.resize(mask, (self.size, self.size))\n        \n        img = torch.from_numpy(img.transpose(2,0,1)).float() / 255.0\n        mask = torch.from_numpy(mask).float().unsqueeze(0) / 255.0\n        return img, mask"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Training\nids = train_df['id'].values\nnp.random.shuffle(ids)\nn_train = int(len(ids) * 0.9)\ntrain_ds = ECGDataset(DATA_DIR, ids[:n_train])\nval_ds = ECGDataset(DATA_DIR, ids[n_train:])\n\ntrain_loader = DataLoader(train_ds, batch_size=8, shuffle=True, num_workers=2)\nval_loader = DataLoader(val_ds, batch_size=8, shuffle=False, num_workers=2)\n\nprint(f'Train: {len(train_ds)}, Val: {len(val_ds)}')"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Loss and optimizer\nclass DiceBCELoss(nn.Module):\n    def forward(self, pred, target):\n        bce = F.binary_cross_entropy_with_logits(pred, target)\n        pred_sig = torch.sigmoid(pred).view(-1)\n        target_flat = target.view(-1)\n        inter = (pred_sig * target_flat).sum()\n        dice = 1 - (2*inter + 1e-6) / (pred_sig.sum() + target_flat.sum() + 1e-6)\n        return 0.5 * bce + 0.5 * dice\n\ncriterion = DiceBCELoss()\noptimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience=3)"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Train loop\nbest_loss = float('inf')\nfor epoch in range(15):\n    model.train()\n    train_loss = 0\n    for imgs, masks in tqdm(train_loader, desc=f'Epoch {epoch+1}'):\n        imgs, masks = imgs.to(device), masks.to(device)\n        optimizer.zero_grad()\n        out = model(imgs)\n        loss = criterion(out, masks)\n        loss.backward()\n        optimizer.step()\n        train_loss += loss.item()\n    train_loss /= len(train_loader)\n    \n    model.eval()\n    val_loss = 0\n    with torch.no_grad():\n        for imgs, masks in val_loader:\n            imgs, masks = imgs.to(device), masks.to(device)\n            out = model(imgs)\n            val_loss += criterion(out, masks).item()\n    val_loss /= len(val_loader)\n    \n    scheduler.step(val_loss)\n    print(f'Epoch {epoch+1}: train={train_loss:.4f}, val={val_loss:.4f}')\n    \n    if val_loss < best_loss:\n        best_loss = val_loss\n        torch.save(model.state_dict(), 'best_unet.pth')\n        print('  Saved best model')"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Inference\nmodel.load_state_dict(torch.load('best_unet.pth'))\nmodel.eval()\n\nLEAD_LAYOUT = [\n    ('I', 0, 0), ('aVR', 0, 1), ('V1', 0, 2), ('V4', 0, 3),\n    ('II', 1, 0), ('aVL', 1, 1), ('V2', 1, 2), ('V5', 1, 3),\n    ('III', 2, 0), ('aVF', 2, 1), ('V3', 2, 2), ('V6', 2, 3),\n]\n\ndef extract_trace(mask):\n    h, w = mask.shape\n    signal = np.full(w, np.nan)\n    for x in range(w):\n        col = mask[:, x]\n        nz = np.where(col > 0)[0]\n        if len(nz) > 0:\n            signal[x] = np.mean(nz)\n    valid = ~np.isnan(signal)\n    if np.any(valid):\n        signal[valid] = h - signal[valid]\n    if np.any(~valid) and np.any(valid):\n        vi = np.where(valid)[0]\n        signal = interp1d(vi, signal[valid], bounds_error=False, fill_value='extrapolate')(np.arange(w))\n    return signal\n\ndef predict_and_extract(image_path, test_info):\n    img = cv2.imread(str(image_path))\n    if img is None:\n        img = np.array(Image.open(image_path))\n        if img.shape[2] == 4:\n            img = cv2.cvtColor(img, cv2.COLOR_RGBA2BGR)\n        else:\n            img = cv2.cvtColor(img, cv2.COLOR_RGB2BGR)\n    \n    oh, ow = img.shape[:2]\n    img_resized = cv2.resize(img, (512, 512))\n    tensor = torch.from_numpy(img_resized.transpose(2,0,1)).float().unsqueeze(0) / 255.0\n    \n    with torch.no_grad():\n        out = model(tensor.to(device))\n        mask = (torch.sigmoid(out) > 0.5).cpu().numpy()[0, 0].astype(np.uint8) * 255\n    \n    mask = cv2.resize(mask, (ow, oh))\n    \n    h, w = mask.shape\n    row_h, col_w = h // 4, w // 4\n    leads = {}\n    \n    for lead, row, col in LEAD_LAYOUT:\n        roi = mask[row*row_h:(row+1)*row_h, col*col_w:(col+1)*col_w]\n        trace = extract_trace(roi)\n        info = test_info[test_info['lead'] == lead]\n        tlen = info['number_of_rows'].values[0] if len(info) > 0 else 2500\n        trace = np.interp(np.linspace(0,1,tlen), np.linspace(0,1,len(trace)), trace)\n        trace = trace - np.median(trace)\n        mad = np.median(np.abs(trace - np.median(trace)))\n        if mad > 0:\n            trace = trace / (mad * 1.4826) * 0.3\n        leads[lead] = trace\n    \n    roi = mask[3*row_h:, :]\n    trace = extract_trace(roi)\n    info = test_info[test_info['lead'] == 'II']\n    tlen = info['number_of_rows'].values[0] if len(info) > 0 else 10000\n    trace = np.interp(np.linspace(0,1,tlen), np.linspace(0,1,len(trace)), trace)\n    trace = trace - np.median(trace)\n    mad = np.median(np.abs(trace - np.median(trace)))\n    if mad > 0:\n        trace = trace / (mad * 1.4826) * 0.3\n    leads['II'] = trace\n    \n    return leads"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"# Generate submission\nrows = []\nfor img_id in tqdm(test_df['id'].unique()):\n    img_path = DATA_DIR / 'test' / f'{img_id}.png'\n    if not img_path.exists():\n        continue\n    img_info = test_df[test_df['id'] == img_id]\n    leads = predict_and_extract(img_path, img_info)\n    \n    for _, row in img_info.iterrows():\n        lead, n = row['lead'], row['number_of_rows']\n        sig = leads.get(lead, np.zeros(n))[:n]\n        for i, v in enumerate(sig):\n            rows.append({'id': f'{img_id}_{i}_{lead}', 'value': float(v) if not np.isnan(v) else 0.0})\n\nsub = pd.DataFrame(rows)\nsub['value'] = sub['value'].fillna(0)\nsub.to_csv('submission.csv', index=False)\nprint(f'Shape: {sub.shape}')\nprint('Done!')"}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python"}},"nbformat":4,"nbformat_minor":4}