{"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,"sourceType":"competition"}],"dockerImageVersionId":31240,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-05T07:34:10.769436Z","iopub.execute_input":"2026-01-05T07:34:10.770040Z","iopub.status.idle":"2026-01-05T07:34:17.861451Z","shell.execute_reply.started":"2026-01-05T07:34:10.770015Z","shell.execute_reply":"2026-01-05T07:34:17.860632Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport cv2\nfrom pathlib import Path\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom scipy import signal, interpolate\nfrom scipy.ndimage import median_filter\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm import tqdm\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# =====================================================================\n# CONFIGURATION\n# =====================================================================\nclass Config:\n    # Paths\n    TRAIN_DIR = '/kaggle/input/physionet-ecg-image-digitization/train'\n    TEST_DIR = '/kaggle/input/physionet-ecg-image-digitization/test'\n    TRAIN_CSV = '/kaggle/input/physionet-ecg-image-digitization/train.csv'\n    TEST_CSV = '/kaggle/input/physionet-ecg-image-digitization/test.csv'\n    OUTPUT_DIR = '/kaggle/working'\n    \n    # Model parameters\n    IMAGE_SIZE = (512, 512)\n    BATCH_SIZE = 4\n    NUM_EPOCHS = 15\n    LEARNING_RATE = 1e-4\n    DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n    NUM_WORKERS = 2\n    \n    # ECG parameters\n    LEADS = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', 'V1', 'V2', 'V3', 'V4', 'V5', 'V6']\n    LEAD_II_DURATION = 10.0  # seconds\n    OTHER_LEAD_DURATION = 2.5  # seconds\n    \n    # Signal processing\n    SAVGOL_WINDOW = 11\n    SAVGOL_POLY = 3\n    BASELINE_WINDOW = 201\n\n# =====================================================================\n# PREPROCESSING UTILITIES\n# =====================================================================\nclass ECGPreprocessor:\n    @staticmethod\n    def apply_clahe(img, clip_limit=2.0, tile_grid_size=(8, 8)):\n        \"\"\"Apply CLAHE for contrast enhancement\"\"\"\n        if len(img.shape) == 3:\n            img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n        clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=tile_grid_size)\n        return clahe.apply(img)\n    \n    @staticmethod\n    def remove_grid(img, freq_threshold=0.85):\n        \"\"\"Remove grid patterns using FFT\"\"\"\n        f_transform = np.fft.fft2(img)\n        f_shift = np.fft.fftshift(f_transform)\n        \n        rows, cols = img.shape\n        crow, ccol = rows // 2, cols // 2\n        \n        # Create mask to remove high frequency components (grid)\n        mask = np.ones((rows, cols), np.uint8)\n        r = int(min(rows, cols) * freq_threshold / 2)\n        center = (crow, ccol)\n        x, y = np.ogrid[:rows, :cols]\n        mask_area = (x - center[0])**2 + (y - center[1])**2 <= r*r\n        mask[~mask_area] = 0\n        \n        f_shift_filtered = f_shift * mask\n        f_ishift = np.fft.ifftshift(f_shift_filtered)\n        img_back = np.fft.ifft2(f_ishift)\n        img_back = np.abs(img_back)\n        \n        return img_back.astype(np.uint8)\n    \n    @staticmethod\n    def preprocess_image(img_path):\n        \"\"\"Complete preprocessing pipeline\"\"\"\n        img = cv2.imread(str(img_path))\n        if img is None:\n            raise ValueError(f\"Could not read image: {img_path}\")\n        \n        # Convert to grayscale\n        gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n        \n        # Apply CLAHE\n        clahe = ECGPreprocessor.apply_clahe(gray)\n        \n        # Remove grid (optional, can be slow)\n        # clahe = ECGPreprocessor.remove_grid(clahe)\n        \n        # Normalize\n        clahe = cv2.normalize(clahe, None, 0, 255, cv2.NORM_MINMAX)\n        \n        return clahe, img\n\n# =====================================================================\n# ATTENTION U-NET++ MODEL\n# =====================================================================\nclass ConvBlock(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    \n    def forward(self, x):\n        return self.conv(x)\n\nclass AttentionBlock(nn.Module):\n    def __init__(self, F_g, F_l, F_int):\n        super().__init__()\n        self.W_g = nn.Sequential(\n            nn.Conv2d(F_g, F_int, 1, stride=1, padding=0, bias=True),\n            nn.BatchNorm2d(F_int)\n        )\n        self.W_x = nn.Sequential(\n            nn.Conv2d(F_l, F_int, 1, stride=1, padding=0, bias=True),\n            nn.BatchNorm2d(F_int)\n        )\n        self.psi = nn.Sequential(\n            nn.Conv2d(F_int, 1, 1, stride=1, padding=0, bias=True),\n            nn.BatchNorm2d(1),\n            nn.Sigmoid()\n        )\n        self.relu = nn.ReLU(inplace=True)\n    \n    def forward(self, g, x):\n        g1 = self.W_g(g)\n        x1 = self.W_x(x)\n        psi = self.relu(g1 + x1)\n        psi = self.psi(psi)\n        return x * psi\n\nclass AttentionUNetPlusPlus(nn.Module):\n    def __init__(self, in_channels=1, out_channels=1, features=[64, 128, 256, 512]):\n        super().__init__()\n        \n        # Encoder\n        self.enc1 = ConvBlock(in_channels, features[0])\n        self.pool1 = nn.MaxPool2d(2, 2)\n        \n        self.enc2 = ConvBlock(features[0], features[1])\n        self.pool2 = nn.MaxPool2d(2, 2)\n        \n        self.enc3 = ConvBlock(features[1], features[2])\n        self.pool3 = nn.MaxPool2d(2, 2)\n        \n        self.enc4 = ConvBlock(features[2], features[3])\n        self.pool4 = nn.MaxPool2d(2, 2)\n        \n        # Bottleneck\n        self.bottleneck = ConvBlock(features[3], features[3] * 2)\n        \n        # Decoder with attention\n        self.up4 = nn.ConvTranspose2d(features[3] * 2, features[3], 2, stride=2)\n        self.att4 = AttentionBlock(features[3], features[3], features[3] // 2)\n        self.dec4 = ConvBlock(features[3] * 2, features[3])\n        \n        self.up3 = nn.ConvTranspose2d(features[3], features[2], 2, stride=2)\n        self.att3 = AttentionBlock(features[2], features[2], features[2] // 2)\n        self.dec3 = ConvBlock(features[2] * 2, features[2])\n        \n        self.up2 = nn.ConvTranspose2d(features[2], features[1], 2, stride=2)\n        self.att2 = AttentionBlock(features[1], features[1], features[1] // 2)\n        self.dec2 = ConvBlock(features[1] * 2, features[1])\n        \n        self.up1 = nn.ConvTranspose2d(features[1], features[0], 2, stride=2)\n        self.att1 = AttentionBlock(features[0], features[0], features[0] // 2)\n        self.dec1 = ConvBlock(features[0] * 2, features[0])\n        \n        self.final = nn.Conv2d(features[0], out_channels, 1)\n    \n    def forward(self, x):\n        # Encoder\n        e1 = self.enc1(x)\n        p1 = self.pool1(e1)\n        \n        e2 = self.enc2(p1)\n        p2 = self.pool2(e2)\n        \n        e3 = self.enc3(p2)\n        p3 = self.pool3(e3)\n        \n        e4 = self.enc4(p3)\n        p4 = self.pool4(e4)\n        \n        # Bottleneck\n        b = self.bottleneck(p4)\n        \n        # Decoder with attention - match sizes with interpolation\n        d4 = self.up4(b)\n        d4 = F.interpolate(d4, size=e4.shape[2:], mode='bilinear', align_corners=False)\n        e4 = self.att4(d4, e4)\n        d4 = torch.cat([d4, e4], dim=1)\n        d4 = self.dec4(d4)\n        \n        d3 = self.up3(d4)\n        d3 = F.interpolate(d3, size=e3.shape[2:], mode='bilinear', align_corners=False)\n        e3 = self.att3(d3, e3)\n        d3 = torch.cat([d3, e3], dim=1)\n        d3 = self.dec3(d3)\n        \n        d2 = self.up2(d3)\n        d2 = F.interpolate(d2, size=e2.shape[2:], mode='bilinear', align_corners=False)\n        e2 = self.att2(d2, e2)\n        d2 = torch.cat([d2, e2], dim=1)\n        d2 = self.dec2(d2)\n        \n        d1 = self.up1(d2)\n        d1 = F.interpolate(d1, size=e1.shape[2:], mode='bilinear', align_corners=False)\n        e1 = self.att1(d1, e1)\n        d1 = torch.cat([d1, e1], dim=1)\n        d1 = self.dec1(d1)\n        \n        return torch.sigmoid(self.final(d1))\n\n# =====================================================================\n# DATASET\n# =====================================================================\nclass ECGDataset(Dataset):\n    def __init__(self, df, image_dir, transform=None, is_train=True):\n        self.df = df\n        self.image_dir = Path(image_dir)\n        self.transform = transform\n        self.is_train = is_train\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_id = row['id']\n        \n        if self.is_train:\n            # Use segment 0001 (original) for training\n            img_path = self.image_dir / str(img_id) / f\"{img_id}-0001.png\"\n            \n            # Load signal\n            signal_path = self.image_dir / str(img_id) / f\"{img_id}.csv\"\n            signal_df = pd.read_csv(signal_path)\n        else:\n            img_path = self.image_dir / f\"{img_id}.png\"\n        \n        # Load and preprocess image\n        img, _ = ECGPreprocessor.preprocess_image(img_path)\n        \n        # Ensure float32 type\n        img = img.astype(np.float32) / 255.0\n        \n        if self.transform:\n            augmented = self.transform(image=img)\n            img = augmented['image']\n        else:\n            img = torch.from_numpy(img).unsqueeze(0).float()\n        \n        if self.is_train:\n            # Create mask (simplified - in practice, you'd need proper lead segmentation)\n            mask = self.create_mask_from_signal(img, signal_df)\n            return img, mask, row['fs']\n        else:\n            return img, row['fs']\n    \n    def create_mask_from_signal(self, img, signal_df):\n        \"\"\"Create a binary mask from signal data (simplified)\"\"\"\n        # This is a placeholder - you'd need proper implementation\n        mask = torch.zeros_like(img)\n        return mask\n\n# =====================================================================\n# SIGNAL EXTRACTION AND POST-PROCESSING\n# =====================================================================\nclass SignalExtractor:\n    @staticmethod\n    def extract_centerline(mask):\n        \"\"\"Extract centerline from binary mask using column-wise median\"\"\"\n        mask_np = mask.squeeze().cpu().numpy()\n        h, w = mask_np.shape\n        centerline = []\n        \n        for col in range(w):\n            column = mask_np[:, col]\n            if column.sum() > 0:\n                # Find median position of signal\n                indices = np.where(column > 0.5)[0]\n                if len(indices) > 0:\n                    median_pos = np.median(indices)\n                    centerline.append(median_pos)\n                else:\n                    centerline.append(np.nan)\n            else:\n                centerline.append(np.nan)\n        \n        # Interpolate missing values\n        centerline = np.array(centerline)\n        valid_idx = ~np.isnan(centerline)\n        if valid_idx.sum() > 0:\n            x = np.arange(len(centerline))\n            centerline = np.interp(x, x[valid_idx], centerline[valid_idx])\n        \n        return centerline\n    \n    @staticmethod\n    def remove_baseline(signal, window_size=201):\n        \"\"\"Remove baseline wander using median filtering\"\"\"\n        baseline = median_filter(signal, size=window_size)\n        return signal - baseline\n    \n    @staticmethod\n    def smooth_signal(signal, window=11, poly=3):\n        \"\"\"Apply Savitzky-Golay smoothing\"\"\"\n        if len(signal) < window:\n            return signal\n        return signal.savgol_filter(signal, window, poly)\n    \n    @staticmethod\n    def resample_signal(signal, original_length, target_length):\n        \"\"\"Resample signal to target length\"\"\"\n        x_old = np.linspace(0, 1, original_length)\n        x_new = np.linspace(0, 1, target_length)\n        f = interpolate.interp1d(x_old, signal, kind='cubic', fill_value='extrapolate')\n        return f(x_new)\n\n# =====================================================================\n# TRAINING\n# =====================================================================\ndef train_model(model, train_loader, criterion, optimizer, device, epoch):\n    model.train()\n    running_loss = 0.0\n    \n    pbar = tqdm(train_loader, desc=f'Epoch {epoch}')\n    for images, masks, fs in pbar:\n        images = images.to(device)\n        masks = masks.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = criterion(outputs, masks)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        pbar.set_postfix({'loss': running_loss / (pbar.n + 1)})\n    \n    return running_loss / len(train_loader)\n\n# =====================================================================\n# INFERENCE\n# =====================================================================\ndef predict_ecg(model, image_path, fs, device, config):\n    \"\"\"Predict ECG signals from image\"\"\"\n    model.eval()\n    \n    # Preprocess image\n    img, _ = ECGPreprocessor.preprocess_image(image_path)\n    img_tensor = torch.from_numpy(img).unsqueeze(0).unsqueeze(0).float() / 255.0\n    img_tensor = img_tensor.to(device)\n    \n    with torch.no_grad():\n        mask = model(img_tensor)\n    \n    # Extract signals for each lead\n    predictions = {}\n    \n    # Split image into 12 lead regions (simplified grid layout)\n    h, w = mask.shape[2:]\n    lead_height = h // 4  # Assuming 4 rows\n    lead_width = w // 3   # Assuming 3 columns\n    \n    for idx, lead in enumerate(config.LEADS):\n        row = idx // 3\n        col = idx % 3\n        \n        # Extract lead region\n        y1, y2 = row * lead_height, (row + 1) * lead_height\n        x1, x2 = col * lead_width, (col + 1) * lead_width\n        lead_mask = mask[:, :, y1:y2, x1:x2]\n        \n        # Extract centerline\n        centerline = SignalExtractor.extract_centerline(lead_mask)\n        \n        # Normalize to mV (approximate conversion)\n        signal_mv = (centerline - centerline.mean()) * 0.01\n        \n        # Post-process\n        signal_mv = SignalExtractor.remove_baseline(signal_mv)\n        signal_mv = signal.savgol_filter(signal_mv, \n                                         min(config.SAVGOL_WINDOW, len(signal_mv)//2*2-1), \n                                         config.SAVGOL_POLY)\n        \n        # Resample to target length\n        if lead == 'II':\n            target_len = int(fs * config.LEAD_II_DURATION)\n        else:\n            target_len = int(fs * config.OTHER_LEAD_DURATION)\n        \n        signal_mv = SignalExtractor.resample_signal(signal_mv, len(signal_mv), target_len)\n        predictions[lead] = signal_mv\n    \n    return predictions\n\n# =====================================================================\n# MAIN EXECUTION\n# =====================================================================\ndef main():\n    config = Config()\n    \n    # Load data\n    print(\"Loading data...\")\n    train_df = pd.read_csv(config.TRAIN_CSV)\n    test_df = pd.read_csv(config.TEST_CSV)\n    \n    # Initialize model\n    print(\"Initializing model...\")\n    model = AttentionUNetPlusPlus(in_channels=1, out_channels=1).to(config.DEVICE)\n    \n    # Training transforms\n    train_transform = A.Compose([\n        A.Resize(config.IMAGE_SIZE[0], config.IMAGE_SIZE[1]),\n        A.RandomBrightnessContrast(p=0.5),\n        A.GaussNoise(p=0.3),\n        A.Rotate(limit=5, p=0.3),\n        A.Normalize(mean=0.0, std=1.0),\n        ToTensorV2()\n    ])\n    \n    # Create datasets\n    train_dataset = ECGDataset(train_df, config.TRAIN_DIR, transform=train_transform)\n    train_loader = DataLoader(train_dataset, batch_size=config.BATCH_SIZE, \n                             shuffle=True, num_workers=config.NUM_WORKERS)\n    \n    # Training setup\n    criterion = nn.BCELoss()\n    optimizer = torch.optim.Adam(model.parameters(), lr=config.LEARNING_RATE)\n    \n    # Train model\n    print(\"Training model...\")\n    for epoch in range(config.NUM_EPOCHS):\n        loss = train_model(model, train_loader, criterion, optimizer, config.DEVICE, epoch)\n        print(f\"Epoch {epoch}: Loss = {loss:.4f}\")\n    \n    # Save model\n    torch.save(model.state_dict(), os.path.join(config.OUTPUT_DIR, 'model.pth'))\n    \n    # Inference\n    print(\"Running inference...\")\n    predictions_list = []\n    \n    for idx, row in tqdm(test_df.iterrows(), total=len(test_df)):\n        img_id = row['id']\n        fs = row['fs']\n        img_path = Path(config.TEST_DIR) / f\"{img_id}.png\"\n        \n        try:\n            predictions = predict_ecg(model, img_path, fs, config.DEVICE, config)\n            \n            # Format predictions\n            for lead in config.LEADS:\n                signal_values = predictions[lead]\n                duration = config.LEAD_II_DURATION if lead == 'II' else config.OTHER_LEAD_DURATION\n                num_samples = int(fs * duration)\n                \n                for row_id in range(num_samples):\n                    predictions_list.append({\n                        'id': f\"{img_id}_{row_id}_{lead}\",\n                        'value': signal_values[row_id] if row_id < len(signal_values) else 0.0\n                    })\n        except Exception as e:\n            print(f\"Error processing {img_id}: {e}\")\n            # Add zeros as fallback\n            for lead in config.LEADS:\n                duration = config.LEAD_II_DURATION if lead == 'II' else config.OTHER_LEAD_DURATION\n                num_samples = int(fs * duration)\n                for row_id in range(num_samples):\n                    predictions_list.append({\n                        'id': f\"{img_id}_{row_id}_{lead}\",\n                        'value': 0.0\n                    })\n    \n    # Create submission\n    submission_df = pd.DataFrame(predictions_list)\n    submission_df.to_parquet(os.path.join(config.OUTPUT_DIR, 'submission.parquet'), index=False)\n    print(\"Submission saved!\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-05T07:35:08.321230Z","iopub.execute_input":"2026-01-05T07:35:08.322017Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}