{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":106809,"databundleVersionId":13056355,"sourceType":"competition"},{"sourceId":13873823,"sourceType":"datasetVersion","datasetId":8839346},{"sourceId":618527,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":465112,"modelId":480951}],"dockerImageVersionId":31153,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Start of Phoneme Retrieval Implementation","metadata":{}},{"cell_type":"code","source":"!pip install redis\n!pip install faiss-cpu\n!pip install fastdtw","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T20:33:40.356741Z","iopub.execute_input":"2025-12-07T20:33:40.357548Z","iopub.status.idle":"2025-12-07T20:34:04.817345Z","shell.execute_reply.started":"2025-12-07T20:33:40.357521Z","shell.execute_reply":"2025-12-07T20:34:04.816607Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# RETRIEVAL-AUGMENTED PHONEME-TO-TEXT SYSTEM\n# Neural signal clustering → Retrieve ground truth phonemes → LLM correction\n# ============================================================================\n\n# ============================================================================\n# CELL 1: IMPORTS AND SETUP\n# ============================================================================\n\nimport os\nimport h5py\nimport torch\nimport torch.nn as nn\nimport numpy as np\nimport pandas as pd\nimport faiss\nfrom omegaconf import OmegaConf\nfrom tqdm.notebook import tqdm\nfrom scipy.spatial.distance import euclidean\nfrom fastdtw import fastdtw\nimport pickle\nfrom typing import List, Dict, Tuple, Optional\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# Configuration\nCOMPETITION_INPUT = '/kaggle/input/brain-to-text-25'\nDATA_DIR = f'{COMPETITION_INPUT}/t15_copyTask_neuralData/hdf5_data_final'\nMODEL_PATH = f'{COMPETITION_INPUT}/t15_pretrained_rnn_baseline/t15_pretrained_rnn_baseline'\nOUTPUT_DIR = '/kaggle/working'\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f'✓ Device: {device}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T20:34:04.818903Z","iopub.execute_input":"2025-12-07T20:34:04.819214Z","iopub.status.idle":"2025-12-07T20:34:07.604447Z","shell.execute_reply.started":"2025-12-07T20:34:04.819183Z","shell.execute_reply":"2025-12-07T20:34:07.603677Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\n\n# Load one trial\ndata_file = \"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final/t15.2023.08.11/data_train.hdf5\"\nwith h5py.File(data_file, 'r') as f:\n    trial_keys = list(f.keys())\n    trial = f[trial_keys[0]]\n    neural_data = trial['input_features'][:]  # (time, 512)\n    print(f\"Shape: {neural_data.shape}\")\n\n# Split TX and SBP\ntx = neural_data[:, :256]       # Threshold crossings\nsbp = neural_data[:, 256:]      # Spike band power\n\n# Plot\nfig, axes = plt.subplots(2, 1, figsize=(14, 10), sharex=True)\n\n# TX heatmap\nim0 = axes[0].imshow(tx.T, aspect='auto', cmap='viridis', interpolation='nearest')\naxes[0].set_ylabel('Channel')\naxes[0].set_title('Threshold Crossings')\nfor boundary in [64, 128, 192]:\n    axes[0].axhline(y=boundary, color='white', linestyle='--', alpha=0.5)\naxes[0].set_yticks([32, 96, 160, 224])\naxes[0].set_yticklabels(['Ventral 6v', 'Area 4', '55b', 'Dorsal 6v'])\nplt.colorbar(im0, ax=axes[0])\n\n# SBP heatmap\nim1 = axes[1].imshow(sbp.T, aspect='auto', cmap='viridis', interpolation='nearest')\naxes[1].set_ylabel('Channel')\naxes[1].set_xlabel('Time (20ms bins)')\naxes[1].set_title('Spike Band Power')\nfor boundary in [64, 128, 192]:\n    axes[1].axhline(y=boundary, color='white', linestyle='--', alpha=0.5)\naxes[1].set_yticks([32, 96, 160, 224])\naxes[1].set_yticklabels(['Ventral 6v', 'Area 4', '55b', 'Dorsal 6v'])\nplt.colorbar(im1, ax=axes[1])\n\nplt.suptitle(f'Trial: {trial_keys[0]}')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T20:36:48.122732Z","iopub.execute_input":"2025-12-07T20:36:48.123384Z","iopub.status.idle":"2025-12-07T20:36:48.965097Z","shell.execute_reply.started":"2025-12-07T20:36:48.123360Z","shell.execute_reply":"2025-12-07T20:36:48.964247Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 2: PHONEME VOCABULARY\n# ============================================================================\n\nPHONEMES = [\n    '<blank>',  # CTC blank token (index 0)\n    'AA', 'AE', 'AH', 'AO', 'AW', 'AY',  # Vowels\n    'B', 'CH', 'D', 'DH',  # Consonants\n    'EH', 'ER', 'EY',\n    'F', 'G', 'HH',\n    'IH', 'IY', 'JH', 'K', 'L', 'M', 'N', 'NG',\n    'OW', 'OY',\n    'P', 'R', 'S', 'SH',\n    'T', 'TH',\n    'UH', 'UW',\n    'V', 'W', 'Y', 'Z', 'ZH',\n    '|'  # Silence token (index 40)\n]\n\nphoneme_to_idx = {p: i for i, p in enumerate(PHONEMES)}\nidx_to_phoneme = {i: p for i, p in enumerate(PHONEMES)}\n\nprint(f\"✓ Loaded {len(PHONEMES)} phonemes\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T05:45:46.217951Z","iopub.execute_input":"2025-11-26T05:45:46.218447Z","iopub.status.idle":"2025-11-26T05:45:46.224374Z","shell.execute_reply.started":"2025-11-26T05:45:46.218423Z","shell.execute_reply":"2025-11-26T05:45:46.223759Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 3: HELPER FUNCTIONS\n# ============================================================================\n\ndef ctc_greedy_decode(logits, blank_idx=0):\n    \"\"\"Simple CTC greedy decoder\"\"\"\n    # Get most likely class at each timestep\n    predictions = torch.argmax(logits, dim=-1)\n    \n    # Remove blanks and consecutive duplicates\n    output = []\n    prev = None\n    for pred in predictions:\n        pred_item = pred.item()\n        if pred_item != blank_idx and pred_item != prev:\n            output.append(pred_item)\n        prev = pred_item\n    \n    return output\n\ndef phoneme_ids_to_string(phoneme_ids, idx_to_phoneme):\n    \"\"\"Convert phoneme IDs to string representation\"\"\"\n    phonemes = [idx_to_phoneme.get(idx, '<unk>') for idx in phoneme_ids]\n    # Remove blank and silence tokens\n    phonemes = [p for p in phonemes if p not in ['<blank>']]\n    return ' '.join(phonemes)\n\ndef phoneme_string_to_ids(phoneme_string, phoneme_to_idx):\n    \"\"\"Convert phoneme string back to IDs\"\"\"\n    phonemes = phoneme_string.split()\n    return [phoneme_to_idx.get(p, 0) for p in phonemes]\n\n# ============================================================================\n# CELL 4: MODEL DEFINITION\n# ============================================================================\n\nclass GRUDecoder(nn.Module):\n    \"\"\"GRU decoder with patching and day-specific transforms\"\"\"\n    \n    def __init__(self, n_days=45, n_input_features=512, n_units=768, \n                 n_layers=5, n_classes=41, patch_size=14, patch_stride=4,\n                 rnn_dropout=0.4, input_dropout=0.2):\n        super().__init__()\n        \n        self.n_days = n_days\n        self.n_input_features = n_input_features\n        self.n_units = n_units\n        self.n_layers = n_layers\n        self.n_classes = n_classes\n        self.patch_size = patch_size\n        self.patch_stride = patch_stride\n        \n        self.gru_input_size = patch_size * n_input_features\n        \n        self.day_weights = nn.ParameterList([\n            nn.Parameter(torch.eye(n_input_features)) for _ in range(n_days)\n        ])\n        self.day_biases = nn.ParameterList([\n            nn.Parameter(torch.zeros(1, n_input_features)) for _ in range(n_days)\n        ])\n        \n        self.h0 = nn.Parameter(torch.zeros(1, 1, n_units))\n        \n        self.gru = nn.GRU(\n            input_size=self.gru_input_size,\n            hidden_size=n_units,\n            num_layers=n_layers,\n            batch_first=True,\n            dropout=rnn_dropout if n_layers > 1 else 0,\n            bidirectional=False\n        )\n        \n        self.out = nn.Linear(n_units, n_classes)\n    \n    def create_patches(self, x):\n        batch_size, seq_len, features = x.shape\n        n_patches = (seq_len - self.patch_size) // self.patch_stride + 1\n        \n        patches = []\n        for i in range(n_patches):\n            start = i * self.patch_stride\n            end = start + self.patch_size\n            patch = x[:, start:end, :].reshape(batch_size, -1)\n            patches.append(patch)\n        \n        return torch.stack(patches, dim=1)\n    \n    def forward(self, x, day_idx=0):\n        batch_size = x.size(0)\n        \n        if day_idx < self.n_days:\n            x = torch.matmul(x, self.day_weights[day_idx].t()) + self.day_biases[day_idx]\n        \n        x_patched = self.create_patches(x)\n        h0 = self.h0.repeat(self.n_layers, batch_size, 1)\n        x_gru, _ = self.gru(x_patched, h0)\n        logits = self.out(x_gru)\n        \n        return logits\n\n# ============================================================================\n# CELL 5: LOAD MODEL\n# ============================================================================\n\ndef load_pretrained_model(model_path, device):\n    \"\"\"Load the pretrained RNN model\"\"\"\n    print(\"Loading pretrained model...\")\n    \n    model_args = OmegaConf.load(os.path.join(model_path, 'checkpoint/args.yaml'))\n    sessions = model_args['dataset']['sessions']\n    \n    model = GRUDecoder(\n        n_input_features=model_args['model']['n_input_features'],\n        n_units=model_args['model']['n_units'],\n        n_days=len(sessions),\n        n_classes=model_args['dataset']['n_classes'],\n        rnn_dropout=model_args['model']['rnn_dropout'],\n        input_dropout=model_args['model']['input_network']['input_layer_dropout'],\n        n_layers=model_args['model']['n_layers'],\n        patch_size=model_args['model']['patch_size'],\n        patch_stride=model_args['model']['patch_stride'],\n    )\n    \n    checkpoint = torch.load(\n        os.path.join(model_path, 'checkpoint/best_checkpoint'),\n        map_location=device,\n        weights_only=False\n    )\n    \n    state_dict = checkpoint['model_state_dict']\n    if any(key.startswith('_orig_mod.') for key in state_dict.keys()):\n        state_dict = {k.replace('_orig_mod.', ''): v for k, v in state_dict.items()}\n    \n    model.load_state_dict(state_dict)\n    model.to(device)\n    model.eval()\n    \n    session_to_day = {session: idx for idx, session in enumerate(sessions)}\n    \n    print(f\"✓ Model loaded: {len(sessions)} sessions\")\n    return model, session_to_day, model_args\n\nmodel, session_to_day, model_args = load_pretrained_model(MODEL_PATH, device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T05:45:48.763666Z","iopub.execute_input":"2025-11-26T05:45:48.763975Z","iopub.status.idle":"2025-11-26T05:45:55.093343Z","shell.execute_reply.started":"2025-11-26T05:45:48.763952Z","shell.execute_reply":"2025-11-26T05:45:55.092644Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 6: EMBEDDING EXTRACTION\n# ============================================================================\n\ndef extract_neural_embedding(model, neural_data, day_idx, pooling='mean', device='cuda'):\n    \"\"\"\n    Extract embedding from neural signals (for clustering/retrieval)\n    \n    Returns:\n        embedding: (768,) - fixed-length vector for similarity search\n        gru_output: (n_patches, 768) - full sequence for optional DTW\n    \"\"\"\n    model.eval()\n    \n    with torch.no_grad():\n        x = torch.FloatTensor(neural_data).unsqueeze(0).to(device)\n        batch_size = x.size(0)\n        \n        # Day-specific transform\n        if day_idx < model.n_days:\n            x = torch.matmul(x, model.day_weights[day_idx].t()) + model.day_biases[day_idx]\n        \n        # Patch and process\n        x_patched = model.create_patches(x)\n        h0 = model.h0.repeat(model.n_layers, batch_size, 1)\n        gru_output, _ = model.gru(x_patched, h0)\n        \n        # Pool for fixed-length embedding\n        if pooling == 'mean':\n            embedding = gru_output.mean(dim=1).squeeze(0)\n        elif pooling == 'max':\n            embedding = gru_output.max(dim=1)[0].squeeze(0)\n        elif pooling == 'last':\n            embedding = gru_output[:, -1, :].squeeze(0)\n        else:\n            raise ValueError(f\"Unknown pooling: {pooling}\")\n    \n    return embedding.cpu().numpy(), gru_output.squeeze(0).cpu().numpy()\n\ndef predict_phonemes_from_neural(model, neural_data, day_idx, device='cuda'):\n    \"\"\"\n    Predict phoneme sequence from neural data using the RNN model\n    \n    Returns:\n        phoneme_ids: List of phoneme class IDs\n        phoneme_string: Space-separated phoneme string\n    \"\"\"\n    model.eval()\n    \n    with torch.no_grad():\n        x = torch.FloatTensor(neural_data).unsqueeze(0).to(device)\n        \n        # Forward pass\n        logits = model(x, day_idx)\n        logits = logits.squeeze(0)  # (n_patches, 41)\n        \n        # CTC decode\n        phoneme_ids = ctc_greedy_decode(logits, blank_idx=0)\n        phoneme_string = phoneme_ids_to_string(phoneme_ids, idx_to_phoneme)\n    \n    return phoneme_ids, phoneme_string\n\n# ============================================================================\n# CELL 7: BUILD TRAINING DATABASE (RUN ONCE)\n# CORRECT DATABASE BUILD - USING seq_class_ids AS PHONEMES\n# ============================================================================\n\ndef build_database_correct(model, data_dir, model_args, pooling='mean', device='cuda'):\n    \"\"\"\n    Build database with CORRECT phoneme extraction from seq_class_ids\n    \"\"\"\n    sessions = model_args['dataset']['sessions']\n    session_to_day = {s: i for i, s in enumerate(sessions)}\n    \n    available = []\n    for session in sessions:\n        path = os.path.join(data_dir, session, 'data_train.hdf5')\n        if os.path.exists(path):\n            available.append(session)\n    \n    print(f\"\\n{'='*70}\")\n    print(f\"Building Training Database (CORRECT)\")\n    print(f\"{'='*70}\")\n    print(f\"Sessions: {len(available)}/{len(sessions)}\")\n    print(f\"Using: seq_class_ids as phoneme IDs\")\n    print(f\"Pooling: {pooling}\\n\")\n    \n    embeddings_list = []\n    metadata_list = []\n    gru_outputs_list = []\n    \n    for session in tqdm(available, desc=\"Sessions\"):\n        day_idx = session_to_day[session]\n        path = os.path.join(data_dir, session, 'data_train.hdf5')\n        \n        with h5py.File(path, 'r') as f:\n            trials = sorted(list(f.keys()))\n            \n            for trial_key in trials:\n                try:\n                    trial = f[trial_key]\n                    \n                    # Neural data\n                    neural_data = trial['input_features'][:]\n                    \n                    # GROUND TRUTH PHONEMES (from seq_class_ids)\n                    seq_class_ids = trial['seq_class_ids'][:]\n                    gt_phoneme_ids = seq_class_ids[seq_class_ids > 0]  # Remove padding\n                    \n                    # Convert to phoneme string\n                    gt_phoneme_string = phoneme_ids_to_string(gt_phoneme_ids, idx_to_phoneme)\n                    \n                    if not gt_phoneme_string.strip():\n                        continue\n                    \n                    # Ground truth sentence (from attributes or transcription)\n                    gt_sentence = trial.attrs.get('sentence_label', None)\n                    if gt_sentence is None or gt_sentence == '':\n                        # Fallback to transcription\n                        transcription = trial['transcription'][:]\n                        valid = transcription[transcription > 0]\n                        gt_sentence = ''.join([chr(int(c)) for c in valid if 32 <= c <= 126])\n                    \n                    # Extract neural embedding\n                    embedding, gru_output = extract_neural_embedding(\n                        model, neural_data, day_idx, pooling, device\n                    )\n                    \n                    embeddings_list.append(embedding)\n                    gru_outputs_list.append(gru_output)\n                    metadata_list.append({\n                        'session': session,\n                        'day_idx': day_idx,\n                        'trial_key': trial_key,\n                        'ground_truth_phonemes': gt_phoneme_string,  # PHONEME STRING\n                        'ground_truth_phoneme_ids': gt_phoneme_ids.tolist(),  # PHONEME IDs\n                        'ground_truth_text': gt_sentence,  # TEXT\n                        'seq_len': len(neural_data),\n                        'n_patches': gru_output.shape[0],\n                        'db_index': len(metadata_list)\n                    })\n                    \n                except Exception as e:\n                    print(f\"Error: {session}/{trial_key}: {e}\")\n    \n    if len(embeddings_list) == 0:\n        print(\"ERROR: No valid trials found!\")\n        return None, None, None\n    \n    embeddings = np.vstack(embeddings_list)\n    \n    print(f\"\\n{'='*70}\")\n    print(f\"✓ Database Complete\")\n    print(f\"{'='*70}\")\n    print(f\"Valid trials: {len(embeddings):,}\")\n    print(f\"Embedding dim: {embeddings.shape[1]}\")\n    \n    # Show samples\n    print(f\"\\nSample entries:\")\n    for i in range(min(3, len(metadata_list))):\n        print(f\"\\n{i+1}.\")\n        print(f\"  Phonemes: {metadata_list[i]['ground_truth_phonemes'][:60]}...\")\n        print(f\"  Text: {metadata_list[i]['ground_truth_text'][:60]}...\")\n    \n    return embeddings, metadata_list, gru_outputs_list\n\ndef save_training_database(embeddings, metadata, gru_outputs, session_to_day, \n                          pooling, output_dir):\n    \"\"\"Save training database\"\"\"\n    print(\"\\nSaving database...\")\n    \n    # Main database\n    db_path = os.path.join(output_dir, 'training_phoneme_db.pkl')\n    db = {\n        'embeddings': embeddings,\n        'metadata': metadata,\n        'session_to_day': session_to_day,\n        'pooling': pooling\n    }\n    with open(db_path, 'wb') as f:\n        pickle.dump(db, f)\n    \n    print(f\"✓ Main DB: {db_path} ({os.path.getsize(db_path)/1e6:.1f} MB)\")\n    \n    # GRU outputs (optional, for DTW)\n    gru_path = os.path.join(output_dir, 'training_phoneme_db_gru.pkl')\n    with open(gru_path, 'wb') as f:\n        pickle.dump(gru_outputs, f)\n    \n    print(f\"✓ GRU DB: {gru_path} ({os.path.getsize(gru_path)/1e6:.1f} MB)\")\n    \n    return db_path, gru_path","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# BUILD DATABASE\nembeddings, metadata, gru_outputs = build_database_correct(\n    model=model,\n    data_dir=DATA_DIR,\n    model_args=model_args,\n    pooling='mean',\n    device=device\n)\n\n# Save if successful\nif embeddings is not None:\n    db_path, gru_path = save_training_database(\n        embeddings=embeddings,\n        metadata=metadata,\n        gru_outputs=gru_outputs,\n        session_to_day=session_to_day,\n        pooling='mean',\n        output_dir=OUTPUT_DIR\n    )\n    print(\"\\n✓ Database saved successfully!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T05:46:00.738958Z","iopub.execute_input":"2025-11-26T05:46:00.739239Z","iopub.status.idle":"2025-11-26T05:46:55.518129Z","shell.execute_reply.started":"2025-11-26T05:46:00.739218Z","shell.execute_reply":"2025-11-26T05:46:55.517164Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Save if successful\nif embeddings is not None:\n    db_path, gru_path = save_training_database(\n        embeddings=embeddings,\n        metadata=metadata,\n        gru_outputs=gru_outputs,\n        session_to_day=session_to_day,\n        pooling='mean',\n        output_dir=OUTPUT_DIR\n    )\n    print(\"\\n✓ Database saved successfully!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-05T21:30:47.403963Z","iopub.execute_input":"2025-11-05T21:30:47.404637Z","iopub.status.idle":"2025-11-05T21:30:57.471218Z","shell.execute_reply.started":"2025-11-05T21:30:47.404611Z","shell.execute_reply":"2025-11-05T21:30:57.469547Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 8: RETRIEVAL SYSTEM\n# ============================================================================\n\nclass NeuralPhonemeRetriever:\n    \"\"\"Retrieve training phonemes based on neural signal similarity\"\"\"\n    \n    def __init__(self, model, device='cuda'):\n        self.model = model\n        self.device = device\n        self.embeddings = None\n        self.metadata = None\n        self.gru_outputs = None\n        self.session_to_day = None\n        self.pooling = None\n        self.index = None\n    \n    def load(self, db_path, load_gru=False):\n        \"\"\"Load pre-built database\"\"\"\n        print(f\"Loading database from {db_path}...\")\n        \n        with open(db_path, 'rb') as f:\n            db = pickle.load(f)\n        \n        self.embeddings = db['embeddings']\n        self.metadata = db['metadata']\n        self.session_to_day = db['session_to_day']\n        self.pooling = db['pooling']\n        \n        print(f\"  Loaded {len(self.embeddings):,} training examples\")\n        \n        if load_gru:\n            gru_path = db_path.replace('.pkl', '_gru.pkl')\n            if os.path.exists(gru_path):\n                with open(gru_path, 'rb') as f:\n                    self.gru_outputs = pickle.load(f)\n                print(f\"  Loaded GRU outputs for DTW\")\n        \n        self._build_index()\n        print(f\"✓ Retriever ready\")\n    \n    def _build_index(self):\n        \"\"\"Build FAISS index\"\"\"\n        embeddings_norm = self.embeddings.astype('float32').copy()\n        faiss.normalize_L2(embeddings_norm)\n        \n        self.index = faiss.IndexFlatIP(embeddings_norm.shape[1])\n        \n        try:\n            if torch.cuda.is_available() and hasattr(faiss, 'StandardGpuResources'):\n                res = faiss.StandardGpuResources()\n                self.index = faiss.index_cpu_to_gpu(res, 0, self.index)\n                print(\"  Using GPU FAISS\")\n        except:\n            print(\"  Using CPU FAISS\")\n        \n        self.index.add(embeddings_norm)\n    \n    def retrieve_similar_phonemes(self, test_neural, test_session, k=5, use_dtw=False):\n        \"\"\"\n        Retrieve k training examples with similar neural signals\n        \n        Returns ground truth phonemes from those examples as \"hints\"\n        \"\"\"\n        day_idx = self.session_to_day.get(test_session, 0)\n        \n        # Extract test neural embedding\n        test_emb, test_gru = extract_neural_embedding(\n            self.model, test_neural, day_idx, self.pooling, self.device\n        )\n        \n        # Search\n        query_norm = test_emb.reshape(1, -1).astype('float32')\n        faiss.normalize_L2(query_norm)\n        \n        k_search = k * 3 if use_dtw else k\n        sims, indices = self.index.search(query_norm, k_search)\n        \n        # Collect results\n        results = []\n        for sim, idx in zip(sims[0], indices[0]):\n            result = self.metadata[idx].copy()\n            result['similarity'] = float(sim)\n            results.append(result)\n        \n        # DTW refinement (optional)\n        if use_dtw and self.gru_outputs is not None:\n            results = self._refine_dtw(test_gru, results, k)\n        else:\n            results = results[:k]\n        \n        return results\n    \n    def _refine_dtw(self, query_gru, candidates, k):\n        \"\"\"Refine with DTW\"\"\"\n        scores = []\n        for c in candidates:\n            cand_gru = self.gru_outputs[c['db_index']]\n            dist, _ = fastdtw(query_gru, cand_gru, dist=euclidean)\n            scores.append(dist)\n        \n        sorted_idx = np.argsort(scores)\n        results = [candidates[i] for i in sorted_idx[:k]]\n        \n        for i, r in enumerate(results):\n            r['dtw_distance'] = scores[sorted_idx[i]]\n        \n        return results\n\n# ============================================================================\n# CELL 9: LLM PROMPT CREATION\n# ============================================================================\n\ndef create_llm_prompt(target_phonemes, retrieved_examples, max_hints=5):\n    \"\"\"\n    Create prompt with:\n    - Target phonemes (RNN prediction, may have errors)\n    - Retrieved phoneme hints (ground truth from similar neural patterns)\n    - Retrieved text hints (for additional context)\n    \"\"\"\n    \n    prompt = f\"\"\"You are a phoneme-to-text decoder. You will receive:\n1. Target phoneme sequence from neural decoder (may contain errors)\n2. Ground truth phonemes and text from training examples with similar neural patterns\n\nYour task: Use the similar examples to correct any errors in the target phonemes and generate accurate text.\n\nTARGET PHONEMES (from neural decoder, may contain errors):\n{target_phonemes}\n\nSIMILAR TRAINING EXAMPLES (hints based on neural similarity):\n\"\"\"\n    \n    for i, ex in enumerate(retrieved_examples[:max_hints], 1):\n        hint_phonemes = ex['ground_truth_phonemes']\n        hint_text = ex['ground_truth_text']\n        similarity = ex['similarity']\n        \n        prompt += f\"\"\"\nExample {i} (neural similarity: {similarity:.3f}):\nPhonemes: {hint_phonemes}\nText: \"{hint_text}\"\n\"\"\"\n    \n    prompt += f\"\"\"\nBased on the target phonemes and the hints above, generate the most accurate text:\n\nFINAL TEXT:\"\"\"\n    \n    return prompt\n\n# Test with sample data\nprint(\"\\nExample prompt structure:\")\nprint(\"=\"*70)\nsample_target = \"AA R Y UW S T IH L DH EH R\"\nsample_hints = [\n    {\n        'ground_truth_phonemes': 'AA R Y UW AH T DH EH R',\n        'ground_truth_text': 'Are you out there?',\n        'similarity': 0.95\n    },\n    {\n        'ground_truth_phonemes': 'HH AW AA R Y UW',\n        'ground_truth_text': 'How are you?',\n        'similarity': 0.92\n    }\n]\n\nsample_prompt = create_llm_prompt(sample_target, sample_hints, max_hints=2)\nprint(sample_prompt)\nprint(\"=\"*70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T05:47:06.867130Z","iopub.execute_input":"2025-11-26T05:47:06.867416Z","iopub.status.idle":"2025-11-26T05:47:06.883031Z","shell.execute_reply.started":"2025-11-26T05:47:06.867394Z","shell.execute_reply":"2025-11-26T05:47:06.882254Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# RUN COMPLETE PIPELINE\n# ============================================================================\n\nDB_LOC = \"/kaggle/input/training-phoneme-db-pkl\"\n\n# Load retriever\nretriever = NeuralPhonemeRetriever(model=model, device=device)\nretriever.load(\n    db_path=os.path.join(DB_LOC, 'training_phoneme_db.pkl'),\n    load_gru=False\n)\n\n# Run inference with CORRECT prompt\ntest_file = os.path.join(DATA_DIR, 't15.2023.11.19', 'data_test.hdf5')\nif not os.path.exists(test_file):\n    test_file = os.path.join(DATA_DIR, 't15.2023.11.19', 'data_val.hdf5')\n\nprint(f\"\\nRunning inference on: {test_file}\\n\")\n\nresults = []\nday_idx = retriever.session_to_day['t15.2023.11.19']\n\nwith h5py.File(test_file, 'r') as f:\n    trials = sorted(list(f.keys()))[:5]  # Test on first 5\n    \n    for trial_key in tqdm(trials, desc=\"Processing\"):\n        neural_data = f[trial_key]['input_features'][:]\n        \n        # Predict target phonemes\n        target_ids, target_phonemes = predict_phonemes_from_neural(\n            model, neural_data, day_idx, device\n        )\n        \n        # Retrieve similar examples\n        retrieved = retriever.retrieve_similar_phonemes(\n            test_neural=neural_data,\n            test_session='t15.2023.11.19',\n            k=5,\n            use_dtw=False\n        )\n        \n        # Create prompt\n        prompt = create_llm_prompt(\n            target_phonemes=target_phonemes,\n            retrieved_examples=retrieved,\n            max_hints=5\n        )\n        \n        results.append({\n            'trial_key': trial_key,\n            'target_phonemes': target_phonemes,\n            'retrieved': retrieved,\n            'prompt': prompt\n        })\n\n# Show example\nprint(\"\\n\" + \"=\"*70)\nprint(\"FINAL PROMPT EXAMPLE\")\nprint(\"=\"*70 + \"\\n\")\nprint(results[0]['prompt'])\nprint(\"\\n\" + \"=\"*70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-26T05:50:34.698648Z","iopub.execute_input":"2025-11-26T05:50:34.699418Z","iopub.status.idle":"2025-11-26T05:50:35.865022Z","shell.execute_reply.started":"2025-11-26T05:50:34.699389Z","shell.execute_reply":"2025-11-26T05:50:35.864278Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# End of Phoneme Retrieval Implementation","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn.functional as F\nimport numpy as np\nfrom scipy.ndimage import gaussian_filter1d\nimport numpy as np\nimport h5py\nimport time\nimport re\n\ndef gauss_smooth(inputs, device, smooth_kernel_std=2, smooth_kernel_size=100,  padding='same'):\n    \"\"\"\n    Applies a 1D Gaussian smoothing operation with PyTorch to smooth the data along the time axis.\n    Args:\n        inputs (tensor : B x T x N): A 3D tensor with batch size B, time steps T, and number of features N.\n                                     Assumed to already be on the correct device (e.g., GPU).\n        kernelSD (float): Standard deviation of the Gaussian smoothing kernel.\n        padding (str): Padding mode, either 'same' or 'valid'.\n        device (str): Device to use for computation (e.g., 'cuda' or 'cpu').\n    Returns:\n        smoothed (tensor : B x T x N): A smoothed 3D tensor with batch size B, time steps T, and number of features N.\n    \"\"\"\n    # Get Gaussian kernel\n    inp = np.zeros(smooth_kernel_size, dtype=np.float32)\n    inp[smooth_kernel_size // 2] = 1\n    gaussKernel = gaussian_filter1d(inp, smooth_kernel_std)\n    validIdx = np.argwhere(gaussKernel > 0.01)\n    gaussKernel = gaussKernel[validIdx]\n    gaussKernel = np.squeeze(gaussKernel / np.sum(gaussKernel))\n\n    # Convert to tensor\n    gaussKernel = torch.tensor(gaussKernel, dtype=torch.float32, device=device)\n    gaussKernel = gaussKernel.view(1, 1, -1)  # [1, 1, kernel_size]\n\n    # Prepare convolution\n    B, T, C = inputs.shape\n    inputs = inputs.permute(0, 2, 1)  # [B, C, T]\n    gaussKernel = gaussKernel.repeat(C, 1, 1)  # [C, 1, kernel_size]\n\n    # Perform convolution\n    smoothed = F.conv1d(inputs, gaussKernel, padding=padding, groups=C)\n    return smoothed.permute(0, 2, 1)  # [B, T, C]\n\n\nLOGIT_TO_PHONEME = [\n    'BLANK',\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    ' | ',\n]\n\ndef _extract_transcription(input):\n    endIdx = np.argwhere(input == 0)[0, 0]\n    trans = ''\n    for c in range(endIdx):\n        trans += chr(input[c])\n    return trans\n\ndef load_h5py_file(file_path, b2txt_csv_df):\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        'corpus': [],\n    }\n    # Open the hdf5 file for that day\n    with h5py.File(file_path, 'r') as f:\n\n        keys = list(f.keys())\n\n        # For each trial in the selected trials in that day\n        for key in keys:\n            g = f[key]\n\n            neural_features = g['input_features'][:]\n            n_time_steps = g.attrs['n_time_steps']\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            session = g.attrs['session']\n            block_num = g.attrs['block_num']\n            trial_num = g.attrs['trial_num']\n\n            # match this trial up with the csv to get the corpus name\n            year, month, day = session.split('.')[1:]\n            date = f'{year}-{month}-{day}'\n            row = b2txt_csv_df[(b2txt_csv_df['Date'] == date) & (b2txt_csv_df['Block number'] == block_num)]\n            corpus_name = row['Corpus'].values[0]\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            data['corpus'].append(corpus_name)\n    return data\n\ndef rearrange_speech_logits_pt(logits):\n    # original order is [BLANK, phonemes..., SIL]\n    # rearrange so the order is [BLANK, SIL, phonemes...]\n    logits = np.concatenate((logits[:, :, 0:1], logits[:, :, -1:], logits[:, :, 1:-1]), axis=-1)\n    return logits\n\n# single decoding step function.\n# smooths data and puts it through the model.\ndef runSingleDecodingStep(x, input_layer, model, model_args, device):\n\n    # Use autocast for efficiency\n    with torch.autocast(device_type = \"cuda\", enabled = model_args['use_amp'], dtype = torch.bfloat16):\n\n        x = gauss_smooth(\n            inputs = x, \n            device = device,\n            smooth_kernel_std = model_args['dataset']['data_transforms']['smooth_kernel_std'],\n            smooth_kernel_size = model_args['dataset']['data_transforms']['smooth_kernel_size'],\n            padding = 'valid',\n        )\n\n        with torch.no_grad():\n            logits, _ = model(\n                x = x,\n                day_idx = torch.tensor([input_layer], device=device),\n                states = None, # no initial states\n                return_state = True,\n            )\n\n    # convert logits from bfloat16 to float32\n    logits = logits.float().cpu().numpy()\n\n    # # original order is [BLANK, phonemes..., SIL]\n    # # rearrange so the order is [BLANK, SIL, phonemes...]\n    # logits = rearrange_speech_logits_pt(logits)\n\n    return logits\n\ndef remove_punctuation(sentence):\n    # Remove punctuation\n    sentence = re.sub(r'[^a-zA-Z\\- \\']', '', sentence)\n    sentence = sentence.replace('- ', ' ').lower()\n    sentence = sentence.replace('--', '').lower()\n    sentence = sentence.replace(\" '\", \"'\").lower()\n\n    sentence = sentence.strip()\n    sentence = ' '.join([word for word in sentence.split() if word != ''])\n\n    return sentence\n\ndef get_current_redis_time_ms(redis_conn):\n    t = redis_conn.time()\n    return int(t[0]*1000 + t[1]/1000)\n\n\n######### language model helper functions ##########\n\ndef reset_remote_language_model(\n        r,\n        remote_lm_done_resetting_lastEntrySeen,\n    ):\n    \n    r.xadd('remote_lm_reset', {'done': 0})\n    time.sleep(0.001)\n    # print('Resetting remote language model before continuing...')\n    remote_lm_done_resetting = []\n    while len(remote_lm_done_resetting) == 0:\n        remote_lm_done_resetting = r.xread(\n            {'remote_lm_done_resetting': remote_lm_done_resetting_lastEntrySeen},\n            count=1,\n            block=10000,\n        )\n        if len(remote_lm_done_resetting) == 0:\n            print(f'Still waiting for remote lm reset from ts {remote_lm_done_resetting_lastEntrySeen}...')\n    for entry_id, entry_data in remote_lm_done_resetting[0][1]:\n        remote_lm_done_resetting_lastEntrySeen = entry_id\n        # print('Remote language model reset.')\n\n    return remote_lm_done_resetting_lastEntrySeen\n\n\ndef update_remote_lm_params(\n        r,\n        remote_lm_done_updating_lastEntrySeen,\n        acoustic_scale=0.35,\n        blank_penalty=90.0,\n        alpha=0.55,\n    ):\n    \n    # update remote lm params\n    entry_dict = {\n        # 'max_active': max_active,\n        # 'min_active': min_active,\n        # 'beam': beam,\n        # 'lattice_beam': lattice_beam,\n        'acoustic_scale': acoustic_scale,\n        # 'ctc_blank_skip_threshold': ctc_blank_skip_threshold,\n        # 'length_penalty': length_penalty,\n        # 'nbest': nbest,\n        'blank_penalty': blank_penalty,\n        'alpha': alpha,\n        # 'do_opt': do_opt,\n        # 'rescore': rescore,\n        # 'top_candidates_to_augment': top_candidates_to_augment,\n        # 'score_penalty_percent': score_penalty_percent,\n        # 'specific_word_bias': specific_word_bias,\n    }\n\n    r.xadd('remote_lm_update_params', entry_dict)\n    time.sleep(0.001)\n    remote_lm_done_updating = []\n    while len(remote_lm_done_updating) == 0:\n        remote_lm_done_updating = r.xread(\n            {'remote_lm_done_updating_params': remote_lm_done_updating_lastEntrySeen},\n            block=10000,\n            count=1,\n        )\n        if len(remote_lm_done_updating) == 0:\n            print(f'Still waiting for remote lm to update parameters from ts {remote_lm_done_updating_lastEntrySeen}...')\n    for entry_id, entry_data in remote_lm_done_updating[0][1]:\n        remote_lm_done_updating_lastEntrySeen = entry_id\n        # print('Remote language model params updated.')\n\n    return remote_lm_done_updating_lastEntrySeen\n\n\ndef send_logits_to_remote_lm(\n        r,\n        remote_lm_input_stream,\n        remote_lm_output_partial_stream,\n        remote_lm_output_partial_lastEntrySeen,\n        logits,\n    ):\n    \n    # put logits into remote lm and get partial output\n    r.xadd(remote_lm_input_stream, {'logits': np.float32(logits).tobytes()})\n    remote_lm_output = []\n    while len(remote_lm_output) == 0:\n        remote_lm_output = r.xread(\n            {remote_lm_output_partial_stream: remote_lm_output_partial_lastEntrySeen},\n            block=10000,\n            count=1,\n        )\n        if len(remote_lm_output) == 0:\n            print(f'Still waiting for remote lm partial output from ts {remote_lm_output_partial_lastEntrySeen}...')\n    for entry_id, entry_data in remote_lm_output[0][1]:\n        remote_lm_output_partial_lastEntrySeen = entry_id\n        decoded = entry_data[b'lm_response_partial'].decode()\n\n    return remote_lm_output_partial_lastEntrySeen, decoded\n\n\ndef finalize_remote_lm(\n        r,\n        remote_lm_output_final_stream,\n        remote_lm_output_final_lastEntrySeen,\n    ):\n    \n    # finalize remote lm\n    r.xadd('remote_lm_finalize', {'done': 0})\n    time.sleep(0.005)\n    remote_lm_output = []\n    while len(remote_lm_output) == 0:\n        remote_lm_output = r.xread(\n            {remote_lm_output_final_stream: remote_lm_output_final_lastEntrySeen},\n            block=10000,\n            count=1,\n        )\n        if len(remote_lm_output) == 0:\n            print(f'Still waiting for remote lm final output from ts {remote_lm_output_final_lastEntrySeen}...')\n    # print('Received remote lm final output.')\n\n    for entry_id, entry_data in remote_lm_output[0][1]:\n        remote_lm_output_final_lastEntrySeen = entry_id\n\n        candidate_sentences = [str(c) for c in entry_data[b'scoring'].decode().split(';')[::5]]\n        candidate_acoustic_scores = [float(c) for c in entry_data[b'scoring'].decode().split(';')[1::5]]\n        candidate_ngram_scores = [float(c) for c in entry_data[b'scoring'].decode().split(';')[2::5]]\n        candidate_llm_scores = [float(c) for c in entry_data[b'scoring'].decode().split(';')[3::5]]\n        candidate_total_scores = [float(c) for c in entry_data[b'scoring'].decode().split(';')[4::5]]\n\n\n    # account for a weird edge case where there are no candidate sentences\n    if len(candidate_sentences) == 0 or len(candidate_total_scores) == 0:\n        print('No candidate sentences were received from the language model.')\n        candidate_sentences = ['']\n        candidate_acoustic_scores = [0]\n        candidate_ngram_scores = [0]\n        candidate_llm_scores = [0]\n        candidate_total_scores = [0]\n\n    else:\n        # sort candidate sentences by total score (higher is better)\n        sort_order = np.argsort(candidate_total_scores)[::-1]\n\n        candidate_sentences = [candidate_sentences[i] for i in sort_order]\n        candidate_acoustic_scores = [candidate_acoustic_scores[i] for i in sort_order]\n        candidate_ngram_scores = [candidate_ngram_scores[i] for i in sort_order]\n        candidate_llm_scores = [candidate_llm_scores[i] for i in sort_order]\n        candidate_total_scores = [candidate_total_scores[i] for i in sort_order]\n\n    # loop through candidates backwards and remove any duplicates\n    for i in range(len(candidate_sentences)-1, 0, -1):\n        if candidate_sentences[i] in candidate_sentences[:i]:\n            candidate_sentences.pop(i)\n            candidate_acoustic_scores.pop(i)\n            candidate_ngram_scores.pop(i)\n            candidate_llm_scores.pop(i)\n            candidate_total_scores.pop(i)\n\n    lm_out = {\n        'candidate_sentences': candidate_sentences,\n        'candidate_acoustic_scores': candidate_acoustic_scores,\n        'candidate_ngram_scores': candidate_ngram_scores,\n        'candidate_llm_scores': candidate_llm_scores,\n        'candidate_total_scores': candidate_total_scores,\n    }\n\n    return remote_lm_output_final_lastEntrySeen, lm_out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-04T20:14:13.390938Z","iopub.execute_input":"2025-11-04T20:14:13.391589Z","iopub.status.idle":"2025-11-04T20:14:20.847391Z","shell.execute_reply.started":"2025-11-04T20:14:13.391556Z","shell.execute_reply":"2025-11-04T20:14:20.846681Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch \nfrom torch import nn\n\nclass GRUDecoder(nn.Module):\n    '''\n    Defines the GRU decoder\n\n    This class combines day-specific input layers, a GRU, and an output classification layer\n    '''\n    def __init__(self,\n                 neural_dim,\n                 n_units,\n                 n_days,\n                 n_classes,\n                 rnn_dropout = 0.0,\n                 input_dropout = 0.0,\n                 n_layers = 5, \n                 patch_size = 0,\n                 patch_stride = 0,\n                 ):\n        '''\n        neural_dim  (int)      - number of channels in a single timestep (e.g. 512)\n        n_units     (int)      - number of hidden units in each recurrent layer - equal to the size of the hidden state\n        n_days      (int)      - number of days in the dataset\n        n_classes   (int)      - number of classes \n        rnn_dropout    (float) - percentage of units to droupout during training\n        input_dropout (float)  - percentage of input units to dropout during training\n        n_layers    (int)      - number of recurrent layers \n        patch_size  (int)      - the number of timesteps to concat on initial input layer - a value of 0 will disable this \"input concat\" step \n        patch_stride(int)      - the number of timesteps to stride over when concatenating initial input \n        '''\n        super(GRUDecoder, self).__init__()\n        \n        self.neural_dim = neural_dim\n        self.n_units = n_units\n        self.n_classes = n_classes\n        self.n_layers = n_layers \n        self.n_days = n_days\n\n        self.rnn_dropout = rnn_dropout\n        self.input_dropout = input_dropout\n        \n        self.patch_size = patch_size\n        self.patch_stride = patch_stride\n\n        # Parameters for the day-specific input layers\n        self.day_layer_activation = nn.Softsign() # basically a shallower tanh \n\n        # Set weights for day layers to be identity matrices so the model can learn its own day-specific transformations\n        self.day_weights = nn.ParameterList(\n            [nn.Parameter(torch.eye(self.neural_dim)) for _ in range(self.n_days)]\n        )\n        self.day_biases = nn.ParameterList(\n            [nn.Parameter(torch.zeros(1, self.neural_dim)) for _ in range(self.n_days)]\n        )\n\n        self.day_layer_dropout = nn.Dropout(input_dropout)\n        \n        self.input_size = self.neural_dim\n\n        # If we are using \"strided inputs\", then the input size of the first recurrent layer will actually be in_size * patch_size\n        if self.patch_size > 0:\n            self.input_size *= self.patch_size\n\n        self.gru = nn.GRU(\n            input_size = self.input_size,\n            hidden_size = self.n_units,\n            num_layers = self.n_layers,\n            dropout = self.rnn_dropout, \n            batch_first = True, # The first dim of our input is the batch dim\n            bidirectional = False,\n        )\n\n        # Set recurrent units to have orthogonal param init and input layers to have xavier init\n        for name, param in self.gru.named_parameters():\n            if \"weight_hh\" in name:\n                nn.init.orthogonal_(param)\n            if \"weight_ih\" in name:\n                nn.init.xavier_uniform_(param)\n\n        # Prediciton head. Weight init to xavier\n        self.out = nn.Linear(self.n_units, self.n_classes)\n        nn.init.xavier_uniform_(self.out.weight)\n\n        # Learnable initial hidden states\n        self.h0 = nn.Parameter(nn.init.xavier_uniform_(torch.zeros(1, 1, self.n_units)))\n\n    def forward(self, x, day_idx, states = None, return_state = False):\n        '''\n        x        (tensor)  - batch of examples (trials) of shape: (batch_size, time_series_length, neural_dim)\n        day_idx  (tensor)  - tensor which is a list of day indexs corresponding to the day of each example in the batch x. \n        '''\n\n        # Apply day-specific layer to (hopefully) project neural data from the different days to the same latent space\n        day_weights = torch.stack([self.day_weights[i] for i in day_idx], dim=0)\n        day_biases = torch.cat([self.day_biases[i] for i in day_idx], dim=0).unsqueeze(1)\n\n        x = torch.einsum(\"btd,bdk->btk\", x, day_weights) + day_biases\n        x = self.day_layer_activation(x)\n\n        # Apply dropout to the ouput of the day specific layer\n        if self.input_dropout > 0:\n            x = self.day_layer_dropout(x)\n\n        # (Optionally) Perform input concat operation\n        if self.patch_size > 0: \n  \n            x = x.unsqueeze(1)                      # [batches, 1, timesteps, feature_dim]\n            x = x.permute(0, 3, 1, 2)               # [batches, feature_dim, 1, timesteps]\n            \n            # Extract patches using unfold (sliding window)\n            x_unfold = x.unfold(3, self.patch_size, self.patch_stride)  # [batches, feature_dim, 1, num_patches, patch_size]\n            \n            # Remove dummy height dimension and rearrange dimensions\n            x_unfold = x_unfold.squeeze(2)           # [batches, feature_dum, num_patches, patch_size]\n            x_unfold = x_unfold.permute(0, 2, 3, 1)  # [batches, num_patches, patch_size, feature_dim]\n\n            # Flatten last two dimensions (patch_size and features)\n            x = x_unfold.reshape(x.size(0), x_unfold.size(1), -1) \n        \n        # Determine initial hidden states\n        if states is None:\n            states = self.h0.expand(self.n_layers, x.shape[0], self.n_units).contiguous()\n\n        # Pass input through RNN \n        output, hidden_states = self.gru(x, states)\n\n        # Compute logits\n        logits = self.out(output)\n        \n        if return_state:\n            return logits, hidden_states\n        \n        return logits\n        ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-04T20:14:20.848666Z","iopub.execute_input":"2025-11-04T20:14:20.849349Z","iopub.status.idle":"2025-11-04T20:14:20.862057Z","shell.execute_reply.started":"2025-11-04T20:14:20.849328Z","shell.execute_reply":"2025-11-04T20:14:20.861440Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport numpy as np\nimport pandas as pd\nimport redis\nfrom omegaconf import OmegaConf\nimport time\nfrom tqdm import tqdm\nimport editdistance\n# import argparse\n\n# # argument parser for command line arguments\n# parser = argparse.ArgumentParser(description='Evaluate a pretrained RNN model on the copy task dataset.')\n# parser.add_argument('--model_path', type=str, default='../data/t15_pretrained_rnn_baseline',\n#                     help='Path to the pretrained model directory (relative to the current working directory).')\n# parser.add_argument('--data_dir', type=str, default='../data/hdf5_data_final',\n#                     help='Path to the dataset directory (relative to the current working directory).')\n# parser.add_argument('--eval_type', type=str, default='test', choices=['val', 'test'],\n#                     help='Evaluation type: \"val\" for validation set, \"test\" for test set. '\n#                          'If \"test\", ground truth is not available.')\n# parser.add_argument('--csv_path', type=str, default='../data/t15_copyTaskData_description.csv',\n#                     help='Path to the CSV file with metadata about the dataset (relative to the current working directory).')\n# parser.add_argument('--gpu_number', type=int, default=1,\n#                     help='GPU number to use for RNN model inference. Set to -1 to use CPU.')\n# args = parser.parse_args()\n\n\nCOMPETITION_INPUT = '/kaggle/input/brain-to-text-25'\ndata_dir = f'{COMPETITION_INPUT}/t15_copyTask_neuralData/hdf5_data_final'\n\n# Pretrained baseline model (from brain-to-text-25 dataset)\nmodel_path = f'{COMPETITION_INPUT}/t15_pretrained_rnn_baseline/t15_pretrained_rnn_baseline'\n\n\n\n# define evaluation type\neval_type = 'val'  # can be 'val' or 'test'. if 'test', ground truth is not available\n\n# load csv file\nb2txt_csv_df = pd.read_csv('https://raw.githubusercontent.com/Neuroprosthetics-Lab/nejm-brain-to-text/refs/heads/main/data/t15_copyTaskData_description.csv')\n\n# load model args\nmodel_args = OmegaConf.load(os.path.join(model_path, 'checkpoint/args.yaml'))\n\nmodel_args['dataset']['sessions'] = model_args['dataset']['sessions'][:5]\n\n# set up gpu device\ngpu_number = 1\n\nif torch.cuda.is_available() and gpu_number >= 0:\n    if gpu_number > torch.cuda.device_count():\n        raise ValueError(f'GPU number {gpu_number} is out of range. Available GPUs: {torch.cuda.device_count()}')\n    device = f'cuda:{gpu_number}'\n    device = torch.device(device)\n    print(f'Using {device} for model inference.')\nelse:\n    if gpu_number >= 0:\n        print(f'GPU number {gpu_number} requested but not available.')\n    print('Using CPU for model inference.')\n    device = torch.device('cpu')\n\n# define model\nmodel = GRUDecoder(\n    neural_dim = model_args['model']['n_input_features'],\n    n_units = model_args['model']['n_units'], \n    n_days = len(model_args['dataset']['sessions']),\n    n_classes = model_args['dataset']['n_classes'],\n    rnn_dropout = model_args['model']['rnn_dropout'],\n    input_dropout = model_args['model']['input_network']['input_layer_dropout'],\n    n_layers = model_args['model']['n_layers'],\n    patch_size = model_args['model']['patch_size'],\n    patch_stride = model_args['model']['patch_stride'],\n)\n\n# load model weights\ncheckpoint = torch.load(os.path.join(model_path, 'checkpoint/best_checkpoint'), weights_only=False)\n# rename keys to not start with \"module.\" (happens if model was saved with DataParallel)\nfor key in list(checkpoint['model_state_dict'].keys()):\n    checkpoint['model_state_dict'][key.replace(\"module.\", \"\")] = checkpoint['model_state_dict'].pop(key)\n    checkpoint['model_state_dict'][key.replace(\"_orig_mod.\", \"\")] = checkpoint['model_state_dict'].pop(key)\nmodel.load_state_dict(checkpoint['model_state_dict'])  \n\n# add model to device\nmodel.to(device) \n\n# set model to eval mode\nmodel.eval()\n\n# load data for each session\ntest_data = {}\ntotal_test_trials = 0\nfor session in model_args['dataset']['sessions']:\n    files = [f for f in os.listdir(os.path.join(data_dir, session)) if f.endswith('.hdf5')]\n    if f'data_{eval_type}.hdf5' in files:\n        eval_file = os.path.join(data_dir, session, f'data_{eval_type}.hdf5')\n\n        data = load_h5py_file(eval_file, b2txt_csv_df)\n        test_data[session] = data\n\n        total_test_trials += len(test_data[session][\"neural_features\"])\n        print(f'Loaded {len(test_data[session][\"neural_features\"])} {eval_type} trials for session {session}.')\nprint(f'Total number of {eval_type} trials: {total_test_trials}')\nprint()\n\n\n# put neural data through the pretrained model to get phoneme predictions (logits)\nwith tqdm(total=total_test_trials, desc='Predicting phoneme sequences', unit='trial') as pbar:\n    for session, data in test_data.items():\n\n        data['logits'] = []\n        data['pred_seq'] = []\n        input_layer = model_args['dataset']['sessions'].index(session)\n        \n        for trial in range(len(data['neural_features'])):\n            # get neural input for the trial\n            neural_input = data['neural_features'][trial]\n\n            # add batch dimension\n            neural_input = np.expand_dims(neural_input, axis=0)\n\n            # convert to torch tensor\n            neural_input = torch.tensor(neural_input, device=device, dtype=torch.bfloat16)\n\n            # run decoding step\n            logits = runSingleDecodingStep(neural_input, input_layer, model, model_args, device)\n            data['logits'].append(logits)\n\n            pbar.update(1)\npbar.close()\n\n\n# convert logits to phoneme sequences and print them out\nfor session, data in test_data.items():\n    data['pred_seq'] = []\n    for trial in range(len(data['logits'])):\n        logits = data['logits'][trial][0]\n        pred_seq = np.argmax(logits, axis=-1)\n        # remove blanks (0)\n        pred_seq = [int(p) for p in pred_seq if p != 0]\n        # remove consecutive duplicates\n        pred_seq = [pred_seq[i] for i in range(len(pred_seq)) if i == 0 or pred_seq[i] != pred_seq[i-1]]\n        # convert to phonemes\n        pred_seq = [LOGIT_TO_PHONEME[p] for p in pred_seq]\n        # add to data\n        data['pred_seq'].append(pred_seq)\n\n        # print out the predicted sequences\n        block_num = data['block_num'][trial]\n        trial_num = data['trial_num'][trial]\n        print(f'Session: {session}, Block: {block_num}, Trial: {trial_num}')\n        if eval_type == 'val':\n            sentence_label = data['sentence_label'][trial]\n            true_seq = data['seq_class_ids'][trial][0:data['seq_len'][trial]]\n            true_seq = [LOGIT_TO_PHONEME[p] for p in true_seq]\n\n            print(f'Sentence label:      {sentence_label}')\n            print(f'True sequence:       {\" \".join(true_seq)}')\n        print(f'Predicted Sequence:  {\" \".join(pred_seq)}')\n        print()\n\n\n# language model inference via redis\n# make sure that the standalone language model is running on the localhost redis ip\n# see README.md for instructions on how to run the language model\nr = redis.Redis(host='localhost', port=6379, db=0)\nr.flushall()  # clear all streams in redis\n\n# define redis streams for the remote language model\nremote_lm_input_stream = 'remote_lm_input'\nremote_lm_output_partial_stream = 'remote_lm_output_partial'\nremote_lm_output_final_stream = 'remote_lm_output_final'\n\n# set timestamps for last entries seen in the redis streams\nremote_lm_output_partial_lastEntrySeen = get_current_redis_time_ms(r)\nremote_lm_output_final_lastEntrySeen = get_current_redis_time_ms(r)\nremote_lm_done_resetting_lastEntrySeen = get_current_redis_time_ms(r)\nremote_lm_done_finalizing_lastEntrySeen = get_current_redis_time_ms(r)\nremote_lm_done_updating_lastEntrySeen = get_current_redis_time_ms(r)\n\nlm_results = {\n    'session': [],\n    'block': [],\n    'trial': [],\n    'true_sentence': [],\n    'pred_sentence': [],\n}\n\n# loop through all trials and put logits into the remote language model to get text predictions\n# note: this takes ~15-20 minutes to run on the entire test split with the 5-gram LM + OPT rescoring (RTX 4090)\nwith tqdm(total=total_test_trials, desc='Running remote language model', unit='trial') as pbar:\n    for session in test_data.keys():\n        for trial in range(len(test_data[session]['logits'])):\n            # get trial logits and rearrange them for the LM\n            logits = rearrange_speech_logits_pt(test_data[session]['logits'][trial])[0]\n\n            # reset language model\n            remote_lm_done_resetting_lastEntrySeen = reset_remote_language_model(r, remote_lm_done_resetting_lastEntrySeen)\n            \n            '''\n            # update language model parameters\n            remote_lm_done_updating_lastEntrySeen = update_remote_lm_params(\n                r,\n                remote_lm_done_updating_lastEntrySeen,\n                acoustic_scale=0.35,\n                blank_penalty=90.0,\n                alpha=0.55,\n            )\n            '''\n\n            # put logits into LM\n            remote_lm_output_partial_lastEntrySeen, decoded = send_logits_to_remote_lm(\n                r,\n                remote_lm_input_stream,\n                remote_lm_output_partial_stream,\n                remote_lm_output_partial_lastEntrySeen,\n                logits,\n            )\n\n            # finalize remote LM\n            remote_lm_output_final_lastEntrySeen, lm_out = finalize_remote_lm(\n                r,\n                remote_lm_output_final_stream,\n                remote_lm_output_final_lastEntrySeen,\n            )\n\n            # get the best candidate sentence\n            best_candidate_sentence = lm_out['candidate_sentences'][0]\n\n            # store results\n            lm_results['session'].append(session)\n            lm_results['block'].append(test_data[session]['block_num'][trial])\n            lm_results['trial'].append(test_data[session]['trial_num'][trial])\n            if eval_type == 'val':\n                lm_results['true_sentence'].append(test_data[session]['sentence_label'][trial])\n            else:\n                lm_results['true_sentence'].append(None)\n            lm_results['pred_sentence'].append(best_candidate_sentence)\n\n            # update progress bar\n            pbar.update(1)\npbar.close()\n\n\n# if using the validation set, lets calculate the aggregate word error rate (WER)\nif eval_type == 'val':\n    total_true_length = 0\n    total_edit_distance = 0\n\n    lm_results['edit_distance'] = []\n    lm_results['num_words'] = []\n\n    for i in range(len(lm_results['pred_sentence'])):\n        true_sentence = remove_punctuation(lm_results['true_sentence'][i]).strip()\n        pred_sentence = remove_punctuation(lm_results['pred_sentence'][i]).strip()\n        ed = editdistance.eval(true_sentence.split(), pred_sentence.split())\n\n        total_true_length += len(true_sentence.split())\n        total_edit_distance += ed\n\n        lm_results['edit_distance'].append(ed)\n        lm_results['num_words'].append(len(true_sentence.split()))\n\n        print(f'{lm_results[\"session\"][i]} - Block {lm_results[\"block\"][i]}, Trial {lm_results[\"trial\"][i]}')\n        print(f'True sentence:       {true_sentence}')\n        print(f'Predicted sentence:  {pred_sentence}')\n        print(f'WER: {ed} / {100 * len(true_sentence.split())} = {ed / len(true_sentence.split()):.2f}%')\n        print()\n\n    print(f'Total true sentence length: {total_true_length}')\n    print(f'Total edit distance: {total_edit_distance}')\n    print(f'Aggregate Word Error Rate (WER): {100 * total_edit_distance / total_true_length:.2f}%')\n\n\n# write predicted sentences to a csv file. put a timestamp in the filename (YYYYMMDD_HHMMSS)\noutput_file = os.path.join(model_path, f'baseline_rnn_{eval_type}_predicted_sentences_{time.strftime(\"%Y%m%d_%H%M%S\")}.csv')\nids = [i for i in range(len(lm_results['pred_sentence']))]\ndf_out = pd.DataFrame({'id': ids, 'text': lm_results['pred_sentence']})\ndf_out.to_csv(output_file, index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T20:42:07.672620Z","iopub.execute_input":"2025-10-29T20:42:07.673120Z","iopub.status.idle":"2025-10-29T20:42:08.907272Z","shell.execute_reply.started":"2025-10-29T20:42:07.673098Z","shell.execute_reply":"2025-10-29T20:42:08.906120Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for x in checkpoint['model_state_dict'].keys():\n    print(f\"{x}: \", checkpoint['model_state_dict'][x].shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-04T20:14:44.914731Z","iopub.execute_input":"2025-11-04T20:14:44.915046Z","iopub.status.idle":"2025-11-04T20:14:44.921170Z","shell.execute_reply.started":"2025-11-04T20:14:44.915025Z","shell.execute_reply":"2025-11-04T20:14:44.920335Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## aa","metadata":{}},{"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    # Print only first 3 files from each directory\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n    if dirname == \"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final/t15.2023.11.19\":\n        break","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-11-04T19:59:29.899139Z","iopub.execute_input":"2025-11-04T19:59:29.899458Z","iopub.status.idle":"2025-11-04T19:59:32.569003Z","shell.execute_reply.started":"2025-11-04T19:59:29.899426Z","shell.execute_reply":"2025-11-04T19:59:32.568098Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import h5py\nimport os\n\n# Competition data (test set)\nCOMPETITION_INPUT = '/kaggle/input/brain-to-text-25'\nDATA_DIR = f'{COMPETITION_INPUT}/t15_copyTask_neuralData/hdf5_data_final'\n\n# Inspect HDF5 file structure\ndef inspect_hdf5(file_path):\n    \"\"\"Print the structure of an HDF5 file\"\"\"\n    print(f\"\\n📂 Inspecting: {file_path}\")\n    print(\"=\" * 60)\n    \n    with h5py.File(file_path, 'r') as f:\n        def print_structure(name, obj):\n            if isinstance(obj, h5py.Dataset):\n                print(f\"  Dataset: {name}\")\n                print(f\"    Shape: {obj.shape}\")\n                print(f\"    Dtype: {obj.dtype}\")\n        \n        print(\"\\nAvailable keys:\")\n        for key in f.keys():\n            print(f\"  - {key}\")\n        \n        print(\"\\nDetailed structure:\")\n        f.visititems(print_structure)\n\n# Check the first available session\nsessions = sorted(os.listdir(DATA_DIR))\nfor session in sessions[:2]:  # Check first 2 sessions\n    session_path = os.path.join(DATA_DIR, session)\n    \n    # Check train file\n    train_path = os.path.join(session_path, 'data_train.hdf5')\n    if os.path.exists(train_path):\n        inspect_hdf5(train_path)\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-04T20:00:52.872107Z","iopub.execute_input":"2025-11-04T20:00:52.872666Z","iopub.status.idle":"2025-11-04T20:00:55.608944Z","shell.execute_reply.started":"2025-11-04T20:00:52.872640Z","shell.execute_reply":"2025-11-04T20:00:55.608008Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Install dependencies, Import libraries, GPU Setting","metadata":{}},{"cell_type":"code","source":"!pip install -q h5py\n!pip install -q g2p-en\n!pip install -q editdistance\n!pip install -q pyyaml\n\nprint(\"✓ Dependencies installed\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-25T00:28:02.684519Z","iopub.execute_input":"2025-10-25T00:28:02.684867Z","iopub.status.idle":"2025-10-25T00:28:17.220145Z","shell.execute_reply.started":"2025-10-25T00:28:02.684847Z","shell.execute_reply":"2025-10-25T00:28:17.219166Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport sys\nimport glob\nimport json\nimport yaml\nimport numpy as np\nimport pandas as pd\nimport h5py\nfrom pathlib import Path\nfrom collections import defaultdict, Counter\nimport pickle\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\n# Verify GPU\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nif torch.cuda.is_available():\n    print(f\"GPU: {torch.cuda.get_device_name(0)}\")\n    print(f\"GPU Memory: {torch.cuda.get_device_properties(0).total_memory / 1024**3:.1f} GB\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-25T00:28:30.181084Z","iopub.execute_input":"2025-10-25T00:28:30.181713Z","iopub.status.idle":"2025-10-25T00:28:34.268578Z","shell.execute_reply.started":"2025-10-25T00:28:30.181686Z","shell.execute_reply":"2025-10-25T00:28:34.267523Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Path Setting","metadata":{}},{"cell_type":"code","source":"import os\n\n# Competition data (test set)\nCOMPETITION_INPUT = '/kaggle/input/brain-to-text-25'\nDATA_DIR = f'{COMPETITION_INPUT}/t15_copyTask_neuralData/hdf5_data_final'\n\n# Pretrained baseline model (from brain-to-text-25 dataset)\nBASELINE_DIR = f'{COMPETITION_INPUT}/t15_pretrained_rnn_baseline/t15_pretrained_rnn_baseline/checkpoint'\n\n# Verify paths\nprint(\"Checking paths...\")\nprint(f\"Competition data exists: {os.path.exists(COMPETITION_INPUT)}\")\nprint(f\"Data directory exists: {os.path.exists(DATA_DIR)}\")\nprint(f\"Baseline model exists: {os.path.exists(BASELINE_DIR)}\")\n\nif os.path.exists(BASELINE_DIR):\n    print(f\"\\n✓ Model found!\")\n    print(f\"  Location: {BASELINE_DIR}\")\n    print(f\"  Files: {os.listdir(BASELINE_DIR)}\")\nelse:\n    print(\"\\n⚠ Model not found at expected path\")\n    print(f\"  Expected: {BASELINE_DIR}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T20:58:02.941084Z","iopub.execute_input":"2025-10-29T20:58:02.941428Z","iopub.status.idle":"2025-10-29T20:58:02.953846Z","shell.execute_reply.started":"2025-10-29T20:58:02.941409Z","shell.execute_reply":"2025-10-29T20:58:02.952744Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Architecture Define","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\n\nclass GRUDecoder(nn.Module):\n    \"\"\"Actual baseline GRU decoder with patching and day-specific transform\"\"\"\n    \n    def __init__(self, n_days=45, n_input_features=512, n_units=768, \n                 n_layers=5, n_outputs=41, patch_size=14, patch_stride=4):\n        super().__init__()\n        \n        self.n_days = n_days\n        self.n_input_features = n_input_features\n        self.n_units = n_units\n        self.n_layers = n_layers\n        self.n_outputs = n_outputs\n        self.patch_size = patch_size\n        self.patch_stride = patch_stride\n        \n        # Calculate input size after patching\n        self.gru_input_size = patch_size * n_input_features\n        \n        # Per-day transformation (affine transform on features)\n        self.day_weights = nn.ParameterList([\n            nn.Parameter(torch.eye(n_input_features)) for _ in range(n_days)\n        ])\n        self.day_biases = nn.ParameterList([\n            nn.Parameter(torch.zeros(1, n_input_features)) for _ in range(n_days)\n        ])\n        \n        # Initial hidden state (unidirectional, so just n_layers)\n        self.h0 = nn.Parameter(torch.zeros(1, 1, n_units))\n        \n        # Unidirectional GRU\n        self.gru = nn.GRU(\n            input_size=self.gru_input_size,\n            hidden_size=n_units,\n            num_layers=n_layers,\n            batch_first=True,\n            bidirectional=False  # UNIDIRECTIONAL!\n        )\n        \n        # Output layer\n        self.out = nn.Linear(n_units, n_outputs)\n    \n    def create_patches(self, x):\n        \"\"\"Create overlapping patches from input sequence\"\"\"\n        # x: (batch, time, features)\n        batch_size, seq_len, features = x.shape\n        \n        # Calculate number of patches\n        n_patches = (seq_len - self.patch_size) // self.patch_stride + 1\n        \n        # Extract patches\n        patches = []\n        for i in range(n_patches):\n            start = i * self.patch_stride\n            end = start + self.patch_size\n            patch = x[:, start:end, :]  # (batch, patch_size, features)\n            patch_flat = patch.reshape(batch_size, -1)  # (batch, patch_size*features)\n            patches.append(patch_flat)\n        \n        # Stack patches: (batch, n_patches, patch_size*features)\n        patches = torch.stack(patches, dim=1)\n        return patches\n    \n    def forward(self, x, day_idx=0):\n        \"\"\"\n        Args:\n            x: (batch, time, features) - raw neural features\n            day_idx: which day's normalization to use (0-44)\n        \"\"\"\n        batch_size = x.size(0)\n        \n        # Apply day-specific linear transformation\n        if day_idx < self.n_days:\n            # x @ W^T + b\n            x = torch.matmul(x, self.day_weights[day_idx].t()) + self.day_biases[day_idx]\n        \n        # Create patches\n        x_patched = self.create_patches(x)\n        \n        # GRU\n        h0 = self.h0.repeat(self.n_layers, batch_size, 1)\n        x_gru, _ = self.gru(x_patched, h0)\n        \n        # Output\n        logits = self.out(x_gru)\n        return logits\n\nprint(\"✓ Correct model architecture defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-25T00:28:37.078941Z","iopub.execute_input":"2025-10-25T00:28:37.079931Z","iopub.status.idle":"2025-10-25T00:28:37.089340Z","shell.execute_reply.started":"2025-10-25T00:28:37.079898Z","shell.execute_reply":"2025-10-25T00:28:37.088459Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Load pretrained weights","metadata":{}},{"cell_type":"code","source":"import yaml\nimport torch\n\n# Load model config\nargs_path = f'{BASELINE_DIR}/args.yaml'\nwith open(args_path, 'r') as f:\n    args = yaml.safe_load(f)\n\nmodel_config = args['model']\nprint(\"Model config:\")\nprint(f\"  n_input_features: {model_config['n_input_features']}\")\nprint(f\"  n_units: {model_config['n_units']}\")\nprint(f\"  n_layers: {model_config['n_layers']}\")\nprint(f\"  bidirectional: {model_config['bidirectional']}\")\nprint(f\"  patch_size: {model_config['patch_size']}\")\nprint(f\"  patch_stride: {model_config['patch_stride']}\")\n\n# Initialize model with correct parameters\nmodel = GRUDecoder(\n    n_days=45,\n    n_input_features=model_config['n_input_features'],\n    n_units=model_config['n_units'],\n    n_layers=model_config['n_layers'],\n    n_outputs=args['dataset']['n_classes'],\n    patch_size=model_config['patch_size'],\n    patch_stride=model_config['patch_stride']\n).to(device)\n\n# Load checkpoint\ncheckpoint_path = f'{BASELINE_DIR}/best_checkpoint'\nprint(f\"\\nLoading checkpoint...\")\n\ncheckpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)\nstate_dict = checkpoint['model_state_dict']\n\n# Remove _orig_mod. prefix (from torch.compile)\nnew_state_dict = {}\nfor key, value in state_dict.items():\n    if key.startswith('_orig_mod.'):\n        new_key = key.replace('_orig_mod.', '')\n        new_state_dict[new_key] = value\n    else:\n        new_state_dict[key] = value\n\n# Load weights\nmodel.load_state_dict(new_state_dict, strict=True)\n\nmodel.eval()\nprint(f\"\\n✓ Model loaded successfully!\")\nprint(f\"Parameters: {sum(p.numel() for p in model.parameters()):,}\")\nprint(f\"Validation loss: {checkpoint.get('val_loss', 'N/A'):.6f}\")\nprint(f\"Validation PER: {checkpoint.get('val_PER', 'N/A'):.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-25T00:28:40.344862Z","iopub.execute_input":"2025-10-25T00:28:40.345127Z","iopub.status.idle":"2025-10-25T00:28:44.928918Z","shell.execute_reply.started":"2025-10-25T00:28:40.345102Z","shell.execute_reply":"2025-10-25T00:28:44.928108Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Utilities","metadata":{}},{"cell_type":"code","source":"# The baseline uses 40 phonemes + blank token\n# Standard English phoneme set (ARPAbet)\nPHONEMES = [\n    '<blank>',  # CTC blank token (index 0)\n    'AA', 'AE', 'AH', 'AO', 'AW', 'AY',  # Vowels\n    'B', 'CH', 'D', 'DH',  # Consonants\n    'EH', 'ER', 'EY',\n    'F', 'G', 'HH',\n    'IH', 'IY', 'JH', 'K', 'L', 'M', 'N', 'NG',\n    'OW', 'OY',\n    'P', 'R', 'S', 'SH',\n    'T', 'TH',\n    'UH', 'UW',\n    'V', 'W', 'Y', 'Z', 'ZH',\n    '<sil>'  # Silence token\n]\n\n# Create index mappings\nphoneme_to_idx = {p: i for i, p in enumerate(PHONEMES)}\nidx_to_phoneme = {i: p for i, p in enumerate(PHONEMES)}\n\nprint(f\"Loaded {len(PHONEMES)} phonemes\")\nprint(f\"Examples: {PHONEMES[:5]}...\")\n\n\nimport torch\n\ndef ctc_greedy_decode(logits, blank_idx=0):\n    \"\"\"Greedy CTC decoding - collapse repeats and remove blanks\"\"\"\n    # Get most likely phoneme at each timestep\n    predictions = torch.argmax(logits, dim=-1)  # Shape: (time_steps,)\n    \n    # Convert to list\n    pred_list = predictions.cpu().numpy().tolist()\n    \n    # Collapse consecutive duplicates and remove blanks\n    decoded = []\n    prev = -1\n    for p in pred_list:\n        if p != prev and p != blank_idx:\n            decoded.append(p)\n        prev = p\n    \n    return decoded\n\n# Cell 9: Simple decoder (no language model)\nclass SimplePhonemeToTextDecoder:\n    \"\"\"Simplified phoneme-to-text decoder - just concatenates phonemes\"\"\"\n    \n    def __init__(self):\n        print(\"Using simple phoneme decoder (no language model)\")\n    \n    def decode(self, phoneme_indices, idx_to_phoneme):\n        \"\"\"Decode phoneme sequence to text\"\"\"\n        # Convert indices to phoneme symbols\n        phonemes = [idx_to_phoneme.get(idx, '') for idx in phoneme_indices]\n        \n        # Filter out blanks and special tokens\n        phonemes = [p for p in phonemes if p not in ['<blank>', '<pad>', '']]\n        \n        # Join with spaces\n        return ' '.join(phonemes)\n\n# Initialize decoder\ntext_decoder = SimplePhonemeToTextDecoder()\nprint(\"✓ Simple decoder initialized\")\nprint(\"\\n⚠ NOTE: This decoder just outputs phoneme sequences\")\nprint(\"   For better results, you would need a phoneme-to-word language model\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-25T00:28:50.045244Z","iopub.execute_input":"2025-10-25T00:28:50.045521Z","iopub.status.idle":"2025-10-25T00:28:50.054023Z","shell.execute_reply.started":"2025-10-25T00:28:50.045502Z","shell.execute_reply":"2025-10-25T00:28:50.053251Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Load data","metadata":{}},{"cell_type":"code","source":"def load_test_data(data_dir, max_sessions=1):\n    \"\"\"Load test HDF5 files\"\"\"\n    test_data = []\n    \n    # Find all test files\n    test_files = glob.glob(f'{data_dir}/t15.*/data_test.hdf5')\n    print(f\"Found {len(test_files)} test files\")\n    \n    # Limit number of sessions if specified\n    if max_sessions:\n        test_files = test_files[:max_sessions]\n        print(f\"Loading only first {len(test_files)} session(s)\")\n    \n    for file_path in test_files:\n        session_name = os.path.basename(os.path.dirname(file_path))\n        print(f\"Loading {session_name}...\")\n        \n        with h5py.File(file_path, 'r') as f:\n            for trial_key in f.keys():\n                trial = f[trial_key]\n                \n                features = np.array(trial['input_features'])\n                \n                metadata = {\n                    'session_date': trial.attrs.get('session', session_name),\n                    'block_number': trial.attrs.get('block_num', 0),\n                    'trial_number': trial.attrs.get('trial_num', 0),\n                }\n                \n                test_data.append({\n                    'features': features,\n                    'metadata': metadata\n                })\n    \n    print(f\"\\nTotal loaded: {len(test_data)} test trials\")\n    return test_data\n\n# # Load just 1 session for quick testing\n# test_data = load_test_data(DATA_DIR, max_sessions=1)\n\n# if len(test_data) > 0:\n#     print(f\"\\nExample trial:\")\n#     print(f\"  Features shape: {test_data[0]['features'].shape}\")\n#     print(f\"  Metadata: {test_data[0]['metadata']}\")\n\n\ndef load_validation_data(data_dir, max_sessions=1):\n    \"\"\"Load validation HDF5 files (these have ground truth)\"\"\"\n    val_data = []\n    \n    # Find validation files\n    val_files = glob.glob(f'{data_dir}/t15.*/data_val.hdf5')\n    print(f\"Found {len(val_files)} validation files\")\n    \n    if max_sessions:\n        val_files = val_files[:max_sessions]\n    \n    for file_path in val_files:\n        session_name = os.path.basename(os.path.dirname(file_path))\n        print(f\"Loading {session_name}...\")\n        \n        with h5py.File(file_path, 'r') as f:\n            for trial_key in list(f.keys())[:20]:  # Just first 20 trials\n                trial = f[trial_key]\n                \n                # Extract features\n                features = np.array(trial['input_features'])\n                \n                # Extract ground truth sentence\n                sentence = trial.attrs.get('sentence_label', 'N/A')\n                \n                # Get phoneme sequence (THIS IS THE KEY CHANGE!)\n                phoneme_indices = np.array(trial['seq_class_ids'])\n                seq_len = trial.attrs.get('seq_len', len(phoneme_indices))\n                phoneme_indices = phoneme_indices[:seq_len]  # Trim to actual length\n                \n                # Convert phoneme indices to phoneme strings\n                phonemes_list = [idx_to_phoneme.get(idx, f'<unk{idx}>') for idx in phoneme_indices]\n                phonemes = ' '.join(phonemes_list)\n                \n                # Extract metadata\n                metadata = {\n                    'session_date': trial.attrs.get('session', session_name),\n                    'block_number': trial.attrs.get('block_num', 0),\n                    'trial_number': trial.attrs.get('trial_num', 0),\n                }\n                \n                val_data.append({\n                    'features': features,\n                    'metadata': metadata,\n                    'ground_truth_sentence': sentence,\n                    'ground_truth_phonemes': phonemes,\n                    'ground_truth_indices': phoneme_indices\n                })\n    \n    print(f\"\\nLoaded {len(val_data)} validation trials\")\n    return val_data\n\n\ndef run_inference_with_gt(model, val_data, text_decoder, device):\n    \"\"\"Run inference and compare with ground truth\"\"\"\n    model.eval()\n    results = []\n\n    print(f\"Running inference on {len(val_data)} validation trials...\")\n\n    with torch.no_grad():\n        for idx, trial in enumerate(val_data):\n            # Prepare input\n            features = torch.FloatTensor(trial['features']).unsqueeze(0).to(device)\n\n            # Forward pass\n            logits = model(features)\n            logits = logits.squeeze(0)\n\n            # CTC decode to phonemes\n            phoneme_indices = ctc_greedy_decode(logits, blank_idx=0)\n\n            # Convert phonemes to text\n            predicted_phonemes = ' '.join([idx_to_phoneme.get(idx, '') for idx in phoneme_indices])\n            predicted_text = text_decoder.decode(phoneme_indices, idx_to_phoneme)\n\n            # Store results\n            results.append({\n                'metadata': trial['metadata'],\n                'predicted_phonemes': predicted_phonemes,\n                'predicted_text': predicted_text,\n                'ground_truth_sentence': trial['ground_truth_sentence'],\n                'ground_truth_phonemes': trial['ground_truth_phonemes']\n            })\n\n    print(f\"✓ Completed inference on {len(results)} trials\")\n    return results","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-25T00:28:53.755019Z","iopub.execute_input":"2025-10-25T00:28:53.755273Z","iopub.status.idle":"2025-10-25T00:28:53.878876Z","shell.execute_reply.started":"2025-10-25T00:28:53.755256Z","shell.execute_reply":"2025-10-25T00:28:53.877845Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Test Inference","metadata":{}},{"cell_type":"code","source":"# Reload validation data with correct phoneme field\nval_data = load_validation_data(DATA_DIR, max_sessions=1)\n\n# Re-run inference\nval_results = run_inference_with_gt(model, val_data, text_decoder, device)\n\n# Print predictions vs ground truth\nprint(\"\\n\" + \"=\"*80)\nprint(\"PREDICTIONS vs GROUND TRUTH\")\nprint(\"=\"*80)\nfor i in range(len(val_results)):\n    result = val_results[i]\n    \n    print(f\"\\n{'='*80}\")\n    print(f\"Trial {i+1}:\")\n    print(f\"{'='*80}\")\n    \n    # Ground truth\n    print(f\"[📝 GROUND TRUTH] Sentence: {result['ground_truth_sentence']}\")\n    print(f\"[📝 GROUND TRUTH] Phonemes: {result['ground_truth_phonemes']}\")\n    \n    # Predictions\n    print(f\"[🤖 PREDICTION] Phonemes: {result['predicted_phonemes']}\")\n    print(f\"[🤖 PREDICTION] Text:     {result['predicted_text']}\")\n    \n    # Clean comparison\n    pred_clean = result['predicted_phonemes'].replace('<sil>', '').split()\n    gt_clean = result['ground_truth_phonemes'].replace('<sil>', '').split()\n    \n    print(f\"\\n📊 Stats:\")\n    print(f\"   Predicted phonemes: {len(pred_clean)}\")\n    print(f\"   Ground truth phonemes: {len(gt_clean)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-25T00:30:04.563327Z","iopub.execute_input":"2025-10-25T00:30:04.563566Z","iopub.status.idle":"2025-10-25T00:30:09.559697Z","shell.execute_reply.started":"2025-10-25T00:30:04.563551Z","shell.execute_reply":"2025-10-25T00:30:09.558805Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# End-to-End Brain-to-Text Decoding with Neural-BERT + Transformer\n\n**Competition:** Brain-to-Text '25  \n**Approach:** Direct neural signal → text with self-supervised pre-training\n\n## Architecture Overview\n\n```\nPhase 1: Neural-BERT Pre-training (Self-Supervised)\n├── Masked Neural Modeling\n├── Contrastive Learning\n└── Learn robust spike representations\n\nPhase 2: End-to-End Fine-tuning\n├── Neural Encoder (Pre-trained)\n├── Cross-Attention Bridge\n├── GPT-2 Decoder (Pre-trained)\n└── Direct text generation\n```\n\n**Key Innovation:** No intermediate phoneme representation - direct neural → text!","metadata":{}},{"cell_type":"markdown","source":"## 1. Setup & Dependencies","metadata":{}},{"cell_type":"code","source":"# Install dependencies\n!pip install -q transformers\n!pip install -q h5py\n!pip install -q editdistance\n!pip install -q einops\n\nprint(\"✓ Dependencies installed\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T18:16:54.988028Z","iopub.execute_input":"2025-10-28T18:16:54.988747Z","iopub.status.idle":"2025-10-28T18:17:09.116499Z","shell.execute_reply.started":"2025-10-28T18:16:54.988716Z","shell.execute_reply":"2025-10-28T18:17:09.115496Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport sys\nimport glob\nimport numpy as np\nimport pandas as pd\nimport h5py\nfrom pathlib import Path\nfrom collections import defaultdict\nfrom tqdm.auto import tqdm\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts\n\nfrom transformers import GPT2Tokenizer, GPT2LMHeadModel, GPT2Config\nimport editdistance\nfrom einops import rearrange, repeat\n\n# Set random seeds\ndef set_seed(seed=42):\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# Device setup\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"🖥️  Using device: {device}\")\nif torch.cuda.is_available():\n    print(f\"   GPU: {torch.cuda.get_device_name(0)}\")\n    print(f\"   Memory: {torch.cuda.get_device_properties(0).total_memory / 1024**3:.1f} GB\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T18:17:09.118275Z","iopub.execute_input":"2025-10-28T18:17:09.118534Z","iopub.status.idle":"2025-10-28T18:17:26.244475Z","shell.execute_reply.started":"2025-10-28T18:17:09.118512Z","shell.execute_reply":"2025-10-28T18:17:26.243690Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Path configuration\nCOMPETITION_INPUT = '/kaggle/input/brain-to-text-25'\nDATA_DIR = f'{COMPETITION_INPUT}/t15_copyTask_neuralData/hdf5_data_final'\n\n# Verify paths\nprint(\"📁 Data paths:\")\nprint(f\"   Competition: {os.path.exists(COMPETITION_INPUT)}\")\nprint(f\"   Data: {os.path.exists(DATA_DIR)}\")\nprint(f\"\\n   Found {len(os.listdir(DATA_DIR))} sessions\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T18:17:26.245290Z","iopub.execute_input":"2025-10-28T18:17:26.245785Z","iopub.status.idle":"2025-10-28T18:17:26.262645Z","shell.execute_reply.started":"2025-10-28T18:17:26.245766Z","shell.execute_reply":"2025-10-28T18:17:26.262001Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Data Loading & Preprocessing","metadata":{}},{"cell_type":"code","source":"class BCIDataset(Dataset):\n    \"\"\"Dataset for loading BCI neural data from HDF5 files\"\"\"\n    \n    def __init__(self, data_dir, split='train', max_len=None, vocab=None):\n        self.data_dir = data_dir\n        self.split = split\n        self.max_len = max_len\n        self.vocab = vocab  # Will be built if None\n        \n        # Load all sessions\n        self.samples = []\n        self.all_tokens = set()\n        self._load_data()\n        \n        # Build vocabulary if not provided\n        if self.vocab is None and self.split == 'train':\n            self._build_vocab()\n        \n        print(f\"✓ Loaded {len(self.samples)} {split} samples\")\n        if len(self.samples) > 0:\n            print(f\"  Feature dim: {self.samples[0]['neural'].shape[-1]}\")\n            print(f\"  Avg length: {np.mean([s['neural'].shape[0] for s in self.samples]):.1f} timesteps\")\n            if self.vocab is not None:\n                print(f\"  Vocab size: {len(self.vocab)}\")\n    \n    def _load_data(self):\n        \"\"\"Load data from all sessions\"\"\"\n        sessions = sorted(os.listdir(self.data_dir))\n        \n        for session in tqdm(sessions, desc=f\"Loading {self.split}\"):\n            session_path = os.path.join(self.data_dir, session)\n            file_path = os.path.join(session_path, f'data_{self.split}.hdf5')\n            \n            if not os.path.exists(file_path):\n                continue\n            \n            with h5py.File(file_path, 'r') as f:\n                # Get all trial keys\n                trial_keys = sorted([k for k in f.keys() if k.startswith('trial_')])\n                \n                for trial_key in trial_keys:\n                    trial_group = f[trial_key]\n                    \n                    # Load neural data\n                    neural_data = trial_group['input_features'][:]  # (time, 512)\n                    \n                    # Load transcription (integer encoded) - only if it exists\n                    transcription = None\n                    if 'transcription' in trial_group:\n                        transcription = trial_group['transcription'][:]\n                        \n                        # Remove padding (assume 0 is padding token)\n                        transcription = transcription[transcription != 0]\n                        \n                        # Store unique tokens for vocab building (only for train)\n                        if self.split == 'train':\n                            self.all_tokens.update(transcription.tolist())\n                    \n                    # Handle max_len\n                    if self.max_len:\n                        if len(neural_data) > self.max_len:\n                            neural_data = neural_data[:self.max_len]\n                        else:\n                            pad = np.zeros((self.max_len - len(neural_data), neural_data.shape[-1]))\n                            neural_data = np.concatenate([neural_data, pad])\n                    \n                    self.samples.append({\n                        'neural': neural_data,\n                        'transcription': transcription,\n                        'session': session,\n                        'trial': trial_key\n                    })\n    \n    def _build_vocab(self):\n        \"\"\"Build vocabulary from tokens\"\"\"\n        # Special tokens\n        self.vocab = {\n            '<PAD>': 0,\n            '<SOS>': 1,\n            '<EOS>': 2,\n            '<UNK>': 3,\n        }\n        \n        # Add all unique tokens\n        for token in sorted(self.all_tokens):\n            if token not in self.vocab.values():\n                self.vocab[token] = len(self.vocab)\n        \n        # Create reverse vocab\n        self.idx2token = {v: k for k, v in self.vocab.items()}\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        \n        return {\n            'neural': torch.FloatTensor(sample['neural']),\n            'transcription': torch.LongTensor(sample['transcription']) if sample['transcription'] is not None else None,\n            'length': sample['neural'].shape[0],\n            'session': sample['session'],\n            'trial': sample['trial']\n        }\n\n\ndef collate_fn(batch):\n    \"\"\"Custom collate function for variable length sequences\"\"\"\n    # Find max length in batch\n    max_len = max([b['neural'].shape[0] for b in batch])\n    feature_dim = batch[0]['neural'].shape[-1]\n    \n    # Pad neural sequences\n    neural_padded = torch.zeros(len(batch), max_len, feature_dim)\n    attention_mask = torch.zeros(len(batch), max_len)\n    \n    # Pad transcriptions\n    has_transcriptions = batch[0]['transcription'] is not None\n    if has_transcriptions:\n        max_transcript_len = max([len(b['transcription']) for b in batch if b['transcription'] is not None])\n        transcriptions = torch.zeros(len(batch), max_transcript_len, dtype=torch.long)\n    else:\n        transcriptions = None\n    \n    sessions = []\n    trials = []\n    \n    for i, b in enumerate(batch):\n        # Pad neural data\n        length = b['neural'].shape[0]\n        neural_padded[i, :length] = b['neural']\n        attention_mask[i, :length] = 1\n        \n        # Pad transcriptions\n        if has_transcriptions and b['transcription'] is not None:\n            trans_len = len(b['transcription'])\n            transcriptions[i, :trans_len] = b['transcription']\n        \n        sessions.append(b['session'])\n        trials.append(b['trial'])\n    \n    return {\n        'neural': neural_padded,\n        'attention_mask': attention_mask,\n        'transcriptions': transcriptions,\n        'sessions': sessions,\n        'trials': trials\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T18:17:26.264223Z","iopub.execute_input":"2025-10-28T18:17:26.264445Z","iopub.status.idle":"2025-10-28T18:17:26.280149Z","shell.execute_reply.started":"2025-10-28T18:17:26.264428Z","shell.execute_reply":"2025-10-28T18:17:26.279365Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load datasets\nprint(\"📚 Loading datasets...\\n\")\ntrain_dataset = BCIDataset(DATA_DIR, split='train')\nval_dataset = BCIDataset(DATA_DIR, split='val', vocab=train_dataset.vocab)\n\n# Create dataloaders\ntrain_loader = DataLoader(\n    train_dataset, \n    batch_size=2, \n    shuffle=True, \n    collate_fn=collate_fn,\n    num_workers=2\n)\nval_loader = DataLoader(\n    val_dataset, \n    batch_size=4, \n    shuffle=False, \n    collate_fn=collate_fn,\n    num_workers=2\n)\n\nprint(f\"\\n✓ DataLoaders created\")\nprint(f\"  Train batches: {len(train_loader)}\")\nprint(f\"  Val batches: {len(val_loader)}\")\n\n# Inspect a batch\nprint(\"\\n\" + \"=\"*60)\nprint(\"📊 Sample batch inspection:\")\nprint(\"=\"*60)\nsample_batch = next(iter(train_loader))\nprint(f\"Neural data shape: {sample_batch['neural'].shape}\")\nprint(f\"Attention mask shape: {sample_batch['attention_mask'].shape}\")\nif sample_batch['transcriptions'] is not None:\n    print(f\"Transcriptions shape: {sample_batch['transcriptions'].shape}\")\n    print(f\"Sample transcription (first 20 tokens): {sample_batch['transcriptions'][0][:20]}\")\nprint(f\"Sessions: {sample_batch['sessions'][:3]}...\")\nprint(f\"Trials: {sample_batch['trials'][:3]}...\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T18:17:26.280907Z","iopub.execute_input":"2025-10-28T18:17:26.281145Z","iopub.status.idle":"2025-10-28T18:20:45.346162Z","shell.execute_reply.started":"2025-10-28T18:17:26.281127Z","shell.execute_reply":"2025-10-28T18:20:45.345199Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Neural-BERT: Self-Supervised Pre-training\n\nBefore training the end-to-end model, we pre-train the neural encoder using self-supervised objectives:\n1. **Masked Neural Modeling:** Predict masked timesteps\n2. **Contrastive Learning:** Distinguish between augmented views","metadata":{}},{"cell_type":"code","source":"class NeuralEncoder(nn.Module):\n    \"\"\"Transformer encoder for neural signals with local attention\"\"\"\n    \n    def __init__(self, input_dim=512, hidden_dim=768, num_layers=6, num_heads=8, dropout=0.1, max_seq_len=5000):\n        super().__init__()\n        \n        self.input_dim = input_dim\n        self.hidden_dim = hidden_dim\n        self.max_seq_len = max_seq_len  # ← NEW: Make it configurable\n        \n        # Input projection with 1D conv (captures local patterns like phonemes)\n        self.input_proj = nn.Sequential(\n            nn.Conv1d(input_dim, hidden_dim, kernel_size=5, padding=2),\n            nn.ReLU(),\n            nn.Conv1d(hidden_dim, hidden_dim, kernel_size=5, padding=2),\n            nn.ReLU()\n        )\n        \n        # Learned positional encoding\n        self.pos_embedding = nn.Parameter(torch.randn(1, max_seq_len, hidden_dim) * 0.02)  # ← CHANGED\n        \n        # Transformer encoder\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=hidden_dim,\n            nhead=num_heads,\n            dim_feedforward=hidden_dim * 4,\n            dropout=dropout,\n            activation='gelu',\n            batch_first=True,\n            norm_first=True\n        )\n        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)\n        \n        # Layer norm\n        self.ln = nn.LayerNorm(hidden_dim)\n    \n    def forward(self, x, attention_mask=None):\n        \"\"\"\n        Args:\n            x: (batch, time, features)\n            attention_mask: (batch, time)\n        Returns:\n            encoded: (batch, time, hidden_dim)\n        \"\"\"\n        batch_size, seq_len, _ = x.shape\n        \n        # Check if sequence is too long\n        if seq_len > self.max_seq_len:  # ← NEW: Safety check\n            raise ValueError(f\"Sequence length {seq_len} exceeds max_seq_len {self.max_seq_len}\")\n        \n        # Convolutional stem\n        x = rearrange(x, 'b t f -> b f t')\n        x = self.input_proj(x)\n        x = rearrange(x, 'b f t -> b t f')\n        \n        # Add positional encoding\n        x = x + self.pos_embedding[:, :seq_len, :]\n        \n        # Create attention mask for transformer\n        if attention_mask is not None:\n            # Convert to transformer format (True = ignore)\n            mask = (attention_mask == 0)\n        else:\n            mask = None\n        \n        # Transform\n        x = self.transformer(x, src_key_padding_mask=mask)\n        x = self.ln(x)\n        \n        return x\n\n\nclass NeuralBERT(nn.Module):\n    \"\"\"Self-supervised pre-training for neural encoder\"\"\"\n    \n    def __init__(self, encoder, input_dim=512, hidden_dim=768, mask_prob=0.15):\n        super().__init__()\n        \n        self.encoder = encoder\n        self.mask_prob = mask_prob\n        \n        # Prediction head for masked modeling\n        self.prediction_head = nn.Sequential(\n            nn.Linear(hidden_dim, hidden_dim),\n            nn.GELU(),\n            nn.LayerNorm(hidden_dim),\n            nn.Linear(hidden_dim, input_dim)\n        )\n        \n        # Projection head for contrastive learning\n        self.projection_head = nn.Sequential(\n            nn.Linear(hidden_dim, hidden_dim),\n            nn.GELU(),\n            nn.Linear(hidden_dim, 256)\n        )\n    \n    def mask_input(self, x, attention_mask):\n        \"\"\"Randomly mask input timesteps\"\"\"\n        batch_size, seq_len, feat_dim = x.shape\n        \n        # Create mask (only mask valid positions)\n        mask = (torch.rand(batch_size, seq_len, device=x.device) < self.mask_prob)\n        mask = mask & (attention_mask.bool())\n        \n        # Create masked input (replace with zeros)\n        x_masked = x.clone()\n        x_masked[mask] = 0\n        \n        return x_masked, mask\n    \n    def forward(self, x, attention_mask=None):\n        \"\"\"\n        Returns:\n            loss: Combined pre-training loss\n            metrics: Dictionary of individual losses\n        \"\"\"\n        # Apply temporal augmentation (jittering)\n        x_aug1 = x + torch.randn_like(x) * 0.1\n        x_aug2 = x + torch.randn_like(x) * 0.1\n        \n        # === Masked Neural Modeling ===\n        x_masked, mask = self.mask_input(x, attention_mask)\n        encoded = self.encoder(x_masked, attention_mask)\n        predictions = self.prediction_head(encoded)\n        \n        # Compute reconstruction loss (only on masked positions)\n        mask_expanded = mask.unsqueeze(-1).expand_as(x)\n        mlm_loss = F.mse_loss(\n            predictions[mask_expanded].view(-1),\n            x[mask_expanded].view(-1)\n        )\n        \n        # === Contrastive Learning ===\n        # Encode augmented views\n        z1 = self.encoder(x_aug1, attention_mask)\n        z2 = self.encoder(x_aug2, attention_mask)\n        \n        # Global average pooling\n        if attention_mask is not None:\n            mask_sum = attention_mask.sum(dim=1, keepdim=True).clamp(min=1)  # Avoid division by zero\n            z1_pooled = (z1 * attention_mask.unsqueeze(-1)).sum(dim=1) / mask_sum\n            z2_pooled = (z2 * attention_mask.unsqueeze(-1)).sum(dim=1) / mask_sum\n        else:\n            z1_pooled = z1.mean(dim=1)\n            z2_pooled = z2.mean(dim=1)\n        \n        # Project to contrastive space\n        p1 = self.projection_head(z1_pooled)\n        p2 = self.projection_head(z2_pooled)\n        \n        # Normalize\n        p1 = F.normalize(p1, dim=-1)\n        p2 = F.normalize(p2, dim=-1)\n        \n        # Contrastive loss (InfoNCE)\n        temperature = 0.07\n        logits = torch.matmul(p1, p2.t()) / temperature\n        labels = torch.arange(len(p1), device=p1.device)\n        contrastive_loss = F.cross_entropy(logits, labels)\n        \n        # Combined loss\n        total_loss = mlm_loss + 0.5 * contrastive_loss\n        \n        metrics = {\n            'mlm_loss': mlm_loss.item(),\n            'contrastive_loss': contrastive_loss.item(),\n            'total_loss': total_loss.item()\n        }\n        \n        return total_loss, metrics\n\n\ndef pretrain_neural_bert(model, train_loader, val_loader, epochs=5, lr=1e-4, device='cuda'):\n    \"\"\"Pre-train the neural encoder\"\"\"\n    \n    optimizer = AdamW(model.parameters(), lr=lr, weight_decay=0.01)\n    scheduler = CosineAnnealingWarmRestarts(optimizer, T_0=len(train_loader), T_mult=2)\n    \n    best_val_loss = float('inf')\n    \n    for epoch in range(epochs):\n        # Training\n        model.train()\n        train_metrics = defaultdict(list)\n        \n        pbar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{epochs} [Train]\")\n        for batch in pbar:\n            neural = batch['neural'].to(device)\n            attention_mask = batch['attention_mask'].to(device)\n            \n            # Forward pass\n            loss, metrics = model(neural, attention_mask)\n            \n            # Backward pass\n            optimizer.zero_grad()\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            scheduler.step()\n            \n            # Track metrics\n            for k, v in metrics.items():\n                train_metrics[k].append(v)\n            \n            pbar.set_postfix({k: f\"{np.mean(v):.4f}\" for k, v in train_metrics.items()})\n        \n        # Validation\n        model.eval()\n        val_metrics = defaultdict(list)\n        \n        with torch.no_grad():\n            for batch in tqdm(val_loader, desc=f\"Epoch {epoch+1}/{epochs} [Val]\"):\n                neural = batch['neural'].to(device)\n                attention_mask = batch['attention_mask'].to(device)\n                \n                loss, metrics = model(neural, attention_mask)\n                \n                for k, v in metrics.items():\n                    val_metrics[k].append(v)\n        \n        # Print epoch summary\n        print(f\"\\n📊 Epoch {epoch+1} Summary:\")\n        print(f\"   Train Loss: {np.mean(train_metrics['total_loss']):.4f}\")\n        print(f\"   Val Loss: {np.mean(val_metrics['total_loss']):.4f}\")\n        \n        # Save best model\n        val_loss = np.mean(val_metrics['total_loss'])\n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            torch.save(model.encoder.state_dict(), 'neural_encoder_pretrained.pth')\n            print(f\"   ✓ Saved best model (val_loss: {val_loss:.4f})\\n\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T18:20:45.347601Z","iopub.execute_input":"2025-10-28T18:20:45.347893Z","iopub.status.idle":"2025-10-28T18:20:45.369370Z","shell.execute_reply.started":"2025-10-28T18:20:45.347866Z","shell.execute_reply":"2025-10-28T18:20:45.368614Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Initialize and pre-train Neural-BERT\nprint(\"🧠 Initializing Neural-BERT...\\n\")\n\n# Get input dimension from data\nsample_batch = next(iter(train_loader))\ninput_dim = sample_batch['neural'].shape[-1]\nprint(f\"Input dimension: {input_dim}\\n\")\n\n# Check maximum sequence length in dataset\nprint(\"📊 Checking sequence lengths...\\n\")\nmax_train_len = max([s['neural'].shape[0] for s in train_dataset.samples])\nmax_val_len = max([s['neural'].shape[0] for s in val_dataset.samples])\nmax_seq_len = max(max_train_len, max_val_len) + 100  # Add buffer\n\nprint(f\"Max train length: {max_train_len}\")\nprint(f\"Max val length: {max_val_len}\")\nprint(f\"Max seq length: {max_seq_len}\")\nprint(f\"Setting max_seq_len to: {max(max_train_len, max_val_len) + 100}\\n\")\n\n# Create encoder\nencoder = NeuralEncoder(\n    input_dim=input_dim,\n    hidden_dim=512,\n    num_layers=4,\n    num_heads=8,\n    dropout=0.1,\n    max_seq_len=max_seq_len\n).to(device)\n\n# Create Neural-BERT model\nneural_bert = NeuralBERT(\n    encoder=encoder,\n    input_dim=input_dim,\n    hidden_dim=512,  # ← CHANGED TO MATCH NEW ENCODER SIZE\n    mask_prob=0.15\n).to(device)\n\nprint(f\"✓ Model created ({sum(p.numel() for p in neural_bert.parameters())/1e6:.1f}M parameters)\\n\")\n\n# Pre-train\nprint(\"🚀 Starting pre-training...\\n\")\npretrain_neural_bert(\n    neural_bert, \n    train_loader, \n    val_loader, \n    epochs=3,  # Adjust based on time\n    lr=1e-4,\n    device=device\n)\n\nprint(\"\\n✅ Pre-training complete!\")\nprint(f\"✓ Best model saved to 'neural_encoder_pretrained.pth'\\n\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T18:20:45.370226Z","iopub.execute_input":"2025-10-28T18:20:45.370489Z","execution_failed":"2025-10-28T20:34:02.103Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Need to check from here!!!\n## Need to check from here!!!\n## Need to check from here!!!\n## Need to check from here!!!\n## Need to check from here!!!\n## Need to check from here!!!\n## Need to check from here!!!","metadata":{}},{"cell_type":"markdown","source":"## 4. End-to-End Model: Neural → Text","metadata":{}},{"cell_type":"code","source":"class BrainToTextModel(nn.Module):\n    \"\"\"End-to-end model: Neural signals → Text\"\"\"\n    \n    def __init__(self, neural_encoder, use_pretrained_gpt2=True):\n        super().__init__()\n        \n        self.neural_encoder = neural_encoder\n        \n        # Initialize GPT-2\n        if use_pretrained_gpt2:\n            self.tokenizer = GPT2Tokenizer.from_pretrained('gpt2')\n            self.tokenizer.pad_token = self.tokenizer.eos_token\n            self.gpt2 = GPT2LMHeadModel.from_pretrained('gpt2')\n        else:\n            config = GPT2Config(\n                vocab_size=50257,\n                n_positions=1024,\n                n_embd=768,\n                n_layer=12,\n                n_head=12\n            )\n            self.tokenizer = GPT2Tokenizer.from_pretrained('gpt2')\n            self.tokenizer.pad_token = self.tokenizer.eos_token\n            self.gpt2 = GPT2LMHeadModel(config)\n        \n        # Cross-attention bridge (project neural features to GPT-2 dimension)\n        gpt2_dim = self.gpt2.config.n_embd\n        neural_dim = neural_encoder.hidden_dim\n        \n        if neural_dim != gpt2_dim:\n            self.bridge = nn.Linear(neural_dim, gpt2_dim)\n        else:\n            self.bridge = nn.Identity()\n        \n        # Modify GPT-2 to accept encoder outputs\n        # We'll use the transformer's cross-attention capability\n        self.use_cross_attention = True\n    \n    def forward(self, neural_data, attention_mask, text_labels=None):\n        \"\"\"\n        Args:\n            neural_data: (batch, time, features)\n            attention_mask: (batch, time)\n            text_labels: (batch, text_len) - tokenized text for training\n        \"\"\"\n        # Encode neural signals\n        neural_features = self.neural_encoder(neural_data, attention_mask)\n        \n        # Project to GPT-2 space\n        encoder_hidden_states = self.bridge(neural_features)\n        \n        if text_labels is not None:\n            # Training: use teacher forcing\n            outputs = self.gpt2(\n                input_ids=text_labels,\n                encoder_hidden_states=encoder_hidden_states,\n                encoder_attention_mask=attention_mask,\n                labels=text_labels\n            )\n            return outputs.loss, outputs.logits\n        else:\n            # Inference: generate text\n            return encoder_hidden_states\n    \n    @torch.no_grad()\n    def generate_text(self, neural_data, attention_mask, max_length=50):\n        \"\"\"Generate text from neural signals\"\"\"\n        # Encode neural signals\n        encoder_hidden_states = self(neural_data, attention_mask, text_labels=None)\n        \n        # Start with BOS token\n        batch_size = neural_data.shape[0]\n        input_ids = torch.full(\n            (batch_size, 1), \n            self.tokenizer.bos_token_id or self.tokenizer.eos_token_id,\n            dtype=torch.long,\n            device=neural_data.device\n        )\n        \n        # Generate tokens\n        generated = self.gpt2.generate(\n            input_ids=input_ids,\n            encoder_hidden_states=encoder_hidden_states,\n            encoder_attention_mask=attention_mask,\n            max_length=max_length,\n            num_beams=4,\n            early_stopping=True,\n            pad_token_id=self.tokenizer.eos_token_id,\n            eos_token_id=self.tokenizer.eos_token_id\n        )\n        \n        # Decode to text\n        texts = self.tokenizer.batch_decode(generated, skip_special_tokens=True)\n        return texts","metadata":{"trusted":true,"execution":{"execution_failed":"2025-10-28T20:34:02.104Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Initialize end-to-end model\nprint(\"🔗 Initializing end-to-end Brain-to-Text model...\\n\")\n\n# Load pre-trained encoder if available\nif os.path.exists('neural_encoder_pretrained.pth'):\n    encoder.load_state_dict(torch.load('neural_encoder_pretrained.pth'))\n    print(\"✓ Loaded pre-trained encoder\\n\")\n\n# Create end-to-end model\nmodel = BrainToTextModel(\n    neural_encoder=encoder,\n    use_pretrained_gpt2=True\n).to(device)\n\nprint(f\"✓ Model created ({sum(p.numel() for p in model.parameters())/1e6:.1f}M parameters)\")\nprint(f\"   Encoder: {sum(p.numel() for p in model.neural_encoder.parameters())/1e6:.1f}M\")\nprint(f\"   GPT-2: {sum(p.numel() for p in model.gpt2.parameters())/1e6:.1f}M\\n\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-10-28T20:34:02.104Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Training Loop","metadata":{}},{"cell_type":"code","source":"def calculate_wer(predictions, references):\n    \"\"\"Calculate Word Error Rate\"\"\"\n    total_words = 0\n    total_errors = 0\n    \n    for pred, ref in zip(predictions, references):\n        if ref is None:\n            continue\n            \n        pred_words = pred.lower().split()\n        ref_words = ref.lower().split()\n        \n        errors = editdistance.eval(pred_words, ref_words)\n        total_errors += errors\n        total_words += len(ref_words)\n    \n    wer = total_errors / max(total_words, 1)\n    return wer * 100","metadata":{"trusted":true,"execution":{"execution_failed":"2025-10-28T20:34:02.104Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_end_to_end(model, train_loader, val_loader, epochs=10, lr=5e-5):\n    \"\"\"Train the end-to-end model\"\"\"\n    \n    # Freeze GPT-2 initially (optional)\n    for param in model.gpt2.parameters():\n        param.requires_grad = False\n    \n    optimizer = AdamW(\n        filter(lambda p: p.requires_grad, model.parameters()),\n        lr=lr,\n        weight_decay=0.01\n    )\n    scheduler = CosineAnnealingWarmRestarts(optimizer, T_0=len(train_loader)*2)\n    \n    best_wer = float('inf')\n    \n    for epoch in range(epochs):\n        # Unfreeze GPT-2 after first epoch (optional)\n        if epoch == 1:\n            print(\"\\n🔓 Unfreezing GPT-2 decoder...\\n\")\n            for param in model.gpt2.parameters():\n                param.requires_grad = True\n            optimizer = AdamW(model.parameters(), lr=lr/2, weight_decay=0.01)\n        \n        # Training\n        model.train()\n        train_losses = []\n        \n        pbar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{epochs} [Train]\")\n        for batch in pbar:\n            neural = batch['neural'].to(device)\n            attention_mask = batch['attention_mask'].to(device)\n            sentences = batch['sentences']\n            \n            # Tokenize sentences\n            tokens = model.tokenizer(\n                sentences,\n                padding=True,\n                truncation=True,\n                max_length=50,\n                return_tensors='pt'\n            ).input_ids.to(device)\n            \n            # Forward pass\n            loss, logits = model(neural, attention_mask, tokens)\n            \n            # Backward pass\n            optimizer.zero_grad()\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            scheduler.step()\n            \n            train_losses.append(loss.item())\n            pbar.set_postfix({'loss': f\"{np.mean(train_losses):.4f}\"})\n        \n        # Validation\n        model.eval()\n        val_predictions = []\n        val_references = []\n        \n        with torch.no_grad():\n            for batch in tqdm(val_loader, desc=f\"Epoch {epoch+1}/{epochs} [Val]\"):\n                neural = batch['neural'].to(device)\n                attention_mask = batch['attention_mask'].to(device)\n                sentences = batch['sentences']\n                \n                # Generate predictions\n                predictions = model.generate_text(neural, attention_mask, max_length=50)\n                \n                val_predictions.extend(predictions)\n                val_references.extend(sentences)\n        \n        # Calculate WER\n        wer = calculate_wer(val_predictions, val_references)\n        \n        # Print epoch summary\n        print(f\"\\n📊 Epoch {epoch+1} Summary:\")\n        print(f\"   Train Loss: {np.mean(train_losses):.4f}\")\n        print(f\"   Val WER: {wer:.2f}%\")\n        \n        # Print sample predictions\n        print(f\"\\n   Sample Predictions:\")\n        for i in range(min(3, len(val_predictions))):\n            print(f\"   [{i+1}] True: {val_references[i]}\")\n            print(f\"       Pred: {val_predictions[i]}\\n\")\n        \n        # Save best model\n        if wer < best_wer:\n            best_wer = wer\n            torch.save(model.state_dict(), 'brain_to_text_best.pth')\n            print(f\"   ✓ Saved best model (WER: {wer:.2f}%)\\n\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-10-28T20:34:02.104Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Train the end-to-end model\nprint(\"🚀 Starting end-to-end training...\\n\")\ntrain_end_to_end(\n    model,\n    train_loader,\n    val_loader,\n    epochs=10,\n    lr=5e-5\n)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-10-28T20:34:02.104Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Generate Test Predictions & Submission","metadata":{}},{"cell_type":"code","source":"test_dataset = BCIDataset(DATA_DIR, split='test', vocab=train_dataset.vocab)\n\n# Inspect test batch\ntest_loader = DataLoader(\n    test_dataset, \n    batch_size=16, \n    shuffle=False, \n    collate_fn=collate_fn,\n    num_workers=2\n)\nprint(f\"  Test batches: {len(test_loader)}\")\n\nprint(\"\\n📊 Test batch inspection:\")\nprint(\"=\"*60)\ntest_batch = next(iter(test_loader))\nprint(f\"Neural data shape: {test_batch['neural'].shape}\")\nprint(f\"Attention mask shape: {test_batch['attention_mask'].shape}\")\nprint(f\"Has transcriptions: {test_batch['transcriptions'] is not None}\")\nprint(f\"Sessions: {test_batch['sessions'][:3]}...\")\nprint(f\"Trials: {test_batch['trials'][:3]}...\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-10-28T20:34:02.104Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load best model\nprint(\"📥 Loading best model...\\n\")\nif os.path.exists('brain_to_text_best.pth'):\n    model.load_state_dict(torch.load('brain_to_text_best.pth'))\n    print(\"✓ Loaded best model\\n\")\n\n# Generate test predictions\nmodel.eval()\ntest_predictions = []\n\nprint(\"🔮 Generating test predictions...\\n\")\nwith torch.no_grad():\n    for batch in tqdm(test_loader, desc=\"Testing\"):\n        neural = batch['neural'].to(device)\n        attention_mask = batch['attention_mask'].to(device)\n        \n        # Generate predictions\n        predictions = model.generate_text(neural, attention_mask, max_length=50)\n        test_predictions.extend(predictions)\n\nprint(f\"\\n✓ Generated {len(test_predictions)} predictions\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-10-28T20:34:02.104Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create submission file\nsubmission = pd.DataFrame({\n    'id': range(len(test_predictions)),\n    'text': test_predictions\n})\n\nsubmission.to_csv('submission.csv', index=False)\nprint(\"\\n💾 Saved submission.csv\")\nprint(f\"\\n📄 Submission preview:\")\nprint(submission.head(10))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7. Analysis & Visualization","metadata":{}},{"cell_type":"code","source":"# Analyze predictions on validation set\nprint(\"📊 Validation Set Analysis:\\n\")\n\n# Re-run validation to get predictions\nmodel.eval()\nval_predictions = []\nval_references = []\n\nwith torch.no_grad():\n    for batch in tqdm(val_loader, desc=\"Analyzing validation\"):\n        neural = batch['neural'].to(device)\n        attention_mask = batch['attention_mask'].to(device)\n        sentences = batch['sentences']\n        \n        predictions = model.generate_text(neural, attention_mask, max_length=50)\n        \n        val_predictions.extend(predictions)\n        val_references.extend(sentences)\n\n# Calculate final WER\nfinal_wer = calculate_wer(val_predictions, val_references)\nprint(f\"\\n📈 Final Validation WER: {final_wer:.2f}%\")\nprint(f\"   Baseline: 6.70%\")\nprint(f\"   Improvement: {6.70 - final_wer:.2f}%\\n\")\n\n# Show best and worst examples\nerrors = []\nfor pred, ref in zip(val_predictions, val_references):\n    if ref is not None:\n        pred_words = pred.lower().split()\n        ref_words = ref.lower().split()\n        error = editdistance.eval(pred_words, ref_words) / max(len(ref_words), 1)\n        errors.append((error, pred, ref))\n\nerrors.sort()\n\nprint(\"✅ Best 5 Predictions:\")\nfor i, (err, pred, ref) in enumerate(errors[:5]):\n    print(f\"\\n{i+1}. Error: {err*100:.1f}%\")\n    print(f\"   True: {ref}\")\n    print(f\"   Pred: {pred}\")\n\nprint(\"\\n\\n❌ Worst 5 Predictions:\")\nfor i, (err, pred, ref) in enumerate(errors[-5:]):\n    print(f\"\\n{i+1}. Error: {err*100:.1f}%\")\n    print(f\"   True: {ref}\")\n    print(f\"   Pred: {pred}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 8. Next Steps & Improvements\n\nIf the baseline approach doesn't beat 6.70% WER, try these improvements:\n\n### **Option 1: Add Wav2Vec2 Pre-training**\n```python\n# Install wav2vec2\n!pip install -q librosa\n\n# Use speech representations to guide neural encoder\nfrom transformers import Wav2Vec2Model\nwav2vec = Wav2Vec2Model.from_pretrained('facebook/wav2vec2-base-960h')\n\n# Align neural representations with speech representations\n# (Add contrastive loss between neural and speech features)\n```\n\n### **Option 2: Enhanced Data Augmentation**\n- Temporal stretching/compression\n- Channel dropout\n- Mixup augmentation\n- Back-translation with GPT-2\n\n### **Option 3: Ensemble Multiple Models**\n- Train 3-5 models with different random seeds\n- Use beam search with different widths\n- Combine predictions with voting\n\n### **Option 4: Fine-tune with CTC Loss**\n- Add auxiliary CTC loss on phoneme predictions\n- Helps encoder learn phonetic structure\n- Improves alignment between neural and text","metadata":{}}]}