{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.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":59093,"databundleVersionId":7469972,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Cell 1: Import Libraries\nimport os\nimport gc\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom tqdm.auto import tqdm\nimport random\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, ConcatDataset\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom torch.nn.parallel import DataParallel\nfrom sklearn.model_selection import StratifiedKFold, GroupKFold\nfrom sklearn.metrics import roc_auc_score, f1_score, accuracy_score\nimport warnings\nimport scipy\nfrom scipy import signal\nimport math\nfrom typing import Dict, List, Tuple\nimport time\nimport glob\nimport pickle\nfrom PIL import Image\nfrom skimage.transform import resize","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-25T01:57:46.170874Z","iopub.execute_input":"2025-03-25T01:57:46.171250Z","iopub.status.idle":"2025-03-25T01:57:51.488635Z","shell.execute_reply.started":"2025-03-25T01:57:46.171218Z","shell.execute_reply":"2025-03-25T01:57:51.487603Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 2: Configuration and Setup\n# Set seeds for reproducibility\ndef seed_everything(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)  # for multi-GPU\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False  # for reproducibility\n\nseed_everything()\nwarnings.filterwarnings('ignore')\n\n# Check if CUDA is available and get GPU count\nprint(f\"CUDA available: {torch.cuda.is_available()}\")\nnum_gpus = torch.cuda.device_count()\nprint(f\"Number of GPUs available: {num_gpus}\")\nfor i in range(num_gpus):\n    print(f\"GPU {i}: {torch.cuda.get_device_name(i)}\")\n\n# Set device\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\n\n# Add spectrogram directory to configuration\nCONFIG = {\n    # Data parameters\n    'target_freq': 100,  \n    'window_size_seconds': 10,  \n    'overlap_ratio': 0.5,\n    'channels': ['Fp1', 'Fp2', 'F7', 'F8', 'F3', 'F4', 'T3', 'T4', 'C3', 'C4', 'T5', 'T6', 'P3', 'P4', 'O1', 'O2', 'Fz', 'Cz', 'Pz'],\n    \n    # Spectrogram parameters\n    'use_spectrograms': True,\n    'spec_height': 128,\n    'spec_width': 256,\n    \n    # Training parameters\n    'batch_size': 32,  # Reduced from 64 to 32 for training batches\n    'num_epochs': 10,\n    'learning_rate': 3e-4,\n    'weight_decay': 1e-5,\n    'fold_count': 3,\n    'patience': 3,\n    \n    # Model parameters\n    'model_name': 'MultimodalEEGNet',\n    'dropout_rate': 0.5,\n    'n_classes': 6,\n    \n    # Multi-GPU settings\n    'use_multi_gpu': True,\n    'num_workers': 1,  # Reduced to 1 worker\n    'pin_memory': True,\n    'mixed_precision': True,\n    \n    # Memory management\n    'batch_processing_size': 100,  # Reduced from 500 to 100\n    'max_file_time': 15,  # Reduced from 30 to 15 seconds\n}\n\n# Set up mixed precision training if available\nif CONFIG['mixed_precision'] and torch.cuda.is_available():\n    try:\n        from torch.cuda.amp import autocast, GradScaler\n        scaler = GradScaler()\n        print(\"Mixed precision training enabled\")\n    except ImportError:\n        CONFIG['mixed_precision'] = False\n        print(\"Mixed precision training not available, disabled\")\n\n# Update paths to include spectrograms\nINPUT_DIR = \"/kaggle/input/hms-harmful-brain-activity-classification\"\nTRAIN_DIR = f\"{INPUT_DIR}/train_eegs\"\nTEST_DIR = f\"{INPUT_DIR}/test_eegs\"\nTRAIN_SPECTROGRAMS_DIR = f\"{INPUT_DIR}/train_spectrograms\" \nTEST_SPECTROGRAMS_DIR = f\"{INPUT_DIR}/test_spectrograms\"\nOUTPUT_DIR = \"/kaggle/working\"\nos.makedirs(OUTPUT_DIR, exist_ok=True)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T01:57:51.489803Z","iopub.execute_input":"2025-03-25T01:57:51.490242Z","iopub.status.idle":"2025-03-25T01:57:51.627903Z","shell.execute_reply.started":"2025-03-25T01:57:51.490208Z","shell.execute_reply":"2025-03-25T01:57:51.626944Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# Cell 3: Data Loading and Exploration\n# Load metadata\ntrain_metadata = pd.read_csv(f\"{INPUT_DIR}/train.csv\")\ntest_metadata = pd.read_csv(f\"{INPUT_DIR}/test.csv\")\nsubmission_df = pd.read_csv(f\"{INPUT_DIR}/sample_submission.csv\")\n\nprint(f\"Train samples: {len(train_metadata)}\")\nprint(f\"Test samples: {len(test_metadata)}\")\n\n# Examine training data\nprint(\"Sample of training metadata:\")\nprint(train_metadata.head())\n\n# Check column names to avoid errors\nprint(f\"\\nAvailable columns: {train_metadata.columns.tolist()}\")\n\n# Determine the label column\nif 'expert_consensus' in train_metadata.columns:\n    label_column = 'expert_consensus'\nelif 'label' in train_metadata.columns:\n    label_column = 'label'\nelse:\n    print(\"Warning: Could not find expected label column. Available columns:\", train_metadata.columns.tolist())\n    # Try to identify the most likely label column\n    possible_label_columns = [col for col in train_metadata.columns if 'label' in col.lower() or 'class' in col.lower() or 'consensus' in col.lower()]\n    if possible_label_columns:\n        label_column = possible_label_columns[0]\n        print(f\"Using '{label_column}' as the label column\")\n    else:\n        raise ValueError(\"No suitable label column found in the data\")\n\n# Create label mapping\nlabels_map = {label: idx for idx, label in enumerate(sorted(train_metadata[label_column].unique()))}\nidx_to_label = {v: k for k, v in labels_map.items()}\n\nprint(f\"\\nClasses: {train_metadata[label_column].unique()}\")\nprint(f\"Class distribution: \\n{train_metadata[label_column].value_counts()}\")\nprint(f\"Label mapping: {labels_map}\")\n\n# Update n_classes in CONFIG based on actual number of classes\nCONFIG['n_classes'] = len(labels_map)\nprint(f\"Number of classes: {CONFIG['n_classes']}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T01:57:51.629735Z","iopub.execute_input":"2025-03-25T01:57:51.629982Z","iopub.status.idle":"2025-03-25T01:57:51.878787Z","shell.execute_reply.started":"2025-03-25T01:57:51.629961Z","shell.execute_reply":"2025-03-25T01:57:51.877968Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 4: Processing Functions (EEG and Spectrogram)\n# EEG data processing functions\ndef load_eeg_file(file_path):\n    \"\"\"Load EEG data from a file and return as a pandas DataFrame.\"\"\"\n    try:\n        eeg_df = pd.read_parquet(file_path)\n        return eeg_df\n    except Exception as e:\n        print(f\"Error loading EEG file {file_path}: {e}\")\n        return None\n\ndef preprocess_eeg(eeg_df, target_freq=100):\n    \"\"\"Preprocess EEG data: handle missing values, filter, and downsample.\"\"\"\n    if eeg_df is None:\n        return None\n        \n    # Determine if there's a time column, if not estimate the sampling frequency\n    if 'time' in eeg_df.columns:\n        time_diff = eeg_df['time'].diff().median()\n        orig_freq = round(1 / time_diff)\n    else:\n        # Assume a standard sampling rate if time column not available\n        orig_freq = 200  # Common EEG sampling frequency\n        print(f\"No time column found. Assuming original sampling frequency of {orig_freq} Hz.\")\n    \n    # Fill missing values (if any)\n    eeg_df = eeg_df.fillna(method='ffill').fillna(method='bfill')\n    \n    # Extract EEG signal columns, excluding non-EEG columns if present\n    available_channels = [ch for ch in CONFIG['channels'] if ch in eeg_df.columns]\n    if len(available_channels) < len(CONFIG['channels']):\n        print(f\"Warning: Only {len(available_channels)} out of {len(CONFIG['channels'])} channels found in data.\")\n    \n    eeg_signals = eeg_df[available_channels].values\n    \n    # Apply bandpass filter (0.5-40 Hz) to remove noise - simplified for speed\n    try:\n        sos = signal.butter(4, [0.5, 40], 'bandpass', fs=orig_freq, output='sos')\n        filtered_signals = np.zeros_like(eeg_signals)\n        for i in range(eeg_signals.shape[1]):\n            filtered_signals[:, i] = signal.sosfilt(sos, eeg_signals[:, i])\n    except Exception as e:\n        print(f\"Error during filtering: {e}\")\n        filtered_signals = eeg_signals  # Use original signals if filtering fails\n    \n    # Downsample if needed\n    if target_freq < orig_freq:\n        # Calculate downsampling factor\n        downsample_factor = orig_freq // target_freq\n        filtered_signals = filtered_signals[::downsample_factor, :]\n    \n    return filtered_signals\n\ndef segment_eeg(eeg_data, window_size, overlap_ratio):\n    \"\"\"Segment EEG data into windows with overlap.\"\"\"\n    if eeg_data is None or eeg_data.shape[0] < window_size:\n        return np.array([])\n        \n    step_size = int(window_size * (1 - overlap_ratio))\n    segments = []\n    \n    # Calculate how many complete segments we can extract\n    num_segments = (eeg_data.shape[0] - window_size) // step_size + 1\n    \n    # Limit to a maximum of 5 segments per recording to save memory\n    max_segments = 5\n    num_segments = min(num_segments, max_segments)\n    \n    for i in range(num_segments):\n        start_idx = i * step_size\n        end_idx = start_idx + window_size\n        segment = eeg_data[start_idx:end_idx, :]\n        segments.append(segment)\n    \n    return np.array(segments)\n\n# Spectrogram processing functions\ndef load_spectrogram(file_path):\n    \"\"\"Load spectrogram data.\"\"\"\n    if file_path is None:\n        return None\n        \n    try:\n        # Assuming the spectrograms are stored as parquet or image files\n        if file_path.endswith('.parquet'):\n            spec_df = pd.read_parquet(file_path)\n            return spec_df\n        elif file_path.endswith(('.png', '.jpg', '.jpeg')):\n            # Load as image\n            img = Image.open(file_path)\n            # Convert to numpy array\n            spec_array = np.array(img)\n            return spec_array\n        else:\n            print(f\"Unsupported spectrogram format: {file_path}\")\n            return None\n    except Exception as e:\n        print(f\"Error loading spectrogram {file_path}: {e}\")\n        return None\n\ndef preprocess_spectrogram(spec_data, target_height=128, target_width=256):\n    \"\"\"Preprocess spectrogram data.\"\"\"\n    if spec_data is None:\n        # Return empty spectrogram if data is None\n        return np.zeros((target_height, target_width))\n        \n    # Handle different input types\n    if isinstance(spec_data, pd.DataFrame):\n        # Convert dataframe to 2D array if needed\n        # This depends on your spectrogram format\n        spec_array = spec_data.values\n    else:\n        spec_array = spec_data\n    \n    # Ensure we have a 2D array\n    if len(spec_array.shape) > 2:\n        # If it's an RGB image, convert to grayscale\n        from skimage.color import rgb2gray\n        spec_array = rgb2gray(spec_array)\n    \n    # Resize if needed and if dimensions are valid\n    try:\n        if spec_array.shape[0] != target_height or spec_array.shape[1] != target_width:\n            spec_array = resize(spec_array, (target_height, target_width), \n                              anti_aliasing=True, preserve_range=True)\n    except Exception as e:\n        print(f\"Error resizing spectrogram: {e}\")\n        return np.zeros((target_height, target_width))\n    \n    # Normalize to [0, 1]\n    try:\n        min_val = np.min(spec_array)\n        max_val = np.max(spec_array)\n        if max_val > min_val:\n            spec_array = (spec_array - min_val) / (max_val - min_val)\n        else:\n            spec_array = np.zeros_like(spec_array)\n    except Exception as e:\n        print(f\"Error normalizing spectrogram: {e}\")\n        return np.zeros((target_height, target_width))\n    \n    return spec_array","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T01:57:51.879866Z","iopub.execute_input":"2025-03-25T01:57:51.880162Z","iopub.status.idle":"2025-03-25T01:57:51.893386Z","shell.execute_reply.started":"2025-03-25T01:57:51.880125Z","shell.execute_reply":"2025-03-25T01:57:51.892436Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MultimodalDataset(Dataset):\n    def __init__(self, eeg_files, spectrogram_files, labels=None, is_test=False):\n        self.eeg_files = eeg_files\n        self.spectrogram_files = spectrogram_files\n        self.labels = labels\n        self.is_test = is_test\n        \n        # Calculate window size in samples\n        self.window_size = int(CONFIG['window_size_seconds'] * CONFIG['target_freq'])\n        \n        # Precompute the segments to save memory during training\n        self.eeg_segments = []\n        self.spectrogram_segments = []\n        self.segment_labels = []\n        self.segment_to_file_idx = []\n        \n        for i, (eeg_file, spec_file) in enumerate(tqdm(zip(eeg_files, spectrogram_files), \n                                                       desc=\"Processing files\", \n                                                       total=len(eeg_files))):\n            try:\n                # Process EEG\n                eeg_df = load_eeg_file(eeg_file)\n                eeg_processed = preprocess_eeg(eeg_df, CONFIG['target_freq'])\n                \n                # Segment the EEG\n                eeg_segments = segment_eeg(\n                    eeg_processed, \n                    self.window_size, \n                    CONFIG['overlap_ratio']\n                )\n                \n                # Process spectrogram if available\n                spec_segments = []\n                if CONFIG['use_spectrograms'] and spec_file is not None:\n                    spec_data = load_spectrogram(spec_file)\n                    if spec_data is not None:\n                        spec_processed = preprocess_spectrogram(\n                            spec_data, \n                            CONFIG['spec_height'], \n                            CONFIG['spec_width']\n                        )\n                        # Create a \"segment\" for each EEG segment\n                        # In a real implementation, you may want to segment spectrograms too\n                        spec_segments = [spec_processed] * len(eeg_segments)\n                    else:\n                        # Create empty spectrogram placeholders if loading failed\n                        empty_spec = np.zeros((CONFIG['spec_height'], CONFIG['spec_width']))\n                        spec_segments = [empty_spec] * len(eeg_segments)\n                else:\n                    # Create empty spectrogram placeholders if not using spectrograms\n                    empty_spec = np.zeros((CONFIG['spec_height'], CONFIG['spec_width']))\n                    spec_segments = [empty_spec] * len(eeg_segments)\n                \n                # Store segments and their corresponding labels\n                for eeg_segment, spec_segment in zip(eeg_segments, spec_segments):\n                    # Normalize each EEG segment\n                    eeg_segment = (eeg_segment - np.mean(eeg_segment, axis=0)) / (np.std(eeg_segment, axis=0) + 1e-8)\n                    \n                    self.eeg_segments.append(eeg_segment)\n                    self.spectrogram_segments.append(spec_segment)\n                    \n                    if not self.is_test:\n                        self.segment_labels.append(labels[i])\n                    self.segment_to_file_idx.append(i)\n            except Exception as e:\n                print(f\"Error processing files {eeg_file} / {spec_file}: {e}\")\n        \n        print(f\"Created dataset with {len(self.eeg_segments)} segments from {len(eeg_files)} files\")\n    \n    def __len__(self):\n        return len(self.eeg_segments)\n    \n    def __getitem__(self, idx):\n        eeg_segment = self.eeg_segments[idx]\n        spec_segment = self.spectrogram_segments[idx]\n        \n        # Convert to tensors\n        eeg_tensor = torch.tensor(eeg_segment, dtype=torch.float32)\n        spec_tensor = torch.tensor(spec_segment, dtype=torch.float32).unsqueeze(0)  # Add channel dim\n        \n        # Permute EEG for channel-first format\n        eeg_tensor = eeg_tensor.permute(1, 0)  # [channels, time]\n        if not self.is_test:\n                    label = self.segment_labels[idx]\n                    return (eeg_tensor, spec_tensor), label, self.segment_to_file_idx[idx]\n        else:\n            return (eeg_tensor, spec_tensor), self.segment_to_file_idx[idx]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T01:57:51.894247Z","iopub.execute_input":"2025-03-25T01:57:51.894515Z","iopub.status.idle":"2025-03-25T01:57:51.913258Z","shell.execute_reply.started":"2025-03-25T01:57:51.894494Z","shell.execute_reply":"2025-03-25T01:57:51.912559Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 5: Memory-Optimized Multimodal Dataset Class\nclass BatchMultimodalDataset(Dataset):\n    def __init__(self, eeg_files, spectrogram_files, labels=None, is_test=False):\n        self.eeg_files = eeg_files\n        self.spectrogram_files = spectrogram_files\n        self.labels = labels\n        self.is_test = is_test\n        \n        # Calculate window size in samples\n        self.window_size = int(CONFIG['window_size_seconds'] * CONFIG['target_freq'])\n        \n        # Precompute the segments to save memory during training\n        self.eeg_segments = []\n        self.spectrogram_segments = []\n        self.segment_labels = []\n        self.segment_to_file_idx = []\n        \n        # Track statistics\n        skipped_files = 0\n        processed_files = 0\n        \n        # For batch processing, process files with a timeout\n        for i, (eeg_file, spec_file) in enumerate(tqdm(zip(eeg_files, spectrogram_files), \n                                                      desc=\"Processing files\", \n                                                      total=len(eeg_files))):\n            start_time = time.time()\n            try:\n                # Check if we should skip this file due to timeout\n                if time.time() - start_time > CONFIG['max_file_time']:\n                    print(f\"Timeout skipping file {i}\")\n                    skipped_files += 1\n                    continue\n                \n                # Load and process EEG\n                eeg_df = load_eeg_file(eeg_file)\n                if eeg_df is None:\n                    skipped_files += 1\n                    continue\n                    \n                eeg_processed = preprocess_eeg(eeg_df, CONFIG['target_freq'])\n                if eeg_processed is None:\n                    skipped_files += 1\n                    continue\n                \n                # Check for timeout again\n                if time.time() - start_time > CONFIG['max_file_time']:\n                    print(f\"Processing taking too long for file {i}\")\n                    skipped_files += 1\n                    continue\n                \n                # Segment the EEG\n                eeg_segments = segment_eeg(eeg_processed, self.window_size, CONFIG['overlap_ratio'])\n                if len(eeg_segments) == 0:\n                    skipped_files += 1\n                    continue\n                \n                # Process spectrogram if available\n                spec_processed = None\n                if CONFIG['use_spectrograms'] and spec_file is not None:\n                    # Check for timeout\n                    if time.time() - start_time > CONFIG['max_file_time'] * 0.7:  # 70% of max time\n                        print(f\"Not enough time for spectrogram, skipping for file {i}\")\n                    else:\n                        try:\n                            spec_data = load_spectrogram(spec_file)\n                            if spec_data is not None:\n                                spec_processed = preprocess_spectrogram(\n                                    spec_data, \n                                    CONFIG['spec_height'], \n                                    CONFIG['spec_width']\n                                )\n                        except Exception as e:\n                            print(f\"Error with spectrogram {spec_file}: {e}\")\n                \n                # If spectrogram processing failed, use empty spectrograms\n                if spec_processed is None:\n                    spec_processed = np.zeros((CONFIG['spec_height'], CONFIG['spec_width']))\n                \n                # Create segments\n                for seg_idx, eeg_segment in enumerate(eeg_segments):\n                    # Normalize each EEG segment\n                    try:\n                        norm_segment = (eeg_segment - np.mean(eeg_segment, axis=0)) / (np.std(eeg_segment, axis=0) + 1e-8)\n                        self.eeg_segments.append(norm_segment)\n                        self.spectrogram_segments.append(spec_processed)\n                        \n                        if not self.is_test:\n                            self.segment_labels.append(labels[i])\n                        self.segment_to_file_idx.append(i)\n                    except Exception as e:\n                        print(f\"Error normalizing segment {seg_idx} from file {i}: {e}\")\n                        continue\n                \n                processed_files += 1\n                \n                # Clean up memory\n                del eeg_df, eeg_processed, eeg_segments\n                if spec_data is not None:\n                    del spec_data\n                \n                # Perform garbage collection periodically\n                if i % 10 == 0:\n                    gc.collect()\n                    if torch.cuda.is_available():\n                        torch.cuda.empty_cache()\n                \n            except Exception as e:\n                print(f\"Error processing file {i} ({eeg_file}): {e}\")\n                skipped_files += 1\n                continue\n                \n            # Stop if we've spent too much time on this batch\n            if time.time() - start_time > CONFIG['max_file_time'] * 2:\n                print(f\"File {i} took too long, forcing garbage collection\")\n                gc.collect()\n                if torch.cuda.is_available():\n                    torch.cuda.empty_cache()\n        \n        print(f\"Created dataset with {len(self.eeg_segments)} segments from {processed_files} files\")\n        print(f\"Skipped {skipped_files} problematic files\")\n    \n    def __len__(self):\n        return len(self.eeg_segments)\n    \n    def __getitem__(self, idx):\n        eeg_segment = self.eeg_segments[idx]\n        spec_segment = self.spectrogram_segments[idx]\n        \n        # Convert to tensors\n        eeg_tensor = torch.tensor(eeg_segment, dtype=torch.float32)\n        spec_tensor = torch.tensor(spec_segment, dtype=torch.float32).unsqueeze(0)  # Add channel dim\n        \n        # Permute EEG for channel-first format\n        eeg_tensor = eeg_tensor.permute(1, 0)  # [channels, time]\n        \n        if not self.is_test:\n            label = self.segment_labels[idx]\n            return (eeg_tensor, spec_tensor), label, self.segment_to_file_idx[idx]\n        else:\n            return (eeg_tensor, spec_tensor), self.segment_to_file_idx[idx]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T01:57:51.913968Z","iopub.execute_input":"2025-03-25T01:57:51.914307Z","iopub.status.idle":"2025-03-25T01:57:51.933887Z","shell.execute_reply.started":"2025-03-25T01:57:51.914276Z","shell.execute_reply":"2025-03-25T01:57:51.933144Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 6: Multimodal Model Architecture\nclass MultimodalEEGNet(nn.Module):\n    def __init__(self, \n                 num_eeg_channels=len(CONFIG['channels']), \n                 num_classes=CONFIG['n_classes'], \n                 dropout_rate=CONFIG['dropout_rate']):\n        super(MultimodalEEGNet, self).__init__()\n        \n        # EEG branch - Same as previous EEGNet\n        self.eeg_conv1 = nn.Conv1d(\n            in_channels=num_eeg_channels, \n            out_channels=32,\n            kernel_size=64, \n            stride=2, \n            padding=32, \n            bias=False\n        )\n        self.eeg_batchnorm1 = nn.BatchNorm1d(32)\n        \n        self.eeg_depthwise_conv = nn.Conv1d(\n            in_channels=32, \n            out_channels=64,\n            kernel_size=16, \n            groups=32,\n            stride=2, \n            padding=8, \n            bias=False\n        )\n        self.eeg_batchnorm2 = nn.BatchNorm1d(64)\n        self.eeg_activation = nn.ELU()\n        self.eeg_avgpool1 = nn.AvgPool1d(kernel_size=4, stride=4)\n        self.eeg_dropout1 = nn.Dropout(dropout_rate)\n        \n        self.eeg_seperable_conv = nn.Conv1d(\n            in_channels=64, \n            out_channels=64, \n            kernel_size=16, \n            padding=8, \n            bias=False\n        )\n        self.eeg_batchnorm3 = nn.BatchNorm1d(64)\n        self.eeg_avgpool2 = nn.AvgPool1d(kernel_size=8, stride=8)\n        self.eeg_dropout2 = nn.Dropout(dropout_rate)\n        \n        # Spectrogram branch - CNN for spectrogram processing\n        self.spec_conv1 = nn.Conv2d(1, 16, kernel_size=3, stride=1, padding=1)\n        self.spec_batchnorm1 = nn.BatchNorm2d(16)\n        self.spec_maxpool1 = nn.MaxPool2d(kernel_size=2, stride=2)\n        \n        self.spec_conv2 = nn.Conv2d(16, 32, kernel_size=3, stride=1, padding=1)\n        self.spec_batchnorm2 = nn.BatchNorm2d(32)\n        self.spec_maxpool2 = nn.MaxPool2d(kernel_size=2, stride=2)\n        \n        self.spec_conv3 = nn.Conv2d(32, 64, kernel_size=3, stride=1, padding=1)\n        self.spec_batchnorm3 = nn.BatchNorm2d(64)\n        self.spec_maxpool3 = nn.MaxPool2d(kernel_size=2, stride=2)\n        self.spec_dropout = nn.Dropout(dropout_rate)\n        \n        # Calculate feature dimensions after pooling operations\n        # EEG branch: 64 channels * 7 time points = 448\n        # Spectrogram branch: 64 * (height/8) * (width/8) = 64 * 16 * 32 = 32,768\n        eeg_features = 64 * 7\n        spec_features = 64 * (CONFIG['spec_height'] // 8) * (CONFIG['spec_width'] // 8)\n        \n        # Fusion and classification layers\n        self.eeg_fc = nn.Linear(eeg_features, 128)\n        self.spec_fc = nn.Linear(spec_features, 128)\n        \n        self.fusion_dropout = nn.Dropout(dropout_rate)\n        self.fusion_fc = nn.Linear(256, 128)  # Combined features\n        self.classifier = nn.Linear(128, num_classes)\n        \n        self.activation = nn.ELU()\n    \n    def forward(self, x):\n        try:\n            # Split inputs\n            eeg_input, spec_input = x\n            \n            # Process EEG branch\n            eeg = self.eeg_conv1(eeg_input)\n            eeg = self.eeg_batchnorm1(eeg)\n            \n            eeg = self.eeg_depthwise_conv(eeg)\n            eeg = self.eeg_batchnorm2(eeg)\n            eeg = self.eeg_activation(eeg)\n            eeg = self.eeg_avgpool1(eeg)\n            eeg = self.eeg_dropout1(eeg)\n            \n            eeg = self.eeg_seperable_conv(eeg)\n            eeg = self.eeg_batchnorm3(eeg)\n            eeg = self.eeg_activation(eeg)\n            eeg = self.eeg_avgpool2(eeg)\n            eeg = self.eeg_dropout2(eeg)\n            \n            eeg = eeg.view(eeg.size(0), -1)  # Flatten\n            eeg = self.eeg_fc(eeg)\n            eeg = self.activation(eeg)\n            \n            # Process Spectrogram branch\n            spec = self.spec_conv1(spec_input)\n            spec = self.spec_batchnorm1(spec)\n            spec = self.activation(spec)\n            spec = self.spec_maxpool1(spec)\n            \n            spec = self.spec_conv2(spec)\n            spec = self.spec_batchnorm2(spec)\n            spec = self.activation(spec)\n            spec = self.spec_maxpool2(spec)\n            \n            spec = self.spec_conv3(spec)\n            spec = self.spec_batchnorm3(spec)\n            spec = self.activation(spec)\n            spec = self.spec_maxpool3(spec)\n            spec = self.spec_dropout(spec)\n            \n            spec = spec.view(spec.size(0), -1)  # Flatten\n            spec = self.spec_fc(spec)\n            spec = self.activation(spec)\n            \n            # Combine features\n            combined = torch.cat((eeg, spec), dim=1)\n            combined = self.fusion_dropout(combined)\n            combined = self.fusion_fc(combined)\n            combined = self.activation(combined)\n            \n            # Classification\n            output = self.classifier(combined)\n            \n            return output\n        except Exception as e:\n            print(f\"Error in model forward pass: {e}\")\n            # Return zeros as a fallback\n            return torch.zeros(eeg_input.size(0), self.classifier.out_features, device=eeg_input.device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T01:57:51.934995Z","iopub.execute_input":"2025-03-25T01:57:51.935323Z","iopub.status.idle":"2025-03-25T01:57:51.949832Z","shell.execute_reply.started":"2025-03-25T01:57:51.935292Z","shell.execute_reply":"2025-03-25T01:57:51.949098Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 7: Training and Evaluation Functions\ndef train_model(model, train_loader, val_loader, criterion, optimizer, scheduler, num_epochs, fold):\n    best_val_loss = float('inf')\n    patience_counter = 0\n    \n    # For storing metrics\n    train_losses = []\n    val_losses = []\n    val_accuracies = []\n    \n    # Wrap model with DataParallel if using multiple GPUs\n    if CONFIG['use_multi_gpu'] and torch.cuda.device_count() > 1:\n        model = DataParallel(model)\n        print(f\"Using {torch.cuda.device_count()} GPUs for training!\")\n    \n    for epoch in range(num_epochs):\n        start_time = time.time()\n        \n        # Training phase\n        model.train()\n        train_loss = 0.0\n        batch_count = 0\n        \n        for batch_idx, ((eeg_data, spec_data), labels, _) in enumerate(tqdm(train_loader, desc=f\"Epoch {epoch+1}/{num_epochs} Training\")):\n            try:\n                eeg_data, spec_data, labels = eeg_data.to(device), spec_data.to(device), labels.to(device)\n                \n                # Zero the parameter gradients\n                optimizer.zero_grad()\n                \n                if CONFIG['mixed_precision']:\n                    # Mixed precision training\n                    with autocast():\n                        outputs = model((eeg_data, spec_data))\n                        loss = criterion(outputs, labels)\n                    \n                    # Scale the loss and backpropagate\n                    scaler.scale(loss).backward()\n                    scaler.step(optimizer)\n                    scaler.update()\n                else:\n                    # Regular training\n                    outputs = model((eeg_data, spec_data))\n                    loss = criterion(outputs, labels)\n                    loss.backward()\n                    optimizer.step()\n                \n                train_loss += loss.item()\n                batch_count += 1\n                \n                # Print statistics\n                if (batch_idx + 1) % 20 == 0 or batch_idx == 0:\n                    print(f\"Batch {batch_idx + 1}/{len(train_loader)}, Loss: {loss.item():.4f}\")\n                \n                # Clear memory\n                del eeg_data, spec_data, outputs, loss\n                if batch_idx % 10 == 0:\n                    torch.cuda.empty_cache()\n                \n            except Exception as e:\n                print(f\"Error in training batch {batch_idx}: {e}\")\n                continue\n        \n        if batch_count > 0:\n            train_loss /= batch_count\n        train_losses.append(train_loss)\n        \n        # Validation phase\n        model.eval()\n        val_loss = 0.0\n        val_preds = []\n        val_true = []\n        val_batch_count = 0\n        \n        with torch.no_grad():\n            for (eeg_data, spec_data), labels, _ in tqdm(val_loader, desc=f\"Epoch {epoch+1}/{num_epochs} Validation\"):\n                try:\n                    eeg_data, spec_data, labels = eeg_data.to(device), spec_data.to(device), labels.to(device)\n                    \n                    if CONFIG['mixed_precision']:\n                        with autocast():\n                            outputs = model((eeg_data, spec_data))\n                            loss = criterion(outputs, labels)\n                    else:\n                        outputs = model((eeg_data, spec_data))\n                        loss = criterion(outputs, labels)\n                    \n                    val_loss += loss.item()\n                    val_batch_count += 1\n                    \n                    _, predicted = torch.max(outputs, 1)\n                    val_preds.extend(predicted.cpu().numpy())\n                    val_true.extend(labels.cpu().numpy())\n                    \n                    # Clear memory\n                    del eeg_data, spec_data, outputs, loss\n                    \n                except Exception as e:\n                    print(f\"Error in validation batch: {e}\")\n                    continue\n        \n        if val_batch_count > 0:\n            val_loss /= val_batch_count\n        val_losses.append(val_loss)\n        \n        # Calculate validation accuracy\n        if len(val_preds) > 0 and len(val_true) > 0:\n            val_accuracy = accuracy_score(val_true, val_preds)\n            val_accuracies.append(val_accuracy)\n        else:\n            val_accuracy = 0.0\n            val_accuracies.append(0.0)\n            print(\"Warning: No valid predictions for validation accuracy calculation\")\n        \n        # Calculate epoch time\n        epoch_time = time.time() - start_time\n        \n        print(f\"Epoch {epoch+1}/{num_epochs}, Train Loss: {train_loss:.4f}, Val Loss: {val_loss:.4f}, Val Accuracy: {val_accuracy:.4f}\")\n        print(f\"Epoch time: {epoch_time:.2f} seconds\")\n        \n        # Check if we should save the model\n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            # Save the model\n            try:\n                if isinstance(model, DataParallel):\n                    torch.save(model.module.state_dict(), f\"{OUTPUT_DIR}/model_fold{fold}_epoch{epoch+1}.pt\")\n                else:\n                    torch.save(model.state_dict(), f\"{OUTPUT_DIR}/model_fold{fold}_epoch{epoch+1}.pt\")\n                print(f\"Model saved at epoch {epoch+1}\")\n            except Exception as e:\n                print(f\"Error saving model: {e}\")\n            \n            patience_counter = 0\n        else:\n            patience_counter += 1\n            \n        # Early stopping\n        if patience_counter >= CONFIG['patience']:\n            print(f\"Early stopping triggered after epoch {epoch+1}\")\n            break\n        \n        # Update learning rate\n        scheduler.step()\n        \n        # Cleanup to prevent memory leaks\n        torch.cuda.empty_cache()\n        gc.collect()\n    \n    # Plot training curves\n    plt.figure(figsize=(12, 5))\n    \n    plt.subplot(1, 2, 1)\n    plt.plot(train_losses, label='Training Loss')\n    plt.plot(val_losses, label='Validation Loss')\n    plt.title(f'Fold {fold} - Loss')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.legend()\n    \n    plt.subplot(1, 2, 2)\n    plt.plot(val_accuracies, label='Validation Accuracy')\n    plt.title(f'Fold {fold} - Accuracy')\n    plt.xlabel('Epoch')\n    plt.ylabel('Accuracy')\n    plt.legend()\n    \n    plt.tight_layout()\n    plt.savefig(f\"{OUTPUT_DIR}/training_curves_fold{fold}.png\")\n    \n    return model\n\ndef predict(model, test_loader):\n    model.eval()\n    all_predictions = []\n    file_indices = []\n    \n    with torch.no_grad():\n        for batch_idx, ((eeg_data, spec_data), file_idx) in enumerate(tqdm(test_loader, desc=\"Generating predictions\")):\n            try:\n                eeg_data, spec_data = eeg_data.to(device), spec_data.to(device)\n                \n                if CONFIG['mixed_precision']:\n                    with autocast():\n                        outputs = model((eeg_data, spec_data))\n                else:\n                    outputs = model((eeg_data, spec_data))\n                    \n                probabilities = F.softmax(outputs, dim=1)\n                all_predictions.append(probabilities.cpu().numpy())\n                file_indices.extend(file_idx.cpu().numpy())\n                \n                # Clean up memory\n                del eeg_data, spec_data, outputs, probabilities\n                if batch_idx % 10 == 0:\n                    torch.cuda.empty_cache()\n                    \n            except Exception as e:\n                print(f\"Error during prediction batch {batch_idx}: {e}\")\n                # Continue with next batch if there's an error\n                continue\n    \n    # Concatenate all predictions\n    if len(all_predictions) > 0:\n        all_predictions = np.vstack(all_predictions)\n    else:\n        print(\"Warning: No predictions were generated\")\n        all_predictions = np.array([])\n    \n    return all_predictions, np.array(file_indices)\n\ndef aggregate_predictions(predictions, file_indices, num_files, num_classes):\n    \"\"\"Aggregate predictions from multiple segments to file-level predictions.\"\"\"\n    # Initialize file-level predictions\n    file_preds = np.zeros((num_files, num_classes))\n    counts = np.zeros(num_files)\n    \n    # Sum up predictions for each file\n    if len(predictions) > 0:\n        for pred, file_idx in zip(predictions, file_indices):\n            file_preds[file_idx] += pred\n            counts[file_idx] += 1\n    \n    # Average predictions\n    for i in range(num_files):\n        if counts[i] > 0:\n            file_preds[i] /= counts[i]\n        else:\n            # If no predictions for this file, use a uniform distribution\n            file_preds[i] = np.ones(num_classes) / num_classes\n            print(f\"Warning: No predictions for file index {i}\")\n    \n    return file_preds","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T01:57:51.951600Z","iopub.execute_input":"2025-03-25T01:57:51.951903Z","iopub.status.idle":"2025-03-25T01:57:51.972451Z","shell.execute_reply.started":"2025-03-25T01:57:51.951882Z","shell.execute_reply":"2025-03-25T01:57:51.971579Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main_test():\n    \"\"\"Run a test with a small dataset to verify the pipeline.\"\"\"\n    print(\"Running test with a smaller dataset...\")\n    \n    # Monitor memory usage\n    if torch.cuda.is_available():\n        print(f\"GPU memory before training: {torch.cuda.memory_allocated() / 1e9:.2f} GB\")\n    \n    # Use only a small subset for testing\n    subset_size = 100  # Just use 100 files for testing\n    \n    # Prepare file paths - use only a subset\n    train_eeg_files = [os.path.join(TRAIN_DIR, f\"{file_id}.parquet\") for file_id in train_metadata['eeg_id'][:subset_size]]\n    test_eeg_files = [os.path.join(TEST_DIR, f\"{file_id}.parquet\") for file_id in test_metadata['eeg_id'][:10]]  # Just a few test files\n    \n    # Prepare spectrogram file paths\n    train_spec_files = [None] * len(train_eeg_files)  # Start with EEG only for testing\n    test_spec_files = [None] * len(test_eeg_files)\n    \n    # Temporarily disable spectrogram usage for faster testing\n    CONFIG['use_spectrograms'] = False\n    \n    # Convert labels to numerical values for the subset\n    train_labels = np.array([labels_map[label] for label in train_metadata[label_column][:subset_size]])\n    \n    # Just use 2 folds for testing\n    CONFIG['fold_count'] = 2\n    \n    # Define cross-validation strategy\n    kfold = StratifiedKFold(n_splits=CONFIG['fold_count'], shuffle=True, random_state=42)\n    \n    # Loop through folds\n    for fold, (train_idx, val_idx) in enumerate(kfold.split(train_eeg_files, train_labels)):\n        print(f\"\\n{'='*20} Test Fold {fold+1}/{CONFIG['fold_count']} {'='*20}\\n\")\n        \n        # Split data for this fold - use fewer files for testing\n        fold_train_eeg_files = [train_eeg_files[i] for i in train_idx][:100]  # Just use 100 files\n        fold_val_eeg_files = [train_eeg_files[i] for i in val_idx][:20]      # Just use 20 files\n        fold_train_spec_files = [train_spec_files[i] for i in train_idx][:100]\n        fold_val_spec_files = [train_spec_files[i] for i in val_idx][:20]\n        fold_train_labels = train_labels[train_idx][:100]\n        fold_val_labels = train_labels[val_idx][:20]\n        \n        # Create datasets with verbose error reporting\n        try:\n            print(f\"Creating training dataset for test fold {fold+1}...\")\n            train_dataset = MultimodalDataset(fold_train_eeg_files, fold_train_spec_files, fold_train_labels)\n            \n            print(f\"Successfully created training dataset with {len(train_dataset)} segments\")\n            print(f\"Creating validation dataset for test fold {fold+1}...\")\n            val_dataset = MultimodalDataset(fold_val_eeg_files, fold_val_spec_files, fold_val_labels)\n            print(f\"Successfully created validation dataset with {len(val_dataset)} segments\")\n            \n            print(\"Test successful! The pipeline is working correctly.\")\n            return\n            \n        except Exception as e:\n            print(f\"Error in test fold {fold+1}: {e}\")\n            import traceback\n            traceback.print_exc()\n            return","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T01:57:51.973501Z","iopub.execute_input":"2025-03-25T01:57:51.973842Z","iopub.status.idle":"2025-03-25T01:57:51.988552Z","shell.execute_reply.started":"2025-03-25T01:57:51.973818Z","shell.execute_reply":"2025-03-25T01:57:51.987831Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 8: Main Competition Function - Batched Processing\ndef main_competition_batched():\n    # Use 30% of the total training data\n    total_samples = len(train_metadata)\n    subset_size = int(total_samples * 0.3)  # ~32,000 samples\n    \n    print(f\"Running competition solution with {subset_size} samples (30% of total data) using batched processing...\")\n    \n    # Monitor memory usage\n    if torch.cuda.is_available():\n        print(f\"GPU memory before training: {torch.cuda.memory_allocated() / 1e9:.2f} GB\")\n    \n    # Prepare file paths\n    train_eeg_files = [os.path.join(TRAIN_DIR, f\"{file_id}.parquet\") for file_id in train_metadata['eeg_id'][:subset_size]]\n    test_eeg_files = [os.path.join(TEST_DIR, f\"{file_id}.parquet\") for file_id in test_metadata['eeg_id']]\n    \n    # Determine if spectrograms can be used\n    has_spectrograms = 'spectrogram_id' in train_metadata.columns\n    if has_spectrograms:\n        print(\"Found spectrogram data, using multimodal approach\")\n        CONFIG['use_spectrograms'] = True\n        train_spec_files = []\n        for i, spec_id in enumerate(train_metadata['spectrogram_id'][:subset_size]):\n            spec_path = os.path.join(TRAIN_SPECTROGRAMS_DIR, f\"{spec_id}.parquet\")\n            if not os.path.exists(spec_path):\n                # Try other possible extensions\n                for ext in ['.png', '.jpg', '.npy']:\n                    alt_path = os.path.join(TRAIN_SPECTROGRAMS_DIR, f\"{spec_id}{ext}\")\n                    if os.path.exists(alt_path):\n                        spec_path = alt_path\n                        break\n            train_spec_files.append(spec_path)\n            \n        test_spec_files = []\n        for i, spec_id in enumerate(test_metadata['spectrogram_id']):\n            spec_path = os.path.join(TEST_SPECTROGRAMS_DIR, f\"{spec_id}.parquet\")\n            if not os.path.exists(spec_path):\n                # Try other possible extensions\n                for ext in ['.png', '.jpg', '.npy']:\n                    alt_path = os.path.join(TEST_SPECTROGRAMS_DIR, f\"{spec_id}{ext}\")\n                    if os.path.exists(alt_path):\n                        spec_path = alt_path\n                        break\n            test_spec_files.append(spec_path)\n    else:\n        print(\"No spectrogram data found, using EEG only\")\n        CONFIG['use_spectrograms'] = False\n        train_spec_files = [None] * len(train_eeg_files)\n        test_spec_files = [None] * len(test_eeg_files)\n    \n    # Convert labels to numerical values\n    train_labels = np.array([labels_map[label] for label in train_metadata[label_column][:subset_size]])\n    \n    # Run a more practical number of folds\n    CONFIG['fold_count'] = 3  # Use fewer folds to save time\n    \n    # Define cross-validation strategy\n    kfold = StratifiedKFold(n_splits=CONFIG['fold_count'], shuffle=True, random_state=42)\n    \n    # Store fold predictions for ensemble\n    fold_test_preds = []\n    \n    # Loop through folds\n    for fold, (train_idx, val_idx) in enumerate(kfold.split(train_eeg_files, train_labels)):\n        print(f\"\\n{'='*20} Fold {fold+1}/{CONFIG['fold_count']} {'='*20}\\n\")\n        \n        # Split data for this fold\n        fold_train_eeg_files = [train_eeg_files[i] for i in train_idx]\n        fold_val_eeg_files = [train_eeg_files[i] for i in val_idx]\n        fold_train_spec_files = [train_spec_files[i] for i in train_idx]\n        fold_val_spec_files = [train_spec_files[i] for i in val_idx]\n        fold_train_labels = train_labels[train_idx]\n        fold_val_labels = train_labels[val_idx]\n        \n        # Process in batches\n        batch_size = CONFIG['batch_processing_size']\n        \n        # Create training datasets in batches\n        train_datasets = []\n        for batch_start in range(0, len(fold_train_eeg_files), batch_size):\n            batch_end = min(batch_start + batch_size, len(fold_train_eeg_files))\n            print(f\"Processing training batch {batch_start//batch_size + 1}/{(len(fold_train_eeg_files) + batch_size - 1)//batch_size}...\")\n            \n            batch_train_eeg = fold_train_eeg_files[batch_start:batch_end]\n            batch_train_spec = fold_train_spec_files[batch_start:batch_end]\n            batch_train_labels = fold_train_labels[batch_start:batch_end]\n            \n            try:\n                batch_dataset = BatchMultimodalDataset(batch_train_eeg, batch_train_spec, batch_train_labels)\n                if len(batch_dataset) > 0:\n                    train_datasets.append(batch_dataset)\n                    \n                # Force garbage collection after each batch\n                gc.collect()\n                torch.cuda.empty_cache()\n                \n            except Exception as e:\n                print(f\"Error processing training batch: {e}\")\n                import traceback\n                traceback.print_exc()\n        \n        # Create validation datasets in batches\n        val_datasets = []\n        for batch_start in range(0, len(fold_val_eeg_files), batch_size):\n            batch_end = min(batch_start + batch_size, len(fold_val_eeg_files))\n            print(f\"Processing validation batch {batch_start//batch_size + 1}/{(len(fold_val_eeg_files) + batch_size - 1)//batch_size}...\")\n            \n            batch_val_eeg = fold_val_eeg_files[batch_start:batch_end]\n            batch_val_spec = fold_val_spec_files[batch_start:batch_end]\n            batch_val_labels = fold_val_labels[batch_start:batch_end]\n            \n            try:\n                batch_dataset = BatchMultimodalDataset(batch_val_eeg, batch_val_spec, batch_val_labels)\n                if len(batch_dataset) > 0:\n                    val_datasets.append(batch_dataset)\n                    \n                # Force garbage collection after each batch\n                gc.collect()\n                torch.cuda.empty_cache()\n                \n            except Exception as e:\n                print(f\"Error processing validation batch: {e}\")\n                import traceback\n                traceback.print_exc()\n        \n        # Combine datasets and create data loaders\n        if len(train_datasets) > 0 and len(val_datasets) > 0:\n            train_dataset = ConcatDataset(train_datasets)\n            val_dataset = ConcatDataset(val_datasets)\n            \n            print(f\"Combined training dataset has {len(train_dataset)} segments\")\n            print(f\"Combined validation dataset has {len(val_dataset)} segments\")\n            \n            # Create data loaders\n            train_loader = DataLoader(\n                train_dataset, \n                batch_size=CONFIG['batch_size'], \n                shuffle=True, \n                num_workers=CONFIG['num_workers'],\n                pin_memory=CONFIG['pin_memory']\n            )\n            \n            val_loader = DataLoader(\n                val_dataset, \n                batch_size=CONFIG['batch_size'], \n                shuffle=False, \n                num_workers=CONFIG['num_workers'],\n                pin_memory=CONFIG['pin_memory']\n            )\n            \n            # Initialize model\n            model = MultimodalEEGNet(\n                num_eeg_channels=len(CONFIG['channels']),\n                num_classes=CONFIG['n_classes'],\n                dropout_rate=CONFIG['dropout_rate']\n            ).to(device)\n            \n            # Define loss function and optimizer\n            criterion = nn.CrossEntropyLoss()\n            optimizer = AdamW(\n                model.parameters(), \n                lr=CONFIG['learning_rate'], \n                weight_decay=CONFIG['weight_decay']\n            )\n            \n            # Define learning rate scheduler\n            scheduler = CosineAnnealingLR(\n                optimizer, \n                T_max=CONFIG['num_epochs']\n            )\n            \n            # Train the model\n            try:\n                model = train_model(\n                    model, \n                    train_loader, \n                    val_loader, \n                    criterion, \n                    optimizer, \n                    scheduler, \n                    CONFIG['num_epochs'],\n                    fold\n                )\n                \n                # Find the best model for this fold\n                model_files = glob.glob(f\"{OUTPUT_DIR}/model_fold{fold}_*.pt\")\n                if model_files:\n                    best_model_path = sorted(model_files)[-1]\n                    print(f\"Using model: {best_model_path}\")\n                    \n                    # Create a clean model instance for inference\n                    inference_model = MultimodalEEGNet(\n                        num_eeg_channels=len(CONFIG['channels']),\n                        num_classes=CONFIG['n_classes'],\n                        dropout_rate=CONFIG['dropout_rate']\n                    ).to(device)\n                    \n                    inference_model.load_state_dict(torch.load(best_model_path))\n                    \n                    # If using multiple GPUs, wrap the inference model\n                    if CONFIG['use_multi_gpu'] and torch.cuda.device_count() > 1:\n                        inference_model = DataParallel(inference_model)\n                    \n                    # Process test data in batches\n                    all_test_preds = []\n                    \n                    # Break test prediction into manageable batches\n                    test_batch_size = min(CONFIG['batch_processing_size'], 100)  # Smaller batch for test files\n                    \n                    for test_batch_start in range(0, len(test_eeg_files), test_batch_size):\n                        test_batch_end = min(test_batch_start + test_batch_size, len(test_eeg_files))\n                        print(f\"Processing test batch {test_batch_start//test_batch_size + 1}/{(len(test_eeg_files) + test_batch_size - 1)//test_batch_size}...\")\n                        \n                        batch_test_eeg = test_eeg_files[test_batch_start:test_batch_end]\n                        batch_test_spec = test_spec_files[test_batch_start:test_batch_end]\n                        \n                        try:\n                            # Create test dataset for this batch\n                            batch_test_dataset = BatchMultimodalDataset(batch_test_eeg, batch_test_spec, is_test=True)\n                            \n                            if len(batch_test_dataset) > 0:\n                                batch_test_loader = DataLoader(\n                                    batch_test_dataset, \n                                    batch_size=CONFIG['batch_size'], \n                                    shuffle=False, \n                                    num_workers=CONFIG['num_workers'],\n                                    pin_memory=CONFIG['pin_memory']\n                                )\n                                \n                                # Generate predictions\n                                batch_preds, batch_indices = predict(inference_model, batch_test_loader)\n                                \n                                # Adjust indices to account for batch offset\n                                adjusted_indices = [idx + test_batch_start for idx in batch_indices]\n                                \n                                # Store predictions\n                                if len(batch_preds) > 0:\n                                    for pred, file_idx in zip(batch_preds, adjusted_indices):\n                                        all_test_preds.append((pred, file_idx))\n                                \n                                # Clean up\n                                del batch_test_dataset, batch_test_loader, batch_preds, batch_indices\n                                gc.collect()\n                                torch.cuda.empty_cache()\n                                \n                        except Exception as e:\n                            print(f\"Error processing test batch: {e}\")\n                            import traceback\n                            traceback.print_exc()\n                    \n                    # Create file-level predictions\n                    if all_test_preds:\n                        # Extract predictions and indices\n                        test_preds = np.array([p[0] for p in all_test_preds])\n                        test_indices = np.array([p[1] for p in all_test_preds])\n                        \n                        # Aggregate to file level\n                        file_preds = aggregate_predictions(\n                            test_preds, \n                            test_indices, \n                            len(test_eeg_files), \n                            CONFIG['n_classes']\n                        )\n                        \n                        fold_test_preds.append(file_preds)\n                    else:\n                        print(f\"No predictions generated for fold {fold}\")\n                else:\n                    print(f\"No model files found for fold {fold}\")\n            except Exception as e:\n                print(f\"Error in training/inference for fold {fold}: {e}\")\n                import traceback\n                traceback.print_exc()\n        else:\n            print(f\"Insufficient data for fold {fold}, skipping\")\n        \n        # Free up memory\n        if 'model' in locals():\n            del model\n        if 'inference_model' in locals():\n            del inference_model\n        if 'train_dataset' in locals():\n            del train_dataset\n        if 'val_dataset' in locals():\n            del val_dataset\n        if 'train_loader' in locals():\n            del train_loader\n        if 'val_loader' in locals():\n            del val_loader\n        \n        # Free up batch datasets\n        del train_datasets, val_datasets\n        \n        # Clean up all other memory\n        gc.collect()\n        torch.cuda.empty_cache()\n        \n        if torch.cuda.is_available():\n            print(f\"GPU memory after fold {fold+1}: {torch.cuda.memory_allocated() / 1e9:.2f} GB\")\n    \n    # Ensemble predictions from all folds\n    if fold_test_preds:\n        ensemble_preds = np.mean(fold_test_preds, axis=0)\n        \n        # Create submission file\n        submission = pd.DataFrame()\n        submission['eeg_id'] = test_metadata['eeg_id']\n        \n        for label, idx in labels_map.items():\n            submission[label] = ensemble_preds[:, idx]\n        \n        # Save submission file\n        submission.to_csv(f\"{OUTPUT_DIR}/submission.csv\", index=False)\n        print(f\"Submission file saved to {OUTPUT_DIR}/submission.csv\")\n    else:\n        print(\"No predictions were generated from any fold\")\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T01:57:51.989373Z","iopub.execute_input":"2025-03-25T01:57:51.989632Z","iopub.status.idle":"2025-03-25T01:57:52.013091Z","shell.execute_reply.started":"2025-03-25T01:57:51.989594Z","shell.execute_reply":"2025-03-25T01:57:52.012363Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 10: Run the Competition Pipeline\n# First, run the test to make sure everything works\nprint(\"Running quick test to verify pipeline...\")\nmain_test()\n\n# If the test is successful, run the full competition solution with batched processing\nprint(\"\\nRunning full competition solution with 30% of the data using batched processing...\")\nmain_competition_batched()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T01:57:52.013915Z","iopub.execute_input":"2025-03-25T01:57:52.014148Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}