{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":106809,"databundleVersionId":13056355,"sourceType":"competition"}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import h5py\nfrom collections import defaultdict\n\ndef load_h5py_file(file_path: str) -> dict:\n    data = defaultdict(list)\n    \n    with h5py.File(file_path, 'r') as f:\n        trials = list(f.keys())\n        for trial in trials:\n            g = f[trial]\n\n            data['neural_features'].append(g['input_features'][:])\n            data['n_time_steps'].append(g.attrs['n_time_steps'])\n            data['seq_class_ids'].append(g['seq_class_ids'][:] if 'seq_class_ids' in g else None)\n            data['seq_len'].append(g.attrs['seq_len'] if 'seq_len' in g.attrs else None)\n            data['transcriptions'].append(g['transcription'][:] if 'transcription' in g else None)\n            data['sentence_label'].append(g.attrs['sentence_label'][:] if 'sentence_label' in g.attrs else None)\n            data['session'].append(g.attrs['session'])\n            data['block_num'].append(g.attrs['block_num'])\n            data['trial_num'].append(g.attrs['trial_num'])\n    \n    return dict(data)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-21T18:55:05.518592Z","iopub.execute_input":"2025-11-21T18:55:05.519359Z","iopub.status.idle":"2025-11-21T18:55:05.524839Z","shell.execute_reply.started":"2025-11-21T18:55:05.519329Z","shell.execute_reply":"2025-11-21T18:55:05.524351Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data = load_h5py_file(\"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final/t15.2023.08.13/data_train.hdf5\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-21T18:55:07.900077Z","iopub.execute_input":"2025-11-21T18:55:07.900804Z","iopub.status.idle":"2025-11-21T18:55:14.461402Z","shell.execute_reply.started":"2025-11-21T18:55:07.900780Z","shell.execute_reply":"2025-11-21T18:55:14.460625Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Utils","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn.functional as F\nimport numpy as np\nfrom scipy.ndimage import gaussian_filter1d\n\ndef gauss_smooth(inputs, device, smooth_kernel_std=2, smooth_kernel_size=100,  padding='same'):\n    \"\"\"\n    Applies a 1D Gaussian smoothing operation with PyTorch to smooth the data along the time axis.\n    Args:\n        inputs (tensor : B x T x N): A 3D tensor with batch size B, time steps T, and number of features N.\n                                     Assumed to already be on the correct device (e.g., GPU).\n        kernelSD (float): Standard deviation of the Gaussian smoothing kernel.\n        padding (str): Padding mode, either 'same' or 'valid'.\n        device (str): Device to use for computation (e.g., 'cuda' or 'cpu').\n    Returns:\n        smoothed (tensor : B x T x N): A smoothed 3D tensor with batch size B, time steps T, and number of features N.\n    \"\"\"\n    # Get Gaussian kernel\n    inp = np.zeros(smooth_kernel_size, dtype=np.float32)\n    inp[smooth_kernel_size // 2] = 1\n    gaussKernel = gaussian_filter1d(inp, smooth_kernel_std)\n    validIdx = np.argwhere(gaussKernel > 0.01)\n    gaussKernel = gaussKernel[validIdx]\n    gaussKernel = np.squeeze(gaussKernel / np.sum(gaussKernel))\n\n    # Convert to tensor\n    gaussKernel = torch.tensor(gaussKernel, dtype=torch.float32, device=device)\n    gaussKernel = gaussKernel.view(1, 1, -1)  # [1, 1, kernel_size]\n\n    # Prepare convolution\n    B, T, C = inputs.shape\n    inputs = inputs.permute(0, 2, 1)  # [B, C, T]\n    gaussKernel = gaussKernel.repeat(C, 1, 1)  # [C, 1, kernel_size]\n\n    # Perform convolution\n    smoothed = F.conv1d(inputs, gaussKernel, padding=padding, groups=C)\n    \n    return smoothed.permute(0, 2, 1)  # [B, T, C]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-21T18:55:17.649403Z","iopub.execute_input":"2025-11-21T18:55:17.649927Z","iopub.status.idle":"2025-11-21T18:55:17.656275Z","shell.execute_reply.started":"2025-11-21T18:55:17.649900Z","shell.execute_reply":"2025-11-21T18:55:17.655589Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def transform_data(features, n_time_steps, mode = 'train', device='cpu'):\n    '''\n    Apply various augmentations and smoothing to data\n    Performing augmentations is much faster on GPU than CPU\n    '''\n\n    transform_args = {\n        'white_noise_std': 1.0, # standard deviation of the white noise added to the data\n        'constant_offset_std': 0.2, # standard deviation of the constant offset added to the data\n        'random_walk_std': 0.0, # standard deviation of the random walk added to the data\n        'random_walk_axis': -1, # axis along which the random walk is applied\n        'static_gain_std': 0.0, # standard deviation of the static gain applied to the data\n        'random_cut': 3, # number of time steps to randomly cut from the beginning of each batch of trials\n        'smooth_kernel_size': 100, # size of the smoothing kernel applied to the data\n        'smooth_data': True, # whether to smooth the data\n        'smooth_kernel_std': 2, # standard deviation of the smoothing kernel applied to the data                     \n    }\n\n    data_shape = features.shape\n    batch_size = data_shape[0]\n    channels = data_shape[-1]\n\n    # We only apply these augmentations in training\n    if mode == 'train':\n\n        # add static gain noise \n        if transform_args['static_gain_std'] > 0:\n            warp_mat = torch.tile(torch.unsqueeze(torch.eye(channels), dim = 0), (batch_size, 1, 1))\n            warp_mat += torch.randn_like(warp_mat, device=device) * transform_args['static_gain_std']\n\n            features = torch.matmul(features, warp_mat)\n\n        # add white noise\n        if transform_args['white_noise_std'] > 0:\n            features += torch.randn(data_shape, device=device) * transform_args['white_noise_std']\n\n        # add constant offset noise \n        if transform_args['constant_offset_std'] > 0:\n            features += torch.randn((batch_size, 1, channels), device=device) * transform_args['constant_offset_std']\n\n        # add random walk noise\n        if transform_args['random_walk_std'] > 0:\n            features += torch.cumsum(torch.randn(data_shape, device=device) * transform_args['random_walk_std'], dim = transform_args['random_walk_axis'])\n\n        # randomly cutoff part of the data timecourse\n        if transform_args['random_cut'] > 0:\n            cut = np.random.randint(0, transform_args['random_cut'])\n            features = features[:, cut:, :]\n            n_time_steps = n_time_steps - cut\n\n    # Apply Gaussian smoothing to data \n    # This is done in both training and validation\n    if transform_args['smooth_data']:\n        features = gauss_smooth(\n            inputs = features, \n            device = device,\n            smooth_kernel_std = transform_args['smooth_kernel_std'],\n            smooth_kernel_size= transform_args['smooth_kernel_size'],\n            padding='valid'\n            )\n    \n    return features, n_time_steps","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# GRU Decoder","metadata":{}},{"cell_type":"code","source":"import torch \nfrom torch import nn\n\nclass GRUDecoder(nn.Module):\n    '''\n    Defines the GRU decoder\n\n    This class combines day-specific input layers, a GRU, and an output classification layer\n    '''\n    def __init__(self,\n                 neural_dim,\n                 n_units,\n                 n_days,\n                 n_classes,\n                 rnn_dropout = 0.0,\n                 input_dropout = 0.0,\n                 n_layers = 5, \n                 patch_size = 0,\n                 patch_stride = 0,\n                 ):\n        '''\n        neural_dim  (int)      - number of channels in a single timestep (e.g. 512)\n        n_units     (int)      - number of hidden units in each recurrent layer - equal to the size of the hidden state\n        n_days      (int)      - number of days in the dataset\n        n_classes   (int)      - number of classes \n        rnn_dropout    (float) - percentage of units to droupout during training\n        input_dropout (float)  - percentage of input units to dropout during training\n        n_layers    (int)      - number of recurrent layers \n        patch_size  (int)      - the number of timesteps to concat on initial input layer - a value of 0 will disable this \"input concat\" step \n        patch_stride(int)      - the number of timesteps to stride over when concatenating initial input \n        '''\n        super(GRUDecoder, self).__init__()\n        \n        self.neural_dim = neural_dim\n        self.n_units = n_units\n        self.n_classes = n_classes\n        self.n_layers = n_layers \n        self.n_days = n_days\n\n        self.rnn_dropout = rnn_dropout\n        self.input_dropout = input_dropout\n        \n        self.patch_size = patch_size\n        self.patch_stride = patch_stride\n\n        # Parameters for the day-specific input layers\n        self.day_layer_activation = nn.Softsign() # basically a shallower tanh \n\n        # Set weights for day layers to be identity matrices so the model can learn its own day-specific transformations\n        self.day_weights = nn.ParameterList(\n            [nn.Parameter(torch.eye(self.neural_dim)) for _ in range(self.n_days)]\n        )\n        self.day_biases = nn.ParameterList(\n            [nn.Parameter(torch.zeros(1, self.neural_dim)) for _ in range(self.n_days)]\n        )\n\n        self.day_layer_dropout = nn.Dropout(input_dropout)\n        \n        self.input_size = self.neural_dim\n\n        # If we are using \"strided inputs\", then the input size of the first recurrent layer will actually be in_size * patch_size\n        if self.patch_size > 0:\n            self.input_size *= self.patch_size\n\n        self.gru = nn.GRU(\n            input_size = self.input_size,\n            hidden_size = self.n_units,\n            num_layers = self.n_layers,\n            dropout = self.rnn_dropout, \n            batch_first = True, # The first dim of our input is the batch dim\n            bidirectional = False,\n        )\n\n        # Set recurrent units to have orthogonal param init and input layers to have xavier init\n        for name, param in self.gru.named_parameters():\n            if \"weight_hh\" in name:\n                nn.init.orthogonal_(param)\n            if \"weight_ih\" in name:\n                nn.init.xavier_uniform_(param)\n\n        # Prediciton head. Weight init to xavier\n        self.out = nn.Linear(self.n_units, self.n_classes)\n        nn.init.xavier_uniform_(self.out.weight)\n\n        # Learnable initial hidden states\n        self.h0 = nn.Parameter(nn.init.xavier_uniform_(torch.zeros(1, 1, self.n_units)))\n\n    def forward(self, x, day_idx, states = None, return_state = False):\n        '''\n        x        (tensor)  - batch of examples (trials) of shape: (batch_size, time_series_length, neural_dim)\n        day_idx  (tensor)  - tensor which is a list of day indexs corresponding to the day of each example in the batch x. \n        '''\n\n        # Apply day-specific layer to (hopefully) project neural data from the different days to the same latent space\n        day_weights = torch.stack([self.day_weights[i] for i in day_idx], dim=0)\n        day_biases = torch.cat([self.day_biases[i] for i in day_idx], dim=0).unsqueeze(1)\n\n        x = torch.einsum(\"btd,bdk->btk\", x, day_weights) + day_biases\n        x = self.day_layer_activation(x)\n\n        # Apply dropout to the ouput of the day specific layer\n        if self.input_dropout > 0:\n            x = self.day_layer_dropout(x)\n\n        # (Optionally) Perform input concat operation\n        if self.patch_size > 0: \n  \n            x = x.unsqueeze(1)                      # [batches, 1, timesteps, feature_dim]\n            x = x.permute(0, 3, 1, 2)               # [batches, feature_dim, 1, timesteps]\n            \n            # Extract patches using unfold (sliding window)\n            x_unfold = x.unfold(3, self.patch_size, self.patch_stride)  # [batches, feature_dim, 1, num_patches, patch_size]\n            \n            # Remove dummy height dimension and rearrange dimensions\n            x_unfold = x_unfold.squeeze(2)           # [batches, feature_dum, num_patches, patch_size]\n            x_unfold = x_unfold.permute(0, 2, 3, 1)  # [batches, num_patches, patch_size, feature_dim]\n\n            # Flatten last two dimensions (patch_size and features)\n            x = x_unfold.reshape(x.size(0), x_unfold.size(1), -1) \n        \n        # Determine initial hidden states\n        if states is None:\n            states = self.h0.expand(self.n_layers, x.shape[0], self.n_units).contiguous()\n\n        # Pass input through RNN \n        output, hidden_states = self.gru(x, states)\n\n        # Compute logits\n        logits = self.out(output)\n        \n        if return_state:\n            return logits, hidden_states\n        \n        return logits","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-21T18:55:28.926417Z","iopub.execute_input":"2025-11-21T18:55:28.927111Z","iopub.status.idle":"2025-11-21T18:55:28.939251Z","shell.execute_reply.started":"2025-11-21T18:55:28.927089Z","shell.execute_reply":"2025-11-21T18:55:28.938399Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# LSTM Decoder","metadata":{}},{"cell_type":"code","source":"class LSTMDecoder(nn.Module):\n    '''\n    Defines the LSTM decoder\n\n    This class combines day-specific input layers, an LSTM, and an output classification layer\n    '''\n    def __init__(self,\n                 neural_dim,\n                 n_units,\n                 n_days,\n                 n_classes,\n                 rnn_dropout=0.0,\n                 input_dropout=0.0,\n                 n_layers=5,\n                 patch_size=0,\n                 patch_stride=0):\n        super(LSTMDecoder, self).__init__()\n        \n        self.neural_dim = neural_dim\n        self.n_units = n_units\n        self.n_classes = n_classes\n        self.n_layers = n_layers \n        self.n_days = n_days\n\n        self.rnn_dropout = rnn_dropout\n        self.input_dropout = input_dropout\n        \n        self.patch_size = patch_size\n        self.patch_stride = patch_stride\n\n        # Day-specific linear transformations\n        self.day_layer_activation = nn.Softsign()\n        self.day_weights = nn.ParameterList(\n            [nn.Parameter(torch.eye(self.neural_dim)) for _ in range(self.n_days)]\n        )\n        self.day_biases = nn.ParameterList(\n            [nn.Parameter(torch.zeros(1, self.neural_dim)) for _ in range(self.n_days)]\n        )\n\n        self.day_layer_dropout = nn.Dropout(input_dropout)\n        \n        self.input_size = self.neural_dim\n        if self.patch_size > 0:\n            self.input_size *= self.patch_size\n\n        # --- LSTM instead of GRU ---\n        self.lstm = nn.LSTM(\n            input_size=self.input_size,\n            hidden_size=self.n_units,\n            num_layers=self.n_layers,\n            dropout=self.rnn_dropout,\n            batch_first=True,\n            bidirectional=False,\n        )\n\n        # Initialize weights (same logic as before)\n        for name, param in self.lstm.named_parameters():\n            if \"weight_hh\" in name:\n                nn.init.orthogonal_(param)\n            if \"weight_ih\" in name:\n                nn.init.xavier_uniform_(param)\n\n        # Output layer\n        self.out = nn.Linear(self.n_units, self.n_classes)\n        nn.init.xavier_uniform_(self.out.weight)\n\n        # Learnable initial hidden and cell states\n        self.h0 = nn.Parameter(torch.zeros(1, 1, self.n_units))\n        self.c0 = nn.Parameter(torch.zeros(1, 1, self.n_units))\n\n    def forward(self, x, day_idx, states=None, return_state=False):\n        '''\n        x        (tensor)  - (batch_size, time_series_length, neural_dim)\n        day_idx  (tensor)  - list of day indices corresponding to each example\n        '''\n        # --- Apply day-specific transformations ---\n        day_weights = torch.stack([self.day_weights[i] for i in day_idx], dim=0)\n        day_biases = torch.cat([self.day_biases[i] for i in day_idx], dim=0).unsqueeze(1)\n        x = torch.einsum(\"btd,bdk->btk\", x, day_weights) + day_biases\n        x = self.day_layer_activation(x)\n\n        if self.input_dropout > 0:\n            x = self.day_layer_dropout(x)\n\n        # --- Optional strided patching ---\n        if self.patch_size > 0:\n            x = x.unsqueeze(1)                      # [b, 1, t, d]\n            x = x.permute(0, 3, 1, 2)               # [b, d, 1, t]\n            x_unfold = x.unfold(3, self.patch_size, self.patch_stride)  # [b, d, 1, n_patches, patch_size]\n            x_unfold = x_unfold.squeeze(2).permute(0, 2, 3, 1)          # [b, n_patches, patch_size, d]\n            x = x_unfold.reshape(x.size(0), x_unfold.size(1), -1)       # [b, n_patches, patch_size*d]\n\n        # --- Initial states ---\n        if states is None:\n            h0 = self.h0.expand(self.n_layers, x.shape[0], self.n_units).contiguous()\n            c0 = self.c0.expand(self.n_layers, x.shape[0], self.n_units).contiguous()\n        else:\n            h0, c0 = states\n\n        # --- Forward pass through LSTM ---\n        output, (hn, cn) = self.lstm(x, (h0, c0))\n\n        # --- Output classification ---\n        logits = self.out(output)\n\n        if return_state:\n            return logits, (hn, cn)\n\n        return logits\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Transformer","metadata":{}},{"cell_type":"code","source":"class Transformer(nn.Module):\n    def __init__(\n        self,\n        neural_dim=512,\n        d_model=256,\n        nhead=8,\n        num_encoder_layers=4,\n        dim_feedforward=1024,\n        dropout=0.1,\n        n_days=45,\n        n_classes=41,\n    ):\n        super().__init__()\n\n        self.d_model = d_model\n\n        self.day_weights = nn.ParameterList(\n            [nn.Parameter(torch.eye(neural_dim)) for _ in range(n_days)]\n        )\n        self.day_biases = nn.ParameterList(\n            [nn.Parameter(torch.zeros(1, neural_dim)) for _ in range(n_days)]\n        )\n        self.day_activation = nn.Softsign()\n\n        self.input_projection = nn.Linear(neural_dim, d_model)\n\n        self.pos_encoder = PositionalEncoding(d_model, dropout)\n\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d_model,\n            nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout,\n            batch_first=True,\n        )\n        self.transformer_encoder = nn.TransformerEncoder(\n            encoder_layer, num_layers=num_encoder_layers\n        )\n\n        self.output_projection = nn.Linear(d_model, n_classes)\n\n    def forward(self, x, day_idx):\n        day_weights = torch.stack([self.day_weights[i] for i in day_idx], dim=0)\n        day_biases = torch.cat([self.day_biases[i] for i in day_idx], dim=0).unsqueeze(\n            1\n        )\n        x = torch.einsum(\"btd,bdk->btk\", x, day_weights) + day_biases\n        x = self.day_activation(x)\n\n        x = self.input_projection(x)\n\n        # add positional encoding\n        x = self.pos_encoder(x)\n        x = self.transformer_encoder(x)\n\n        # project to phoneme logits\n        logits = self.output_projection(x)\n\n        return logits\n\n\nclass PositionalEncoding(nn.Module):\n    def __init__(self, d_model, dropout=0.1, max_len=5000):\n        super().__init__()\n        self.dropout = nn.Dropout(p=dropout)\n\n        position = torch.arange(max_len).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, d_model, 2) * (-np.log(10000.0) / d_model))\n        pe = torch.zeros(1, max_len, d_model)\n        pe[0, :, 0::2] = torch.sin(position * div_term)\n        pe[0, :, 1::2] = torch.cos(position * div_term)\n        self.register_buffer(\"pe\", pe)\n\n    def forward(self, x):\n        x = x + self.pe[:, : x.size(1), :]\n        return self.dropout(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-21T18:54:35.917749Z","iopub.status.idle":"2025-11-21T18:54:35.918392Z","shell.execute_reply.started":"2025-11-21T18:54:35.918021Z","shell.execute_reply":"2025-11-21T18:54:35.918033Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CNN + Transformer","metadata":{}},{"cell_type":"code","source":"class EEGFeatureExtractor(nn.Module):\n    \"\"\"\n    CNN-based feature extractor for EEG data to learn EEG patterns, oscillations, etc.\n    Reduces dimensionality before transformer input by mapping raw signal to a lower dimension feature space.\n    \"\"\"\n\n    def __init__(self, in_features, feature_dim=256):\n        super().__init__()\n        network = nn.Sequential(\n            # layer 1: downsample by 2, output 128 channels - capture local patterns\n            nn.Conv1d(in_features, 128, kernel_size=5, stride=2, padding=2),\n            nn.BatchNorm1d(128), # batch normalization for stable training\n            nn.GELU(), # smooth non-linearity - better than Relu for continuous EEG signal\n\n            # layer 2: downsample by 2, output 256 channels - capture broader patterns\n            nn.Conv1d(128, 256, kernel_size=5, stride=2, padding=2),\n            nn.BatchNorm1d(256), # batch normalization for stable training\n            nn.GELU(), # smooth non-linearity - better than Relu for continuous EEG signal\n\n            # layer 3: no downsampling, output feature_dim channels - final feature representation\n            nn.Conv1d(256, feature_dim, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm1d(feature_dim), # batch normalization for stable training\n            nn.GELU(), # smooth non-linearity - better than Relu for continuous EEG signal\n        )\n\n        self.net = network\n\n    def forward(self, x):\n        # input: (B, T, D) \n        x = x.permute(0, 2, 1)  # → (B, D, T) network swaps time and feature dims\n        x = self.net(x) \n        x = x.permute(0, 2, 1)  # → (B, reduced_T, feature_dim) to original format\n        return x\n    \n\nclass CTCDecoder(nn.Module):\n    def __init__(self, d_model=256, vocab_size=60):\n        super().__init__()\n        self.fc = nn.Linear(d_model, vocab_size) # fully connected layer to map to vocab size\n\n    def forward(self, x):\n        # x: (B, T, d_model)\n        return self.fc(x)  # (B, T, vocab_size)\n    \nclass LearnedPositionalEncoding(nn.Module):\n    \"\"\"Learned positional embeddings for temporal sequences.\"\"\"\n    def __init__(self, max_len: int, d_model: int):\n        super().__init__()\n        self.pe = nn.Parameter(torch.randn(1, max_len, d_model) * 0.02)\n        self.max_len = max_len\n        self.d_model = d_model\n\n    def forward(self, x):\n        # Extend positional embeddings if sequence is longer than expected\n        if x.size(1) > self.max_len:\n            extra = x.size(1) - self.max_len\n            extra_pe = nn.Parameter(torch.randn(1, extra, self.d_model) * 0.02).to(x.device)\n            self.pe = nn.Parameter(torch.cat([self.pe, extra_pe], dim=1))\n            self.max_len = x.size(1)\n        return x + self.pe[:, :x.size(1), :]\n\n\nclass TemporalTransformer(nn.Module):\n    \"\"\"\n    Transformer model for EEG feature modeling.\n\n    Args:\n      neural_dim:   EEG feature dimension per timestep\n      d_model:      transformer embedding dimension\n      n_days:       number of days data was recorded on\n      n_tokens:     number of output tokens\n      n_layers:     number of transformer layers\n      n_heads:      number of attention heads - to sum to d_model, d_model must be dividable by n_heads\n      ff_mult:      feedforward expansion factor\n      dropout:      dropout probability\n      patch_stride: stride for patching input (0 for no stride)\n      max_seq_len:  maximum sequence length to tokenize\n    \"\"\"\n\n    def __init__(self,\n                 neural_dim,\n                 d_model,\n                 n_days,\n                 n_tokens,\n                 n_layers=6,\n                 n_heads=8,\n                 ff_mult=4,\n                 dropout=0.1,\n                 patch_stride=0,\n                 max_seq_len=2048):\n\n        super().__init__()\n\n        self.neural_dim = neural_dim\n        self.d_model = d_model\n        self.n_days = n_days\n        self.n_tokens = n_tokens\n        self.patch_stride = patch_stride\n        self.input_size = self.neural_dim\n\n        # need to allow for cross session variability and training by day \n        # day specific data input layers\n        self.day_layer_activation = nn.Softsign() # non linear scaling for day specific layers\n        self.day_weights = nn.ParameterList(\n            [nn.Parameter(torch.eye(self.neural_dim)) for _ in range(self.n_days)]\n        )\n        self.day_biases = nn.ParameterList(\n            [nn.Parameter(torch.zeros(1, self.neural_dim)) for _ in range(self.n_days)]\n        )\n        # dropout for day-specific layers (in between layers) - prevent overfitting to day specific features\n        self.day_dropout = nn.Dropout(dropout)\n        \n        # input projection to d_model dimensions\n        self.input_proj = nn.Linear(self.input_size, self.d_model)\n        nn.init.xavier_uniform_(self.input_proj.weight)\n\n        # encoding of positional information\n        self.pos_encoder = LearnedPositionalEncoding(max_seq_len, self.d_model) # maps time steps to d_model dimension space\n\n        # transformer encoder\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=self.d_model,\n            nhead=n_heads,\n            dim_feedforward=self.d_model * ff_mult,\n            dropout=dropout,\n            activation=\"gelu\", # gaussian error linear unit - transformers use soft attention, so relu can be too harsh (hard attention)\n            batch_first=True\n        )\n        self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=n_layers)\n\n        # output classification\n        self.out = nn.Linear(self.d_model, self.n_tokens)\n\n        # initialize output layer to be uniform (based on number of input and output units layers)\n        nn.init.xavier_uniform_(self.out.weight)\n\n        # add final dropout layer if we have dropout\n        self.final_dropout = nn.Dropout(dropout)\n\n    def forward(self, x, day_idx): # should work with batch input from nejm\n        \"\"\"\n        x: (batch, time, neural_dim)\n        day_idx: tensor of shape (batch,) with day indices [0..n_days-1]\n        \"\"\"\n        device = x.device # set device to use for tensors\n\n        day_weights = torch.stack([self.day_weights[int(i)] for i in day_idx], dim=0).to(device)\n        day_biases = torch.cat([self.day_biases[int(i)] for i in day_idx], dim=0).unsqueeze(1).to(device)\n        \n        x = torch.einsum(\"btd,bdk->btk\", x, day_weights) + day_biases # eigensum for batch matrix multiplication\n        x = self.day_layer_activation(x) \n        x = self.day_dropout(x)\n\n        # project batch dimension input to d_model (embedder) dimension\n        x = self.input_proj(x)\n\n        # add positional encoding to map time steps to d_model space\n        x = self.pos_encoder(x)\n\n        # pass to transformer encoder\n        x = self.encoder(x)\n        x = self.final_dropout(x) # dropout in final fully connected layer\n\n        # output classification - map to output tokens\n        logits = self.out(x)  # (batch, time, n_classes)\n        return logits\n    \nclass EEGToTextModel(nn.Module):\n    def __init__(self,\n                 in_features,\n                 vocab_size,\n                 d_model=256,\n                 n_days=45,\n                 n_tokens=None,\n                 n_layers=6,\n                 n_heads=8,\n                 ff_mult=4,\n                 dropout=0.1,\n                 patch_stride=0,\n                 max_seq_len=2048,\n                 feature_dim=256):\n        super().__init__()\n        \n        # Feature extractor maps EEG features to transformer embedding space\n        self.feature_extractor = EEGFeatureExtractor(in_features, feature_dim=feature_dim)\n\n        # Temporal encoder models temporal dependencies across timesteps\n        self.temporal_encoder = TemporalTransformer(\n            neural_dim=feature_dim,\n            d_model=d_model,\n            n_days=n_days,\n            n_tokens=n_tokens if n_tokens is not None else vocab_size,\n            n_layers=n_layers,\n            n_heads=n_heads,\n            ff_mult=ff_mult,\n            dropout=dropout,\n            patch_stride=patch_stride,\n            max_seq_len=max_seq_len\n        )\n\n        # Decoder converts encoded sequence into text logits\n        self.decoder = CTCDecoder(d_model=d_model, vocab_size=vocab_size)\n\n    def forward(self, input_features, day_idx):\n        x = self.feature_extractor(input_features)\n        x = self.temporal_encoder(x, day_idx)\n        logits = self.decoder(x)\n        return logits","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# EEG Conformer","metadata":{}},{"cell_type":"code","source":"from einops.layers.torch import Rearrange\n\n# Convolution module\n# use conv to capture local features, instead of postion embedding.\nclass PatchEmbedding(nn.Module):\n    \"\"\"\n    Convert (B, T, F) -> (B, T//time_pool, E)\n    Designed to preserve time resolution by using convs that operate across features\n    and downsample the time axis in a controlled way.\n    \"\"\"\n    def __init__(self, n_features=512, emb_size=64, time_pool=4, kernel_size=[7,3]):\n        super().__init__()\n        self.time_pool = time_pool\n        layers = []\n\n        layers.extend([\n            nn.Conv1d(n_features, emb_size, kernel_size=kernel_size[0], stride=1, padding=3),\n            nn.BatchNorm1d(emb_size),\n            nn.ELU(),\n        ])\n        \n        for size in kernel_size[1:-1]:\n            layers.extend([\n                nn.Conv1d(emb_size, emb_size, kernel_size=size, stride=1, \n                         padding=size//2),\n                nn.BatchNorm1d(emb_size),\n                nn.ELU(),\n            ])\n\n        layers.append(\n            nn.Conv1d(emb_size, emb_size, kernel_size=kernel_size[-1], stride=time_pool, padding=1)\n        )\n        self.proj = nn.Sequential(*layers)\n\n        print(f\"PatchEmbedding with {len(layers)} layers\")\n\n    def forward(self, x):\n        # x: (B, T, F)\n        x = x.transpose(1, 2)  # (B, F, T)\n        x = self.proj(x)       # (B, E, T//time_pool)\n        x = x.transpose(1, 2)  # (B, T//time_pool, E)\n        return x\n\nclass EEGConformer(nn.Module):\n    def __init__(\n        self,\n        neural_dim=512,\n        d_model=64,\n        nhead=1,\n        num_encoder_layers=1,\n        dim_feedforward=256,\n        dropout=0.1,\n        n_days=45,\n        n_classes=41,\n        time_pool=4,\n        kernel_size=[7,3],\n    ):\n        super().__init__()\n        self.d_model = d_model\n        self.time_pool = time_pool\n        self.n_classes = n_classes\n        self.patch = PatchEmbedding(n_features=neural_dim, emb_size=d_model, time_pool=time_pool, kernel_size=kernel_size)\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d_model,\n            nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout,\n            batch_first=True,\n        )\n        self.transformer_encoder = nn.TransformerEncoder(\n            encoder_layer, num_layers=num_encoder_layers\n        )\n        self.output_projection = nn.Linear(d_model, n_classes)\n\n    def forward(self, x, day_idx):\n        # x: (B, T, 512)\n        x = self.patch(x)  # (B, T//time_pool, d_model)\n        x = self.transformer_encoder(x)\n        logits = self.output_projection(x)  # (B, T//time_pool, n_classes)\n        return logits","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-21T19:02:21.453992Z","iopub.execute_input":"2025-11-21T19:02:21.454495Z","iopub.status.idle":"2025-11-21T19:02:21.462765Z","shell.execute_reply.started":"2025-11-21T19:02:21.454472Z","shell.execute_reply":"2025-11-21T19:02:21.461932Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\n\nimport h5py\nimport torch\nfrom torch.nn.utils.rnn import pad_sequence\nfrom torch.utils.data import Dataset\n\n\nclass BrainToTextDataset(Dataset):\n    def __init__(\n        self,\n        data_dir: str,\n        dataset_key: str = \"train\",\n        max_sessions: int = None,\n    ):\n        if dataset_key in [\"train\", \"val\", \"test\"]:\n            all_paths = sorted(Path(data_dir).rglob(f\"*/data_{dataset_key}.hdf5\"))\n            \n            if max_sessions is not None:\n                self.file_paths = all_paths[:max_sessions]\n            else:\n                self.file_paths = all_paths\n        else:\n            self.file_paths = []\n\n        self.session_info = []\n        self.total_trials = None\n        self._load_session_info()\n\n    def _load_session_info(self) -> None:\n        total_trials = 0\n\n        for day_idx, path in enumerate(self.file_paths):\n            with h5py.File(path, \"r\") as f:\n                num_trials = len(list(f.keys()))\n                self.session_info.append(\n                    {\n                        \"day_idx\": day_idx,\n                        \"file_path\": path,\n                        \"num_trials\": num_trials,\n                        \"trial_offset\": total_trials,\n                    }\n                )\n                total_trials += num_trials\n\n        self.total_trials = total_trials\n\n    def __len__(self):\n        return self.total_trials\n\n    def __getitem__(self, idx):\n        for s in self.session_info:\n            if idx < s[\"trial_offset\"] + s[\"num_trials\"]:\n                trial_idx = idx - s[\"trial_offset\"]\n                day_idx = s[\"day_idx\"]\n                file_path = s[\"file_path\"]\n                break\n\n        with h5py.File(file_path, \"r\") as f:\n            trial_key = f\"trial_{trial_idx:04d}\"\n            g = f[trial_key]\n\n            return {\n                \"input_features\": torch.from_numpy(g[\"input_features\"][:]).float(),\n                \"seq_class_ids\": torch.from_numpy(g[\"seq_class_ids\"][:]).long(),\n                \"phone_seq_lens\": g.attrs[\"seq_len\"],\n                \"n_time_steps\": g.attrs[\"n_time_steps\"],\n                \"day_idx\": day_idx,\n            }\n\n\ndef collate_fn(batch: list) -> dict:\n    return {\n        \"input_features\": pad_sequence(\n            [b[\"input_features\"] for b in batch], batch_first=True\n        ),\n        \"seq_class_ids\": pad_sequence(\n            [b[\"seq_class_ids\"] for b in batch], batch_first=True\n        ),\n        \"phone_seq_lens\": torch.tensor([b[\"phone_seq_lens\"] for b in batch]),\n        \"n_time_steps\": torch.tensor([b[\"n_time_steps\"] for b in batch]),\n        \"day_idxs\": torch.tensor([b[\"day_idx\"] for b in batch]),\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-21T19:02:23.709864Z","iopub.execute_input":"2025-11-21T19:02:23.710128Z","iopub.status.idle":"2025-11-21T19:02:23.721854Z","shell.execute_reply.started":"2025-11-21T19:02:23.710108Z","shell.execute_reply":"2025-11-21T19:02:23.721088Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train loop","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\n\n\ndef train_step(\n    model: nn.Module,\n    dataloader: DataLoader,\n    optimizer: torch.optim,\n    device: torch.device,\n) -> float:\n    model.train()\n    total_loss = 0\n    ctc_loss = nn.CTCLoss(blank=0, zero_infinity=True)\n\n    pbar = tqdm(dataloader, desc=\"Training\", leave=False)\n\n    for batch in pbar:\n        features = batch[\"input_features\"].to(device)\n        labels = batch[\"seq_class_ids\"].to(device)\n        n_time_steps = batch[\"n_time_steps\"].to(device)\n        phone_seq_lens = batch[\"phone_seq_lens\"].to(device)\n        day_indicies = batch[\"day_idxs\"].to(device)\n\n        # smooth neural data\n        #features_smooth = gauss_smooth(\n        #    features,\n        #    device,\n        #    smooth_kernel_std=2,\n        #    smooth_kernel_size=100,\n        #    padding=\"valid\",\n        #)\n\n        features_aug, n_time_steps = transform_data(features, n_time_steps, 'train', device=device)\n\n        if isinstance(model,EEGConformer):\n            adjusted_lens = (n_time_steps - 100 + 1) // model.time_pool\n        else:\n            adjusted_lens = n_time_steps - 100 + 1\n\n        #logits = model(features_smooth, day_idx=day_indicies)\n        logits = model(features_aug, day_idx=day_indicies)\n        \n\n        # CTC Loss\n        log_probs = logits.log_softmax(2).permute(1, 0, 2)\n        loss = ctc_loss(log_probs, labels, adjusted_lens, phone_seq_lens)\n\n        optimizer.zero_grad()\n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0)\n        optimizer.step()\n\n        batch_loss = loss.item()\n        total_loss += batch_loss\n        pbar.set_postfix({\"loss\": f\"{batch_loss:.4f}\"})\n\n    return total_loss / len(dataloader)\n\n\ndef save_checkpoint(\n    model: nn.Module,\n    optimizer: torch.optim,\n    epoch: int,\n    val_per: float,\n    save_path: Path,\n) -> None:\n    checkpoint = {\n        \"epoch\": epoch,\n        \"model_state_dict\": model.state_dict(),\n        \"optimizer_state_dict\": optimizer.state_dict(),\n        \"val_per\": val_per,\n    }\n    torch.save(checkpoint, save_path)\n\n    print(f\"Checkpoint saved to {save_path}\")\n\n\ndef load_checkpoint(\n    model: nn.Module,\n    optimizer: torch.optim,\n    checkpoint_path: Path,\n) -> tuple:\n    checkpoint = torch.load(checkpoint_path)\n    model.load_state_dict(checkpoint[\"model_state_dict\"])\n    optimizer.load_state_dict(checkpoint[\"optimizer_state_dict\"])\n    epoch = checkpoint[\"epoch\"]\n    val_per = checkpoint[\"val_per\"]\n\n    print(f\"Loaded checkpoint from epoch {epoch}, PER: {val_per:.4f}\")\n\n    return epoch, val_per\n\n\ndef evaluate(model: nn.Module, \n             dataloader: DataLoader, \n             device: torch.device,\n            ) -> tuple:\n    model.eval()\n    total_loss = 0\n    total_edit_distance = 0\n    total_phonemes = 0\n    ctc_loss = nn.CTCLoss(blank=0, zero_infinity=True)\n\n    pbar = tqdm(dataloader, desc=\"Evaluating\", leave=False)\n\n    with torch.no_grad():\n        for batch in pbar:\n            features = batch[\"input_features\"].to(device)\n            labels = batch[\"seq_class_ids\"].to(device)\n            n_time_steps = batch[\"n_time_steps\"].to(device)\n            phone_seq_lens = batch[\"phone_seq_lens\"].to(device)\n            day_indicies = batch[\"day_idxs\"].to(device)\n\n            features_smooth = gauss_smooth(\n                features,\n                device,\n                smooth_kernel_std=2,\n                smooth_kernel_size=100,\n                padding=\"valid\",\n            )\n\n            if isinstance(model,EEGConformer):\n                adjusted_lens = (n_time_steps - 100 + 1) // model.time_pool\n            else:\n                adjusted_lens = n_time_steps - 100 + 1\n\n\n            logits = model(features_smooth, day_idx=day_indicies)\n\n            log_probs = logits.log_softmax(2).permute(1, 0, 2)\n            loss = ctc_loss(log_probs, labels, adjusted_lens, phone_seq_lens)\n            total_loss += loss.item()\n\n            for i in range(logits.shape[0]):\n                pred_seq = torch.argmax(logits[i, : adjusted_lens[i]], dim=-1)\n                pred_seq = torch.unique_consecutive(pred_seq)\n                pred_seq = pred_seq[pred_seq != 0]\n\n                true_seq = labels[i, : phone_seq_lens[i]]\n\n                edit_dist = edit_distance(\n                    pred_seq.cpu().numpy(), true_seq.cpu().numpy()\n                )\n\n                total_edit_distance += edit_dist\n                total_phonemes += phone_seq_lens[i].item()\n\n            current_per = (\n                total_edit_distance / total_phonemes if total_phonemes > 0 else 0\n            )\n            pbar.set_postfix({\"PER\": f\"{current_per:.4f}\"})\n\n    avg_loss = total_loss / len(dataloader)\n    per = total_edit_distance / total_phonemes\n\n    return avg_loss, per\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-21T19:04:34.266214Z","iopub.execute_input":"2025-11-21T19:04:34.266998Z","iopub.status.idle":"2025-11-21T19:04:34.283015Z","shell.execute_reply.started":"2025-11-21T19:04:34.266967Z","shell.execute_reply":"2025-11-21T19:04:34.282443Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def edit_distance(seq1: np.ndarray, seq2: np.ndarray) -> int:\n    len1, len2 = len(seq1), len(seq2)\n\n    dp = np.zeros((len1 + 1, len2 + 1), dtype=int)\n\n    for i in range(len1 + 1):\n        dp[i][0] = i\n    for j in range(len2 + 1):\n        dp[0][j] = j\n\n    for i in range(1, len1 + 1):\n        for j in range(1, len2 + 1):\n            if seq1[i - 1] == seq2[j - 1]:\n                dp[i][j] = dp[i - 1][j - 1]\n            else:\n                dp[i][j] = 1 + min(dp[i - 1][j], dp[i][j - 1], dp[i - 1][j - 1])\n\n    return dp[len1][len2]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-21T19:03:24.322627Z","iopub.execute_input":"2025-11-21T19:03:24.323249Z","iopub.status.idle":"2025-11-21T19:03:24.328593Z","shell.execute_reply.started":"2025-11-21T19:03:24.323223Z","shell.execute_reply":"2025-11-21T19:03:24.327886Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Loading","metadata":{}},{"cell_type":"code","source":"import random\n\ndef set_all_seeds(seed: int) -> None:\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    np.random.seed(seed)\n    random.seed(seed)\n\n\nset_all_seeds(42)\n\ntrain_dataset = BrainToTextDataset(\n    data_dir=\"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final\",\n    dataset_key=\"train\",\n    #max_sessions=11,\n)\nval_dataset = BrainToTextDataset(\n    data_dir=\"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final\",\n    dataset_key=\"val\",\n    #max_sessions=10,\n)\n\ntrain_loader = DataLoader(\n    train_dataset, batch_size=32, shuffle=True, collate_fn=collate_fn\n)\nval_loader = DataLoader(val_dataset, batch_size=32, collate_fn=collate_fn)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-21T19:09:50.738305Z","iopub.execute_input":"2025-11-21T19:09:50.738999Z","iopub.status.idle":"2025-11-21T19:09:55.177996Z","shell.execute_reply.started":"2025-11-21T19:09:50.738974Z","shell.execute_reply":"2025-11-21T19:09:55.177413Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Result","metadata":{}},{"cell_type":"code","source":"from torch.optim.lr_scheduler import ReduceLROnPlateau\n\nn_sessions = len(train_dataset.session_info)\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\"\"\"\nmodel = GRUDecoder(\n    neural_dim=512,\n    n_units=512,\n    n_days=n_sessions,\n    n_classes=41,\n    rnn_dropout=0.4,\n    input_dropout=0.4,\n    n_layers=3,\n    patch_size=0,\n    patch_stride=0,\n).to(device)\n\"\"\"\n\n\"\"\"\nmodel = LSTMDecoder(\n    neural_dim=512,\n    n_units=512,\n    n_days=n_sessions,\n    n_classes=41,\n    rnn_dropout=0.3,\n    input_dropout=0.3,\n    n_layers=2,\n    patch_size=0,\n    patch_stride=0,\n).to(device)\n\"\"\"\n\n\"\"\"\nmodel = Transformer(\n    neural_dim=512,\n    d_model=384,\n    nhead=8,\n    num_encoder_layers=4,\n    dim_feedforward=1024,\n    dropout=0.15,\n    n_days=n_sessions,\n    n_classes=41,\n).to(device)\n\"\"\"\n\n\"\"\"\nmodel = EEGToTextModel(\n    in_features=64,\n    vocab_size=30,\n    n_layers=4,     \n    n_heads=8,\n    d_model=256\n)\n\"\"\"\n\nmodel = EEGConformer(\n    neural_dim=512,\n    d_model=128,\n    nhead=4,\n    num_encoder_layers=2,\n    dim_feedforward=256,\n    dropout=0.1,\n    n_days=45,\n    n_classes=41,\n    time_pool=4,\n    kernel_size=[9,7,5,3,3],\n).to(device)\n\noptimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)\nscheduler = ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=3)\n\nprint(f\"Model has {sum(p.numel() for p in model.parameters()):,} parameters\")\nprint(f\"Training on {n_sessions} sessions\")\nprint(f\"Total train trials: {train_dataset.total_trials}\")\nprint(f\"Total val trials: {val_dataset.total_trials}\")\n\ncheckpoint_dir = Path(\"/kaggle/working/\")\ncheckpoint_dir.mkdir(exist_ok=True, parents=True)\n\nbest_per = float('inf')\nstart_epoch = 0\ntrain_losses = []\nval_losses = []\nval_pers = []\n\nresume_from = None\n\nif resume_from and resume_from.exists():\n    start_epoch, best_per = load_checkpoint(model, optimizer, resume_from)\n    start_epoch += 1\n\nfor epoch in range(start_epoch, 50):\n    train_loss = train_step(model, train_loader, optimizer, device)\n    val_loss, val_per = evaluate(model, val_loader, device)\n\n    scheduler.step(val_per)\n\n    print(\n        f\"epoch {epoch+1} train Loss={train_loss:.4f} val Loss={val_loss:.4f} PER={val_per:.4f}\"\n    )\n    train_losses.append(train_loss)\n    val_losses.append(val_loss)\n    val_pers.append(val_per)\n\n    if val_per < best_per:\n        best_per = val_per\n        save_checkpoint(\n            model, optimizer, epoch, val_per, checkpoint_dir / \"best_model.pt\"\n        )\n        print(f\"    new best PER: {best_per:.4f}\")\n\n    if (epoch + 1) % 5 == 0:\n        save_checkpoint(\n            model,\n            optimizer,\n            epoch,\n            val_per,\n            checkpoint_dir / f\"checkpoint_epoch_{epoch+1}.pt\",\n        )\n\n    save_checkpoint(\n        model, optimizer, epoch, val_per, checkpoint_dir / \"last_checkpoint.pt\"\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-21T19:09:58.592666Z","iopub.execute_input":"2025-11-21T19:09:58.592934Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Plot loss","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\nmin_val_i = np.argmin(val_losses)\nbest_per_i = np.argmin(val_pers)\n\nfig, ax1 = plt.subplots(figsize=(8, 5))\n\nax1.plot(train_losses, label=\"Train Loss\", color=\"tab:blue\", linewidth=2)\nax1.plot(val_losses, label=\"Val Loss\", color=\"tab:orange\", linewidth=2)\nax1.scatter(min_val_i, val_losses[min_val_i], label=f\"Val Min: {val_losses[min_val_i]}\", \n            color=\"tab:orange\", marker='o', s=50)\nax1.set_xlabel(\"Epoch\", fontsize=12)\nax1.set_ylabel(\"Loss\", color=\"tab:blue\", fontsize=12)\nax1.tick_params(axis='y', labelcolor=\"tab:blue\")\n\nax2 = ax1.twinx()\nax2.plot(val_pers, label=\"Val PER\", color=\"tab:green\", linewidth=2, linestyle=\"--\")\nax2.scatter(best_per_i, val_pers[best_per_i], label=f\"Best PER: {val_pers[best_per_i]}\", \n            color=\"tab:green\", marker='*', s=50)\nax2.set_ylabel(\"PER\", color=\"tab:green\", fontsize=12)\nax2.tick_params(axis='y', labelcolor=\"tab:green\")\n\nlines_1, labels_1 = ax1.get_legend_handles_labels()\nlines_2, labels_2 = ax2.get_legend_handles_labels()\nax1.legend(lines_1 + lines_2, labels_1 + labels_2, loc=\"upper right\")\n\nplt.title(\"Training Progress\", fontsize=14)\nplt.grid(True, linestyle=\"--\", alpha=0.5)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-21T18:54:35.927027Z","iopub.status.idle":"2025-11-21T18:54:35.927407Z","shell.execute_reply.started":"2025-11-21T18:54:35.927223Z","shell.execute_reply":"2025-11-21T18:54:35.927237Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"try:\n    num_epochs = len(train_losses)\n    \n    # Calculate key metrics\n    min_val_i = np.argmin(val_losses)\n    best_per_i = np.argmin(val_pers)\n    \n    print(f\"{'Epoch':<5} | {'Train Loss':<12} | {'Val Loss':<10} | {'Val PER':<10}\")\n    print(\"-\" * 43)\n\n    # Iterate through the data and print each row\n    for epoch in range(num_epochs):\n        # Use f-strings for clean alignment and precision\n        print(f\"{epoch:<5} | {train_losses[epoch]:<12.6f} | {val_losses[epoch]:<10.6f} | {val_pers[epoch]:<10.6f}\")\n\n    print(\"-\" * 43)\n    \n    # Print the summary statistics\n    print(f\"Minimum Validation Loss: {val_losses[min_val_i]:.6f} (Epoch {min_val_i})\")\n    print(f\"Best Validation PER:     {val_pers[best_per_i]:.6f} (Epoch {best_per_i})\")\n\nexcept NameError:\n    print(\"Error: Please ensure 'train_losses', 'val_losses', and 'val_pers' are defined and accessible.\")\nexcept Exception as e:\n    print(f\"An error occurred: {e}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Final Evaluation","metadata":{}},{"cell_type":"code","source":"def levenshtein_align(pred, true):\n    \"\"\"\n    Returns lists of aligned pairs (p, t) where:\n      p = predicted phoneme or None (for deletion)\n      t = true phoneme or None (for insertion)\n    \"\"\"\n    import numpy as np\n\n    n, m = len(pred), len(true)\n    dp = np.zeros((n+1, m+1), dtype=int)\n\n    # initialize borders\n    for i in range(1, n+1): dp[i, 0] = i\n    for j in range(1, m+1): dp[0, j] = j\n\n    # fill DP matrix\n    for i in range(1, n+1):\n        for j in range(1, m+1):\n            cost = 0 if pred[i-1] == true[j-1] else 1\n            dp[i, j] = min(\n                dp[i-1, j] + 1,        # deletion\n                dp[i, j-1] + 1,        # insertion\n                dp[i-1, j-1] + cost    # match/substitution\n            )\n\n    # backtrace\n    i, j = n, m\n    alignment = []\n    while i > 0 or j > 0:\n        if i > 0 and dp[i, j] == dp[i-1, j] + 1:\n            alignment.append((pred[i-1], None))\n            i -= 1\n        elif j > 0 and dp[i, j] == dp[i, j-1] + 1:\n            alignment.append((None, true[j-1]))\n            j -= 1\n        else:\n            alignment.append((pred[i-1], true[j-1]))\n            i -= 1\n            j -= 1\n\n    return alignment[::-1]\n\ndef plot_confusion_matrix(cm, phoneme_labels=None, normalize=\"true\"):\n    \"\"\"\n    Visualize a phoneme confusion matrix.\n\n    Args:\n        cm (Tensor or ndarray): Confusion matrix (pred x true)\n        phoneme_labels (list of str): Optional list mapping class IDs → phoneme symbols.\n                                      Length must be num_classes+1 (including blank at index 0).\n        normalize:\n            None        → raw counts\n            \"true\"      → normalize each column (per true phoneme)\n            \"pred\"      → normalize each row (per predicted phoneme)\n            \"all\"       → normalize entire matrix\n    \"\"\"\n    import seaborn as sns\n\n    cm = cm.cpu().numpy() if hasattr(cm, \"cpu\") else cm\n\n    # normalization\n    cm_norm = cm.astype(float)\n    if normalize == \"true\":\n        cm_norm = cm_norm / (cm_norm.sum(axis=1, keepdims=True) + 1e-9)\n    elif normalize == \"pred\":\n        cm_norm = cm_norm / (cm_norm.sum(axis=0, keepdims=True) + 1e-9)\n    elif normalize == \"all\":\n        cm_norm = cm_norm / cm_norm.sum()\n\n    # labels\n    if phoneme_labels is None:\n        phoneme_labels = [f\"{i}\" for i in range(cm.shape[0])]\n    else:\n        print(f'len phon: {len(phoneme_labels)}, cm.shape[0]: {cm.shape[0]}')\n        assert len(phoneme_labels) == cm.shape[0], \\\n            \"phoneme_labels must have length num_classes+1 (including blank=0).\"\n\n    plt.figure(figsize=(14, 12))\n\n    sns.heatmap(\n        cm_norm,\n        square=True,\n        cmap=\"viridis\",\n        xticklabels=phoneme_labels,\n        yticklabels=phoneme_labels,\n        cbar=True,\n        linewidths=0.05\n    )\n\n    plt.title(\"Phoneme Confusion Matrix\")\n    plt.xlabel(\"Predicted phoneme\")\n    plt.ylabel(\"True phoneme\")\n    plt.tight_layout()\n    plt.show()\n\n\ndef final_evaluation(model: nn.Module, testing_dataloader: DataLoader, device: torch.device):\n    \"\"\"\n    Evaluate the best model on the testing set and compute:\n      - CTC loss\n      - PER\n      - Full phoneme confusion matrix\n\n    model.num_classes must be defined (excluding blank).\n    \"\"\"\n\n    model.eval()\n    ctc_loss = nn.CTCLoss(blank=0, zero_infinity=True)\n\n    total_loss = 0\n    total_edit_distance = 0\n    total_phonemes = 0\n\n    # Confusion matrix: rows = predicted, columns = true\n    num_classes = model.n_classes   # e.g., 39 phoneme classes (blank=0)\n    cm = torch.zeros((num_classes, num_classes), dtype=torch.int64)\n\n    with torch.no_grad():\n        pbar = tqdm(testing_dataloader, desc=\"Final Evaluation\", leave=True)\n\n        for batch in pbar:\n            features = batch[\"input_features\"].to(device)\n            labels = batch[\"seq_class_ids\"].to(device)\n            n_time_steps = batch[\"n_time_steps\"].to(device)\n            phone_seq_lens = batch[\"phone_seq_lens\"].to(device)\n            day_indicies = batch[\"day_idxs\"].to(device)\n\n            # smoothing for consistency\n            features_smooth = gauss_smooth(\n                features,\n                device,\n                smooth_kernel_std=2,\n                smooth_kernel_size=100,\n                padding=\"valid\",\n            )\n            if isinstance(model,EEGConformer):\n                adjusted_lens = (n_time_steps - 100 + 1) // model.time_pool\n            else:\n                adjusted_lens = n_time_steps - 100 + 1\n\n            # forward pass\n            logits = model(features_smooth, day_idx=day_indicies)\n\n            log_probs = logits.log_softmax(2).permute(1, 0, 2)\n            loss = ctc_loss(log_probs, labels, adjusted_lens, phone_seq_lens)\n            total_loss += loss.item()\n\n            # sequence-level evaluation\n            for i in range(logits.shape[0]):\n\n                # CTC greedy decode\n                pred_seq = torch.argmax(logits[i, : adjusted_lens[i]], dim=-1)\n                pred_seq = torch.unique_consecutive(pred_seq)\n                pred_seq = pred_seq[pred_seq != 0]  # remove CTC blank\n                pred_seq = pred_seq.cpu().tolist()\n\n                # true sequence\n                true_seq = labels[i, : phone_seq_lens[i]].cpu().tolist()\n\n                # edit distance\n                edit_dist = edit_distance(\n                    np.array(pred_seq), np.array(true_seq)\n                )\n                total_edit_distance += edit_dist\n                total_phonemes += len(true_seq)\n\n                # update confusion matrix\n                alignment = levenshtein_align(pred_seq, true_seq)\n                for p, t in alignment:\n                    if t is None:\n                        cm[0, p] += 1      # deletion\n                    elif p is None:\n                        cm[t, 0] += 1      # insertion\n                    else:\n                        cm[t, p] += 1      # match/substitution\n\n            current_per = (\n                total_edit_distance / total_phonemes if total_phonemes > 0 else 0\n            )\n            pbar.set_postfix({\"PER\": f\"{current_per:.4f}\"})\n\n    avg_loss = total_loss / len(testing_dataloader)\n    per = total_edit_distance / total_phonemes\n\n    return avg_loss, per, cm","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"phoneme_labels = [\n    \"BLANK\", # for CTC loss\n    \"SIL\",  # silence\n    \"AA\", \"AE\", \"AH\",\"AO\", \"AW\", \"AY\",\n    \"B\", \"CH\", \"D\", \"DH\", \"EH\", \"ER\",\n    \"EY\", \"F\", \"G\", \"HH\", \"IH\", \"IY\",\n    \"JH\", \"K\", \"L\", \"M\", \"N\", \"NG\",\n    \"OW\", \"OY\", \"P\", \"R\", \"S\", \"SH\",\n    \"T\", \"TH\", \"UH\", \"UW\", \"V\", \"W\",\n    \"Y\", \"Z\", \"ZH\",\n]","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_model = EEGConformer(\n    neural_dim=512,\n    d_model=128,\n    nhead=4,\n    num_encoder_layers=2,\n    dim_feedforward=256,\n    dropout=0.1,\n    n_days=45,\n    n_classes=41,\n    time_pool=4,\n    kernel_size=[9,7,5,3,3],\n).to(device)\n\ncheckpoint = torch.load(\"/kaggle/working/best_model.pt\", weights_only=False)  # include metadata\nbest_model.load_state_dict(checkpoint[\"model_state_dict\"])\n\navg_loss, per, cm = final_evaluation(best_model, val_loader, device)\n\n# Optionally remove the BLANK class from the confusion matrix and labels\n# Note: BLANK (index 0) is still included in CTC loss calculations,\n#       but we may want to exclude it from visualization for clarity\ntrim_BLANK = True\nif trim_BLANK:\n    cm_ = cm[1:, 1:].clone()\n    phoneme_labels_ = phoneme_labels[1:] \nelse:\n    cm_ = cm.clone()\n    phoneme_labels_ = phoneme_labels.copy()\n\n# cm_trimmed = cm[1:,1:] #remove blank for insertion/deletion\nprint(\"Test CTC Loss:\", avg_loss)\nprint(\"Test PER:\", per)\nprint(\"Confusion matrix shape:\", cm_.shape)\n\nplot_confusion_matrix(cm_, phoneme_labels_)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport gc\n\ntorch.cuda.empty_cache()\ngc.collect()\n\nprint(torch.cuda.memory_summary())","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}