{"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":"gpu","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7970005,"sourceType":"datasetVersion","datasetId":4689524},{"sourceId":7997245,"sourceType":"datasetVersion","datasetId":4659732},{"sourceId":8007490,"sourceType":"datasetVersion","datasetId":4476729},{"sourceId":8015911,"sourceType":"datasetVersion","datasetId":4469088},{"sourceId":8041644,"sourceType":"datasetVersion","datasetId":4679791},{"sourceId":8041978,"sourceType":"datasetVersion","datasetId":4716292},{"sourceId":8049430,"sourceType":"datasetVersion","datasetId":4564661},{"sourceId":8049601,"sourceType":"datasetVersion","datasetId":4515496}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Ensemble Learning for EEG Classification\n\nThis notebook implements a comprehensive ensemble learning framework designed to enhance EEG classification accuracy. By fusing predictions from multiple specialized models, our approach leverages the complementary strengths of different feature extraction strategies to capture the rich, multi-faceted nature of EEG signals. Ensemble methods have long been recognized for their ability to reduce model variance and improve robustness, which is particularly valuable in the context of complex biological signals like EEG.\n\n## Ensemble Branches\n\nThe ensemble is built on two distinct yet complementary branches:\n\n- **X3D Branch:**  \n  This branch converts EEG signals into spectrograms and processes them using a state-of-the-art X3D model, which is tailored to capture spatiotemporal patterns. By operating on multi-resolution representations of the EEG data, the X3D branch is adept at identifying both fine temporal details and broad spectral trends. This dual focus is crucial for distinguishing subtle differences in neural activity across different classes.\n\n- **Double-Head Butter Filter Branch:**  \n  This branch takes a different approach by first applying a Butterworth filter to the raw EEG signal, thereby transforming it into a spectrogram representation that emphasizes lower frequency components. It then extracts features from both the filtered (spectrogram) and unfiltered (raw EEG) data streams. The dual-head strategy enables the model to integrate complementary information from the time and frequency domains, enhancing its ability to capture the nuanced characteristics inherent in EEG signals.\n\n## Detailed Workflow Overview\n\n1. **Data Preparation:**\n   - **Loading:** EEG data is loaded from the dataset, ensuring that raw signals are correctly imported for further processing.\n   - **Preprocessing:** The data undergoes essential preprocessing steps such as filtering, normalization, and segmentation. These steps are vital for reducing noise and preparing the data for robust feature extraction.\n\n2. **Model Inference:**\n   - **Multi-Resolution Analysis:**  \n     The X3D branch computes spectrograms at various temporal resolutions. This multi-scale analysis allows the model to capture both micro and macro-level features of the EEG signal.\n   - **Butter Filter Transformation:**  \n     The Double-Head Butter Filter branch applies a Butterworth filter to focus on key frequency bands (typically lower frequencies). Subsequent logarithmic scaling and clipping are performed to emphasize critical spectral features.\n   - **Data Augmentation via Flipping:**  \n     To further enhance model robustness, each model weight is used to perform inference twice: once with the original data and once with a horizontally flipped version. This augmentation helps capture different perspectives of the data variability.\n   - **GPU Memory Management:**  \n     Efficient GPU memory handling is a priority. After processing each model weight, GPU memory is explicitly cleared to avoid memory buildup and ensure smooth operation during the inference process.\n\n3. **Ensemble Aggregation:**\n   - **Averaging Predictions:**  \n     The predictions obtained from the different models and augmentation configurations are aggregated by averaging. This process mitigates individual model biases and results in a more stable and accurate final prediction.\n   - **Robust Decision Making:**  \n     The ensemble approach capitalizes on the diverse strengths of each branch, leading to improved generalization and reliability in classifying EEG data.\n\n*Note: GPU memory is cleared after each model weight inference to ensure efficient resource utilization and prevent potential memory overflow issues.*\n\nBy combining these advanced techniques into a unified framework, our ensemble learning method not only achieves higher accuracy in EEG classification but also provides deeper insights into the underlying neural patterns, paving the way for more robust and interpretable models in neuroscience applications.\n","metadata":{}},{"cell_type":"code","source":"import random\nimport cv2\nimport json\nimport numpy as np\nimport copy\nimport pandas as pd\nimport torch\nfrom torch import Tensor\nimport gc\nfrom typing import Any, List, Tuple, Union\n\nimport albumentations as A\nimport os\nimport librosa\nimport pickle\nimport timm\nfrom tqdm import tqdm\nimport matplotlib.pyplot as pl\n\nimport mne\n\nimport torch\nimport torchaudio\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader\nfrom scipy.signal import butter, lfilter","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-08T11:28:51.478276Z","iopub.execute_input":"2025-04-08T11:28:51.478599Z","iopub.status.idle":"2025-04-08T11:28:51.483474Z","shell.execute_reply.started":"2025-04-08T11:28:51.478572Z","shell.execute_reply":"2025-04-08T11:28:51.482679Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!cp /kaggle/input/torchvideo/x3d.py .\n!cp /kaggle/input/torchvideo/hgnet.py .\nfrom x3d import create_x3d\nfrom hgnet import hgnetv2_b5","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-08T10:26:41.701975Z","iopub.execute_input":"2025-04-08T10:26:41.702406Z","iopub.status.idle":"2025-04-08T10:26:42.014948Z","shell.execute_reply.started":"2025-04-08T10:26:41.702385Z","shell.execute_reply":"2025-04-08T10:26:42.013794Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"config = {\n    'batch_size': 32,\n    'num_worker': 4,\n    'data': '/kaggle/input/hms-harmful-brain-activity-classification/test.csv',\n    'weights_x3d': '/kaggle/input/hms-x3d',\n    'weights_doubleheadbutterfilter': '/kaggle/input/hms-doublehead-butter-filter',\n    'flip': True\n}\n\nkeys_to_update = [\n    'weights_x3d',\n    'weights_doubleheadbutterfilter',\n]\n\nfor key in keys_to_update:\n    config[key] = [os.path.join(config[key], fname) for fname in sorted(os.listdir(config[key]))]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-08T11:59:32.368630Z","iopub.execute_input":"2025-04-08T11:59:32.369000Z","iopub.status.idle":"2025-04-08T11:59:32.375049Z","shell.execute_reply.started":"2025-04-08T11:59:32.368972Z","shell.execute_reply":"2025-04-08T11:59:32.374152Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Preparation","metadata":{}},{"cell_type":"code","source":"class DataProcessor:\n    \"\"\"\n    Iterator for processing brain activity data, supporting EEG, spectrogram,\n    and mixed modalities. It handles data augmentation (e.g., flipping), filtering,\n    and channel differencing.\n    \"\"\"\n\n    def __init__(self,\n         dataframe: pd.DataFrame,\n         training_flag: bool = False,\n         shuffle: bool = False,\n         use_spec: bool = False,\n         use_eeg: bool = False,\n         use_mix: bool = False,\n         lower_cut: float = 0,\n         upper_cut: float = 20,\n         flip: bool = False,\n         use_mne_filter: bool = True,\n         use_18_lead: bool = False) -> None:\n        \"\"\"\n        Initialize the data iterator with configuration options.\n\n        Args:\n            dataframe (pd.DataFrame): Dataframe with metadata for each sample.\n            training_flag (bool): Set to True if used for training.\n            shuffle (bool): Set to True to shuffle the data.\n            use_spec (bool): Use spectrogram data if True.\n            use_eeg (bool): Use EEG data if True.\n            use_mix (bool): Use both EEG and spectrogram data if True.\n            lower_cut (float): Lower frequency cutoff for filtering.\n            upper_cut (float): Upper frequency cutoff for filtering.\n            flip (bool): If True, apply mirroring to the data.\n            use_mne_filter (bool): If True, use MNE filtering; otherwise use a Butterworth filter.\n            use_18_lead (bool): Whether to use all 18 EEG leads.\n        \"\"\"\n        self.flip_eeg: bool = flip\n        self.lower_cut: float = lower_cut\n        self.upper_cut: float = upper_cut\n        self.use_18_lead: bool = use_18_lead\n\n        print(self.lower_cut, self.upper_cut, 'with mne filter:', use_mne_filter, 'use 18 lead:', use_18_lead)\n\n        self.training_flag: bool = training_flag\n        self.shuffle: bool = shuffle\n        self.dataframe: pd.DataFrame = dataframe\n\n        # Mapping of brain activity classes to integer labels\n        activity_to_label = {'Seizure': 0, 'LPD': 1, 'GPD': 2, 'LRDA': 3, 'GRDA': 4, 'Other': 5}\n        self.target_mapping: dict = activity_to_label\n        self.target_mapping_inv: dict = {label: activity for activity, label in activity_to_label.items()}\n\n        # List of EEG channel names (including an EKG channel)\n        self.eeg_channel_names: List[str] = [\n            'Fp1', 'F3', 'C3', 'P3', 'F7', 'T3', 'T5', 'O1',\n            'Fz', 'Cz', 'Pz',\n            'Fp2', 'F4', 'C4', 'P4', 'F8', 'T4', 'T6', 'O2', 'EKG'\n        ]\n\n        # Define channel groups for differential computations\n        self.left_lateral: List[str] = ['Fp1', 'F7', 'T3', 'T5', 'O1']\n        self.right_lateral: List[str] = ['Fp2', 'F8', 'T4', 'T6', 'O2']\n        self.left_parietal: List[str] = ['Fp1', 'F3', 'C3', 'P3', 'O1']\n        self.right_parietal: List[str] = ['Fp2', 'F4', 'C4', 'P4', 'O2']\n        self.midline: List[str] = ['Fz', 'Cz', 'Pz']\n\n        # Map channel names to their indices for quick lookup\n        self.channel_index: dict = {name: idx for idx, name in enumerate(self.eeg_channel_names)}\n\n        self.use_eeg: bool = use_eeg\n        self.use_spec: bool = use_spec\n        self.use_mix: bool = use_mix\n        self.use_mne_filter: bool = use_mne_filter\n\n    def __getitem__(self, index: int) -> Union[np.ndarray, Tuple[np.ndarray, np.ndarray]]:\n        \"\"\"\n        Retrieve a processed data sample.\n\n        Args:\n            index (int): Index of the sample to retrieve.\n\n        Returns:\n            Union[np.ndarray, Tuple[np.ndarray, np.ndarray]]:\n                Processed EEG or spectrogram data. In mixed mode, returns a tuple (EEG, spectrogram).\n        \"\"\"\n        data_point = self.dataframe.iloc[index]\n        return self._process_single_data(data_point, self.training_flag)\n\n    def __len__(self) -> int:\n        \"\"\"Return the number of samples in the dataset.\"\"\"\n        return len(self.dataframe)\n\n    def compute_brain_leads(self, waves: np.ndarray) -> np.ndarray:\n        \"\"\"\n        Compute differential signals (brain leads) from raw EEG data.\n\n        Args:\n            waves (np.ndarray): Raw EEG data with shape (channels, time).\n\n        Returns:\n            np.ndarray: Concatenated differential brain lead signals.\n        \"\"\"\n        waves_copy = copy.deepcopy(waves)\n        # Groups of channels used for differential calculation\n        brain_lead_groups = [self.left_lateral, self.right_lateral, self.left_parietal, self.right_parietal]\n        differential_leads: List[np.ndarray] = []\n\n        for group in brain_lead_groups:\n            for i in range(len(group) - 1):\n                diff_signal = waves_copy[self.channel_index[group[i]]] - waves_copy[self.channel_index[group[i + 1]]]\n                differential_leads.append(diff_signal)\n\n        return np.stack(differential_leads, axis=0)\n\n    def mirror_spectrogram(self, spectrogram: np.ndarray) -> np.ndarray:\n        \"\"\"\n        Apply mirroring transformation to a spectrogram using a fixed index permutation.\n\n        Args:\n            spectrogram (np.ndarray): Spectrogram data.\n\n        Returns:\n            np.ndarray: Mirrored spectrogram.\n        \"\"\"\n        index_order = [1, 0, 3, 2]\n        return spectrogram[..., index_order]\n\n    def mirror_eeg(self, eeg_data: np.ndarray) -> np.ndarray:\n        \"\"\"\n        Mirror EEG data by swapping left-side channels with corresponding right-side channels.\n\n        Args:\n            eeg_data (np.ndarray): EEG data.\n\n        Returns:\n            np.ndarray: EEG data after applying mirror swap.\n        \"\"\"\n        # Define indices for left and right channels (based on your data ordering)\n        left_indices = [0, 1, 2, 3, 4, 5, 6, 7]\n        right_indices = [11, 12, 13, 14, 15, 16, 17, 18]\n        eeg_data[left_indices, ...], eeg_data[right_indices, ...] = eeg_data[right_indices, ...], eeg_data[left_indices, ...]\n        return eeg_data\n\n    def butter_bandpass(self, lowcut: float, highcut: float, fs: float, order: int = 5) -> Tuple[np.ndarray, np.ndarray]:\n        \"\"\"\n        Generate Butterworth bandpass filter coefficients.\n\n        Args:\n            lowcut (float): Lower frequency cutoff.\n            highcut (float): Upper frequency cutoff.\n            fs (float): Sampling frequency.\n            order (int): Filter order.\n\n        Returns:\n            Tuple[np.ndarray, np.ndarray]: Filter coefficients (b, a).\n        \"\"\"\n        return butter(order, [lowcut, highcut], fs=fs, btype=\"band\")\n\n    def butter_bandpass_filter(self, data: np.ndarray, lowcut: float, highcut: float, fs: float, order: int = 5) -> np.ndarray:\n        \"\"\"\n        Apply a Butterworth bandpass filter to the data.\n\n        Args:\n            data (np.ndarray): Data to filter.\n            lowcut (float): Lower frequency cutoff.\n            highcut (float): Upper frequency cutoff.\n            fs (float): Sampling frequency.\n            order (int): Filter order.\n\n        Returns:\n            np.ndarray: Filtered data.\n        \"\"\"\n        b, a = self.butter_bandpass(lowcut, highcut, fs, order=order)\n        return lfilter(b, a, data)\n\n    def load_eeg_data(self, data_point: pd.Series, is_training: bool, flip: bool = False) -> np.ndarray:\n        \"\"\"\n        Load and process EEG data from a single data record.\n\n        Args:\n            data_point (pd.Series): A row from the dataframe containing metadata.\n            is_training (bool): Flag indicating training mode.\n            flip (bool): If True, apply mirroring to the EEG data.\n\n        Returns:\n            np.ndarray: Processed EEG data.\n        \"\"\"\n        if is_training:\n            eeg_file_path = f\"/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/{data_point['eeg_id']}.parquet\"\n        else:\n            eeg_file_path = f\"/kaggle/input/hms-harmful-brain-activity-classification/test_eegs/{data_point['eeg_id']}.parquet\"\n        eeg_df = pd.read_parquet(eeg_file_path)\n\n        offset = 0\n        eeg_df = eeg_df.iloc[int(offset * 200):int(offset * 200) + 10000]\n        waves = eeg_df.values\n        waves = np.transpose(waves, axes=[1, 0])\n\n        # Handle NaN values for each channel\n        for channel in range(waves.shape[0]):\n            channel_mean = np.nanmean(waves[channel])\n            if np.isnan(waves[channel]).mean() < 1:\n                waves[channel] = np.nan_to_num(waves[channel], nan=channel_mean)\n            else:\n                waves[channel] = 0\n\n        if flip:\n            waves = self.mirror_eeg(waves)\n\n        # Compute differential leads and apply filtering\n        waves = self.compute_brain_leads(waves)\n        waves = np.array(waves, dtype=np.float64)\n        waves = np.clip(waves, -1024, 1024)\n        if self.use_mne_filter:\n            waves = mne.filter.filter_data(waves, 200, self.lower_cut, self.upper_cut, verbose=False)\n        else:\n            waves = self.butter_bandpass_filter(waves, 0.5, 20, 200, order=2)\n        return waves\n\n    def load_spectrogram(self, data_point: pd.Series, is_training: bool, flip: bool = False) -> np.ndarray:\n        \"\"\"\n        Load and process spectrogram data from a single data record.\n\n        Args:\n            data_point (pd.Series): A row from the dataframe containing metadata.\n            is_training (bool): Flag indicating training mode.\n            flip (bool): If True, apply mirroring to the spectrogram.\n\n        Returns:\n            np.ndarray: Processed spectrogram data.\n        \"\"\"\n        if is_training:\n            spec_file_path = f\"/kaggle/input/hms-harmful-brain-activity-classification/train_spectrograms/{data_point['spectrogram_id']}.parquet\"\n        else:\n            spec_file_path = f\"/kaggle/input/hms-harmful-brain-activity-classification/test_spectrograms/{data_point['spectrogram_id']}.parquet\"\n        spec_df = pd.read_parquet(spec_file_path)\n        spec_values = spec_df.values[:, 1:]  # Exclude the first column\n\n        spectrogram_images: List[np.ndarray] = []\n        row_start = 0\n\n        for region in range(4):\n            # Extract and process region-specific image\n            image = spec_values[row_start:row_start + 300, region * 100:(region + 1) * 100].T\n            image = np.clip(image, np.exp(-4), np.exp(8))\n            image = np.log(image)\n            image = np.nan_to_num(image, nan=0.0)\n            spectrogram_images.append(image)\n\n        stacked_images = np.stack(spectrogram_images, axis=-1)\n        if flip:\n            stacked_images = self.mirror_spectrogram(stacked_images)\n        processed_spec = np.transpose(stacked_images, [2, 0, 1])\n        return processed_spec\n\n    def load_mixed_data(self, data_point: pd.Series, is_training: bool, flip: bool = False) -> Tuple[np.ndarray, np.ndarray]:\n        \"\"\"\n        Load and process both EEG and spectrogram data for a single data record.\n\n        Args:\n            data_point (pd.Series): A row from the dataframe containing metadata.\n            is_training (bool): Flag indicating training mode.\n            flip (bool): If True, apply mirroring to the data.\n\n        Returns:\n            Tuple[np.ndarray, np.ndarray]: A tuple containing processed EEG data and spectrogram data.\n        \"\"\"\n        if is_training:\n            eeg_file_path = f\"/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/{data_point['eeg_id']}.parquet\"\n        else:\n            eeg_file_path = f\"/kaggle/input/hms-harmful-brain-activity-classification/test_eegs/{data_point['eeg_id']}.parquet\"\n        eeg_df = pd.read_parquet(eeg_file_path)\n        offset = 0\n        eeg_df = eeg_df.iloc[int(offset * 200):int(offset * 200) + 10000]\n        waves = eeg_df.values\n        waves = np.transpose(waves, axes=[1, 0])\n        for channel in range(waves.shape[0]):\n            channel_mean = np.nanmean(waves[channel])\n            if np.isnan(waves[channel]).mean() < 1:\n                waves[channel] = np.nan_to_num(waves[channel], nan=channel_mean)\n            else:\n                waves[channel] = 0\n        if flip:\n            waves = self.mirror_eeg(waves)\n        waves = self.compute_brain_leads(waves)\n        waves = np.array(waves, dtype=np.float64)\n        waves = np.clip(waves, -1024, 1024)\n        waves = mne.filter.filter_data(waves, 200, self.lower_cut, self.upper_cut, verbose=False)\n\n        # Load spectrogram data\n        row_start = 0\n        spec_file_path = f\"/kaggle/input/hms-harmful-brain-activity-classification/test_spectrograms/{data_point['spectrogram_id']}.parquet\"\n        spec_df = pd.read_parquet(spec_file_path)\n        spec_values = spec_df.values[:, 1:]\n        spectrogram_images: List[np.ndarray] = []\n        for region in range(4):\n            image = spec_values[row_start:row_start + 300, region * 100:(region + 1) * 100].T\n            image = np.clip(image, np.exp(-4), np.exp(8))\n            image = np.log(image)\n            image = np.nan_to_num(image, nan=0.0)\n            spectrogram_images.append(image)\n        stacked_images = np.stack(spectrogram_images, axis=-1)\n        if flip:\n            stacked_images = self.mirror_spectrogram(stacked_images)\n        processed_spec = np.transpose(stacked_images, [2, 0, 1])\n\n        return waves, processed_spec\n\n    def _process_single_data(self, data_point: pd.Series, is_training: bool) -> Union[np.ndarray, Tuple[np.ndarray, np.ndarray]]:\n        \"\"\"\n        Process a single data record into the desired modality (EEG, spectrogram, or mixed).\n\n        Args:\n            data_point (pd.Series): A row from the dataframe.\n            is_training (bool): Flag indicating training mode.\n\n        Returns:\n            Union[np.ndarray, Tuple[np.ndarray, np.ndarray]]:\n                Processed EEG or spectrogram data, or both in mixed mode.\n        \"\"\"\n        if self.use_eeg:\n            eeg_data = self.load_eeg_data(data_point, is_training, self.flip_eeg)\n            return eeg_data.astype(np.float32)\n        elif self.use_spec:\n            spec_data = self.load_spectrogram(data_point, is_training, self.flip_eeg)\n            return spec_data.astype(np.float32)\n        elif self.use_mix:\n            eeg_data, spec_data = self.load_mixed_data(data_point, is_training, self.flip_eeg)\n            return eeg_data.astype(np.float32), spec_data.astype(np.float32)\n        \n        # Default: return EEG data if no specific modality is selected\n        default_data = self.load_eeg_data(data_point, is_training, self.flip_eeg)\n        return default_data.astype(np.float32)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-08T10:26:42.082570Z","iopub.execute_input":"2025-04-08T10:26:42.082802Z","iopub.status.idle":"2025-04-08T10:26:42.108581Z","shell.execute_reply.started":"2025-04-08T10:26:42.082781Z","shell.execute_reply":"2025-04-08T10:26:42.107787Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Inference","metadata":{}},{"cell_type":"code","source":"class Transform50s(nn.Module):\n    \"\"\"\n    Applies a spectrogram transformation tailored for a 50s view of the EEG signal.\n\n    This transform computes a spectrogram from the input signal using:\n        - FFT size: 512\n        - Window lencgth: 128\n        - Hop length: 50\n        - Power: 1 (magnitude spectrogram)\n    It then clips the values, scales them down, and slices the frequency axis\n    to select a specific region.\n    \"\"\"\n    def __init__(self) -> None:\n        super().__init__()\n        self.wave_transform = torchaudio.transforms.Spectrogram(\n            n_fft=512,\n            win_length=128,\n            hop_length=50,\n            power=1\n        )\n\n    def forward(self, x: Tensor) -> Tensor:\n        \"\"\"\n        Forward pass of the Transform50s.\n\n        Args:\n            x (Tensor): Input tensor of shape (batch_size, time) or (batch_size, channels, time).\n\n        Returns:\n            Tensor: Processed spectrogram with shape (batch_size, channels, selected_freq, time).\n        \"\"\"\n        # Compute spectrogram.\n        image = self.wave_transform(x)\n        # Clip values and scale down.\n        image = torch.clip(image, min=0, max=10000) / 1000\n        # Get dimensions and slice along frequency (height) dimension.\n        n, c, h, w = image.size()\n        image = image[:, :, :int(20 / 100 * h + 10), :]\n        return image\n\n\nclass Transform10s(nn.Module):\n    \"\"\"\n    Applies a spectrogram transformation tailored for a 10s slice of the EEG signal.\n\n    This transform computes a spectrogram with:\n        - FFT size: 512\n        - Window length: 128\n        - Hop length: 10\n        - Power: 1\n    Similar to Transform50s, it clips and scales the output and then selects a portion\n    of the frequency axis.\n    \"\"\"\n    def __init__(self) -> None:\n        super().__init__()\n        self.wave_transform = torchaudio.transforms.Spectrogram(\n            n_fft=512,\n            win_length=128,\n            hop_length=10,\n            power=1\n        )\n\n    def forward(self, x: Tensor) -> Tensor:\n        \"\"\"\n        Forward pass of the Transform10s.\n\n        Args:\n            x (Tensor): Input tensor of shape (batch_size, time) or (batch_size, channels, time).\n\n        Returns:\n            Tensor: Processed spectrogram with shape (batch_size, channels, selected_freq, time).\n        \"\"\"\n        image = self.wave_transform(x)\n        image = torch.clip(image, min=0, max=10000) / 1000\n        n, c, h, w = image.size()\n        image = image[:, :, :int(20 / 100 * h + 10), :]\n        return image\n\n\nclass Modelx3d(nn.Module):\n    \"\"\"\n    Wrapper for an X3D model configured for processing video-like data.\n\n    The X3D model is created with:\n        - Input clip length: 16 frames\n        - Input crop size: 312\n        - Depth factor: 5.0\n\n    Certain layers in block 5 are replaced by identity modules to adjust the network's behavior.\n    \"\"\"\n    def __init__(self) -> None:\n        super().__init__()\n        model_name = \"x3d_l\"\n        self.net = create_x3d(\n            input_clip_length=16,\n            input_crop_size=312,\n            depth_factor=5.0,\n        )\n        # Modify block 5: remove dropout, projection, activation, and output pooling.\n        self.net.blocks[5].dropout = nn.Identity()\n        self.net.blocks[5].proj = nn.Identity()\n        self.net.blocks[5].activation = nn.Identity()\n        self.net.blocks[5].output_pool = nn.Identity()\n\n    def forward(self, x: Tensor) -> Tensor:\n        \"\"\"\n        Forward pass for the X3D model.\n\n        Args:\n            x (Tensor): Input tensor of shape compatible with X3D (typically video-like).\n\n        Returns:\n            Tensor: Feature representation produced by the X3D model.\n        \"\"\"\n        x = self.net(x)\n        return x\n\n\nclass Netx3d(nn.Module):\n    \"\"\"\n    Combines two spectrogram transforms with an X3D backbone for EEG classification.\n\n    The network processes EEG data in two temporal resolutions:\n        - A \"50s\" view using Transform50s.\n        - A \"10s\" view (a slice from the EEG) using Transform10s.\n\n    The outputs from these transforms are concatenated, converted to a 3-channel input,\n    and then passed through the X3D model. Finally, a fully connected layer (with softmax)\n    produces the class probabilities.\n    \"\"\"\n    def __init__(self, num_classes: int = 6) -> None:\n        \"\"\"\n        Initialize the Netx3d model.\n\n        Args:\n            num_classes (int): The number of output classes. Default is 6.\n        \"\"\"\n        super().__init__()\n        self.preprocess50s = Transform50s()\n        self.preprocess10s = Transform10s()\n        self.model = Modelx3d()\n        self.pool = nn.AdaptiveAvgPool3d(1)\n        self.fc = nn.Linear(2048, num_classes, bias=True)\n    \n    def forward(self, eeg: Tensor) -> Tensor:\n        \"\"\"\n        Forward pass of the Netx3d model.\n\n        Args:\n            eeg (Tensor): Input EEG data tensor. Expected shape is at least (batch_size, channels, time).\n                          The full EEG is used for the 50s transform, and a temporal slice is used for the 10s transform.\n\n        Returns:\n            Tensor: The output probabilities for each class, with shape (batch_size, num_classes).\n        \"\"\"\n        bs = eeg.size(0)\n\n        # For the 50s transform, use the entire EEG.\n        eeg_50s = eeg\n        # For the 10s transform, select a temporal slice (e.g., from time index 4000 to 6000).\n        eeg_10s = eeg[:, :, 4000:6000]\n        x_50 = self.preprocess50s(eeg_50s)\n        x_10 = self.preprocess10s(eeg_10s)\n        # Concatenate along the channel dimension.\n        x = torch.cat([x_10, x_50], dim=1)\n\n        # The resulting tensor x is then unsqueezed to add a new dimension.\n        x = torch.unsqueeze(x, dim=1)\n        # Replicate the tensor across three channels to match the expected input for the X3D model.\n        x = torch.cat([x, x, x], dim=1)\n            \n        # Pass through the X3D model.\n        x = self.model(x)\n        x = x.view(bs, -1)\n        \n        # Apply the fully connected layer.\n        x = self.fc(x)\n        # Apply softmax to produce probability distributions over classes.\n        x = torch.softmax(x, dim=-1)\n        return x\n\n    def forward_ablate(self, eeg: torch.Tensor, ablate: str = None) -> torch.Tensor:\n        \"\"\"\n        Forward pass that can \"ablate\" (zero-out) one of the branches.\n        \n        Args:\n            eeg (torch.Tensor): Input EEG data, shape (batch_size, channels, time).\n            ablate (str, optional): \n                - '10s' to ablate (zero out) the 10s branch,\n                - '50s' to ablate the 50s branch,\n                - None to use the full model.\n        \n        Returns:\n            torch.Tensor: The output probabilities.\n        \"\"\"\n        bs = eeg.size(0)\n        eeg_50s = eeg\n        eeg_10s = eeg[:, :, 4000:6000]\n        x_50 = self.preprocess50s(eeg_50s)\n        x_10 = self.preprocess10s(eeg_10s)\n        \n        # Ablate one branch if specified.\n        if ablate == '10s':\n            x_10 = torch.zeros_like(x_10)\n        elif ablate == '50s':\n            x_50 = torch.zeros_like(x_50)\n        \n        x = torch.cat([x_10, x_50], dim=1)\n        x = torch.unsqueeze(x, dim=1)\n        x = torch.cat([x, x, x], dim=1)\n        x = self.model(x)\n        x = x.view(bs, -1)\n        x = self.fc(x)\n        x = torch.softmax(x, dim=-1)\n        return x\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-08T10:26:42.109810Z","iopub.execute_input":"2025-04-08T10:26:42.110112Z","iopub.status.idle":"2025-04-08T10:26:42.126535Z","shell.execute_reply.started":"2025-04-08T10:26:42.110080Z","shell.execute_reply":"2025-04-08T10:26:42.125816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Transformdoubleheadbutterfilter(nn.Module):\n    \"\"\"\n    Transforms raw EEG data into a spectrogram representation using a Butterworth filtering approach.\n    \n    This module computes a spectrogram via the STFT using torchaudio.transforms.Spectrogram, applies a logarithmic scaling,\n    clips the values, and then slices the frequency axis to retain the lower frequencies (roughly 0-20 Hz).\n    \n    Expected input:\n        Tensor of shape (batch_size, time)\n    \n    Returns:\n        Tensor of shape (batch_size, channels, selected_freq, time_frames)\n    \"\"\"\n    def __init__(self) -> None:\n        super().__init__()\n        self.wave_transform = torchaudio.transforms.Spectrogram(\n            n_fft=512,\n            hop_length=50,\n            power=1\n        )\n\n    def forward(self, x: Tensor) -> Tensor:\n        \"\"\"\n        Forward pass for the spectrogram transformation.\n        \n        Args:\n            x (Tensor): Raw EEG signal of shape (batch_size, time).\n            \n        Returns:\n            Tensor: Processed spectrogram of shape (batch_size, channels, selected_freq, time_frames).\n        \"\"\"\n        bs: int = x.size(0)\n        image: Tensor = self.wave_transform(x)  # Compute spectrogram\n        image = torch.log10(image)              # Logarithmic scaling\n        image = torch.clip(image, min=0)         # Clip negative values\n        n, c, h, w = image.size()\n        # Retain only the lower frequencies (approximately 0-20 Hz)\n        image = image[:, :, :int(20 / 100 * h + 5), :]\n        return image\n\nclass Modeleegbutterfilter(nn.Module):\n    \"\"\"\n    Feature extraction branch for raw EEG data using EfficientNet_B5.\n    \n    This module reshapes a flattened EEG signal into a pseudo-image.\n    The raw EEG (assumed to be a flattened tensor of shape (batch_size, 160000))\n    is reshaped to (batch_size, 16, 1000, 10), then permuted and flattened into \n    (batch_size, 16*10, 1000). It then unsqueezes and replicates the single-channel \n    data into 3 channels to form a pseudo-RGB image for EfficientNet_B5.\n    \n    Expected input:\n        Tensor of shape (batch_size, 160000)\n    \n    Returns:\n        Tensor: Feature vector of shape (batch_size, feature_dim) extracted by EfficientNet_B5.\n    \"\"\"\n    def __init__(self) -> None:\n        super(Modeleegbutterfilter, self).__init__()\n        self.model = timm.create_model('efficientnet_b5', pretrained=False, in_chans=3)\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.fc = nn.Linear(2048, out_features=6, bias=True)\n\n    def extract_features(self, x: Tensor) -> Tensor:\n        \"\"\"\n        Extract features using EfficientNet_B5's forward_features method.\n        \n        Args:\n            x (Tensor): Input tensor of shape (batch_size, 3, H, W).\n            \n        Returns:\n            Tensor: Extracted feature map.\n        \"\"\"\n        x = self.model.forward_features(x)\n        return x\n\n    def forward(self, x: Tensor) -> Tensor:\n        \"\"\"\n        Forward pass for the raw EEG branch.\n        \n        Steps:\n          1. Reshape the flattened EEG vector from (batch_size, 160000) to (batch_size, 16, 1000, 10).\n          2. Permute dimensions to (batch_size, 16, 10, 1000) and flatten the channel and segment dimensions.\n          3. Unsqueeze to add a channel dimension and replicate to form a 3-channel image.\n          4. Extract features using EfficientNet_B5 and apply adaptive pooling.\n        \n        Args:\n            x (Tensor): Input tensor of shape (batch_size, 160000).\n        \n        Returns:\n            Tensor: Extracted features of shape (batch_size, feature_dim).\n        \"\"\"\n        bs: int = x.size(0)\n        reshaped_tensor: Tensor = x.view(bs, 16, 1000, 10)          # (bs, 16, 1000, 10)\n        reshaped_and_permuted_tensor: Tensor = reshaped_tensor.permute(0, 1, 3, 2)  # (bs, 16, 10, 1000)\n        reshaped_and_permuted_tensor = reshaped_and_permuted_tensor.reshape(bs, 16 * 10, 1000)  # (bs, 160, 1000)\n        x = torch.unsqueeze(reshaped_and_permuted_tensor, dim=1)       # (bs, 1, 160, 1000)\n        x = torch.cat([x, x, x], dim=1)                                # (bs, 3, 160, 1000)\n        bs = x.size(0)\n        x = self.extract_features(x)\n        x = self.pool(x)\n        x = x.view(bs, -1)\n        return x\n\nclass Modelspecbutterfilter(nn.Module):\n    \"\"\"\n    Feature extraction branch for EEG spectrogram data using an X3D backbone.\n    \n    This branch processes the spectrogram output from Transformdoubleheadbutterfilter.\n    The spectrogram is expanded to 3 channels and passed through an X3D model (with modifications\n    in block 5, where dropout, projection, activation, and output pooling are replaced with identity).\n    \n    Expected input:\n        Tensor of shape (batch_size, H, W) (i.e., the spectrogram).\n    \n    Returns:\n        Tensor: Flattened feature vector of shape (batch_size, feature_dim).\n    \"\"\"\n    def __init__(self, num_classes: int = 1) -> None:\n        super().__init__()\n        model_name: str = \"x3d_l\"\n        self.net = create_x3d(input_clip_length=16, input_crop_size=312, depth_factor=5.0)\n        # Replace specific components in block 5 with Identity to modify network behavior\n        self.net.blocks[5].dropout = nn.Identity()\n        self.net.blocks[5].proj = nn.Identity()\n        self.net.blocks[5].activation = nn.Identity()\n        self.net.blocks[5].output_pool = nn.Identity()\n        self.avg = nn.AdaptiveAvgPool2d(1)\n\n    def forward(self, x: Tensor) -> Tensor:\n        \"\"\"\n        Forward pass for the spectrogram branch.\n        \n        Args:\n            x (Tensor): Input spectrogram of shape (batch_size, H, W).\n            \n        Returns:\n            Tensor: Flattened feature vector.\n        \"\"\"\n        # Add channel dimension and replicate to 3 channels: (bs, 3, H, W)\n        x = torch.unsqueeze(x, dim=1)\n        x = torch.cat([x, x, x], dim=1)\n        x = self.net(x)\n        bs: int = x.size(0)\n        x = x.view(bs, -1)\n        return x\n\nclass Netdoubleheadbutterfilter(nn.Module):\n    \"\"\"\n    Double-head EEG classification model that fuses features from both the raw EEG branch and the\n    spectrogram branch (using a Butterworth filter variant).\n    \n    The raw EEG branch (Modeleegbutterfilter) processes the flattened EEG data to extract features,\n    while the spectrogram branch (Modelspecbutterfilter) processes the log-scaled, clipped spectrogram.\n    Their outputs are concatenated and passed through a dropout layer and a fully connected layer to\n    generate class probabilities over 6 classes.\n    \n    Expected input:\n        Tensor of shape (batch_size, time) where time corresponds to the raw EEG signal length.\n    \n    Returns:\n        Tensor: Class probabilities of shape (batch_size, 6).\n    \"\"\"\n    def __init__(self) -> None:\n        super(Netdoubleheadbutterfilter, self).__init__()\n        self.transform: nn.Module = Transformdoubleheadbutterfilter()  # Converts raw EEG to spectrogram\n        self.model_wave: nn.Module = Modeleegbutterfilter()             # Extracts features from raw EEG\n        self.model_spec: nn.Module = Modelspecbutterfilter()            # Extracts features from the spectrogram\n        self.pool: nn.Module = nn.AdaptiveAvgPool3d(1)\n        self.fc: nn.Linear = nn.Linear(2048 * 2, 6, bias=True)\n        self.droup: nn.Dropout = nn.Dropout(0.5)\n\n    def forward(self, eeg: Tensor) -> Tensor:\n        \"\"\"\n        Forward pass for the double-head model.\n        \n        Steps:\n          1. Convert the raw EEG signal into a spectrogram using Transformdoubleheadbutterfilter.\n          2. Extract features from the raw EEG branch (Modeleegbutterfilter).\n          3. Extract features from the spectrogram branch (Modelspecbutterfilter).\n          4. Concatenate the two feature vectors.\n          5. Apply dropout and a fully connected layer.\n          6. Apply softmax activation to produce class probabilities.\n        \n        Args:\n            eeg (Tensor): Raw EEG data of shape (batch_size, time).\n            \n        Returns:\n            Tensor: Class probability tensor of shape (batch_size, 6).\n        \"\"\"\n        bs: int = eeg.size(0)\n        # Generate spectrogram from raw EEG using the transform branch\n        eeg_spec: Tensor = self.transform(eeg)\n        # Extract features from raw EEG\n        x: Tensor = self.model_wave(eeg)\n        # Extract features from spectrogram\n        y: Tensor = self.model_spec(eeg_spec)\n        # Concatenate features from both branches\n        x = torch.cat([x, y], dim=1)\n        # Apply dropout for regularization\n        x = self.droup(x)\n        # Fully connected layer to map features to 6 classes\n        x = self.fc(x)\n        # Softmax activation to output probabilities\n        x = torch.softmax(x, dim=-1)\n        return x\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-08T10:26:42.127431Z","iopub.execute_input":"2025-04-08T10:26:42.127632Z","iopub.status.idle":"2025-04-08T10:26:42.142501Z","shell.execute_reply.started":"2025-04-08T10:26:42.127615Z","shell.execute_reply":"2025-04-08T10:26:42.141884Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def inference_function(\n    test_loader: DataLoader,\n    model: torch.nn.Module,\n    device: torch.device,\n    double_input: bool = False\n) -> dict[str, np.ndarray]:\n    \"\"\"\n    Run inference on the provided DataLoader using the given model.\n    \n    Args:\n        test_loader (DataLoader): DataLoader that provides test batches.\n        model (torch.nn.Module): The model to run inference on.\n        device (torch.device): Device (CPU or GPU) on which inference is performed.\n        double_input (bool): If True, expects each batch to be a tuple (wave, spec).\n                             Otherwise, expects a single tensor.\n                             \n    Returns:\n        Dict[str, np.ndarray]: Dictionary with key \"predictions\" containing a numpy array\n                               of concatenated predictions.\n    \"\"\"\n    model.eval()\n    preds = []\n    \n    with tqdm(test_loader, unit=\"test_batch\", desc=\"Inference\") as t_loader:\n        for batch in t_loader:\n            if double_input:\n                wave, spec = batch\n                wave = wave.to(device)\n                spec = spec.to(device)\n                with torch.no_grad():\n                    y_preds = model(wave, spec)\n            else:\n                batch = batch.to(device)\n                with torch.no_grad():\n                    y_preds = model(batch)\n            preds.append(y_preds.cpu().numpy())\n    \n    concatenated_preds = np.concatenate(preds, axis=0)\n    return {\"predictions\": concatenated_preds}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-08T11:47:12.051510Z","iopub.execute_input":"2025-04-08T11:47:12.051972Z","iopub.status.idle":"2025-04-08T11:47:12.058289Z","shell.execute_reply.started":"2025-04-08T11:47:12.051937Z","shell.execute_reply":"2025-04-08T11:47:12.057304Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_weight_x3d(data_frame: pd.DataFrame) -> np.ndarray:\n    \"\"\"\n    Run inference using the Netx3d model with multiple weights on the given data.\n    \n    This function iterates over a list of model weight paths provided in the global config\n    under 'weights_x3d'. For each weight, it:\n    \n      1. Creates a DataProcessor instance (without flipping) and a corresponding DataLoader.\n      2. Loads the model weights into a Netx3d model, runs inference, and collects predictions.\n      3. Repeats the above steps with flipping enabled.\n    \n    Finally, it averages the predictions from all runs.\n    \n    Args:\n        data_frame (pd.DataFrame): DataFrame containing metadata for each EEG sample.\n    \n    Returns:\n        np.ndarray: Averaged predictions from the ensemble of models.\n    \"\"\"\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    print(\"Inference with weights_x3d\")\n    all_predictions = []\n    \n    # Iterate over each model weight in the global config\n    for model_weight in config['weights_x3d']:\n        # --- Inference without flipping ---\n        test_dataset = DataProcessor(\n            data_frame,\n            training_flag=False,\n            shuffle=False,\n            use_eeg=True,\n            lower_cut=0.5,\n            upper_cut=20,\n            use_mne_filter=True\n        )\n        test_loader = DataLoader(\n            test_dataset,\n            batch_size=config['batch_size'],\n            num_workers=config['num_worker'],\n            shuffle=False\n        )\n        model = Netx3d()\n        state_dict = torch.load(model_weight, map_location=device)\n        model.load_state_dict(state_dict, strict=True)\n        model.to(device)\n        pred_dict = inference_function(test_loader, model, device)\n        all_predictions.append(pred_dict[\"predictions\"])\n        \n        # --- Inference with flipping enabled ---\n        test_dataset = DataProcessor(\n            data_frame,\n            training_flag=False,\n            shuffle=False,\n            use_eeg=True,\n            flip=True,\n            lower_cut=0.5,\n            upper_cut=20,\n            use_mne_filter=True\n        )\n        test_loader = DataLoader(\n            test_dataset,\n            batch_size=config['batch_size'],\n            num_workers=config['num_worker'],\n            shuffle=False\n        )\n        model = Netx3d()\n        state_dict = torch.load(model_weight, map_location=device)\n        model.load_state_dict(state_dict, strict=True)\n        model.to(device)\n        pred_dict = inference_function(test_loader, model, device)\n        all_predictions.append(pred_dict[\"predictions\"])\n        \n        torch.cuda.empty_cache()\n        gc.collect()\n    \n    # Average predictions across all runs\n    all_predictions = np.array(all_predictions)\n    averaged_predictions = np.mean(all_predictions, axis=0)\n    return averaged_predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-08T10:26:42.160448Z","iopub.execute_input":"2025-04-08T10:26:42.160701Z","iopub.status.idle":"2025-04-08T10:26:42.176189Z","shell.execute_reply.started":"2025-04-08T10:26:42.160665Z","shell.execute_reply":"2025-04-08T10:26:42.175454Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_weight_double_headbutterfilter(data_frame: pd.DataFrame) -> np.ndarray:\n    \"\"\"\n    Run inference using the Netdoubleheadbutterfilter model with multiple weight files on the given data.\n    \n    This function iterates over each model weight path provided in the global configuration\n    under CFG['weights_doubleheadbutterfilter']. For each weight, it performs inference twice:\n      1. Once with the default configuration (without flipping).\n      2. Once with flipping enabled.\n    \n    In each case, an AlaskaDataIter instance is created with the appropriate parameters, a DataLoader\n    is built, the model weights are loaded into a new Netdoubleheadbutterfilter instance, and inference\n    is performed via the inference_function. The predictions from both runs for all weight files are collected\n    and then averaged to produce the final output.\n    \n    Args:\n        data_frame (pd.DataFrame): DataFrame containing metadata for each EEG sample.\n    \n    Returns:\n        np.ndarray: Averaged predictions from the ensemble of model runs as a NumPy array.\n    \"\"\"\n    device: torch.device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    print(\"Inference with weights_doubleheadbutterfilter\")\n    \n    predictions: List[np.ndarray] = []\n    \n    # Iterate over each model weight in the configuration\n    for model_weight in config['weights_doubleheadbutterfilter']:\n        # --- Inference without flipping ---\n        test_dataset = DataProcessor(\n            data_frame,\n            training_flag=False,\n            shuffle=False,\n            use_eeg=True,\n            lower_cut=0.5,\n            upper_cut=20,\n            use_mne_filter=False\n        )\n        test_loader: DataLoader = DataLoader(\n            test_dataset,\n            batch_size=config['batch_size'] // 2,\n            num_workers=config['num_worker'],\n            shuffle=False\n        )\n        \n        model = Netdoubleheadbutterfilter()\n        state_dict = torch.load(model_weight, map_location=device)\n        model.load_state_dict(state_dict, strict=True)\n        model.to(device)\n        pred_dict: Dict[str, np.ndarray] = inference_function(test_loader, model, device)\n        predictions.append(pred_dict[\"predictions\"])\n        \n        # --- Inference with flipping enabled ---\n        test_dataset = DataProcessor(\n            data_frame,\n            training_flag=False,\n            shuffle=False,\n            flip=True,\n            use_eeg=True,\n            lower_cut=0.5,\n            upper_cut=20,\n            use_mne_filter=False\n        )\n        test_loader = DataLoader(\n            test_dataset,\n            batch_size=config['batch_size'] // 2,\n            num_workers=config['num_worker'],\n            shuffle=False\n        )\n        \n        model = Netdoubleheadbutterfilter()\n        state_dict = torch.load(model_weight, map_location=device)\n        model.load_state_dict(state_dict, strict=True)\n        model.to(device)\n        pred_dict = inference_function(test_loader, model, device)\n        predictions.append(pred_dict[\"predictions\"])\n        \n        # Clear GPU memory after processing each weight\n        torch.cuda.empty_cache()\n        gc.collect()\n    \n    # Convert the list of predictions to a NumPy array and average across all runs\n    predictions_np: np.ndarray = np.array(predictions)\n    averaged_predictions: np.ndarray = np.mean(predictions_np, axis=0)\n    \n    return averaged_predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-08T10:26:49.382961Z","iopub.execute_input":"2025-04-08T10:26:49.383242Z","iopub.status.idle":"2025-04-08T10:26:49.390695Z","shell.execute_reply.started":"2025-04-08T10:26:49.383223Z","shell.execute_reply":"2025-04-08T10:26:49.389841Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Ensemble Aggregation","metadata":{}},{"cell_type":"code","source":"# Initialize a list to collect predictions from different ensemble branches.\npredictions = []\n\n# -------------------------------\n# Inference using the X3D branch:\n# -------------------------------\nprediction_x3d = run_weight_x3d(test_df)\npredictions.append(prediction_x3d)\n\n# --------------------------------------------------------------\n# Inference using the Double-Head Butter Filter branch:\n# --------------------------------------------------------------\nprediction_double_net = run_weight_double_headbutterfilter(test_df)\npredictions.append(prediction_double_net)\n\n# ---------------------------------------\n# Ensemble Aggregation:\n# ---------------------------------------\nweights_by_score = np.array([0.2, 0.8])\n\n# Combine the predictions from both models using a weighted sum.\nfinal_predictions = predictions[0] * weights_by_score[0] + predictions[1] * weights_by_score[1]\n\n# Output the final ensemble predictions.\nfinal_predictions\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Prepare Submission\nTARGETS = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\nsub = pd.DataFrame({'eeg_id': test_df.eeg_id.values})\nsub[TARGETS] = predictions\nsub.to_csv('submission.csv',index=False)\nprint(f'Submissionn shape: {sub.shape}')\nsub.head()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install zennit","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-08T10:26:57.630343Z","iopub.execute_input":"2025-04-08T10:26:57.630701Z","iopub.status.idle":"2025-04-08T10:27:01.842158Z","shell.execute_reply.started":"2025-04-08T10:26:57.630658Z","shell.execute_reply":"2025-04-08T10:27:01.841018Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## XAI Analysis","metadata":{}},{"cell_type":"code","source":"from zennit.composites import EpsilonGammaBox\nfrom zennit.canonizers import SequentialMergeBatchNorm\nfrom zennit.attribution import Gradient","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-08T11:29:02.627067Z","iopub.execute_input":"2025-04-08T11:29:02.627348Z","iopub.status.idle":"2025-04-08T11:29:02.631051Z","shell.execute_reply.started":"2025-04-08T11:29:02.627327Z","shell.execute_reply":"2025-04-08T11:29:02.630249Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# Utility Functions\n# =============================================================================\n\ndef load_data_and_model(model, model_weight, device):\n    \"\"\"\n    Loads training and test data, sets up the model with weights, and moves it to the specified device.\n    \n    Args:\n        model (nn.Module): Model instance (e.g., Netx3d).\n        model_weight (str): Path to the model weights file.\n        device (torch.device): Device on which to load the model.\n        config (dict): Dictionary containing data paths.\n        \n    Returns:\n        tuple: (train_df, test_df, model)\n    \"\"\"\n    # Load data.\n    train_path = '/kaggle/input/hms-harmful-brain-activity-classification/train.csv'\n    train_df = pd.read_csv(train_path)\n    test_df = pd.read_csv(config['data'])\n    \n    # Load weights.\n    state_dict = torch.load(model_weight, map_location=device)\n    model.load_state_dict(state_dict, strict=True)\n    model.to(device)\n    model.eval()\n    \n    return train_df, test_df, model\n\ndef prepare_sample(test_df, sample_index, device):\n    \"\"\"\n    Creates a test dataset instance, selects a sample, and converts it to a torch.Tensor.\n    \n    Args:\n        test_df (pd.DataFrame): DataFrame with test metadata.\n        sample_index (int): Index of the sample to explain.\n        device (torch.device): Device for the tensor.\n        \n    Returns:\n        torch.Tensor: Input tensor of shape (1, channels, time)\n    \"\"\"\n    test_dataset = DataProcessor(\n        test_df,\n        training_flag=False,\n        shuffle=False,\n        use_eeg=True,\n        lower_cut=0.5,\n        upper_cut=20,\n        use_mne_filter=False\n    )\n    \n    sample_data = test_dataset[sample_index]\n    if isinstance(sample_data, tuple):\n        sample_data = sample_data[0]\n    \n    # Convert sample to tensor and add a batch dimension.\n    input_tensor = torch.tensor(sample_data).unsqueeze(0).to(device)\n    print(\"Input tensor shape:\", input_tensor.shape)\n    return input_tensor\n\ndef compute_lrp_attributions(model, input_tensor, composite):\n    \"\"\"\n    Computes relevance scores on the input using Gradient-based LRP.\n    \n    Args:\n        model (nn.Module): The model on which to compute attributions.\n        input_tensor (torch.Tensor): Model input with shape (1, channels, time).\n        composite: Composite rule for Zennit (e.g., EpsilonGammaBox).\n        \n    Returns:\n        tuple: (output, relevance)\n            - output: Model's output tensor.\n            - relevance: Relevance tensor with the same shape as the input.\n    \"\"\"\n    with Gradient(model=model, composite=composite) as attributor:\n        output = model(input_tensor)\n        target_class = output.argmax(dim=1)\n        one_hot_target = torch.zeros_like(output)\n        one_hot_target.scatter_(1, target_class.unsqueeze(1), 1.0)\n        # Compute relevance.\n        _, relevance = attributor(input_tensor, one_hot_target)\n    return output, relevance\n\ndef visualize_relevance_overlay(input_np, relevance_np, channel_idx):\n    \"\"\"\n    Overlays the relevance map on the raw EEG signal for a specified channel.\n    \n    Args:\n        input_np (np.ndarray): Raw EEG input, shape (channels, time).\n        relevance_np (np.ndarray): Relevance map, same shape as input_np.\n        channel_idx (int): Index of the channel to visualize.\n    \"\"\"\n    time_axis = np.arange(input_np.shape[1])\n    plt.figure(figsize=(12, 4))\n    plt.plot(time_axis, input_np[channel_idx], color='black', label='Raw EEG')\n    plt.imshow(relevance_np[channel_idx][np.newaxis, :],\n               aspect='auto', cmap='seismic', alpha=0.5,\n               extent=[0, input_np.shape[1], input_np[channel_idx].min(), input_np[channel_idx].max()])\n    plt.xlabel('Time')\n    plt.ylabel('Amplitude')\n    plt.title(f'Channel {channel_idx} - Raw EEG with LRP Relevance Overlay')\n    plt.legend()\n    plt.show()\n\ndef plot_avg_relevance(avg_relevance_per_channel):\n    \"\"\"\n    Plots a bar chart of the average relevance per channel.\n    \n    Args:\n        avg_relevance_per_channel (np.ndarray): Array with average relevance for each channel.\n    \"\"\"\n    plt.figure(figsize=(8, 4))\n    plt.bar(np.arange(len(avg_relevance_per_channel)), avg_relevance_per_channel)\n    plt.xlabel('Channel Index')\n    plt.ylabel('Average Relevance')\n    plt.title('Average LRP Relevance per Channel')\n    plt.show()\n\ndef compute_class_relevance(model, input_tensor, composite, num_classes):\n    \"\"\"\n    Computes LRP relevance scores for each class and aggregates mean relevance per channel.\n    \n    Args:\n        model (nn.Module): The model on which to compute attributions.\n        input_tensor (torch.Tensor): Input with shape (1, channels, time).\n        composite: Composite rule for LRP (e.g., EpsilonGammaBox).\n        num_classes (int): Total number of classes.\n        \n    Returns:\n        dict: Mapping class index -> numpy array of mean relevance per channel.\n    \"\"\"\n    class_relevance = {}\n    # Get output for reference.\n    with Gradient(model=model, composite=composite) as attributor:\n        output = model(input_tensor)\n    for cls in range(num_classes):\n        one_hot = torch.zeros_like(output)\n        one_hot[:, cls] = 1.0\n        with Gradient(model=model, composite=composite) as attributor_cls:\n            _, rel_cls = attributor_cls(input_tensor, one_hot)\n        rel_cls_np = rel_cls.squeeze(0).cpu().detach().numpy()\n        class_relevance[cls] = np.mean(rel_cls_np, axis=1)  # Average over time.\n    return class_relevance\n\ndef plot_class_relevance_grouped(class_relevance, num_channels):\n    \"\"\"\n    Plots a grouped bar chart of mean LRP relevance per channel across classes.\n    \n    Args:\n        class_relevance (dict): Mapping class index -> per-channel relevance (numpy array).\n        num_channels (int): Number of EEG channels.\n    \"\"\"\n    channels = np.arange(num_channels)\n    plt.figure(figsize=(10, 6))\n    width = 0.1\n    num_classes = len(class_relevance)\n    for cls in range(num_classes):\n        plt.bar(channels + cls * width, class_relevance[cls],\n                width=width, label=f'Class {cls}')\n    plt.xlabel('Channel Index')\n    plt.ylabel('Mean Relevance')\n    plt.title('Mean LRP Relevance per Channel across Classes')\n    plt.legend()\n    plt.show()\n\ndef compute_avg_class_relevance_over_dataset(model, dataset, composite, num_classes, device, sample_limit=None):\n    \"\"\"\n    Iterates over a dataset and computes the average LRP relevance per channel for each class.\n    \n    Args:\n        model (nn.Module): Model on which to compute LRP.\n        dataset (iterable): Dataset instance (e.g., from DataProcessor) returning raw EEG samples.\n        composite: Composite rule for LRP (e.g., EpsilonGammaBox).\n        num_classes (int): Total number of classes.\n        device (torch.device): Device for computation.\n        sample_limit (int, optional): Optional limit on the number of samples to process.\n        \n    Returns:\n        dict: Mapping class index -> average per-channel relevance (numpy array).\n    \"\"\"\n    sums = {cls: None for cls in range(num_classes)}\n    count = 0\n    for idx, sample in enumerate(dataset):\n        if sample_limit is not None and idx >= sample_limit:\n            break\n        if isinstance(sample, tuple):\n            sample = sample[0]\n        input_tensor = torch.tensor(sample).unsqueeze(0).to(device)\n        with Gradient(model=model, composite=composite) as attributor:\n            output = model(input_tensor)\n            for cls in range(num_classes):\n                one_hot = torch.zeros_like(output)\n                one_hot[:, cls] = 1.0\n                _, rel_cls = attributor(input_tensor, one_hot)\n                rel_cls_np = rel_cls.squeeze(0).cpu().detach().numpy()\n                sample_rel = np.mean(rel_cls_np, axis=1)  # Mean over time.\n                if sums[cls] is None:\n                    sums[cls] = sample_rel\n                else:\n                    sums[cls] += sample_rel\n        count += 1\n        if (idx + 1) % 50 == 0:\n            print(f\"Processed {idx + 1} samples...\")\n    \n    averages = {cls: sums[cls] / count for cls in range(num_classes)}\n    return averages\n\ndef plot_dataset_class_relevance(avg_class_relevance, num_channels):\n    \"\"\"\n    Plots grouped bar charts of average LRP relevance per channel across classes computed over a dataset.\n    \n    Args:\n        avg_class_relevance (dict): Mapping class index -> average per-channel relevance.\n        num_channels (int): Number of EEG channels.\n    \"\"\"\n    channels = np.arange(num_channels)\n    plt.figure(figsize=(10, 6))\n    width = 0.1\n    num_classes = len(avg_class_relevance)\n    for cls in range(num_classes):\n        plt.bar(channels + cls * width, avg_class_relevance[cls],\n                width=width, label=f'Class {cls}')\n    plt.xlabel('Channel Index')\n    plt.ylabel('Average Relevance (Dataset)')\n    plt.title('Average LRP Relevance per Channel across Classes (Dataset Aggregation)')\n    plt.legend()\n    plt.show()\n\n# =============================================================================\n# Main Workflow\n# =============================================================================\n\ndef main():\n    # Set up device.\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    \n    # Load data and model.\n    model_weight = config['weights_x3d'][0]\n    model = Netx3d()  \n    train_df, test_df, model = load_data_and_model(model, model_weight, device)\n    \n    # Prepare a test sample and compute single-sample LRP.\n    sample_index = 0\n    input_tensor = prepare_sample(test_df, sample_index, device)\n    composite = EpsilonGammaBox(low=-3., high=3., canonizers=[SequentialMergeBatchNorm()])\n    output, relevance = compute_lrp_attributions(model, input_tensor, composite)\n    \n    # Convert tensors to numpy arrays.\n    relevance_np = relevance.squeeze(0).cpu().detach().numpy()  # (channels, time)\n    input_np = input_tensor.squeeze(0).cpu().detach().numpy()     # (channels, time)\n    \n    # Visualize raw EEG overlay with relevance for a specified channel.\n    for channel_index in range(len(relevance_np)):\n        visualize_relevance_overlay(input_np, relevance_np, channel_index)\n    \n    # Plot average relevance per channel for the single sample.\n    avg_relevance_per_channel = np.mean(relevance_np, axis=1)\n    plot_avg_relevance(avg_relevance_per_channel)\n    \n    # Compute and plot per-class relevance for the single sample.\n    num_classes = output.shape[1]\n    class_relevance = compute_class_relevance(model, input_tensor, composite, num_classes)\n    num_channels = input_np.shape[0]\n    plot_class_relevance_grouped(class_relevance, num_channels)\n    \n    # =============================================================================\n    # Dataset-Aggregated Analysis\n    # =============================================================================\n    sample_limit = 200  \n    \n    # Create a dataset instance from the training DataFrame.\n    train_dataset = DataProcessor(\n        test_df,\n        training_flag=False,\n        shuffle=False,\n        use_eeg=True,\n        lower_cut=0.5,\n        upper_cut=20,\n        use_mne_filter=False\n    )\n    \n    avg_class_relevance = compute_avg_class_relevance_over_dataset(model, train_dataset, composite, num_classes, device, sample_limit)\n    plot_dataset_class_relevance(avg_class_relevance, num_channels)\n    \n    print(\"XAI analysis and dataset-aggregated relevance complete.\")\n\nif __name__ == \"__main__\":\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-08T12:06:33.019323Z","iopub.execute_input":"2025-04-08T12:06:33.019646Z","iopub.status.idle":"2025-04-08T12:06:49.780125Z","shell.execute_reply.started":"2025-04-08T12:06:33.019617Z","shell.execute_reply":"2025-04-08T12:06:49.779245Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ------------------------------\n# Model Wrapper for Logits\n# ------------------------------\n\nclass Netx3dLogits(Netx3d):\n    \"\"\"\n    A wrapper for Netx3d that returns logits instead of probabilities.\n    \"\"\"\n    def forward(self, eeg: torch.Tensor) -> torch.Tensor:\n        bs = eeg.size(0)\n        # Process full EEG for 50s branch and a slice for 10s branch.\n        eeg_50s = eeg\n        eeg_10s = eeg[:, :, 4000:6000]\n        x_50 = self.preprocess50s(eeg_50s)\n        x_10 = self.preprocess10s(eeg_10s)\n        x = torch.cat([x_10, x_50], dim=1)\n        x = torch.unsqueeze(x, dim=1)\n        x = torch.cat([x, x, x], dim=1)\n        x = self.model(x)\n        x = x.view(bs, -1)\n        logits = self.fc(x)  # do not apply softmax\n        return logits\n\n# ------------------------------\n# Counterfactual Generation Function\n# ------------------------------\n\ndef generate_counterfactual(model, original_input, target_class, \n                            num_iterations=100, learning_rate=0.01,\n                            regularization_weight=0.1):\n    \"\"\"\n    Generates a counterfactual input that forces the model to predict the target class.\n    \n    Optimization minimizes a loss combining cross-entropy (to match the target) and an L2 penalty\n    (to keep the perturbation small).\n    \n    Args:\n        model (nn.Module): Model that outputs logits.\n        original_input (torch.Tensor): Input of shape (1, channels, time).\n        target_class (int): Target class index.\n        num_iterations (int): Maximum steps for optimization.\n        learning_rate (float): Learning rate.\n        regularization_weight (float): Weight for the perturbation loss.\n        \n    Returns:\n        torch.Tensor: Counterfactual input (same shape as original_input).\n    \"\"\"\n    cf_input = original_input.clone().detach().requires_grad_(True)\n    optimizer = torch.optim.Adam([cf_input], lr=learning_rate)\n    \n    # Use a fixed target tensor.\n    target = torch.tensor([target_class], device=original_input.device)\n    \n    for i in range(num_iterations):\n        logits = model(cf_input)\n        classification_loss = nn.functional.cross_entropy(logits, target)\n        perturbation_loss = torch.norm(cf_input - original_input)\n        loss = classification_loss + regularization_weight * perturbation_loss\n        \n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        \n        # Clamp values to remain within valid EEG ranges (adjust as needed).\n        with torch.no_grad():\n            cf_input.clamp_(min=original_input.min(), max=original_input.max())\n        \n        if i % 10 == 0:\n            pred_class = logits.argmax(dim=1).item()\n            print(f\"Iteration {i}: Loss = {loss.item():.4f}, Predicted = {pred_class}\")\n            # Optionally, break early if the prediction matches the target.\n            if pred_class == target_class:\n                print(f\"Counterfactual achieved target class: {target_class}.\")\n                break\n    \n    return cf_input.detach()\n\n# ------------------------------\n# Visualization Functions for Counterfactuals\n# ------------------------------\n\ndef plot_best_counterfactuals_side_by_side(original_np, counterfactuals, predictions, channels_to_plot):\n    \"\"\"\n    For each channel in channels_to_plot, select the counterfactual signal that has the largest mean absolute\n    difference from the original signal, and plot the original signal and that counterfactual side by side.\n    \n    Args:\n        original_np (np.ndarray): Original EEG signal of shape (channels, time).\n        counterfactuals (dict): Mapping from target class (int) to counterfactual numpy array (channels, time).\n        predictions (dict): Mapping from target class (int) to final predicted class (int) for the counterfactual.\n        channels_to_plot (list): List of channel indices to plot.\n    \"\"\"\n    time_axis = np.arange(original_np.shape[1])\n    num_channels = len(channels_to_plot)\n    \n    # Create a subplot with two columns: one for original and one for counterfactual\n    fig, axs = plt.subplots(num_channels, 2, figsize=(16, num_channels * 4), sharex=True)\n    # Ensure axs is 2D even for a single channel.\n    if num_channels == 1:\n        axs = np.expand_dims(axs, axis=0)\n    \n    for i, ch in enumerate(channels_to_plot):\n        original_channel = original_np[ch]\n        best_diff = -np.inf\n        best_cf = None\n        best_cls = None\n        \n        # Find the counterfactual with the highest mean absolute difference for this channel\n        for cls, cf_np in counterfactuals.items():\n            cf_channel = cf_np[ch]\n            diff = np.mean(np.abs(cf_channel - original_channel))\n            if diff > best_diff:\n                best_diff = diff\n                best_cf = cf_channel\n                best_cls = cls\n        \n        # Left subplot: plot original signal\n        axs[i, 0].plot(time_axis, original_channel, color='black', linewidth=2)\n        axs[i, 0].set_title(f'Channel {ch} Original')\n        axs[i, 0].set_ylabel('Amplitude')\n        axs[i, 0].grid(True)\n        \n        # Right subplot: plot best counterfactual signal\n        axs[i, 1].plot(time_axis, best_cf, color='red', linestyle='--', linewidth=2)\n        axs[i, 1].set_title(f'Channel {ch} CF Target {best_cls} (Pred {predictions[best_cls]}), Diff={best_diff:.4f}')\n        axs[i, 1].grid(True)\n    \n    # Set x-axis label on the bottom row of subplots.\n    for ax in axs[-1, :]:\n        ax.set_xlabel('Time')\n    \n    plt.suptitle('Original vs Best Counterfactual (Side-by-Side) per Channel', fontsize=18)\n    plt.tight_layout(rect=[0, 0, 1, 0.96])\n    plt.show()\n\n\n# ------------------------------\n# Main Workflow\n# ------------------------------\n\ndef main():\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    \n    # Load data and original model.\n    model_weight = config['weights_x3d'][0]\n    base_model = Netx3d()\n    _, test_df, base_model = load_data_and_model(base_model, model_weight, device)\n    \n    # Create the logits wrapper model.\n    logits_model = Netx3dLogits()\n    logits_model.load_state_dict(base_model.state_dict(), strict=True)\n    logits_model.to(device)\n    logits_model.eval()\n    \n    # Prepare a sample from the test data.\n    sample_index = 0  \n    original_input = prepare_sample(test_df, sample_index, device)\n    \n    # Generate counterfactuals for each class.\n    logits = logits_model(original_input)\n    num_classes = logits.shape[1]\n    \n    counterfactuals = {}\n    predictions = {}\n    for cls in range(num_classes):\n        print(f\"\\nGenerating counterfactual for target class {cls}:\")\n        cf_input = generate_counterfactual(logits_model, original_input, target_class=cls,\n                                           num_iterations=100, learning_rate=0.01,\n                                           regularization_weight=0.01)\n        pred_logits = logits_model(cf_input)\n        pred = pred_logits.argmax(dim=1).item()\n        predictions[cls] = pred\n        counterfactuals[cls] = cf_input.squeeze(0).cpu().detach().numpy()\n    \n    # Convert the original input to numpy.\n    original_np = original_input.squeeze(0).cpu().detach().numpy()\n    \n    # Instead of plotting all counterfactuals per channel, select only the most different one for each channel.\n    channels_to_plot = list(range(original_np.shape[0]))  # Plot all channels.\n    plot_best_counterfactuals_side_by_side(original_np, counterfactuals, predictions, channels_to_plot)\n    \n    print(\"Counterfactual generation complete.\")\n\nif __name__ == \"__main__\":\n    main()\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-08T12:45:32.927815Z","iopub.execute_input":"2025-04-08T12:45:32.928144Z","iopub.status.idle":"2025-04-08T12:45:57.531139Z","shell.execute_reply.started":"2025-04-08T12:45:32.928122Z","shell.execute_reply":"2025-04-08T12:45:57.530348Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Sensitivity Analysis","metadata":{}},{"cell_type":"code","source":"# ------------------------------\n# Sensitivity Analysis Functions\n# ------------------------------\n\ndef add_noise(input_tensor, noise_std=0.01):\n    \"\"\"\n    Adds Gaussian noise to the input tensor.\n    \"\"\"\n    noise = torch.randn_like(input_tensor) * noise_std\n    return input_tensor + noise\n\ndef occlude_region(input_tensor, mask):\n    \"\"\"\n    Applies occlusion on the input_tensor using a mask.\n    'mask' is a tensor with 1s for regions to keep and 0s for regions to occlude.\n    \n    Here, occluding a channel by setting it to a baseline value (e.g., zeros).\n    \"\"\"\n    baseline = torch.zeros_like(input_tensor)  # or use the mean value of input_tensor\n    return input_tensor * mask + baseline * (1 - mask)\n\n\ndef sensitivity_analysis_on_dataset(model, dataset, noise_std, device, sample_limit=100):\n    \"\"\"\n    Performs sensitivity analysis over all samples in the dataset.\n    \n    For each sample:\n      - Noise Analysis: Adds Gaussian noise and computes the average absolute difference in logits.\n      - Occlusion Analysis: For each channel, occludes that channel and computes:\n            (a) The average absolute difference in logits.\n            (b) Whether the predicted class (argmax) changes.\n            \n    Returns:\n        avg_noise_diff (float): Average difference in logits due to noise over samples.\n        avg_occlusion_diffs (np.ndarray): Array of average differences (L1 norm on logits) per channel.\n        decision_changes (np.ndarray): Array (per channel) of the fraction of samples where occlusion caused a decision change.\n    \"\"\"\n    noise_diffs = []\n    num_channels = None\n    occlusion_diffs_sum = None\n    decision_change_counts = None\n    num_samples = 0\n    \n    for idx, sample in enumerate(dataset):\n        if idx >= sample_limit:\n            break\n        # Assume sample is either raw EEG or (data, ...) tuple.\n        if isinstance(sample, tuple):\n            sample = sample[0]\n        input_tensor = torch.tensor(sample).unsqueeze(0).to(device)\n        original_pred = model(input_tensor)\n        orig_class = original_pred.argmax(dim=1).item()\n        \n        # Noise analysis.\n        noisy_input = add_noise(input_tensor, noise_std=noise_std)\n        noisy_pred = model(noisy_input)\n        noise_diff = torch.mean(torch.abs(original_pred - noisy_pred)).item()\n        noise_diffs.append(noise_diff)\n        \n        # Initialize channel count and occlusion accumulators.\n        if num_channels is None:\n            num_channels = input_tensor.shape[1]\n            occlusion_diffs_sum = np.zeros(num_channels)\n            decision_change_counts = np.zeros(num_channels)\n        \n        # Occlusion analysis per channel.\n        for ch in range(num_channels):\n            mask = torch.ones_like(input_tensor)\n            mask[:, ch, :] = 0  # occlude channel ch\n            occluded_input = occlude_region(input_tensor, mask)\n            occluded_pred = model(occluded_input)\n            diff = torch.mean(torch.abs(original_pred - occluded_pred)).item()\n            occlusion_diffs_sum[ch] += diff\n            \n            # Check if the occlusion changed the predicted class.\n            occluded_class = occluded_pred.argmax(dim=1).item()\n            if occluded_class != orig_class:\n                decision_change_counts[ch] += 1\n        \n        num_samples += 1\n        if (idx + 1) % 20 == 0:\n            print(f\"Processed {idx + 1} samples...\")\n    \n    avg_noise_diff = np.mean(noise_diffs) if noise_diffs else 0\n    avg_occlusion_diffs = occlusion_diffs_sum / num_samples if num_samples > 0 else None\n    decision_changes = decision_change_counts / num_samples if num_samples > 0 else None\n    \n    return avg_noise_diff, avg_occlusion_diffs, decision_changes\n\ndef plot_occlusion_results(avg_occlusion_diffs, decision_changes):\n    \"\"\"\n    Plots the occlusion sensitivity results.\n    \n    Two bar charts are produced:\n      1. The average change in logits per channel.\n      2. The fraction of samples where occluding a channel changed the predicted class.\n    \n    Args:\n        avg_occlusion_diffs (np.ndarray): Array of average differences per channel.\n        decision_changes (np.ndarray): Array of fraction (or percentage) of decision changes per channel.\n    \"\"\"\n    channels = np.arange(len(avg_occlusion_diffs))\n    \n    fig, axs = plt.subplots(1, 2, figsize=(14, 6))\n    \n    # Plot average change in logits.\n    axs[0].bar(channels, avg_occlusion_diffs, color='salmon')\n    axs[0].set_xlabel('Channel Index')\n    axs[0].set_ylabel('Average Change in Logits')\n    axs[0].set_title('Occlusion Sensitivity (Logits Difference)')\n    axs[0].grid(True)\n    \n    # Plot decision change fractions.\n    axs[1].bar(channels, decision_changes * 100, color='skyblue')\n    axs[1].set_xlabel('Channel Index')\n    axs[1].set_ylabel('Percentage of Decision Changes (%)')\n    axs[1].set_title('Occlusion Sensitivity (Decision Changes)')\n    axs[1].grid(True)\n    \n    plt.tight_layout()\n    plt.show()\n\ndef temporal_occlusion_analysis(model, input_tensor, segments):\n    \"\"\"\n    Performs temporal occlusion analysis on a single input_tensor over given segments.\n    For each segment, occludes that time window and computes:\n        - The average difference in logits.\n        - Whether the predicted class changed.\n    \n    Args:\n        model (nn.Module): The model (logits_model).\n        input_tensor (torch.Tensor): Input of shape (1, channels, time).\n        segments (list of tuple): Each tuple is (start, end) indices.\n    \n    Returns:\n        dict: Mapping from segment (start, end) to a dict with keys:\n            'diff', 'decision_change', 'orig_class', 'new_class'\n    \"\"\"\n    original_pred = model(input_tensor)\n    orig_class = original_pred.argmax(dim=1).item()\n    results = {}\n    for (start, end) in segments:\n        occluded_input = occlude_temporal_segments(input_tensor, start, end)\n        occluded_pred = model(occluded_input)\n        diff = torch.mean(torch.abs(original_pred - occluded_pred)).item()\n        new_class = occluded_pred.argmax(dim=1).item()\n        decision_change = (new_class != orig_class)\n        results[(start, end)] = {\n            \"diff\": diff,\n            \"decision_change\": decision_change,\n            \"orig_class\": orig_class,\n            \"new_class\": new_class\n        }\n    return results\n\ndef occlude_temporal_segments(input_tensor, start, end):\n    \"\"\"\n    Occludes (zeros) a temporal segment from start to end.\n    \"\"\"\n    mask = torch.ones_like(input_tensor)\n    mask[:, :, start:end] = 0\n    return occlude_region(input_tensor, mask)\n\ndef plot_temporal_occlusion_results(results):\n    \"\"\"\n    Plots the differences for temporal occlusion analysis.\n    \n    Args:\n        results (dict): Dictionary returned by temporal_occlusion_analysis.\n    \"\"\"\n    segments = list(results.keys())\n    diffs = [results[seg][\"diff\"] for seg in segments]\n    decision_changes = [results[seg][\"decision_change\"] for seg in segments]\n    \n    # Plot the average difference per segment.\n    labels = [f\"{seg[0]}-{seg[1]}\" for seg in segments]\n    plt.figure(figsize=(10, 6))\n    plt.bar(labels, diffs, color='lightgreen')\n    plt.xlabel('Time Segment (start-end)')\n    plt.ylabel('Average Logit Difference')\n    plt.title('Temporal Occlusion Analysis: Logit Differences')\n    plt.grid(True)\n    plt.show()\n    \n    # Print decision change results.\n    for seg in segments:\n        res = results[seg]\n        print(f\"Segment {seg}: Diff={res['diff']:.4f}, Decision changed: {res['decision_change']}, Orig class: {res['orig_class']}, New class: {res['new_class']}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-08T14:49:21.248332Z","iopub.execute_input":"2025-04-08T14:49:21.248690Z","iopub.status.idle":"2025-04-08T14:49:21.265026Z","shell.execute_reply.started":"2025-04-08T14:49:21.248661Z","shell.execute_reply":"2025-04-08T14:49:21.264033Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ------------------------------\n# Main Workflow\n# ------------------------------\n\ndef main():\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    \n    # Load data and original model.\n    model_weight = config['weights_x3d'][0]\n    base_model = Netx3d()\n    _, test_df, base_model = load_data_and_model(base_model, model_weight, device)\n    \n    # Create logits wrapper model.\n    logits_model = Netx3dLogits()\n    logits_model.load_state_dict(base_model.state_dict(), strict=True)\n    logits_model.to(device)\n    logits_model.eval()\n    \n\n    # ------------------------------\n    # Sensitivity Analysis Over the Test Dataset\n    # ------------------------------\n    print(\"\\nPerforming Sensitivity Analysis on Test Dataset:\")\n    test_dataset = DataProcessor(\n        test_df, \n        training_flag=False,\n        shuffle=False,\n        use_eeg=True,\n        lower_cut=0.5,\n        upper_cut=20,\n        use_mne_filter=False\n    )\n    \n    # You can adjust sample_limit as needed.\n    avg_noise_diff, avg_occlusion_diffs, decision_changes = sensitivity_analysis_on_dataset(\n        logits_model, test_dataset, noise_std=0.01, device=device, sample_limit=100)\n    \n    print(f\"Average noise-induced difference in logits: {avg_noise_diff:.4f}\")\n    print(\"Average occlusion-induced difference per channel:\")\n    for ch, diff in enumerate(avg_occlusion_diffs):\n        print(f\"  Channel {ch}: {diff:.4f}\")\n    print(\"Percentage of decision changes per channel due to occlusion:\")\n    for ch, change in enumerate(decision_changes):\n        print(f\"  Channel {ch}: {change*100:.2f}%\")\n    \n    plot_occlusion_results(avg_occlusion_diffs, decision_changes)\n    \n    print(\"Sensitivity analysis complete.\")\n\n    # ------------------------------\n    # Temporal Occlusion Analysis on a single sample\n    # ------------------------------\n    print(\"\\nPerforming Temporal Occlusion Analysis on selected sample:\")\n    # Define segments for occlusion (e.g., segments of 100 time points).\n    # Adjust these ranges to fit the time axis of your EEG data.\n    segments = [(0, 100), (100, 200), (200, 300), (300, 400)]\n    temporal_results = temporal_occlusion_analysis(logits_model, original_input, segments)\n    \n    print(\"Temporal Occlusion Analysis Results:\")\n    for seg, res in temporal_results.items():\n        print(f\"Segment {seg}: Diff={res['diff']:.4f}, Decision changed: {res['decision_change']}, Orig class: {res['orig_class']}, New class: {res['new_class']}\")\n    \n    plot_temporal_occlusion_results(temporal_results)\n\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-08T14:49:22.894858Z","iopub.execute_input":"2025-04-08T14:49:22.895165Z","iopub.status.idle":"2025-04-08T14:49:24.958752Z","shell.execute_reply.started":"2025-04-08T14:49:22.895144Z","shell.execute_reply":"2025-04-08T14:49:24.958086Z"}},"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,"execution":{"iopub.status.busy":"2025-04-08T10:27:03.908357Z","iopub.execute_input":"2025-04-08T10:27:03.908739Z","iopub.status.idle":"2025-04-08T10:27:04.139330Z","shell.execute_reply.started":"2025-04-08T10:27:03.908695Z","shell.execute_reply":"2025-04-08T10:27:04.138646Z"}},"outputs":[],"execution_count":null}]}