{"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":"none","dataSources":[{"sourceId":106809,"databundleVersionId":13056355,"sourceType":"competition"},{"sourceId":13873823,"sourceType":"datasetVersion","datasetId":8839346},{"sourceId":618527,"sourceType":"modelInstanceVersion","modelInstanceId":465112,"modelId":480951}],"dockerImageVersionId":31154,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"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-03T23:17:42.875069Z","iopub.execute_input":"2025-12-03T23:17:42.875465Z","iopub.status.idle":"2025-12-03T23:18:14.058251Z","shell.execute_reply.started":"2025-12-03T23:17:42.875437Z","shell.execute_reply":"2025-12-03T23:18:14.056919Z"}},"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-03T23:18:14.060603Z","iopub.execute_input":"2025-12-03T23:18:14.060968Z","iopub.status.idle":"2025-12-03T23:18:19.974945Z","shell.execute_reply.started":"2025-12-03T23:18:14.060933Z","shell.execute_reply":"2025-12-03T23:18:19.973493Z"}},"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-12-03T23:18:19.975677Z","iopub.execute_input":"2025-12-03T23:18:19.976166Z","iopub.status.idle":"2025-12-03T23:18:19.983594Z","shell.execute_reply.started":"2025-12-03T23:18:19.976118Z","shell.execute_reply":"2025-12-03T23:18:19.982684Z"}},"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-12-03T23:18:19.984727Z","iopub.execute_input":"2025-12-03T23:18:19.985050Z","iopub.status.idle":"2025-12-03T23:18:26.322046Z","shell.execute_reply.started":"2025-12-03T23:18:19.985018Z","shell.execute_reply":"2025-12-03T23:18:26.321064Z"}},"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 embe]dding 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,"execution":{"iopub.status.busy":"2025-12-03T23:18:26.324572Z","iopub.execute_input":"2025-12-03T23:18:26.324936Z","iopub.status.idle":"2025-12-03T23:18:26.347834Z","shell.execute_reply.started":"2025-12-03T23:18:26.324914Z","shell.execute_reply":"2025-12-03T23:18:26.346664Z"}},"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-12-03T23:18:26.348786Z","iopub.execute_input":"2025-12-03T23:18:26.349091Z","iopub.status.idle":"2025-12-03T23:29:01.893116Z","shell.execute_reply.started":"2025-12-03T23:18:26.349069Z","shell.execute_reply":"2025-12-03T23:29:01.891467Z"}},"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-12-03T23:29:01.893958Z","iopub.status.idle":"2025-12-03T23:29:01.894410Z","shell.execute_reply.started":"2025-12-03T23:29:01.894182Z","shell.execute_reply":"2025-12-03T23:29:01.894201Z"}},"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-12-03T23:29:11.299376Z","iopub.execute_input":"2025-12-03T23:29:11.299706Z","iopub.status.idle":"2025-12-03T23:29:11.320527Z","shell.execute_reply.started":"2025-12-03T23:29:11.299685Z","shell.execute_reply":"2025-12-03T23:29:11.319476Z"}},"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-12-03T23:30:15.288430Z","iopub.execute_input":"2025-12-03T23:30:15.288769Z","iopub.status.idle":"2025-12-03T23:30:18.890617Z","shell.execute_reply.started":"2025-12-03T23:30:15.288745Z","shell.execute_reply":"2025-12-03T23:30:18.889482Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# End of Phoneme Retrieval Implementation","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"# Construct the Training Dataset","metadata":{}},{"cell_type":"code","source":"results[0].keys()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T23:30:22.760047Z","iopub.execute_input":"2025-12-03T23:30:22.760521Z","iopub.status.idle":"2025-12-03T23:30:22.768529Z","shell.execute_reply.started":"2025-12-03T23:30:22.760479Z","shell.execute_reply":"2025-12-03T23:30:22.767589Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"results[0]['retrieved'][0]['ground_truth_text']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T23:30:23.334795Z","iopub.execute_input":"2025-12-03T23:30:23.335536Z","iopub.status.idle":"2025-12-03T23:30:23.342050Z","shell.execute_reply.started":"2025-12-03T23:30:23.335508Z","shell.execute_reply":"2025-12-03T23:30:23.340786Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_complete_training_dataset(retriever, data_dir, model_args, k_similar=3):\n    \"\"\"遍历所有training trials,为每个生成 {target, similar1, similar2, similar3} 格式的数据\"\"\"\n    \n    import pandas as pd\n    from tqdm.notebook import tqdm\n    \n    sessions = model_args['dataset']['sessions']\n    session_to_day = {s: i for i, s in enumerate(sessions)}\n    \n    training_samples = []\n    \n    print(f\"\\n{'='*70}\")\n    print(f\"构建训练数据集: 遍历所有phoneme sequences\")\n    print(f\"{'='*70}\\n\")\n    \n    for session in tqdm(sessions, desc=\"Sessions\"):\n        train_file = os.path.join(data_dir, session, 'data_train.hdf5')\n        \n        if not os.path.exists(train_file):\n            continue\n            \n        with h5py.File(train_file, '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                    # 1. Neural data\n                    neural_data = trial['input_features'][:]\n                    \n                    # 2. Ground truth phonemes\n                    seq_class_ids = trial['seq_class_ids'][:]\n                    gt_phoneme_ids = seq_class_ids[seq_class_ids > 0]\n                    target_phoneme = phoneme_ids_to_string(gt_phoneme_ids, idx_to_phoneme)\n                    \n                    # 3. Ground truth text\n                    target_text = trial.attrs.get('sentence_label', '')\n                    if not target_text:\n                        trans = trial['transcription'][:]\n                        valid = trans[trans > 0]\n                        target_text = ''.join([chr(int(c)) for c in valid if 32 <= c <= 126])\n                    \n                    if not target_phoneme.strip() or not target_text:\n                        continue\n                    \n                    # 4. 检索相似样本\n                    retrieved = retriever.retrieve_similar_phonemes(\n                        test_neural=neural_data,\n                        test_session=session,\n                        k=k_similar,\n                        use_dtw=False\n                    )\n                    \n                    # 5. 构建样本\n                    sample = {\n                        'target_phoneme': target_phoneme,\n                        'target_text': target_text,\n                    }\n                    \n                    # 添加检索到的相似样本\n                    for i, ret in enumerate(retrieved[:k_similar], 1):\n                        sample[f'similar_phoneme_{i}'] = ret['ground_truth_phonemes']\n                        sample[f'similar_text_{i}'] = ret['ground_truth_text']\n                        sample[f'similarity_score_{i}'] = ret.get('similarity', 0.0)\n                    \n                    # 填充空缺\n                    for i in range(len(retrieved) + 1, k_similar + 1):\n                        sample[f'similar_phoneme_{i}'] = ''\n                        sample[f'similar_text_{i}'] = ''\n                        sample[f'similarity_score_{i}'] = 0.0\n                    \n                    training_samples.append(sample)\n                    \n                except Exception as e:\n                    print(f\"Error {session}/{trial_key}: {e}\")\n                    continue\n    \n    df = pd.DataFrame(training_samples)\n    \n    print(f\"\\n✓ 完成! 总样本数: {len(df):,}\")\n    print(f\"列: {list(df.columns)}\")\n    \n    return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T23:30:23.786025Z","iopub.execute_input":"2025-12-03T23:30:23.786497Z","iopub.status.idle":"2025-12-03T23:30:23.804551Z","shell.execute_reply.started":"2025-12-03T23:30:23.786459Z","shell.execute_reply":"2025-12-03T23:30:23.803283Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"training_df = build_complete_training_dataset(\n    retriever=retriever,\n    data_dir=DATA_DIR,\n    model_args=model_args,\n    k_similar=3\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T23:30:26.405205Z","iopub.execute_input":"2025-12-03T23:30:26.405619Z","iopub.status.idle":"2025-12-03T23:30:31.703583Z","shell.execute_reply.started":"2025-12-03T23:30:26.405596Z","shell.execute_reply":"2025-12-03T23:30:31.702174Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n样本预览:\")\ndisplay(training_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T23:30:32.809248Z","iopub.execute_input":"2025-12-03T23:30:32.810177Z","iopub.status.idle":"2025-12-03T23:30:32.826732Z","shell.execute_reply.started":"2025-12-03T23:30:32.810118Z","shell.execute_reply":"2025-12-03T23:30:32.825577Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"training_df.to_csv(os.path.join(OUTPUT_DIR, 'phoneme_training_dataset.csv'), index=False)\nprint(\"✓ Saved!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T23:30:34.173842Z","iopub.execute_input":"2025-12-03T23:30:34.174510Z","iopub.status.idle":"2025-12-03T23:30:34.188842Z","shell.execute_reply.started":"2025-12-03T23:30:34.174475Z","shell.execute_reply":"2025-12-03T23:30:34.187200Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# GENERATE TEST DATASET CSV\n# For sessions with data_test.hdf5 files\n# ============================================================================\n\n# Add this code to your Kaggle notebook after loading the model and retriever\n\ndef build_test_dataset(retriever, model, data_dir, model_args, device, k_similar=3):\n    \"\"\"\n    Build test dataset by:\n    1. Finding all data_test.hdf5 files\n    2. Predicting phonemes from neural data using GRU\n    3. Retrieving similar examples from training database\n    4. Saving to CSV format matching training data (without target_text)\n    \"\"\"\n    \n    import pandas as pd\n    from tqdm.notebook import tqdm\n    \n    sessions = model_args['dataset']['sessions']\n    session_to_day = {s: i for i, s in enumerate(sessions)}\n    \n    test_samples = []\n    \n    print(f\"\\n{'='*70}\")\n    print(f\"Building Test Dataset: Processing all data_test.hdf5 files\")\n    print(f\"{'='*70}\\n\")\n    \n    # Find sessions with test files\n    test_sessions = []\n    for session in sessions:\n        test_file = os.path.join(data_dir, session, 'data_test.hdf5')\n        if os.path.exists(test_file):\n            test_sessions.append(session)\n    \n    print(f\"Found {len(test_sessions)} sessions with test data:\")\n    for s in test_sessions:\n        print(f\"  - {s}\")\n    print()\n    \n    for session in tqdm(test_sessions, desc=\"Sessions\"):\n        test_file = os.path.join(data_dir, session, 'data_test.hdf5')\n        day_idx = session_to_day.get(session, 0)\n        \n        with h5py.File(test_file, 'r') as f:\n            trials = sorted(list(f.keys()))\n            \n            for trial_key in tqdm(trials, desc=f\"{session}\", leave=False):\n                try:\n                    trial = f[trial_key]\n                    \n                    # 1. Get neural data\n                    neural_data = trial['input_features'][:]\n                    \n                    # 2. Predict phonemes using GRU model (no ground truth for test)\n                    target_ids, target_phoneme = predict_phonemes_from_neural(\n                        model, neural_data, day_idx, device\n                    )\n                    \n                    if not target_phoneme.strip():\n                        continue\n                    \n                    # 3. Retrieve similar examples from training database\n                    retrieved = retriever.retrieve_similar_phonemes(\n                        test_neural=neural_data,\n                        test_session=session,\n                        k=k_similar,\n                        use_dtw=False\n                    )\n                    \n                    # 4. Build sample\n                    sample = {\n                        'session': session,\n                        'trial_key': trial_key,\n                        'target_phoneme': target_phoneme,\n                        # Note: No target_text for test data - that's what we want to generate!\n                    }\n                    \n                    # Add retrieved similar examples\n                    for i, ret in enumerate(retrieved[:k_similar], 1):\n                        sample[f'similar_phoneme_{i}'] = ret['ground_truth_phonemes']\n                        sample[f'similar_text_{i}'] = ret['ground_truth_text']\n                        sample[f'similarity_score_{i}'] = ret.get('similarity', 0.0)\n                    \n                    # Fill empty slots if fewer than k_similar retrieved\n                    for i in range(len(retrieved) + 1, k_similar + 1):\n                        sample[f'similar_phoneme_{i}'] = ''\n                        sample[f'similar_text_{i}'] = ''\n                        sample[f'similarity_score_{i}'] = 0.0\n                    \n                    test_samples.append(sample)\n                    \n                except Exception as e:\n                    print(f\"Error {session}/{trial_key}: {e}\")\n                    continue\n    \n    df = pd.DataFrame(test_samples)\n    \n    print(f\"\\n✓ Done! Total test samples: {len(df):,}\")\n    print(f\"Columns: {list(df.columns)}\")\n    \n    return df\n\n\n# ============================================================================\n# USAGE EXAMPLE (add to your notebook)\n# ============================================================================\n\"\"\"\n# After loading model and retriever:\n\ntest_df = build_test_dataset(\n    retriever=retriever,\n    model=model,\n    data_dir=DATA_DIR,\n    model_args=model_args,\n    device=device,\n    k_similar=3\n)\n\n# Preview\nprint(\"\\nTest Dataset Preview:\")\ndisplay(test_df.head())\n\n# Save to CSV\ntest_df.to_csv(os.path.join(OUTPUT_DIR, 'phoneme_test_dataset.csv'), index=False)\nprint(\"✓ Saved to phoneme_test_dataset.csv\")\n\"\"\"\n\n\n# ============================================================================\n# ALTERNATIVE: If you also want to process val data\n# ============================================================================\n\ndef build_val_dataset(retriever, model, data_dir, model_args, device, k_similar=3):\n    \"\"\"Same as test but for data_val.hdf5 files\"\"\"\n    \n    import pandas as pd\n    from tqdm.notebook import tqdm\n    \n    sessions = model_args['dataset']['sessions']\n    session_to_day = {s: i for i, s in enumerate(sessions)}\n    \n    val_samples = []\n    \n    print(f\"\\n{'='*70}\")\n    print(f\"Building Validation Dataset\")\n    print(f\"{'='*70}\\n\")\n    \n    for session in tqdm(sessions, desc=\"Sessions\"):\n        val_file = os.path.join(data_dir, session, 'data_val.hdf5')\n        \n        if not os.path.exists(val_file):\n            continue\n            \n        day_idx = session_to_day.get(session, 0)\n        \n        with h5py.File(val_file, 'r') as f:\n            trials = sorted(list(f.keys()))\n            \n            for trial_key in trials:\n                try:\n                    trial = f[trial_key]\n                    neural_data = trial['input_features'][:]\n                    \n                    # For validation, we might have ground truth - check if available\n                    has_gt = 'seq_class_ids' in trial.keys()\n                    \n                    if has_gt:\n                        # Use ground truth phonemes\n                        seq_class_ids = trial['seq_class_ids'][:]\n                        gt_phoneme_ids = seq_class_ids[seq_class_ids > 0]\n                        target_phoneme = phoneme_ids_to_string(gt_phoneme_ids, idx_to_phoneme)\n                        \n                        # Get ground truth text\n                        target_text = trial.attrs.get('sentence_label', '')\n                        if not target_text:\n                            trans = trial['transcription'][:]\n                            valid = trans[trans > 0]\n                            target_text = ''.join([chr(int(c)) for c in valid if 32 <= c <= 126])\n                    else:\n                        # Predict phonemes\n                        target_ids, target_phoneme = predict_phonemes_from_neural(\n                            model, neural_data, day_idx, device\n                        )\n                        target_text = ''  # No ground truth\n                    \n                    if not target_phoneme.strip():\n                        continue\n                    \n                    # Retrieve similar examples\n                    retrieved = retriever.retrieve_similar_phonemes(\n                        test_neural=neural_data,\n                        test_session=session,\n                        k=k_similar,\n                        use_dtw=False\n                    )\n                    \n                    sample = {\n                        'session': session,\n                        'trial_key': trial_key,\n                        'target_phoneme': target_phoneme,\n                        'target_text': target_text,  # May be empty for true test\n                    }\n                    \n                    for i, ret in enumerate(retrieved[:k_similar], 1):\n                        sample[f'similar_phoneme_{i}'] = ret['ground_truth_phonemes']\n                        sample[f'similar_text_{i}'] = ret['ground_truth_text']\n                        sample[f'similarity_score_{i}'] = ret.get('similarity', 0.0)\n                    \n                    for i in range(len(retrieved) + 1, k_similar + 1):\n                        sample[f'similar_phoneme_{i}'] = ''\n                        sample[f'similar_text_{i}'] = ''\n                        sample[f'similarity_score_{i}'] = 0.0\n                    \n                    val_samples.append(sample)\n                    \n                except Exception as e:\n                    continue\n    \n    df = pd.DataFrame(val_samples)\n    print(f\"\\n✓ Done! Total validation samples: {len(df):,}\")\n    \n    return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T23:34:26.531203Z","iopub.execute_input":"2025-12-03T23:34:26.531585Z","iopub.status.idle":"2025-12-03T23:34:26.555723Z","shell.execute_reply.started":"2025-12-03T23:34:26.531562Z","shell.execute_reply":"2025-12-03T23:34:26.554681Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# BUILD TEST DATASET\n# ============================================================================\n\ntest_df = build_test_dataset(\n    retriever=retriever,\n    model=model,\n    data_dir=DATA_DIR,\n    model_args=model_args,\n    device=device,\n    k_similar=3\n)\n\n# Preview\nprint(\"\\nTest Dataset Preview:\")\ndisplay(test_df.head())\n\n# Save\ntest_df.to_csv(os.path.join(OUTPUT_DIR, 'real_phoneme_test_dataset.csv'), index=False)\nprint(\"✓ Saved!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T23:34:28.953293Z","iopub.execute_input":"2025-12-03T23:34:28.953622Z","iopub.status.idle":"2025-12-03T23:52:29.625896Z","shell.execute_reply.started":"2025-12-03T23:34:28.953601Z","shell.execute_reply":"2025-12-03T23:52:29.624864Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}