{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","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":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-23T16:00:42.723938Z","iopub.execute_input":"2025-12-23T16:00:42.724254Z","iopub.status.idle":"2025-12-23T16:00:43.072448Z","shell.execute_reply.started":"2025-12-23T16:00:42.72422Z","shell.execute_reply":"2025-12-23T16:00:43.070178Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install numpy==1.26.4 scipy==1.12.0 --force-reinstall","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-23T16:00:43.073629Z","iopub.execute_input":"2025-12-23T16:00:43.074231Z","iopub.status.idle":"2025-12-23T16:00:53.590586Z","shell.execute_reply.started":"2025-12-23T16:00:43.074187Z","shell.execute_reply":"2025-12-23T16:00:53.58959Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Install required packages\n!pip install jiwer pyctcdecode scikit-learn\n\nimport os\nimport h5py\nimport numpy as np\nfrom pathlib import Path\nfrom tqdm import tqdm\nfrom sklearn.decomposition import IncrementalPCA\n\nimport torch\nimport torch.nn as nn  \nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.nn.utils.rnn import pad_sequence\nfrom torch.utils.data import Dataset, DataLoader\nfrom jiwer import wer\n\n# Use GPU\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\n\n# Configuration\nCONFIG = {\n    'data_dir': '/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final',\n    'batch_size': 16,\n    'num_epochs': 30,\n    'input_dim': 128,\n    'hidden_dim': 512,\n    'ff_expansion': 4,\n    'num_heads': 8,\n    'num_layers': 3,\n    'conv_kernel_size': 15,\n    'dropout': 0.1,\n    'learning_rate': 3e-4,\n    'blank_index': 0,\n    'time_mask_width': 20,\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-23T16:00:53.592233Z","iopub.execute_input":"2025-12-23T16:00:53.592636Z","iopub.status.idle":"2025-12-23T16:01:01.393031Z","shell.execute_reply.started":"2025-12-23T16:00:53.592596Z","shell.execute_reply":"2025-12-23T16:01:01.392323Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_hdf5_files(file_list):\n    data = {'neural': [], 'n_steps': [], 'sent': []}\n    for filepath in file_list:\n        with h5py.File(filepath, 'r') as hf:\n            for trial_name in hf:\n                trial = hf[trial_name]\n                neural = np.array(trial['input_features'])  # ← fix here\n                n_steps = trial.attrs['n_time_steps']\n                sentence = trial.attrs.get('sentence_label')\n                if sentence is None:\n                    sentence = \"\"\n                elif isinstance(sentence, bytes):\n                    sentence = sentence.decode('utf-8')\n                data['neural'].append(neural)\n                data['n_steps'].append(n_steps)\n                data['sent'].append(sentence.lower())\n    print(f\"Loaded {len(data['neural'])} samples from {len(file_list)} files.\")\n    return data\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-23T16:01:01.394215Z","iopub.execute_input":"2025-12-23T16:01:01.394716Z","iopub.status.idle":"2025-12-23T16:01:01.400665Z","shell.execute_reply.started":"2025-12-23T16:01:01.394688Z","shell.execute_reply":"2025-12-23T16:01:01.399927Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import h5py\nfrom pathlib import Path\n\ndata_path = Path(CONFIG['data_dir'])\ntrain_files = sorted(data_path.glob(\"**/data_train.hdf5\"))\nval_files = sorted(data_path.glob(\"**/data_val.hdf5\"))\ntest_files = sorted(data_path.glob(\"**/data_test.hdf5\"))\n\n# Build sentence dictionary and vocabulary\ntrain_sentences_dict = {}\nall_texts = []\n\nprint(\"Scanning files for labels...\")\nfor filepath in train_files:\n    session_id = filepath.parent.name\n    with h5py.File(filepath, 'r') as hf:\n        for trial_id in hf.keys():\n            # Check if sentence exists in attributes\n            if 'sentence_label' in hf[trial_id].attrs:\n                sentence = hf[trial_id].attrs['sentence_label']\n            elif 'sentence' in hf[trial_id].attrs:\n                sentence = hf[trial_id].attrs['sentence']\n            else:\n                continue  # Skip if no sentence\n            \n            if isinstance(sentence, bytes):\n                sentence = sentence.decode('utf-8')\n            \n            sentence = sentence.lower().strip()\n            \n            # Create unique key\n            unique_key = f\"{session_id}_{trial_id}\"\n            train_sentences_dict[unique_key] = sentence\n            all_texts.append(sentence)\n\n# Create vocabulary\nunique_chars = sorted(list(set(\"\".join(all_texts))))\nchar2idx = {'<BLANK>': CONFIG['blank_index']}\nfor i, char in enumerate(unique_chars, start=1):\n    char2idx[char] = i\nidx2char = {v: k for k, v in char2idx.items()}\n\nprint(f\"✓ Mapped {len(train_sentences_dict)} sentences.\")\nprint(f\"✓ Vocab size: {len(char2idx)}\")\nprint(f\"✓ Sample characters: {list(char2idx.keys())[:10]}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-23T16:01:01.401654Z","iopub.execute_input":"2025-12-23T16:01:01.401897Z","iopub.status.idle":"2025-12-23T16:01:23.072608Z","shell.execute_reply.started":"2025-12-23T16:01:01.401874Z","shell.execute_reply":"2025-12-23T16:01:23.071916Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.decomposition import IncrementalPCA\n\nprint(\"\\nTraining PCA...\")\npca = IncrementalPCA(n_components=CONFIG['input_dim'], batch_size=1000)\n\nfor filepath in tqdm(train_files, desc=\"PCA Training\"):\n    with h5py.File(filepath, 'r') as hf:\n        for trial_name in hf.keys():\n            neural_data = hf[trial_name]['input_features'][()]\n            pca.partial_fit(neural_data)\n\nprint(\"✓ PCA training complete.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-23T16:01:23.074598Z","iopub.execute_input":"2025-12-23T16:01:23.074866Z","iopub.status.idle":"2025-12-23T16:26:06.933757Z","shell.execute_reply.started":"2025-12-23T16:01:23.074822Z","shell.execute_reply":"2025-12-23T16:26:06.932691Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BrainToTextDataset(Dataset):\n    def __init__(self, file_paths, pca, char2idx, sentences=None, normalize=True, train=True):\n        self.file_paths = file_paths\n        self.pca = pca\n        self.char2idx = char2idx\n        self.sentences = sentences  # Dictionary mapping session_trial -> sentence\n        self.normalize = normalize\n        self.train = train\n        self.data_keys = []\n        \n        # Scan files to get trial names\n        for filepath in self.file_paths:\n            with h5py.File(filepath, 'r') as hf:\n                session_name = filepath.parent.name\n                for trial_name in sorted(hf.keys()):\n                    self.data_keys.append((filepath, trial_name, session_name))\n\n    def __len__(self):\n        return len(self.data_keys)\n\n    def __getitem__(self, idx):\n        filepath, trial_name, session_name = self.data_keys[idx]\n        \n        with h5py.File(filepath, 'r') as hf:\n            x = hf[trial_name]['input_features'][()]  # Load neural data\n            \n        # 1. PCA transform\n        if self.pca is not None:\n            x = self.pca.transform(x)\n            \n        # 2. Normalization\n        if self.normalize:\n            mean = x.mean()\n            std = x.std() + 1e-8\n            x = (x - mean) / std\n            \n        x = torch.FloatTensor(x)\n        \n        # 3. Time Masking Augmentation (only during training)\n        if self.train and np.random.rand() < 0.5:\n            t = x.size(0)\n            mask_len = np.random.randint(0, CONFIG['time_mask_width'] + 1)\n            if t > mask_len:\n                start = np.random.randint(0, t - mask_len)\n                x[start:start+mask_len, :] = 0.0\n        \n        # 4. Target Sentence\n        full_id = f\"{session_name}_{trial_name}\"\n        sentence = self.sentences.get(full_id, \"\") if self.sentences else \"\"\n        target = [self.char2idx.get(ch, CONFIG['blank_index']) for ch in sentence]\n        \n        return {\n            'neural': x,\n            'target': torch.LongTensor(target),\n            'sentence': sentence\n        }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-23T16:26:06.935164Z","iopub.execute_input":"2025-12-23T16:26:06.935483Z","iopub.status.idle":"2025-12-23T16:26:06.958982Z","shell.execute_reply.started":"2025-12-23T16:26:06.935444Z","shell.execute_reply":"2025-12-23T16:26:06.957795Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def collate_fn(batch):\n    \"\"\"\n    Custom collate function for CTC training.\n    Handles variable-length sequences by padding.\n    \"\"\"\n    MAX_LENGTH = 800\n    \n    batch_inputs = []\n    input_lengths = []\n    batch_targets = []\n    target_lengths = []\n\n    for sample in batch:\n        neural = sample['neural']\n        \n        # Truncate if too long\n        if neural.shape[0] > MAX_LENGTH:\n            neural = neural[:MAX_LENGTH, :]\n        \n        batch_inputs.append(neural)\n        input_lengths.append(neural.shape[0])\n\n        target = sample['target']\n        batch_targets.append(target)\n        target_lengths.append(len(target))\n\n    # Pad inputs to same length: (B, T, D)\n    padded_inputs = nn.utils.rnn.pad_sequence(batch_inputs, batch_first=True)\n    \n    # Concatenate targets (required for CTCLoss)\n    flattened_targets = torch.cat(batch_targets)\n\n    return (\n        padded_inputs, \n        torch.tensor(input_lengths), \n        flattened_targets, \n        torch.tensor(target_lengths)\n    )\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-23T16:26:06.960343Z","iopub.execute_input":"2025-12-23T16:26:06.96071Z","iopub.status.idle":"2025-12-23T16:26:06.986807Z","shell.execute_reply.started":"2025-12-23T16:26:06.960675Z","shell.execute_reply":"2025-12-23T16:26:06.985699Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\nCreating datasets...\")\n\n# Training dataset\ntrain_ds = BrainToTextDataset(\n    file_paths=train_files, \n    pca=pca,                 \n    char2idx=char2idx, \n    sentences=train_sentences_dict, \n    normalize=True,\n    train=True\n)\n\n# Validation dataset (reuse training sentences)\nval_ds = BrainToTextDataset(\n    file_paths=val_files, \n    pca=pca, \n    char2idx=char2idx, \n    sentences=train_sentences_dict,  # Use same dict\n    normalize=True,\n    train=False\n)\n\n# Test dataset (no sentences)\ntest_ds = BrainToTextDataset(\n    file_paths=test_files, \n    pca=pca, \n    char2idx=char2idx, \n    sentences=None,  # No labels for test\n    normalize=True,\n    train=False\n)\n\n# Create data loaders\ntrain_loader = DataLoader(\n    train_ds, \n    batch_size=CONFIG['batch_size'], \n    shuffle=True, \n    collate_fn=collate_fn,\n    num_workers=2\n)\n\nval_loader = DataLoader(\n    val_ds, \n    batch_size=CONFIG['batch_size'], \n    shuffle=False, \n    collate_fn=collate_fn,\n    num_workers=2\n)\n\ntest_loader = DataLoader(\n    test_ds, \n    batch_size=CONFIG['batch_size'], \n    shuffle=False, \n    collate_fn=collate_fn,\n    num_workers=2\n)\n\nprint(f\"✓ Train: {len(train_ds)} samples\")\nprint(f\"✓ Val: {len(val_ds)} samples\")\nprint(f\"✓ Test: {len(test_ds)} samples\")\nprint(f\"✓ Vocab size: {len(char2idx)}\")\n\n# Test one batch\ntest_batch = next(iter(train_loader))\nprint(f\"\\n✓ Batch test successful!\")\nprint(f\"  - Input shape: {test_batch[0].shape}\")\nprint(f\"  - Input lengths: {test_batch[1][:5]}\")\nprint(f\"  - Target shape: {test_batch[2].shape}\")\nprint(f\"  - Target lengths: {test_batch[3][:5]}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-23T16:26:06.98828Z","iopub.execute_input":"2025-12-23T16:26:06.989755Z","iopub.status.idle":"2025-12-23T16:26:15.641641Z","shell.execute_reply.started":"2025-12-23T16:26:06.989721Z","shell.execute_reply":"2025-12-23T16:26:15.640658Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class FeedForward(nn.Module):\n    \"\"\"Position-wise Feed-Forward layer (2 linear layers with SiLU and dropout).\"\"\"\n    def __init__(self, dim, expansion_factor=4, dropout=0.1):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Linear(dim, dim * expansion_factor),\n            nn.SiLU(),\n            nn.Dropout(dropout),\n            nn.Linear(dim * expansion_factor, dim),\n            nn.Dropout(dropout)\n        )\n    def forward(self, x):\n        return self.net(x)\n\nclass ConvolutionModule(nn.Module):\n    \"\"\"Convolution Module in Conformer block (GLU, depthwise conv, etc.).\"\"\"\n    def __init__(self, dim, kernel_size=15, dropout=0.1):\n        super().__init__()\n        self.layer_norm = nn.LayerNorm(dim)\n        self.pointwise_conv1 = nn.Conv1d(dim, dim*2, kernel_size=1)\n        self.depthwise_conv = nn.Conv1d(dim, dim, kernel_size=kernel_size, groups=dim,\n                                        padding=(kernel_size-1)//2)\n        self.batch_norm = nn.BatchNorm1d(dim)\n        self.activation = nn.SiLU()\n        self.pointwise_conv2 = nn.Conv1d(dim, dim, kernel_size=1)\n        self.dropout = nn.Dropout(dropout)\n    def forward(self, x):\n        \"\"\"\n        x: (B, T, D)\n        returns: (B, T, D) after conv and residual\n        \"\"\"\n        res = x\n        # LayerNorm over feature dim\n        x = self.layer_norm(x)\n        # Conv1d expects (B, D, T)\n        x = x.transpose(1, 2)\n        # Pointwise conv + GLU\n        x = self.pointwise_conv1(x)  # (B, 2*D, T)\n        x = F.glu(x, dim=1)      # (B, D, T)\n        # Depthwise conv + BN + SiLU\n        x = self.depthwise_conv(x)\n        x = self.batch_norm(x)\n        x = self.activation(x)\n        # Pointwise conv\n        x = self.pointwise_conv2(x)\n        x = self.dropout(x)\n        x = x.transpose(1, 2)  # back to (B, T, D)\n        return res + x\n\nclass ConformerBlock(nn.Module):\n    \"\"\"One Conformer block (FF - MHA - Conv - FF, with residuals).\"\"\"\n    def __init__(self, dim, num_heads, ff_expansion, conv_kernel, dropout=0.1):\n        super().__init__()\n        self.ff1 = FeedForward(dim, expansion_factor=ff_expansion, dropout=dropout)\n        self.attn = nn.MultiheadAttention(dim, num_heads, dropout=dropout, batch_first=True)\n        self.conv = ConvolutionModule(dim, kernel_size=conv_kernel, dropout=dropout)\n        self.ff2 = FeedForward(dim, expansion_factor=ff_expansion, dropout=dropout)\n        # Layer norms for pre-norm, and final post-norm\n        self.norm_ff1 = nn.LayerNorm(dim)\n        self.norm_attn = nn.LayerNorm(dim)\n        self.norm_conv = nn.LayerNorm(dim)\n        self.norm_ff2 = nn.LayerNorm(dim)\n        self.norm_final = nn.LayerNorm(dim)\n    \n    def forward(self, x, attn_mask=None):\n        # x: (B, T, D)\n        # Feed-forward (1st half, Macaron)\n        res = x\n        x = self.norm_ff1(x)\n        x = res + 0.5 * self.ff1(x)\n        # Self-attention\n        res = x\n        x = self.norm_attn(x)\n        # MultiheadAttention with shape (B,T,D)\n        attn_out, _ = self.attn(x, x, x, attn_mask=attn_mask)\n        x = res + attn_out\n        # Convolution module\n        res = x\n        x = self.norm_conv(x)\n        x = res + self.conv(x)\n        # Feed-forward (2nd half)\n        res = x\n        x = self.norm_ff2(x)\n        x = res + 0.5 * self.ff2(x)\n        # Final LayerNorm\n        x = self.norm_final(x)\n        return x\n\nclass ConformerCTC(nn.Module):\n    \"\"\"Conformer Encoder for CTC (no decoder, outputs char probabilities).\"\"\"\n    def __init__(self, input_dim, hidden_dim, vocab_size, \n                 num_layers=3, num_heads=8, ff_expansion=4, \n                 conv_kernel=31, dropout=0.1):\n        super().__init__()\n        self.input_proj = nn.Linear(input_dim, hidden_dim)\n        self.pos_enc = nn.Sequential(\n            nn.Dropout(dropout),\n            # Positional encoding could be added here if needed (sinusoidal/learnable)\n        )\n        self.layers = nn.ModuleList([\n            ConformerBlock(hidden_dim, num_heads, ff_expansion, conv_kernel, dropout)\n            for _ in range(num_layers)\n        ])\n        self.fc_out = nn.Linear(hidden_dim, vocab_size)  # to character logits\n\n    def forward(self, x, lengths=None):\n        \"\"\"\n        x: (B, T, input_dim)\n        lengths: (B,) lengths of each sequence\n        Returns: log_probs (T, B, C), output_lengths (no change here)\n        \"\"\"\n        # Input projection\n        x = self.input_proj(x)  # (B, T, hidden_dim)\n        x = self.pos_enc(x)     # (B, T, hidden_dim)\n        \n        # NOTE: For padded input, one could create an attention mask here to ignore padding in attention.\n        attn_mask = None\n        \n        # Pass through Conformer layers\n        for layer in self.layers:\n            x = layer(x, attn_mask=attn_mask)\n        \n        # Final linear + log softmax\n        logits = self.fc_out(x)             # (B, T, C)\n        log_probs = nn.functional.log_softmax(logits, dim=-1)\n        # Permute to (T, B, C) for CTC loss\n        log_probs = log_probs.transpose(0, 1)\n        return log_probs, lengths\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-23T16:26:15.642917Z","iopub.execute_input":"2025-12-23T16:26:15.643194Z","iopub.status.idle":"2025-12-23T16:26:15.659215Z","shell.execute_reply.started":"2025-12-23T16:26:15.643153Z","shell.execute_reply":"2025-12-23T16:26:15.658464Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"vocab_size = len(char2idx)\nmodel = ConformerCTC(\n    input_dim=CONFIG['input_dim'], \n    hidden_dim=CONFIG['hidden_dim'],\n    vocab_size=vocab_size, \n    num_layers=CONFIG['num_layers'],\n    num_heads=CONFIG['num_heads'],\n    ff_expansion=CONFIG['ff_expansion'],\n    conv_kernel=CONFIG['conv_kernel_size'],\n    dropout=CONFIG['dropout']\n).to(device)\n\nctc_loss = nn.CTCLoss(blank=CONFIG['blank_index'], zero_infinity=True)\noptimizer = optim.AdamW(model.parameters(), lr=CONFIG['learning_rate'], weight_decay=1e-4)\nscheduler = optim.lr_scheduler.OneCycleLR(\n    optimizer, max_lr=CONFIG['learning_rate'],\n    steps_per_epoch=len(train_loader),\n    epochs=CONFIG['num_epochs']\n)\nscaler = torch.cuda.amp.GradScaler()\n\n# ============================================================================\n# Training Loop\n# ============================================================================\nbest_wer = float('inf')\n\nfor epoch in range(1, CONFIG['num_epochs']+1):\n    # ===== TRAINING =====\n    model.train()\n    total_loss = 0\n    \n    pbar = tqdm(train_loader, desc=f\"Epoch {epoch}/{CONFIG['num_epochs']}\")\n    for batch in pbar:\n        inputs, input_lengths, targets, target_lengths = batch\n        \n        inputs = inputs.to(device)\n        input_lengths = input_lengths.to(device)\n        targets = targets.to(device)\n        target_lengths = target_lengths.to(device)\n        \n        optimizer.zero_grad()\n        \n        with torch.cuda.amp.autocast():\n            log_probs, output_lengths = model(inputs, input_lengths)\n            loss = ctc_loss(log_probs, targets, output_lengths, target_lengths)\n        \n        scaler.scale(loss).backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        scaler.step(optimizer)\n        scaler.update()\n        scheduler.step()\n        \n        total_loss += loss.item()\n        pbar.set_postfix({'loss': f'{loss.item():.4f}'})\n\n    avg_loss = total_loss / len(train_loader)\n    print(f\"\\n✓ Epoch {epoch} Training Loss: {avg_loss:.4f}\")\n    \n    # ===== VALIDATION =====\n    model.eval()\n    preds = []\n    refs = []\n    \n    with torch.no_grad():\n        for batch_idx, batch in enumerate(tqdm(val_loader, desc=\"Validation\")):\n            inputs, input_lengths, targets, target_lengths = batch\n            inputs = inputs.to(device)\n            input_lengths = input_lengths.to(device)\n            \n            # Forward pass\n            log_probs, output_lengths = model(inputs, input_lengths)\n            \n            # Greedy Decode\n            probs = log_probs.cpu().transpose(0, 1)  # (B, T, C)\n            \n            for i, seq in enumerate(probs):\n                seq_idx = seq.argmax(dim=-1).tolist()\n                \n                # CTC collapse\n                decoded_sent = []\n                prev = CONFIG['blank_index']\n                for idx in seq_idx:\n                    if idx != CONFIG['blank_index'] and idx != prev:\n                        decoded_sent.append(idx2char.get(idx, ''))\n                    prev = idx\n                \n                pred_text = \"\".join(decoded_sent)\n                preds.append(pred_text)\n                \n                # Get reference from dataset\n                data_idx = batch_idx * CONFIG['batch_size'] + i\n                if data_idx < len(val_ds):\n                    ref_text = val_ds[data_idx]['sentence']\n                    refs.append(ref_text)\n    \n    # Calculate WER\n    if len(refs) > 0:\n        val_wer = wer(refs, preds[:len(refs)])\n        print(f\"✓ Epoch {epoch} Validation WER: {val_wer*100:.2f}%\")\n        \n        # Show sample predictions\n        print(\"\\nSample Predictions:\")\n        for i in range(min(3, len(refs))):\n            print(f\"  Ref: {refs[i]}\")\n            print(f\"  Pred: {preds[i]}\")\n            print()\n        \n        # Save best model\n        if val_wer < best_wer:\n            best_wer = val_wer\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'wer': val_wer,\n                'char2idx': char2idx,\n                'idx2char': idx2char,\n            }, \"best_model.pth\")\n            print(f\"✓ Saved New Best Model! (WER: {val_wer*100:.2f}%)\")\n    \n    print(\"-\" * 80)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-23T16:26:15.660271Z","iopub.execute_input":"2025-12-23T16:26:15.660492Z","iopub.status.idle":"2025-12-23T19:05:34.105767Z","shell.execute_reply.started":"2025-12-23T16:26:15.660471Z","shell.execute_reply":"2025-12-23T19:05:34.104855Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nfrom tqdm import tqdm\n\n# ============================================================================\n# STEP 1: Load Best Model\n# ============================================================================\nprint(\"Loading best model...\")\ncheckpoint = torch.load(\"best_model.pth\")\nmodel.load_state_dict(checkpoint['model_state_dict'])\nmodel.eval()\nprint(f\"✓ Loaded model with best WER: {checkpoint['wer']*100:.2f}%\")\n\n# ============================================================================\n# STEP 2: Generate Predictions on Test Set\n# ============================================================================\nprint(\"\\nGenerating predictions on test set...\")\ntest_predictions = []\ntest_ids = []\n\nwith torch.no_grad():\n    for batch_idx, batch in enumerate(tqdm(test_loader, desc=\"Test Inference\")):\n        # Unpack batch\n        inputs, input_lengths, _, _ = batch\n        inputs = inputs.to(device)\n        input_lengths = input_lengths.to(device)\n        \n        # Forward pass\n        log_probs, output_lengths = model(inputs, input_lengths)\n        \n        # Greedy Decoding\n        # Transpose to (B, T, C)\n        probs = log_probs.cpu().transpose(0, 1)\n        \n        # Decode each sequence in the batch\n        for i, seq in enumerate(probs):\n            # Get argmax indices\n            seq_idx = seq.argmax(dim=-1).tolist()\n            \n            # CTC Collapse: remove repeats and blanks\n            decoded_sent = []\n            prev = CONFIG['blank_index']\n            for idx in seq_idx:\n                if idx != CONFIG['blank_index'] and idx != prev:\n                    # Map index to character\n                    char = idx2char.get(idx, '')\n                    decoded_sent.append(char)\n                prev = idx\n            \n            # Join characters to form sentence\n            predicted_sentence = \"\".join(decoded_sent)\n            test_predictions.append(predicted_sentence)\n            \n            # Get the trial ID for this prediction\n            data_idx = batch_idx * CONFIG['batch_size'] + i\n            if data_idx < len(test_ds):\n                filepath, trial_name, session_name = test_ds.data_keys[data_idx]\n                # Create ID in format: session_trialname\n                test_id = f\"{session_name}_{trial_name}\"\n                test_ids.append(test_id)\n\nprint(f\"✓ Generated {len(test_predictions)} predictions.\")\nprint(f\"\\nSample Predictions:\")\nfor i in range(min(5, len(test_predictions))):\n    print(f\"  {test_ids[i]}: {test_predictions[i]}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-23T19:05:34.107245Z","iopub.execute_input":"2025-12-23T19:05:34.107901Z","iopub.status.idle":"2025-12-23T19:06:13.594851Z","shell.execute_reply.started":"2025-12-23T19:05:34.107866Z","shell.execute_reply":"2025-12-23T19:06:13.593795Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\nCreating submission file...\")\n\nsubmission_df = pd.DataFrame({\n    'ID': test_ids,\n    'sentence': test_predictions\n})\n\n# Check for duplicates\nif submission_df['ID'].duplicated().any():\n    print(\"⚠ Warning: Duplicate IDs found in submission!\")\n    print(f\"  Duplicate IDs: {submission_df[submission_df['ID'].duplicated()]['ID'].tolist()}\")\n\n# Display submission info\nprint(f\"\\nSubmission Info:\")\nprint(f\"  - Total predictions: {len(submission_df)}\")\nprint(f\"  - Unique IDs: {submission_df['ID'].nunique()}\")\nprint(f\"  - Empty predictions: {(submission_df['sentence'] == '').sum()}\")\nprint(f\"  - Average sentence length: {submission_df['sentence'].str.len().mean():.2f} chars\")\n\n# Show sample\nprint(f\"\\nSubmission Preview:\")\nprint(submission_df.head(10))\n\n# ============================================================================\n# STEP 4: Save Submission File\n# ============================================================================\nsubmission_df.to_csv('submission.csv', index=False)\nprint(f\"\\n✓ Submission saved to 'submission.csv'\")\n\n# ============================================================================\n# STEP 5: Validation Checks\n# ============================================================================\nprint(\"\\n\" + \"=\"*80)\nprint(\"SUBMISSION VALIDATION\")\nprint(\"=\"*80)\n\n# Check file format\nsubmission_check = pd.read_csv('submission.csv')\nprint(f\"✓ File readable\")\nprint(f\"✓ Shape: {submission_check.shape}\")\nprint(f\"✓ Columns: {list(submission_check.columns)}\")\n\n# Check for required columns\nrequired_cols = ['ID', 'sentence']\nif all(col in submission_check.columns for col in required_cols):\n    print(f\"✓ All required columns present\")\nelse:\n    print(f\"⚠ Missing columns: {set(required_cols) - set(submission_check.columns)}\")\n\n# Check for null values\nif submission_check.isnull().any().any():\n    print(f\"⚠ Warning: Null values found!\")\n    print(submission_check.isnull().sum())\nelse:\n    print(f\"✓ No null values\")\n\n# Statistics\nprint(f\"\\nPrediction Statistics:\")\nprint(f\"  - Min length: {submission_check['sentence'].str.len().min()}\")\nprint(f\"  - Max length: {submission_check['sentence'].str.len().max()}\")\nprint(f\"  - Median length: {submission_check['sentence'].str.len().median():.0f}\")\n\n# Character distribution\nall_chars = ''.join(submission_check['sentence'].fillna('').astype(str).tolist())\nunique_pred_chars = set(all_chars)\nprint(f\"  - Unique characters used: {len(unique_pred_chars)}\")\nprint(f\"  - Characters: {sorted(unique_pred_chars)[:30]}\")  # Show first 30\n\nprint(\"\\n\" + \"=\"*80)\nprint(\"READY TO SUBMIT! 🚀\")\nprint(\"=\"*80)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-23T19:06:13.596228Z","iopub.execute_input":"2025-12-23T19:06:13.59649Z","iopub.status.idle":"2025-12-23T19:06:13.649039Z","shell.execute_reply.started":"2025-12-23T19:06:13.596462Z","shell.execute_reply":"2025-12-23T19:06:13.64844Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}