{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"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 + SPECAUGMENT\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).\n2. Added SpecAugment (Time & Feature Masking) - SOTA regularization technique.\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, # Increased due to SpecAugment\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\n# ============================================================================\n# 3. PHASE 4 MODEL ARCHITECTURE (LSTM + SpecAugment)\n# ============================================================================\nclass SpecAugment(nn.Module):\n    \"\"\"\n    SpecAugment-style regularization:\n    - Time masking\n    - Feature (channel) masking\n    Applied only during training.\n    \"\"\"\n    def __init__(self, prob=0.5, time_mask_param=40, feature_mask_param=30):\n        super().__init__()\n        self.prob = prob\n        self.time_mask_param = time_mask_param\n        self.feature_mask_param = feature_mask_param\n\n    def forward(self, x):\n        # Input shape: (Batch, Channels, Time)\n        if not self.training or torch.rand(1) > self.prob:\n            return x\n            \n        b, c, t = x.size()\n        x_aug = x.clone()\n        \n        # Time masking\n        mask_len = torch.randint(0, self.time_mask_param, (1,)).item()\n        t0 = torch.randint(0, max(1, t - mask_len), (1,)).item()\n        x_aug[:, :, t0:t0 + mask_len] = 0\n        \n        # Feature masking\n        mask_feat = torch.randint(0, self.feature_mask_param, (1,)).item()\n        f0 = torch.randint(0, max(1, c - mask_feat), (1,)).item()\n        x_aug[:, f0:f0 + mask_feat, :] = 0\n        \n        return x_aug\n\nclass ImprovedBrainCTCModel(nn.Module):\n    \"\"\"\n    End-to-end neural-to-text model using CTC loss.\n    \"\"\"\n    def __init__(self, input_dim=512, hidden_dim=512, num_layers=3, \n                 vocab_size=50, dropout=0.5):\n        super().__init__()\n        \n        # SpecAugment regularization\n        self.spec_aug = SpecAugment(prob=0.6)\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        # Bidirectional LSTM encoder\n        self.lstm = nn.LSTM(\n            256, hidden_dim, num_layers,\n            batch_first=True, bidirectional=True,\n            dropout=dropout if num_layers > 1 else 0\n        )\n\n        # Projection to vocabulary logits\n        self.fc = nn.Linear(hidden_dim * 2, vocab_size)\n    \n    def forward(self, x, lengths):\n        # Input: (Batch, Time, Channels)\n        \n        # Convert to (Batch, Channels, Time) for CNN\n        x = x.transpose(1, 2) \n        \n        # Apply SpecAugment during training\n        if self.training:\n            x = self.spec_aug(x)\n            \n        # 3. CNN feature extraction\n        x = self.cnn(x)\n        \n        # Convert back to (Batch, Time, Features)\n        x = x.transpose(1, 2)\n        \n        # Adjust sequence lengths after striding\n        cnn_lengths = (lengths.cpu() // 4).clamp(min=1)\n        \n        # Pack sequences for LSTM\n        x_packed = pack_padded_sequence(x, cnn_lengths, batch_first=True, enforce_sorted=False)\n        lstm_out, _ = self.lstm(x_packed)\n        lstm_out, _ = pad_packed_sequence(lstm_out, batch_first=True)\n        \n        # Compute log-probabilities for CTC\n        logits = self.fc(lstm_out)\n        log_probs = torch.log_softmax(logits, dim=-1)\n        \n        # Return in (Time, Batch, Vocab) format\n        return log_probs.transpose(0, 1), cnn_lengths\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# ============================================================================\n\ndef train_model(train_loader, val_loader, char2idx, config):\n    model = ImprovedBrainCTCModel(\n        input_dim=512,\n        hidden_dim=config['hidden_dim'],\n        num_layers=config['num_layers'],\n        vocab_size=len(char2idx),\n        dropout=config['dropout']\n    ).to(config['device'])\n\n    criterion = nn.CTCLoss(blank=0, zero_infinity=True)\n    optimizer = optim.AdamW(model.parameters(),\n                            lr=config['learning_rate'],\n                            weight_decay=1e-4)\n\n    scheduler = optim.lr_scheduler.OneCycleLR(\n        optimizer,\n        max_lr=config['learning_rate'],\n        epochs=config['num_epochs'],\n        steps_per_epoch=len(train_loader)\n    )\n\n    idx2char = {v: k for k, v in char2idx.items()}\n\n    train_losses = []\n    train_wers = []\n    val_losses = []\n    val_wers = []\n\n    best_wer = float('inf')\n\n    print(\"\\nStarting training (LSTM + SpecAugment)...\")\n\n    for epoch in range(config['num_epochs']):\n        model.train()\n        epoch_loss = 0\n        all_preds = []\n        all_targets = []\n\n        for batch in tqdm(train_loader,\n                          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\n            loss = criterion(log_probs, target,\n                             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            scheduler.step()\n\n            epoch_loss += loss.item()\n\n            # Collect predictions for train WER\n            _, max_indices = log_probs.max(dim=-1)\n            for b_idx in range(max_indices.size(1)):\n                valid_len = cnn_lengths[b_idx].item()\n                seq = max_indices[:valid_len, b_idx].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_train_loss = epoch_loss / len(train_loader)\n        train_losses.append(avg_train_loss)\n\n        # Compute train WER for the epoch\n        from jiwer import wer\n        train_wer = wer([t.lower() for t in all_targets], [p.lower() for p in all_preds]) * 100\n        train_wers.append(train_wer)\n\n        # Validate every 5 epochs\n        if (epoch + 1) % 1 == 0:\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 Loss={avg_train_loss:.4f}, Train WER={train_wer:.2f}%, Val Loss={val_loss:.4f}, Val WER={val_wer:.2f}%\")\n\n            if val_wer < best_wer:\n                best_wer = val_wer\n                torch.save({\n                    'model_state_dict': model.state_dict(),\n                    'char2idx': char2idx,\n                    'config': config,\n                    'wer': val_wer\n                }, 'best_model.pt')\n                print(f\"✓ Saved (WER: {val_wer:.2f}%)\")\n\n    # =======================\n    # Corrected Plots\n    # =======================\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('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\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},"outputs":[],"execution_count":null}]}