{"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,"isSourceIdPinned":false,"sourceType":"competition"}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Brain-to-Text 2.5 Competition - OPTIMIZED VERSION\n# Key Improvements:\n# 1. Multi-session training (uses ALL available data folders)\n# 2. CTC Loss for better sequence alignment\n# 3. Conformer-style architecture (Conv + Attention)\n# 4. Bidirectional processing\n# 5. Data augmentation (time warping, noise injection)\n# 6. Proper train/val split\n# 7. Label smoothing + mixup\n# 8. Learning rate warmup with cosine decay\n# 9. Test-time augmentation (TTA)\n# 10. Gradient accumulation for larger effective batch size\n\nimport h5py\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, ConcatDataset\nfrom torch.nn.utils.rnn import pad_sequence\nimport numpy as np\nfrom tqdm import tqdm\nimport os\nimport math\nimport pandas as pd\nimport glob\nimport random\nfrom collections import defaultdict\n\n# Set device and seeds\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n\nset_seed(42)\n\n# ============================================================================\n# CONFIGURATION - Optimized hyperparameters\n# ============================================================================\nclass Config:\n    # Model Architecture\n    input_size = 512\n    d_model = 512\n    nhead = 8\n    num_encoder_layers = 6  # Increased depth\n    num_decoder_layers = 2  # Add decoder for better seq2seq\n    dim_feedforward = 2048  # Larger FFN\n    dropout = 0.15\n    conv_kernel_size = 31  # For conformer-style conv\n    \n    # Training\n    batch_size = 16\n    accumulation_steps = 2  # Effective batch = 32\n    learning_rate = 3e-4\n    warmup_steps = 1000\n    num_epochs = 50\n    max_seq_len = 500\n    grad_clip = 1.0\n    weight_decay = 0.01\n    label_smoothing = 0.1\n    \n    # Data\n    vocab_size = 500\n    use_ctc = True\n    ctc_weight = 0.3  # Hybrid CTC + CE loss\n    \n    # Augmentation\n    use_augmentation = True\n    time_mask_prob = 0.3\n    time_mask_len = 20\n    noise_std = 0.1\n    mixup_alpha = 0.2\n    \n    # TTA\n    use_tta = True\n    tta_samples = 5\n\nconfig = Config()\n\n# ============================================================================\n# DATA PATHS - Use ALL available sessions\n# ============================================================================\nBASE_PATH = \"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final\"\n\ndef get_all_data_files():\n    \"\"\"Find all training and test files across all sessions\"\"\"\n    train_files = []\n    test_files = []\n    \n    # Get all session folders\n    session_folders = glob.glob(os.path.join(BASE_PATH, \"t15.*\"))\n    print(f\"Found {len(session_folders)} session folders\")\n    \n    for folder in sorted(session_folders):\n        train_path = os.path.join(folder, \"data_train.hdf5\")\n        test_path = os.path.join(folder, \"data_test.hdf5\")\n        \n        if os.path.exists(train_path):\n            train_files.append(train_path)\n            print(f\"  Train: {train_path}\")\n        if os.path.exists(test_path):\n            test_files.append(test_path)\n            print(f\"  Test: {test_path}\")\n    \n    return train_files, test_files\n\n# ============================================================================\n# DATA AUGMENTATION\n# ============================================================================\nclass DataAugmentation:\n    def __init__(self, config):\n        self.config = config\n    \n    def time_mask(self, x):\n        \"\"\"Mask random time steps (SpecAugment-style)\"\"\"\n        if random.random() > self.config.time_mask_prob:\n            return x\n        \n        seq_len = x.shape[0]\n        mask_len = min(self.config.time_mask_len, seq_len // 4)\n        start = random.randint(0, seq_len - mask_len)\n        x[start:start + mask_len] = 0\n        return x\n    \n    def add_noise(self, x):\n        \"\"\"Add Gaussian noise\"\"\"\n        noise = torch.randn_like(x) * self.config.noise_std\n        return x + noise\n    \n    def time_warp(self, x, max_warp=5):\n        \"\"\"Simple time warping via interpolation\"\"\"\n        if random.random() > 0.3:\n            return x\n        \n        seq_len = x.shape[0]\n        if seq_len < 10:\n            return x\n        \n        # Simple stretch/compress\n        factor = 1.0 + random.uniform(-0.1, 0.1)\n        new_len = int(seq_len * factor)\n        new_len = max(10, min(new_len, seq_len + 20))\n        \n        x_interp = F.interpolate(\n            x.unsqueeze(0).transpose(1, 2),\n            size=new_len,\n            mode='linear',\n            align_corners=False\n        ).transpose(1, 2).squeeze(0)\n        \n        return x_interp\n    \n    def __call__(self, x, training=True):\n        if not training or not self.config.use_augmentation:\n            return x\n        \n        x = self.time_mask(x.clone())\n        x = self.add_noise(x)\n        # x = self.time_warp(x)  # Can be slow, enable if needed\n        return x\n\n# ============================================================================\n# DATASET - Multi-session support\n# ============================================================================\nclass BrainDataset(Dataset):\n    def __init__(self, hdf5_files, is_test=False, max_len=500, augmentation=None):\n        \"\"\"\n        Args:\n            hdf5_files: List of HDF5 file paths or single path\n        \"\"\"\n        if isinstance(hdf5_files, str):\n            hdf5_files = [hdf5_files]\n        \n        self.files = hdf5_files\n        self.is_test = is_test\n        self.max_len = max_len\n        self.augmentation = augmentation\n        \n        # Build index mapping (file_idx, trial_key)\n        self.trial_index = []\n        for file_idx, file_path in enumerate(self.files):\n            with h5py.File(file_path, 'r') as f:\n                trials = list(f.keys())\n                for trial in trials:\n                    self.trial_index.append((file_idx, trial))\n        \n        print(f\"Total trials across {len(self.files)} files: {len(self.trial_index)}\")\n    \n    def __len__(self):\n        return len(self.trial_index)\n    \n    def __getitem__(self, idx):\n        file_idx, trial_key = self.trial_index[idx]\n        \n        with h5py.File(self.files[file_idx], 'r') as f:\n            trial = f[trial_key]\n            neural_data = torch.tensor(trial['input_features'][:], dtype=torch.float32)\n            \n            if not self.is_test:\n                labels = torch.tensor(trial['seq_class_ids'][:], dtype=torch.long)\n                labels = torch.clamp(labels, 0, config.vocab_size - 1)\n                \n                # Truncate/pad labels\n                if len(labels) > self.max_len:\n                    labels = labels[:self.max_len]\n                label_len = len(labels)\n            else:\n                labels = torch.zeros(self.max_len, dtype=torch.long)\n                label_len = 0\n        \n        # Apply augmentation\n        if self.augmentation is not None and not self.is_test:\n            neural_data = self.augmentation(neural_data, training=True)\n        \n        return neural_data, labels, label_len, trial_key\n\n# ============================================================================\n# CONFORMER-STYLE ARCHITECTURE\n# ============================================================================\nclass ConvModule(nn.Module):\n    \"\"\"Conformer-style convolution module\"\"\"\n    def __init__(self, d_model, kernel_size=31, dropout=0.1):\n        super().__init__()\n        self.layer_norm = nn.LayerNorm(d_model)\n        self.pointwise_conv1 = nn.Linear(d_model, 2 * d_model)\n        self.glu = nn.GLU(dim=-1)\n        self.depthwise_conv = nn.Conv1d(\n            d_model, d_model, kernel_size,\n            padding=kernel_size // 2, groups=d_model\n        )\n        self.batch_norm = nn.BatchNorm1d(d_model)\n        self.swish = nn.SiLU()\n        self.pointwise_conv2 = nn.Linear(d_model, d_model)\n        self.dropout = nn.Dropout(dropout)\n    \n    def forward(self, x):\n        # x: (batch, seq, d_model)\n        residual = x\n        x = self.layer_norm(x)\n        x = self.pointwise_conv1(x)\n        x = self.glu(x)\n        \n        # Depthwise conv expects (batch, channels, seq)\n        x = x.transpose(1, 2)\n        x = self.depthwise_conv(x)\n        x = self.batch_norm(x)\n        x = x.transpose(1, 2)\n        \n        x = self.swish(x)\n        x = self.pointwise_conv2(x)\n        x = self.dropout(x)\n        \n        return residual + x\n\nclass FeedForwardModule(nn.Module):\n    \"\"\"Feed-forward module with pre-norm\"\"\"\n    def __init__(self, d_model, dim_feedforward, dropout=0.1):\n        super().__init__()\n        self.layer_norm = nn.LayerNorm(d_model)\n        self.linear1 = nn.Linear(d_model, dim_feedforward)\n        self.swish = nn.SiLU()\n        self.dropout1 = nn.Dropout(dropout)\n        self.linear2 = nn.Linear(dim_feedforward, d_model)\n        self.dropout2 = nn.Dropout(dropout)\n    \n    def forward(self, x):\n        residual = x\n        x = self.layer_norm(x)\n        x = self.linear1(x)\n        x = self.swish(x)\n        x = self.dropout1(x)\n        x = self.linear2(x)\n        x = self.dropout2(x)\n        return residual + 0.5 * x\n\nclass ConformerBlock(nn.Module):\n    \"\"\"Single Conformer block\"\"\"\n    def __init__(self, d_model, nhead, dim_feedforward, conv_kernel_size, dropout):\n        super().__init__()\n        self.ff1 = FeedForwardModule(d_model, dim_feedforward, dropout)\n        \n        self.self_attn_layer_norm = nn.LayerNorm(d_model)\n        self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout, batch_first=True)\n        self.self_attn_dropout = nn.Dropout(dropout)\n        \n        self.conv = ConvModule(d_model, conv_kernel_size, dropout)\n        self.ff2 = FeedForwardModule(d_model, dim_feedforward, dropout)\n        self.final_layer_norm = nn.LayerNorm(d_model)\n    \n    def forward(self, x, src_key_padding_mask=None):\n        # First FFN\n        x = self.ff1(x)\n        \n        # Multi-head attention\n        residual = x\n        x = self.self_attn_layer_norm(x)\n        x, _ = self.self_attn(x, x, x, key_padding_mask=src_key_padding_mask)\n        x = self.self_attn_dropout(x)\n        x = residual + x\n        \n        # Convolution module\n        x = self.conv(x)\n        \n        # Second FFN\n        x = self.ff2(x)\n        \n        # Final layer norm\n        x = self.final_layer_norm(x)\n        \n        return x\n\nclass PositionalEncoding(nn.Module):\n    def __init__(self, d_model, dropout=0.1, max_len=5000):\n        super().__init__()\n        self.dropout = nn.Dropout(p=dropout)\n        \n        pe = torch.zeros(max_len, d_model)\n        position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-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)\n    \n    def forward(self, x):\n        x = x + self.pe[:x.size(1)]\n        return self.dropout(x)\n\nclass BrainConformer(nn.Module):\n    \"\"\"Conformer-based Brain-to-Text model\"\"\"\n    def __init__(self, config):\n        super().__init__()\n        self.config = config\n        \n        # Input projection with conv subsampling\n        self.input_conv = nn.Sequential(\n            nn.Conv1d(config.input_size, config.d_model, kernel_size=3, padding=1),\n            nn.BatchNorm1d(config.d_model),\n            nn.SiLU(),\n            nn.Dropout(config.dropout),\n        )\n        \n        self.pos_encoding = PositionalEncoding(config.d_model, config.dropout)\n        \n        # Conformer encoder blocks\n        self.encoder_blocks = nn.ModuleList([\n            ConformerBlock(\n                config.d_model,\n                config.nhead,\n                config.dim_feedforward,\n                config.conv_kernel_size,\n                config.dropout\n            )\n            for _ in range(config.num_encoder_layers)\n        ])\n        \n        # Output projection\n        self.output_proj = nn.Linear(config.d_model, config.vocab_size)\n        \n        # CTC output (if using hybrid loss)\n        if config.use_ctc:\n            self.ctc_proj = nn.Linear(config.d_model, config.vocab_size + 1)  # +1 for blank\n        \n        self._init_weights()\n    \n    def _init_weights(self):\n        for p in self.parameters():\n            if p.dim() > 1:\n                nn.init.xavier_uniform_(p)\n    \n    def forward(self, x, src_key_padding_mask=None):\n        # x: (batch, seq, features)\n        \n        # Conv input projection\n        x = x.transpose(1, 2)  # (batch, features, seq)\n        x = self.input_conv(x)\n        x = x.transpose(1, 2)  # (batch, seq, d_model)\n        \n        # Add positional encoding\n        x = self.pos_encoding(x)\n        \n        # Pass through conformer blocks\n        for block in self.encoder_blocks:\n            x = block(x, src_key_padding_mask)\n        \n        # Output projections\n        logits = self.output_proj(x)\n        \n        if self.config.use_ctc:\n            ctc_logits = self.ctc_proj(x)\n            return logits, ctc_logits\n        \n        return logits, None\n\n# ============================================================================\n# SIMPLER BUT EFFECTIVE: BiLSTM + Transformer Hybrid\n# ============================================================================\nclass BrainHybridModel(nn.Module):\n    \"\"\"BiLSTM + Transformer hybrid - often works better than pure Transformer\"\"\"\n    def __init__(self, config):\n        super().__init__()\n        self.config = config\n        \n        # Input projection\n        self.input_proj = nn.Sequential(\n            nn.Linear(config.input_size, config.d_model),\n            nn.LayerNorm(config.d_model),\n            nn.Dropout(config.dropout),\n        )\n        \n        # Bidirectional LSTM for local context\n        self.bilstm = nn.LSTM(\n            config.d_model,\n            config.d_model // 2,\n            num_layers=2,\n            batch_first=True,\n            bidirectional=True,\n            dropout=config.dropout\n        )\n        \n        # Transformer for global context\n        self.pos_encoding = PositionalEncoding(config.d_model, config.dropout)\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=config.d_model,\n            nhead=config.nhead,\n            dim_feedforward=config.dim_feedforward,\n            dropout=config.dropout,\n            batch_first=True,\n            activation='gelu'\n        )\n        self.transformer = nn.TransformerEncoder(encoder_layer, config.num_encoder_layers)\n        \n        # Output\n        self.output_norm = nn.LayerNorm(config.d_model)\n        self.output_proj = nn.Linear(config.d_model, config.vocab_size)\n        \n        if config.use_ctc:\n            self.ctc_proj = nn.Linear(config.d_model, config.vocab_size + 1)\n        \n        self._init_weights()\n    \n    def _init_weights(self):\n        for name, p in self.named_parameters():\n            if 'weight' in name and p.dim() > 1:\n                nn.init.xavier_uniform_(p)\n    \n    def forward(self, x, src_key_padding_mask=None):\n        # Input projection\n        x = self.input_proj(x)\n        \n        # BiLSTM\n        x, _ = self.bilstm(x)\n        \n        # Transformer\n        x = self.pos_encoding(x)\n        x = self.transformer(x, src_key_padding_mask=src_key_padding_mask)\n        \n        # Output\n        x = self.output_norm(x)\n        logits = self.output_proj(x)\n        \n        ctc_logits = None\n        if self.config.use_ctc:\n            ctc_logits = self.ctc_proj(x)\n        \n        return logits, ctc_logits\n\n# ============================================================================\n# LOSS FUNCTIONS\n# ============================================================================\nclass HybridLoss(nn.Module):\n    \"\"\"Hybrid CTC + Cross-Entropy loss\"\"\"\n    def __init__(self, vocab_size, ctc_weight=0.3, label_smoothing=0.1, blank_id=0):\n        super().__init__()\n        self.ctc_weight = ctc_weight\n        self.ce_loss = nn.CrossEntropyLoss(\n            ignore_index=0,\n            label_smoothing=label_smoothing\n        )\n        self.ctc_loss = nn.CTCLoss(blank=blank_id, zero_infinity=True)\n    \n    def forward(self, logits, ctc_logits, labels, input_lengths, label_lengths):\n        # Cross-entropy loss\n        batch_size, seq_len, vocab_size = logits.shape\n        label_len = labels.shape[1]\n        \n        if seq_len >= label_len:\n            logits_aligned = logits[:, :label_len, :]\n        else:\n            logits_aligned = F.pad(logits, (0, 0, 0, label_len - seq_len), value=0)\n        \n        ce_loss = self.ce_loss(\n            logits_aligned.reshape(-1, vocab_size),\n            labels.reshape(-1)\n        )\n        \n        # CTC loss (optional)\n        if ctc_logits is not None and self.ctc_weight > 0:\n            # CTC expects (seq, batch, vocab)\n            ctc_logits = ctc_logits.transpose(0, 1)\n            ctc_log_probs = F.log_softmax(ctc_logits, dim=-1)\n            \n            ctc_loss = self.ctc_loss(\n                ctc_log_probs,\n                labels,\n                input_lengths,\n                label_lengths\n            )\n            \n            total_loss = (1 - self.ctc_weight) * ce_loss + self.ctc_weight * ctc_loss\n        else:\n            total_loss = ce_loss\n        \n        return total_loss\n\n# ============================================================================\n# COLLATE FUNCTION\n# ============================================================================\ndef collate_fn(batch):\n    neural_data = [item[0] for item in batch]\n    labels = [item[1] for item in batch]\n    label_lengths = torch.tensor([item[2] for item in batch], dtype=torch.long)\n    trial_keys = [item[3] for item in batch]\n    \n    # Pad neural data\n    neural_padded = pad_sequence(neural_data, batch_first=True)\n    input_lengths = torch.tensor([len(x) for x in neural_data], dtype=torch.long)\n    \n    # Create padding mask\n    max_len = neural_padded.shape[1]\n    padding_mask = torch.arange(max_len).expand(len(input_lengths), max_len) >= input_lengths.unsqueeze(1)\n    \n    # Pad labels to same length\n    max_label_len = max(len(l) for l in labels)\n    labels_padded = torch.zeros(len(labels), max_label_len, dtype=torch.long)\n    for i, l in enumerate(labels):\n        labels_padded[i, :len(l)] = l\n    \n    return neural_padded, labels_padded, padding_mask, input_lengths, label_lengths, trial_keys\n\n# ============================================================================\n# LEARNING RATE SCHEDULER WITH WARMUP\n# ============================================================================\nclass WarmupCosineScheduler:\n    def __init__(self, optimizer, warmup_steps, total_steps, min_lr=1e-6):\n        self.optimizer = optimizer\n        self.warmup_steps = warmup_steps\n        self.total_steps = total_steps\n        self.min_lr = min_lr\n        self.base_lr = optimizer.param_groups[0]['lr']\n        self.current_step = 0\n    \n    def step(self):\n        self.current_step += 1\n        lr = self.get_lr()\n        for param_group in self.optimizer.param_groups:\n            param_group['lr'] = lr\n    \n    def get_lr(self):\n        if self.current_step < self.warmup_steps:\n            return self.base_lr * self.current_step / self.warmup_steps\n        else:\n            progress = (self.current_step - self.warmup_steps) / (self.total_steps - self.warmup_steps)\n            return self.min_lr + (self.base_lr - self.min_lr) * 0.5 * (1 + math.cos(math.pi * progress))\n\n# ============================================================================\n# TRAINING FUNCTIONS\n# ============================================================================\ndef train_epoch(model, dataloader, criterion, optimizer, scheduler, device, config):\n    model.train()\n    total_loss = 0\n    total_tokens = 0\n    \n    optimizer.zero_grad()\n    \n    pbar = tqdm(dataloader, desc=\"Training\")\n    for batch_idx, (neural_data, labels, padding_mask, input_lengths, label_lengths, _) in enumerate(pbar):\n        neural_data = neural_data.to(device)\n        labels = labels.to(device)\n        padding_mask = padding_mask.to(device)\n        input_lengths = input_lengths.to(device)\n        label_lengths = label_lengths.to(device)\n        \n        logits, ctc_logits = model(neural_data, src_key_padding_mask=padding_mask)\n        \n        loss = criterion(logits, ctc_logits, labels, input_lengths, label_lengths)\n        loss = loss / config.accumulation_steps\n        loss.backward()\n        \n        if (batch_idx + 1) % config.accumulation_steps == 0:\n            torch.nn.utils.clip_grad_norm_(model.parameters(), config.grad_clip)\n            optimizer.step()\n            scheduler.step()\n            optimizer.zero_grad()\n        \n        batch_tokens = (labels != 0).sum().item()\n        total_loss += loss.item() * config.accumulation_steps * batch_tokens\n        total_tokens += batch_tokens\n        \n        pbar.set_postfix({\n            'loss': f'{loss.item() * config.accumulation_steps:.4f}',\n            'lr': f'{scheduler.get_lr():.2e}'\n        })\n    \n    return total_loss / total_tokens if total_tokens > 0 else 0\n\ndef validate(model, dataloader, criterion, device):\n    model.eval()\n    total_loss = 0\n    total_tokens = 0\n    correct = 0\n    \n    with torch.no_grad():\n        for neural_data, labels, padding_mask, input_lengths, label_lengths, _ in dataloader:\n            neural_data = neural_data.to(device)\n            labels = labels.to(device)\n            padding_mask = padding_mask.to(device)\n            input_lengths = input_lengths.to(device)\n            label_lengths = label_lengths.to(device)\n            \n            logits, ctc_logits = model(neural_data, src_key_padding_mask=padding_mask)\n            \n            loss = criterion(logits, ctc_logits, labels, input_lengths, label_lengths)\n            \n            # Accuracy\n            batch_size, seq_len, vocab_size = logits.shape\n            label_len = labels.shape[1]\n            \n            if seq_len >= label_len:\n                logits_aligned = logits[:, :label_len, :]\n            else:\n                logits_aligned = F.pad(logits, (0, 0, 0, label_len - seq_len), value=0)\n            \n            preds = torch.argmax(logits_aligned, dim=-1)\n            mask = labels != 0\n            correct += ((preds == labels) & mask).sum().item()\n            \n            batch_tokens = mask.sum().item()\n            total_loss += loss.item() * batch_tokens\n            total_tokens += batch_tokens\n    \n    return total_loss / total_tokens, correct / total_tokens if total_tokens > 0 else 0\n\n# ============================================================================\n# TEST-TIME AUGMENTATION\n# ============================================================================\ndef predict_with_tta(model, neural_data, padding_mask, config, device):\n    \"\"\"Predict with test-time augmentation\"\"\"\n    model.eval()\n    \n    all_logits = []\n    \n    with torch.no_grad():\n        # Original prediction\n        logits, _ = model(neural_data, src_key_padding_mask=padding_mask)\n        all_logits.append(logits)\n        \n        if config.use_tta:\n            for _ in range(config.tta_samples - 1):\n                # Add small noise\n                noisy_data = neural_data + torch.randn_like(neural_data) * 0.05\n                logits, _ = model(noisy_data, src_key_padding_mask=padding_mask)\n                all_logits.append(logits)\n    \n    # Average predictions\n    avg_logits = torch.stack(all_logits).mean(dim=0)\n    return avg_logits\n\n# ============================================================================\n# SUBMISSION GENERATION\n# ============================================================================\ndef generate_submission(model, test_files, output_file, config, device):\n    \"\"\"Generate submission file\"\"\"\n    model.eval()\n    \n    trial_ids = []\n    predictions_list = []\n    \n    for test_file in test_files:\n        print(f\"Processing: {test_file}\")\n        \n        with h5py.File(test_file, 'r') as f:\n            trials = list(f.keys())\n        \n        for trial_id in tqdm(trials, desc=\"Predicting\"):\n            with h5py.File(test_file, 'r') as f:\n                trial = f[trial_id]\n                neural_data = torch.tensor(trial['input_features'][:], dtype=torch.float32)\n                neural_data = neural_data.unsqueeze(0).to(device)\n            \n            # Create padding mask\n            padding_mask = torch.zeros(1, neural_data.shape[1], dtype=torch.bool, device=device)\n            \n            # Predict with TTA\n            if config.use_tta:\n                logits = predict_with_tta(model, neural_data, padding_mask, config, device)\n            else:\n                with torch.no_grad():\n                    logits, _ = model(neural_data, src_key_padding_mask=padding_mask)\n            \n            # Get predictions\n            preds = torch.argmax(logits[0, :config.max_seq_len, :], dim=-1).cpu().numpy()\n            \n            # Ensure exactly 500 predictions\n            if len(preds) > config.max_seq_len:\n                preds = preds[:config.max_seq_len]\n            elif len(preds) < config.max_seq_len:\n                preds = np.pad(preds, (0, config.max_seq_len - len(preds)), mode='constant')\n            \n            pred_str = ' '.join(map(str, preds))\n            trial_ids.append(trial_id)\n            predictions_list.append(pred_str)\n    \n    # Create submission\n    submission_df = pd.DataFrame({\n        'id': trial_ids,\n        'predictions': predictions_list\n    })\n    \n    submission_df.to_csv(output_file, index=False)\n    print(f\"Submission saved: {output_file}\")\n    print(f\"Total predictions: {len(submission_df)}\")\n    return submission_df\n\n# ============================================================================\n# MAIN TRAINING SCRIPT\n# ============================================================================\ndef main():\n    print(\"=\" * 60)\n    print(\"Brain-to-Text 2.5 - OPTIMIZED VERSION\")\n    print(\"=\" * 60)\n    \n    # Get all data files\n    train_files, test_files = get_all_data_files()\n    \n    if not train_files:\n        print(\"No training files found! Using default path...\")\n        train_files = [\"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final/t15.2025.03.14/data_train.hdf5\"]\n        test_files = [\"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final/t15.2025.03.14/data_test.hdf5\"]\n    \n    # Create augmentation\n    augmentation = DataAugmentation(config)\n    \n    # Create dataset with ALL training files\n    print(\"\\nCreating datasets...\")\n    full_dataset = BrainDataset(\n        train_files,\n        is_test=False,\n        max_len=config.max_seq_len,\n        augmentation=augmentation\n    )\n    \n    # Split into train/val (90/10)\n    val_size = int(0.1 * len(full_dataset))\n    train_size = len(full_dataset) - val_size\n    train_dataset, val_dataset = torch.utils.data.random_split(\n        full_dataset, [train_size, val_size]\n    )\n    \n    print(f\"Train samples: {train_size}, Val samples: {val_size}\")\n    \n    # Create dataloaders\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=config.batch_size,\n        shuffle=True,\n        collate_fn=collate_fn,\n        num_workers=2,\n        pin_memory=True\n    )\n    \n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=config.batch_size,\n        shuffle=False,\n        collate_fn=collate_fn,\n        num_workers=2,\n        pin_memory=True\n    )\n    \n    # Create model\n    print(\"\\nCreating model...\")\n    # Use Hybrid model (often better than pure Conformer)\n    model = BrainHybridModel(config).to(device)\n    print(f\"Model parameters: {sum(p.numel() for p in model.parameters()):,}\")\n    \n    # Loss and optimizer\n    criterion = HybridLoss(\n        config.vocab_size,\n        ctc_weight=config.ctc_weight if config.use_ctc else 0,\n        label_smoothing=config.label_smoothing\n    )\n    \n    optimizer = torch.optim.AdamW(\n        model.parameters(),\n        lr=config.learning_rate,\n        weight_decay=config.weight_decay,\n        betas=(0.9, 0.98)\n    )\n    \n    # Learning rate scheduler\n    total_steps = config.num_epochs * len(train_loader) // config.accumulation_steps\n    scheduler = WarmupCosineScheduler(\n        optimizer,\n        warmup_steps=config.warmup_steps,\n        total_steps=total_steps\n    )\n    \n    # Training loop\n    print(\"\\n\" + \"=\" * 60)\n    print(\"Starting training...\")\n    print(\"=\" * 60)\n    \n    best_val_loss = float('inf')\n    patience = 10\n    patience_counter = 0\n    \n    for epoch in range(config.num_epochs):\n        print(f\"\\nEpoch {epoch + 1}/{config.num_epochs}\")\n        \n        train_loss = train_epoch(model, train_loader, criterion, optimizer, scheduler, device, config)\n        val_loss, val_acc = validate(model, val_loader, criterion, device)\n        \n        print(f\"Train Loss: {train_loss:.4f}, Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.4f}\")\n        \n        # Save best model\n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            patience_counter = 0\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'val_loss': val_loss,\n                'val_acc': val_acc,\n            }, 'best_model.pth')\n            print(f\"★ Saved best model (val_loss: {val_loss:.4f})\")\n        else:\n            patience_counter += 1\n            if patience_counter >= patience:\n                print(f\"Early stopping at epoch {epoch + 1}\")\n                break\n    \n    # Load best model\n    print(\"\\nLoading best model for inference...\")\n    checkpoint = torch.load('best_model.pth')\n    model.load_state_dict(checkpoint['model_state_dict'])\n    print(f\"Best model from epoch {checkpoint['epoch'] + 1}, val_loss: {checkpoint['val_loss']:.4f}\")\n    \n    # Generate submission\n    print(\"\\n\" + \"=\" * 60)\n    print(\"Generating submission...\")\n    print(\"=\" * 60)\n    \n    # Filter test files to only existing ones\n    existing_test_files = [f for f in test_files if os.path.exists(f)]\n    if existing_test_files:\n        submission_df = generate_submission(model, existing_test_files, 'submission.csv', config, device)\n    else:\n        print(\"No test files found!\")\n    \n    print(\"\\n\" + \"=\" * 60)\n    print(\"DONE!\")\n    print(\"=\" * 60)\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-05T17:12:54.496509Z","iopub.execute_input":"2025-12-05T17:12:54.497221Z","execution_failed":"2025-12-06T01:14:00.358Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Visualization and Analysis Cell ---\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom sklearn.metrics import confusion_matrix, classification_report\nimport pandas as pd\nfrom collections import Counter\nimport numpy as np\n\n# Set style for plots\nplt.style.use('seaborn-v0_8')\nsns.set_palette(\"husl\")\n\n# 1. Data Distribution Analysis\nprint(\"=== DATA DISTRIBUTION ANALYSIS ===\")\n\n# Analyze neural sequence lengths\nneural_lengths = []\nlabel_sequences = []\n\nwith h5py.File(TRAIN_FILE, 'r') as f:\n    for trial_key in f.keys():\n        trial = f[trial_key]\n        neural_lengths.append(trial['input_features'].shape[0])\n        label_sequences.append(trial['seq_class_ids'][:])\n\nneural_lengths = np.array(neural_lengths)\nprint(f\"Neural sequence length statistics:\")\nprint(f\"  Min: {neural_lengths.min()}\")\nprint(f\"  Max: {neural_lengths.max()}\")\nprint(f\"  Mean: {neural_lengths.mean():.1f}\")\nprint(f\"  Std: {neural_lengths.std():.1f}\")\n\n# Plot neural sequence length distribution\nplt.figure(figsize=(15, 12))\n\nplt.subplot(2, 3, 1)\nplt.hist(neural_lengths, bins=30, alpha=0.7, edgecolor='black')\nplt.xlabel('Sequence Length')\nplt.ylabel('Frequency')\nplt.title('Distribution of Neural Sequence Lengths')\nplt.grid(True, alpha=0.3)\n\n# 2. Label Distribution Analysis\nall_labels = np.concatenate(label_sequences)\nlabel_counts = Counter(all_labels)\n\nplt.subplot(2, 3, 2)\ntop_labels = dict(sorted(label_counts.items(), key=lambda x: x[1], reverse=True)[:20])\nplt.bar(top_labels.keys(), top_labels.values(), alpha=0.7)\nplt.xlabel('Label ID')\nplt.ylabel('Frequency')\nplt.title('Top 20 Most Frequent Labels')\nplt.xticks(rotation=45)\nplt.grid(True, alpha=0.3)\n\n# 3. Label Sequence Length Analysis\nlabel_lengths = [len(seq) for seq in label_sequences]\nplt.subplot(2, 3, 3)\nplt.hist(label_lengths, bins=20, alpha=0.7, edgecolor='black', color='orange')\nplt.xlabel('Label Sequence Length')\nplt.ylabel('Frequency')\nplt.title('Distribution of Label Sequence Lengths')\nplt.grid(True, alpha=0.3)\n\n# 4. Model Performance Analysis\nprint(\"\\n=== MODEL PERFORMANCE ANALYSIS ===\")\n\n# Get sample predictions for analysis - FIXED VERSION\nmodel.eval()\nsample_predictions = []\nsample_targets = []\n\nwith torch.no_grad():\n    batch_count = 0\n    for neural_data, labels, padding_mask, _ in train_loader:\n        if batch_count >= 3:  # Use first 3 batches for analysis\n            break\n            \n        neural_data = neural_data.to(device)\n        labels = labels.to(device)\n        padding_mask = padding_mask.to(device)\n        \n        logits = model(neural_data, src_key_padding_mask=padding_mask)\n        \n        batch_size, neural_len, vocab_size = logits.shape\n        label_len = labels.shape[1]\n        \n        # Align logits with labels - take first label_len time steps\n        if neural_len >= label_len:\n            logits_aligned = logits[:, :label_len, :]\n        else:\n            # If neural sequence is shorter, we can't make all predictions\n            logits_aligned = logits\n        \n        preds = torch.argmax(logits_aligned, dim=-1)\n        \n        # Process each sequence in the batch individually\n        for i in range(batch_size):\n            # Get the actual sequence length for this sample\n            actual_neural_len = neural_lengths[batch_count * batch_size + i] if (batch_count * batch_size + i) < len(neural_lengths) else neural_len\n            actual_label_len = min(actual_neural_len, label_len)\n            \n            # Take predictions for the actual sequence length\n            seq_preds = preds[i, :actual_label_len].cpu().numpy()\n            seq_labels = labels[i, :actual_label_len].cpu().numpy()\n            \n            # Only add non-zero labels (assuming 0 is padding)\n            non_zero_mask = seq_labels != 0\n            if non_zero_mask.sum() > 0:\n                sample_predictions.extend(seq_preds[non_zero_mask])\n                sample_targets.extend(seq_labels[non_zero_mask])\n        \n        batch_count += 1\n\nsample_predictions = np.array(sample_predictions)\nsample_targets = np.array(sample_targets)\n\nprint(f\"Collected {len(sample_predictions)} samples for analysis\")\n\nif len(sample_predictions) > 0:\n    # Calculate accuracy per class\n    correct_predictions = (sample_predictions == sample_targets)\n    accuracy_per_class = {}\n    for label in np.unique(sample_targets):\n        mask = sample_targets == label\n        if mask.sum() > 0:\n            accuracy_per_class[label] = correct_predictions[mask].mean()\n\n    # Plot accuracy per class\n    plt.subplot(2, 3, 4)\n    if accuracy_per_class:\n        top_classes = dict(sorted(accuracy_per_class.items(), key=lambda x: x[1], reverse=True)[:20])\n        plt.bar(top_classes.keys(), top_classes.values(), alpha=0.7, color='green')\n        plt.xlabel('Class ID')\n        plt.ylabel('Accuracy')\n        plt.title('Top 20 Classes by Accuracy')\n        plt.xticks(rotation=45)\n        plt.grid(True, alpha=0.3)\n    else:\n        plt.text(0.5, 0.5, 'No accuracy data available', ha='center', va='center', transform=plt.gca().transAxes)\n        plt.title('Accuracy per Class (No Data)')\n\n    # 5. Confusion Matrix (for top classes)\n    plt.subplot(2, 3, 5)\n    if len(sample_targets) > 0:\n        top_n_classes = 10\n        # Get most frequent classes in our sample\n        sample_label_counts = Counter(sample_targets)\n        top_class_indices = np.array([label for label, _ in sample_label_counts.most_common(top_n_classes)])\n        \n        # Filter predictions for top classes\n        mask = np.isin(sample_targets, top_class_indices)\n        if mask.sum() > 0:\n            cm = confusion_matrix(sample_targets[mask], sample_predictions[mask], labels=top_class_indices)\n            sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', \n                        xticklabels=top_class_indices, yticklabels=top_class_indices)\n            plt.xlabel('Predicted')\n            plt.ylabel('Actual')\n            plt.title(f'Confusion Matrix (Top {top_n_classes} Classes)')\n        else:\n            plt.text(0.5, 0.5, 'No data for confusion matrix', ha='center', va='center', transform=plt.gca().transAxes)\n            plt.title('Confusion Matrix (No Data)')\n    else:\n        plt.text(0.5, 0.5, 'No prediction data available', ha='center', va='center', transform=plt.gca().transAxes)\n        plt.title('Confusion Matrix (No Data)')\n\nelse:\n    # Placeholder plots if no prediction data\n    for i in [4, 5]:\n        plt.subplot(2, 3, i)\n        plt.text(0.5, 0.5, 'No prediction data available', ha='center', va='center', transform=plt.gca().transAxes)\n        plt.title('No Data Available')\n\n# 6. Prediction Confidence Analysis\nplt.subplot(2, 3, 6)\ntry:\n    with torch.no_grad():\n        # Get confidence scores for a single batch\n        neural_data, labels, padding_mask, _ = next(iter(train_loader))\n        neural_data = neural_data.to(device)\n        padding_mask = padding_mask.to(device)\n        \n        logits = model(neural_data, src_key_padding_mask=padding_mask)\n        probabilities = F.softmax(logits, dim=-1)\n        max_probs, _ = torch.max(probabilities, dim=-1)\n        \n        # Process each sequence individually\n        confidence_scores = []\n        for i in range(neural_data.shape[0]):\n            seq_len = min(neural_data.shape[1], labels.shape[1])\n            seq_probs = max_probs[i, :seq_len].cpu().numpy()\n            seq_labels = labels[i, :seq_len].cpu().numpy()\n            \n            # Only keep non-padded positions\n            non_zero_mask = seq_labels != 0\n            confidence_scores.extend(seq_probs[non_zero_mask])\n        \n        if confidence_scores:\n            plt.hist(confidence_scores, bins=30, alpha=0.7, edgecolor='black', color='purple')\n            plt.xlabel('Maximum Probability')\n            plt.ylabel('Frequency')\n            plt.title('Distribution of Prediction Confidence')\n            plt.grid(True, alpha=0.3)\n        else:\n            plt.text(0.5, 0.5, 'No confidence data', ha='center', va='center', transform=plt.gca().transAxes)\n            plt.title('Prediction Confidence (No Data)')\n            \nexcept Exception as e:\n    plt.text(0.5, 0.5, f'Error: {str(e)}', ha='center', va='center', transform=plt.gca().transAxes, fontsize=8)\n    plt.title('Prediction Confidence (Error)')\n\nplt.tight_layout()\nplt.show()\n\n# 7. Detailed Performance Metrics\nprint(\"\\n=== DETAILED PERFORMANCE METRICS ===\")\nif len(sample_predictions) > 0:\n    overall_accuracy = (sample_predictions == sample_targets).mean()\n    print(f\"Overall Accuracy: {overall_accuracy:.4f}\")\n    \n    # Per-class metrics for top classes\n    print(\"\\nPer-class metrics for top 10 most frequent classes in sample:\")\n    sample_label_counts = Counter(sample_targets)\n    top_10_classes = [label for label, _ in sample_label_counts.most_common(10)]\n    \n    for class_id in top_10_classes:\n        class_mask = sample_targets == class_id\n        if class_mask.sum() > 0:\n            class_accuracy = (sample_predictions[class_mask] == class_id).mean()\n            class_frequency = class_mask.mean()\n            support = class_mask.sum()\n            print(f\"  Class {class_id:3d}: Accuracy={class_accuracy:.4f}, Frequency={class_frequency:.4f}, Support={support}\")\nelse:\n    print(\"No prediction data available for detailed metrics\")\n\n# 8. Training Dynamics Analysis\nprint(\"\\n=== TRAINING DYNAMICS ===\")\nprint(\"Model Weights Analysis:\")\nweight_stats = []\nfor name, param in model.named_parameters():\n    if param.requires_grad and param.numel() > 0:\n        weight_stats.append({\n            'Layer': name,\n            'Mean': param.data.mean().item(),\n            'Std': param.data.std().item(),\n            'Min': param.data.min().item(),\n            'Max': param.data.max().item(),\n            'Parameters': param.numel()\n        })\n\nif weight_stats:\n    # Create a summary DataFrame\n    weight_df = pd.DataFrame(weight_stats)\n    print(weight_df.to_string(index=False))\nelse:\n    print(\"No weight statistics available\")\n\n# 9. Additional Visualizations\nplt.figure(figsize=(15, 5))\n\n# Neural vs Label length correlation\nplt.subplot(1, 3, 1)\nplt.scatter(neural_lengths, label_lengths, alpha=0.6)\nplt.xlabel('Neural Sequence Length')\nplt.ylabel('Label Sequence Length')\nplt.title('Neural vs Label Sequence Lengths')\nplt.grid(True, alpha=0.3)\n\n# Class distribution (full dataset)\nplt.subplot(1, 3, 2)\ntop_30_labels = dict(sorted(label_counts.items(), key=lambda x: x[1], reverse=True)[:30])\nplt.bar(range(len(top_30_labels)), list(top_30_labels.values()), alpha=0.7)\nplt.xlabel('Class Rank')\nplt.ylabel('Frequency')\nplt.title('Class Frequency Distribution (Top 30)')\nplt.xticks(range(len(top_30_labels)), list(top_30_labels.keys()), rotation=45)\n\n# Sequence length over trials\nplt.subplot(1, 3, 3)\nplt.plot(neural_lengths, marker='o', alpha=0.7, markersize=3)\nplt.xlabel('Trial Index')\nplt.ylabel('Sequence Length')\nplt.title('Sequence Length by Trial')\nplt.grid(True, alpha=0.3)\n\nplt.tight_layout()\nplt.show()\n\n# 10. Data Quality Check\nprint(\"\\n=== DATA QUALITY CHECK ===\")\nprint(f\"Total trials: {len(neural_lengths)}\")\nprint(f\"Total label tokens: {len(all_labels)}\")\nprint(f\"Unique classes: {len(np.unique(all_labels))}\")\nprint(f\"Label value range: {all_labels.min()} to {all_labels.max()}\")\n\n# Check for any anomalies\nzero_labels = (all_labels == 0).sum()\nprint(f\"Zero labels (potential padding): {zero_labels} ({zero_labels/len(all_labels):.2%})\")\n\n# Check sequence length consistency\nneural_to_label_ratio = neural_lengths.mean() / np.mean(label_lengths)\nprint(f\"Neural-to-label length ratio: {neural_to_label_ratio:.2f}\")\n\n# 11. Model Capacity Analysis\nprint(\"\\n=== MODEL CAPACITY ANALYSIS ===\")\ntotal_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f\"Total trainable parameters: {total_params:,}\")\nprint(f\"Parameters per class: {total_params / config.vocab_size:,.0f}\")\n\n# Estimate model capacity vs data size\ndata_points = len(all_labels)\nprint(f\"Total data points: {data_points:,}\")\nprint(f\"Parameters per data point: {total_params / data_points:.2f}\")\n\nprint(\"\\n=== ANALYSIS COMPLETE ===\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Conclusion:\n\nThe bidirectional LSTM architecture effectively captures temporal dependencies in neural data for text generation. Key successes include robust handling of variable-length inputs, stable training through gradient management, and generation of competition-ready submissions. Future improvements could explore transformer architectures, attention mechanisms, and more sophisticated sequence alignment strategies to enhance prediction accuracy.","metadata":{}}]}