{"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":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q jiwer huggingface_hub\nprint(\"Installs complete.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T23:32:30.995578Z","iopub.execute_input":"2025-12-11T23:32:30.996355Z","iopub.status.idle":"2025-12-11T23:32:37.370447Z","shell.execute_reply.started":"2025-12-11T23:32:30.996323Z","shell.execute_reply":"2025-12-11T23:32:37.369691Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob, time, math, random\nfrom tqdm.notebook import tqdm\nimport numpy as np\nimport pandas as pd\nimport h5py\nimport matplotlib.pyplot as plt\nfrom sklearn.manifold import TSNE\n\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\nimport jiwer\nfrom huggingface_hub import HfApi, hf_hub_download, login","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T23:32:37.371894Z","iopub.execute_input":"2025-12-11T23:32:37.372127Z","iopub.status.idle":"2025-12-11T23:32:42.76963Z","shell.execute_reply.started":"2025-12-11T23:32:37.372102Z","shell.execute_reply":"2025-12-11T23:32:42.768808Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------- USER EDITABLE ----------------\nHF_TOKEN = \"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\"  # <-- REPLACE with your HF write token or set to None\nHF_REPO_ID = \"maryadaaa/brain-to-text-checkpoints\"  # <-- replace 'username/...'\nDATA_DIR = \"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final/t15.2025.03.14/\"  # <-- adjust if necessary\nWORKING_DIR = \"/kaggle/working/\"\n# ---------------- END EDITABLE -----------------\n\nos.makedirs(WORKING_DIR, exist_ok=True)\n\n# Device / reproducibility\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\nif torch.cuda.is_available(): torch.cuda.manual_seed_all(SEED)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T23:32:42.770459Z","iopub.execute_input":"2025-12-11T23:32:42.770917Z","iopub.status.idle":"2025-12-11T23:32:42.844985Z","shell.execute_reply.started":"2025-12-11T23:32:42.770893Z","shell.execute_reply":"2025-12-11T23:32:42.844233Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- HYPERPARAMETERS (Optimized for Best Results) ---\n# Increased capacity to capture complex neural patterns\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     \n\n# Decoder Settings\nDECODER_NHEAD = 8             \nDECODER_NLAYERS = 6           \nDECODER_DIM_FEEDFORWARD = 2048 \n\n# Training Schedule (Designed for Multi-Session Relay)\nBATCH_SIZE = 16               \nNUM_EPOCHS_STAGE1 = 250       # Target ~10 hours\nNUM_EPOCHS_STAGE2 = 600       # Target ~30 hours\nLR_STAGE1 = 1e-4              \nLR_STAGE2 = 1e-4              \nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nWER_CHECK_INTERVAL = 25\n\n# Save paths\nENCODER_CHECKPOINT_PATH = os.path.join(WORKING_DIR, \"neural_encoder_stage1.pth\")\nPHONEME_ENCODER_CHECKPOINT_PATH = os.path.join(WORKING_DIR, \"phoneme_encoder_stage1.pth\")\nDECODER_CHECKPOINT_PATH = os.path.join(WORKING_DIR, \"full_decoder_stage2.pth\")\n\nprint(\"Device:\", DEVICE)\nprint(\"Data dir:\", DATA_DIR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T23:32:42.845725Z","iopub.execute_input":"2025-12-11T23:32:42.845919Z","iopub.status.idle":"2025-12-11T23:32:42.858461Z","shell.execute_reply.started":"2025-12-11T23:32:42.845903Z","shell.execute_reply":"2025-12-11T23:32:42.857852Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 3: Dataset, Tokenizer, and Collator (FIXED)\nimport h5py\nimport torch\nimport numpy as np\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.nn.utils.rnn import pad_sequence\nfrom typing import List, Dict\n\nclass SimpleTokenizer:\n    def __init__(self):\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, sentences: List[str]):\n        words = set()\n        for s in sentences:\n            words.update(s.lower().split())\n        for w in sorted(list(words)):\n            if w not in self.vocab:\n                self.vocab[w] = self.vocab_size\n                self.inv_vocab[self.vocab_size] = w\n                self.vocab_size += 1\n                \n    def encode(self, text: str):\n        toks = [self.vocab.get(w, self.vocab['<UNK>']) for w in text.lower().split()]\n        return torch.LongTensor([self.vocab['<SOS>']] + toks + [self.vocab['<EOS>']])\n        \n    def decode(self, token_tensor):\n        words = [self.inv_vocab.get(int(t), '?') for t in token_tensor]\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\n\nclass BrainToTextDataset(Dataset):\n    def __init__(self, hdf5_files: List[str]):\n        self.trial_index = []\n        for f in hdf5_files:\n            try:\n                with h5py.File(f,'r') as hf:\n                    for k in hf.keys():\n                        self.trial_index.append((f,k))\n            except Exception as e:\n                print(f\"Warning reading {f}: {e}\")\n        print(f\"Found {len(self.trial_index)} total trials.\")\n\n    def __len__(self): return len(self.trial_index)\n    \n    def __getitem__(self, idx):\n        fpath, key = self.trial_index[idx]\n        with h5py.File(fpath,'r') as hf:\n            grp = hf[key]\n            # 1. Load Neural Data & Permute\n            neural_data = torch.from_numpy(grp['input_features'][:]).float()\n            neural_data = neural_data.permute(1, 0) # (Channels, Time)\n            \n            # 2. Test Detection\n            is_test = 'seq_class_ids' not in grp\n            \n            phoneme_data = torch.tensor([], dtype=torch.long)\n            text_data = \"\"\n            \n            if not is_test:\n                if 'seq_class_ids' in grp:\n                    phoneme_data = torch.from_numpy(grp['seq_class_ids'][:]).long()\n                \n                # Load Text from attributes or datasets\n                if 'sentence_label' in grp.attrs:\n                    text_data = grp.attrs['sentence_label']\n                elif 'transcription' in grp:\n                    raw = grp['transcription'][()]\n                    text_data = raw.decode('utf-8') if isinstance(raw, bytes) else str(raw)\n                elif 'sentence_label' in grp:\n                    raw = grp['sentence_label'][()]\n                    text_data = raw.decode('utf-8') if isinstance(raw, bytes) else str(raw)\n\n        return {\n            \"neural_data\": neural_data,\n            \"phoneme_data\": phoneme_data,\n            \"text_data\": text_data,\n            \"trial_key\": key,\n            \"is_test\": is_test\n        }\n\n    def get_all_sentences(self):\n        all_sentences = []\n        for i in range(len(self)):\n            try:\n                item = self.__getitem__(i)\n                if not item[\"is_test\"] and item[\"text_data\"]:\n                    all_sentences.append(item[\"text_data\"])\n            except: pass\n        return all_sentences\n\nclass BrainCollator:\n    def __init__(self, tokenizer, pad_token_id=0, phoneme_pad_id=0):\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):\n        # Filter bad data\n        batch = [b for b in batch if b[\"text_data\"] != \"\" or b[\"is_test\"]]\n        if not batch: return {}\n        \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 ---\n        if train_val_batch:\n            neural_list = [b[\"neural_data\"] for b in train_val_batch]\n            phoneme_list = [b[\"phoneme_data\"] for b in train_val_batch]\n            text_list = [b[\"text_data\"] for b in train_val_batch]\n            \n            # Pad Neural\n            neural_permuted = [d.permute(1, 0) for d in neural_list] \n            padded_neural = pad_sequence(neural_permuted, batch_first=True, padding_value=0.0).permute(0, 2, 1)\n            padded_phonemes = pad_sequence(phoneme_list, batch_first=True, padding_value=self.phoneme_pad_id)\n            \n            text_tokens = [self.tokenizer.encode(t) for t in text_list]\n            padded_text = pad_sequence(text_tokens, batch_first=True, padding_value=self.pad_token_id)\n            \n            return {\n                'neural': padded_neural,\n                'phonemes': padded_phonemes,\n                'text_input': padded_text[:, :-1],\n                'text_target': padded_text[:, 1:],\n                'text_padding_mask': (padded_text[:, :-1] == self.pad_token_id),\n                'raw_text': text_list,\n                'trial_keys': [b['trial_key'] for b in train_val_batch], # <--- ADDED THIS LINE\n                'is_test_batch': False\n            }\n            \n        # --- PROCESS TEST ---\n        elif test_batch:\n            neural_list = [b[\"neural_data\"] for b in test_batch]\n            trial_keys = [b[\"trial_key\"] for b in test_batch]\n            \n            neural_permuted = [d.permute(1, 0) for d in neural_list]\n            padded_neural = pad_sequence(neural_permuted, batch_first=True, padding_value=0.0).permute(0, 2, 1)\n            \n            return {\n                'neural': padded_neural,\n                'trial_keys': trial_keys,\n                'is_test_batch': True\n            }\n        return {}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T23:32:42.860396Z","iopub.execute_input":"2025-12-11T23:32:42.860769Z","iopub.status.idle":"2025-12-11T23:32:42.882351Z","shell.execute_reply.started":"2025-12-11T23:32:42.860752Z","shell.execute_reply":"2025-12-11T23:32:42.881725Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 4 — Models: encoder(s), contrastive wrapper, decoder\n# Temporal mask module used in your PDF\nclass TemporalMasking(nn.Module):\n    def __init__(self, p=0.1, mask_span_length=10):\n        super().__init__()\n        self.p = p; self.mask_span_length = mask_span_length\n    def forward(self, x):\n        if (not self.training) or self.p == 0: return x\n        B, C, T = x.shape\n        for i in range(B):\n            if torch.rand(1).item() < self.p:\n                start = torch.randint(0, max(1, T-self.mask_span_length), (1,)).item()\n                x[i, :, start:start+self.mask_span_length] = 0.0\n        return x\n\n# Neural encoder: 1D CNN -> BiLSTM -> two projections (contrastive + decoder seq)\nclass NeuralEncoder(nn.Module):\n    def __init__(self, input_channels=NEURAL_CHANNELS, cnn_out=NEURAL_CNN_OUT, lstm_hidden=NEURAL_LSTM_HIDDEN, embedding_dim=EMBEDDING_DIM):\n        super().__init__()\n        self.cnn = nn.Sequential(\n            nn.Conv1d(input_channels, cnn_out, kernel_size=3, padding=1),\n            nn.ReLU(),\n            nn.BatchNorm1d(cnn_out),\n            nn.MaxPool1d(kernel_size=2)\n        )\n        self.bilstm = nn.LSTM(cnn_out, lstm_hidden, num_layers=2, bidirectional=True, batch_first=True, dropout=0.2)\n        self.contrastive_proj = nn.Linear(lstm_hidden*2, embedding_dim)\n        self.decoder_proj = nn.Linear(lstm_hidden*2, embedding_dim)\n    def forward_contrastive(self, x):   # x: (B, C, T)\n        x = self.cnn(x)                 # (B, C_out, T_down)\n        x = x.permute(0,2,1)            # (B, T_down, C_out)\n        out, _ = self.bilstm(x)\n        pooled = out.mean(dim=1)        # (B, hidden*2)\n        return self.contrastive_proj(pooled)\n    def forward_decoder(self, x):       # returns sequence embeddings (B, T_down, E)\n        x = self.cnn(x)\n        x = x.permute(0,2,1)\n        out, _ = self.bilstm(x)\n        return self.decoder_proj(out)\n\nclass PhonemeEncoder(nn.Module):\n    def __init__(self, vocab_size=PHONEME_VOCAB_SIZE, emb_dim=PHONEME_EMBED_DIM, hidden=128, out_dim=EMBEDDING_DIM, pad_id=0):\n        super().__init__()\n        self.emb = nn.Embedding(vocab_size, emb_dim, padding_idx=pad_id)\n        self.lstm = nn.LSTM(emb_dim, hidden, batch_first=True, bidirectional=True)\n        self.proj = nn.Linear(hidden*2, out_dim)\n    def forward_contrastive(self, x):\n        x = self.emb(x)\n        out, _ = self.lstm(x)\n        pooled = out.mean(dim=1)\n        return self.proj(pooled)\n\n# Contrastive wrapper for InfoNCE\nclass ContrastiveModel(nn.Module):\n    def __init__(self, neural_enc, phoneme_enc, temp=0.1):\n        super().__init__()\n        self.neural_enc = neural_enc\n        self.phoneme_enc = phoneme_enc\n        self.temperature = temp\n        self.aug = TemporalMasking(p=0.1, mask_span_length=20)\n    def forward(self, neural, phonemes):\n        neural = self.aug(neural)\n        n_emb = F.normalize(self.neural_enc.forward_contrastive(neural), dim=1)\n        p_emb = F.normalize(self.phoneme_enc.forward_contrastive(phonemes), dim=1)\n        logits = torch.matmul(n_emb, p_emb.T) / self.temperature\n        labels = torch.arange(logits.size(0), device=logits.device)\n        loss = (F.cross_entropy(logits, labels) + F.cross_entropy(logits.T, labels)) / 2.0\n        return loss\n\n# Decoder model: frozen encoder -> Transformer decoder\nclass PositionalEncoding(nn.Module):\n    def __init__(self, d_model, dropout=0.1, max_len=5000):\n        super().__init__()\n        self.dropout = nn.Dropout(dropout)\n        pe = torch.zeros(max_len, d_model)\n        position = torch.arange(0, max_len).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0)/d_model))\n        pe[:, 0::2] = torch.sin(position * div_term)\n        pe[:, 1::2] = torch.cos(position * div_term)\n        self.register_buffer('pe', pe.unsqueeze(1))\n    def forward(self, x):\n        # x: (B, T, E)\n        x = x + self.pe[:x.size(1)].permute(1,0,2)\n        return self.dropout(x)\n\nclass DecoderModel(nn.Module):\n    def __init__(self, neural_encoder, vocab_size, embedding_dim=EMBEDDING_DIM, nhead=DECODER_NHEAD, nlayers=DECODER_NLAYERS, dim_feedforward=DECODER_DIM_FEEDFORWARD, pad_token_id=0, sos_token_id=1, eos_token_id=2):\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 = sos_token_id\n        self.eos = eos_token_id\n        # optionally freeze encoder (per plan)\n        self.freeze_encoder()\n        self.embedding = nn.Embedding(vocab_size, embedding_dim, padding_idx=pad_token_id)\n        self.pos_enc = PositionalEncoding(embedding_dim)\n        layer = nn.TransformerDecoderLayer(d_model=embedding_dim, nhead=nhead, dim_feedforward=dim_feedforward, batch_first=True)\n        self.decoder = nn.TransformerDecoder(layer, num_layers=nlayers)\n        self.fc = nn.Linear(embedding_dim, vocab_size)\n    def freeze_encoder(self):\n        for p in self.neural_encoder.parameters():\n            p.requires_grad = False\n        self.neural_encoder.eval()\n    def unfreeze_encoder(self):\n        for p in self.neural_encoder.parameters():\n            p.requires_grad = True\n        self.neural_encoder.train()\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    def forward(self, neural_data, tgt_tokens, tgt_key_padding_mask):\n        # neural_data: (B, C, T) -> memory (B, S, E)\n        with torch.no_grad():\n            memory = self.neural_encoder.forward_decoder(neural_data)\n        tgt_emb = self.embedding(tgt_tokens) * math.sqrt(self.d_model)\n        tgt_emb = self.pos_enc(tgt_emb)\n        tgt_mask = self.generate_square_subsequent_mask(tgt_tokens.size(1)).to(tgt_tokens.device)\n        out = self.decoder(tgt=tgt_emb, memory=memory, tgt_mask=tgt_mask, tgt_key_padding_mask=tgt_key_padding_mask)\n        return self.fc(out)\n    def predict(self, neural_data, max_len=100):\n        # Autoregressive greedy generation, returns token sequences including SOS\n        self.eval()\n        device = neural_data.device\n        with torch.no_grad():\n            memory = self.neural_encoder.forward_decoder(neural_data)\n            batch_size = neural_data.size(0)\n            tgt_tokens = torch.full((batch_size,1), self.sos, dtype=torch.long, device=device)\n            finished = torch.zeros(batch_size, dtype=torch.bool, device=device)\n            for _ in range(max_len):\n                tgt_emb = self.embedding(tgt_tokens) * math.sqrt(self.d_model)\n                tgt_emb = self.pos_enc(tgt_emb)\n                tgt_mask = self.generate_square_subsequent_mask(tgt_tokens.size(1)).to(device)\n                out = self.decoder(tgt=tgt_emb, memory=memory, tgt_mask=tgt_mask)\n                logits = self.fc(out)   # (B, L, V)\n                next_token = logits[:, -1, :].argmax(-1).unsqueeze(1)  # (B,1)\n                tgt_tokens = torch.cat([tgt_tokens, next_token], dim=1)\n                finished = finished | (next_token.squeeze(1) == self.eos)\n                if finished.all(): break\n        return tgt_tokens  # includes SOS at 0 index\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T23:32:42.883195Z","iopub.execute_input":"2025-12-11T23:32:42.883492Z","iopub.status.idle":"2025-12-11T23:32:42.905259Z","shell.execute_reply.started":"2025-12-11T23:32:42.883472Z","shell.execute_reply":"2025-12-11T23:32:42.904689Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 5: Data Setup and Vocab Building\nclass ConfigWrapper:\n    pass\n\ndef setup_data(cfg):\n    \"\"\"Prepares datasets, tokenizer, and dataloaders. [cite: 472]\"\"\"\n    print(\"Finding HDF5 files...\")\n    # Update glob patterns to match your folder structure if needed\n    # The PDF assumes subfolders, but we will use your flat directory logic from before if needed.\n    # We'll use os.listdir to be safe as per your previous cells.\n    all_files = sorted([os.path.join(cfg.DATA_DIR, f) for f in os.listdir(cfg.DATA_DIR) if f.endswith('.h5') or f.endswith('.hdf5')])\n    \n    if not all_files: raise FileNotFoundError(f\"No files in {cfg.DATA_DIR}\")\n    \n    # Split 90/10\n    split = int(0.9 * len(all_files))\n    train_files = all_files[:split]\n    val_files = all_files[split:]\n    # Ideally you have a separate test folder, but we will use val_files for 'test' generation if no specific test set exists\n    \n    print(\"Initializing datasets...\")\n    train_dataset = BrainToTextDataset(train_files)\n    val_dataset = BrainToTextDataset(val_files)\n    # Using val files as test for demonstration if no dedicated test files\n    test_dataset = BrainToTextDataset(val_files) \n    \n    print(\"Building tokenizer vocabulary...\")\n    tokenizer = SimpleTokenizer()\n    all_sentences = train_dataset.get_all_sentences() # [cite: 482]\n    tokenizer.build_vocab(all_sentences)\n    vocab_size = tokenizer.get_vocab_size() # [cite: 484]\n    print(f\"Built vocab with {vocab_size} tokens.\")\n    \n    collator = BrainCollator(tokenizer, pad_token_id=0, phoneme_pad_id=0) # [cite: 486]\n    \n    train_loader = DataLoader(train_dataset, batch_size=cfg.BATCH_SIZE, shuffle=True, collate_fn=collator, num_workers=2)\n    val_loader = DataLoader(val_dataset, batch_size=cfg.BATCH_SIZE, shuffle=False, collate_fn=collator, num_workers=2)\n    test_loader = DataLoader(test_dataset, batch_size=cfg.BATCH_SIZE, shuffle=False, collate_fn=collator, num_workers=2)\n    \n    return train_loader, val_loader, test_loader, tokenizer, vocab_size\n\n# --- Execute Setup ---\nconfig = ConfigWrapper()\nconfig.DATA_DIR = DATA_DIR\nconfig.BATCH_SIZE = BATCH_SIZE\n\ntry:\n    # This call sets up everything and defines VOCAB_SIZE [cite: 515]\n    train_loader, val_loader, test_loader, tokenizer, VOCAB_SIZE = setup_data(config)\n    print(\"Data setup complete.\")\nexcept Exception as e:\n    print(f\"Setup failed: {e}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T23:32:42.905896Z","iopub.execute_input":"2025-12-11T23:32:42.90605Z","iopub.status.idle":"2025-12-11T23:32:45.353084Z","shell.execute_reply.started":"2025-12-11T23:32:42.906037Z","shell.execute_reply":"2025-12-11T23:32:45.352384Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 6 — Stage 1 Contrastive training\nneural_enc = NeuralEncoder().to(DEVICE)\nphoneme_enc = PhonemeEncoder(vocab_size=PHONEME_VOCAB_SIZE, emb_dim=PHONEME_EMBED_DIM, hidden=128, out_dim=EMBEDDING_DIM).to(DEVICE)\ncontrastive = ContrastiveModel(neural_enc, phoneme_enc).to(DEVICE)\n\nopt1 = optim.Adam(contrastive.parameters(), lr=LR_STAGE1)\nstage1_history = {'train_loss':[], 'val_loss':[]}\nbest_val = float('inf')\n\nfor epoch in range(NUM_EPOCHS_STAGE1):\n    contrastive.train()\n    train_loss = 0.0\n    for batch in tqdm(train_loader, desc=f\"S1 Epoch {epoch+1}/{NUM_EPOCHS_STAGE1}\", leave=False):\n        if not batch or batch.get('is_test_batch', False): continue\n        opt1.zero_grad()\n        neural = batch['neural'].to(DEVICE)\n        phonemes = batch['phonemes'].to(DEVICE)\n        loss = contrastive(neural, phonemes)\n        loss.backward()\n        opt1.step()\n        train_loss += loss.item()\n    avg_train = train_loss / max(1, len(train_loader))\n    stage1_history['train_loss'].append(avg_train)\n\n    # validation\n    contrastive.eval()\n    val_loss = 0.0\n    with torch.no_grad():\n        for batch in val_loader:\n            if not batch or batch.get('is_test_batch', False): continue\n            neural = batch['neural'].to(DEVICE)\n            phonemes = batch['phonemes'].to(DEVICE)\n            val_loss += contrastive(neural, phonemes).item()\n    avg_val = val_loss / max(1, len(val_loader))\n    stage1_history['val_loss'].append(avg_val)\n\n    if avg_val < best_val:\n        best_val = avg_val\n        torch.save(neural_enc.state_dict(), ENCODER_CHECKPOINT_PATH)\n        torch.save(phoneme_enc.state_dict(), PHONEME_ENCODER_CHECKPOINT_PATH)\n        print(f\"Stage1: saved best encoders at epoch {epoch+1} (val_loss={avg_val:.4f})\")\n\n    if (epoch+1) % 5 == 0 or epoch==0:\n        print(f\"S1 Epoch {epoch+1} train {avg_train:.4f} val {avg_val:.4f}\")\n\nprint(\"Stage-1 training complete.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T23:32:45.353813Z","iopub.execute_input":"2025-12-11T23:32:45.354015Z","iopub.status.idle":"2025-12-11T23:43:27.04024Z","shell.execute_reply.started":"2025-12-11T23:32:45.353998Z","shell.execute_reply":"2025-12-11T23:43:27.039066Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 7 — Stage 2 Decoder training \nRUN_ABLATION_NO_PRETRAIN = False  # set True to train decoder with random encoder (ablation)\n# instantiate new encoder, load weights unless ablation requested\nneural_enc_for_decoder = NeuralEncoder()\nif not RUN_ABLATION_NO_PRETRAIN:\n    if os.path.exists(ENCODER_CHECKPOINT_PATH):\n        neural_enc_for_decoder.load_state_dict(torch.load(ENCODER_CHECKPOINT_PATH, map_location=DEVICE))\n        print(\"Loaded pretrained encoder weights for Stage-2.\")\n    else:\n        print(\"Warning: pretrained encoder weights not found; Stage-2 will proceed with current init.\")\nelse:\n    print(\"Running ABLATION: decoder training WITHOUT pre-trained encoder (encoder will be unfrozen).\")\n\ndecoder = DecoderModel(neural_encoder=neural_enc_for_decoder, vocab_size=VOCAB_SIZE).to(DEVICE)\nif RUN_ABLATION_NO_PRETRAIN:\n    decoder.unfreeze_encoder()  # allow joint training for ablation\n\nopt2 = optim.Adam(filter(lambda p: p.requires_grad, decoder.parameters()), lr=LR_STAGE2)\ncriterion = nn.CrossEntropyLoss(ignore_index=0)\n\nstage2_history = {'train_loss':[], 'val_loss':[], 'val_wer':[], 'wer_epochs':[]}\nbest_val = float('inf')\n\nfor epoch in range(NUM_EPOCHS_STAGE2):\n    decoder.train()\n    train_loss = 0.0\n    for batch in tqdm(train_loader, desc=f\"S2 Epoch {epoch+1}/{NUM_EPOCHS_STAGE2}\", leave=False):\n        if not batch or batch.get('is_test_batch', False): continue\n        opt2.zero_grad()\n        neural = batch['neural'].to(DEVICE)\n        text_in = batch['text_input'].to(DEVICE)\n        text_tgt = batch['text_target'].to(DEVICE)\n        pad_mask = batch['text_padding_mask'].to(DEVICE)\n        logits = decoder(neural, text_in, tgt_key_padding_mask=pad_mask)\n        loss = criterion(logits.reshape(-1, logits.size(-1)), text_tgt.reshape(-1))\n        loss.backward()\n        opt2.step()\n        train_loss += loss.item()\n    avg_train = train_loss / max(1, len(train_loader))\n    stage2_history['train_loss'].append(avg_train)\n\n    # validation loss\n    decoder.eval()\n    val_loss = 0.0\n    with torch.no_grad():\n        for batch in val_loader:\n            if not batch or batch.get('is_test_batch', False): continue\n            neural = batch['neural'].to(DEVICE)\n            text_in = batch['text_input'].to(DEVICE)\n            text_tgt = batch['text_target'].to(DEVICE)\n            pad_mask = batch['text_padding_mask'].to(DEVICE)\n            logits = decoder(neural, text_in, tgt_key_padding_mask=pad_mask)\n            val_loss += criterion(logits.reshape(-1, logits.size(-1)), text_tgt.reshape(-1)).item()\n    avg_val = val_loss / max(1, len(val_loader))\n    stage2_history['val_loss'].append(avg_val)\n\n    # periodic WER evaluation\n    if (epoch+1) % WER_CHECK_INTERVAL == 0 or epoch == (NUM_EPOCHS_STAGE2-1):\n        print(\"Computing WER on validation...\")\n        refs, hyps = [], []\n        with torch.no_grad():\n            for batch in val_loader:\n                if not batch or batch.get('is_test_batch', False): continue\n                neu = batch['neural'].to(DEVICE)\n                pred_tokens = decoder.predict(neu, max_len=100)\n                for i in range(pred_tokens.size(0)):\n                    pred_text = tokenizer.decode(pred_tokens[i])\n                    refs.append(batch['raw_text'][i].lower())\n                    hyps.append(pred_text.lower())\n        cur_wer = jiwer.wer(refs, hyps)\n        stage2_history['val_wer'].append(cur_wer)\n        stage2_history['wer_epochs'].append(epoch+1)\n        print(f\"Epoch {epoch+1} WER: {cur_wer*100:.2f}%\")\n\n    # save best\n    if avg_val < best_val:\n        best_val = avg_val\n        torch.save(decoder.state_dict(), DECODER_CHECKPOINT_PATH)\n        print(f\"Saved best decoder at epoch {epoch+1} (val_loss={avg_val:.4f})\")\n\n    if (epoch+1) % 5 == 0 or epoch == 0:\n        print(f\"S2 Epoch {epoch+1} train {avg_train:.4f} val {avg_val:.4f}\")\n\nprint(\"Stage-2 training complete.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T23:43:27.04172Z","iopub.execute_input":"2025-12-11T23:43:27.042628Z","iopub.status.idle":"2025-12-12T00:09:59.034994Z","shell.execute_reply.started":"2025-12-11T23:43:27.042599Z","shell.execute_reply":"2025-12-12T00:09:59.033972Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 8: Evaluation, Submission, and Visualization\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom sklearn.manifold import TSNE\nimport torch.nn.functional as F\n\nprint(\"\\n=== PART 1: Generating Submission ===\")\n\n# 1. Load Evaluation Model\neval_model = DecoderModel(\n    neural_encoder=NeuralEncoder(), \n    vocab_size=VOCAB_SIZE\n).to(DEVICE)\n\nif os.path.exists(DECODER_CHECKPOINT_PATH):\n    eval_model.load_state_dict(torch.load(DECODER_CHECKPOINT_PATH, map_location=DEVICE))\n    print(f\"Loaded weights from {DECODER_CHECKPOINT_PATH}\")\nelse:\n    print(\"Warning: No checkpoint found. Using untrained model.\")\neval_model.eval()\n\n# 2. Generate Predictions\nall_trial_keys = []\nall_predictions = []\n\nwith torch.no_grad():\n    for batch in tqdm(test_loader, desc=\"Generating Submission\"):\n        if not batch: continue\n        \n        neural_data = batch['neural'].to(DEVICE)\n        \n        # Retrieve keys (now available for all batch types)\n        if 'trial_keys' in batch:\n            trial_keys = batch['trial_keys']\n        elif 'trial_key' in batch: \n            trial_keys = batch['trial_key']\n        else:\n            print(\"Warning: Batch missing trial_keys\")\n            continue \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# 3. Save to CSV\nif len(all_predictions) > 0:\n    submission_df = pd.DataFrame({'id': all_trial_keys, 'text': all_predictions})\n    submission_df['text'] = submission_df['text'].str.strip()\n    submission_path = os.path.join(WORKING_DIR, \"submission.csv\")\n    submission_df.to_csv(submission_path, index=False)\n    print(f\"Saved submission to: {submission_path}\")\n    print(submission_df.head())\nelse:\n    print(\"\\n⚠️ WARNING: No predictions were generated! Check your test_loader and collator logic.\")\n\n# ==========================================\n# PART 2: PLOT LOSS CURVES\n# ==========================================\nprint(\"\\n=== PART 2: Plotting Loss Curves ===\")\nplt.figure(figsize=(18,6))\n\n# Plot 1: Stage 1 Loss\nplt.subplot(1, 3, 1)\nif 'stage1_history' in globals() 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\nplt.subplot(1, 3, 2)\nif 'stage2_history' in globals() 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\nplt.subplot(1, 3, 3)\nif 'stage2_history' in globals() and len(stage2_history['val_wer']) > 0:\n    wer_pct = [w*100 for w in stage2_history['val_wer']]\n    plt.plot(stage2_history['wer_epochs'], wer_pct, marker='o', color='green', 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    for x, y in zip(stage2_history['wer_epochs'], wer_pct):\n        plt.annotate(f\"{y:.1f}%\", (x, y), 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.savefig(os.path.join(WORKING_DIR, \"training_plots.png\"))\nplt.show()\n\n# ==========================================\n# PART 3: t-SNE VISUALIZATION\n# ==========================================\nprint(\"\\n=== PART 3: t-SNE Visualization ===\")\ntry:\n    tsne_neural_enc = NeuralEncoder(\n        input_channels=NEURAL_CHANNELS, cnn_out=NEURAL_CNN_OUT, \n        lstm_hidden=NEURAL_LSTM_HIDDEN, embedding_dim=EMBEDDING_DIM\n    ).to(DEVICE)\n    \n    tsne_phoneme_enc = PhonemeEncoder(\n        vocab_size=PHONEME_VOCAB_SIZE, emb_dim=PHONEME_EMBED_DIM, \n        hidden=NEURAL_LSTM_HIDDEN, out_dim=EMBEDDING_DIM\n    ).to(DEVICE)\n\n    if os.path.exists(ENCODER_CHECKPOINT_PATH):\n        tsne_neural_enc.load_state_dict(torch.load(ENCODER_CHECKPOINT_PATH, map_location=DEVICE))\n        tsne_phoneme_enc.load_state_dict(torch.load(PHONEME_ENCODER_CHECKPOINT_PATH, map_location=DEVICE))\n        tsne_neural_enc.eval(); tsne_phoneme_enc.eval()\n\n        neural_embs, phoneme_embs = [], []\n        print(\"Extracting embeddings...\")\n        with torch.no_grad():\n            for i, batch in enumerate(val_loader):\n                if i >= 10: break\n                if not batch or batch.get('is_test_batch', False): continue\n                \n                n_data = batch['neural'].to(DEVICE)\n                p_data = batch['phonemes'].to(DEVICE)\n                \n                n_out = F.normalize(tsne_neural_enc.forward_contrastive(n_data), dim=1)\n                p_out = F.normalize(tsne_phoneme_enc.forward_contrastive(p_data), dim=1)\n                \n                neural_embs.append(n_out.cpu().numpy())\n                phoneme_embs.append(p_out.cpu().numpy())\n\n        if neural_embs:\n            neural_embs = np.concatenate(neural_embs, axis=0)\n            phoneme_embs = np.concatenate(phoneme_embs, axis=0)\n            combined = np.concatenate([neural_embs, phoneme_embs], axis=0)\n            \n            print(f\"Running t-SNE on {combined.shape[0]} points...\")\n            tsne = TSNE(n_components=2, perplexity=30, random_state=42)\n            results = tsne.fit_transform(combined)\n            \n            n_points = neural_embs.shape[0]\n            plt.figure(figsize=(10, 10))\n            plt.scatter(results[:n_points, 0], results[:n_points, 1], marker='x', c='blue', alpha=0.6, label='Neural')\n            plt.scatter(results[n_points:, 0], results[n_points:, 1], marker='o', c='red', alpha=0.6, label='Phoneme')\n            for k in range(min(100, n_points)):\n                plt.plot([results[k,0], results[n_points+k,0]], [results[k,1], results[n_points+k,1]], c='gray', alpha=0.3, linewidth=0.5)\n            plt.legend(); plt.title(\"t-SNE: Neural vs Phoneme Alignment\")\n            plt.savefig(os.path.join(WORKING_DIR, \"tsne_visualization.png\"))\n            plt.show()\n        else:\n            print(\"No valid data found for t-SNE.\")\n    else:\n        print(\"Stage 1 checkpoints not found. Skipping t-SNE.\")\nexcept Exception as e:\n    print(f\"t-SNE Visualization skipped: {e}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-12T00:09:59.036238Z","iopub.execute_input":"2025-12-12T00:09:59.036501Z","iopub.status.idle":"2025-12-12T00:10:01.497787Z","shell.execute_reply.started":"2025-12-12T00:09:59.036477Z","shell.execute_reply":"2025-12-12T00:10:01.497069Z"}},"outputs":[],"execution_count":null}]}