{"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":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"18cfac8f","cell_type":"code","source":"# Install memory monitoring\n!pip install -q psutil\n\n# Core libraries\nimport os\nimport gc\nimport psutil\nimport numpy as np\nimport pandas as pd\nimport h5py\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# Deep learning\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.nn.utils.rnn import pad_sequence, pack_padded_sequence, pad_packed_sequence\n\n# Visualization\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nsns.set_style('whitegrid')\n\n# Set random seeds for reproducibility\nSEED = 42\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(SEED)\n\n# Device configuration\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T07:07:26.080462Z","iopub.execute_input":"2026-01-01T07:07:26.081218Z","iopub.status.idle":"2026-01-01T07:07:37.061874Z","shell.execute_reply.started":"2026-01-01T07:07:26.081171Z","shell.execute_reply":"2026-01-01T07:07:37.061159Z"}},"outputs":[],"execution_count":null},{"id":"6a2d4cbc","cell_type":"code","source":"# Timing utilities\nimport time\nfrom datetime import datetime\n\nclass Timer:\n    \"\"\"Context manager for timing code blocks.\"\"\"\n    def __init__(self, name=\"\"):\n        self.name = name\n        \n    def __enter__(self):\n        self.start = time.time()\n        if self.name:\n            print(f\"⏱️  Starting: {self.name}...\")\n        return self\n    \n    def __exit__(self, *args):\n        self.end = time.time()\n        self.duration = self.end - self.start\n        if self.name:\n            print(f\"✓ Completed: {self.name} in {self.duration:.1f}s ({self.duration/60:.1f} min)\")\n\n# Global timer for notebook\nnotebook_start_time = time.time()\n\ndef print_elapsed_time():\n    \"\"\"Print total elapsed time since notebook start.\"\"\"\n    elapsed = time.time() - notebook_start_time\n    print(f\"\\n⏱️  Total elapsed time: {elapsed:.0f}s ({elapsed/60:.1f} min / {elapsed/3600:.2f} hrs)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T07:07:37.063405Z","iopub.execute_input":"2026-01-01T07:07:37.063753Z","iopub.status.idle":"2026-01-01T07:07:37.070121Z","shell.execute_reply.started":"2026-01-01T07:07:37.063727Z","shell.execute_reply":"2026-01-01T07:07:37.069434Z"}},"outputs":[],"execution_count":null},{"id":"1a31855e","cell_type":"code","source":"# Phoneme mapping (40 classes including BLANK and silence)\nLOGIT_TO_PHONEME = [\n    'BLANK',    # CTC blank symbol\n    'AA', 'AE', 'AH', 'AO', 'AW',\n    'AY', 'B', 'CH', 'D', 'DH',\n    'EH', 'ER', 'EY', 'F', 'G',\n    'HH', 'IH', 'IY', 'JH', 'K',\n    'L', 'M', 'N', 'NG', 'OW',\n    'OY', 'P', 'R', 'S', 'SH',\n    'T', 'TH', 'UH', 'UW', 'V',\n    'W', 'Y', 'Z', 'ZH',\n    ' | ',    # silence token\n]\n\nPHONEME_TO_LOGIT = {p: i for i, p in enumerate(LOGIT_TO_PHONEME)}\n\n# Data paths\nDATA_ROOT = Path('/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final')\n\n# Model hyperparameters - OPTIMIZED FOR <10 MIN RUNTIME\nCONFIG = {\n    'input_dim': 512,          # Neural features\n    'hidden_dim': 256,         # ⚡ Reduced for speed\n    'num_layers': 2,           # ⚡ Reduced to 2 layers\n    'num_classes': 40,         # Number of phonemes\n    'dropout': 0.2,            # ⚡ Lower dropout\n    'bidirectional': True,\n    'rnn_type': 'GRU',         # GRU or LSTM\n    \n    # Training - ULTRA FAST\n    'batch_size': 64,          # ⚡ Maximum batch size\n    'learning_rate': 2e-3,     # ⚡ Higher LR for faster convergence\n    'num_epochs': 3,           # ⚡ Minimal epochs\n    'grad_clip': 5.0,\n    'patience': 2,             # ⚡ Very aggressive early stopping\n    \n    # Data augmentation - DISABLED\n    'temporal_mask_prob': 0.0,  # ⚡ Disabled\n    'feature_mask_prob': 0.0,   # ⚡ Disabled\n}\n\nprint(\"Configuration loaded successfully!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T07:07:37.070776Z","iopub.execute_input":"2026-01-01T07:07:37.071037Z","iopub.status.idle":"2026-01-01T07:07:37.086918Z","shell.execute_reply.started":"2026-01-01T07:07:37.071013Z","shell.execute_reply":"2026-01-01T07:07:37.086225Z"}},"outputs":[],"execution_count":null},{"id":"0fad404f","cell_type":"code","source":"def load_h5py_file(file_path, include_labels=True):\n    \"\"\"\n    Load data from HDF5 file.\n    \n    Args:\n        file_path: Path to HDF5 file\n        include_labels: Whether to include labels (False for test set)\n    \n    Returns:\n        Dictionary containing all trial data\n    \"\"\"\n    data = {\n        'neural_features': [],\n        'n_time_steps': [],\n        'seq_class_ids': [],\n        'seq_len': [],\n        'transcriptions': [],\n        'sentence_label': [],\n        'session': [],\n        'block_num': [],\n        'trial_num': [],\n    }\n    \n    with h5py.File(file_path, 'r') as f:\n        keys = list(f.keys())\n        \n        for key in keys:\n            g = f[key]\n            \n            # Neural data (always present) - load efficiently\n            neural_features = np.array(g['input_features'], dtype=np.float32)\n            n_time_steps = g.attrs['n_time_steps']\n            session = g.attrs['session']\n            block_num = g.attrs['block_num']\n            trial_num = g.attrs['trial_num']\n            \n            # Labels (only in train/val)\n            if include_labels:\n                seq_class_ids = g['seq_class_ids'][:] if 'seq_class_ids' in g else None\n                seq_len = g.attrs['seq_len'] if 'seq_len' in g.attrs else None\n                transcription = g['transcription'][:] if 'transcription' in g else None\n                sentence_label = g.attrs['sentence_label'] if 'sentence_label' in g.attrs else None\n            else:\n                seq_class_ids = None\n                seq_len = None\n                transcription = None\n                sentence_label = None\n            \n            data['neural_features'].append(neural_features)\n            data['n_time_steps'].append(n_time_steps)\n            data['seq_class_ids'].append(seq_class_ids)\n            data['seq_len'].append(seq_len)\n            data['transcriptions'].append(transcription)\n            data['sentence_label'].append(sentence_label)\n            data['session'].append(session)\n            data['block_num'].append(block_num)\n            data['trial_num'].append(trial_num)\n    \n    return data\n\n\ndef load_all_sessions(data_root, split='train', include_labels=True):\n    \"\"\"\n    Load all sessions for a given split.\n    \n    Args:\n        data_root: Root directory containing session folders\n        split: 'train', 'val', or 'test'\n        include_labels: Whether to include labels\n    \n    Returns:\n        Combined dictionary with all data\n    \"\"\"\n    all_data = {\n        'neural_features': [],\n        'n_time_steps': [],\n        'seq_class_ids': [],\n        'seq_len': [],\n        'transcriptions': [],\n        'sentence_label': [],\n        'session': [],\n        'block_num': [],\n        'trial_num': [],\n    }\n    \n    # Get all session directories\n    session_dirs = sorted([d for d in data_root.iterdir() if d.is_dir()])\n    \n    print(f\"Loading {split} data...\")\n    for session_dir in tqdm(session_dirs):\n        file_path = session_dir / f'data_{split}.hdf5'\n        \n        if file_path.exists():\n            session_data = load_h5py_file(file_path, include_labels)\n            \n            # Append to combined data\n            for key in all_data.keys():\n                all_data[key].extend(session_data[key])\n    \n    print(f\"Loaded {len(all_data['neural_features'])} trials for {split} split\")\n    return all_data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T07:07:37.088445Z","iopub.execute_input":"2026-01-01T07:07:37.088714Z","iopub.status.idle":"2026-01-01T07:07:37.108804Z","shell.execute_reply.started":"2026-01-01T07:07:37.088694Z","shell.execute_reply":"2026-01-01T07:07:37.108291Z"}},"outputs":[],"execution_count":null},{"id":"a1125683","cell_type":"code","source":"class BrainToTextDataset(Dataset):\n    \"\"\"\n    Dataset for brain-to-text neural decoding.\n    \"\"\"\n    \n    def __init__(self, data, augment=False, temporal_mask_prob=0.1, feature_mask_prob=0.1):\n        self.neural_features = data['neural_features']\n        self.seq_class_ids = data['seq_class_ids']\n        self.n_time_steps = data['n_time_steps']\n        self.seq_len = data['seq_len']\n        self.augment = augment\n        self.temporal_mask_prob = temporal_mask_prob\n        self.feature_mask_prob = feature_mask_prob\n        \n    def __len__(self):\n        return len(self.neural_features)\n    \n    def __getitem__(self, idx):\n        # Get neural features - check shape first\n        neural_raw = self.neural_features[idx]\n        \n        # Ensure correct shape (T, 512)\n        # Data is stored as (512, T) so we transpose\n        if neural_raw.shape[0] == 512:\n            neural = neural_raw.T  # (512, T) -> (T, 512)\n        elif neural_raw.shape[1] == 512:\n            neural = neural_raw  # Already (T, 512)\n        else:\n            # Handle unexpected shape by padding/truncating to 512 features\n            if neural_raw.shape[0] < neural_raw.shape[1]:\n                neural = neural_raw.T\n            else:\n                neural = neural_raw\n            \n            # Ensure exactly 512 features\n            if neural.shape[1] != 512:\n                if neural.shape[1] > 512:\n                    neural = neural[:, :512]  # Truncate\n                else:\n                    # Pad with zeros\n                    pad_width = ((0, 0), (0, 512 - neural.shape[1]))\n                    neural = np.pad(neural, pad_width, mode='constant', constant_values=0)\n        \n        neural = torch.FloatTensor(neural)\n        \n        # Normalize features (z-score normalization)\n        neural = (neural - neural.mean(dim=0, keepdim=True)) / (neural.std(dim=0, keepdim=True) + 1e-8)\n        \n        # Data augmentation (only during training)\n        if self.augment:\n            neural = self.apply_augmentation(neural)\n        \n        # Get labels\n        if self.seq_class_ids[idx] is not None:\n            labels = torch.LongTensor(self.seq_class_ids[idx])\n        else:\n            labels = torch.LongTensor([])  # Empty for test set\n        \n        return {\n            'neural': neural,\n            'labels': labels,\n            'neural_len': self.n_time_steps[idx],\n            'label_len': self.seq_len[idx] if self.seq_len[idx] is not None else 0\n        }\n    \n    def apply_augmentation(self, neural):\n        \"\"\"\n        Apply temporal and feature masking augmentation.\n        \"\"\"\n        T, F = neural.shape\n        \n        # Temporal masking\n        if np.random.rand() < self.temporal_mask_prob:\n            mask_len = int(T * 0.1)  # Mask 10% of time steps\n            mask_start = np.random.randint(0, max(1, T - mask_len))\n            neural[mask_start:mask_start + mask_len] = 0\n        \n        # Feature masking\n        if np.random.rand() < self.feature_mask_prob:\n            num_mask = int(F * 0.1)  # Mask 10% of features\n            mask_features = np.random.choice(F, num_mask, replace=False)\n            neural[:, mask_features] = 0\n        \n        return neural\n\n\ndef collate_fn(batch):\n    \"\"\"\n    Collate function for DataLoader to handle variable-length sequences.\n    \"\"\"\n    neural = [item['neural'] for item in batch]\n    labels = [item['labels'] for item in batch]\n    neural_lens = torch.LongTensor([item['neural_len'] for item in batch])\n    label_lens = torch.LongTensor([item['label_len'] for item in batch])\n    \n    # Pad sequences\n    neural_padded = pad_sequence(neural, batch_first=True)  # (B, T, 512)\n    labels_padded = pad_sequence(labels, batch_first=True, padding_value=0)  # (B, L)\n    \n    return neural_padded, labels_padded, neural_lens, label_lens","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T07:07:37.109551Z","iopub.execute_input":"2026-01-01T07:07:37.109732Z","iopub.status.idle":"2026-01-01T07:07:37.124354Z","shell.execute_reply.started":"2026-01-01T07:07:37.109715Z","shell.execute_reply":"2026-01-01T07:07:37.123758Z"}},"outputs":[],"execution_count":null},{"id":"a10d5836","cell_type":"code","source":"class BrainToTextRNN(nn.Module):\n    \"\"\"\n    Bidirectional RNN model for phoneme prediction from neural activity.\n    Uses CTC loss for sequence-to-sequence learning without alignment.\n    \"\"\"\n    \n    def __init__(self, config):\n        super(BrainToTextRNN, self).__init__()\n        \n        self.input_dim = config['input_dim']\n        self.hidden_dim = config['hidden_dim']\n        self.num_layers = config['num_layers']\n        self.num_classes = config['num_classes']\n        self.dropout = config['dropout']\n        self.bidirectional = config['bidirectional']\n        \n        # Input projection\n        self.input_proj = nn.Sequential(\n            nn.Linear(self.input_dim, self.hidden_dim),\n            nn.LayerNorm(self.hidden_dim),\n            nn.ReLU(),\n            nn.Dropout(self.dropout)\n        )\n        \n        # RNN layers\n        if config['rnn_type'] == 'GRU':\n            self.rnn = nn.GRU(\n                self.hidden_dim,\n                self.hidden_dim,\n                num_layers=self.num_layers,\n                batch_first=True,\n                dropout=self.dropout if self.num_layers > 1 else 0,\n                bidirectional=self.bidirectional\n            )\n        else:\n            self.rnn = nn.LSTM(\n                self.hidden_dim,\n                self.hidden_dim,\n                num_layers=self.num_layers,\n                batch_first=True,\n                dropout=self.dropout if self.num_layers > 1 else 0,\n                bidirectional=self.bidirectional\n            )\n        \n        # Output projection\n        rnn_output_dim = self.hidden_dim * 2 if self.bidirectional else self.hidden_dim\n        \n        self.output_proj = nn.Sequential(\n            nn.Linear(rnn_output_dim, self.hidden_dim),\n            nn.LayerNorm(self.hidden_dim),\n            nn.ReLU(),\n            nn.Dropout(self.dropout),\n            nn.Linear(self.hidden_dim, self.num_classes)\n        )\n    \n    def forward(self, x, lengths):\n        \"\"\"\n        Forward pass.\n        \n        Args:\n            x: (B, T, input_dim) neural features\n            lengths: (B,) actual sequence lengths\n        \n        Returns:\n            log_probs: (T, B, num_classes) log probabilities for CTC\n        \"\"\"\n        # Input projection\n        x = self.input_proj(x)  # (B, T, hidden_dim)\n        \n        # Pack padded sequence for efficient RNN processing\n        x_packed = pack_padded_sequence(x, lengths.cpu(), batch_first=True, enforce_sorted=False)\n        \n        # RNN\n        rnn_out, _ = self.rnn(x_packed)\n        \n        # Unpack\n        rnn_out, _ = pad_packed_sequence(rnn_out, batch_first=True)\n        \n        # Output projection\n        logits = self.output_proj(rnn_out)  # (B, T, num_classes)\n        \n        # CTC expects (T, B, C)\n        log_probs = F.log_softmax(logits, dim=-1).transpose(0, 1)\n        \n        return log_probs","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T07:07:37.125236Z","iopub.execute_input":"2026-01-01T07:07:37.125547Z","iopub.status.idle":"2026-01-01T07:07:37.143397Z","shell.execute_reply.started":"2026-01-01T07:07:37.125526Z","shell.execute_reply":"2026-01-01T07:07:37.142795Z"}},"outputs":[],"execution_count":null},{"id":"4e580042","cell_type":"code","source":"# Memory monitoring helper\ndef print_memory_usage():\n    \"\"\"Print current memory usage.\"\"\"\n    if torch.cuda.is_available():\n        gpu_mem = torch.cuda.memory_allocated() / 1024**3\n        gpu_mem_cached = torch.cuda.memory_reserved() / 1024**3\n        print(f\"  GPU Memory: {gpu_mem:.2f} GB allocated, {gpu_mem_cached:.2f} GB cached\")\n    import psutil\n    ram = psutil.Process().memory_info().rss / 1024**3\n    print(f\"  RAM Usage: {ram:.2f} GB\")\n\nprint(\"\\n📊 Memory usage before loading:\")\nprint_memory_usage()\n\n# Load training data\nwith Timer(\"Load training data\"):\n    train_data = load_all_sessions(DATA_ROOT, split='train', include_labels=True)\n    print_memory_usage()\n\n# Create train dataset and clear raw data\ntrain_dataset = BrainToTextDataset(train_data, augment=True, \n                                   temporal_mask_prob=CONFIG['temporal_mask_prob'],\n                                   feature_mask_prob=CONFIG['feature_mask_prob'])\ndel train_data  # Free memory\ngc.collect()\nprint(\"✓ Train dataset created, raw data cleared\")\nprint_memory_usage()\n\n# Load validation data\nwith Timer(\"Load validation data\"):\n    val_data = load_all_sessions(DATA_ROOT, split='val', include_labels=True)\n    print_memory_usage()\n\n# Create val dataset and clear raw data\nval_dataset = BrainToTextDataset(val_data, augment=False)\ndel val_data  # Free memory\ngc.collect()\nprint(\"✓ Val dataset created, raw data cleared\")\nprint_memory_usage()\n\n# Load test data\nwith Timer(\"Load test data\"):\n    test_data = load_all_sessions(DATA_ROOT, split='test', include_labels=False)\n    print_memory_usage()\n\n# Create test dataset and clear raw data\ntest_dataset = BrainToTextDataset(test_data, augment=False)\ndel test_data  # Free memory\ngc.collect()\nprint(\"✓ Test dataset created, raw data cleared\")\nprint_memory_usage()\n\n# Create dataloaders (num_workers=0 to avoid multiprocessing issues)\ntrain_loader = DataLoader(train_dataset, batch_size=CONFIG['batch_size'], \n                         shuffle=True, collate_fn=collate_fn, num_workers=0)\nval_loader = DataLoader(val_dataset, batch_size=CONFIG['batch_size'], \n                       shuffle=False, collate_fn=collate_fn, num_workers=0)\ntest_loader = DataLoader(test_dataset, batch_size=CONFIG['batch_size'], \n                        shuffle=False, collate_fn=collate_fn, num_workers=0)\n\nprint(f\"\\n✓ Dataset sizes:\")\nprint(f\"  Train: {len(train_dataset)} trials\")\nprint(f\"  Val: {len(val_dataset)} trials\")\nprint(f\"  Test: {len(test_dataset)} trials\")\nprint(f\"  Batches per epoch: {len(train_loader)}\")\n\nprint_elapsed_time()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T07:07:37.144097Z","iopub.execute_input":"2026-01-01T07:07:37.144306Z","iopub.status.idle":"2026-01-01T07:11:28.263309Z","shell.execute_reply.started":"2026-01-01T07:07:37.144288Z","shell.execute_reply":"2026-01-01T07:11:28.262630Z"}},"outputs":[],"execution_count":null},{"id":"d3d5b69d","cell_type":"code","source":"def train_epoch(model, loader, optimizer, criterion, device, grad_clip):\n    \"\"\"\n    Train for one epoch.\n    \"\"\"\n    model.train()\n    total_loss = 0\n    \n    pbar = tqdm(loader, desc='Training')\n    for neural, labels, neural_lens, label_lens in pbar:\n        neural = neural.to(device)\n        labels = labels.to(device)\n        neural_lens = neural_lens.to(device)\n        label_lens = label_lens.to(device)\n        \n        # Forward pass\n        log_probs = model(neural, neural_lens)  # (T, B, C)\n        \n        # CTC Loss\n        loss = criterion(log_probs, labels, neural_lens, label_lens)\n        \n        # Backward pass\n        optimizer.zero_grad()\n        loss.backward()\n        \n        # Gradient clipping\n        torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip)\n        \n        optimizer.step()\n        \n        total_loss += loss.item()\n        pbar.set_postfix({'loss': loss.item()})\n    \n    return total_loss / len(loader)\n\n\ndef validate_epoch(model, loader, criterion, device):\n    \"\"\"\n    Validate for one epoch.\n    \"\"\"\n    model.eval()\n    total_loss = 0\n    \n    with torch.no_grad():\n        pbar = tqdm(loader, desc='Validation')\n        for neural, labels, neural_lens, label_lens in pbar:\n            neural = neural.to(device)\n            labels = labels.to(device)\n            neural_lens = neural_lens.to(device)\n            label_lens = label_lens.to(device)\n            \n            # Forward pass\n            log_probs = model(neural, neural_lens)\n            \n            # CTC Loss\n            loss = criterion(log_probs, labels, neural_lens, label_lens)\n            \n            total_loss += loss.item()\n            pbar.set_postfix({'loss': loss.item()})\n    \n    return total_loss / len(loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T07:11:28.264200Z","iopub.execute_input":"2026-01-01T07:11:28.264433Z","iopub.status.idle":"2026-01-01T07:11:28.272610Z","shell.execute_reply.started":"2026-01-01T07:11:28.264413Z","shell.execute_reply":"2026-01-01T07:11:28.271992Z"}},"outputs":[],"execution_count":null},{"id":"47fb8233","cell_type":"code","source":"# Initialize model\nmodel = BrainToTextRNN(CONFIG).to(device)\n\n# Count parameters\ntotal_params = sum(p.numel() for p in model.parameters())\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f\"Total parameters: {total_params:,}\")\nprint(f\"Trainable parameters: {trainable_params:,}\")\n\n# Loss function (CTC Loss)\ncriterion = nn.CTCLoss(blank=0, zero_infinity=True)\n\n# Optimizer with weight decay\noptimizer = torch.optim.AdamW(model.parameters(), lr=CONFIG['learning_rate'], weight_decay=1e-5)\n\n# Learning rate scheduler\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer, mode='min', factor=0.5, patience=5\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T07:11:28.273632Z","iopub.execute_input":"2026-01-01T07:11:28.273892Z","iopub.status.idle":"2026-01-01T07:11:31.752046Z","shell.execute_reply.started":"2026-01-01T07:11:28.273861Z","shell.execute_reply":"2026-01-01T07:11:31.751459Z"}},"outputs":[],"execution_count":null},{"id":"2ead9e0b","cell_type":"code","source":"# Training loop with early stopping\ntrain_losses = []\nval_losses = []\nbest_val_loss = float('inf')\npatience_counter = 0\n\nprint(\"Starting training...\\n\")\nprint(f\"⏱️  Estimated training time: ~{CONFIG['num_epochs'] * 3} minutes\\n\")\n\ntraining_start = time.time()\n\nfor epoch in range(CONFIG['num_epochs']):\n    epoch_start = time.time()\n    \n    print(f\"\\nEpoch {epoch + 1}/{CONFIG['num_epochs']}\")\n    print(\"=\" * 50)\n    \n    # Train\n    train_loss = train_epoch(model, train_loader, optimizer, criterion, device, CONFIG['grad_clip'])\n    train_losses.append(train_loss)\n    \n    # Validate\n    val_loss = validate_epoch(model, val_loader, criterion, device)\n    val_losses.append(val_loss)\n    \n    # Learning rate scheduling\n    scheduler.step(val_loss)\n    \n    # Clear GPU cache to prevent memory issues\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    gc.collect()\n    \n    epoch_time = time.time() - epoch_start\n    \n    print(f\"\\nTrain Loss: {train_loss:.4f}\")\n    print(f\"Val Loss: {val_loss:.4f}\")\n    print(f\"Learning Rate: {optimizer.param_groups[0]['lr']:.6f}\")\n    print(f\"Epoch Time: {epoch_time:.1f}s ({epoch_time/60:.1f} min)\")\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            'config': CONFIG,\n        }, 'best_model.pt')\n        print(f\"✓ Best model saved! (Val Loss: {val_loss:.4f})\")\n    else:\n        patience_counter += 1\n        print(f\"Patience: {patience_counter}/{CONFIG['patience']}\")\n    \n    # Early stopping\n    if patience_counter >= CONFIG['patience']:\n        print(f\"\\nEarly stopping triggered after {epoch + 1} epochs\")\n        break\n\ntraining_time = time.time() - training_start\nprint(f\"\\n✓ Training completed in {training_time:.0f}s ({training_time/60:.1f} min)\")\nprint_elapsed_time()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T07:11:31.753881Z","iopub.execute_input":"2026-01-01T07:11:31.754277Z","iopub.status.idle":"2026-01-01T07:18:09.378821Z","shell.execute_reply.started":"2026-01-01T07:11:31.754255Z","shell.execute_reply":"2026-01-01T07:18:09.378128Z"}},"outputs":[],"execution_count":null},{"id":"a9445155","cell_type":"code","source":"def predict_test_set_fast(model, test_loader, device):\n  \n    model.eval()\n    all_predictions = []\n    \n    with torch.no_grad():\n        pbar = tqdm(test_loader, desc='⚡ Fast greedy decoding')\n        for neural, _, neural_lens, _ in pbar:\n            neural = neural.to(device)\n            neural_lens = neural_lens.to(device)\n            \n            # Forward pass\n            log_probs = model(neural, neural_lens)  # (T, B, C)\n            \n            # Greedy decode (fastest method)\n            decoded_phonemes = ctc_greedy_decode(log_probs)\n            \n            # Convert to simple text\n            for phoneme_ids in decoded_phonemes:\n                # Simple phoneme-to-text conversion\n                text = phoneme_to_text_simple(phoneme_ids)\n                all_predictions.append(text if len(text) > 0 else \"the\")\n    \n    return all_predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T07:18:09.379892Z","iopub.execute_input":"2026-01-01T07:18:09.380205Z","iopub.status.idle":"2026-01-01T07:18:09.386079Z","shell.execute_reply.started":"2026-01-01T07:18:09.380170Z","shell.execute_reply":"2026-01-01T07:18:09.385428Z"}},"outputs":[],"execution_count":null},{"id":"f6f4eb0d","cell_type":"code","source":"def ctc_greedy_decode(log_probs):\n \n    # Get most likely class at each time step\n    _, max_indices = torch.max(log_probs, dim=2)  # (T, B)\n    max_indices = max_indices.transpose(0, 1)  # (B, T)\n    \n    decoded = []\n    for sequence in max_indices:\n        # Remove consecutive duplicates\n        prev = -1\n        result = []\n        for idx in sequence.cpu().numpy():\n            if idx != prev:\n                result.append(idx)\n                prev = idx\n        \n        # Remove blanks (index 0)\n        result = [idx for idx in result if idx != 0]\n        decoded.append(result)\n    \n    return decoded\n\n\ndef phoneme_to_text_simple(phoneme_ids):\n    \n    phonemes = [LOGIT_TO_PHONEME[idx] for idx in phoneme_ids if idx < len(LOGIT_TO_PHONEME)]\n    \n    # Simple phoneme-to-word mapping (this is very basic)\n    # In practice, you would use a proper phoneme-to-grapheme model or language model\n    phoneme_str = ' '.join(phonemes)\n    \n    # Remove silence tokens\n    phoneme_str = phoneme_str.replace(' | ', ' ')\n    \n    # This is a placeholder - ideally use a proper decoder\n    # For now, just return phoneme sequence as proxy\n    return phoneme_str.strip()\n\n\ndef predict_test_set(model, test_loader, device):\n    \"\"\"\n    Generate predictions for test set (baseline greedy decoder).\n    \n    Returns:\n        List of predicted texts\n    \"\"\"\n    model.eval()\n    all_predictions = []\n    \n    with torch.no_grad():\n        pbar = tqdm(test_loader, desc='Generating predictions (greedy)')\n        for neural, _, neural_lens, _ in pbar:\n            neural = neural.to(device)\n            neural_lens = neural_lens.to(device)\n            \n            # Forward pass\n            log_probs = model(neural, neural_lens)  # (T, B, C)\n            \n            # Decode\n            decoded_phonemes = ctc_greedy_decode(log_probs)\n            \n            # Convert to text\n            for phoneme_ids in decoded_phonemes:\n                text = phoneme_to_text_simple(phoneme_ids)\n                all_predictions.append(text)\n    \n    return all_predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T07:18:09.386787Z","iopub.execute_input":"2026-01-01T07:18:09.386990Z","iopub.status.idle":"2026-01-01T07:18:09.406579Z","shell.execute_reply.started":"2026-01-01T07:18:09.386949Z","shell.execute_reply":"2026-01-01T07:18:09.405977Z"}},"outputs":[],"execution_count":null},{"id":"799eafa7","cell_type":"code","source":"# Load best model and generate predictions\ncheckpoint = torch.load('best_model.pt')\nmodel.load_state_dict(checkpoint['model_state_dict'])\nprint(f\"✓ Loaded best model (Val Loss: {checkpoint['val_loss']:.4f})\")\n\n# Generate predictions\npredictions = predict_test_set_fast(model, test_loader, device)\nprint(f\"✓ Generated {len(predictions)} predictions\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T07:18:09.407326Z","iopub.execute_input":"2026-01-01T07:18:09.407601Z","iopub.status.idle":"2026-01-01T07:18:20.017541Z","shell.execute_reply.started":"2026-01-01T07:18:09.407573Z","shell.execute_reply":"2026-01-01T07:18:20.016794Z"}},"outputs":[],"execution_count":null},{"id":"33e8dfee","cell_type":"code","source":"# Create and save submission\nsubmission = pd.DataFrame({'id': range(len(predictions)), 'text': predictions})\nsubmission.to_csv('submission.csv', index=False)\nprint(f\"✓ Submission saved: {len(submission)} predictions → submission.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T07:18:20.018584Z","iopub.execute_input":"2026-01-01T07:18:20.019064Z","iopub.status.idle":"2026-01-01T07:18:20.051161Z","shell.execute_reply.started":"2026-01-01T07:18:20.019040Z","shell.execute_reply":"2026-01-01T07:18:20.050579Z"}},"outputs":[],"execution_count":null},{"id":"aad9cb3c","cell_type":"code","source":"# Save final model\ntorch.save({\n    'model_state_dict': model.state_dict(),\n    'config': CONFIG,\n    'train_losses': train_losses,\n    'val_losses': val_losses,\n    'best_val_loss': best_val_loss,\n}, 'final_model.pt')\n\nprint(\"✓ Model saved: final_model.pt\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T07:18:21.645166Z","iopub.execute_input":"2026-01-01T07:18:21.645461Z","iopub.status.idle":"2026-01-01T07:18:21.666573Z","shell.execute_reply.started":"2026-01-01T07:18:21.645436Z","shell.execute_reply":"2026-01-01T07:18:21.665829Z"}},"outputs":[],"execution_count":null}]}