{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":39272,"databundleVersionId":4629629,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # Clean up existing packages\n# !pip uninstall -y numpy pandas scikit-learn torch torchvision albumentations opencv-python-headless\n# !pip cache purge\n\n# # Install exact versions in correct order\n# !pip install numpy==1.23.5\n# !pip install pandas==1.5.3\n# !pip install scikit-learn==1.0.2\n# !pip install torch==2.0.1 torchvision==0.15.2\n# !pip install timm==0.9.2\n# !pip install albumentations==1.3.1\n# !pip install opencv-python-headless==4.8.0.76\n\n# # Install DICOM handling packages\n# !apt-get update && apt-get install -y libgdcm-dev\n# !pip install gdcm==3.0.10\n# !pip install pydicom==2.3.1\n# !pip install pylibjpeg==1.4.0 pylibjpeg-libjpeg==1.3.0\n\n# # Test imports\n# print(\"----- Testing imports -----\")\n# import numpy as np\n# print(f\"NumPy: {np.__version__}\")\n# import pandas as pd\n# print(f\"Pandas: {pd.__version__}\")\n# import torch\n# print(f\"PyTorch: {torch.__version__}\")\n# import timm\n# print(f\"Timm: {timm.__version__}\")\n# import albumentations as A\n# print(f\"Albumentations: {A.__version__}\")\n# print(\"All imports successful!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-11T10:00:35.636958Z","iopub.execute_input":"2024-11-11T10:00:35.637694Z","iopub.status.idle":"2024-11-11T10:00:35.644341Z","shell.execute_reply.started":"2024-11-11T10:00:35.637614Z","shell.execute_reply":"2024-11-11T10:00:35.643405Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Basic imports\nimport os\nos.environ['NO_ALBUMENTATIONS_UPDATE'] = '1'\n\n# Core data processing\nimport numpy as np\nimport pandas as pd\n\n# Machine learning\nimport sklearn\nfrom sklearn.model_selection import GroupKFold, StratifiedGroupKFold\nfrom sklearn.metrics import roc_auc_score\nfrom torch.utils.data.sampler import WeightedRandomSampler\nfrom sklearn.utils.class_weight import compute_class_weight\n\n\n# Computer vision\nimport cv2\nimport pydicom\nfrom pydicom.pixel_data_handlers import gdcm_handler, pillow_handler\npydicom.config.image_handlers = [gdcm_handler, pillow_handler]\n\n# Deep learning\nimport torch\nimport torch.nn.functional as F\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport timm\n\n# Image augmentation\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# Monitoring\nimport matplotlib.pyplot as plt\nimport traceback\nfrom datetime import datetime\n\n# Logging\nimport sys\nimport logging\n\n# Utilities\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\nimport gc\nimport warnings\nwarnings.filterwarnings('ignore')\n\nprint(\"All imports successful!\")\n\nclass CFG:\n    \"\"\"Configuration class containing all parameters\"\"\"\n    # Debug mode\n    debug = True  # Set to False for full training\n    \n    # Paths\n    base_path = Path(\"/kaggle/input/rsna-breast-cancer-detection\")\n    processed_dir = Path(\"/kaggle/working/processed_images\")\n    model_dir = Path(\"/kaggle/working/models\")\n    \n    # Preprocessing\n    image_size = 512  # Single value for square images\n    target_size = (512, 512)  # Reduced from 2048 for memory efficiency\n    output_format = 'png'\n    \n    # Training parameters\n    seed = 42\n    epochs = 2 if debug else 10\n    train_batch_size = 8\n    valid_batch_size = 32\n    num_workers = 0\n    num_folds = 5\n    patience = 3  # Added for early stopping\n    # Add safety flags\n    persistent_workers = False\n    pin_memory = True\n    \n    # Model\n    model_name = 'efficientnet_b3'\n    pretrained = True\n    num_classes = 1 \n    \n    # Optimizer\n    optimizer = 'AdamW'\n    learning_rate = 1e-4\n    weight_decay = 1e-6\n    \n    # Scheduler\n    scheduler = 'CosineAnnealingLR'\n    min_lr = 1e-7\n    T_max = int(epochs * 0.7)\n    \n    # Class balancing parameters\n    focal_loss_alpha = 0.25\n    focal_loss_gamma = 2.0\n    use_class_weights = True\n    oversample_minority = True\n    \n    # Augmentations\n    train_aug_list = [\n        A.RandomRotate90(p=0.5),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.ShiftScaleRotate(p=0.5),\n        A.OneOf([\n            A.GaussNoise(var_limit=[10, 50]),\n            A.GaussianBlur(),\n            A.MotionBlur(),\n        ], p=0.3),\n        A.GridDistortion(p=0.3),\n        A.CoarseDropout(max_holes=8, max_width=20, max_height=20, p=0.3),\n        A.Normalize(),\n        ToTensorV2(),\n    ]\n    \n    valid_aug_list = [\n        A.Normalize(),\n        ToTensorV2(),\n    ]\n\nclass BalancedRSNADataset(Dataset):  # Changed to inherit from Dataset\n    def __init__(self, df, transform=None, is_train=True):\n        self.df = df\n        self.transform = transform\n        self.is_train = is_train\n        self.image_cache = {}  # Add image caching\n        \n        # Calculate class weights if training\n        if is_train:\n            self.class_weights = compute_class_weight(\n                class_weight='balanced',\n                classes=np.unique(df['cancer']),\n                y=df['cancer']\n            )\n            self.class_weights = torch.FloatTensor(self.class_weights)\n            \n            # Calculate sample weights for WeightedRandomSampler\n            self.sample_weights = [\n                self.class_weights[int(label)] for label in df['cancer']\n            ]\n    \n    def _load_image(self, img_path):\n        \"\"\"Load image with caching and error handling\"\"\"\n        if img_path in self.image_cache:\n            return self.image_cache[img_path].copy()\n            \n        try:\n            img = cv2.imread(str(img_path))\n            if img is not None:\n                img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n                # Cache image if not in training mode (to save memory)\n                if not self.is_train:\n                    self.image_cache[img_path] = img.copy()\n                return img\n        except Exception as e:\n            print(f\"Error loading image {img_path}: {str(e)}\")\n        \n        return None\n    \n    def _get_blank_image(self):\n        \"\"\"Create blank image with proper dimensions\"\"\"\n        return np.zeros((CFG.target_size[0], CFG.target_size[1], 3), dtype=np.uint8)\n             \n    def get_sampler(self):\n        \"\"\"Returns WeightedRandomSampler for balanced batches\"\"\"\n        if self.is_train:\n            return WeightedRandomSampler(\n                self.sample_weights,\n                len(self.sample_weights),\n                replacement=True\n            )\n        return None\n    \n    def __getitem__(self, idx):\n        # Same as your original RSNADataset __getitem__\n        row = self.df.iloc[idx]\n        img_path = CFG.processed_dir / row['view'] / row['laterality'] / f\"{row['patient_id']}_{row['image_id']}.png\"\n        \n        img = self._load_image(img_path)\n        if img is None:\n            img = self._get_blank_image()\n            print(f\"Warning: Using blank image for: {img_path}\")\n        \n        if self.transform:\n            try:\n                transformed = self.transform(image=img)\n                img = transformed['image']\n            except Exception as e:\n                print(f\"Error in transformation: {str(e)}\")\n                img = self.transform(image=self._get_blank_image())['image']\n        \n        if self.is_train:\n            label = torch.tensor(row['cancer'], dtype=torch.float32)\n            return img, label\n        else:\n            return img\n    \n    def __len__(self):\n        return len(self.df)\n\n    def get_class_distribution(self):\n        \"\"\"Calculate current class distribution\"\"\"\n        return self.df['cancer'].value_counts(normalize=True).to_dict()\n\n    def get_class_ratio(self):\n        \"\"\"Calculate ratio between classes\"\"\"\n        dist = self.get_class_distribution()\n        return dist[0] / dist[1] if 1 in dist else float('inf')\n\n    def get_all_images(self):\n        \"\"\"Get all images as a tensor\"\"\"\n        images = []\n        for idx in range(len(self)):\n            if self.is_train:\n                img, _ = self[idx]\n            else:\n                img = self[idx]\n            images.append(img)\n        return torch.stack(images)\n    \n    def get_all_labels(self):\n        \"\"\"Get all labels as a tensor\"\"\"\n        if not self.is_train:\n            return None\n        return torch.tensor(self.df['cancer'].values, dtype=torch.float32)\n\nclass FocalLoss(nn.Module):\n    \"\"\"Focal Loss for dealing with class imbalance\"\"\"\n    def __init__(self, alpha=0.25, gamma=2.0):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        \n    def forward(self, inputs, targets):\n        bce_loss = F.binary_cross_entropy_with_logits(\n            inputs, targets, reduction='none'\n        )\n        pt = torch.exp(-bce_loss)\n        focal_loss = self.alpha * (1-pt)**self.gamma * bce_loss\n        return focal_loss.mean()\n      \nclass RSNAPreprocessor:\n    \"\"\"Handles preprocessing of DICOM images\"\"\"\n    \"\"\"Enhanced preprocessing with parallel processing and better error handling\"\"\"\n    def __init__(self, **kwargs):\n        self.base_path = kwargs.get('base_path', CFG.base_path)\n        self.target_size = kwargs.get('target_size', CFG.target_size)\n        self.output_format = kwargs.get('output_format', CFG.output_format)\n        \n        # Initialize paths\n        self.train_images_path = self.base_path / \"train_images\"\n        self.test_images_path = self.base_path / \"test_images\"\n        \n        # Validate output format\n        if self.output_format not in ['png', 'jpg', 'jpeg']:\n            raise ValueError(\"output_format must be 'png' or 'jpg'/'jpeg'\")\n        \n        # Initialize error logging\n        self.error_log = []\n        \n        # Initialize CLAHE object once\n        self.clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8,8))\n        \n        # Verify paths exist\n        if not self.train_images_path.exists():\n            raise ValueError(f\"Train images path does not exist: {self.train_images_path}\")\n        if not self.test_images_path.exists():\n            raise ValueError(f\"Test images path does not exist: {self.test_images_path}\")\n            \n        # Set image handlers for pydicom\n        pydicom.config.image_handlers = [gdcm_handler, pillow_handler]\n\n    def read_dicom(self, patient_id, image_id, is_train=True):\n        \"\"\"Enhanced DICOM reading with better error handling and performance\"\"\"\n        try:\n            images_path = self.train_images_path if is_train else self.test_images_path\n            dicom_path = images_path / str(patient_id) / f\"{image_id}.dcm\"\n            \n            if not dicom_path.exists():\n                error_msg = f\"File not found: {dicom_path}\"\n                self.error_log.append({\n                    'patient_id': patient_id,\n                    'image_id': image_id,\n                    'error': error_msg\n                })\n                print(error_msg)\n                return None\n            \n            try:\n                # Use primary reading method\n                dicom = pydicom.dcmread(str(dicom_path), force=True)\n                if hasattr(dicom, 'file_meta') and hasattr(dicom.file_meta, 'TransferSyntaxUID'):\n                    if dicom.file_meta.TransferSyntaxUID.is_compressed:\n                        dicom.decompress()\n                img = dicom.pixel_array\n                \n            except Exception as e:\n                # Try alternate reading methods if primary fails\n                dicom = self._try_alternate_reading(dicom_path)\n                if dicom is None:\n                    error_msg = f\"Failed to read DICOM with all methods: {str(e)}\"\n                    self.error_log.append({\n                        'patient_id': patient_id,\n                        'image_id': image_id,\n                        'error': error_msg\n                    })\n                    print(error_msg)\n                    return None\n                img = dicom.pixel_array\n                \n            # Convert to float and normalize\n            img = img.astype(float)\n            if img.max() != img.min():\n                img = (img - img.min()) / (img.max() - img.min())\n            img = (img * 255).astype(np.uint8)\n            \n            # Apply CLAHE for better contrast - use instance variable\n            img = self.clahe.apply(img)\n            \n            # Resize with padding\n            img = self._resize_with_padding(img)\n            \n            return img\n                \n        except Exception as e:\n            error_msg = f\"Error processing image {image_id} for patient {patient_id}: {str(e)}\"\n            self.error_log.append({\n                'patient_id': patient_id,\n                'image_id': image_id,\n                'error': error_msg\n            })\n            print(error_msg)\n            return None\n\n    def _try_alternate_reading(self, dicom_path):\n        \"\"\"Try different methods to read problematic DICOM files\"\"\"\n        try:\n            # Try GDCM first\n            pydicom.config.image_handlers = [pydicom.pixel_data_handlers.gdcm_handler]\n            dicom = pydicom.dcmread(str(dicom_path), force=True)\n            return dicom\n        except:\n            try:\n                # Try PyLibJPEG\n                pydicom.config.image_handlers = [pydicom.pixel_data_handlers.pillow_handler]\n                dicom = pydicom.dcmread(str(dicom_path), force=True)\n                return dicom\n            except:\n                try:\n                    # Try without any specific handler\n                    pydicom.config.image_handlers = [None]\n                    dicom = pydicom.dcmread(str(dicom_path), force=True)\n                    return dicom\n                except:\n                    return None\n\n    def _process_dicom_image(self, dicom):\n        \"\"\"Enhanced DICOM processing with better error handling\"\"\"\n        try:\n            img = dicom.pixel_array.astype(np.float32)\n            \n            # Basic normalization\n            if img.max() != img.min():\n                img = (img - img.min()) / (img.max() - img.min())\n            img = (img * 255).astype(np.uint8)\n            \n            # Use instance CLAHE object\n            img = self.clahe.apply(img)\n            \n            return img\n                \n        except Exception as e:\n            print(f\"Error in _process_dicom_image: {str(e)}\")\n            return None\n\n    def _resize_with_padding(self, img):\n        \"\"\"Enhanced resize with better error checking\"\"\"\n        if img is None:\n            return None\n            \n        try:\n            aspect = img.shape[0] / img.shape[1]\n            if aspect > 1:\n                new_height = self.target_size[0]\n                new_width = int(new_height / aspect)\n            else:\n                new_width = self.target_size[1]\n                new_height = int(new_width * aspect)\n            \n            # Use INTER_AREA for downscaling, INTER_LINEAR for upscaling\n            if img.shape[0] > new_height or img.shape[1] > new_width:\n                interpolation = cv2.INTER_AREA\n            else:\n                interpolation = cv2.INTER_LINEAR\n                \n            img = cv2.resize(img, (new_width, new_height), interpolation=interpolation)\n            \n            # Add padding\n            top_pad = (self.target_size[0] - img.shape[0]) // 2\n            bottom_pad = self.target_size[0] - img.shape[0] - top_pad\n            left_pad = (self.target_size[1] - img.shape[1]) // 2\n            right_pad = self.target_size[1] - img.shape[1] - left_pad\n            \n            return cv2.copyMakeBorder(\n                img, top_pad, bottom_pad, left_pad, right_pad,\n                cv2.BORDER_CONSTANT, value=0\n            )\n        except Exception as e:\n            print(f\"Error in resize_with_padding: {str(e)}\")\n            return None\n\n    def save_image(self, img, output_path):\n        \"\"\"Save image with error checking\"\"\"\n        try:\n            if img is not None and img.size > 0:\n                if self.output_format == 'png':\n                    success = cv2.imwrite(str(output_path.with_suffix('.png')), img)\n                else:\n                    success = cv2.imwrite(str(output_path.with_suffix('.jpg')), img, \n                                        [cv2.IMWRITE_JPEG_QUALITY, 100])\n                return success\n            return False\n        except Exception as e:\n            print(f\"Error saving image to {output_path}: {str(e)}\")\n            return False\n     \n    def _validate_dicom(self, dicom, patient_id, image_id):\n        \"\"\"Validates DICOM file and logs metadata\"\"\"\n        try:\n            # Required DICOM attributes for validation\n            required_attributes = [\n                'PatientID', \n                'StudyInstanceUID', \n                'SeriesInstanceUID',\n                'Rows', \n                'Columns'\n            ]\n            \n            # Check for required attributes\n            missing_attributes = [\n                attr for attr in required_attributes \n                if not hasattr(dicom, attr)\n            ]\n            \n            if missing_attributes:\n                self.error_log.append({\n                    'patient_id': patient_id,\n                    'image_id': image_id,\n                    'error': f'Missing required DICOM attributes: {missing_attributes}',\n                    'type': 'validation_error'\n                })\n                return False\n                \n            # Validate image dimensions\n            if dicom.Rows == 0 or dicom.Columns == 0:\n                self.error_log.append({\n                    'patient_id': patient_id,\n                    'image_id': image_id,\n                    'error': f'Invalid image dimensions: {dicom.Rows}x{dicom.Columns}',\n                    'type': 'dimension_error'\n                })\n                return False\n                \n            # Log metadata for analysis\n            metadata = {\n                'patient_id': patient_id,\n                'image_id': image_id,\n                'rows': dicom.Rows,\n                'columns': dicom.Columns,\n                'bits_allocated': getattr(dicom, 'BitsAllocated', None),\n                'bits_stored': getattr(dicom, 'BitsStored', None),\n                'pixel_representation': getattr(dicom, 'PixelRepresentation', None),\n                'window_center': getattr(dicom, 'WindowCenter', None),\n                'window_width': getattr(dicom, 'WindowWidth', None),\n                'modality': getattr(dicom, 'Modality', None)\n            }\n            \n            # Save metadata\n            if not hasattr(self, 'metadata_log'):\n                self.metadata_log = []\n            self.metadata_log.append(metadata)\n            \n            return True\n            \n        except Exception as e:\n            self.error_log.append({\n                'patient_id': patient_id,\n                'image_id': image_id,\n                'error': str(e),\n                'type': 'validation_exception'\n            })\n            return False\n\n    def process_and_save(self, metadata_df, output_dir, num_samples=None):\n        \"\"\"Enhanced data processing and saving with detailed logging\"\"\"\n        try:\n            if num_samples:\n                metadata_df = metadata_df.head(num_samples)\n            \n            output_dir = Path(output_dir)\n            self._create_directory_structure(output_dir)\n            \n            # Initialize counters and logs\n            processed_count = 0\n            failed_count = 0\n            batch_stats = []\n            \n            # Process in smaller batches to manage memory\n            batch_size = 50\n            num_batches = (len(metadata_df) + batch_size - 1) // batch_size\n            \n            for batch_idx in range(num_batches):\n                start_idx = batch_idx * batch_size\n                end_idx = min((batch_idx + 1) * batch_size, len(metadata_df))\n                batch_df = metadata_df.iloc[start_idx:end_idx]\n                \n                batch_processed = 0\n                batch_failed = 0\n                batch_start_time = pd.Timestamp.now()\n                \n                # Process batch\n                for idx, row in tqdm(batch_df.iterrows(), \n                                total=len(batch_df),\n                                desc=f'Processing batch {batch_idx + 1}/{num_batches}'):\n                    try:\n                        # Read DICOM\n                        img = self.read_dicom(\n                            patient_id=str(row['patient_id']),\n                            image_id=str(row['image_id'])\n                        )\n                        \n                        if img is not None and img.size > 0:\n                            # Prepare output paths\n                            output_path = (output_dir / row['view'] / row['laterality'] / \n                                        f\"{row['patient_id']}_{row['image_id']}\")\n                            \n                            # Save main image\n                            success = self.save_image(img, output_path)\n                            \n                            if success:\n                                # Create and save thumbnail\n                                try:\n                                    thumbnail = cv2.resize(img, (512, 512))\n                                    thumbnail_path = output_path.with_name(f\"{output_path.stem}_thumb\")\n                                    thumb_success = self.save_image(thumbnail, thumbnail_path)\n                                    \n                                    if thumb_success:\n                                        processed_count += 1\n                                        batch_processed += 1\n                                    else:\n                                        failed_count += 1\n                                        batch_failed += 1\n                                        self.error_log.append({\n                                            'patient_id': row['patient_id'],\n                                            'image_id': row['image_id'],\n                                            'error': 'Failed to save thumbnail',\n                                            'batch': batch_idx + 1,\n                                            'timestamp': pd.Timestamp.now().strftime('%Y-%m-%d %H:%M:%S')\n                                        })\n                                except Exception as thumb_error:\n                                    failed_count += 1\n                                    batch_failed += 1\n                                    self.error_log.append({\n                                        'patient_id': row['patient_id'],\n                                        'image_id': row['image_id'],\n                                        'error': f'Thumbnail error: {str(thumb_error)}',\n                                        'batch': batch_idx + 1,\n                                        'timestamp': pd.Timestamp.now().strftime('%Y-%m-%d %H:%M:%S')\n                                    })\n                            else:\n                                failed_count += 1\n                                batch_failed += 1\n                                self.error_log.append({\n                                    'patient_id': row['patient_id'],\n                                    'image_id': row['image_id'],\n                                    'error': 'Failed to save main image',\n                                    'batch': batch_idx + 1,\n                                    'timestamp': pd.Timestamp.now().strftime('%Y-%m-%d %H:%M:%S')\n                                })\n                        else:\n                            failed_count += 1\n                            batch_failed += 1\n                            self.error_log.append({\n                                'patient_id': row['patient_id'],\n                                'image_id': row['image_id'],\n                                'error': 'Invalid image data',\n                                'batch': batch_idx + 1,\n                                'timestamp': pd.Timestamp.now().strftime('%Y-%m-%d %H:%M:%S')\n                            })\n                            \n                    except Exception as e:\n                        failed_count += 1\n                        batch_failed += 1\n                        self.error_log.append({\n                            'patient_id': row['patient_id'],\n                            'image_id': row['image_id'],\n                            'error': str(e),\n                            'batch': batch_idx + 1,\n                            'timestamp': pd.Timestamp.now().strftime('%Y-%m-%d %H:%M:%S')\n                        })\n                        print(f\"Error processing row {idx}: {str(e)}\")\n                \n                # Record batch statistics\n                batch_end_time = pd.Timestamp.now()\n                batch_duration = (batch_end_time - batch_start_time).total_seconds()\n                batch_stats.append({\n                    'batch': batch_idx + 1,\n                    'processed': batch_processed,\n                    'failed': batch_failed,\n                    'duration': batch_duration,\n                    'start_time': batch_start_time.strftime('%Y-%m-%d %H:%M:%S'),\n                    'end_time': batch_end_time.strftime('%Y-%m-%d %H:%M:%S')\n                })\n                \n                # Clear memory after each batch\n                gc.collect()\n            \n            # Save logs and statistics\n            try:\n                # Save error log with additional details\n                if self.error_log:\n                    error_df = pd.DataFrame(self.error_log)\n                    error_summary = error_df.groupby('error')['patient_id'].count()\n                    \n                    # Save detailed error log\n                    error_df.to_csv(output_dir / 'processing_errors.csv', index=False)\n                    \n                    # Save error summary\n                    with open(output_dir / 'error_summary.txt', 'w') as f:\n                        f.write(\"Error Summary:\\n\")\n                        f.write(str(error_summary))\n                    \n                    print(f\"\\nError log saved with {len(self.error_log)} entries\")\n                    print(\"\\nError Summary:\")\n                    print(error_summary)\n                \n                # Save batch statistics\n                if batch_stats:\n                    stats_df = pd.DataFrame(batch_stats)\n                    stats_df.to_csv(output_dir / 'batch_statistics.csv', index=False)\n                    \n                    # Calculate and save processing summary\n                    total_duration = sum(stat['duration'] for stat in batch_stats)\n                    avg_time_per_image = total_duration / (processed_count + failed_count)\n                    \n                    with open(output_dir / 'processing_summary.txt', 'w') as f:\n                        f.write(f\"Processing Summary:\\n\")\n                        f.write(f\"Total Images Processed: {processed_count}\\n\")\n                        f.write(f\"Total Images Failed: {failed_count}\\n\")\n                        f.write(f\"Total Processing Time: {total_duration:.2f} seconds\\n\")\n                        f.write(f\"Average Time per Image: {avg_time_per_image:.2f} seconds\\n\")\n                        f.write(f\"Success Rate: {(processed_count / (processed_count + failed_count) * 100):.2f}%\\n\")\n                \n                # Save DICOM metadata if collected\n                if hasattr(self, 'metadata_log') and self.metadata_log:\n                    metadata_df = pd.DataFrame(self.metadata_log)\n                    metadata_df.to_csv(output_dir / 'dicom_metadata.csv', index=False)\n                    print(f\"\\nDICOM metadata saved for {len(self.metadata_log)} files\")\n                    \n            except Exception as log_error:\n                print(f\"Error saving logs: {str(log_error)}\")\n            \n            return processed_count, failed_count\n            \n        except Exception as e:\n            print(f\"Fatal error in process_and_save: {str(e)}\")\n            return 0, 0\n\n    def _create_directory_structure(self, output_dir):\n        output_dir.mkdir(exist_ok=True)\n        for view in ['CC', 'MLO']:\n            (output_dir / view).mkdir(exist_ok=True)\n            (output_dir / view / 'L').mkdir(exist_ok=True)\n            (output_dir / view / 'R').mkdir(exist_ok=True)\n\nclass RSNADataset(Dataset):\n    \"\"\"Enhanced PyTorch Dataset for RSNA images\"\"\"\n    def __init__(self, df, transform=None, is_train=True):\n        self.df = df\n        self.transform = transform\n        self.is_train = is_train\n        self.image_cache = {}  # Add image caching\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def _load_image(self, img_path):\n        \"\"\"Load image with caching and error handling\"\"\"\n        if img_path in self.image_cache:\n            return self.image_cache[img_path].copy()\n            \n        try:\n            img = cv2.imread(str(img_path))\n            if img is not None:\n                img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n                # Cache image if not in training mode (to save memory)\n                if not self.is_train:\n                    self.image_cache[img_path] = img.copy()\n                return img\n        except Exception as e:\n            print(f\"Error loading image {img_path}: {str(e)}\")\n        \n        return None\n    \n    def _get_blank_image(self):\n        \"\"\"Create blank image with proper dimensions\"\"\"\n        return np.zeros((CFG.target_size[0], CFG.target_size[1], 3), dtype=np.uint8)\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = CFG.processed_dir / row['view'] / row['laterality'] / f\"{row['patient_id']}_{row['image_id']}.png\"\n        \n        # Load image or get blank if failed\n        img = self._load_image(img_path)\n        if img is None:\n            img = self._get_blank_image()\n            print(f\"Warning: Using blank image for: {img_path}\")\n        \n        # Apply augmentations\n        if self.transform:\n            try:\n                transformed = self.transform(image=img)\n                img = transformed['image']\n            except Exception as e:\n                print(f\"Error in transformation: {str(e)}\")\n                img = self.transform(image=self._get_blank_image())['image']\n        \n        if self.is_train:\n            label = torch.tensor(row['cancer'], dtype=torch.float32)\n            return img, label\n        else:\n            return img\n    \n    def clear_cache(self):\n        \"\"\"Clear the image cache\"\"\"\n        self.image_cache.clear()\n\nclass RSNAModel(nn.Module):\n    \"\"\"Enhanced model architecture with attention and feature extraction\"\"\"\n    def __init__(self, model_name=CFG.model_name, pretrained=CFG.pretrained):\n        super().__init__()\n        \n        # Initialize base model\n        self.base_model = timm.create_model(\n            model_name,\n            pretrained=pretrained,\n            num_classes=0  # Remove classifier to add custom head\n        )\n        \n        # Get number of features from base model\n        self.num_features = self.base_model.num_features\n        \n        # Add attention mechanism\n        self.attention = nn.Sequential(\n            nn.Linear(self.num_features, self.num_features // 16),\n            nn.ReLU(),\n            nn.Linear(self.num_features // 16, self.num_features),\n            nn.Sigmoid()\n        )\n        \n        # Add classifier head\n        self.classifier = nn.Sequential(\n            nn.Dropout(0.5),\n            nn.Linear(self.num_features, CFG.num_classes)\n        )\n        \n        # Initialize metrics tracking\n        self.batch_predictions = []\n        self.batch_targets = []\n        \n    def forward(self, x):\n        \"\"\"Forward pass with attention mechanism\"\"\"\n        # Get features from base model\n        features = self.base_model.forward_features(x)\n        \n        # Apply global average pooling if needed\n        if len(features.shape) > 2:\n            features = nn.functional.adaptive_avg_pool2d(features, 1).squeeze(-1).squeeze(-1)\n        \n        # Apply attention\n        attention_weights = self.attention(features)\n        features = features * attention_weights\n        \n        # Final classification\n        output = self.classifier(features)\n        return output\n    \n    def get_attention_maps(self, x):\n        \"\"\"Get attention maps for visualization\"\"\"\n        self.eval()\n        with torch.no_grad():\n            features = self.base_model.forward_features(x)\n            attention = self.attention(features.mean((-2, -1)))\n            attention = attention.view(attention.size(0), -1, 1, 1)\n            attention_maps = features * attention\n        return attention_maps\n    \n    def get_features(self, x):\n        \"\"\"Extract features for analysis\"\"\"\n        self.eval()\n        with torch.no_grad():\n            features = self.base_model.forward_features(x)\n            if len(features.shape) > 2:\n                features = nn.functional.adaptive_avg_pool2d(features, 1).squeeze(-1).squeeze(-1)\n        return features\n    \n    def count_parameters(self):\n        \"\"\"Count trainable parameters\"\"\"\n        return sum(p.numel() for p in self.parameters() if p.requires_grad)\n    \n    def get_layer_info(self):\n        \"\"\"Get information about model layers\"\"\"\n        layers_info = []\n        for name, module in self.named_modules():\n            if len(list(module.children())) == 0:  # Leaf module\n                # num_params = sum(p.numel() for p in module.parameters(if_exists=True))\n                num_params = sum(p.numel() for p in module.parameters())\n                layers_info.append({\n                    'name': name,\n                    'type': module.__class__.__name__,\n                    'parameters': num_params,\n                })\n        return pd.DataFrame(layers_info)\n    \n    @torch.no_grad()\n    def update_metrics(self, outputs, targets):\n        \"\"\"Update batch-wise metrics\"\"\"\n        predictions = torch.sigmoid(outputs).cpu().numpy()\n        targets = targets.cpu().numpy()\n        \n        self.batch_predictions.extend(predictions)\n        self.batch_targets.extend(targets)\n    \n    def get_metrics(self):\n        \"\"\"Calculate current metrics\"\"\"\n        if not self.batch_predictions:\n            return {}\n            \n        predictions = np.array(self.batch_predictions)\n        targets = np.array(self.batch_targets)\n        \n        try:\n            auc_score = roc_auc_score(targets, predictions)\n        except:\n            auc_score = float('nan')\n            \n        metrics = {\n            'auc_score': auc_score,\n            'avg_prediction': predictions.mean(),\n            'pos_ratio': (targets == 1).mean(),\n        }\n        \n        # Reset tracking\n        self.batch_predictions = []\n        self.batch_targets = []\n        \n        return metrics\n\n    def freeze_backbone(self, freeze=True):\n        \"\"\"Freeze/unfreeze backbone for transfer learning\"\"\"\n        for param in self.base_model.parameters():\n            param.requires_grad = not freeze\n            \n    def load_pretrained(self, checkpoint_path):\n        \"\"\"Safe loading of pretrained weights\"\"\"\n        try:\n            checkpoint = torch.load(checkpoint_path, map_location='cpu')\n            if 'model_state_dict' in checkpoint:\n                self.load_state_dict(checkpoint['model_state_dict'])\n            else:\n                self.load_state_dict(checkpoint)\n            return True\n        except Exception as e:\n            print(f\"Error loading pretrained weights: {str(e)}\")\n            return False\n        \nclass AverageMeter:\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n        \n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n        \n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\nclass TrainingMonitor:\n    \"\"\"Training monitoring and visualization utilities\"\"\"\n    def __init__(self, cfg):\n        self.cfg = cfg\n        \n        # Import required libraries\n        import matplotlib.pyplot as plt\n        self.plt = plt\n        \n        # Create directories for saving monitoring data\n        self.monitor_dir = cfg.model_dir / 'monitoring'\n        self.monitor_dir.mkdir(exist_ok=True)\n        \n    def visualize_training_progress(self, fold_history):\n        \"\"\"Plot training progress for a fold\"\"\"\n        epochs = range(1, len(fold_history['train_loss']) + 1)\n        \n        fig, (ax1, ax2) = self.plt.subplots(1, 2, figsize=(15, 5))\n        \n        # Plot losses\n        ax1.plot(epochs, fold_history['train_loss'], 'b-', label='Train Loss')\n        ax1.plot(epochs, fold_history['valid_loss'], 'r-', label='Valid Loss')\n        ax1.set_title('Training and Validation Loss')\n        ax1.set_xlabel('Epoch')\n        ax1.set_ylabel('Loss')\n        ax1.legend()\n        ax1.grid(True)\n        \n        # Plot validation score\n        ax2.plot(epochs, fold_history['valid_score'], 'g-', label='Valid Score')\n        ax2.set_title('Validation Score')\n        ax2.set_xlabel('Epoch')\n        ax2.set_ylabel('Score')\n        ax2.legend()\n        ax2.grid(True)\n        \n        self.plt.tight_layout()\n        return fig\n    \n    def print_model_summary(self, model):\n        \"\"\"Print detailed model summary\"\"\"\n        total_params = sum(p.numel() for p in model.parameters())\n        trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n        \n        print(f\"\\nModel Summary:\")\n        print(f\"Architecture: {self.cfg.model_name}\")\n        print(f\"Total parameters: {total_params:,}\")\n        print(f\"Trainable parameters: {trainable_params:,}\")\n        print(f\"Non-trainable parameters: {total_params - trainable_params:,}\")\n        \n        # Layer-wise summary\n        layer_info = model.get_layer_info()\n        print(\"\\nLayer-wise parameter distribution:\")\n        for _, row in layer_info.iterrows():\n            if row['parameters'] > 0:\n                print(f\"{row['name']}: {row['type']} - {row['parameters']:,} parameters\")\n    \n    def log_fold_metrics(self, fold, fold_history, best_score, best_loss):\n        \"\"\"Log fold metrics to file\"\"\"\n        metrics_path = self.monitor_dir / 'training_metrics.txt'\n        \n        with open(metrics_path, 'a') as f:\n            f.write(f\"\\nFold {fold + 1} Results:\\n\")\n            f.write(f\"Best Score: {best_score:.4f}\\n\")\n            f.write(f\"Best Loss: {best_loss:.4f}\\n\")\n            f.write(f\"Final Train Loss: {fold_history['train_loss'][-1]:.4f}\\n\")\n            f.write(f\"Final Valid Loss: {fold_history['valid_loss'][-1]:.4f}\\n\")\n            f.write(\"-\" * 50 + \"\\n\")\n     \n    def save_fold_predictions(self, fold, valid_predictions, valid_targets, fold_df, class_weights=None):\n        pred_df = pd.DataFrame({\n            'patient_id': fold_df['patient_id'],\n            'image_id': fold_df['image_id'],\n            'true_label': valid_targets,\n            'prediction': valid_predictions,\n            'fold': fold\n        })\n        if class_weights is not None:\n            pred_df['class_weight_0'] = class_weights[0]\n            pred_df['class_weight_1'] = class_weights[1]\n        \n        pred_path = self.monitor_dir / f'fold{fold}_predictions.csv'\n        pred_df.to_csv(pred_path, index=False)\n    \n    def analyze_predictions(self, predictions_path):\n        \"\"\"Analyze model predictions\"\"\"\n        pred_df = pd.read_csv(predictions_path)\n        \n        # Calculate metrics\n        metrics = {\n            'auc_score': roc_auc_score(pred_df['true_label'], pred_df['prediction']),\n            'avg_prediction': pred_df['prediction'].mean(),\n            'std_prediction': pred_df['prediction'].std(),\n            'positive_rate': (pred_df['prediction'] > 0.5).mean(),\n            'true_positive_rate': (\n                (pred_df['prediction'] > 0.5) & \n                (pred_df['true_label'] == 1)\n            ).mean()\n        }\n        \n        return metrics\n    \n    def monitor_gpu_usage(self):\n        \"\"\"Monitor GPU memory usage\"\"\"\n        if torch.cuda.is_available():\n            gpu_memory = []\n            for i in range(torch.cuda.device_count()):\n                memory_allocated = torch.cuda.memory_allocated(i) / 1024**2\n                memory_reserved = torch.cuda.memory_reserved(i) / 1024**2\n                gpu_memory.append({\n                    'device': i,\n                    'allocated_mb': memory_allocated,\n                    'reserved_mb': memory_reserved\n                })\n            return gpu_memory\n        return None\n    \n    def save_batch_metrics(self, batch_idx, metrics, fold, epoch):\n        \"\"\"Save batch-level metrics\"\"\"\n        metrics_file = self.monitor_dir / f'fold{fold}_epoch{epoch}_batch_metrics.csv'\n        \n        pd.DataFrame([{\n            'batch': batch_idx,\n            **metrics\n        }]).to_csv(metrics_file, mode='a', header=not metrics_file.exists(), index=False)\n\n\ndef train_one_epoch(model, train_loader, criterion, optimizer, scheduler, device):\n    \"\"\"Trains the model for one epoch\"\"\"\n    model.train()\n    scaler = GradScaler()\n    losses = AverageMeter()\n    \n    pbar = tqdm(train_loader, desc='Training')\n    \n    for images, labels in pbar:\n        images = images.to(device)\n        labels = labels.to(device)\n        \n        with autocast():\n            y_preds = model(images).squeeze(1)\n            loss = criterion(y_preds, labels)\n        \n        losses.update(loss.item(), labels.size(0))\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        optimizer.zero_grad()\n        \n        if scheduler is not None:\n            scheduler.step()\n            \n        pbar.set_postfix({'train_loss': losses.avg})\n    \n    return losses.avg\n\ndef valid_one_epoch(model, valid_loader, criterion, device):\n    \"\"\"Validates the model for one epoch with safe metric calculation\"\"\"\n    model.eval()\n    losses = AverageMeter()\n    preds = []\n    targets = []\n    \n    pbar = tqdm(valid_loader, desc='Validation')\n    \n    with torch.no_grad():\n        for images, labels in pbar:\n            images = images.to(device)\n            labels = labels.to(device)\n            \n            y_preds = model(images).squeeze(1)\n            loss = criterion(y_preds, labels)\n            \n            losses.update(loss.item(), labels.size(0))\n            preds.append(y_preds.sigmoid().cpu().numpy())\n            targets.append(labels.cpu().numpy())\n            \n            pbar.set_postfix({'valid_loss': losses.avg})\n    \n    preds = np.concatenate(preds)\n    targets = np.concatenate(targets)\n    \n    # Safe ROC AUC calculation\n    try:\n        if len(np.unique(targets)) > 1:\n            score = roc_auc_score(targets, preds)\n        else:\n            print(\"Warning: Only one class present in validation set. Using loss as score.\")\n            score = -losses.avg  # Use negative loss as score\n    except Exception as e:\n        print(f\"Error calculating score: {str(e)}\")\n        score = -losses.avg\n    \n    return losses.avg, score\n\ndef get_dataloader(dataset, batch_size, shuffle=True, is_train=True):\n    \"\"\"Enhanced DataLoader creation with better error handling\"\"\"\n    try:\n        return DataLoader(\n            dataset,\n            batch_size=batch_size,\n            shuffle=shuffle,\n            num_workers=0,\n            pin_memory=True,\n            drop_last=is_train,\n            persistent_workers=False,\n            timeout=60,  # Add timeout\n            prefetch_factor=2 if CFG.num_workers > 0 else None,\n        )\n    except Exception as e:\n        print(f\"Error creating DataLoader: {str(e)}\")\n        # Fallback to most basic configuration\n        return DataLoader(\n            dataset,\n            batch_size=batch_size,\n            shuffle=shuffle,\n            num_workers=0,\n            pin_memory=False\n        )\n\ndef get_predictions(model, dataset, device, batch_size=32):\n    \"\"\"Get predictions using batched inference\"\"\"\n    dataloader = DataLoader(\n        dataset, \n        batch_size=batch_size, \n        shuffle=False,\n        num_workers=0,\n        pin_memory=True\n    )\n    predictions = []\n    model.eval()\n    with torch.no_grad():\n        for images in tqdm(dataloader, desc='Getting predictions'):\n            if isinstance(images, (tuple, list)):\n                images = images[0]\n            images = images.to(device)\n            with autocast():\n                preds = model(images).sigmoid().cpu().numpy()\n            predictions.append(preds)\n    return np.concatenate(predictions)\n\ndef train_model():\n    \"\"\"Main training loop with enhanced error handling, monitoring and class balance handling\"\"\"\n    # Set seeds for reproducibility\n    torch.manual_seed(CFG.seed)\n    np.random.seed(CFG.seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    \n    # Initialize training metrics and monitor\n    fold_scores = []\n    best_scores = []\n    monitor = TrainingMonitor(CFG)\n    \n    try:\n        # Read and prepare data\n        train_df = pd.read_csv(CFG.base_path / 'train.csv')\n        if CFG.debug:\n            train_df = train_df.head(100)\n            print(\"Debug mode: Using only 100 training samples\")\n        \n        # Print initial class distribution and imbalance ratio\n        print(\"\\nOverall class distribution:\")\n        class_dist = train_df['cancer'].value_counts(normalize=True)\n        print(class_dist)\n        imbalance_ratio = class_dist[0] / class_dist[1]\n        print(f\"Imbalance ratio (negative:positive): {imbalance_ratio:.2f}:1\")\n        \n        # Create stratified folds\n        skf = StratifiedGroupKFold(\n            n_splits=CFG.num_folds, \n            shuffle=True, \n            random_state=CFG.seed\n        )\n        \n        # Create folds while maintaining patient groups\n        train_df['fold'] = -1\n        for fold, (train_idx, val_idx) in enumerate(\n            skf.split(train_df, train_df['cancer'], groups=train_df['patient_id'])\n        ):\n            train_df.loc[val_idx, 'fold'] = fold\n        \n        # Save fold assignments for reproducibility\n        train_df.to_csv(CFG.processed_dir / 'fold_assignments.csv', index=False)\n        \n        # Training loop for each fold\n        for fold in range(CFG.num_folds):\n            print(f'\\n{\"=\"*20} Fold {fold + 1}/{CFG.num_folds} {\"=\"*20}')\n            \n            # Clear memory before each fold\n            if torch.cuda.is_available():\n                torch.cuda.empty_cache()\n            gc.collect()\n            \n            try:\n                # Prepare fold data\n                train_fold = train_df[train_df.fold != fold].reset_index(drop=True)\n                valid_fold = train_df[train_df.fold == fold].reset_index(drop=True)\n                \n                # Print fold-specific class distributions\n                print(\"\\nTrain fold class distribution:\")\n                train_dist = train_fold['cancer'].value_counts(normalize=True)\n                print(train_dist)\n                print(\"\\nValid fold class distribution:\")\n                print(valid_fold['cancer'].value_counts(normalize=True))\n                \n                # Create datasets with balanced sampling\n                train_dataset = BalancedRSNADataset(\n                    train_fold, \n                    transform=A.Compose(CFG.train_aug_list),\n                    is_train=True\n                )\n                valid_dataset = BalancedRSNADataset(\n                    valid_fold, \n                    transform=A.Compose(CFG.valid_aug_list),\n                    is_train=False\n                )\n                \n                # Create dataloaders with balanced sampling for training\n                train_loader = get_balanced_dataloader(\n                    train_dataset, \n                    CFG.train_batch_size, \n                    is_train=True\n                )\n                valid_loader = get_balanced_dataloader(\n                    valid_dataset, \n                    CFG.valid_batch_size, \n                    shuffle=False, \n                    is_train=False\n                )\n                \n                # Initialize model and move to device\n                device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n                model = RSNAModel().to(device)\n                \n                # Print model summary\n                monitor.print_model_summary(model)\n                \n                # Initialize Focal Loss with class balancing\n                criterion = FocalLoss(\n                    alpha=CFG.focal_loss_alpha,\n                    gamma=CFG.focal_loss_gamma\n                ).to(device)\n                \n                # Initialize optimizer\n                optimizer = getattr(torch.optim, CFG.optimizer)(\n                    model.parameters(),\n                    lr=CFG.learning_rate,\n                    weight_decay=CFG.weight_decay\n                )\n                \n                # Initialize scheduler\n                scheduler = getattr(torch.optim.lr_scheduler, CFG.scheduler)(\n                    optimizer,\n                    T_max=CFG.T_max,\n                    eta_min=CFG.min_lr\n                )\n                \n                # Training loop\n                best_score = float('-inf')\n                best_loss = float('inf')\n                patience_counter = 0\n                fold_history = {\n                    'train_loss': [],\n                    'valid_loss': [],\n                    'valid_score': [],\n                    'class_distribution': []  # Track class distribution\n                }\n                \n                for epoch in range(CFG.epochs):\n                    print(f'\\nEpoch {epoch + 1}/{CFG.epochs}')\n                    \n                    # Monitor GPU usage and class distribution\n                    gpu_stats = monitor.monitor_gpu_usage()\n                    if gpu_stats:\n                        print(\"\\nGPU Memory Usage:\")\n                        for stat in gpu_stats:\n                            print(f\"Device {stat['device']}: \"\n                                  f\"{stat['allocated_mb']:.0f}MB allocated, \"\n                                  f\"{stat['reserved_mb']:.0f}MB reserved\")\n                    \n                    try:\n                        # Training phase with class balance monitoring\n                        train_loss = train_one_epoch(\n                            model, train_loader, criterion,\n                            optimizer, scheduler, device\n                        )\n                        \n                        # Validation phase\n                        valid_loss, valid_score = valid_one_epoch(\n                            model, valid_loader, criterion, device\n                        )\n                        \n                        # Update history with class distribution\n                        fold_history['train_loss'].append(train_loss)\n                        fold_history['valid_loss'].append(valid_loss)\n                        fold_history['valid_score'].append(valid_score)\n                        fold_history['class_distribution'].append(\n                            train_dataset.get_class_distribution()\n                        )\n                        \n                        # Save batch metrics with class balance info\n                        monitor.save_batch_metrics(\n                            epoch, \n                            {\n                                'train_loss': train_loss,\n                                'valid_loss': valid_loss,\n                                'valid_score': valid_score,\n                                'class_ratio': train_dataset.get_class_ratio()\n                            },\n                            fold,\n                            epoch\n                        )\n                        \n                        # Print metrics with class balance info\n                        print(\n                            f'Train Loss: {train_loss:.4f} '\n                            f'Valid Loss: {valid_loss:.4f} '\n                            f'Valid Score: {valid_score:.4f} '\n                            f'Class Ratio: {train_dataset.get_class_ratio():.2f}'\n                        )\n                        \n                        # Save best model and visualize\n                        if valid_score > best_score:\n                            best_score = valid_score\n                            best_loss = valid_loss\n                            \n                            # Save model with class balance info\n                            torch.save(\n                                {\n                                    'epoch': epoch,\n                                    'model_state_dict': model.state_dict(),\n                                    'optimizer_state_dict': optimizer.state_dict(),\n                                    'scheduler_state_dict': scheduler.state_dict(),\n                                    'best_score': best_score,\n                                    'best_loss': best_loss,\n                                    'fold_history': fold_history,\n                                    'class_weights': train_dataset.class_weights.cpu().numpy()\n                                },\n                                CFG.model_dir / f'fold{fold}_best.pth'\n                            )\n                            print(f'Best model saved! Score: {best_score:.4f}')\n                            \n                            # Visualization and logging\n                            fig = monitor.visualize_training_progress(fold_history)\n                            fig.savefig(monitor.monitor_dir / f'fold{fold}_training_progress.png')\n                            plt.close(fig)\n                            \n                            monitor.log_fold_metrics(fold, fold_history, best_score, best_loss)\n                            patience_counter = 0\n                        else:\n                            patience_counter += 1\n                            \n                        # Early stopping check\n                        if patience_counter >= CFG.patience:\n                            print(f'Early stopping triggered after {epoch + 1} epochs')\n                            break\n                            \n                    except Exception as e:\n                        print(f\"Error in epoch {epoch + 1}: {str(e)}\")\n                        print(traceback.format_exc())\n                        break\n                \n                # Store fold results and save predictions\n                fold_scores.append(best_score)\n                best_scores.append({\n                    'fold': fold,\n                    'score': best_score,\n                    'loss': best_loss,\n                    'final_class_ratio': train_dataset.get_class_ratio()\n                })\n                \n                # Save validation predictions with class balance metrics\n                valid_preds = get_predictions(model, valid_dataset, device, CFG.valid_batch_size)\n                valid_labels = valid_dataset.get_all_labels()\n                monitor.save_fold_predictions(\n                    fold, \n                    valid_preds, \n                    valid_labels, \n                    valid_fold, \n                    class_weights=train_dataset.class_weights.cpu().numpy()\n                )\n                \n            except Exception as e:\n                print(f\"Error in fold {fold + 1}: {str(e)}\")\n                print(traceback.format_exc())\n                continue\n                \n            finally:\n                # Cleanup\n                try:\n                    del model, train_loader, valid_loader\n                    del train_dataset, valid_dataset\n                    del optimizer, scheduler\n                    gc.collect()\n                    if torch.cuda.is_available():\n                        torch.cuda.empty_cache()\n                except Exception as e:\n                    print(f\"Error in cleanup: {str(e)}\")\n        \n        # Print final results with class balance metrics\n        print(\"\\nTraining completed!\")\n        print(\"\\nBest scores per fold:\")\n        for score_dict in best_scores:\n            print(\n                f\"Fold {score_dict['fold'] + 1}: \"\n                f\"Score = {score_dict['score']:.4f}, \"\n                f\"Loss = {score_dict['loss']:.4f}, \"\n                f\"Class Ratio = {score_dict['final_class_ratio']:.2f}\"\n            )\n        \n        print(f\"\\nMean CV score: {np.mean(fold_scores):.4f}\")\n        print(f\"Std CV score: {np.std(fold_scores):.4f}\")\n        \n        # Analyze overall predictions with class balance consideration\n        for fold in range(CFG.num_folds):\n            pred_path = monitor.monitor_dir / f'fold{fold}_predictions.csv'\n            if pred_path.exists():\n                metrics = monitor.analyze_predictions(pred_path)\n                print(f\"\\nFold {fold + 1} Prediction Analysis:\")\n                for metric_name, value in metrics.items():\n                    print(f\"{metric_name}: {value:.4f}\")\n        \n        return best_scores\n        \n    except Exception as e:\n        print(f\"Fatal error in training: {str(e)}\")\n        print(traceback.format_exc())\n        return None\n\ndef inference():\n    \"\"\"Performs inference using trained models\"\"\"\n    print(\"\\nStarting inference...\")\n    \n    try:\n        # Read test data\n        test_df = pd.read_csv(CFG.base_path / 'test.csv')\n        if CFG.debug:\n            test_df = test_df.head(100)\n            print(\"Debug mode: Using only 100 test samples\")\n        \n        device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n        predictions = []\n        \n        # Create test dataset\n        test_dataset = RSNADataset(\n            test_df,\n            transform=A.Compose(CFG.valid_aug_list),\n            is_train=False\n        )\n        test_loader = DataLoader(\n            test_dataset,\n            batch_size=CFG.valid_batch_size,\n            shuffle=False,\n            num_workers=0,  # Set to 0 to avoid multiprocessing issues\n            pin_memory=True\n        )\n        \n        # Inference with all folds\n        for fold in range(CFG.num_folds):\n            print(f'Inferencing fold {fold + 1}/{CFG.num_folds}')\n            model = RSNAModel().to(device)\n            \n            try:\n                # Load the saved model state\n                checkpoint = torch.load(\n                    CFG.model_dir / f'fold{fold}_best.pth',\n                    map_location=device\n                )\n                model.load_state_dict(checkpoint['model_state_dict'])\n                model.eval()\n                \n                fold_preds = []\n                with torch.no_grad():\n                    for images in tqdm(test_loader, desc=f'Fold {fold + 1}'):\n                        images = images.to(device)\n                        with autocast():  # Add mixed precision inference\n                            y_preds = model(images).squeeze(1)\n                        fold_preds.append(y_preds.sigmoid().cpu().numpy())\n                \n                fold_preds = np.concatenate(fold_preds)\n                predictions.append(fold_preds)\n                print(f\"Fold {fold + 1} inference completed\")\n                \n            except Exception as e:\n                print(f\"Error in fold {fold + 1} inference: {str(e)}\")\n                print(traceback.format_exc())\n                continue\n            \n            finally:\n                # Cleanup\n                try:\n                    del model\n                    gc.collect()\n                    torch.cuda.empty_cache()\n                except Exception as e:\n                    print(f\"Error in cleanup: {str(e)}\")\n        \n        if not predictions:\n            raise ValueError(\"No valid predictions from any fold\")\n            \n        # Average predictions from all folds\n        predictions = np.mean(predictions, axis=0)\n        \n        # Create submission\n        submission = pd.DataFrame({\n            'prediction_id': test_df['prediction_id'],\n            'cancer': predictions\n        })\n        \n        # Save submission with timestamp\n        timestamp = pd.Timestamp.now().strftime(\"%Y%m%d_%H%M%S\")\n        submission_path = f'submission_{timestamp}.csv'\n        submission.to_csv(submission_path, index=False)\n        print(f'Submission saved to {submission_path}!')\n        \n        return submission\n        \n    except Exception as e:\n        print(f\"Fatal error in inference: {str(e)}\")\n        print(traceback.format_exc())\n        return None\n\ndef process_test_data():\n    \"\"\"Processes test data for inference\"\"\"\n    print(\"\\nProcessing test data...\")\n    test_df = pd.read_csv(CFG.base_path / 'test.csv')\n    if CFG.debug:\n        test_df = test_df.head(100)\n    \n    preprocessor = RSNAPreprocessor(\n        base_path=CFG.base_path,\n        target_size=CFG.target_size,\n        output_format=CFG.output_format\n    )\n    \n    processed_count, failed_count = preprocessor.process_and_save(\n        test_df,\n        CFG.processed_dir,\n        num_samples=None if not CFG.debug else 100\n    )\n    print(f\"Test data processing completed. Processed: {processed_count}, Failed: {failed_count}\")\n\ndef get_balanced_dataloader(dataset, batch_size, shuffle=True, is_train=True):\n    \"\"\"Creates DataLoader with balanced sampling if needed\"\"\"\n    if is_train:\n        sampler = dataset.get_sampler()\n        return DataLoader(\n            dataset,\n            batch_size=batch_size,\n            sampler=sampler,  # Use sampler instead of shuffle\n            num_workers=CFG.num_workers,\n            pin_memory=CFG.pin_memory,\n            drop_last=is_train,\n            persistent_workers=CFG.persistent_workers\n        )\n    else:\n        return DataLoader(\n            dataset,\n            batch_size=batch_size,\n            shuffle=shuffle,\n            num_workers=CFG.num_workers,\n            pin_memory=CFG.pin_memory,\n            drop_last=is_train,\n            persistent_workers=CFG.persistent_workers\n        )\n\ndef save_run_config(cfg, run_dir):\n    \"\"\"Save training configuration and run info\"\"\"\n    timestamp = datetime.now().strftime(\"%Y%m%d_%H%M%S\")\n    config_path = run_dir / f'run_config_{timestamp}.txt'\n    \n    with open(config_path, 'w') as f:\n        f.write(f\"Run started at: {timestamp}\\n\\n\")\n        f.write(\"Training Configuration:\\n\")\n        f.write(\"-\" * 50 + \"\\n\")\n        \n        # Save all CFG attributes\n        for attr_name in dir(cfg):\n            if not attr_name.startswith('__'):\n                attr_value = getattr(cfg, attr_name)\n                if isinstance(attr_value, Path):\n                    attr_value = str(attr_value)\n                f.write(f\"{attr_name}: {attr_value}\\n\")\n        \n        # Save system info\n        f.write(\"\\nSystem Information:\\n\")\n        f.write(\"-\" * 50 + \"\\n\")\n        f.write(f\"Python version: {sys.version}\\n\")\n        f.write(f\"PyTorch version: {torch.__version__}\\n\")\n        f.write(f\"CUDA available: {torch.cuda.is_available()}\\n\")\n        if torch.cuda.is_available():\n            f.write(f\"CUDA version: {torch.version.cuda}\\n\")\n            f.write(f\"GPU: {torch.cuda.get_device_name(0)}\\n\")\n            \n        return config_path\n\ndef main():\n    \"\"\"Main execution function with enhanced logging and error handling\"\"\"\n    start_time = datetime.now()\n    print(f\"Starting RSNA Mammography Pipeline at {start_time.strftime('%Y-%m-%d %H:%M:%S')}\")\n    \n    try:\n        # Create run directory with timestamp\n        timestamp = start_time.strftime(\"%Y%m%d_%H%M%S\")\n        run_dir = CFG.model_dir / f'run_{timestamp}'\n        run_dir.mkdir(parents=True, exist_ok=True)\n        \n        # Save run configuration\n        config_path = save_run_config(CFG, run_dir)\n        print(f\"Run configuration saved to: {config_path}\")\n        \n        # Create necessary directories\n        CFG.processed_dir.mkdir(parents=True, exist_ok=True)\n        \n        # Initialize log file\n        log_path = run_dir / 'run_log.txt'\n        logging.basicConfig(\n            filename=str(log_path),\n            level=logging.INFO,\n            format='%(asctime)s - %(levelname)s - %(message)s'\n        )\n        logging.info(\"Pipeline started\")\n        \n        try:\n            # Step 1: Process training data\n            print(\"\\nStep 1: Processing training data...\")\n            logging.info(\"Starting training data processing\")\n            \n            train_df = pd.read_csv(CFG.base_path / 'train.csv')\n            if CFG.debug:\n                train_df = train_df.head(100)\n                print(\"Debug mode: Using only 100 training samples\")\n                logging.info(\"Running in debug mode with 100 samples\")\n            \n            preprocessor = RSNAPreprocessor(\n                base_path=CFG.base_path,\n                target_size=CFG.target_size,\n                output_format=CFG.output_format\n            )\n            \n            processed_count, failed_count = preprocessor.process_and_save(\n                train_df,\n                CFG.processed_dir,\n                num_samples=None if not CFG.debug else 100\n            )\n            \n            processing_msg = f\"Training data processing completed. Processed: {processed_count}, Failed: {failed_count}\"\n            print(processing_msg)\n            logging.info(processing_msg)\n            \n            # Step 2: Train models\n            print(\"\\nStep 2: Training models...\")\n            logging.info(\"Starting model training\")\n            best_scores = train_model()\n            \n            if best_scores:\n                mean_cv = np.mean([score['score'] for score in best_scores])\n                logging.info(f\"Training completed. Mean CV score: {mean_cv:.4f}\")\n            \n            # Step 3: Process test data\n            print(\"\\nStep 3: Processing test data...\")\n            logging.info(\"Starting test data processing\")\n            process_test_data()\n            \n            # Step 4: Generate predictions\n            print(\"\\nStep 4: Generating predictions...\")\n            logging.info(\"Starting inference\")\n            submission = inference()\n            \n            if submission is not None:\n                submission_stats = {\n                    'mean': submission['cancer'].mean(),\n                    'std': submission['cancer'].std(),\n                    'min': submission['cancer'].min(),\n                    'max': submission['cancer'].max()\n                }\n                logging.info(f\"Submission statistics: {submission_stats}\")\n            \n            # Calculate and log total runtime\n            end_time = datetime.now()\n            runtime = end_time - start_time\n            runtime_msg = f\"\\nPipeline completed successfully! Total runtime: {runtime}\"\n            print(runtime_msg)\n            logging.info(runtime_msg)\n            \n            if CFG.debug:\n                debug_msg = \"\\nNote: This was run in debug mode. Set CFG.debug = False for full training.\"\n                print(debug_msg)\n                logging.info(debug_msg)\n            \n            # Save final summary\n            with open(run_dir / 'run_summary.txt', 'w') as f:\n                f.write(f\"Run Summary\\n\")\n                f.write(f\"===========\\n\")\n                f.write(f\"Start time: {start_time}\\n\")\n                f.write(f\"End time: {end_time}\\n\")\n                f.write(f\"Total runtime: {runtime}\\n\")\n                f.write(f\"Processed images: {processed_count}\\n\")\n                f.write(f\"Failed images: {failed_count}\\n\")\n                if best_scores:\n                    f.write(f\"Mean CV score: {mean_cv:.4f}\\n\")\n                if submission is not None:\n                    f.write(f\"\\nSubmission Statistics:\\n\")\n                    for stat, value in submission_stats.items():\n                        f.write(f\"{stat}: {value:.4f}\\n\")\n            \n            return submission\n            \n        except Exception as e:\n            error_msg = f\"Pipeline step error: {str(e)}\"\n            print(error_msg)\n            logging.error(error_msg)\n            logging.error(traceback.format_exc())\n            return None\n            \n    except Exception as e:\n        error_msg = f\"Fatal error in pipeline initialization: {str(e)}\"\n        print(error_msg)\n        print(traceback.format_exc())\n        return None\n\nif __name__ == \"__main__\":\n    # Ensure proper multiprocessing behavior\n    torch.multiprocessing.set_start_method('spawn', force=True)\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-11T10:53:05.676749Z","iopub.execute_input":"2024-11-11T10:53:05.677076Z","iopub.status.idle":"2024-11-11T10:55:25.390034Z","shell.execute_reply.started":"2024-11-11T10:53:05.677041Z","shell.execute_reply":"2024-11-11T10:55:25.389008Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}