{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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":31234,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Setting Kaggle Environment for PhysioNet - ECG image problem","metadata":{}},{"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-01T14:02:22.649937Z","iopub.execute_input":"2026-01-01T14:02:22.650285Z","iopub.status.idle":"2026-01-01T14:02:30.433291Z","shell.execute_reply.started":"2026-01-01T14:02:22.650257Z","shell.execute_reply":"2026-01-01T14:02:30.43267Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Import Required Libraries\nLoading all necessary packages for data processing, visualization, and deep learning with PyTorch.","metadata":{}},{"cell_type":"code","source":"# Import libraries\nimport os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom PIL import Image, ImageOps, ImageFilter, ImageEnhance\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm import tqdm\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T14:02:30.434395Z","iopub.execute_input":"2026-01-01T14:02:30.434741Z","iopub.status.idle":"2026-01-01T14:02:37.870177Z","shell.execute_reply.started":"2026-01-01T14:02:30.434718Z","shell.execute_reply":"2026-01-01T14:02:37.869316Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Configuration and Device Setup\nSetting up data directories and detecting available GPU/CPU for training.","metadata":{}},{"cell_type":"code","source":"# Paths and config\nDATA_DIR = '/kaggle/input/physionet-ecg-image-digitization'\nTRAIN_DIR = f'{DATA_DIR}/train'\nTEST_DIR = f'{DATA_DIR}/test'\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f'Using device: {DEVICE}')\n\nBATCH_SIZE = 8\nEPOCHS = 5\nLR = 0.0001","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T14:02:37.871677Z","iopub.execute_input":"2026-01-01T14:02:37.872173Z","iopub.status.idle":"2026-01-01T14:02:37.964737Z","shell.execute_reply.started":"2026-01-01T14:02:37.872138Z","shell.execute_reply":"2026-01-01T14:02:37.963777Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Load Metadata and Split Dataset","metadata":{}},{"cell_type":"code","source":"# Load metadata\nfrom sklearn.model_selection import train_test_split\n\ntrain_df = pd.read_csv(f'{DATA_DIR}/train.csv')\ntest_df = pd.read_csv(f'{DATA_DIR}/test.csv')\n\n# Get available record IDs from disk\navailable_train_ids = [int(d) for d in os.listdir(TRAIN_DIR) if os.path.isdir(os.path.join(TRAIN_DIR, d))]\navailable_test_ids = [int(f.replace('.png', '')) for f in os.listdir(TEST_DIR) if f.endswith('.png')]\n\nprint(f'Original train samples in csv: {len(train_df)}, available on disk: {len(available_train_ids)}')\nprint(f'Original test samples in csv: {len(test_df)}, available on disk: {len(available_test_ids)}')\n\n# Filter to only use available samples\ntrain_df = train_df[train_df['id'].isin(available_train_ids)].reset_index(drop=True)\ntest_df = test_df[test_df['id'].isin(available_test_ids)].reset_index(drop=True)\n\nprint(f'\\nFiltered train samples: {len(train_df)}')\nprint(f'Filtered test samples: {len(test_df)}')\n\n# Split train into train and validation (80/20)\nif len(train_df) > 0:\n    train_df, val_df = train_test_split(train_df, test_size=0.2, random_state=42)\n    train_df = train_df.reset_index(drop=True)\n    val_df = val_df.reset_index(drop=True)\n    print(f'\\nAfter 80/20 split:')\n    print(f'Train samples: {len(train_df)}')\n    print(f'Validation samples: {len(val_df)}')\nelse:\n    val_df = pd.DataFrame()\n    print('\\nWARNING: No training data available!')\n\ntrain_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T14:02:37.965927Z","iopub.execute_input":"2026-01-01T14:02:37.966236Z","iopub.status.idle":"2026-01-01T14:02:39.045834Z","shell.execute_reply.started":"2026-01-01T14:02:37.966212Z","shell.execute_reply":"2026-01-01T14:02:39.044948Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# ECG Dataset Class\nCustom PyTorch dataset to load ECG images and ground truth signals. Images are resized from original size (varies) to 256x256 pixels, then normalized to [0,1]. Signals are interpolated to a standard length for consistent batching.","metadata":{}},{"cell_type":"code","source":"# Old Implementation\n# Dataset class for ECG images\nclass ECGDataset(Dataset):\n    def __init__(self, df, data_dir, is_train=True, standard_len=10000):\n        self.df = df\n        self.data_dir = data_dir\n        self.is_train = is_train\n        #self.img_size = img_size\n        self.standard_len = standard_len  # Standard length for all leads via interpolation\n        self.leads = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', 'V1', 'V2', 'V3', 'V4', 'V5', 'V6']\n        self.sig_row = ['I, aVR, V1, V4', 'II, aVL, V2, V5', 'III, aVF, V3, V6', 'II']\n\n    def __len__(self):\n        return len(self.df)\n    \n    def interpolate_signal(self, signal, target_length):\n        \"\"\"Interpolate signal to target length using linear interpolation\"\"\"\n        if len(signal) == target_length:\n            return signal\n        \n        # Create indices for interpolation\n        old_indices = np.linspace(0, len(signal) - 1, len(signal))\n        new_indices = np.linspace(0, len(signal) - 1, target_length)\n        \n        # Linear interpolation\n        interpolated = np.interp(new_indices, old_indices, signal)\n        return interpolated\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        record_id = str(row['id'])\n\n        try:\n            if self.is_train:\n                # Train: images are in subdirectories with multiple pages\n                img_path = f'{self.data_dir}/{record_id}/{record_id}-0001.png'\n            else:\n                # Test: images are directly in the test directory\n                img_path = f'{self.data_dir}/{record_id}.png'\n\n            # Load and preprocess image\n            if not os.path.exists(img_path):\n                raise FileNotFoundError(f\"Image not found: {img_path}\")\n                \n            img = Image.open(img_path)\n            img = img.resize(self.img_size)\n            img = np.array(img) / 255.0  # Normalize to [0, 1]\n            img = torch.FloatTensor(img).permute(2, 0, 1)  # (C, H, W)\n\n            # 1. Load image\n            img = Image.open(img_path)\n            img_height, img_width = img.shape[:2]\n\n            x_start, x_end = 50, img_width - 50  # crop horizontally (width)\n            y_start, y_end = 550, img_height - 75   # crop vertically (height)\n\n            # Cropping\n            img_crop = img[y_start:y_end, x_start:x_end]\n        \n            # 2. Convert to grayscale\n            gray = ImageOps.grayscale(img_crop)\n        \n            # 3. Noise reduction (median filter)\n            denoised = gray.filter(ImageFilter.MedianFilter(size=3))\n        \n            # 4. Contrast enhancement\n            enhancer = ImageEnhance.Contrast(denoised)\n            enhanced = enhancer.enhance(2.0)  # adjust factor as needed\n        \n            # 5. Resize to target size\n            # resized = enhanced.resize(img_size)\n        \n            # 6. Convert to numpy and normalize\n            arr = np.array(enhanced).astype(np.float32) / 255.0\n        \n            # 7. Add channel dimension (grayscale → 1 channel)\n            arr = np.expand_dims(arr, axis=-1)  # shape (H, W, 1)\n        \n            # 8. Convert to torch tensor and permute to (C, H, W)\n            img_tensor = torch.FloatTensor(arr).permute(2, 0, 1)\n\n            if self.is_train:\n                # Load ground truth CSV\n                csv_path = f'{self.data_dir}/{record_id}/{record_id}.csv'\n                if not os.path.exists(csv_path):\n                    raise FileNotFoundError(f\"CSV not found: {csv_path}\")\n                    \n                signal_df = pd.read_csv(csv_path)\n\n                # Process each lead separately and interpolate to standard_len\n                processed_leads = []\n                for lead in self.leads:\n                    if lead not in signal_df.columns:\n                        raise ValueError(f\"Lead {lead} not found in CSV for record {record_id}\")\n                        \n                    # Get non-null values for this lead\n                    lead_signal = signal_df[lead].dropna().values\n                    \n                    # Interpolate to standard length\n                    interpolated = self.interpolate_signal(lead_signal, self.standard_len)\n                    processed_leads.append(interpolated)\n\n                # Stack all leads into shape (12, standard_len)\n                signal = np.stack(processed_leads, axis=0)\n                signal = torch.FloatTensor(signal)\n\n                # Get metadata\n                fs = row['fs']\n                sig_len = row['sig_len']\n\n                return img_tensor, signal, fs, sig_len, record_id\n            else:\n                # For test, just return image and record_id\n                return img_tensor, record_id\n                \n        except Exception as e:\n            print(f\"Error loading record {record_id}: {str(e)}\")\n            raise","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-03T10:20:23.643093Z","iopub.execute_input":"2026-01-03T10:20:23.643739Z","iopub.status.idle":"2026-01-03T10:20:23.665937Z","shell.execute_reply.started":"2026-01-03T10:20:23.643706Z","shell.execute_reply":"2026-01-03T10:20:23.665023Z"},"jupyter":{"source_hidden":true,"outputs_hidden":true},"collapsed":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## My Implementation","metadata":{}},{"cell_type":"code","source":"class ECGDataset(Dataset):\n    def __init__(self, df, data_dir, img_size=(256,256), is_train=True, standard_len=10000):\n        self.df = df\n        self.data_dir = data_dir\n        self.img_size = img_size\n        self.is_train = is_train\n        self.standard_len = standard_len\n        self.ratios = [8,7,7,5]\n        self.lead_groups = [\n            ['I','aVR','V1','V4'],\n            ['II','aVL','V2','V5'],\n            ['III','aVF','V3','V6'],\n            ['II']\n        ]\n\n    def __len__(self):\n        return len(self.df)\n\n    def interpolate_signal(self, signal, target_length):\n        if len(signal) == target_length:\n            return signal\n        old_indices = np.linspace(0, len(signal)-1, len(signal))\n        new_indices = np.linspace(0, len(signal)-1, target_length)\n        return np.interp(new_indices, old_indices, signal)\n\n    def preprocess_strip(self, arr_strip):\n        # Convert to PIL for filters\n        img_strip = Image.fromarray(arr_strip)        \n        gray = ImageOps.grayscale(img_strip)\n        denoised = gray.filter(ImageFilter.MedianFilter(size=3))\n        enhanced = ImageEnhance.Contrast(denoised).enhance(2.0)\n        #resized = enhanced.resize(self.img_size)\n        arr_final = np.array(enhanced).astype(np.float32)/255.0\n        arr_final = np.expand_dims(arr_final, axis=-1)  # (H,W,1)\n        tensor = torch.FloatTensor(arr_final).permute(2,0,1)  # (C,H,W)\n        return tensor\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        record_id = str(row['id'])\n\n        if self.is_train:\n            img_path = f\"{self.data_dir}/{record_id}/{record_id}-0001.png\"\n            csv_path = f\"{self.data_dir}/{record_id}/{record_id}.csv\"\n        else:\n            img_path = f\"{self.data_dir}/{record_id}.png\"\n            csv_path = None\n\n        if not os.path.exists(img_path):\n            raise FileNotFoundError(f\"Image not found: {img_path}\")\n\n        # --- Loading full image ---\n        img = Image.open(img_path).convert(\"RGB\")\n        img_width, img_height = img.size\n        x_start, x_end = 50, img_width - 50   # horizontal crop\n        y_start, y_end = 550, img_height - 75 # vertical crop\n        img_crop = img.crop((x_start, y_start, x_end, y_end))\n\n        arr = np.array(img_crop)\n        h, w, _ = arr.shape\n\n        # --- Split into 4 strips ---\n        h_total = sum(self.ratios)\n        split_heights = [int(r/h_total * h) for r in self.ratios]\n\n        img_tensors = []\n        y = 0\n        for sh in split_heights:\n            y1, y2 = y, y+sh\n            arr_strip = arr[y1:y2,:]\n            img_tensors.append(self.preprocess_strip(arr_strip))\n            y = y2\n\n        # --- Signals ---\n        signal_groups = []\n        if self.is_train and csv_path and os.path.exists(csv_path):\n            signal_df = pd.read_csv(csv_path)\n            for leads in self.lead_groups:\n                # processed_leads = []\n                strip_sig = np.array([])\n                \n                if leads == self.lead_groups[1]:\n                    for lead in leads:\n                        if lead not in signal_df.columns:\n                            raise ValueError(f\"Lead {lead} not found in CSV for record {record_id}\")\n                        if lead == 'II':\n                            sig_II = signal_df[lead].dropna().values\n                            n = len(sig_II)\n                            quarter = n // 4\n                            sig = sig_II[:quarter]\n                        else:\n                            sig = signal_df[lead].dropna().values\n                        strip_sig = np.concatenate([strip_sig, sig], axis = 0)\n\n                elif leads == self.lead_groups[3]:\n                    for lead in leads:\n                        if lead not in signal_df.columns:\n                            raise ValueError(f\"Lead {lead} not found in CSV for record {record_id}\")\n                        strip_sig = signal_df[lead].dropna().values\n\n                else:\n                    for lead in leads:\n                        if lead not in signal_df.columns:\n                            raise ValueError(f\"Lead {lead} not found in CSV for record {record_id}\")\n                        sig = signal_df[lead].dropna().values\n                        strip_sig = np.concatenate([strip_sig, sig], axis = 0)\n\n                interpolated = self.interpolate_signal(strip_sig, self.standard_len)\n                group_tensor = torch.FloatTensor(interpolated)\n                signal_groups.append(group_tensor)\n                            \n                # for lead in leads:\n                #     if lead not in signal_df.columns:\n                #         raise ValueError(f\"Lead {lead} not found in CSV for record {record_id}\")\n                #     sig = signal_df[lead].dropna().values\n                    \n                #     interpolated = self.interpolate_signal(sig, self.standard_len)\n                #     processed_leads.append(interpolated)\n                \n                # group_tensor = torch.FloatTensor(np.stack(processed_leads, axis=0))  # (n_leads, L)\n                # signal_groups.append(group_tensor)\n\n            fs = row['fs']\n            sig_len = row['sig_len']\n            return img_tensors, signal_groups, fs, sig_len, record_id\n        else:\n            return img_tensors, record_id","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T14:03:14.745566Z","iopub.execute_input":"2026-01-01T14:03:14.745859Z","iopub.status.idle":"2026-01-01T14:03:14.761586Z","shell.execute_reply.started":"2026-01-01T14:03:14.745834Z","shell.execute_reply":"2026-01-01T14:03:14.760931Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Verify Dataset Loading\nCreating train/validation datasets with standard length of 2500 samples per lead. Testing data loading and checking shapes are correct.","metadata":{}},{"cell_type":"code","source":"# Test the updated dataset with interpolation\ntrain_subset = train_df.copy()\nval_subset = val_df.copy()\nprint(f'Training on {len(train_subset)} samples')\nprint(f'Validating on {len(val_subset)} samples')\n\n# Create datasets - both use TRAIN_DIR since validation is split from training data\nSTANDARD_LEN = 10000  # Standard length for all leads via interpolation (most efficient)\ntrain_dataset = ECGDataset(train_subset, TRAIN_DIR, is_train=True, standard_len=STANDARD_LEN)\nval_dataset = ECGDataset(val_subset, TRAIN_DIR, is_train=True, standard_len=STANDARD_LEN)\n\n# Check one sample from train\nif len(train_dataset) > 0:\n    img, signal, fs, sig_len, record_id = train_dataset[0]\n    print(f'\\nTrain Sample:')\n    print(f'Image shape: {img[2].shape}')\n    print(f'Signal shape: {signal[0].shape}')  # Should be (12, STANDARD_LEN)\n    print(f'Sampling frequency: {fs} Hz')\n    print(f'Original signal length: {sig_len}')\n    print(f'Record ID: {record_id}')\n    print(f'\\nNote: All leads interpolated to {STANDARD_LEN} for efficient training')\n    print(f'      Short leads (2500) kept as-is, long leads (10000) downsampled')\nelse:\n    print('No training samples available!')\n\n# Check one sample from val\nif len(val_dataset) > 0:\n    img, signal, fs, sig_len, record_id = val_dataset[0]\n    print(f'\\nVal Sample:')\n    print(f'Image shape: {img[3].shape}')\n    print(f'Signal shape: {signal[0].shape}')\n    print(f'Record ID: {record_id}')\nelse:\n    print('No validation samples available!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T14:15:40.06401Z","iopub.execute_input":"2026-01-01T14:15:40.064343Z","iopub.status.idle":"2026-01-01T14:15:40.454351Z","shell.execute_reply.started":"2026-01-01T14:15:40.064317Z","shell.execute_reply":"2026-01-01T14:15:40.453595Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Visualize Training Sample\nDisplaying input ECG image (256x256 resized) along with ground truth signals for all 12 leads to verify data quality.","metadata":{}},{"cell_type":"code","source":"# Visualize one training sample - 4 strips and grouped signals\nif len(train_dataset) > 0:\n    img_strips, signal_groups, fs, sig_len, record_id = train_dataset[10]\n\n    print(f'Training Sample Visualization')\n    print(f'Record ID: {record_id}')\n    print(f'Input Image: {record_id}-0001.png (first page)')\n    print(f'Sampling frequency: {fs} Hz')\n    print(f'Original signal length: {sig_len}')\n    print(f'Interpolated to: {STANDARD_LEN}')\n    print(f'Number of strips: {len(img_strips)}')\n    print(f'Number of signal groups: {len(signal_groups)}\\n')\n\n    # Lead groups for labeling\n    lead_groups = [\n        [\"I, aVR, V1, V4\"],\n        [\"II, aVL, V2, V5\"],\n        [\"III, aVF, V3, V6\"],\n        [\"II\"]\n    ]\n\n    # Create figure: each strip at top of its section, signals below\n    fig = plt.figure(figsize=(12, 16))\n    outer = fig.add_gridspec(4, 1, hspace=0.4)\n\n    for strip_idx in range(4):\n        # Sub-gridspec: one strip image + its leads\n        inner = outer[strip_idx].subgridspec(len(lead_groups[strip_idx])+1, 1,\n                                             height_ratios=[4] + [1]*len(lead_groups[strip_idx]),\n                                             hspace=0.3)\n\n        # Plot strip image\n        ax_img = fig.add_subplot(inner[0])\n        # img_display = img_strips[strip_idx].permute(1,2,0).numpy()\n        img_display = img_strips[strip_idx]\n        if isinstance(img_display, torch.Tensor):\n            if img_display.dim() == 3:  # (C,H,W)\n                img_display = img_display.permute(1,2,0).numpy()\n            elif img_display.dim() == 2:  # (H,W)\n                img_display = img_display.numpy()\n        else:\n            img_display = np.array(img_display)\n\n        ax_img.imshow(img_display.squeeze(), cmap='gray')\n        ax_img.set_title(f'Strip {strip_idx+1} - Leads {\",\".join(lead_groups[strip_idx])}',\n                         fontsize=10, fontweight='bold', pad=8)\n        ax_img.axis('off')\n\n        # Plot leads strip\n        lead_name = lead_groups[strip_idx]\n        ax = fig.add_subplot(inner[1])\n        lead_signal = signal_groups[strip_idx].numpy()\n        time_axis = np.arange(len(lead_signal)) / fs\n        ax.plot(time_axis, lead_signal, linewidth=0.8, color='blue')\n        ax.set_ylabel(lead_name, fontweight='bold', fontsize=8)\n        ax.grid(True, alpha=0.3)\n        ax.set_xlabel('Time (seconds)', fontsize=8)\n\n    plt.suptitle(f'Training Sample - Record {record_id}\\n4 Strips and Grouped Signals',\n                 fontsize=12, fontweight='bold', y=0.995)\n    plt.tight_layout()\n    plt.show()\n\n    # Print statistics per strip\n    print(f'\\nSignal statistics by strip:')\n    for strip_idx, leads in enumerate(lead_groups):\n        print(f'Strip {strip_idx+1} ({\",\".join(leads)}):')\n        for j, lead_name in enumerate(leads):\n            lead_signal = signal_groups[strip_idx][j].numpy()\n            print(f' {lead_name:>3}: min={lead_signal.min():7.2f}, max={lead_signal.max():7.2f}, '\n                  f'mean={lead_signal.mean():7.2f}, std={lead_signal.std():7.2f}')\nelse:\n    print('No training samples available for visualization!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T14:03:19.400257Z","iopub.execute_input":"2026-01-01T14:03:19.401043Z","iopub.status.idle":"2026-01-01T14:03:20.750944Z","shell.execute_reply.started":"2026-01-01T14:03:19.401012Z","shell.execute_reply":"2026-01-01T14:03:20.75009Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Previously separated signals - Old","metadata":{}},{"cell_type":"code","source":"# Visualize one training sample - all 12 leads from CSV with input image\nif len(train_dataset) > 0:\n    img, signal, fs, sig_len, record_id = train_dataset[10]\n    \n    print(f'Training Sample Visualization')\n    print(f'Record ID: {record_id}')\n    print(f'Input Image: {record_id}-0001.png (first page)')\n    print(f'Sampling frequency: {fs} Hz')\n    print(f'Original signal length: {sig_len}')\n    print(f'Interpolated to: {STANDARD_LEN}')\n    print(f'Signal shape: {signal.shape} (12 leads, {STANDARD_LEN} samples each)\\n')\n    \n    # Create a figure with input image at top and 12 lead plots below\n    leads = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', 'V1', 'V2', 'V3', 'V4', 'V5', 'V6']\n    \n    fig = plt.figure(figsize=(12, 12))\n    gs = fig.add_gridspec(13, 1, height_ratios=[10] + [1]*12, hspace=0.3)\n    \n    # Plot the input image at the top\n    ax_img = fig.add_subplot(gs[0])\n    img_display = img.permute(1, 2, 0).numpy()  # Convert from (C, H, W) to (H, W, C)\n    ax_img.imshow(img_display)\n    ax_img.set_title(f'Input ECG Image - Record {record_id}-0001.png', \n                     fontsize=10, fontweight='bold', pad=8)\n    ax_img.axis('off')\n    \n    # Plot all 12 leads\n    for i, lead_name in enumerate(leads):\n        ax = fig.add_subplot(gs[i+1])\n        lead_signal = signal[i].numpy()\n        time_axis = np.arange(len(lead_signal)) / fs  # Convert to seconds\n        \n        ax.plot(time_axis, lead_signal, linewidth=0.8, color='blue')\n        ax.set_ylabel(lead_name, fontweight='bold', fontsize=8)\n        ax.grid(True, alpha=0.3)\n        \n        if i == len(leads) - 1:\n            ax.set_xlabel('Time (seconds)', fontsize=8)\n    \n    plt.suptitle(f'Training Sample - Record {record_id}\\nInput Image (0001) and Ground Truth Signals (all 12 leads)', \n                 fontsize=12, fontweight='bold', y=0.995)\n    plt.tight_layout()\n    plt.show()\n    \n    print(f'\\nSignal statistics:')\n    for i, lead_name in enumerate(leads):\n        lead_signal = signal[i].numpy()\n        print(f'{lead_name:>3}: min={lead_signal.min():7.2f}, max={lead_signal.max():7.2f}, mean={lead_signal.mean():7.2f}, std={lead_signal.std():7.2f}')\nelse:\n    print('No training samples available for visualization!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T08:39:32.465653Z","iopub.execute_input":"2025-12-30T08:39:32.466265Z","iopub.status.idle":"2025-12-30T08:39:34.965602Z","shell.execute_reply.started":"2025-12-30T08:39:32.466236Z","shell.execute_reply":"2025-12-30T08:39:34.9646Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CNN Model Architecture - 1","metadata":{}},{"cell_type":"code","source":"# Simple CNN model\nclass SimpleCNN(nn.Module):\n    def __init__(self, standard_len=10000):\n        super(SimpleCNN, self).__init__()\n        self.standard_len = standard_len\n        \n        self.conv = nn.Sequential(\n            nn.Conv2d(3, 32, 3, padding=1),\n            nn.ReLU(),\n            nn.MaxPool2d(2),\n            \n            nn.Conv2d(32, 64, 3, padding=1),\n            nn.ReLU(),\n            nn.MaxPool2d(2),\n            \n            nn.Conv2d(64, 128, 3, padding=1),\n            nn.ReLU(),\n            nn.MaxPool2d(2),\n        )\n        \n        # After 3 maxpool layers: 256/8 = 32, 32/8 = 32 -> 32x32x128\n        self.fc = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(32*32*128, 2048),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(2048, standard_len)\n        )\n        \n    def forward(self, x):\n        x = self.conv(x)\n        x = self.fc(x)\n        # Reshape to (batch, num_leads, standard_len)\n        x = x.view(-1, self.standard_len)\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T14:07:01.797413Z","iopub.execute_input":"2026-01-01T14:07:01.798001Z","iopub.status.idle":"2026-01-01T14:07:01.804407Z","shell.execute_reply.started":"2026-01-01T14:07:01.797971Z","shell.execute_reply":"2026-01-01T14:07:01.80367Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# New CNN using Residual and temporal blocks","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\n\nclass ResidualBlock(nn.Module):\n    def __init__(self, in_ch, out_ch, stride=1):\n        super().__init__()\n        self.conv1 = nn.Conv2d(in_ch, out_ch, kernel_size=3, stride=stride, padding=1, bias=False)\n        self.bn1   = nn.BatchNorm2d(out_ch)\n        self.relu  = nn.ReLU(inplace=True)\n        self.conv2 = nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1, bias=False)\n        self.bn2   = nn.BatchNorm2d(out_ch)\n\n        self.skip = nn.Identity()\n        if in_ch != out_ch or stride != 1:\n            self.skip = nn.Sequential(\n                nn.Conv2d(in_ch, out_ch, kernel_size=1, stride=stride, bias=False),\n                nn.BatchNorm2d(out_ch)\n            )\n\n    def forward(self, x):\n        out = self.relu(self.bn1(self.conv1(x)))\n        out = self.bn2(self.conv2(out))\n        out = self.relu(out + self.skip(x))\n        return out\n\n\nclass ImageToSignalCNN(nn.Module):\n    def __init__(self, standard_len=10000, in_channels=1):\n        super().__init__()\n        self.standard_len = standard_len\n\n        # Encoder\n        self.encoder = nn.Sequential(\n            nn.Conv2d(in_channels, 32, kernel_size=7, stride=2, padding=3, bias=False),\n            nn.BatchNorm2d(32),\n            nn.ReLU(inplace=True),\n            nn.MaxPool2d(2),\n\n            ResidualBlock(32, 64, stride=2),\n            ResidualBlock(64, 128, stride=2),\n            ResidualBlock(128, 256, stride=2),\n\n            nn.AdaptiveAvgPool2d((8, 8))  # compress to fixed grid\n        )\n\n        flat_dim = 256 * 8 * 8\n\n        # MLP\n        self.mlp = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(flat_dim, 2048),\n            nn.ReLU(inplace=True),\n            nn.Dropout(0.5),\n            nn.Linear(2048, standard_len)\n        )\n\n        # 🔑 Temporal block\n        self.temporal = nn.Sequential(\n            nn.Conv1d(1, 32, kernel_size=7, padding=3),\n            nn.ReLU(inplace=True),\n            nn.Conv1d(32, 32, kernel_size=7, padding=3),\n            nn.ReLU(inplace=True),\n            nn.Conv1d(32, 1, kernel_size=7, padding=3)\n        )\n\n    def forward(self, x):\n        f = self.encoder(x)              # (B,256,8,8)\n        seq = self.mlp(f)                # (B,L)\n        seq = seq.unsqueeze(1)           # (B,1,L)\n        seq = self.temporal(seq)         # (B,1,L)\n        return seq.squeeze(1)            # (B,L)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T15:11:04.504206Z","iopub.execute_input":"2026-01-01T15:11:04.504845Z","iopub.status.idle":"2026-01-01T15:11:04.51526Z","shell.execute_reply.started":"2026-01-01T15:11:04.504817Z","shell.execute_reply":"2026-01-01T15:11:04.514618Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def loss_with_shape(pred, target, alpha=0.05):\n    # Standard MSE\n    mse = torch.mean((pred - target) ** 2)\n\n    # Derivative MSE (differences between consecutive points)\n    dp = pred[:, 1:] - pred[:, :-1]\n    dt = target[:, 1:] - target[:, :-1]\n    grad_mse = torch.mean((dp - dt) ** 2)\n\n    return mse + alpha * grad_mse","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T15:13:11.370367Z","iopub.execute_input":"2026-01-01T15:13:11.371019Z","iopub.status.idle":"2026-01-01T15:13:11.375749Z","shell.execute_reply.started":"2026-01-01T15:13:11.370991Z","shell.execute_reply":"2026-01-01T15:13:11.375039Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.optim as optim\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nSTANDARD_LEN = 10000\n\nmodel = ImageToSignalCNN(standard_len=STANDARD_LEN, in_channels=1).to(DEVICE)\noptimizer = optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-5)\n\n# Scheduler: reduce LR if val loss plateaus\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, factor=0.5, patience=3)\n\nEPOCHS = 50\n\nfor epoch in range(EPOCHS):\n    model.train()\n    train_loss, train_batches = 0.0, 0\n\n    for batch in train_loader:\n        images, signals, fs_batch, sig_len_batch, record_ids = batch\n        images = [img.to(DEVICE) for img in images]                     # list of 4 (B,1,H,W)\n        targets = torch.stack([sig.to(DEVICE) for sig in signals], 1)  # (B,4,L)\n\n        optimizer.zero_grad()\n        preds_list = [model(img) for img in images]                    # list of 4 (B,L)\n        preds = torch.stack(preds_list, 1)                             # (B,4,L)\n\n        loss = loss_with_shape(preds, targets, alpha=0.05)\n        loss.backward()\n        optimizer.step()\n\n        train_loss += loss.item()\n        train_batches += 1\n\n    avg_train = train_loss / max(train_batches, 1)\n\n    # Validation\n    model.eval()\n    val_loss, val_batches = 0.0, 0\n    with torch.no_grad():\n        for batch in val_loader:\n            images, signals, fs_batch, sig_len_batch, record_ids = batch\n            images = [img.to(DEVICE) for img in images]\n            targets = torch.stack([sig.to(DEVICE) for sig in signals], 1)\n\n            preds_list = [model(img) for img in images]\n            preds = torch.stack(preds_list, 1)\n\n            loss = loss_with_shape(preds, targets, alpha=0.05)\n            val_loss += loss.item()\n            val_batches += 1\n\n    avg_val = val_loss / max(val_batches, 1)\n    scheduler.step(avg_val)  # 🔑 adjust LR if plateau\n\n    print(f\"Epoch {epoch+1}/{EPOCHS}  Train: {avg_train:.6f}  Val: {avg_val:.6f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T15:13:29.755281Z","iopub.execute_input":"2026-01-01T15:13:29.755864Z","iopub.status.idle":"2026-01-01T16:23:42.87644Z","shell.execute_reply.started":"2026-01-01T15:13:29.755836Z","shell.execute_reply":"2026-01-01T16:23:42.875573Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Compare Ground Truth vs Prediction\nVisualizing one validation sample showing how well the model predictions match the actual ECG signals across all 12 leads.","metadata":{}},{"cell_type":"code","source":"# Validation sample comparison: Ground Truth vs Prediction (4 strips)\nif len(val_dataset) > 0:\n    # Get one validation sample\n    imgs, gt_signals, fs, sig_len, record_id = val_dataset[0]  \n    # imgs: list of 4 (1,H,W), gt_signals: list of 4 (L,)\n\n    # Get prediction for this sample\n    model.eval()\n    with torch.no_grad():\n        img_batch = [img.unsqueeze(0).to(DEVICE) for img in imgs]   # add batch dim -> (1,1,H,W) x4\n        pred_signals = [model(img).cpu().squeeze(0) for img in img_batch]  # list of 4 (L,)\n\n    print(f'Validation Sample Comparison - GT vs Prediction')\n    print(f'Record ID: {record_id}')\n    print(f'Sampling frequency: {fs} Hz')\n    print(f'Signal length: {STANDARD_LEN}')\n    print(f'GT shapes: {[sig.shape for sig in gt_signals]}')\n    print(f'Prediction shapes: {[pred.shape for pred in pred_signals]}\\n')\n\n    # Create a figure with 4 subplots (one per strip)\n    fig, axes = plt.subplots(4, 1, figsize=(18, 12))\n    fig.suptitle(f'Validation Sample Comparison - Record {record_id}\\nGround Truth (Blue) vs Prediction (Red)',\n                 fontsize=16, fontweight='bold')\n\n    for i, ax in enumerate(axes):\n        gt_lead = gt_signals[i].numpy()\n        pred_lead = pred_signals[i].numpy()\n        time_axis = np.arange(len(gt_lead)) / fs  # Convert to seconds\n\n        # Plot both signals\n        ax.plot(time_axis, gt_lead, linewidth=1.2, color='blue', alpha=0.7, label='Ground Truth')\n        ax.plot(time_axis, pred_lead, linewidth=1.0, color='red', alpha=0.6, label='Prediction')\n\n        ax.set_ylabel(f'Strip {i+1}', fontweight='bold', fontsize=12)\n        ax.grid(True, alpha=0.3)\n\n        if i == 0:\n            ax.legend(loc='upper right')\n        if i == len(axes) - 1:\n            ax.set_xlabel('Time (seconds)', fontsize=12)\n\n    plt.tight_layout()\n    plt.show()\n\n    # Calculate simple MSE for each strip\n    print(f'\\nMean Squared Error (MSE) per strip:')\n    for i in range(4):\n        gt_lead = gt_signals[i].numpy()\n        pred_lead = pred_signals[i].numpy()\n        mse = np.mean((gt_lead - pred_lead) ** 2)\n        print(f'Strip {i+1}: MSE = {mse:.6f}')\nelse:\n    print('No validation samples available for comparison!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T16:23:42.87792Z","iopub.execute_input":"2026-01-01T16:23:42.878225Z","iopub.status.idle":"2026-01-01T16:23:43.841982Z","shell.execute_reply.started":"2026-01-01T16:23:42.8782Z","shell.execute_reply":"2026-01-01T16:23:43.841373Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Validation sample comparison: Ground Truth vs Prediction (all 12 leads)\nif len(val_dataset) > 0:\n    # Get one validation sample\n    img, gt_signal, fs, sig_len, record_id = val_dataset[0]\n    \n    # Get prediction for this sample\n    model.eval()\n    with torch.no_grad():\n        img_batch = img.unsqueeze(0).to(DEVICE)  # Add batch dimension\n        pred_signal = model(img_batch).cpu().squeeze(0)  # Remove batch dimension\n    \n    print(f'Validation Sample Comparison - GT vs Prediction')\n    print(f'Record ID: {record_id}')\n    print(f'Sampling frequency: {fs} Hz')\n    print(f'Signal length: {STANDARD_LEN}')\n    print(f'GT shape: {gt_signal.shape}')\n    print(f'Prediction shape: {pred_signal.shape}\\n')\n    \n    # Create a figure with 12 subplots (one per lead)\n    leads = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF', 'V1', 'V2', 'V3', 'V4', 'V5', 'V6']\n    \n    fig, axes = plt.subplots(12, 1, figsize=(18, 22))\n    fig.suptitle(f'Validation Sample Comparison - Record {record_id}\\\\nGround Truth (Blue) vs Prediction (Red)', \n                 fontsize=16, fontweight='bold')\n    \n    for i, (ax, lead_name) in enumerate(zip(axes, leads)):\n        gt_lead = gt_signal[i].numpy()\n        pred_lead = pred_signal[i].numpy()\n        time_axis = np.arange(len(gt_lead)) / fs  # Convert to seconds\n        \n        # Plot both signals with different colors and some transparency\n        ax.plot(time_axis, gt_lead, linewidth=1.2, color='blue', alpha=0.7, label='Ground Truth')\n        ax.plot(time_axis, pred_lead, linewidth=1.0, color='red', alpha=0.6, label='Prediction')\n        \n        ax.set_ylabel(lead_name, fontweight='bold', fontsize=12)\n        ax.grid(True, alpha=0.3)\n        \n        # Add legend only to first subplot\n        if i == 0:\n            ax.legend(loc='upper right')\n        \n        if i == len(axes) - 1:\n            ax.set_xlabel('Time (seconds)', fontsize=12)\n    \n    plt.tight_layout()\n    plt.show()\n    \n    # Calculate simple MSE for each lead\n    print(f'\\\\nMean Squared Error (MSE) per lead:')\n    for i, lead_name in enumerate(leads):\n        gt_lead = gt_signal[i].numpy()\n        pred_lead = pred_signal[i].numpy()\n        mse = np.mean((gt_lead - pred_lead) ** 2)\n        print(f'{lead_name:>3}: MSE = {mse:.6f}')\nelse:\n    print('No validation samples available for comparison!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T14:58:11.083507Z","iopub.execute_input":"2026-01-01T14:58:11.083774Z","iopub.status.idle":"2026-01-01T14:58:11.291215Z","shell.execute_reply.started":"2026-01-01T14:58:11.083752Z","shell.execute_reply":"2026-01-01T14:58:11.29006Z"},"jupyter":{"source_hidden":true,"outputs_hidden":true},"collapsed":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Prepare Test Dataset\nCreating test dataset from unique record IDs (test.csv has one row per lead). Model outputs standard length which will be interpolated to required lengths.","metadata":{}},{"cell_type":"code","source":"# Create test dataset and predict\n# Get unique record IDs from test_df (test_df has one row per lead, not per record)\nunique_test_ids = test_df['id'].unique()\ntest_df_unique = pd.DataFrame({'id': unique_test_ids})\n\n# Filter to available test images\ntest_df_unique = test_df_unique[test_df_unique['id'].isin(available_test_ids)].reset_index(drop=True)\n\nprint(f'Unique test records: {len(test_df_unique)}')\ntest_dataset = ECGDataset(test_df_unique, TEST_DIR, is_train=False, standard_len=STANDARD_LEN)\ntest_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)\n\nprint(f'Generating predictions for {len(test_dataset)} test samples...')\nprint(f'Model will output {STANDARD_LEN} values per lead')\nprint(f'Then interpolate to required lengths per test.csv (flexible for any dimension)')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T14:56:09.943236Z","iopub.execute_input":"2026-01-01T14:56:09.943762Z","iopub.status.idle":"2026-01-01T14:56:09.957862Z","shell.execute_reply.started":"2026-01-01T14:56:09.943735Z","shell.execute_reply":"2026-01-01T14:56:09.957182Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Generate predictions","metadata":{}},{"cell_type":"code","source":"# Generate predictions\nmodel.eval()\npredictions = {}\n\nwith torch.no_grad():\n    for images, record_ids in tqdm(test_loader):\n        images = [img.to(DEVICE) for img in images]\n        outputs = model(images)  # Shape: (batch, 12, standard_len)\n        \n        # Store predictions for each record\n        for i, record_id in enumerate(record_ids):\n            pred_signal = outputs[i].cpu().numpy()  # Shape: (12, standard_len)\n            predictions[record_id] = pred_signal\n\nprint(f'Generated predictions for {len(predictions)} records')\nprint(f'Each prediction has shape (12, {STANDARD_LEN})')\nprint(f'Will be interpolated to required dimensions during submission generation')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T14:56:59.399058Z","iopub.execute_input":"2026-01-01T14:56:59.399602Z","iopub.status.idle":"2026-01-01T14:56:59.623561Z","shell.execute_reply.started":"2026-01-01T14:56:59.399571Z","shell.execute_reply":"2026-01-01T14:56:59.622381Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null}]}