{"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"},{"sourceId":14141820,"sourceType":"datasetVersion","datasetId":9008191}],"dockerImageVersionId":31193,"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":"markdown","source":"#### How to Run the Notebook Using the “Relay” Training Strategy\n\nKaggle notebooks have a 12-hour runtime limit, so this 40+ hours of training is split across three separate sessions using the same notebook. Each session loads the checkpoint produced by the previous one.","metadata":{}},{"cell_type":"code","source":"# --- imports and libraries ---\n\n!pip install jiwer\n\n# --- 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\nimport time\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},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Configuration & Hyperparameters\n\nThis block defines all major settings for the model, including optimized hyperparameters, training schedules, and file paths used throughout the notebook.\n\nThis is the section where you control the behavior of the entire “Relay” training strategy.","metadata":{}},{"cell_type":"code","source":"# -- Data Config --\nDATA_DIR = \"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final/\"\n\n# --- RELAY TRAINING CONFIG (CRITICAL) ---\n# 1. Set this to None for the very first run.\n\n# Phase 2: We resume from the Part 1 dataset\nRESUME_PATH = \"/kaggle/input/brain-training-part1/relay_checkpoint.pth\"\n\n# 3. How long (in hours) to run before safely saving and stopping?\n#    Kaggle limit is 12 hours. We set safety buffer to 11.5 hours.\nMAX_RUNTIME_HOURS = 11.5 \nSTART_TIME = time.time()\n\n# -- Model Config (Optimized) --\nEMBEDDING_DIM = 512           \nNEURAL_CHANNELS = 512         \nNEURAL_CNN_OUT = 128          \nNEURAL_LSTM_HIDDEN = 256      \nPHONEME_VOCAB_SIZE = 42       \nPHONEME_EMBED_DIM = 128       \nPHONEME_LSTM_HIDDEN = 256     \nDECODER_NHEAD = 8             \nDECODER_NLAYERS = 6           \nDECODER_DIM_FEEDFORWARD = 2048\n\n# -- Training Schedule --\nBATCH_SIZE = 16\nNUM_EPOCHS_STAGE1 = 250       \nNUM_EPOCHS_STAGE2 = 15       \nLR_STAGE1 = 1e-4\nLR_STAGE2 = 1e-4\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nWER_CHECK_INTERVAL = 25       \n\n# -- Checkpoint Paths --\n# This is the single \"Relay baton\" file we will save at the end of every run\nRELAY_CHECKPOINT_PATH = \"/kaggle/working/relay_checkpoint.pth\"\n# We also save standalone files for t-SNE compatibility\nENCODER_STANDALONE_PATH = \"/kaggle/working/neural_encoder_stage1.pth\"\nPHONEME_STANDALONE_PATH = \"/kaggle/working/phoneme_encoder_stage1.pth\"\nSUBMISSION_PATH = \"/kaggle/working/submission.csv\"\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}\")\nprint(f\"Resume Path: {RESUME_PATH}\")\nprint(f\"Max Runtime: {MAX_RUNTIME_HOURS} hours\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Tokenizer Class\n\nCreates a simple word-level tokenizer that converts text into integer token IDs (and back). It also manages special tokens like <SOS> and <EOS>.","metadata":{}},{"cell_type":"code","source":"class 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},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Data Loading\n\nDefines the BrainToTextDataset for reading HDF5 files and the BrainCollator for padding batches.\n\nIt also handles the difference between training data (with labels) and test data (which requires trial_keys for submission).","metadata":{}},{"cell_type":"code","source":"class BrainToTextDataset(Dataset):\n    def __init__(self, hdf5_files: List[str]):\n        super().__init__()\n        self.trial_index = [] \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    def __len__(self) -> int:\n        return len(self.trial_index)\n\n    def __getitem__(self, idx: int) -> Dict:\n        file_path, trial_key = self.trial_index[idx]\n        with h5py.File(file_path, 'r') as f:\n            trial_group = f[trial_key]\n            neural_data = torch.from_numpy(trial_group['input_features'][:]).float()\n            is_test = 'seq_class_ids' not in trial_group\n\n            if not is_test:\n                phoneme_data = torch.from_numpy(trial_group['seq_class_ids'][:]).long()\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 = \"\"\n            else:\n                phoneme_data = torch.tensor([], dtype=torch.long)\n                text_data = \"\"\n\n        neural_data = neural_data.permute(1, 0)\n        return {\n            \"neural_data\": neural_data,\n            \"phoneme_data\": phoneme_data,\n            \"text_data\": text_data,\n            \"trial_key\": trial_key, \n            \"is_test\": is_test \n        }\n\n    def get_all_sentences(self) -> List[str]:\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                pass\n        return all_sentences\n\n\nclass BrainCollator:\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        batch = [b for b in batch if b[\"text_data\"].strip() != \"\" or b[\"is_test\"]]\n        if not batch: return {}\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        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            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            padded_phonemes = pad_sequence(phoneme_data, batch_first=True, padding_value=self.phoneme_pad_id)\n\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        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] \n\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, \n                'is_test_batch': True\n            }\n        else:\n            return {}","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Model Architectures\nThis block defines the neural networks:\n\n* NeuralEncoder: CNN + LSTM to process brain signals.\n\n* PhonemeEncoder: Processes phonemes for Stage 1.\n\n* ContrastiveModel: Ties them together for Stage 1 training.\n\n* DecoderModel: The Transformer that generates text in Stage 2.","metadata":{}},{"cell_type":"code","source":"class TemporalMasking(nn.Module):\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        batch_size, num_channels, seq_len = x.shape\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    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        x = x.permute(1, 0, 2)\n        x = x + self.pe[:x.size(0)]\n        x = self.dropout(x)\n        return x.permute(1, 0, 2)\n\nclass NeuralEncoder(nn.Module):\n    def __init__(self, input_channels, cnn_out_channels, lstm_hidden, embedding_dim):\n        super().__init__()\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) \n        )\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        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 \n\n    def forward_contrastive(self, x):\n        x = self.cnn(x) \n        x = x.permute(0, 2, 1) \n        lstm_out, _ = self.bilstm(x)\n        pooled_out = torch.mean(lstm_out, dim=1) \n        embedding = self.contrastive_projection(pooled_out)\n        return embedding\n\n    def forward_decoder(self, x):\n        x = self.cnn(x) \n        x = x.permute(0, 2, 1) \n        lstm_out, _ = self.bilstm(x)\n        sequence_embeddings = self.decoder_projection(lstm_out)\n        return sequence_embeddings\n\nclass PhonemeEncoder(nn.Module):\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): \n        emb = self.embedding(x) \n        lstm_out, _ = self.lstm(emb) \n        pooled = torch.mean(lstm_out, dim=1) \n        projected = self.projection(pooled) \n        return projected\n\nclass ContrastiveModel(nn.Module):\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        neural_data = self.augmentation(neural_data)\n        neural_emb = self.neural_encoder.forward_contrastive(neural_data)\n        phoneme_emb = self.phoneme_encoder.forward_contrastive(phoneme_data)\n        neural_emb = F.normalize(neural_emb, p=2, dim=1)\n        phoneme_emb = F.normalize(phoneme_emb, p=2, dim=1)\n        batch_size = neural_emb.shape[0]\n        logits = torch.matmul(neural_emb, phoneme_emb.T) / self.temperature\n        labels = torch.arange(batch_size, device=neural_emb.device)\n        loss_n_to_p = F.cross_entropy(logits, labels)\n        loss_p_to_n = F.cross_entropy(logits.T, labels)\n        loss = (loss_n_to_p + loss_p_to_n) / 2\n        return loss\n\nclass DecoderModel(nn.Module):\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        self.freeze_encoder()\n        self.decoder_embedding = nn.Embedding(decoder_vocab_size, embedding_dim, padding_idx=pad_token_id)\n        self.pos_encoder = PositionalEncoding(embedding_dim)\n        decoder_layer = nn.TransformerDecoderLayer(\n            d_model=embedding_dim,\n            nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            batch_first=True \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        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() \n\n    def generate_square_subsequent_mask(self, sz):\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        with torch.no_grad():\n            memory = self.neural_encoder.forward_decoder(neural_data)\n        tgt_emb = self.decoder_embedding(target_text_tokens) * math.sqrt(self.d_model)\n        tgt_emb = self.pos_encoder(tgt_emb)\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        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        logits = self.final_fc(output)\n        return logits\n\n    def predict(self, neural_data, max_len=50):\n        self.eval() \n        device = neural_data.device\n        with torch.no_grad():\n            batch_size = neural_data.shape[0]\n            memory = self.neural_encoder.forward_decoder(neural_data)\n            tgt_tokens = torch.full((batch_size, 1), self.sos_token_id, dtype=torch.long, device=device)\n            for _ in range(max_len):\n                tgt_emb = self.decoder_embedding(tgt_tokens) * math.sqrt(self.d_model)\n                tgt_emb = self.pos_encoder(tgt_emb)\n                tgt_seq_len = tgt_tokens.size(1)\n                tgt_mask = self.generate_square_subsequent_mask(tgt_seq_len).to(device)\n                output = self.transformer_decoder(\n                    tgt=tgt_emb,\n                    memory=memory,\n                    tgt_mask=tgt_mask\n                ) \n                last_token_logits = self.final_fc(output[:, -1, :]) \n                pred_token = torch.argmax(last_token_logits, dim=-1) \n                tgt_tokens = torch.cat((tgt_tokens, pred_token.unsqueeze(1)), dim=1) \n                if (pred_token == self.eos_token_id).all():\n                    break\n        self.train() \n        return tgt_tokens","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Environment Setup\n\nScans the directories, initializes the Datasets and DataLoaders, and prepares the Tokenizer.","metadata":{}},{"cell_type":"code","source":"def setup_data(cfg):\n    print(\"Finding HDF5 files...\")\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    train_dataset = BrainToTextDataset(train_files)\n    val_dataset = BrainToTextDataset(val_files)\n    test_dataset = BrainToTextDataset(test_files)\n\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\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    return train_loader, val_loader, test_loader, tokenizer, vocab_size\n\ntry:\n    class ConfigWrapper: pass\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\n    train_loader, val_loader, test_loader, tokenizer, VOCAB_SIZE = setup_data(config)\n    print(\"\\nData setup complete.\")\nexcept Exception as e:\n    print(f\"Error: {e}\")\n    raise e","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Resume Logic & State Recovery\n\nThis is the Brain of the relay strategy. It checks RESUME_PATH.\n\nIf RESUME_PATH exists, it loads the file and figures out: \"Did we finish Stage 1? If yes, skip to Stage 2. If no, what epoch of Stage 1 were we at?\"","metadata":{}},{"cell_type":"code","source":"def load_checkpoint_if_exists(path, device):\n    if path and os.path.exists(path):\n        print(f\"--- Loading Checkpoint from {path} ---\")\n        checkpoint = torch.load(path, map_location=device)\n        return checkpoint\n    return None\n\ndef save_relay_checkpoint(state_dict, path):\n    print(f\"--- Saving Relay Checkpoint to {path} ---\")\n    torch.save(state_dict, path)\n\n# Global State Variables for Resume Logic\nstart_stage1 = True\nstart_stage2 = True\nstart_epoch_s1 = 0\nstart_epoch_s2 = 0\ns1_train_losses = []\ns1_val_losses = []\ns2_train_losses = []\ns2_val_losses = []\n\n# --- Load Previous State ---\ncheckpoint = load_checkpoint_if_exists(RESUME_PATH, DEVICE)\n\nif checkpoint:\n    print(\"Checkpoint found. Checking progress...\")\n    \n    # Recover History\n    s1_train_losses = checkpoint.get('s1_train_losses', [])\n    s1_val_losses = checkpoint.get('s1_val_losses', [])\n    s2_train_losses = checkpoint.get('s2_train_losses', [])\n    s2_val_losses = checkpoint.get('s2_val_losses', [])\n    \n    # Determine Stage\n    if checkpoint.get('stage1_complete', False):\n        print(\"Stage 1 was marked complete in checkpoint. Skipping Stage 1.\")\n        start_stage1 = False\n    else:\n        print(f\"Resuming Stage 1 from epoch {checkpoint.get('epoch', 0)}\")\n        start_epoch_s1 = checkpoint.get('epoch', 0)\n        \n    if not start_stage1:\n        if checkpoint.get('stage2_complete', False):\n            print(\"Stage 2 also complete. Going straight to evaluation.\")\n            start_stage2 = False\n        else:\n            print(f\"Resuming Stage 2 from epoch {checkpoint.get('epoch', 0)}\")\n            start_epoch_s2 = checkpoint.get('epoch', 0)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Stage 1 Training (Contrastive)\nThis block runs the InfoNCE contrastive training.\n\n#### Relay Logic:\n\nChecks the elapsed_hours. If > 11.5, it saves and stops execution immediately.\n\nIf it finishes all epochs, it saves the model and also creates standalone copies for t-SNE.","metadata":{}},{"cell_type":"code","source":"# --- STAGE 1: CONTRASTIVE PRE-TRAINING ---\n\nif start_stage1:\n    print(f\"\\n--- Starting/Resuming Stage 1: Contrastive Pre-training ---\")\n    \n    neural_enc = NeuralEncoder(NEURAL_CHANNELS, NEURAL_CNN_OUT, NEURAL_LSTM_HIDDEN, EMBEDDING_DIM).to(DEVICE)\n    phoneme_enc = PhonemeEncoder(PHONEME_VOCAB_SIZE, PHONEME_EMBED_DIM, PHONEME_LSTM_HIDDEN, EMBEDDING_DIM, PHONEME_PAD_ID).to(DEVICE)\n    contrastive_model = ContrastiveModel(neural_enc, phoneme_enc).to(DEVICE)\n    optimizer = optim.Adam(contrastive_model.parameters(), lr=LR_STAGE1)\n\n    # Load weights if resuming mid-stage\n    if checkpoint and not checkpoint.get('stage1_complete', False):\n        neural_enc.load_state_dict(checkpoint['neural_encoder_state'])\n        phoneme_enc.load_state_dict(checkpoint['phoneme_encoder_state'])\n        optimizer.load_state_dict(checkpoint['optimizer_state'])\n        print(\"Loaded Stage 1 weights and optimizer state.\")\n\n    stop_early = False\n    \n    for epoch in range(start_epoch_s1, NUM_EPOCHS_STAGE1):\n        # Time Check\n        elapsed_hours = (time.time() - START_TIME) / 3600\n        if elapsed_hours > MAX_RUNTIME_HOURS:\n            print(f\"\\n!!! TIME LIMIT REACHED ({elapsed_hours:.2f} hrs). Saving and Stopping. !!!\")\n            stop_early = True\n            save_relay_checkpoint({\n                'stage1_complete': False,\n                'epoch': epoch, # Resume from this epoch next time\n                'neural_encoder_state': neural_enc.state_dict(),\n                'phoneme_encoder_state': phoneme_enc.state_dict(),\n                'optimizer_state': optimizer.state_dict(),\n                's1_train_losses': s1_train_losses,\n                's1_val_losses': s1_val_losses\n            }, RELAY_CHECKPOINT_PATH)\n            break\n\n        print(f\"\\nEpoch {epoch+1}/{NUM_EPOCHS_STAGE1}\")\n        \n        # Train\n        contrastive_model.train()\n        train_loss = 0\n        for batch in tqdm(train_loader, desc=\"Training S1\", leave=False):\n            if not batch: continue\n            neural_data = batch['neural_data'].to(DEVICE)\n            phoneme_data = batch['phoneme_data'].to(DEVICE)\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        s1_train_losses.append(avg_train)\n        \n        # Val\n        contrastive_model.eval()\n        val_loss = 0\n        with torch.no_grad():\n            for batch in tqdm(val_loader, desc=\"Validating S1\", 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        s1_val_losses.append(avg_val)\n        print(f\"Train Loss: {avg_train:.4f} | Val Loss: {avg_val:.4f}\")\n\n    if not stop_early:\n        print(\"Stage 1 Finished. Saving complete state.\")\n        # Save state marking Stage 1 as complete\n        save_relay_checkpoint({\n            'stage1_complete': True,\n            'epoch': 0, # Reset epoch for Stage 2\n            'neural_encoder_state': neural_enc.state_dict(),\n            'phoneme_encoder_state': phoneme_enc.state_dict(),\n            # We don't need optimizer for stage 2 transition\n            's1_train_losses': s1_train_losses,\n            's1_val_losses': s1_val_losses\n        }, RELAY_CHECKPOINT_PATH)\n        \n        # ALSO save standalone weights so t-SNE can pick them up later easily\n        torch.save(neural_enc.state_dict(), ENCODER_STANDALONE_PATH)\n        torch.save(phoneme_enc.state_dict(), PHONEME_STANDALONE_PATH)\n        print(\"Saved standalone encoder weights for t-SNE.\")\n        \n    else:\n        print(\"Stopping execution due to time limit.\")\n        # Verify the file exists\n        if os.path.exists(RELAY_CHECKPOINT_PATH):\n            print(\"Checkpoint verified.\")\n        exit() # Stop the notebook here","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Stage 2 Training (Decoder)\nThis block runs the Transformer Decoder training.\n\n#### Relay Logic:\n\nLoads the frozen encoder.\n\nResumes the optimizer state (very important for stability) if continuing from a previous session.\n\nAlso has the 11.5-hour safety stop mechanism.","metadata":{}},{"cell_type":"code","source":"# --- STAGE 2: DECODER TRAINING ---\n\nif not start_stage1 and start_stage2:\n    print(f\"\\n--- Starting/Resuming Stage 2: Decoder Training ---\")\n\n    # 1. Initialize Neural Encoder\n    neural_enc = NeuralEncoder(NEURAL_CHANNELS, NEURAL_CNN_OUT, NEURAL_LSTM_HIDDEN, EMBEDDING_DIM).to(DEVICE)\n    \n    # 2. Load Encoder Weights\n    if checkpoint and 'neural_encoder_state' in checkpoint:\n        print(\"Loading frozen neural encoder from checkpoint.\")\n        neural_enc.load_state_dict(checkpoint['neural_encoder_state'])\n    elif os.path.exists(ENCODER_STANDALONE_PATH):\n        # Fallback if we are running this in a new session where relay dict isn't loaded but files are\n        print(\"Loading encoder from standalone file.\")\n        neural_enc.load_state_dict(torch.load(ENCODER_STANDALONE_PATH, map_location=DEVICE))\n    else:\n        # If we just finished Stage 1 in this same memory space, we already have weights.\n        pass \n\n    # 3. Initialize Decoder\n    decoder_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\n    # 4. Optimizer\n    optimizer = optim.Adam(filter(lambda p: p.requires_grad, decoder_model.parameters()), lr=LR_STAGE2)\n    criterion = nn.CrossEntropyLoss(ignore_index=PAD_TOKEN_ID)\n\n    # 5. Load Optimizer if resuming Stage 2\n    if checkpoint and not checkpoint.get('stage1_complete', False) == False and 'optimizer_state' in checkpoint and not checkpoint.get('stage2_complete', False):\n        # We only load optimizer if we are strictly RESUMING stage 2, not starting it fresh\n        if start_epoch_s2 > 0:\n            print(\"Resuming Stage 2 Optimizer state.\")\n            optimizer.load_state_dict(checkpoint['optimizer_state'])\n            decoder_model.load_state_dict(checkpoint['decoder_model_state'])\n\n    stop_early = False\n    \n    for epoch in range(start_epoch_s2, NUM_EPOCHS_STAGE2):\n        elapsed_hours = (time.time() - START_TIME) / 3600\n        if elapsed_hours > MAX_RUNTIME_HOURS:\n            print(f\"\\n!!! TIME LIMIT REACHED ({elapsed_hours:.2f} hrs). Saving and Stopping. !!!\")\n            stop_early = True\n            save_relay_checkpoint({\n                'stage1_complete': True,\n                'stage2_complete': False,\n                'epoch': epoch, # Resume here\n                'neural_encoder_state': neural_enc.state_dict(), # Keep passing this along\n                'decoder_model_state': decoder_model.state_dict(),\n                'optimizer_state': optimizer.state_dict(),\n                's1_train_losses': s1_train_losses,\n                's1_val_losses': s1_val_losses,\n                's2_train_losses': s2_train_losses,\n                's2_val_losses': s2_val_losses\n            }, RELAY_CHECKPOINT_PATH)\n            break\n\n        print(f\"\\nEpoch {epoch+1}/{NUM_EPOCHS_STAGE2}\")\n        \n        # Train\n        decoder_model.train()\n        train_loss = 0\n        for batch in tqdm(train_loader, desc=\"Training S2\", 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        s2_train_losses.append(avg_train)\n\n        # Val\n        decoder_model.eval()\n        val_loss = 0\n        with torch.no_grad():\n            for batch in tqdm(val_loader, desc=\"Validating S2\", 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                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        s2_val_losses.append(avg_val)\n        print(f\"Train Loss: {avg_train:.4f} | Val Loss: {avg_val:.4f}\")\n\n    if not stop_early:\n        print(\"Stage 2 Finished. Saving complete state.\")\n        save_relay_checkpoint({\n            'stage1_complete': True,\n            'stage2_complete': True,\n            'epoch': 0,\n            'neural_encoder_state': neural_enc.state_dict(),\n            'decoder_model_state': decoder_model.state_dict(),\n            's1_train_losses': s1_train_losses,\n            's1_val_losses': s1_val_losses,\n            's2_train_losses': s2_train_losses,\n            's2_val_losses': s2_val_losses\n        }, RELAY_CHECKPOINT_PATH)\n    else:\n        print(\"Stopping execution due to time limit.\")\n        exit()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### Evaluation & Submission\n\n* Plots the loss curves (restored from checkpoint history).\n\n* Calculates the final Word Error Rate (WER) on the validation set.\n\n* Generates submission.csv for the test set.","metadata":{}},{"cell_type":"code","source":"# --- PLOTTING & EVALUATION (Runs only when finished) ---\n\nif not start_stage1 and not start_stage2:\n    print(\"\\n--- Training Complete. Generating Final Outputs ---\")\n    \n    # 1. Plot Combined Losses\n    plt.figure(figsize=(12, 5))\n    plt.subplot(1, 2, 1)\n    if s1_train_losses:\n        plt.plot(s1_train_losses, label='Train')\n        plt.plot(s1_val_losses, label='Val')\n        plt.title('Stage 1 (Contrastive) Loss')\n        plt.legend()\n    \n    plt.subplot(1, 2, 2)\n    if s2_train_losses:\n        plt.plot(s2_train_losses, label='Train')\n        plt.plot(s2_val_losses, label='Val')\n        plt.title('Stage 2 (Decoder) Loss')\n        plt.legend()\n    plt.show()\n\n    # 2. Load Model for Eval\n    neural_enc = NeuralEncoder(NEURAL_CHANNELS, NEURAL_CNN_OUT, NEURAL_LSTM_HIDDEN, EMBEDDING_DIM).to(DEVICE)\n    neural_enc.load_state_dict(checkpoint['neural_encoder_state'])\n    \n    eval_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    eval_model.load_state_dict(checkpoint['decoder_model_state'])\n    eval_model.eval()\n\n    # 3. Calculate WER\n    print(\"Calculating WER...\")\n    ground_truths = []\n    predictions = []\n    with torch.no_grad():\n        for batch in tqdm(val_loader, desc=\"WER Valid\"):\n            if not batch: continue\n            neural_data = batch['neural_data'].to(DEVICE)\n            raw_text_gt = batch['raw_text'] \n            predicted_tokens = eval_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    wer = jiwer.wer(ground_truths, predictions)\n    print(f\"\\nFinal WER: {wer * 100:.2f}%\")\n\n    # 4. Generate Submission\n    print(\"Generating Submission...\")\n    all_trial_keys = []\n    all_predictions = []\n    with torch.no_grad():\n        for batch in tqdm(test_loader, desc=\"Submission\"):\n            if not batch: continue\n            neural_data = batch['neural_data'].to(DEVICE)\n            trial_keys = batch['trial_keys']\n            predicted_tokens = eval_model.predict(neural_data)\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    submission_df = pd.DataFrame({'id': all_trial_keys, 'text': all_predictions})\n    submission_df['text'] = submission_df['text'].str.strip()\n    submission_df.to_csv(SUBMISSION_PATH, index=False)\n    print(f\"Saved submission to {SUBMISSION_PATH}\")\n    print(submission_df.head())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### t-SNE Visualization\nThis block creates the visualization to verify if the Stage 1 pre-training aligned the neural signals with the phonemes. This only runs at the very end.","metadata":{}},{"cell_type":"code","source":"# --- t-SNE VISUALIZATION ---\nif not start_stage1 and not start_stage2:\n    print(\"--- Starting t-SNE Visualization ---\")\n    \n    # 1. Load Encoders\n    try:\n        neural_enc = NeuralEncoder(NEURAL_CHANNELS, NEURAL_CNN_OUT, NEURAL_LSTM_HIDDEN, EMBEDDING_DIM).to(DEVICE)\n        phoneme_enc = PhonemeEncoder(PHONEME_VOCAB_SIZE, PHONEME_EMBED_DIM, PHONEME_LSTM_HIDDEN, EMBEDDING_DIM, PHONEME_PAD_ID).to(DEVICE)\n        \n        # Load from the relay checkpoint dict\n        neural_enc.load_state_dict(checkpoint['neural_encoder_state'])\n        \n        # Try loading phoneme encoder from standalone file or checkpoint\n        if os.path.exists(PHONEME_STANDALONE_PATH):\n             phoneme_enc.load_state_dict(torch.load(PHONEME_STANDALONE_PATH, map_location=DEVICE))\n        elif 'phoneme_encoder_state' in checkpoint:\n             phoneme_enc.load_state_dict(checkpoint['phoneme_encoder_state'])\n        else:\n             print(\"Warning: Phoneme encoder weights not found. Skipping t-SNE.\")\n             raise FileNotFoundError\n             \n        neural_enc.eval()\n        phoneme_enc.eval()\n        \n        # 2. Get Embeddings\n        neural_embeddings = []\n        phoneme_embeddings = []\n        labels = [] \n        num_batches_to_plot = 10\n        \n        with torch.no_grad():\n            for i, batch in enumerate(val_loader):\n                if not batch: continue\n                if i >= num_batches_to_plot: break\n\n                neural_data = batch['neural_data'].to(DEVICE)\n                phoneme_data = batch['phoneme_data'].to(DEVICE)\n                neu_emb = neural_enc.forward_contrastive(neural_data)\n                pho_emb = phoneme_enc.forward_contrastive(phoneme_data)\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                labels.extend([f\"sent_{i*BATCH_SIZE + j}\" for j in range(len(batch['raw_text']))])\n\n        neural_embeddings = torch.cat(neural_embeddings, dim=0).numpy()\n        phoneme_embeddings = torch.cat(phoneme_embeddings, dim=0).numpy()\n        all_embeddings = np.concatenate([neural_embeddings, phoneme_embeddings], axis=0)\n\n        # 3. Run t-SNE\n        tsne = TSNE(n_components=2, perplexity=30, n_iter=1000, random_state=42)\n        tsne_results = tsne.fit_transform(all_embeddings)\n\n        # 4. Plot\n        num_points = len(labels)\n        tsne_neural = tsne_results[:num_points]\n        tsne_phoneme = tsne_results[num_points:]\n\n        plt.figure(figsize=(14, 10))\n        plt.scatter(tsne_neural[:, 0], tsne_neural[:, 1], marker='x', c='blue', label='Neural Embeddings')\n        plt.scatter(tsne_phoneme[:, 0], tsne_phoneme[:, 1], marker='o', c='red', s=50, alpha=0.6, label='Phoneme Embeddings')\n        for 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        plt.title('t-SNE Visualization of Aligned Neural and Phoneme Embedding Space')\n        plt.xlabel('t-SNE Component 1')\n        plt.ylabel('t-SNE Component 2')\n        plt.legend()\n        plt.show()\n        print(\"If pre-training was successful, there should be blue 'x's (Neural) \"\n              \"and red 'o's (Phoneme) forming small, tight clusters, \"\n              \"connected by gray lines. This indicates the space is aligned.\")\n\n    except Exception as e:\n        print(f\"Skipping t-SNE due to error: {e}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}