{"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\n#Install jiwer for Word Error Rate computation\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\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nfrom tqdm import tqdm\nimport json\nimport warnings\nimport random\nimport matplotlib.pyplot as plt\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': 60,\n    'hidden_dim': 512, # LSTM hidden size\n    'num_layers': 3, # Number of BiLSTM layers\n    'dropout': 0.5, # Strong regularization\n    'learning_rate': 1e-3,\n}\n\nprint(f\"Device: {CONFIG['device']}\")\nprint(f\"PyTorch version: {torch.__version__}\")\n\n# ============================================================================\n# 2. DATA LOADING & DATASET\n# ============================================================================\ndef load_split(data_dir, split='train'):\n    \"\"\"\n    Load HDF5 files for a given split (train / val).\n    Returns neural features, sequence lengths, and sentence labels.\n    \"\"\"\n    from glob import glob\n    \n    pattern = f'{data_dir}/**/data_{split}.hdf5'\n    files = sorted(glob(pattern, recursive=True))\n    \n    print(f\"\\nLoading {split} split...\")\n    \n    all_data = {k: [] for k in ['neural', 'n_steps', 'sentence']}\n    \n    for filepath in tqdm(files):\n        with h5py.File(filepath, 'r') as f:\n            for trial_key in f.keys():\n                trial = f[trial_key]\n                \n                # Neural time-series features\n                neural = trial['input_features'][:]\n                n_steps = trial.attrs['n_time_steps']\n                \n                # Sentence label (string)\n                sentence = trial.attrs.get('sentence_label')\n                if sentence and 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    print(f\"✓ Loaded {len(all_data['neural'])} samples\")\n    return all_data\n\nclass BrainToTextDataset(Dataset):\n    \"\"\"\n    PyTorch Dataset for neural-to-text CTC training.\n    Handles normalization and character-level tokenization.\n    \"\"\"\n    def __init__(self, data, char2idx=None, normalize=True):\n        self.neural = data['neural']\n        self.n_steps = data['n_steps']\n        self.sentences = data['sentence']\n        self.normalize = normalize\n        \n        if char2idx is None:\n            self.char2idx = self._build_vocab()\n        else:\n            self.char2idx = char2idx\n        \n        self.idx2char = {v: k for k, v in self.char2idx.items()}\n        self.vocab_size = len(self.char2idx)\n    \n    def _build_vocab(self):\n        \"\"\"Create character-to-index mapping.\"\"\"\n        chars = set()\n        for sent in self.sentences:\n            if sent: chars.update(sent.lower())\n        chars = sorted(list(chars))\n        char2idx = {'<BLANK>': 0}\n        for i, ch in enumerate(chars, start=1):\n            char2idx[ch] = i\n        return char2idx\n    \n    def __len__(self): return len(self.neural)\n    \n    def __getitem__(self, idx):\n        # Trim sequence to valid time steps\n        neural = self.neural[idx][:self.n_steps[idx]]\n        \n        # Z-score normalization (per-sample)\n        if self.normalize:\n            neural = (neural - neural.mean()) / (neural.std() + 1e-8)\n        \n        sentence = self.sentences[idx] if self.sentences[idx] else \"\"\n        target = [self.char2idx.get(ch.lower(), 0) for ch in sentence]\n        \n        return {\n            'neural': torch.FloatTensor(neural),\n            'target': torch.LongTensor(target),\n            'length': len(neural),\n            'target_length': len(target),\n            'sentence': sentence\n        }\n\ndef collate_fn(batch):\n    \"\"\"\n    Collate function for variable-length batching.\n    Sorts sequences for efficient LSTM packing.\n    \"\"\"\n    batch = sorted(batch, key=lambda x: x['length'], reverse=True)\n    neurals = [item['neural'] for item in batch]\n    targets = [item['target'] for item in batch]\n    \n    neural_padded = pad_sequence(neurals, batch_first=True)\n    target_padded = pad_sequence(targets, batch_first=True)\n    lengths = torch.LongTensor([item['length'] for item in batch])\n    target_lengths = torch.LongTensor([item['target_length'] for item in batch])\n    \n    return {\n        'neural': neural_padded,\n        'target': target_padded,\n        'lengths': lengths,\n        'target_lengths': target_lengths,\n        'sentences': [item['sentence'] for item in batch]\n    }\n\nclass SimpleBrainCTCModel(nn.Module):\n    \"\"\"\n    Simple CNN + BiLSTM CTC model (NO SpecAugment)\n    \"\"\"\n    def __init__(\n        self,\n        input_dim=512,\n        hidden_dim=256,\n        num_layers=2,\n        vocab_size=50,\n        dropout=0.3\n    ):\n        super().__init__()\n\n        # Temporal downsampling CNN\n        self.cnn = nn.Sequential(\n            nn.Conv1d(input_dim, 256, kernel_size=11, stride=2, padding=5),\n            nn.BatchNorm1d(256),\n            nn.ReLU(),\n            nn.Dropout(dropout),\n\n            nn.Conv1d(256, 256, kernel_size=5, stride=2, padding=2),\n            nn.BatchNorm1d(256),\n            nn.ReLU(),\n            nn.Dropout(dropout),\n        )\n\n        # BiLSTM encoder\n        self.lstm = nn.LSTM(\n            input_size=256,\n            hidden_size=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\n        # Projection head\n        self.fc = nn.Linear(hidden_dim * 2, vocab_size)\n\n    def forward(self, x, lengths):\n        # x: (B, T, C)\n        x = x.transpose(1, 2)  # (B, C, T)\n\n        x = self.cnn(x)\n\n        x = x.transpose(1, 2)  # (B, T', F)\n\n        # adjust lengths (stride=2 twice)\n        cnn_lengths = (lengths // 4).clamp(min=1)\n\n        x = pack_padded_sequence(\n            x, cnn_lengths.cpu(), batch_first=True, enforce_sorted=False\n        )\n\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\n        return log_probs.transpose(0, 1), cnn_lengths\n\n\n# ============================================================================\n# 4. VALIDATION FUNCTION (computes both loss and WER)\n# ============================================================================\n\ndef validate_model(model, val_loader, idx2char, device):\n    from jiwer import wer\n    model.eval()\n    all_preds = []\n    all_targets = []\n    val_loss_total = 0\n    criterion = nn.CTCLoss(blank=0, zero_infinity=True)\n\n    with torch.no_grad():\n        for batch in val_loader:\n            neural = batch['neural'].to(device)\n            target = batch['target'].to(device)\n            lengths = batch['lengths']\n            target_lengths = batch['target_lengths']\n\n            log_probs, cnn_lengths = model(neural, lengths)\n\n            # Compute CTC loss for validation batch\n            loss = criterion(log_probs, target, cnn_lengths, target_lengths)\n            val_loss_total += loss.item()\n\n            _, max_indices = log_probs.max(dim=-1)\n            for b in range(max_indices.size(1)):\n                valid_len = cnn_lengths[b].item()\n                seq = max_indices[:valid_len, b].cpu().numpy()\n                decoded = []\n                prev = None\n                for token in seq:\n                    if token != 0 and token != prev:\n                        decoded.append(idx2char.get(token, ''))\n                    prev = token\n                all_preds.append(''.join(decoded))\n            all_targets.extend(batch['sentences'])\n\n    avg_val_loss = val_loss_total / len(val_loader)\n    val_wer = wer([t.lower() for t in all_targets], [p.lower() for p in all_preds]) * 100\n\n    return avg_val_loss, val_wer\n\n# ============================================================================\n# 5. TRAINING LOOP WITH VAL LOSS AND TRAIN WER\n# ============================================================================\ndef train_model(train_loader, val_loader, char2idx, config):\n    model = SimpleBrainCTCModel(\n        input_dim=512,\n        hidden_dim=256,\n        num_layers=2,\n        vocab_size=len(char2idx),\n        dropout=0.3\n    ).to(config['device'])\n\n    criterion = nn.CTCLoss(blank=0, zero_infinity=True)\n    optimizer = optim.Adam(model.parameters(), lr=1e-3)\n\n    idx2char = {v: k for k, v in char2idx.items()}\n\n    # === Tracking ===\n    train_losses = []\n    val_losses = []\n    train_wers = []\n    val_wers = []\n\n    best_wer = float('inf')\n\n    for epoch in range(config[\"num_epochs\"]):\n        model.train()\n        epoch_loss = 0.0\n        all_preds, all_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            lengths = batch['lengths']\n            target_lengths = batch['target_lengths']\n\n            log_probs, cnn_lengths = model(neural, lengths)\n            loss = criterion(log_probs, target, cnn_lengths, 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            # Greedy decoding (train WER)\n            _, max_indices = log_probs.max(dim=-1)\n            for b in range(max_indices.size(1)):\n                valid_len = cnn_lengths[b].item()\n                seq = max_indices[:valid_len, 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                all_preds.append(\"\".join(decoded))\n\n            all_targets.extend(batch['sentences'])\n\n        avg_train_loss = epoch_loss / len(train_loader)\n        train_losses.append(avg_train_loss)\n\n        from jiwer import wer\n        train_wer = wer(\n            [t.lower() for t in all_targets],\n            [p.lower() for p in all_preds]\n        ) * 100\n        train_wers.append(train_wer)\n\n        val_loss, val_wer = validate_model(\n            model, val_loader, idx2char, config['device']\n        )\n        val_losses.append(val_loss)\n        val_wers.append(val_wer)\n\n        print(\n            f\"Epoch {epoch+1}: \"\n            f\"Train Loss={avg_train_loss:.4f}, \"\n            f\"Train WER={train_wer:.2f}%, \"\n            f\"Val Loss={val_loss:.4f}, \"\n            f\"Val WER={val_wer:.2f}%\"\n        )\n\n        if val_wer < best_wer:\n            best_wer = val_wer\n            # torch.save(model.state_dict(), \"simple_best_model.pt\")\n            \n            torch.save({\n                \"model_state_dict\": model.state_dict(),\n                \"char2idx\": char2idx\n            }, \"best_model.pt\")\n            print(\"✓ Saved best model\")\n\n    # =======================\n    # Corrected Plots\n    # =======================\n    epochs = np.arange(1, config[\"num_epochs\"] + 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(\"Training vs Validation 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(\"Training vs Validation 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 load_test_data_for_submission(data_dir):\n    from glob import glob\n    print(\"\\nLoading test data...\")\n    pattern = f'{data_dir}/**/data_test.hdf5'\n    files = sorted(glob(pattern, recursive=True))\n    all_samples = []\n    sample_id = 0\n    for filepath in tqdm(files):\n        with h5py.File(filepath, 'r') as f:\n            trial_keys = [k for k in f.keys() if 'trial' in k.lower()]\n            for trial_key in trial_keys:\n                trial = f[trial_key]\n                if 'input_features' not in trial: continue\n                features = trial['input_features'][:]\n                n_steps = trial.attrs['n_time_steps']\n                features = features[:n_steps]\n                features = (features - features.mean(axis=0)) / (features.std(axis=0) + 1e-8)\n                all_samples.append({'id': sample_id, 'features': torch.FloatTensor(features)})\n                sample_id += 1\n    return all_samples\n\ndef generate_predictions(model, test_samples, idx2char, device, batch_size=64):\n    model.eval()\n    predictions = []\n    print(\"\\nGenerating predictions...\")\n    with torch.no_grad():\n        for i in tqdm(range(0, len(test_samples), batch_size)):\n            batch = test_samples[i:i+batch_size]\n            features = [s['features'] for s in batch]\n            lengths = torch.LongTensor([len(f) for f in features])\n            features_padded = pad_sequence(features, batch_first=True).to(device)\n            sorted_lengths, sorted_idx = lengths.sort(descending=True)\n            features_sorted = features_padded[sorted_idx]\n            \n            log_probs, cnn_lengths = model(features_sorted, sorted_lengths)\n            \n            _, max_indices = log_probs.max(dim=-1)\n            batch_preds = []\n            for b in range(max_indices.size(1)):\n                valid_len = cnn_lengths[b].item()\n                seq = max_indices[:valid_len, b].cpu().numpy()\n                decoded = []\n                prev = None\n                for token in seq:\n                    if token != 0 and token != prev:\n                        decoded.append(idx2char.get(token, ''))\n                    prev = token\n                batch_preds.append(''.join(decoded))\n            \n            unsorted_preds = [''] * len(batch_preds)\n            for i, pred in zip(sorted_idx.tolist(), batch_preds):\n                unsorted_preds[i] = pred\n            predictions.extend(unsorted_preds)\n    return predictions\n\n# ============================================================================\n# 6. MAIN PIPELINE\n# ============================================================================\ndef main():\n    print(\"=\"*80)\n    print(\"BRAIN-TO-TEXT '25 - PHASE 4 (TUNED LSTM + SPECAUGMENT)\")\n    print(\"=\"*80)\n    \n    # 1. Load Data\n    train_data = load_split(CONFIG['data_dir'], 'train')\n    val_data = load_split(CONFIG['data_dir'], 'val')\n    \n    # 2. Datasets\n    train_dataset = BrainToTextDataset(train_data)\n    val_dataset = BrainToTextDataset(val_data, char2idx=train_dataset.char2idx)\n    \n    train_loader = DataLoader(train_dataset, batch_size=CONFIG['batch_size'], \n                              shuffle=True, collate_fn=collate_fn, num_workers=2)\n    val_loader = DataLoader(val_dataset, batch_size=CONFIG['batch_size'],\n                            shuffle=False, collate_fn=collate_fn, num_workers=2)\n    \n    # 3. Train\n    model, best_wer = train_model(train_loader, val_loader, \n                                   train_dataset.char2idx, CONFIG)\n    \n    print(f\"\\n✓ Training complete! Best WER: {best_wer:.2f}%\")\n    \n    # 4. Predict\n    test_samples = load_test_data_for_submission(CONFIG['data_dir'])\n    checkpoint = torch.load('best_model.pt')\n    model.load_state_dict(checkpoint['model_state_dict'])\n    idx2char = {v: k for k, v in checkpoint['char2idx'].items()}\n    \n    predictions = generate_predictions(model, test_samples, idx2char, CONFIG['device'])\n    \n    # 5. Submit\n    df = pd.DataFrame({'id': [s['id'] for s in test_samples], 'text': predictions})\n    df = df.sort_values('id')\n    df.to_csv('submission.csv', index=False)\n    print(\"\\n✓ submission.csv ready!\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T14:47:45.775870Z","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}]}