{"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":"## 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\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-10T15:56:05.282615Z","iopub.execute_input":"2026-01-10T15:56:05.282881Z","iopub.status.idle":"2026-01-10T15:56:09.333806Z","shell.execute_reply.started":"2026-01-10T15:56:05.282857Z","shell.execute_reply":"2026-01-10T15:56:09.333202Z"}},"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_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.MAX_TRAIN_IMAGES = 500  # Reduced for testing\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-10T15:56:09.334983Z","iopub.execute_input":"2026-01-10T15:56:09.335289Z","iopub.status.idle":"2026-01-10T15:56:09.345116Z","shell.execute_reply.started":"2026-01-10T15:56:09.335267Z","shell.execute_reply":"2026-01-10T15:56:09.344581Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 3: Data Finder","metadata":{}},{"cell_type":"code","source":"class DataFinder:\n    \"\"\"Finds and organizes 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 load_all_data(self):\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            else:\n                print(\"  No training images found!\")\n        else:\n            print(\"  No training directory found!\")\n        \n        \n        return self.train_images, self.val_images\n\n# Initialize and load data\ndata_finder = DataFinder(config)\ntrain_images, val_images = data_finder.load_all_data()\nprint(f\"  Data loading complete!\")\nprint(f\"  Total train images: {len(train_images)}\")\nprint(f\"  Total val images: {len(val_images)}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T15:56:09.345971Z","iopub.execute_input":"2026-01-10T15:56:09.346228Z","iopub.status.idle":"2026-01-10T15:56:23.328906Z","shell.execute_reply.started":"2026-01-10T15:56:09.346208Z","shell.execute_reply":"2026-01-10T15:56:23.328318Z"}},"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\nprocessor = GridAwareProcessor(config)\nprint(\"Grid_Aware Processor Initialized\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T15:56:23.329698Z","iopub.execute_input":"2026-01-10T15:56:23.330205Z","iopub.status.idle":"2026-01-10T15:56:23.347095Z","shell.execute_reply.started":"2026-01-10T15:56:23.330183Z","shell.execute_reply":"2026-01-10T15:56:23.346506Z"}},"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-10T15:56:23.348450Z","iopub.execute_input":"2026-01-10T15:56:23.348703Z","iopub.status.idle":"2026-01-10T15:56:23.364388Z","shell.execute_reply.started":"2026-01-10T15:56:23.348684Z","shell.execute_reply":"2026-01-10T15:56:23.363818Z"}},"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-10T15:56:23.365082Z","iopub.execute_input":"2026-01-10T15:56:23.365531Z","iopub.status.idle":"2026-01-10T15:56:23.383019Z","shell.execute_reply.started":"2026-01-10T15:56:23.365511Z","shell.execute_reply":"2026-01-10T15:56:23.382492Z"}},"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-10T15:56:23.383778Z","iopub.execute_input":"2026-01-10T15:56:23.383995Z","iopub.status.idle":"2026-01-10T15:56:23.400395Z","shell.execute_reply.started":"2026-01-10T15:56:23.383978Z","shell.execute_reply":"2026-01-10T15:56:23.399675Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 8. Execute Training","metadata":{}},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    print(\"=\" * 60)\n    print(\"GRID-AWARE ECG MODEL TRAINING\")\n    print(\"=\" * 60)\n    \n    model, processor, history = main_training()\n    \n    # List output files\n    print(\"\\n\" + \"=\"*50)\n    print(\"TRAINING NOTEBOOK OUTPUT FILES:\")\n    print(\"=\"*50)\n    output_files = []\n    for item in config.OUTPUT_DIR.iterdir():\n        if item.is_file():\n            size_kb = item.stat().st_size / 1024\n            output_files.append((item.name, size_kb))\n    \n    # Sort by size\n    output_files.sort(key=lambda x: x[1], reverse=True)\n    \n    for name, size in output_files:\n        print(f\"  • {name:30} ({size:.1f} KB)\")\n    \n    print(f\"\\nModel saved to: {config.MODEL_SAVE_PATH}\")\n    print(f\"Processor saved to: {config.PROCESSOR_SAVE_PATH}\")\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T15:56:23.401225Z","iopub.execute_input":"2026-01-10T15:56:23.401475Z","iopub.status.idle":"2026-01-10T16:09:00.349258Z","shell.execute_reply.started":"2026-01-10T15:56:23.401457Z","shell.execute_reply":"2026-01-10T16:09:00.348690Z"}},"outputs":[],"execution_count":null}]}