{"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":31236,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Cell 1: Imports","metadata":{}},{"cell_type":"code","source":"# =============================================\n# TRAINING NOTEBOOK (Train_Model.ipynb)\n# Purpose: Train model, save weights (.pth) and config.\n# =============================================\n\nimport os\nimport sys\nimport warnings\n\n# SUPPRESS ALL WARNINGS\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"expandable_segments:True\"\nos.environ['OPENCV_LOG_LEVEL'] = 'FATAL'  # No OpenCV warnings\nwarnings.filterwarnings('ignore')  # No Python warnings\n\nimport cv2\nimport torch, gc\nimport numpy as np\nimport pandas as pd\nimport pickle\nimport matplotlib.pyplot as plt\nimport matplotlib.image as mpimg\nimport random\nfrom pathlib import Path\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.optim as optim\nfrom tqdm.auto import tqdm\n\n\ngc.collect()\ntorch.cuda.empty_cache()\n\nprint(\"All files imported\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-11T09:45:05.382889Z","iopub.execute_input":"2026-01-11T09:45:05.383201Z","iopub.status.idle":"2026-01-11T09:45:05.482424Z","shell.execute_reply.started":"2026-01-11T09:45:05.383177Z","shell.execute_reply":"2026-01-11T09:45:05.481802Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 2: Configuration (Training Focused)","metadata":{}},{"cell_type":"code","source":"class TrainConfig:\n    \"\"\"Configuration for training only\"\"\"\n    def __init__(self):\n        self.KAGGLE_INPUT = Path('/kaggle/input')\n        self.TRAIN_DATA_DIR = None\n        for item in self.KAGGLE_INPUT.iterdir():\n            if item.is_dir():\n                # Look for train folders\n                train_folders = ['train', 'training', 'train_images']\n                for train_name in train_folders:\n                    train_path = item / train_name\n                    if train_path.exists():\n                        self.TRAIN_DATA_DIR = train_path\n                        print(f\"Found train data at: {train_path}\")\n                        break\n                if self.TRAIN_DATA_DIR:\n                    break\n        \n        # If not found, use a default path\n        if not self.TRAIN_DATA_DIR:\n            self.TRAIN_DATA_DIR = self.KAGGLE_INPUT / 'ecg-images' / 'train'\n            print(f\"Using default path: {self.TRAIN_DATA_DIR}\")\n        \n        self.OUTPUT_DIR = Path('/kaggle/working/')\n        self.MODEL_SAVE_PATH = self.OUTPUT_DIR / 'grid_aware_ecg_weights_model.pth'\n        self.PROCESSOR_SAVE_PATH = self.OUTPUT_DIR / 'processor_config.pkl'\n        \n        # Model & Training Params (FIXED: Using TARGET_WIDTH/TARGET_HEIGHT)\n        self.TARGET_WIDTH = 512\n        self.TARGET_HEIGHT = 512\n        self.BATCH_SIZE = 8\n        self.NUM_EPOCHS = 8\n        self.LEARNING_RATE = 0.0001\n\n        \n        # Data sampling\n        self.MAX_TRAIN_FOLDERS = 20\n        self.IMAGES_PER_FOLDER = 80\n        self.TRAIN_VAL_SPLIT = 0.85\n        # Signal parameters for processor\n        self.SAMPLE_RATE = 1000\n        self.DURATION_SEC = 10.0\n        self.SIGNAL_LENGTH = 10000\n        self.MV_SCALE = 1.5\n        self.MV_OFFSET = 0.5\n        self.ECG_LEADS = ['I', 'II', 'III', 'aVR', 'aVL', 'aVF',\n                         'V1', 'V2', 'V3', 'V4', 'V5', 'V6']\n        \n\nconfig = TrainConfig()\nprint(f\"Target size: {config.TARGET_WIDTH}x{config.TARGET_HEIGHT}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-11T09:45:05.483518Z","iopub.execute_input":"2026-01-11T09:45:05.483717Z","iopub.status.idle":"2026-01-11T09:45:05.507769Z","shell.execute_reply.started":"2026-01-11T09:45:05.483699Z","shell.execute_reply":"2026-01-11T09:45:05.507071Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 3: Data Finder and Image Uploader","metadata":{}},{"cell_type":"code","source":"class DataFinder:\n    \"\"\"Finds, organizes, and visualizes ECG image data\"\"\"\n    \n    def __init__(self, config):\n        self.config = config\n        self.train_images = []\n        self.val_images = []\n        self.test_images = []\n    \n    def find_images(self, base_path, max_folders=None, images_per_folder=None):\n        \"\"\"Find images in a directory structure\"\"\"\n        if not base_path or not base_path.exists():\n            return []\n        \n        all_images = []\n        \n        # Check if base_path contains images directly\n        image_files = list(base_path.glob('*.png')) + list(base_path.glob('*.jpg')) + \\\n                     list(base_path.glob('*.jpeg')) + list(base_path.glob('*.bmp'))\n        \n        if image_files:\n            # This is a folder with images\n            all_images = image_files\n            print(f\"  Found {len(all_images)} images directly in {base_path.name}\")\n        else:\n            # Look for subfolders\n            folders = []\n            for item in base_path.iterdir():\n                if item.is_dir():\n                    img_files = list(item.glob('*.png')) + list(item.glob('*.jpg')) + \\\n                               list(item.glob('*.jpeg')) + list(item.glob('*.bmp'))\n                    if img_files:\n                        folders.append((item, len(img_files)))\n            \n            folders.sort(key=lambda x: x[1], reverse=True)\n            \n            if max_folders:\n                folders = folders[:max_folders]\n            \n            for folder, count in folders:\n                img_files = list(folder.glob('*.png')) + list(folder.glob('*.jpg')) + \\\n                           list(folder.glob('*.jpeg')) + list(folder.glob('*.bmp'))\n                \n                if images_per_folder and len(img_files) > images_per_folder:\n                    # Sample random images\n                    import random\n                    random.seed(42)\n                    img_files = random.sample(img_files, images_per_folder)\n                \n                all_images.extend(img_files)\n                print(f\"  {folder.name}: {len(img_files)} images\")\n        \n        return all_images\n    \n    def display_sample_images(self, images, num_samples=2):\n        \"\"\"Display and save sample images from the dataset directly to /kaggle/working/\"\"\"\n        if not images:\n            print(\"No images to display!\")\n            return\n        \n        # Select random samples\n        random.seed(42)\n        sample_images = random.sample(images, min(num_samples, len(images)))\n        \n        # Create figure\n        fig, axes = plt.subplots(1, len(sample_images), figsize=(15, 5))\n        if len(sample_images) == 1:\n            axes = [axes]\n        \n        for idx, img_path in enumerate(sample_images):\n            # Read and display image\n            img = mpimg.imread(str(img_path))\n            \n            axes[idx].imshow(img, cmap='gray' if len(img.shape) == 2 else None)\n            axes[idx].set_title(f\"Sample {idx+1}\\n{img_path.parent.name}/{img_path.name}\")\n            axes[idx].axis('off')\n            \n            # Print image info\n            print(f\"Sample {idx+1}:\")\n            print(f\"  Path: {img_path}\")\n            print(f\"  Shape: {img.shape}\")\n            print(f\"  Folder: {img_path.parent.name}\")\n            print(f\"  Type: {img.dtype}\")\n            print()\n        \n        plt.tight_layout()\n        plt.show()\n        \n        # Save directly to /kaggle/working/\n        for idx, img_path in enumerate(sample_images):\n            # Save a copy of the image\n            img = mpimg.imread(str(img_path))\n            sample_save_path = Path(\"/kaggle/working\") / f\"sample_{idx+1}_{img_path.name}\"\n            mpimg.imsave(sample_save_path, img)\n            print(f\"✓ Saved sample {idx+1} to: {sample_save_path}\")\n        \n        # Also save the figure directly to /kaggle/working/\n        fig_save_path = Path(\"/kaggle/working\") / \"sample_images_grid.png\"\n        fig.savefig(fig_save_path, dpi=100, bbox_inches='tight')\n        print(f\"✓ Saved grid image to: {fig_save_path}\")\n        \n        # Close figure to free memory\n        plt.close(fig)\n    \n    def load_all_data(self, display_samples=False):\n        \"\"\"Load train, validation, and test data\"\"\"\n        print(\"\\nLoading training data...\")\n        if config.TRAIN_DATA_DIR:\n            all_train = self.find_images(\n                config.TRAIN_DATA_DIR, \n                max_folders=config.MAX_TRAIN_FOLDERS,\n                images_per_folder=config.IMAGES_PER_FOLDER\n            )\n            \n            if all_train:\n                # Split into train/val\n                split_idx = int(len(all_train) * config.TRAIN_VAL_SPLIT)\n                self.train_images = all_train[:split_idx]\n                self.val_images = all_train[split_idx:]\n                print(f\"✓ Train: {len(self.train_images)} images\")\n                print(f\"✓ Val: {len(self.val_images)} images\")\n                \n                # Display sample images ONLY if requested\n                if display_samples:\n                    print(\"\\nDisplaying sample training images...\")\n                    self.display_sample_images(self.train_images, num_samples=2)\n            else:\n                print(\"✗ No training images found!\")\n        else:\n            print(\"✗ No training directory found!\")\n        \n        return self.train_images, self.val_images","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-11T09:45:05.508579Z","iopub.execute_input":"2026-01-11T09:45:05.508772Z","iopub.status.idle":"2026-01-11T09:45:05.525820Z","shell.execute_reply.started":"2026-01-11T09:45:05.508754Z","shell.execute_reply":"2026-01-11T09:45:05.525117Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 4: Grid Aware Processing ","metadata":{}},{"cell_type":"code","source":"class GridAwareProcessor:\n    \"\"\"Grid-aware preprocessing - ADAPTED VERSION\"\"\"\n    \n    def __init__(self, config):\n        self.config = config\n        # Set OpenCV to silent mode\n        import os\n        os.environ['OPENCV_LOG_LEVEL'] = 'FATAL'\n    \n    def detect_grid(self, image):\n        \"\"\"Detect grid lines\"\"\"\n        edges = cv2.Canny(image, 50, 150)\n        \n        h_kernel = cv2.getStructuringElement(cv2.MORPH_RECT, (40, 1))\n        h_lines = cv2.morphologyEx(edges, cv2.MORPH_OPEN, h_kernel)\n        \n        v_kernel = cv2.getStructuringElement(cv2.MORPH_RECT, (1, 40))\n        v_lines = cv2.morphologyEx(edges, cv2.MORPH_OPEN, v_kernel)\n        \n        grid_mask = cv2.add(h_lines, v_lines)\n        grid_mask = cv2.dilate(grid_mask, np.ones((3, 3), np.uint8), iterations=1)\n        \n        return grid_mask\n    \n    def remove_grid_fft(self, image):\n        \"\"\"Remove grid using FFT\"\"\"\n        f = np.fft.fft2(image)\n        fshift = np.fft.fftshift(f)\n        \n        rows, cols = image.shape\n        crow, ccol = rows // 2, cols // 2\n        \n        mask = np.ones((rows, cols), np.uint8)\n        d = 30\n        mask[crow-d:crow+d, :] = 0\n        mask[:, ccol-d:ccol+d] = 0\n        \n        fshift = fshift * mask\n        f_ishift = np.fft.ifftshift(fshift)\n        img_back = np.fft.ifft2(f_ishift)\n        img_back = np.abs(img_back)\n        img_back = cv2.normalize(img_back, None, 0, 255, cv2.NORM_MINMAX)\n        \n        return img_back.astype(np.uint8)\n    \n    def preprocess(self, image_path):\n        \"\"\"Preprocess with grid awareness - FIXED\"\"\"\n        try:\n            # Try to read as image file\n            img = cv2.imread(str(image_path))\n            \n            # If image_path is a string (not a real file), create synthetic data\n            if img is None and isinstance(image_path, (str, Path)) and 'dummy' in str(image_path):\n                # Create synthetic ECG image for training\n                img = self._create_synthetic_ecg_image()\n            elif img is None:\n                # Try grayscale\n                img = cv2.imread(str(image_path), cv2.IMREAD_GRAYSCALE)\n                if img is not None:\n                    img = cv2.cvtColor(img, cv2.COLOR_GRAY2BGR)\n            \n            if img is None:\n                # Create fallback image\n                print(f\"Creating fallback for: {image_path}\")\n                img = self._create_synthetic_ecg_image()\n            \n            # Process the image\n            gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n            gray = cv2.resize(gray, (self.config.TARGET_WIDTH, self.config.TARGET_HEIGHT))\n            \n            # Grid processing\n            grid_mask = self.detect_grid(gray)\n            gray_no_grid = self.remove_grid_fft(gray)\n            \n            if grid_mask.max() > 0:\n                gray_no_grid = cv2.inpaint(gray_no_grid, grid_mask, 3, cv2.INPAINT_TELEA)\n            \n            # Enhancement\n            clahe = cv2.createCLAHE(clipLimit=3.0, tileGridSize=(8, 8))\n            enhanced = clahe.apply(gray_no_grid)\n            blurred = cv2.GaussianBlur(enhanced, (5, 5), 0)\n            \n            normalized = blurred.astype(np.float32) / 255.0\n            grid_mask_norm = (grid_mask > 0).astype(np.float32)\n            \n            return normalized, grid_mask_norm\n            \n        except Exception as e:\n            print(f\"Error in preprocess for {image_path}: {e}\")\n            # Return default\n            normalized = np.ones((self.config.TARGET_HEIGHT, self.config.TARGET_WIDTH), dtype=np.float32) * 0.5\n            grid_mask_norm = np.zeros((self.config.TARGET_HEIGHT, self.config.TARGET_WIDTH), dtype=np.float32)\n            return normalized, grid_mask_norm\n    \n    def _create_synthetic_ecg_image(self):\n        \"\"\"Create a synthetic ECG image for training\"\"\"\n        height, width = 512, 512\n        img = np.ones((height, width, 3), dtype=np.uint8) * 255  # White background\n        \n        # Add grid lines (simulated ECG paper)\n        grid_spacing = 20\n        line_color = (220, 220, 220)  # Light gray\n        \n        # Horizontal lines\n        for y in range(0, height, grid_spacing):\n            cv2.line(img, (0, y), (width, y), line_color, 1)\n        \n        # Vertical lines (thicker every 5 lines)\n        for x in range(0, width, grid_spacing):\n            thickness = 1 if (x // grid_spacing) % 5 != 0 else 2\n            line_color = (200, 200, 200) if thickness == 2 else line_color\n            cv2.line(img, (x, 0), (x, height), line_color, thickness)\n        \n        # Add synthetic ECG signal\n        x_points = np.linspace(50, width-50, 1000)\n        y_points = height//2 + 100 * np.sin(2*np.pi*x_points/200) * np.exp(-0.0001*(x_points-256)**2)\n        \n        points = np.array([(int(x), int(y)) for x, y in zip(x_points, y_points)], np.int32)\n        cv2.polylines(img, [points], isClosed=False, color=(0, 0, 0), thickness=2)\n        \n        return img\n    \n    def create_labels(self, image, grid_mask):\n        \"\"\"Create signal labels from processed image\"\"\"\n        img_8bit = (image * 255).astype(np.uint8)\n        \n        # Threshold to find signal\n        binary = cv2.adaptiveThreshold(\n            img_8bit, 255,\n            cv2.ADAPTIVE_THRESH_GAUSSIAN_C,\n            cv2.THRESH_BINARY_INV, 11, 2\n        )\n        \n        # Remove grid from binary\n        grid_8bit = (grid_mask * 255).astype(np.uint8)\n        binary = cv2.subtract(binary, grid_8bit)\n        \n        # Clean up\n        kernel = np.ones((3, 3), np.uint8)\n        cleaned = cv2.morphologyEx(binary, cv2.MORPH_CLOSE, kernel)\n        cleaned = cv2.morphologyEx(cleaned, cv2.MORPH_OPEN, kernel)\n        \n        return cleaned.astype(np.float32) / 255.0\n    \n    def visualize_processing_pipeline(self, image_path):\n        \"\"\"Visualize all processing steps for a single image\"\"\"\n        if train_images and len(train_images) > 0:\n            # Use the first training image\n            sample_path = train_images[0] if isinstance(train_images[0], (str, Path)) else train_images[0]\n        else:\n            # Use synthetic image\n            sample_path = \"dummy_synthetic\"\n        \n        print(f\"\\n{'='*60}\")\n        print(f\"VISUALIZING PROCESSING PIPELINE FOR: {Path(sample_path).name}\")\n        print(f\"{'='*60}\")\n        \n        try:\n            # Read and process image through all steps\n            img = cv2.imread(str(sample_path)) if Path(sample_path).exists() else self._create_synthetic_ecg_image()\n            original = img.copy()\n            \n            # Step 1: Convert to grayscale\n            gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n            gray_resized = cv2.resize(gray, (self.config.TARGET_WIDTH, self.config.TARGET_HEIGHT))\n            \n            # Step 2: Detect grid\n            grid_mask = self.detect_grid(gray_resized)\n            \n            # Step 3: Remove grid using FFT\n            gray_no_grid = self.remove_grid_fft(gray_resized)\n            \n            # Step 4: Inpainting for remaining grid\n            if grid_mask.max() > 0:\n                gray_no_grid = cv2.inpaint(gray_no_grid, grid_mask, 3, cv2.INPAINT_TELEA)\n            \n            # Step 5: Enhancement\n            clahe = cv2.createCLAHE(clipLimit=3.0, tileGridSize=(8, 8))\n            enhanced = clahe.apply(gray_no_grid)\n            \n            # Step 6: Gaussian blur\n            blurred = cv2.GaussianBlur(enhanced, (5, 5), 0)\n            \n            # Step 7: Normalization\n            normalized = blurred.astype(np.float32) / 255.0\n            grid_mask_norm = (grid_mask > 0).astype(np.float32)\n            \n            # Step 8: Create labels\n            labels = self.create_labels(normalized, grid_mask_norm)\n            \n            # Create visualization figure\n            fig, axes = plt.subplots(3, 4, figsize=(16, 12))\n            axes = axes.flatten()\n            \n            # Plot each step\n            steps = [\n                (\"Original Image\", original, 'BGR'),\n                (\"Grayscale\", gray_resized, 'gray'),\n                (\"Grid Detection\", grid_mask, 'gray'),\n                (\"After FFT\", gray_no_grid, 'gray'),\n                (\"CLAHE Enhanced\", enhanced, 'gray'),\n                (\"Gaussian Blur\", blurred, 'gray'),\n                (\"Normalized\", normalized, 'gray'),\n                (\"Grid Mask\", grid_mask_norm, 'gray'),\n                (\"Binary Labels\", (labels * 255).astype(np.uint8), 'gray'),\n                (\"Final Signal\", labels, 'gray'),\n            ]\n            \n            for idx, (title, img_data, cmap) in enumerate(steps):\n                if idx < len(axes):\n                    if cmap == 'BGR':\n                        axes[idx].imshow(cv2.cvtColor(img_data, cv2.COLOR_BGR2RGB))\n                    else:\n                        axes[idx].imshow(img_data, cmap=cmap)\n                    axes[idx].set_title(title, fontsize=10, fontweight='bold')\n                    axes[idx].axis('off')\n                    \n                    # Add shape info\n                    if len(img_data.shape) == 2:\n                        axes[idx].text(0.5, -0.1, f\"{img_data.shape}\", \n                                      transform=axes[idx].transAxes,\n                                      ha='center', fontsize=8)\n            \n            # Hide empty subplots\n            for idx in range(len(steps), len(axes)):\n                axes[idx].axis('off')\n            \n            plt.suptitle(\"Grid-Aware ECG Processing Pipeline\", fontsize=14, fontweight='bold', y=0.98)\n            plt.tight_layout()\n            plt.show()\n            \n            # Save the visualization\n            save_path = self.config.OUTPUT_DIR / \"processing_pipeline.png\"\n            fig.savefig(save_path, dpi=150, bbox_inches='tight')\n            print(f\"\\n✓ Processing pipeline visualization saved to: {save_path}\")\n            \n            # Print summary statistics\n            print(f\"\\nProcessing Pipeline Summary:\")\n            print(f\"  • Input shape: {gray_resized.shape}\")\n            print(f\"  • Grid pixels detected: {grid_mask.sum()} ({grid_mask_norm.sum()*100:.1f}%)\")\n            print(f\"  • Signal pixels extracted: {(labels > 0.5).sum()} ({(labels > 0.5).mean()*100:.1f}%)\")\n            print(f\"  • Final normalized range: [{normalized.min():.3f}, {normalized.max():.3f}]\")\n            \n        except Exception as e:\n            print(f\"Error in visualization: {e}\")\n\nprocessor = GridAwareProcessor(config)\nprint(\"Grid_Aware Processor Initialized\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-11T09:45:05.527376Z","iopub.execute_input":"2026-01-11T09:45:05.527568Z","iopub.status.idle":"2026-01-11T09:45:05.553953Z","shell.execute_reply.started":"2026-01-11T09:45:05.527551Z","shell.execute_reply":"2026-01-11T09:45:05.553250Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 5: Dataset Class (Training only)","metadata":{}},{"cell_type":"code","source":"class TrainDataset(Dataset):\n    def __init__(self, image_paths, processor):\n        \n                # ---------- ASSERTIONS (EARLY FAILURE CHECKS) ----------\n        assert image_paths is not None, \"image_paths is None\"\n        assert isinstance(image_paths, (list, tuple)), (\n            f\"image_paths must be a list or tuple, got {type(image_paths)}\"\n        )\n        assert len(image_paths) > 0, \"image_paths is empty\"\n        assert processor is not None, \"processor is None\"\n        \n        # ---------- ASSIGNMENTS ----------\n        self.paths = image_paths\n        self.processor = processor\n    \n    def __len__(self):\n        return len(self.paths)\n    \n    def __getitem__(self, idx):\n        path = self.paths[idx]\n        try:\n            # Get processed image and grid mask\n            processed, grid_mask = self.processor.preprocess(path)\n            \n            # Stack as 2-channel input\n            img_tensor = torch.FloatTensor(np.stack([processed, grid_mask], axis=0))\n            \n            # Create labels\n            label = self.processor.create_labels(processed, grid_mask)\n            label_tensor = torch.FloatTensor(label).unsqueeze(0)\n            \n            return img_tensor, label_tensor\n        except Exception as e:\n            print(f\"Error in dataset for {path}: {e}\")\n            # Return dummy tensors\n            dummy_img = torch.zeros((2, config.TARGET_HEIGHT, config.TARGET_WIDTH))\n            dummy_label = torch.zeros((1, config.TARGET_HEIGHT, config.TARGET_WIDTH))\n            return dummy_img, dummy_label\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-11T09:45:05.554753Z","iopub.execute_input":"2026-01-11T09:45:05.555224Z","iopub.status.idle":"2026-01-11T09:45:05.569608Z","shell.execute_reply.started":"2026-01-11T09:45:05.555204Z","shell.execute_reply":"2026-01-11T09:45:05.569113Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 6. Grid-Aware Attention-Based U-Net Architecture ","metadata":{}},{"cell_type":"code","source":"class GridAttention(nn.Module):\n    \"\"\"Attention for grid awareness\"\"\"\n    def __init__(self, channels):\n        super().__init__()\n        self.conv1 = nn.Conv2d(channels, channels // 8, 1)\n        self.conv2 = nn.Conv2d(channels // 8, channels, 1)\n        self.sigmoid = nn.Sigmoid()\n    \n    def forward(self, x):\n        att = self.conv1(x)\n        att = F.relu(att)\n        att = self.conv2(att)\n        att = self.sigmoid(att)\n        return x * att\n\nclass DoubleConv(nn.Module):\n    \"\"\"Double convolution with attention\"\"\"\n    def __init__(self, in_ch, out_ch, use_att=False):\n        super().__init__()\n        self.conv = nn.Sequential(\n            nn.Conv2d(in_ch, out_ch, 3, padding=1, bias=False),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_ch, out_ch, 3, padding=1, bias=False),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True)\n        )\n        self.att = GridAttention(out_ch) if use_att else None\n    \n    def forward(self, x):\n        x = self.conv(x)\n        if self.att:\n            x = self.att(x)\n        return x\n\nclass GridAwareUNet(nn.Module):\n    \"\"\"Grid-Aware U-Net with attention\"\"\"\n    def __init__(self):\n        super().__init__()\n        \n        # Encoder (2-channel input: image + grid mask)\n        self.inc = DoubleConv(2, 64, use_att=True)\n        self.down1 = nn.Sequential(nn.MaxPool2d(2), DoubleConv(64, 128, use_att=True))\n        self.down2 = nn.Sequential(nn.MaxPool2d(2), DoubleConv(128, 256, use_att=True))\n        self.down3 = nn.Sequential(nn.MaxPool2d(2), DoubleConv(256, 512))\n        self.bottleneck = nn.Sequential(nn.MaxPool2d(2), DoubleConv(512, 1024))\n        \n        # Decoder\n        self.up1 = nn.ConvTranspose2d(1024, 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        \n        self.outc = nn.Conv2d(64, 1, 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.bottleneck(x4)\n        \n        x = self.up1(x5)\n        x = torch.cat([x, x4], dim=1)\n        x = self.conv1(x)\n        \n        x = self.up2(x)\n        x = torch.cat([x, x3], dim=1)\n        x = self.conv2(x)\n        \n        x = self.up3(x)\n        x = torch.cat([x, x2], dim=1)\n        x = self.conv3(x)\n        \n        x = self.up4(x)\n        x = torch.cat([x, x1], dim=1)\n        x = self.conv4(x)\n        \n        return torch.sigmoid(self.outc(x))\n        \nprint(\"Grid-Aware U-Net architecture defined!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-11T09:45:05.570398Z","iopub.execute_input":"2026-01-11T09:45:05.570647Z","iopub.status.idle":"2026-01-11T09:45:05.585390Z","shell.execute_reply.started":"2026-01-11T09:45:05.570630Z","shell.execute_reply":"2026-01-11T09:45:05.584846Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 7. Data Preparation & Training Loop","metadata":{}},{"cell_type":"code","source":"def main_training():\n\n    \"\"\"Find real image files in the input directory\"\"\"\n    \n    print(\"Initializing training pipeline...\")\n    finder = DataFinder(config)\n    train_images, val_images = finder.load_all_data()\n    \n    #Initialize components\n    processor = GridAwareProcessor(config)\n    train_dataset = TrainDataset(train_images, processor)\n    val_dataset = TrainDataset(val_images, processor)\n    \n    # Use num_workers=0 to avoid multiprocessing issues with synthetic data\n    train_loader = DataLoader(train_dataset, batch_size=config.BATCH_SIZE,\n                             shuffle=True, num_workers=0, pin_memory=False)\n    val_loader = DataLoader(val_dataset, batch_size=config.BATCH_SIZE,\n                             shuffle=False, num_workers=0, pin_memory=False)\n    \n    # Setup model, optimizer, loss\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    print(f\"Using device: {device}\")\n    \n    model = GridAwareUNet().to(device)\n    optimizer = optim.AdamW(model.parameters(), lr=config.LEARNING_RATE)\n    scheduler = optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, T_0=5)\n    criterion = nn.BCELoss()\n\n    \n    # Training loop\n    best_val_loss = float('inf')\n    history = {'train': [], 'val': []}\n    \n    print(f\"\\nStarting training for {config.NUM_EPOCHS} epochs...\")\n    \n    for epoch in range(config.NUM_EPOCHS):\n        # Training phase\n        model.train()\n        train_loss = 0.0\n        \n        train_bar = tqdm(train_loader, desc=f'Epoch {epoch+1}/{config.NUM_EPOCHS} [Train]')\n        for batch_idx, (images, labels) in enumerate(train_bar):\n            images, labels = images.to(device), labels.to(device)\n            \n            optimizer.zero_grad()\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            optimizer.step()\n            \n            train_loss += loss.item()\n            train_bar.set_postfix({'loss': f'{loss.item():.4f}'})\n            \n            # Break early for testing\n            if batch_idx >= 5 and epoch == 0:  # Just a few batches for first epoch test\n                break\n        \n        avg_train_loss = train_loss / len(train_loader)\n        history['train'].append(avg_train_loss)\n        \n        # Validation phase\n        model.eval()\n        val_loss = 0.0\n        \n        with torch.no_grad():\n            val_bar = tqdm(val_loader, desc=f'Epoch {epoch+1}/{config.NUM_EPOCHS} [Val]')\n            for batch_idx, (images, labels) in enumerate(val_bar):\n                images, labels = images.to(device), labels.to(device)\n                outputs = model(images)\n                val_loss += criterion(outputs, labels).item()\n                \n                # Break early for testing\n                if batch_idx >= 3:  # Just a few batches\n                    break\n        \n        avg_val_loss = val_loss / len(val_loader)\n        history['val'].append(avg_val_loss)\n        \n        scheduler.step()\n        \n        print(f\"Epoch {epoch+1}: Train Loss = {avg_train_loss:.4f}, Val Loss = {avg_val_loss:.4f}\")\n        \n        # Save best model\n        if avg_val_loss < best_val_loss:\n            best_val_loss = avg_val_loss\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'val_loss': best_val_loss,\n                'config': {k: v for k, v in config.__dict__.items() \n                          if not k.startswith('_') and not callable(v)}\n            }, config.MODEL_SAVE_PATH)\n            \n            # Save processor\n            with open(config.PROCESSOR_SAVE_PATH, 'wb') as f:\n                pickle.dump(processor, f)\n                \n            print(f\"  ✓ Saved best model (val_loss={best_val_loss:.4f})\")\n    \n    # Plot training history\n    plt.figure(figsize=(10, 5))\n    plt.plot(history['train'], label='Train Loss', marker='o')\n    plt.plot(history['val'], label='Val Loss', marker='s')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.title('Training History')\n    plt.legend()\n    plt.grid(True)\n    plt.savefig(config.OUTPUT_DIR / 'training_history.png', dpi=100, bbox_inches='tight')\n    \n    print(f\"\\nTraining complete! Best validation loss: {best_val_loss:.4f}\")\n    return model, processor, history\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-11T09:45:05.586292Z","iopub.execute_input":"2026-01-11T09:45:05.586520Z","iopub.status.idle":"2026-01-11T09:45:05.602372Z","shell.execute_reply.started":"2026-01-11T09:45:05.586496Z","shell.execute_reply":"2026-01-11T09:45:05.601752Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 8. Training Execution ","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# EXECUTION BLOCK \n# ============================================================\n\nif __name__ == \"__main__\":\n    print(\"=\" * 60)\n    print(\"GRID-AWARE ECG MODEL TRAINING\")\n    print(\"=\" * 60)\n    \n    # Step 1: Initialize processor\n    print(\"\\nStep 1: Initializing Processor\")\n    processor = GridAwareProcessor(config)\n    print(\"✓ Grid_Aware Processor Initialized\")\n    \n    # Step 2: Load data (ONCE) - WITHOUT auto-display\n    print(\"\\nStep 2: Loading Training Data\")\n    data_finder = DataFinder(config)\n    train_images, val_images = data_finder.load_all_data(display_samples=True)  # <-- TRUE to show samples\n    print(f\"✓ Loaded {len(train_images)} training images\")\n    \n    # Step 3: Visualize processing pipeline\n    print(\"\\nStep 3: Visualizing Processing Pipeline\")\n    if train_images and len(train_images) > 0:\n        processor.visualize_processing_pipeline(train_images[0])\n    else:\n        print(\"No training images found, using synthetic example\")\n        processor.visualize_processing_pipeline()\n    \n    # Step 4: Run main training\n    print(\"\\nStep 4: Starting Model Training\")\n    model, processor, history = main_training()\n    \n    # Step 5: List output files\n    print(\"\\n\" + \"=\"*50)\n    print(\"TRAINING NOTEBOOK OUTPUT FILES:\")\n    print(\"=\"*50)\n    \n    # Get all files in output directory\n    all_files = list(config.OUTPUT_DIR.glob(\"*\"))\n    if all_files:\n        # Group by file type\n        sample_files = [f for f in all_files if f.name.startswith('sample_')]\n        processing_files = [f for f in all_files if 'processing' in f.name.lower()]\n        model_files = [f for f in all_files if f.suffix in ['.pth', '.pkl', '.pt']]\n        other_files = [f for f in all_files if f not in sample_files + processing_files + model_files]\n        \n        if sample_files:\n            print(\"\\nSample Data Files:\")\n            for file in sorted(sample_files):\n                size_kb = file.stat().st_size / 1024\n                print(f\"  • {file.name:35} ({size_kb:.1f} KB)\")\n        \n        if processing_files:\n            print(\"\\nProcessing Pipeline Files:\")\n            for file in sorted(processing_files):\n                size_kb = file.stat().st_size / 1024\n                print(f\"  • {file.name:35} ({size_kb:.1f} KB)\")\n        \n        if model_files:\n            print(\"\\nModel Files:\")\n            for file in sorted(model_files):\n                size_kb = file.stat().st_size / 1024\n                print(f\"  • {file.name:35} ({size_kb:.1f} KB)\")\n        \n        if other_files:\n            print(\"\\nOther Output Files:\")\n            for file in sorted(other_files):\n                size_kb = file.stat().st_size / 1024\n                print(f\"  • {file.name:35} ({size_kb:.1f} KB)\")\n    else:\n        print(\"No output files found!\")\n    \n    print(f\"\\nModel saved to: {config.MODEL_SAVE_PATH}\")\n    print(f\"Processor saved to: {config.PROCESSOR_SAVE_PATH}\")\n    \n    # Final summary\n    print(f\"\\n{'='*50}\")\n    print(\"EXECUTION COMPLETE!\")\n    print('='*50)\n    print(f\"✓ Total training images: {len(train_images)}\")\n    print(f\"✓ Total validation images: {len(val_images)}\")\n    print(f\"✓ Total output files created: {len(all_files)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-11T09:45:05.603226Z","iopub.execute_input":"2026-01-11T09:45:05.603533Z","iopub.status.idle":"2026-01-11T09:58:52.995901Z","shell.execute_reply.started":"2026-01-11T09:45:05.603512Z","shell.execute_reply":"2026-01-11T09:58:52.995317Z"}},"outputs":[],"execution_count":null}]}