{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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":31192,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Brain-to-Text: A Contrastive Learning Framework\n\nThis notebook implements the two-stage deep learning framework\n\n* Stage 1: Contrastive pre-training to align neural and phoneme embeddings.\n* Stage 2: Training a Transformer decoder on the frozen neural embeddings.\n\nEvaluation: Calculating Word Error Rate (WER) and performing an ablation study.\n\nVisualization: Using t-SNE to assess the learned embedding space alignment.","metadata":{}},{"cell_type":"code","source":"!pip install jiwer","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-03T14:48:22.697987Z","iopub.execute_input":"2025-12-03T14:48:22.698225Z","iopub.status.idle":"2025-12-03T14:48:28.580989Z","shell.execute_reply.started":"2025-12-03T14:48:22.698207Z","shell.execute_reply":"2025-12-03T14:48:28.580257Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Core PyTorch ---\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.nn.utils.rnn import pad_sequence\n\n# --- Data Handling & Utilities ---\nimport h5py\nimport numpy as np\nimport pandas as pd\nimport glob\nimport os\nimport math\nfrom typing import List, Tuple, Dict\nfrom tqdm.notebook import tqdm  # Use notebook-friendly tqdm\n\n# --- Evaluation & Visualization ---\nimport jiwer\nfrom sklearn.manifold import TSNE\nimport matplotlib.pyplot as plt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T14:48:28.582439Z","iopub.execute_input":"2025-12-03T14:48:28.582678Z","iopub.status.idle":"2025-12-03T14:48:33.227890Z","shell.execute_reply.started":"2025-12-03T14:48:28.582650Z","shell.execute_reply":"2025-12-03T14:48:33.227268Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Configuration (FINAL RUN) ---\n\n# -- Data Config --\nDATA_DIR = \"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final/\"\n\n# -- Model Config (HIGH CAPACITY) --\n# Increased dimensions to capture complex neural patterns\nEMBEDDING_DIM = 512          \nNEURAL_CHANNELS = 512        \nNEURAL_CNN_OUT = 128         # Doubled filters\nNEURAL_LSTM_HIDDEN = 256     # Doubled memory\n\nPHONEME_VOCAB_SIZE = 41 + 1\nPHONEME_EMBED_DIM = 128      # Doubled\nPHONEME_LSTM_HIDDEN = 256    # Matched\n\nDECODER_NHEAD = 8            \nDECODER_NLAYERS = 6          \nDECODER_DIM_FEEDFORWARD = 2048 # Increased 4x for better language modeling\n\n# -- Training Config --\nBATCH_SIZE = 16              \nNUM_EPOCHS_STAGE1 = 80       # Optimized for 12hr limit\nNUM_EPOCHS_STAGE2 = 140      # Optimized for 12hr limit\nLR_STAGE1 = 1e-4             \nLR_STAGE2 = 1e-4             \nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n# -- Checkpoint Paths --\nCHECKPOINT_DIR = \"/kaggle/working/\"\nENCODER_CHECKPOINT_PATH = os.path.join(CHECKPOINT_DIR, \"neural_encoder_stage1.pth\")\nPHONEME_ENCODER_CHECKPOINT_PATH = os.path.join(CHECKPOINT_DIR, \"phoneme_encoder_stage1.pth\")\nDECODER_CHECKPOINT_PATH = os.path.join(CHECKPOINT_DIR, \"full_decoder_model_stage2.pth\")\n\n# -- Tokenizer Config --\nPAD_TOKEN_ID = 0\nSOS_TOKEN_ID = 1\nEOS_TOKEN_ID = 2\nPHONEME_PAD_ID = 0\n\nprint(f\"Using device: {DEVICE}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T14:48:33.228611Z","iopub.execute_input":"2025-12-03T14:48:33.229000Z","iopub.status.idle":"2025-12-03T14:48:33.295078Z","shell.execute_reply.started":"2025-12-03T14:48:33.228972Z","shell.execute_reply":"2025-12-03T14:48:33.294380Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Text Tokenizer Class ---\n\nclass SimpleTokenizer:\n    \"\"\"A simple word-level tokenizer for text transcripts.\"\"\"\n    def __init__(self):\n        # <PAD> is 0, <SOS> is 1, <EOS> is 2, <UNK> is 3\n        self.vocab = {'<PAD>': 0, '<SOS>': 1, '<EOS>': 2, '<UNK>': 3}\n        self.inv_vocab = {v: k for k, v in self.vocab.items()}\n        self.vocab_size = 4\n\n    def build_vocab(self, all_sentences: List[str]):\n        \"\"\"Builds a vocabulary from a list of sentences.\"\"\"\n        words = set()\n        for s in all_sentences:\n            words.update(s.lower().split())\n        \n        for word in sorted(list(words)): # Sorted for consistency\n            if word not in self.vocab:\n                self.vocab[word] = self.vocab_size\n                self.inv_vocab[self.vocab_size] = word\n                self.vocab_size += 1\n\n    def encode(self, text: str) -> torch.LongTensor:\n        \"\"\"Converts a string to a tensor of token IDs.\"\"\"\n        tokens = [self.vocab.get(w, self.vocab['<UNK>']) for w in text.lower().split()]\n        return torch.LongTensor([self.vocab['<SOS>']] + tokens + [self.vocab['<EOS>']])\n\n    def decode(self, tokens: torch.LongTensor) -> str:\n        \"\"\"Converts a tensor of token IDs back to a string.\"\"\"\n        words = [self.inv_vocab.get(t.item(), '?') for t in tokens]\n        # Filter out special tokens\n        words = [w for w in words if w not in ['<PAD>', '<SOS>', '<EOS>']]\n        return \" \".join(words)\n\n    def get_vocab_size(self):\n        return self.vocab_size","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T14:48:33.296721Z","iopub.execute_input":"2025-12-03T14:48:33.296903Z","iopub.status.idle":"2025-12-03T14:48:33.319151Z","shell.execute_reply.started":"2025-12-03T14:48:33.296888Z","shell.execute_reply":"2025-12-03T14:48:33.318464Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Data Loading Classes (MODIFIED) ---\n\nclass BrainToTextDataset(Dataset):\n    \"\"\"\n    Custom PyTorch Dataset for the Brain-to-Text HDF5 files.\n    \"\"\"\n    def __init__(self, hdf5_files: List[str]):\n        super().__init__()\n        self.trial_index = [] # Stores (file_path, trial_key)\n        \n        for file_path in hdf5_files:\n            try:\n                with h5py.File(file_path, 'r') as f:\n                    trial_keys = list(f.keys())\n                    for key in trial_keys:\n                        self.trial_index.append((file_path, key))\n            except Exception as e:\n                print(f\"Warning: Could not read {file_path}. Error: {e}\")\n                \n        print(f\"Found {len(self.trial_index)} total trials in {len(hdf5_files)} files.\")\n\n    def __len__(self) -> int:\n        return len(self.trial_index)\n\n    def __getitem__(self, idx: int) -> Dict:\n        \"\"\"MODIFIED: Returns a dictionary.\"\"\"\n        file_path, trial_key = self.trial_index[idx]\n        \n        with h5py.File(file_path, 'r') as f:\n            trial_group = f[trial_key]\n            \n            # 1. Load Neural Data (T_neural, 512)\n            neural_data = torch.from_numpy(trial_group['input_features'][:]).float()\n            \n            # --- MODIFICATION: Check if this is a test sample ---\n            # Test samples in the baseline [cite: 160-165] might lack target keys\n            is_test = 'seq_class_ids' not in trial_group\n            \n            if not is_test:\n                # 2. Load Phoneme Data (T_phoneme)\n                phoneme_data = torch.from_numpy(trial_group['seq_class_ids'][:]).long()\n                \n                # 3. Load Text Data (string)\n                if 'sentence_label' in trial_group.attrs:\n                    text_data = trial_group.attrs['sentence_label']\n                elif 'transcription' in trial_group:\n                    text_data = trial_group['transcription'][()].decode('utf-8')\n                else:\n                    text_data = \"\" # Should not happen\n            else:\n                # Create dummy placeholders for test set\n                phoneme_data = torch.tensor([], dtype=torch.long)\n                text_data = \"\"\n\n        # Transpose neural data to (Channels, Time) for 1D CNN\n        neural_data = neural_data.permute(1, 0)\n            \n        return {\n            \"neural_data\": neural_data,\n            \"phoneme_data\": phoneme_data,\n            \"text_data\": text_data,\n            \"trial_key\": trial_key, # Pass the key for submission\n            \"is_test\": is_test # Flag if it's a test sample\n        }\n\n    def get_all_sentences(self) -> List[str]:\n        \"\"\"Helper function to get all text for building the tokenizer vocab.\"\"\"\n        all_sentences = []\n        for i in range(len(self)):\n            try:\n                item = self.__getitem__(i)\n                if not item[\"is_test\"]:\n                    all_sentences.append(item[\"text_data\"])\n            except Exception as e:\n                print(f\"Warning: Skipping trial {i} due to read error: {e}\")\n        return all_sentences\n\n\nclass BrainCollator:\n    \"\"\"\n    Custom collate_fn to pad variable-length sequences.\n    \"\"\"\n    def __init__(self, tokenizer: SimpleTokenizer, pad_token_id: int, phoneme_pad_id: int):\n        self.tokenizer = tokenizer\n        self.pad_token_id = pad_token_id\n        self.phoneme_pad_id = phoneme_pad_id\n\n    def __call__(self, batch: List[Dict]) -> Dict:\n        \"\"\"MODIFIED: Takes a list of dictionaries.\"\"\"\n        \n        # Filter out bad data\n        batch = [b for b in batch if b[\"text_data\"].strip() != \"\" or b[\"is_test\"]]\n        if not batch:\n            return {}\n            \n        # Separate train/val and test samples\n        train_val_batch = [b for b in batch if not b[\"is_test\"]]\n        test_batch = [b for b in batch if b[\"is_test\"]]\n\n        # --- Process Train/Val data (has labels) ---\n        if train_val_batch:\n            neural_data = [b[\"neural_data\"] for b in train_val_batch]\n            phoneme_data = [b[\"phoneme_data\"] for b in train_val_batch]\n            text_data = [b[\"text_data\"] for b in train_val_batch]\n            \n            # 1. Pad Neural Data (B, 512, T_max)\n            neural_data_permuted = [d.permute(1, 0) for d in neural_data]\n            padded_neural = pad_sequence(neural_data_permuted, batch_first=True, padding_value=0.0)\n            padded_neural = padded_neural.permute(0, 2, 1)\n            \n            # 2. Pad Phoneme Data (B, T_phon_max)\n            padded_phonemes = pad_sequence(phoneme_data, batch_first=True, padding_value=self.phoneme_pad_id)\n            \n            # 3. Tokenize and Pad Text Data\n            text_tokens = [self.tokenizer.encode(t) for t in text_data]\n            padded_text = pad_sequence(text_tokens, batch_first=True, padding_value=self.pad_token_id)\n            text_input = padded_text[:, :-1]\n            text_target = padded_text[:, 1:]\n            text_padding_mask = (text_input == self.pad_token_id)\n\n            return {\n                'neural_data': padded_neural,\n                'phoneme_data': padded_phonemes,\n                'text_input': text_input,\n                'text_target': text_target,\n                'text_padding_mask': text_padding_mask,\n                'raw_text': text_data,\n                'is_test_batch': False\n            }\n        \n        # --- Process Test data (no labels) ---\n        elif test_batch:\n            neural_data = [b[\"neural_data\"] for b in test_batch]\n            trial_keys = [b[\"trial_key\"] for b in test_batch] # [cite: 352]\n            \n            # 1. Pad Neural Data (B, 512, T_max)\n            neural_data_permuted = [d.permute(1, 0) for d in neural_data]\n            padded_neural = pad_sequence(neural_data_permuted, batch_first=True, padding_value=0.0)\n            padded_neural = padded_neural.permute(0, 2, 1)\n            \n            return {\n                'neural_data': padded_neural,\n                'trial_keys': trial_keys, # Pass keys for submission [cite: 366]\n                'is_test_batch': True\n            }\n        \n        else:\n            return {}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T14:48:33.319827Z","iopub.execute_input":"2025-12-03T14:48:33.320902Z","iopub.status.idle":"2025-12-03T14:48:33.337948Z","shell.execute_reply.started":"2025-12-03T14:48:33.320879Z","shell.execute_reply":"2025-12-03T14:48:33.337276Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Model Architecture Classes ---\n\nclass TemporalMasking(nn.Module):\n    \"\"\"\n    Applies temporal masking augmentation as described in the proposal.\n    Zeros out a random continuous segment of the time series.\n    \"\"\"\n    def __init__(self, p=0.1, mask_span_length=10):\n        super().__init__()\n        self.p = p\n        self.mask_span_length = mask_span_length\n\n    def forward(self, x):\n        if not self.training or self.p == 0:\n            return x\n        \n        batch_size, num_channels, seq_len = x.shape\n        \n        for i in range(batch_size):\n            if torch.rand(1).item() < self.p:\n                mask_start = torch.randint(0, max(1, seq_len - self.mask_span_length), (1,)).item()\n                x[i, :, mask_start:mask_start + self.mask_span_length] = 0.0\n        return x\n\nclass PositionalEncoding(nn.Module):\n    \"\"\"Standard Positional Encoding for Transformer models.\"\"\"\n    def __init__(self, d_model: int, dropout: float = 0.1, max_len: int = 5000):\n        super().__init__()\n        self.dropout = nn.Dropout(p=dropout)\n        position = torch.arange(max_len).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))\n        pe = torch.zeros(max_len, 1, 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: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Args:\n            x: Tensor, shape [batch_size, seq_len, embedding_dim]\n        \"\"\"\n        # (B, T, E) -> (T, B, E) for PE compatibility\n        x = x.permute(1, 0, 2)\n        x = x + self.pe[:x.size(0)]\n        x = self.dropout(x)\n        # (T, B, E) -> (B, T, E)\n        return x.permute(1, 0, 2)\n\n# --- 1. The Encoders (Neural and Phoneme) ---\n\nclass NeuralEncoder(nn.Module):\n    \"\"\"\n    Implements the 1D CNN + BiLSTM encoder for neural data.\n    \"\"\"\n    def __init__(self, input_channels, cnn_out_channels, lstm_hidden, embedding_dim):\n        super().__init__()\n        \n        # 1D CNN layers to capture local features\n        self.cnn = nn.Sequential(\n            nn.Conv1d(input_channels, cnn_out_channels, kernel_size=3, padding=1),\n            nn.ReLU(),\n            nn.BatchNorm1d(cnn_out_channels),\n            nn.MaxPool1d(kernel_size=2, stride=2) # Downsample time\n        )\n        \n        # BiLSTM to capture long-range dependencies\n        self.bilstm = nn.LSTM(\n            input_size=cnn_out_channels,\n            hidden_size=lstm_hidden,\n            num_layers=2,\n            bidirectional=True,\n            batch_first=True,\n            dropout=0.2\n        )\n        \n        # Projection heads for contrastive learning and decoder\n        self.contrastive_projection = nn.Linear(lstm_hidden * 2, embedding_dim)\n        self.decoder_projection = nn.Linear(lstm_hidden * 2, embedding_dim) \n        self.d_model = embedding_dim # To match decoder\n\n    def forward_contrastive(self, x):\n        \"\"\"For Stage 1: Returns a single vector embedding (B, E)\"\"\"\n        x = self.cnn(x) # (B, C_out, T_down)\n        x = x.permute(0, 2, 1) # (B, T_down, C_out)\n        lstm_out, _ = self.bilstm(x)\n        # Mean pooling for a single sentence-level embedding\n        pooled_out = torch.mean(lstm_out, dim=1) # (B, lstm_hidden * 2)\n        embedding = self.contrastive_projection(pooled_out)\n        return embedding\n\n    def forward_decoder(self, x):\n        \"\"\"For Stage 2: Returns a sequence of embeddings (B, T_down, E)\"\"\"\n        x = self.cnn(x) # (B, C_out, T_down)\n        x = x.permute(0, 2, 1) # (B, T_down, C_out)\n        lstm_out, _ = self.bilstm(x)\n        # Project the entire sequence for the decoder\n        sequence_embeddings = self.decoder_projection(lstm_out)\n        return sequence_embeddings\n\nclass PhonemeEncoder(nn.Module):\n    \"\"\"\n    Implements the Phoneme Encoder  (adapted for phoneme IDs).\n    \"\"\"\n    def __init__(self, vocab_size, embedding_dim, hidden_dim, output_dim, pad_id):\n        super().__init__()\n        self.embedding = nn.Embedding(vocab_size, embedding_dim, padding_idx=pad_id)\n        self.lstm = nn.LSTM(embedding_dim, hidden_dim, batch_first=True, bidirectional=True)\n        self.projection = nn.Linear(hidden_dim * 2, output_dim)\n\n    def forward_contrastive(self, x): # x is (B, T_phoneme)\n        emb = self.embedding(x) # (B, T, emb_dim)\n        lstm_out, _ = self.lstm(emb) # (B, T, hidden*2)\n        pooled = torch.mean(lstm_out, dim=1) # (B, hidden*2)\n        projected = self.projection(pooled) # (B, output_dim)\n        return projected\n\n# --- 2. Stage 1: Contrastive Pre-training Model ---\n\nclass ContrastiveModel(nn.Module):\n    \"\"\"\n    Wraps the two encoders for the InfoNCE loss stage.\n    \"\"\"\n    def __init__(self, neural_encoder, phoneme_encoder, temperature=0.1):\n        super().__init__()\n        self.neural_encoder = neural_encoder\n        self.phoneme_encoder = phoneme_encoder\n        self.temperature = temperature\n        self.augmentation = TemporalMasking(p=0.1, mask_span_length=20) # \n\n    def forward(self, neural_data, phoneme_data):\n        # Apply temporal masking augmentation \n        neural_data = self.augmentation(neural_data)\n        \n        # Get embeddings\n        neural_emb = self.neural_encoder.forward_contrastive(neural_data)\n        phoneme_emb = self.phoneme_encoder.forward_contrastive(phoneme_data)\n        \n        # Normalize embeddings (crucial for contrastive loss)\n        neural_emb = F.normalize(neural_emb, p=2, dim=1)\n        phoneme_emb = F.normalize(phoneme_emb, p=2, dim=1)\n        \n        # Calculate InfoNCE Loss\n        batch_size = neural_emb.shape[0]\n        # Calculate cosine similarity matrix\n        logits = torch.matmul(neural_emb, phoneme_emb.T) / self.temperature\n        \n        # Labels are the diagonal (i,i)\n        labels = torch.arange(batch_size, device=neural_emb.device)\n        \n        # Symmetric loss (Neural->Phoneme and Phoneme->Neural)\n        loss_n_to_p = F.cross_entropy(logits, labels)\n        loss_p_to_n = F.cross_entropy(logits.T, labels)\n        \n        loss = (loss_n_to_p + loss_p_to_n) / 2\n        return loss\n\n# --- 3. Stage 2: Decoder Model ---\n\nclass DecoderModel(nn.Module):\n    \"\"\"\n    Uses the frozen neural encoder and trains a Transformer Decoder.\n    \"\"\"\n    def __init__(self, neural_encoder, decoder_vocab_size, embedding_dim, \n                 nhead, num_decoder_layers, dim_feedforward, pad_token_id,\n                 sos_token_id, eos_token_id):\n        super().__init__()\n        self.neural_encoder = neural_encoder\n        self.d_model = embedding_dim\n        self.pad_token_id = pad_token_id\n        self.sos_token_id = sos_token_id\n        self.eos_token_id = eos_token_id\n        \n        # Freeze the encoder as per the plan \n        self.freeze_encoder()\n            \n        self.decoder_embedding = nn.Embedding(decoder_vocab_size, embedding_dim, padding_idx=pad_token_id)\n        self.pos_encoder = PositionalEncoding(embedding_dim)\n\n        # Standard Transformer Decoder\n        decoder_layer = nn.TransformerDecoderLayer(\n            d_model=embedding_dim, \n            nhead=nhead, \n            dim_feedforward=dim_feedforward,\n            batch_first=True # We use batch_first for easier data handling\n        )\n        self.transformer_decoder = nn.TransformerDecoder(decoder_layer, num_layers=num_decoder_layers)\n        self.final_fc = nn.Linear(embedding_dim, decoder_vocab_size)\n\n    def freeze_encoder(self):\n        print(\"Freezing neural encoder weights.\")\n        for param in self.neural_encoder.parameters():\n            param.requires_grad = False\n        self.neural_encoder.eval()\n\n    def unfreeze_encoder(self):\n        \"\"\"Used for the ablation study.\"\"\"\n        print(\"Unfreezing neural encoder weights for joint training.\")\n        for param in self.neural_encoder.parameters():\n            param.requires_grad = True\n        self.neural_encoder.train() # Set back to train mode\n\n    def generate_square_subsequent_mask(self, sz):\n        \"\"\"Generates a mask to prevent attention to future tokens.\"\"\"\n        mask = (torch.triu(torch.ones(sz, sz)) == 1).transpose(0, 1)\n        mask = mask.float().masked_fill(mask == 0, float('-inf')).masked_fill(mask == 1, float(0.0))\n        return mask\n\n    def forward(self, neural_data, target_text_tokens, target_padding_mask):\n        \"\"\"\n        Forward pass for training (teacher forcing).\n        \"\"\"\n        # 1. Get memory from frozen encoder (B, T_neural, E)\n        with torch.no_grad():\n            memory = self.neural_encoder.forward_decoder(neural_data)\n\n        # 2. Prepare target text (B, T_text, E)\n        tgt_emb = self.decoder_embedding(target_text_tokens) * math.sqrt(self.d_model)\n        tgt_emb = self.pos_encoder(tgt_emb)\n        \n        # 3. Create target mask (T_text, T_text)\n        tgt_seq_len = target_text_tokens.size(1)\n        tgt_mask = self.generate_square_subsequent_mask(tgt_seq_len).to(neural_data.device)\n        \n        # 4. Run through decoder\n        output = self.transformer_decoder(\n            tgt=tgt_emb, \n            memory=memory, \n            tgt_mask=tgt_mask,\n            tgt_key_padding_mask=target_padding_mask\n        )\n        \n        # 5. Final projection\n        logits = self.final_fc(output)\n        return logits\n        \n    def predict(self, neural_data, max_len=50):\n        \"\"\"\n        Autoregressive generation (greedy search) for inference/evaluation.\n        \"\"\"\n        self.eval() # Set model to evaluation mode\n        device = neural_data.device\n        \n        with torch.no_grad():\n            batch_size = neural_data.shape[0]\n            \n            # 1. Get memory from encoder (B, T_neural, E)\n            memory = self.neural_encoder.forward_decoder(neural_data)\n            \n            # 2. Start with <SOS> token (B, 1)\n            tgt_tokens = torch.full((batch_size, 1), self.sos_token_id, dtype=torch.long, device=device)\n            \n            for _ in range(max_len):\n                # 3. Prepare input for decoder\n                tgt_emb = self.decoder_embedding(tgt_tokens) * math.sqrt(self.d_model)\n                tgt_emb = self.pos_encoder(tgt_emb)\n                \n                tgt_seq_len = tgt_tokens.size(1)\n                tgt_mask = self.generate_square_subsequent_mask(tgt_seq_len).to(device)\n\n                # 4. Get decoder output\n                output = self.transformer_decoder(\n                    tgt=tgt_emb,\n                    memory=memory,\n                    tgt_mask=tgt_mask\n                ) # (B, T_current, E)\n                \n                # 5. Get logits for the *last* token only\n                last_token_logits = self.final_fc(output[:, -1, :]) # (B, Vocab)\n                \n                # 6. Get greedy prediction\n                pred_token = torch.argmax(last_token_logits, dim=-1) # (B)\n                \n                # 7. Append predicted token to target sequence\n                tgt_tokens = torch.cat((tgt_tokens, pred_token.unsqueeze(1)), dim=1) # (B, T+1)\n                \n                # 8. Stop if all sequences in batch ended\n                if (pred_token == self.eos_token_id).all():\n                    break\n        \n        self.train() # Set model back to train mode\n        return tgt_tokens","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T14:48:33.338610Z","iopub.execute_input":"2025-12-03T14:48:33.338777Z","iopub.status.idle":"2025-12-03T14:48:33.363300Z","shell.execute_reply.started":"2025-12-03T14:48:33.338763Z","shell.execute_reply":"2025-12-03T14:48:33.362623Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Setup DataLoaders and Tokenizer ---\n\ndef setup_data(cfg):\n    \"\"\"Prepares datasets, tokenizer, and dataloaders.\"\"\"\n    print(\"Finding HDF5 files...\")\n    # This function expects a 'cfg' object\n    train_files = sorted(glob.glob(os.path.join(cfg.DATA_DIR, \"*/*_train.hdf5\")))\n    val_files = sorted(glob.glob(os.path.join(cfg.DATA_DIR, \"*/*_val.hdf5\")))\n    test_files = sorted(glob.glob(os.path.join(cfg.DATA_DIR, \"*/*_test.hdf5\")))\n\n    print(\"Initializing training dataset...\")\n    train_dataset = BrainToTextDataset(train_files)\n    print(\"Initializing validation dataset...\")\n    val_dataset = BrainToTextDataset(val_files)\n    print(\"Initializing test dataset...\")\n    test_dataset = BrainToTextDataset(test_files) \n    \n    print(\"Building tokenizer vocabulary from training sentences...\")\n    tokenizer = SimpleTokenizer()\n    all_sentences = train_dataset.get_all_sentences()\n    tokenizer.build_vocab(all_sentences)\n    vocab_size = tokenizer.get_vocab_size()\n    print(f\"Built vocab with {vocab_size} tokens.\")\n    \n    collator = BrainCollator(tokenizer, cfg.PAD_TOKEN_ID, cfg.PHONEME_PAD_ID)\n    \n    train_loader = DataLoader(train_dataset, batch_size=cfg.BATCH_SIZE, \n                              shuffle=True, collate_fn=collator, num_workers=2)\n    val_loader = DataLoader(val_dataset, batch_size=cfg.BATCH_SIZE, \n                            shuffle=False, collate_fn=collator, num_workers=2)\n    test_loader = DataLoader(test_dataset, batch_size=cfg.BATCH_SIZE, \n                             shuffle=False, collate_fn=collator, num_workers=2)\n    \n    print(f\"\\nTotal Train samples: {len(train_dataset)}\")\n    print(f\"Total Val samples: {len(val_dataset)}\")\n    print(f\"Total Test samples: {len(test_dataset)}\")\n\n    return train_loader, val_loader, test_loader, tokenizer, vocab_size\n\n# --- Global setup (run once) ---\ntry:\n    # --- THIS IS THE FIX ---\n    # Create a simple 'config' object from the global variables defined in Cell 3\n    class ConfigWrapper:\n        pass\n    \n    config = ConfigWrapper()\n    config.DATA_DIR = DATA_DIR\n    config.PAD_TOKEN_ID = PAD_TOKEN_ID\n    config.PHONEME_PAD_ID = PHONEME_PAD_ID\n    config.BATCH_SIZE = BATCH_SIZE\n    # --- END OF FIX ---\n\n    # Now this call will work\n    train_loader, val_loader, test_loader, tokenizer, VOCAB_SIZE = setup_data(config)\n    print(\"\\nData setup complete.\")\n\nexcept Exception as e:\n    print(f\"\\nData setup failed. This is common on Kaggle if the data path is wrong.\")\n    print(f\"Error: {e}\")\n    # Re-raise to show the user the full traceback\n    raise e","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T14:48:33.364244Z","iopub.execute_input":"2025-12-03T14:48:33.364539Z","iopub.status.idle":"2025-12-03T14:52:50.436948Z","shell.execute_reply.started":"2025-12-03T14:48:33.364517Z","shell.execute_reply":"2025-12-03T14:52:50.436258Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- STAGE 1: CONTRASTIVE PRE-TRAINING (RESUMABLE) ---\nimport os\n\nprint(f\"--- Starting Stage 1: Contrastive Pre-training ---\")\n\n# 1. Set up Models (High Capacity)\nneural_enc = NeuralEncoder(\n    input_channels=NEURAL_CHANNELS,\n    cnn_out_channels=NEURAL_CNN_OUT,\n    lstm_hidden=NEURAL_LSTM_HIDDEN,\n    embedding_dim=EMBEDDING_DIM\n).to(DEVICE)\n\nphoneme_enc = PhonemeEncoder(\n    vocab_size=PHONEME_VOCAB_SIZE,\n    embedding_dim=PHONEME_EMBED_DIM,\n    hidden_dim=PHONEME_LSTM_HIDDEN,\n    output_dim=EMBEDDING_DIM,\n    pad_id=PHONEME_PAD_ID\n).to(DEVICE)\n\n# Note: Use temperature=0.07 for sharper alignment!\ncontrastive_model = ContrastiveModel(neural_enc, phoneme_enc, temperature=0.07).to(DEVICE)\noptimizer = optim.Adam(contrastive_model.parameters(), lr=LR_STAGE1)\n\n# 2. Check for existing checkpoint\nstart_epoch = 0\nstage1_history = {'train_loss': [], 'val_loss': []}\n\n# Simple resume logic: look for the latest epoch checkpoint\n# For simplicity in this final run, we overwrite the main checkpoint\nif os.path.exists(ENCODER_CHECKPOINT_PATH):\n    print(\"Found existing Stage 1 weights. Skipping Stage 1 training to save time.\")\n    # If you want to force re-train, comment out the break or set start_epoch manually\n    # But for a 12h run, if it exists, assume it's done.\n    start_epoch = NUM_EPOCHS_STAGE1 \nelse:\n    print(\"No checkpoint found. Starting Stage 1 from scratch.\")\n\n# 3. Training Loop\nif start_epoch < NUM_EPOCHS_STAGE1:\n    best_val_loss = float('inf')\n    \n    for epoch in range(start_epoch, NUM_EPOCHS_STAGE1):\n        contrastive_model.train()\n        train_loss = 0\n        for batch in tqdm(train_loader, desc=f\"Stage 1 Epoch {epoch+1}\", leave=False):\n            if not batch: continue\n            neural_data = batch['neural_data'].to(DEVICE)\n            phoneme_data = batch['phoneme_data'].to(DEVICE)\n            \n            optimizer.zero_grad()\n            loss = contrastive_model(neural_data, phoneme_data)\n            loss.backward()\n            optimizer.step()\n            train_loss += loss.item()\n\n        avg_train = train_loss / len(train_loader)\n        stage1_history['train_loss'].append(avg_train)\n\n        # Validation\n        contrastive_model.eval()\n        val_loss = 0\n        with torch.no_grad():\n            for batch in tqdm(val_loader, desc=\"Validating\", leave=False):\n                if not batch: continue\n                neural_data = batch['neural_data'].to(DEVICE)\n                phoneme_data = batch['phoneme_data'].to(DEVICE)\n                loss = contrastive_model(neural_data, phoneme_data)\n                val_loss += loss.item()\n\n        avg_val = val_loss / len(val_loader)\n        stage1_history['val_loss'].append(avg_val)\n        \n        if (epoch + 1) % 5 == 0:\n            print(f\"Epoch {epoch+1}/{NUM_EPOCHS_STAGE1} | Train: {avg_train:.4f} | Val: {avg_val:.4f}\")\n        \n        # Save \"Best\" Checkpoint\n        if avg_val < best_val_loss:\n            best_val_loss = avg_val\n            torch.save(neural_enc.state_dict(), ENCODER_CHECKPOINT_PATH)\n            torch.save(phoneme_enc.state_dict(), PHONEME_ENCODER_CHECKPOINT_PATH)\n            \n        # Save \"Latest\" Checkpoint (Safety) every 10 epochs\n        if (epoch + 1) % 10 == 0:\n             torch.save(neural_enc.state_dict(), f\"stage1_latest.pth\")\n\nprint(\"--- Stage 1 Training Complete ---\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T14:52:50.437764Z","iopub.execute_input":"2025-12-03T14:52:50.438038Z","iopub.status.idle":"2025-12-03T22:40:51.471358Z","shell.execute_reply.started":"2025-12-03T14:52:50.438017Z","shell.execute_reply":"2025-12-03T22:40:51.470482Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- STAGE 2: DECODER TRAINING (RESUMABLE) ---\n\nRUN_ABLATION_NO_PRETRAIN = False \nWER_CHECK_INTERVAL = 20 # Check every 20 epochs\n\nprint(f\"--- Starting Stage 2: Decoder Training ---\")\n\n# 1. Setup Models\nneural_enc = NeuralEncoder(\n    input_channels=NEURAL_CHANNELS,\n    cnn_out_channels=NEURAL_CNN_OUT,\n    lstm_hidden=NEURAL_LSTM_HIDDEN,\n    embedding_dim=EMBEDDING_DIM\n)\n\n# Always load the BEST Stage 1 weights\nprint(f\"Loading pretrained encoder from {ENCODER_CHECKPOINT_PATH}\")\ntry:\n    neural_enc.load_state_dict(torch.load(ENCODER_CHECKPOINT_PATH, map_location=DEVICE))\nexcept FileNotFoundError:\n    print(f\"!!! ERROR: Pretrained encoder file not found. !!!\")\n    raise\n\ndecoder_model = DecoderModel(\n    neural_encoder=neural_enc,\n    decoder_vocab_size=VOCAB_SIZE,\n    embedding_dim=EMBEDDING_DIM,\n    nhead=DECODER_NHEAD,\n    num_decoder_layers=DECODER_NLAYERS,\n    dim_feedforward=DECODER_DIM_FEEDFORWARD,\n    pad_token_id=PAD_TOKEN_ID,\n    sos_token_id=SOS_TOKEN_ID,\n    eos_token_id=EOS_TOKEN_ID\n).to(DEVICE)\n\noptimizer = optim.Adam(filter(lambda p: p.requires_grad, decoder_model.parameters()), lr=LR_STAGE2)\ncriterion = nn.CrossEntropyLoss(ignore_index=PAD_TOKEN_ID)\n\n# 2. Check for Resume\nstart_epoch = 0\nstage2_history = {'train_loss': [], 'val_loss': [], 'val_wer': [], 'wer_epochs': []}\n\nif os.path.exists(DECODER_CHECKPOINT_PATH):\n    print(\"Found existing Stage 2 checkpoint. Loading...\")\n    # NOTE: Ideally you load full state dict. For this run, we just overwrite.\n    # If the file exists, it means a previous run finished OR crashed after saving.\n    # To be safe, we will just start from 0 unless you manually implement complex state loading.\n    # But we WILL save frequently so you can manually load if needed.\n    pass \n\n# 3. Training Loop\nbest_val_loss = float('inf')\n\nfor epoch in range(start_epoch, NUM_EPOCHS_STAGE2):\n    decoder_model.train()\n    train_loss = 0\n    for batch in tqdm(train_loader, desc=f\"Stage 2 Epoch {epoch+1}\", leave=False):\n        if not batch: continue\n        neural_data = batch['neural_data'].to(DEVICE)\n        text_input = batch['text_input'].to(DEVICE)\n        text_target = batch['text_target'].to(DEVICE)\n        text_padding_mask = batch['text_padding_mask'].to(DEVICE)\n        \n        optimizer.zero_grad()\n        logits = decoder_model(neural_data, text_input, text_padding_mask)\n        loss = criterion(logits.reshape(-1, logits.shape[-1]), text_target.reshape(-1))\n        loss.backward()\n        optimizer.step()\n        train_loss += loss.item()\n        \n    avg_train = train_loss / len(train_loader)\n    stage2_history['train_loss'].append(avg_train)\n\n    # Validation Loss\n    decoder_model.eval()\n    val_loss = 0\n    with torch.no_grad():\n        for batch in tqdm(val_loader, desc=\"Validating\", leave=False):\n            if not batch: continue\n            neural_data = batch['neural_data'].to(DEVICE)\n            text_input = batch['text_input'].to(DEVICE)\n            text_target = batch['text_target'].to(DEVICE)\n            text_padding_mask = batch['text_padding_mask'].to(DEVICE)\n\n            logits = decoder_model(neural_data, text_input, text_padding_mask)\n            loss = criterion(logits.reshape(-1, logits.shape[-1]), text_target.reshape(-1))\n            val_loss += loss.item()\n\n    avg_val = val_loss / len(val_loader)\n    stage2_history['val_loss'].append(avg_val)\n    \n    if (epoch+1) % 10 == 0:\n        print(f\"Epoch {epoch+1}/{NUM_EPOCHS_STAGE2} | Train: {avg_train:.4f} | Val: {avg_val:.4f}\")\n\n    # Validation WER (Every 20 epochs)\n    if (epoch + 1) % WER_CHECK_INTERVAL == 0:\n        print(\"  Calculating WER...\")\n        ground_truths = []\n        predictions = []\n        with torch.no_grad():\n            for batch in tqdm(val_loader, desc=\"Calculating WER\", leave=False):\n                if not batch: continue\n                neural_data = batch['neural_data'].to(DEVICE)\n                raw_text_gt = batch['raw_text']\n                predicted_tokens = decoder_model.predict(neural_data)\n                for i in range(len(raw_text_gt)):\n                    pred_text = tokenizer.decode(predicted_tokens[i])\n                    ground_truths.append(raw_text_gt[i].lower())\n                    predictions.append(pred_text.lower())\n        \n        current_wer = jiwer.wer(ground_truths, predictions)\n        stage2_history['val_wer'].append(current_wer)\n        stage2_history['wer_epochs'].append(epoch + 1)\n        print(f\"  >> Validation WER: {current_wer*100:.2f}%\")\n\n    # Save Checkpoints\n    if avg_val < best_val_loss:\n        best_val_loss = avg_val\n        torch.save(decoder_model.state_dict(), DECODER_CHECKPOINT_PATH)\n        \n    # Safety Save every 10 epochs\n    if (epoch + 1) % 10 == 0:\n        torch.save(decoder_model.state_dict(), f\"decoder_latest.pth\")\n            \nprint(\"--- Stage 2 Training Complete ---\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-04T03:49:05.970472Z","iopub.execute_input":"2025-12-04T03:49:05.970990Z","iopub.status.idle":"2025-12-04T03:49:06.044909Z","shell.execute_reply.started":"2025-12-04T03:49:05.970966Z","shell.execute_reply":"2025-12-04T03:49:06.044010Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- EVALUATION (WER) & FINAL SUBMISSION ---\n\nprint(f\"--- Starting Evaluation & Submission ---\")\n\n# 1. Load Model\n# We must re-define the structure first\nneural_enc_eval = NeuralEncoder(\n    input_channels=NEURAL_CHANNELS,\n    cnn_out_channels=NEURAL_CNN_OUT,\n    lstm_hidden=NEURAL_LSTM_HIDDEN,\n    embedding_dim=EMBEDDING_DIM\n)\neval_model = DecoderModel(\n    neural_encoder=neural_enc_eval,\n    decoder_vocab_size=VOCAB_SIZE,\n    embedding_dim=EMBEDDING_DIM,\n    nhead=DECODER_NHEAD,\n    num_decoder_layers=DECODER_NLAYERS,\n    dim_feedforward=DECODER_DIM_FEEDFORWARD,\n    pad_token_id=PAD_TOKEN_ID,\n    sos_token_id=SOS_TOKEN_ID,\n    eos_token_id=EOS_TOKEN_ID\n).to(DEVICE)\n\n# --- CHOOSE WHICH MODEL TO EVALUATE ---\n# Set to True if you are evaluating the ablation model\nEVAL_ABLATION_MODEL = False\n\nmodel_load_path = f\"/kaggle/working/ablation_model_stage2.pth\" if EVAL_ABLATION_MODEL else DECODER_CHECKPOINT_PATH\nprint(f\"Loading full model weights from {model_load_path}\")\ntry:\n    eval_model.load_state_dict(torch.load(model_load_path, map_location=DEVICE))\nexcept FileNotFoundError:\n    print(\"!!! ERROR: Trained decoder file not found. Run the Stage 2 cell. !!!\")\n    raise\n\neval_model.eval()\n\n# --- PART 1: VALIDATION (Calculate WER) ---\nprint(\"\\n--- Running Validation (to calculate WER) ---\")\nground_truths = []\npredictions = []\n\nwith torch.no_grad():\n    # We use the val_loader created in Cell 7\n    for batch in tqdm(val_loader, desc=\"Validating\"):\n        if not batch: continue\n        neural_data = batch['neural_data'].to(DEVICE)\n        raw_text_gt = batch['raw_text'] # List of ground truth strings\n        \n        predicted_tokens = eval_model.predict(neural_data)\n        \n        for i in range(len(raw_text_gt)):\n            pred_text = tokenizer.decode(predicted_tokens[i])\n            ground_truths.append(raw_text_gt[i].lower())\n            predictions.append(pred_text.lower())\n\n# Calculate WER [cite: 500]\nwer = jiwer.wer(ground_truths, predictions)\nprint(\"\\n\" + \"=\"*30)\nif EVAL_ABLATION_MODEL:\n    print(\"--- Ablation Model Validation Results ---\")\nelse:\n    print(\"--- Final Model Validation Results ---\")\nprint(f\"Overall Word Error Rate (WER): {wer * 100:.2f}%\")\nprint(\"=\"*30 + \"\\n\")\n\n# Show examples\nprint(\"--- Example Validation Predictions ---\")\ndf = pd.DataFrame({'Ground Truth': ground_truths, 'Prediction': predictions})\nprint(df.head(10).to_string())\n\n\n# --- PART 2: TEST (Generate Submission.csv) ---\nprint(\"\\n--- Running Test (to generate submission.csv) ---\")\n\nall_trial_keys = []\nall_predictions = []\n\nwith torch.no_grad():\n    # We use the test_loader created in Cell 7\n    for batch in tqdm(test_loader, desc=\"Generating Submission\"):\n        if not batch: continue\n        \n        # The baseline collator [cite: 348-352] returns 'keys'\n        # Our collator (Cell 5) does not, let's fix that.\n        # FOR NOW, we assume the test_loader is just data.\n        # ---\n        # ** Correction: My `BrainToTextDataset` (Cell 5) and `BrainCollator` (Cell 5)\n        # ** need to be updated to handle the `trial_key`.\n        # ** Let's assume for now the loader is correct.\n        # **\n        # ** I will go back and modify Cell 5 to add the trial_key,\n        # ** then this code will work.\n        \n        # ---\n        # The following code assumes Cell 5 & 7 are updated to pass `trial_key`.\n        # I will provide the fixes for Cell 5 & 7 right after this.\n        # ---\n        \n        neural_data = batch['neural_data'].to(DEVICE)\n        \n        # The baseline [cite: 585] stores trial keys\n        # My collator (Cell 5) needs to be updated to pass this.\n        # Assuming it's in batch['trial_keys'] for now.\n        trial_keys = batch['trial_keys'] \n        \n        predicted_tokens = eval_model.predict(neural_data)\n        \n        for i in range(len(trial_keys)):\n            pred_text = tokenizer.decode(predicted_tokens[i])\n            all_trial_keys.append(trial_keys[i])\n            all_predictions.append(pred_text)\n\n# Create submission DataFrame, matching the baseline [cite: 597]\nsubmission_df = pd.DataFrame({\n    'id': all_trial_keys,\n    'text': all_predictions\n})\n\n# Format text for submission [cite: 599]\nsubmission_df['text'] = submission_df['text'].str.strip()\n\n# Save to submission.csv [cite: 601, 741]\nsubmission_path = \"/kaggle/working/submission.csv\"\nsubmission_df.to_csv(submission_path, index=False)\n\nprint(f\"\\nSuccessfully generated submission file!\")\nprint(f\"Saved to: {submission_path}\")\nprint(\"--- Final Submission File Head ---\")\nprint(submission_df.head())","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-04T02:47:41.639Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- VISUALIZATION: LOSS CURVES & WER ---\nimport matplotlib.pyplot as plt\n\nplt.figure(figsize=(18, 6))\n\n# Plot 1: Stage 1 Loss (Contrastive)\nplt.subplot(1, 3, 1)\nif 'stage1_history' in locals() and len(stage1_history['train_loss']) > 0:\n    plt.plot(stage1_history['train_loss'], label='Training Loss', color='blue')\n    plt.plot(stage1_history['val_loss'], label='Validation Loss', color='orange')\n    plt.title('Stage 1: Contrastive Loss')\n    plt.xlabel('Epoch')\n    plt.ylabel('InfoNCE Loss')\n    plt.legend()\n    plt.grid(True, alpha=0.3)\nelse:\n    plt.text(0.5, 0.5, 'Stage 1 data not found', ha='center')\n\n# Plot 2: Stage 2 Loss (Decoder)\nplt.subplot(1, 3, 2)\nif 'stage2_history' in locals() and len(stage2_history['train_loss']) > 0:\n    plt.plot(stage2_history['train_loss'], label='Training Loss', color='blue')\n    plt.plot(stage2_history['val_loss'], label='Validation Loss', color='orange')\n    plt.title('Stage 2: Decoder Loss')\n    plt.xlabel('Epoch')\n    plt.ylabel('Cross Entropy Loss')\n    plt.legend()\n    plt.grid(True, alpha=0.3)\nelse:\n    plt.text(0.5, 0.5, 'Stage 2 data not found', ha='center')\n\n# Plot 3: Stage 2 WER Progress\nplt.subplot(1, 3, 3)\nif 'stage2_history' in locals() and len(stage2_history['val_wer']) > 0:\n    plt.plot(stage2_history['wer_epochs'], [w*100 for w in stage2_history['val_wer']], \n             marker='o', color='green', linestyle='-', linewidth=2)\n    plt.title('Stage 2: Word Error Rate (WER)')\n    plt.xlabel('Epoch')\n    plt.ylabel('WER (%)')\n    plt.grid(True, alpha=0.3)\n    # Add labels to points\n    for x, y in zip(stage2_history['wer_epochs'], stage2_history['val_wer']):\n        plt.annotate(f\"{y*100:.1f}%\", (x, y*100), xytext=(0, 5), textcoords='offset points', ha='center')\nelse:\n    plt.text(0.5, 0.5, 'WER data not found', ha='center')\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-04T03:48:26.614059Z","iopub.execute_input":"2025-12-04T03:48:26.614945Z","iopub.status.idle":"2025-12-04T03:48:27.020734Z","shell.execute_reply.started":"2025-12-04T03:48:26.614916Z","shell.execute_reply":"2025-12-04T03:48:27.020049Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- t-SNE VISUALIZATION ---\nprint(\"--- Starting t-SNE Visualization ---\")\nprint(\"This will assess the alignment of the Stage 1 embedding space.\")\n\n# 1. Load Stage 1 Encoders\ntry:\n    neural_enc = NeuralEncoder(\n        input_channels=NEURAL_CHANNELS,\n        cnn_out_channels=NEURAL_CNN_OUT,\n        lstm_hidden=NEURAL_LSTM_HIDDEN,\n        embedding_dim=EMBEDDING_DIM\n    ).to(DEVICE)\n    neural_enc.load_state_dict(torch.load(ENCODER_CHECKPOINT_PATH, map_location=DEVICE))\n\n    phoneme_enc = PhonemeEncoder(\n        vocab_size=PHONEME_VOCAB_SIZE,\n        embedding_dim=PHONEME_EMBED_DIM,\n        hidden_dim=PHONEME_LSTM_HIDDEN,\n        output_dim=EMBEDDING_DIM,\n        pad_id=PHONEME_PAD_ID\n    ).to(DEVICE)\n    phoneme_enc.load_state_dict(torch.load(PHONEME_ENCODER_CHECKPOINT_PATH, map_location=DEVICE))\n    \n    neural_enc.eval()\n    phoneme_enc.eval()\n    print(\"Loaded Stage 1 encoders successfully.\")\n\nexcept FileNotFoundError:\n    print(\"!!! ERROR: Could not load Stage 1 encoders. Run Cell 8 first. !!!\")\n    # Don't proceed if files are missing\n    raise\n\n# 2. Get Embeddings\nneural_embeddings = []\nphoneme_embeddings = []\nlabels = [] # We can use the raw text as labels\n\n# We'll just use a few batches to make plotting faster\nnum_batches_to_plot = 10\nprint(f\"Generating embeddings from {num_batches_to_plot} validation batches...\")\n\nwith torch.no_grad():\n    for i, batch in enumerate(val_loader):\n        if not batch: continue\n        if i >= num_batches_to_plot:\n            break\n            \n        neural_data = batch['neural_data'].to(DEVICE)\n        phoneme_data = batch['phoneme_data'].to(DEVICE)\n        \n        # Get embeddings using the *contrastive* forward pass\n        neu_emb = neural_enc.forward_contrastive(neural_data)\n        pho_emb = phoneme_enc.forward_contrastive(phoneme_data)\n        \n        # Normalize them just like in training\n        neu_emb = F.normalize(neu_emb, p=2, dim=1)\n        pho_emb = F.normalize(pho_emb, p=2, dim=1)\n        \n        neural_embeddings.append(neu_emb.cpu())\n        phoneme_embeddings.append(pho_emb.cpu())\n        # Use a unique ID for each sample in the batch\n        labels.extend([f\"sent_{i*BATCH_SIZE + j}\" for j in range(len(batch['raw_text']))])\n\n# Concatenate all batches\nneural_embeddings = torch.cat(neural_embeddings, dim=0).numpy()\nphoneme_embeddings = torch.cat(phoneme_embeddings, dim=0).numpy()\n\n# Combine for t-SNE\nall_embeddings = np.concatenate([neural_embeddings, phoneme_embeddings], axis=0)\nprint(f\"Total embeddings to plot: {all_embeddings.shape[0]}\")\n\n# 3. Run t-SNE\nprint(\"Running t-SNE... (this may take a minute)\")\ntsne = TSNE(n_components=2, perplexity=30, n_iter=1000, random_state=42)\ntsne_results = tsne.fit_transform(all_embeddings)\n\n# 4. Plot Results\nprint(\"Plotting results...\")\nnum_points = len(labels)\ntsne_neural = tsne_results[:num_points]\ntsne_phoneme = tsne_results[num_points:]\n\nplt.figure(figsize=(14, 10))\n\n# Plot neural embeddings as 'x'\nplt.scatter(tsne_neural[:, 0], tsne_neural[:, 1], marker='x', c='blue', label='Neural Embeddings')\n# Plot phoneme embeddings as 'o'\nplt.scatter(tsne_phoneme[:, 0], tsne_phoneme[:, 1], marker='o', c='red', s=50, alpha=0.6, label='Phoneme Embeddings')\n\n# Draw lines connecting aligned pairs\nfor i in range(num_points):\n    plt.plot(\n        [tsne_neural[i, 0], tsne_phoneme[i, 0]],\n        [tsne_neural[i, 1], tsne_phoneme[i, 1]],\n        c='gray', linestyle='--', linewidth=0.5, alpha=0.5\n    )\n\nplt.title('t-SNE Visualization of Aligned Neural and Phoneme Embedding Space')\nplt.xlabel('t-SNE Component 1')\nplt.ylabel('t-SNE Component 2')\nplt.legend()\nplt.show()\n\nprint(\"\\n--- Visualization Complete ---\")\nprint(\"If pre-training was successful, there will be blue 'x's (Neural) \"\n      \"and red 'o's (Phoneme) forming small, tight clusters, \"\n      \"connected by gray lines.\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-12-04T02:47:41.639Z"}},"outputs":[],"execution_count":null}]}