{"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":"nvidiaTeslaT4","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":"\"\"\"\n================================================================================\nALL-IN-ONE NOTEBOOK: PHASE 4 - TUNED LSTM\n================================================================================\nArchitecture: Strided CNN + BiLSTM + SpecAugment\nGoal: Push WER below 25% by using advanced regularization.\nChanges:\n1. Reverted to LSTM (Transformer proved too data-hungry).\n3. Increased epochs to 60 (as augmentation makes training harder initially).\n================================================================================\n\"\"\"\n\n# ============================================================================\n# 1. SETUP & IMPORTS\n# ============================================================================\nimport os\nos.system('pip install jiwer')\n\nimport h5py\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.nn.utils.rnn import pad_sequence, pack_padded_sequence, pad_packed_sequence\n\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nimport warnings\n\nwarnings.filterwarnings('ignore')\n\nCONFIG = {\n    'data_dir': '/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final/',\n    'device': 'cuda' if torch.cuda.is_available() else 'cpu',\n    'batch_size': 64,\n    'num_epochs': 40,\n    'hidden_dim': 256,\n    'num_layers': 1,\n    'dropout': 0.3,\n    'learning_rate': 1e-3,\n}\n\nprint(\"Device:\", CONFIG['device'])\n\n# ============================================================================\n# 2. DATA LOADING & DATASET\n# ============================================================================\ndef load_split(data_dir, split='train'):\n    from glob import glob\n    files = sorted(glob(f'{data_dir}/**/data_{split}.hdf5', recursive=True))\n\n    all_data = {'neural': [], 'n_steps': [], 'sentence': []}\n\n    for filepath in tqdm(files, desc=f\"Loading {split}\"):\n        with h5py.File(filepath, 'r') as f:\n            for trial_key in f.keys():\n                trial = f[trial_key]\n                neural = trial['input_features'][:]\n                n_steps = trial.attrs['n_time_steps']\n                sentence = trial.attrs.get('sentence_label')\n                if isinstance(sentence, bytes):\n                    sentence = sentence.decode('utf-8')\n\n                all_data['neural'].append(neural)\n                all_data['n_steps'].append(n_steps)\n                all_data['sentence'].append(sentence)\n\n    return all_data\n\n\nclass BrainToTextDataset(Dataset):\n    def __init__(self, data, char2idx=None):\n        self.neural = data['neural']\n        self.n_steps = data['n_steps']\n        self.sentences = data['sentence']\n\n        self.char2idx = char2idx if char2idx else self._build_vocab()\n        self.idx2char = {v: k for k, v in self.char2idx.items()}\n\n    def _build_vocab(self):\n        chars = set()\n        for s in self.sentences:\n            if s:\n                chars.update(s.lower())\n        chars = sorted(chars)\n        return {'<BLANK>': 0, **{c: i+1 for i, c in enumerate(chars)}}\n\n    def __len__(self):\n        return len(self.neural)\n\n    def __getitem__(self, idx):\n        x = self.neural[idx][:self.n_steps[idx]]\n        x = (x - x.mean()) / (x.std() + 1e-8)\n\n        sentence = self.sentences[idx] or \"\"\n        target = [self.char2idx[c] for c in sentence.lower()]\n\n        return {\n            'neural': torch.FloatTensor(x),\n            'target': torch.LongTensor(target),\n            'length': len(x),\n            'target_length': len(target),\n            'sentence': sentence\n        }\n\n\ndef collate_fn(batch):\n    batch.sort(key=lambda x: x['length'], reverse=True)\n    neural = pad_sequence([b['neural'] for b in batch], batch_first=True)\n    target = pad_sequence([b['target'] for b in batch], batch_first=True)\n\n    return {\n        'neural': neural,\n        'target': target,\n        'lengths': torch.LongTensor([b['length'] for b in batch]),\n        'target_lengths': torch.LongTensor([b['target_length'] for b in batch]),\n        'sentences': [b['sentence'] for b in batch]\n    }\n\nclass LSTMOnlyCTC(nn.Module):\n    def __init__(self, input_dim=512, hidden_dim=256, num_layers=1, vocab_size=50, dropout=0.3):\n        super().__init__()\n        self.lstm = nn.LSTM(\n            input_dim,\n            hidden_dim,\n            num_layers=num_layers,\n            batch_first=True,\n            bidirectional=True,\n            dropout=dropout if num_layers > 1 else 0\n        )\n        self.fc = nn.Linear(hidden_dim * 2, vocab_size)\n\n    def forward(self, x, lengths):\n        x = pack_padded_sequence(x, lengths.cpu(), batch_first=True, enforce_sorted=False)\n        x, _ = self.lstm(x)\n        x, _ = pad_packed_sequence(x, batch_first=True)\n\n        logits = self.fc(x)\n        log_probs = torch.log_softmax(logits, dim=-1)\n        return log_probs.transpose(0, 1), lengths\n        \n# ============================================================================\n# 4. VALIDATION FUNCTION (computes both loss and WER)\n# ============================================================================\n\ndef validate_model(model, loader, idx2char, device):\n    from jiwer import wer\n    model.eval()\n    criterion = nn.CTCLoss(blank=0, zero_infinity=True)\n\n    total_loss = 0\n    preds, targets = [], []\n\n    with torch.no_grad():\n        for batch in loader:\n            neural = batch['neural'].to(device)\n            target = batch['target'].to(device)\n\n            log_probs, lengths = model(neural, batch['lengths'])\n            loss = criterion(log_probs, target, lengths, batch['target_lengths'])\n            total_loss += loss.item()\n\n            _, max_idx = log_probs.max(dim=-1)\n            for b in range(max_idx.size(1)):\n                seq = max_idx[:lengths[b], b].cpu().numpy()\n                decoded, prev = [], None\n                for t in seq:\n                    if t != 0 and t != prev:\n                        decoded.append(idx2char[t])\n                    prev = t\n                preds.append(\"\".join(decoded))\n\n            targets.extend(batch['sentences'])\n\n    return total_loss / len(loader), wer(\n        [t.lower() for t in targets],\n        [p.lower() for p in preds]\n    ) * 100\n\n# ============================================================================\n# 5. TRAINING LOOP WITH VAL LOSS AND TRAIN WER\n# ============================================================================\ndef train_model(train_loader, val_loader, char2idx, config):\n    model = LSTMOnlyCTC(\n        vocab_size=len(char2idx),\n        hidden_dim=config['hidden_dim'],\n        num_layers=config['num_layers'],\n        dropout=config['dropout']\n    ).to(config['device'])\n\n    optimizer = optim.Adam(model.parameters(), lr=config['learning_rate'])\n    criterion = nn.CTCLoss(blank=0, zero_infinity=True)\n\n    idx2char = {v: k for k, v in char2idx.items()}\n\n    train_losses, val_losses = [], []\n    train_wers, val_wers = [], []\n\n    best_wer = float('inf')\n\n    for epoch in range(config['num_epochs']):\n        model.train()\n        epoch_loss = 0\n        preds, targets = [], []\n\n        for batch in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{config['num_epochs']}\"):\n            neural = batch['neural'].to(config['device'])\n            target = batch['target'].to(config['device'])\n\n            log_probs, lengths = model(neural, batch['lengths'])\n            loss = criterion(log_probs, target, lengths, batch['target_lengths'])\n\n            optimizer.zero_grad()\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)\n            optimizer.step()\n\n            epoch_loss += loss.item()\n\n            _, max_idx = log_probs.max(dim=-1)\n            for b in range(max_idx.size(1)):\n                seq = max_idx[:lengths[b], b].cpu().numpy()\n                decoded, prev = [], None\n                for t in seq:\n                    if t != 0 and t != prev:\n                        decoded.append(idx2char[t])\n                    prev = t\n                preds.append(\"\".join(decoded))\n\n            targets.extend(batch['sentences'])\n\n        from jiwer import wer\n        train_losses.append(epoch_loss / len(train_loader))\n        train_wers.append(wer(\n            [t.lower() for t in targets],\n            [p.lower() for p in preds]\n        ) * 100)\n\n        val_loss, val_wer = validate_model(model, val_loader, idx2char, config['device'])\n        val_losses.append(val_loss)\n        val_wers.append(val_wer)\n\n        print(f\"Epoch {epoch+1}: Train WER={train_wers[-1]:.2f}%, Val WER={val_wer:.2f}%\")\n\n        if val_wer < best_wer:\n            best_wer = val_wer\n            #torch.save(model.state_dict(), \"lstm_only_best.pt\")\n\n            torch.save({\n                \"model_state_dict\": model.state_dict(),\n                \"char2idx\": char2idx\n            }, \"lstm_only_best.pt\")\n\n    # ===== Plots =====\n    epochs = np.arange(1, len(train_losses) + 1)\n\n    plt.figure(figsize=(8, 5))\n    plt.plot(epochs, train_losses, label=\"Train Loss\")\n    plt.plot(epochs, val_losses, label=\"Val Loss\")\n    plt.xlabel(\"Epoch\")\n    plt.ylabel(\"CTC Loss\")\n    plt.title(\"Loss\")\n    plt.legend()\n    plt.grid(True)\n    plt.show()\n\n    plt.figure(figsize=(8, 5))\n    plt.plot(epochs, train_wers, label=\"Train WER\")\n    plt.plot(epochs, val_wers, label=\"Val WER\")\n    plt.xlabel(\"Epoch\")\n    plt.ylabel(\"WER (%)\")\n    plt.title(\"WER\")\n    plt.legend()\n    plt.grid(True)\n    plt.show()\n\n    return model, best_wer\n\n# ============================================================================\n# 5. SUBMISSION GENERATION (Greedy Decoding)\n# ============================================================================\ndef main():\n    train_data = load_split(CONFIG['data_dir'], 'train')\n    val_data = load_split(CONFIG['data_dir'], 'val')\n\n    train_ds = BrainToTextDataset(train_data)\n    val_ds = BrainToTextDataset(val_data, char2idx=train_ds.char2idx)\n\n    train_loader = DataLoader(train_ds, batch_size=CONFIG['batch_size'],\n                              shuffle=True, collate_fn=collate_fn)\n    val_loader = DataLoader(val_ds, batch_size=CONFIG['batch_size'],\n                            shuffle=False, collate_fn=collate_fn)\n\n    model, best_wer = train_model(train_loader, val_loader, train_ds.char2idx, CONFIG)\n    print(f\"Best Val WER: {best_wer:.2f}%\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T14:47:45.77587Z","iopub.execute_input":"2025-12-13T14:47:45.776441Z","iopub.status.idle":"2025-12-13T14:56:24.406732Z","shell.execute_reply.started":"2025-12-13T14:47:45.776413Z","shell.execute_reply":"2025-12-13T14:56:24.405589Z"}},"outputs":[],"execution_count":null}]}