{"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,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":14126434,"sourceType":"datasetVersion","datasetId":9000582},{"sourceId":14351060,"sourceType":"datasetVersion","datasetId":9163438},{"sourceId":14352461,"sourceType":"datasetVersion","datasetId":9164425},{"sourceId":14358826,"sourceType":"datasetVersion","datasetId":9168820},{"sourceId":14359723,"sourceType":"datasetVersion","datasetId":9169343},{"sourceId":28785,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":8318,"modelId":3301},{"sourceId":681234,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":516905,"modelId":531567},{"sourceId":704129,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":534444,"modelId":548105},{"sourceId":704564,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":534794,"modelId":548417}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Bi-GRU Baseline + LLM Post-Processing (simple experiment Notebook)\n\nHi everyone,\nI’m sharing a small experiment notebook for the Brain-to-Text task.\n\nI had no prior background in BCI or brain-signal decoding before this competition. This work mainly reflects a top-down learning process: starting from AI-assisted exploration and community-shared ideas, understanding the overall decoding pipeline first, and then running small experiments on top of it.\n\nI’d like to thank the organizers for providing a dataset and task that contribute to BCI research and real-world human assistance.\n\nWhat’s included:\n\nA standard Bi-GRU acoustic model trained on Kaggle (2× T4 GPUs)\n\n4-gram KenLM decoding (pruned, ~4.4GB ,downloaded via a link provided by Gemini)\n\nAn experimental LLM post-processing step (DeepSeek V3 via API) to explore semantic refinement\n\nThanks to the 7th place solution author for sharing their work — before seeing it, I didn’t realize how much 4-gram KenLM size can vary, which helped me better understand the role of the language model in this task.\n\nNote: the LLM post-processing step requires internet access. In a strict Internet Off setup, this would need to be replaced by a local quantized model.\n\nThis notebook is shared for learning and discussion, not as a finalized or competition-optimized solution.\nFeedback is very welcome.","metadata":{}},{"cell_type":"code","source":"# ==================================================================================\n# 🚀 ULTRA-PIPELINE: GRU Training + 3-Gram Decoding + LLM N-Best Rescoring\n# ==================================================================================\n\nimport os, sys, glob, re, random, math, importlib.util, subprocess\nimport h5py\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.nn.utils.rnn import pad_sequence, pack_padded_sequence, pad_packed_sequence\nfrom torch.cuda.amp import autocast, GradScaler\nimport gc\n\n# --- 0. 自动安装依赖 (Flashlight 用于快速 CTC 解码) ---\nif importlib.util.find_spec(\"flashlight\") is None:\n    print(\"🔧 Installing flashlight-text for fast decoding...\")\n    subprocess.check_call([sys.executable, \"-m\", \"pip\", \"-q\", \"install\", \"flashlight-text\"])\n\n# 引入解码器\nfrom torchaudio.models.decoder import ctc_decoder\n# 引入 LLM\nfrom transformers import AutoTokenizer, AutoModelForCausalLM\n\n# --------------------------\n# 1. 全局配置 (Config)\n# --------------------------\nDATA_DIR    = \"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final\"\nOUTPUT_DIR  = \"/kaggle/working\"\nLEXICON_PATH= \"/kaggle/input/lexicon/librispeech-lexicon.txt\"\nLM_PATH     = \"/kaggle/input/3gm/pytorch/default/1/3-gram.pruned.1e-7.arpa\"\n# LLM 路径 (请确保已 Add Input)\nLLM_PATH    = \"/kaggle/input/gemma/transformers/2b-it/3\"  \n\n# 训练参数\nBATCH_SIZE  = 64    # 双卡或大显存设 64，否则 32\nNUM_EPOCHS  = 120   # 增加轮数配合 SpecAugment\nNUM_WORKERS = 2\nLR_MAX      = 5e-4\nWEIGHT_DECAY= 1e-2\n\n# 解码参数 (生成 N-Best)\nBEAM_SIZE      = 150\nBEAM_THRESHOLD = 15.0\nLM_WEIGHT      = 1.5\nWORD_SCORE     = 1.5\nNBEST_COUNT    = 10  # 生成前 10 个候选给 LLM 挑\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"🚀 Device: {DEVICE} (GPUs: {torch.cuda.device_count()})\")\n\n# --------------------------\n# 2. 字符表 & 文本处理\n# --------------------------\nVOCAB = [\n    ' ', \"'\", 'a', 'b', 'c', 'd', 'e', 'f', 'g', 'h', 'i', 'j', 'k', 'l', 'm',\n    'n', 'o', 'p', 'q', 'r', 's', 't', 'u', 'v', 'w', 'x', 'y', 'z']\nCHAR2IDX = {c: i+1 for i, c in enumerate(VOCAB)}\nIDX2CHAR = {i+1: c for i, c in enumerate(VOCAB)}\nTOKENS = [\"BLANK\"] + VOCAB\n\ndef normalize_text(s: str) -> str:\n    s = (s or \"\").strip().lower()\n    s = re.sub(r\"[^a-z0-9'\\s]\", \" \", s)\n    s = re.sub(r\"\\s+\", \" \", s).strip()\n    return s\n\n# --------------------------\n# 3. 数据集 (BrainDataset)\n# --------------------------\nclass BrainDataset(Dataset):\n    def __init__(self, data_dir, split=\"train\", session_map=None):\n        pattern = \"data_train.hdf5\" # 训练/验证都用带标签的数据\n        if split == \"test\": pattern = \"data_test.hdf5\"\n        \n        search_pattern = os.path.join(data_dir, \"**\", pattern)\n        files = sorted(glob.glob(search_pattern, recursive=True))\n        \n        self.data = []\n        self.session_map = session_map if session_map is not None else {}\n        self.is_test = (split == \"test\")\n        \n        print(f\"📥 Loading {split} data... files={len(files)}\")\n        \n        # 建立 Session Map\n        if not self.session_map:\n            sessions = set()\n            for fp in files:\n                try:\n                    with h5py.File(fp, \"r\") as f:\n                        k0 = list(f.keys())[0]\n                        sess = f[k0].attrs.get(\"session\", \"\")\n                        if isinstance(sess, bytes): sess = sess.decode(\"utf-8\")\n                        sessions.add(sess)\n                except: pass\n            for i, s in enumerate(sorted(list(sessions))):\n                self.session_map[s] = i\n            print(f\"🧩 Sessions found: {len(self.session_map)}\")\n\n        # 加载数据索引\n        for fp in tqdm(files, desc=f\"Scan {split}\"):\n            try:\n                with h5py.File(fp, \"r\") as f:\n                    for k in f.keys():\n                        g = f[k]\n                        \n                        # 简单的哈希划分 Train/Val\n                        is_val = (hash(k) % 10 == 0)\n                        if split == \"train\" and is_val: continue\n                        if split == \"val\" and not is_val: continue\n\n                        feats = torch.from_numpy(g[\"input_features\"][:]).float()\n                        if feats.ndim == 2 and feats.shape[0] == 512:\n                            feats = feats.transpose(0, 1)\n\n                        sess = g.attrs.get(\"session\", \"\")\n                        if isinstance(sess, bytes): sess = sess.decode(\"utf-8\")\n                        sess_id = self.session_map.get(sess, 0)\n\n                        label = torch.tensor([], dtype=torch.long)\n                        if not self.is_test:\n                            sent = g.attrs.get(\"sentence_label\", \"\")\n                            if isinstance(sent, bytes): sent = sent.decode(\"utf-8\")\n                            sent = (sent or \"\").lower()\n                            # 过滤无效字符\n                            idxs = [CHAR2IDX[c] for c in sent if c in CHAR2IDX]\n                            label = torch.tensor(idxs, dtype=torch.long)\n                            \n                            # 过滤掉空标签的样本，防止 CTC Loss 报错\n                            if len(idxs) == 0: continue\n\n                        self.data.append({\"key\": k, \"feats\": feats, \"label\": label, \"sess_id\": sess_id})\n            except: pass\n        print(f\"✅ {split}: trials={len(self.data)}\")\n\n    def __len__(self): return len(self.data)\n    def __getitem__(self, idx): return self.data[idx]\n\n    @staticmethod\n    def collate_fn(batch):\n        batch.sort(key=lambda x: x[\"feats\"].shape[0], reverse=True)\n        feats   = [x[\"feats\"] for x in batch]\n        labels  = [x[\"label\"] for x in batch]\n        sess_ids = torch.tensor([x[\"sess_id\"] for x in batch], dtype=torch.long)\n        keys    = [x[\"key\"] for x in batch]\n        \n        in_len = torch.tensor([f.shape[0] for f in feats], dtype=torch.long)\n        padded_feats = pad_sequence(feats, batch_first=True)\n        \n        tgt_len = torch.tensor([l.numel() for l in labels], dtype=torch.long)\n        if int(tgt_len.sum().item()) > 0:\n            flat_targets = torch.cat([l for l in labels if l.numel() > 0], dim=0)\n        else:\n            flat_targets = torch.empty(0, dtype=torch.long)\n            \n        return padded_feats, flat_targets, in_len, tgt_len, sess_ids, keys\n\n# --------------------------\n# 4. 模型：SpecAugment + RNN\n# --------------------------\nclass SpecAugment(nn.Module):\n    def __init__(self, freq_mask=20, time_mask=30):\n        super().__init__()\n        self.freq_mask = freq_mask\n        self.time_mask = time_mask\n\n    def forward(self, x):\n        if not self.training: return x\n        x_aug = x.clone()\n        B, T, C = x_aug.shape\n        # 随机 Mask 两次\n        for _ in range(2):\n            f = np.random.randint(0, self.freq_mask)\n            f0 = np.random.randint(0, C - f)\n            x_aug[:, :, f0:f0+f] = 0\n        for _ in range(2):\n            t = np.random.randint(0, self.time_mask)\n            t0 = np.random.randint(0, T - t)\n            x_aug[:, t0:t0+t, :] = 0\n        return x_aug\n\nclass OfficialRNN(nn.Module):\n    def __init__(self, n_days, n_chars=len(VOCAB)+1):\n        super().__init__()\n        # 1. Day Adapter (Vectorized)\n        self.day_W = nn.Parameter(torch.stack([torch.eye(512) for _ in range(n_days)], dim=0))\n        self.day_b = nn.Parameter(torch.zeros(n_days, 512))\n        \n        # 2. SpecAugment (Robustness)\n        self.augment = SpecAugment()\n\n        # 3. Conv Feature Extractor\n        self.conv = nn.Sequential(\n            nn.Conv1d(512, 512, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm1d(512),\n            nn.ReLU(),\n            nn.Dropout(0.4)\n        )\n\n        # 4. Bi-GRU\n        self.gru = nn.GRU(512, 512, num_layers=5, bidirectional=True, batch_first=True, dropout=0.4)\n        self.fc = nn.Linear(1024, n_chars)\n\n    def forward(self, x, lengths, day_ids):\n        # Apply Adapter\n        day_ids = day_ids.clamp(min=0, max=self.day_W.shape[0]-1)\n        W = self.day_W[day_ids] \n        b = self.day_b[day_ids] \n        x = torch.einsum(\"btc,bcd->btd\", x, W) + b[:, None, :]\n        \n        # Apply Augmentation\n        x = self.augment(x)\n\n        # Conv\n        x = x.transpose(1, 2)\n        x = self.conv(x)\n        x = x.transpose(1, 2)\n        \n        # DataParallel Fix: Force output length alignment\n        total_length = x.size(1) \n\n        # RNN\n        packed = pack_padded_sequence(x, lengths.cpu(), batch_first=True, enforce_sorted=True)\n        out, _ = self.gru(packed)\n        out, _ = pad_packed_sequence(out, batch_first=True, total_length=total_length)\n        \n        return self.fc(out)\n\n# --------------------------\n# 5. 解码器构建 (Flashlight)\n# --------------------------\ndef build_char_lexicon(src_lex, out_lex):\n    # 清洗并构建字符级 Lexicon\n    seen = set()\n    with open(src_lex, \"r\") as f, open(out_lex, \"w\") as o:\n        for line in f:\n            parts = line.strip().split()\n            if len(parts) < 2: continue\n            w = parts[0].lower().split('(')[0]\n            if w in seen: continue\n            \n            # 只保留纯字母单词\n            if not all(c in \"abcdefghijklmnopqrstuvwxyz'\" for c in w): continue\n            \n            # 写入: \"word w o r d\"\n            o.write(f\"{w} {' '.join(list(w))}\\n\")\n            seen.add(w)\n    print(f\"📖 Lexicon built at {out_lex}\")\n\ndef build_decoder(work_dir):\n    lex_out = os.path.join(work_dir, \"lexicon_char.txt\")\n    build_char_lexicon(LEXICON_PATH, lex_out)\n    \n    print(\"🔧 Building Beam Search Decoder...\")\n    return ctc_decoder(\n        lexicon=lex_out,\n        tokens=TOKENS,\n        lm=LM_PATH,\n        nbest=NBEST_COUNT,  # 生成 Top-N\n        beam_size=BEAM_SIZE,\n        beam_threshold=BEAM_THRESHOLD,\n        lm_weight=LM_WEIGHT,\n        word_score=WORD_SCORE,\n        blank_token=\"BLANK\",\n        sil_token=\" \",\n        unk_word=\"<unk>\"\n    )\n\n# --------------------------\n# 6. LLM 打分器 (PPL Scorer)\n# --------------------------\nclass LLMScorer:\n    def __init__(self, model_path):\n        print(f\"🧠 Loading LLM for Rescoring: {model_path}\")\n        self.tokenizer = AutoTokenizer.from_pretrained(model_path)\n        self.model = AutoModelForCausalLM.from_pretrained(\n            model_path, \n            device_map=\"cuda\", \n            torch_dtype=torch.float16\n        )\n        self.model.eval()\n        \n    def get_best_sentence(self, candidates):\n        # 如果只有一个候选，直接返回\n        if not candidates: return \"\"\n        if len(candidates) == 1: return candidates[0]\n        \n        # 计算 PPL\n        scores = []\n        with torch.no_grad():\n            for text in candidates:\n                if not text.strip(): \n                    scores.append(9999.0)\n                    continue\n                inputs = self.tokenizer(text, return_tensors=\"pt\").to(\"cuda\")\n                outputs = self.model(**inputs, labels=inputs[\"input_ids\"])\n                loss = outputs.loss.item()\n                scores.append(loss)\n        \n        # 选 Loss 最小的\n        best_idx = np.argmin(scores)\n        return candidates[best_idx]\n\n# --------------------------\n# 7. 主流程 (Train -> Decode -> Rescore)\n# --------------------------\ndef main():\n    # --- A. 训练阶段 ---\n    train_ds = BrainDataset(DATA_DIR, split=\"train\")\n    train_loader = DataLoader(\n        train_ds, batch_size=BATCH_SIZE, shuffle=True,\n        collate_fn=BrainDataset.collate_fn, num_workers=NUM_WORKERS, pin_memory=True\n    )\n    \n    n_days = max(len(train_ds.session_map), 45)\n    model = OfficialRNN(n_days).to(DEVICE)\n    \n    # 自动多卡\n    if torch.cuda.device_count() > 1:\n        print(f\"⚡ Using {torch.cuda.device_count()} GPUs\")\n        model = nn.DataParallel(model)\n        \n    optimizer = optim.AdamW(model.parameters(), lr=LR_MAX, weight_decay=WEIGHT_DECAY)\n    scheduler = optim.lr_scheduler.OneCycleLR(\n        optimizer, max_lr=LR_MAX, steps_per_epoch=len(train_loader), epochs=NUM_EPOCHS\n    )\n    criterion = nn.CTCLoss(blank=0, zero_infinity=True)\n    scaler = GradScaler()\n    \n    best_loss = float(\"inf\")\n    best_path = os.path.join(OUTPUT_DIR, \"best_gru_sota.pth\")\n    \n    print(\"🏋️ Start Training...\")\n    for epoch in range(NUM_EPOCHS):\n        model.train()\n        total_loss = 0.0\n        pbar = tqdm(train_loader, desc=f\"Ep {epoch+1}/{NUM_EPOCHS}\")\n        \n        for feats, targets, in_len, tgt_len, sess_ids, _ in pbar:\n            feats, targets, sess_ids = feats.to(DEVICE), targets.to(DEVICE), sess_ids.to(DEVICE)\n            \n            optimizer.zero_grad()\n            with autocast():\n                logits = model(feats, in_len, sess_ids)\n                log_probs = logits.log_softmax(2).transpose(0, 1) # (T,B,C)\n                loss = criterion(log_probs, targets, in_len, tgt_len)\n                \n            scaler.scale(loss).backward()\n            scaler.unscale_(optimizer)\n            nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            scaler.step(optimizer)\n            scaler.update()\n            scheduler.step()\n            \n            total_loss += loss.item()\n            pbar.set_postfix(loss=f\"{loss.item():.4f}\")\n            \n        avg_loss = total_loss / len(train_loader)\n        \n        # 简单保存逻辑 (只看 Train Loss 收敛情况，因为验证集太小)\n        if avg_loss < best_loss:\n            best_loss = avg_loss\n            # 处理 DataParallel\n            save_model = model.module if isinstance(model, nn.DataParallel) else model\n            torch.save(save_model.state_dict(), best_path)\n            print(f\"💾 Saved Best: {best_loss:.4f}\")\n            \n    # --- B. 准备推理 ---\n    print(\"🧹 Cleaning up training memory...\")\n    del model, optimizer, scheduler, scaler\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    # 加载最佳模型\n    model = OfficialRNN(n_days).to(DEVICE)\n    model.load_state_dict(torch.load(best_path))\n    model.eval()\n    \n    # 构建解码器\n    decoder = build_decoder(OUTPUT_DIR)\n    \n    # --- C. 解码阶段 (生成 N-Best) ---\n    test_ds = BrainDataset(DATA_DIR, split=\"test\", session_map=train_ds.session_map)\n    test_loader = DataLoader(\n        test_ds, batch_size=BATCH_SIZE, shuffle=False,\n        collate_fn=BrainDataset.collate_fn, num_workers=NUM_WORKERS\n    )\n    \n    print(\"🌊 Decoding N-Best Lists...\")\n    all_candidates = [] # Stores: {\"key\": key, \"hyps\": [cand1, cand2...]}\n    \n    with torch.no_grad():\n        for feats, _, in_len, _, sess_ids, keys in tqdm(test_loader, desc=\"Decoding\"):\n            feats, sess_ids = feats.to(DEVICE), sess_ids.to(DEVICE)\n            \n            logits = model(feats, in_len, sess_ids)\n            logp = logits.log_softmax(2).cpu() # Decoder needs CPU\n            \n            # Beam Search returns list of Hypotheses per sample\n            hypos_batch = decoder(logp, in_len.cpu()) # List[List[Hypothesis]]\n            \n            for i, hyps in enumerate(hypos_batch):\n                # 提取 Top-N 文本\n                candidates = []\n                for h in hyps:\n                    text = \" \".join(h.words).strip()\n                    candidates.append(normalize_text(text))\n                \n                # 去重\n                unique_cands = []\n                seen = set()\n                for c in candidates:\n                    if c and c not in seen:\n                        unique_cands.append(c)\n                        seen.add(c)\n                \n                if not unique_cands: unique_cands = [\" \"]\n                all_candidates.append({\"key\": keys[i], \"hyps\": unique_cands})\n\n    # --- D. 释放显存，加载 LLM ---\n    print(\"🗑️ Freeing RNN memory...\")\n    del model, decoder\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    # --- E. LLM 重打分 ---\n    try:\n        scorer = LLMScorer(LLM_PATH)\n        llm_available = True\n    except Exception as e:\n        print(f\"⚠️ LLM Load Failed ({e}), falling back to Top-1.\")\n        llm_available = False\n        \n    print(\"🧠 Rescoring with LLM...\")\n    final_results = []\n    \n    for item in tqdm(all_candidates, desc=\"Rescoring\"):\n        candidates = item[\"hyps\"]\n        \n        if llm_available and len(candidates) > 1:\n            best_text = scorer.get_best_sentence(candidates)\n        else:\n            best_text = candidates[0] # Fallback to RNN Top-1\n            \n        final_results.append({\"key\": item[\"key\"], \"text\": best_text})\n        \n    # --- F. 保存提交 ---\n    df = pd.DataFrame(final_results)\n    # Kaggle 格式要求 id 从 0 开始\n    df = df.sort_values(\"key\").reset_index(drop=True) # 按 key 排序比较稳妥\n    df[\"id\"] = range(len(df))\n    \n    sub_path = os.path.join(OUTPUT_DIR, \"submission.csv\")\n    df[[\"id\", \"text\"]].to_csv(sub_path, index=False)\n    \n    print(f\"✅ Submission Saved: {sub_path}\")\n    print(df.head(10))\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-24T13:04:21.315902Z","iopub.execute_input":"2025-12-24T13:04:21.316786Z","iopub.status.idle":"2025-12-24T17:00:56.961425Z","shell.execute_reply.started":"2025-12-24T13:04:21.316753Z","shell.execute_reply":"2025-12-24T17:00:56.960655Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================================================================================\n# 🚀 最终究极版: 鲁棒加载 + 4-gram LM + 大小写自动对齐 + 自动转小写\n  **Robust loading + 4-gram LM + automatic case alignment + automatic lowercasing**\n# ==================================================================================\nimport os, sys, glob, h5py, re\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.nn.utils.rnn import pad_sequence, pack_padded_sequence, pad_packed_sequence\n\n# --- 0. 环境准备 ---\ntry:\n    from pyctcdecode import build_ctcdecoder\nexcept ImportError:\n    print(\"🔧 Installing pyctcdecode...\")\n    import subprocess\n    subprocess.check_call([sys.executable, \"-m\", \"pip\", \"-q\", \"install\", \"pyctcdecode\"])\n    subprocess.check_call([sys.executable, \"-m\", \"pip\", \"-q\", \"install\", \"https://github.com/kpu/kenlm/archive/master.zip\"])\n    from pyctcdecode import build_ctcdecoder\n\n# --------------------------\n# 1. 路径与配置\n# --------------------------\n# 🔥 你的最佳模型权重\nMODEL_PATH   = \"/kaggle/input/my-best-gru/pytorch/default/1/best_gru_sota (1).pth\"\n\n# 🔥 你的 4-gram LM (如果路径不对，请改为你实际上传的路径)\n# 如果找不到 4-gram，脚本会自动降级到 3-gram\nLM_PATH_4G   = \"/kaggle/input/4gm/pytorch/default/1/4-gram.arpa\"\nLM_PATH_3G   = \"/kaggle/input/3gm/pytorch/default/1/3-gram.pruned.1e-7.arpa\"\n\nif os.path.exists(LM_PATH_4G):\n    LM_PATH = LM_PATH_4G\n    print(f\"✅ 使用强力 LM: {LM_PATH}\")\nelse:\n    LM_PATH = LM_PATH_3G\n    print(f\"⚠️ 未找到 4-gram，回退到: {LM_PATH}\")\n\nDATA_DIR     = \"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final\"\nOUTPUT_DIR   = \"/kaggle/working\"\nDEVICE       = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# 原始字符表 (模型输出的 logits 对应这个)\nVOCAB = [' ', \"'\", 'a', 'b', 'c', 'd', 'e', 'f', 'g', 'h', 'i', 'j', 'k', 'l', 'm', 'n', 'o', 'p', 'q', 'r', 's', 't', 'u', 'v', 'w', 'x', 'y', 'z']\nvocab_list = [\"\"] + VOCAB \n\n# --------------------------\n# 2. 模型定义\n# --------------------------\nclass SpecAugment(nn.Module): \n    def forward(self, x): return x\n\nclass OfficialRNN(nn.Module):\n    def __init__(self, n_days, n_chars=len(vocab_list)):\n        super().__init__()\n        self.day_W = nn.Parameter(torch.stack([torch.eye(512) for _ in range(n_days)], dim=0))\n        self.day_b = nn.Parameter(torch.zeros(n_days, 512))\n        self.augment = SpecAugment()\n        self.conv = nn.Sequential(nn.Conv1d(512, 512, 3, 1, 1), nn.BatchNorm1d(512), nn.ReLU(), nn.Dropout(0.4))\n        self.gru = nn.GRU(512, 512, 5, bidirectional=True, batch_first=True, dropout=0.4)\n        self.fc = nn.Linear(1024, n_chars)\n\n    def forward(self, x, lengths, day_ids):\n        day_ids = day_ids.clamp(min=0, max=self.day_W.shape[0]-1)\n        x = torch.einsum(\"btc,bcd->btd\", x, self.day_W[day_ids]) + self.day_b[day_ids][:, None, :]\n        x = self.augment(x)\n        x = self.conv(x.transpose(1, 2)).transpose(1, 2)\n        total_len = x.size(1)\n        packed = pack_padded_sequence(x, lengths.cpu(), batch_first=True, enforce_sorted=True)\n        out, _ = self.gru(packed)\n        out, _ = pad_packed_sequence(out, batch_first=True, total_length=total_len)\n        return self.fc(out)\n\n# --------------------------\n# 3. 鲁棒数据加载\n# --------------------------\ndef get_session_map(data_dir):\n    print(\"🧩 Rebuilding Session Map...\")\n    sessions = set()\n    files = glob.glob(os.path.join(data_dir, \"**\", \"data_train.hdf5\"), recursive=True)\n    for fp in files:\n        try:\n            with h5py.File(fp, \"r\") as f:\n                k = list(f.keys())[0]\n                s = f[k].attrs.get(\"session\", \"\")\n                if isinstance(s, bytes): s = s.decode(\"utf-8\")\n                sessions.add(s)\n        except: pass\n    return {s: i for i, s in enumerate(sorted(list(sessions)))}\n\nclass RobustTestDataset(Dataset):\n    def __init__(self, data_dir, session_map):\n        self.files = []\n        files = sorted(glob.glob(os.path.join(data_dir, \"**\", \"data_test.hdf5\"), recursive=True))\n        print(f\"📥 Found {len(files)} test files. Scanning...\")\n        \n        valid_cnt = 0\n        for fp in files:\n            try:\n                with h5py.File(fp, \"r\") as f:\n                    for k in f.keys():\n                        g = f[k]\n                        sess = g.attrs.get(\"session\", \"\")\n                        if isinstance(sess, bytes): sess = sess.decode(\"utf-8\")\n                        sess_id = session_map.get(sess, 0)\n                        self.files.append({\"fp\": fp, \"key\": k, \"sess_id\": sess_id})\n                        valid_cnt += 1\n            except Exception as e:\n                pass\n        print(f\"✅ Loaded {valid_cnt} samples successfully.\")\n\n    def __len__(self): return len(self.files)\n    def __getitem__(self, idx):\n        item = self.files[idx]\n        with h5py.File(item['fp'], 'r') as f:\n            feats = torch.from_numpy(f[item['key']]['input_features'][:]).float()\n            if feats.ndim == 2 and feats.shape[0] == 512: feats = feats.transpose(0, 1)\n        return feats, feats.shape[0], item['sess_id'], idx\n\n    @staticmethod\n    def collate_fn(batch):\n        batch.sort(key=lambda x: x[1], reverse=True)\n        feats, lens, sess_ids, idxs = zip(*batch)\n        padded = pad_sequence(feats, batch_first=True)\n        return padded, torch.tensor(lens), torch.tensor(sess_ids), torch.tensor(idxs)\n\n# --------------------------\n# 4. 主流程\n# --------------------------\ndef main():\n    # A. 准备数据\n    session_map = get_session_map(DATA_DIR)\n    n_days = max(len(session_map), 45)\n    \n    # B. 加载模型\n    print(f\"📥 Loading Model: {MODEL_PATH}\")\n    model = OfficialRNN(n_days).to(DEVICE)\n    try:\n        ckpt = torch.load(MODEL_PATH, map_location=DEVICE)\n        sd = ckpt['state_dict'] if 'state_dict' in ckpt else ckpt\n        sd = {k.replace(\"module.\", \"\"): v for k, v in sd.items()}\n        model.load_state_dict(sd, strict=False)\n        print(\"✅ Model weights loaded.\")\n    except Exception as e:\n        print(f\"❌ Load Failed: {e}\")\n        return\n    model.eval()\n\n    # C. 智能构建解码器 (自动检测大小写)\n    print(\"🕵️ Checking LM case sensitivity...\")\n    lm_is_uppercase = False\n    try:\n        with open(LM_PATH, 'r') as f:\n            for i, line in enumerate(f):\n                if i > 200: break\n                parts = line.strip().split()\n                if len(parts) >= 2:\n                    word = parts[1]\n                    if word.isalpha() and word.isupper() and len(word) > 1:\n                        lm_is_uppercase = True\n                        print(f\"💡 Detected UPPERCASE LM (e.g., '{word}')\")\n                        break\n    except: pass\n    \n    # 调整字符表\n    decoder_vocab = list(vocab_list)\n    if lm_is_uppercase:\n        print(\"🔧 Converting vocab to UPPERCASE to match LM...\")\n        decoder_vocab = [c.upper() for c in decoder_vocab]\n\n    print(\"🔧 Building Decoder (alpha=2.2, beta=-1.0)...\")\n    # beta 不用 -5.0 了，因为 LM 生效后会自动纠错，-1.0 比较自然\n    decoder = build_ctcdecoder(\n        labels=decoder_vocab,\n        kenlm_model_path=LM_PATH,\n        alpha=2.2, \n        beta=-1.0 \n    )\n    \n    # D. 推理循环\n    test_ds = RobustTestDataset(DATA_DIR, session_map)\n    loader = DataLoader(test_ds, batch_size=32, shuffle=False, collate_fn=RobustTestDataset.collate_fn)\n    \n    N = len(test_ds)\n    preds = [\"\"] * N\n    \n    print(\"🚀 Decoding Started...\")\n    with torch.no_grad():\n        for feats, lens, sess_ids, idxs in tqdm(loader):\n            feats, sess_ids = feats.to(DEVICE), sess_ids.to(DEVICE)\n            \n            with torch.cuda.amp.autocast():\n                logits = model(feats, lens, sess_ids)\n            \n            logits_np = logits.float().cpu().numpy()\n            \n            for i, original_idx in enumerate(idxs):\n                text = decoder.decode(logits_np[i, :lens[i], :], beam_width=100)\n                # 🔥 关键：转回小写\n                preds[original_idx.item()] = text.lower()\n\n    # E. 保存\n    preds = [p if p.strip() else \"the\" for p in preds]\n    sub = pd.DataFrame({\"id\": range(N), \"text\": preds})\n    out_file = os.path.join(OUTPUT_DIR, \"submission.csv\")\n    sub.to_csv(out_file, index=False)\n    print(f\"🏆 Final Submission Saved: {out_file}\")\n    print(sub.head(10))\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T06:41:29.073690Z","iopub.execute_input":"2025-12-31T06:41:29.074450Z","iopub.status.idle":"2025-12-31T06:44:40.505642Z","shell.execute_reply.started":"2025-12-31T06:41:29.074424Z","shell.execute_reply":"2025-12-31T06:44:40.504829Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================================================================================\n# 🧪 TTA 极致优化版:低学习率微调 + 强力 LM 约束 (目标: 0.26 -> 0.23)\n**Low learning-rate fine-tuning + strong LM constraints (target: 0.26 → 0.23)**\n# ==================================================================================\nimport os, sys, glob, h5py, re\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.nn.utils.rnn import pad_sequence, pack_padded_sequence, pad_packed_sequence\n\n# --- 0. 环境依赖 ---\ntry:\n    from pyctcdecode import build_ctcdecoder\nexcept ImportError:\n    import subprocess\n    subprocess.check_call([sys.executable, \"-m\", \"pip\", \"-q\", \"install\", \"pyctcdecode\"])\n    subprocess.check_call([sys.executable, \"-m\", \"pip\", \"-q\", \"install\", \"https://github.com/kpu/kenlm/archive/master.zip\"])\n    from pyctcdecode import build_ctcdecoder\n\n# --------------------------\n# 1. 配置与路径\n# --------------------------\nMODEL_PATH    = \"/kaggle/input/my-best-gru/pytorch/default/1/best_gru_sota (1).pth\"\nSUBMISSION_CSV= \"/kaggle/input/0-25csv/submission_tta.csv\" # 你的高分伪标签\nDATA_DIR      = \"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final\"\nOUTPUT_DIR    = \"/kaggle/working\"\nDEVICE        = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# LM 配置\nLM_PATH_4G = \"/kaggle/input/4gm/pytorch/default/1/4-gram.arpa\"\nLM_PATH_3G = \"/kaggle/input/3gm/pytorch/default/1/3-gram.pruned.1e-7.arpa\"\nLM_PATH = LM_PATH_4G if os.path.exists(LM_PATH_4G) else LM_PATH_3G\n\nVOCAB = [' ', \"'\", 'a', 'b', 'c', 'd', 'e', 'f', 'g', 'h', 'i', 'j', 'k', 'l', 'm', 'n', 'o', 'p', 'q', 'r', 's', 't', 'u', 'v', 'w', 'x', 'y', 'z']\nvocab_list = [\"\"] + VOCAB\nCHAR2IDX = {c: i+1 for i, c in enumerate(VOCAB)}\n\n# --------------------------\n# 2. 模型定义\n# --------------------------\nclass SpecAugment(nn.Module): \n    def forward(self, x): return x\n\nclass OfficialRNN(nn.Module):\n    def __init__(self, n_days, n_chars=len(vocab_list)):\n        super().__init__()\n        self.day_W = nn.Parameter(torch.stack([torch.eye(512) for _ in range(n_days)], dim=0))\n        self.day_b = nn.Parameter(torch.zeros(n_days, 512))\n        self.augment = SpecAugment()\n        self.conv = nn.Sequential(nn.Conv1d(512, 512, 3, 1, 1), nn.BatchNorm1d(512), nn.ReLU(), nn.Dropout(0.4))\n        self.gru = nn.GRU(512, 512, 5, bidirectional=True, batch_first=True, dropout=0.4)\n        self.fc = nn.Linear(1024, n_chars)\n\n    def forward(self, x, lengths, day_ids):\n        day_ids = day_ids.clamp(min=0, max=self.day_W.shape[0]-1)\n        x = torch.einsum(\"btc,bcd->btd\", x, self.day_W[day_ids]) + self.day_b[day_ids][:, None, :]\n        x = self.augment(x)\n        x = self.conv(x.transpose(1, 2)).transpose(1, 2)\n        total_len = x.size(1)\n        packed = pack_padded_sequence(x, lengths.cpu(), batch_first=True, enforce_sorted=True)\n        out, _ = self.gru(packed)\n        out, _ = pad_packed_sequence(out, batch_first=True, total_length=total_len)\n        return self.fc(out)\n\n# --------------------------\n# 3. 数据集\n# --------------------------\ndef get_session_map(data_dir):\n    print(\"🧩 Rebuilding Session Map...\")\n    sessions = set()\n    files = glob.glob(os.path.join(data_dir, \"**\", \"data_train.hdf5\"), recursive=True)\n    for fp in files:\n        try:\n            with h5py.File(fp, \"r\") as f:\n                s = f[list(f.keys())[0]].attrs.get(\"session\", \"\")\n                if isinstance(s, bytes): s = s.decode(\"utf-8\")\n                sessions.add(s)\n        except: pass\n    return {s: i for i, s in enumerate(sorted(list(sessions)))}\n\nclass PseudoLabelDataset(Dataset):\n    def __init__(self, data_dir, pseudo_map, session_map):\n        self.files = []\n        files = sorted(glob.glob(os.path.join(data_dir, \"**\", \"data_test.hdf5\"), recursive=True))\n        print(f\"📥 Scanning {len(files)} files for TTA...\")\n        self.session_map = session_map\n        self.pseudo_map = pseudo_map\n        valid_cnt = 0\n        for fp in files:\n            try:\n                with h5py.File(fp, \"r\") as f:\n                    for k in f.keys():\n                        g = f[k]\n                        sess = g.attrs.get(\"session\", \"\")\n                        if isinstance(sess, bytes): sess = sess.decode(\"utf-8\")\n                        sess_id = session_map.get(sess, 0)\n                        self.files.append({\"fp\": fp, \"key\": k, \"sess_id\": sess_id})\n                        valid_cnt += 1\n            except Exception as e:\n                print(f\"⚠️ Error reading {fp}: {e}\")\n        print(f\"✅ Loaded {valid_cnt} samples for Fine-tuning.\")\n        if valid_cnt == 0: raise ValueError(\"❌ No samples loaded!\")\n\n    def __len__(self): return len(self.files)\n    def __getitem__(self, idx):\n        item = self.files[idx]\n        with h5py.File(item['fp'], 'r') as f:\n            feats = torch.from_numpy(f[item['key']]['input_features'][:]).float()\n            if feats.ndim == 2 and feats.shape[0] == 512: feats = feats.transpose(0, 1)\n        text = self.pseudo_map.get(idx, \"\")\n        if not isinstance(text, str): text = \"\"\n        label_idxs = [CHAR2IDX[c] for c in text if c in CHAR2IDX]\n        label = torch.tensor(label_idxs, dtype=torch.long)\n        return feats, label, item['sess_id'], idx\n\n    @staticmethod\n    def collate_fn(batch):\n        batch.sort(key=lambda x: x[0].shape[0], reverse=True)\n        feats, labels, sess_ids, idxs = zip(*batch)\n        padded = pad_sequence(feats, batch_first=True)\n        in_len = torch.tensor([f.shape[0] for f in feats], dtype=torch.long)\n        tgt_len = torch.tensor([l.numel() for l in labels], dtype=torch.long)\n        flat_targets = torch.cat(labels) if int(tgt_len.sum()) > 0 else torch.empty(0, dtype=torch.long)\n        return padded, flat_targets, in_len, tgt_len, torch.tensor(sess_ids), torch.tensor(idxs)\n\n# --------------------------\n# 4. 主流程 (参数调优版)\n# --------------------------\ndef main():\n    # A. 准备\n    if not os.path.exists(SUBMISSION_CSV):\n        print(f\"❌ Error: {SUBMISSION_CSV} not found.\")\n        return\n\n    df = pd.read_csv(SUBMISSION_CSV)\n    df = df.sort_values(\"id\").reset_index(drop=True)\n    pseudo_map = dict(zip(df.id, df.text))\n    print(f\"📖 Loaded {len(pseudo_map)} pseudo-labels (Source: {os.path.basename(SUBMISSION_CSV)})\")\n\n    session_map = get_session_map(DATA_DIR)\n    n_days = max(len(session_map), 45)\n    \n    # B. 加载模型\n    print(f\"📥 Reloading Base Model: {MODEL_PATH}\")\n    model = OfficialRNN(n_days).to(DEVICE)\n    # Weights Only = False 修复\n    try:\n        ckpt = torch.load(MODEL_PATH, map_location=DEVICE, weights_only=False) \n    except TypeError:\n        ckpt = torch.load(MODEL_PATH, map_location=DEVICE) # 旧版 pytorch 兼容\n        \n    sd = ckpt['state_dict'] if 'state_dict' in ckpt else ckpt\n    sd = {k.replace(\"module.\", \"\"): v for k, v in sd.items()}\n    model.load_state_dict(sd, strict=False)\n    \n    # C. TTA 微调设置\n    model.train()\n    print(\"❄️ Freezing layers (Training ONLY Day Adapters)...\")\n    trainable_params = []\n    for name, param in model.named_parameters():\n        if \"day_\" in name: \n            param.requires_grad = True\n            trainable_params.append(param)\n        else:\n            param.requires_grad = False\n    \n    # 🔥 优化点 1: 降低学习率，增加权重衰减 (更稳健)\n    # LR: 1e-4 -> 5e-5 (防止改过头)\n    # Weight Decay: 0.01 (防止过拟合噪音标签)\n    optimizer = optim.AdamW(trainable_params, lr=5e-5, weight_decay=0.01) \n    criterion = nn.CTCLoss(blank=0, zero_infinity=True)\n    scaler = torch.cuda.amp.GradScaler()\n\n    tta_ds = PseudoLabelDataset(DATA_DIR, pseudo_map, session_map)\n    tta_loader = DataLoader(tta_ds, batch_size=32, shuffle=True, collate_fn=PseudoLabelDataset.collate_fn)\n\n    # 🔥 优化点 2: 增加 Epoch 到 4 (因为学习率低了，可以多跑一轮)\n    EPOCHS = 4\n    print(f\"🚀 Starting Test-Time Adaptation ({EPOCHS} Epochs)...\")\n    for epoch in range(EPOCHS):\n        total_loss = 0\n        steps = 0\n        pbar = tqdm(tta_loader, desc=f\"TTA Ep {epoch+1}/{EPOCHS}\")\n        for feats, targets, in_len, tgt_len, sess_ids, _ in pbar:\n            feats, targets, sess_ids = feats.to(DEVICE), targets.to(DEVICE), sess_ids.to(DEVICE)\n            if targets.numel() == 0: continue\n            \n            optimizer.zero_grad()\n            with torch.cuda.amp.autocast():\n                logits = model(feats, in_len, sess_ids)\n                log_probs = logits.log_softmax(2).transpose(0, 1)\n                loss = criterion(log_probs, targets, in_len, tgt_len)\n            \n            scaler.scale(loss).backward()\n            \n            # 🔥 优化点 3: 梯度裁剪 (防止梯度爆炸)\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(trainable_params, max_norm=1.0)\n            \n            scaler.step(optimizer)\n            scaler.update()\n            \n            total_loss += loss.item()\n            steps += 1\n            pbar.set_postfix(loss=f\"{loss.item():.4f}\")\n\n    print(\"✅ TTA Fine-tuning Done!\")\n    \n    # E. 最终推理\n    model.eval()\n    if torch.cuda.device_count() > 1:\n        model = nn.DataParallel(model)\n        \n    print(\"🔧 Building Decoder...\")\n    lm_is_upper = False\n    try:\n        with open(LM_PATH, 'r') as f:\n            for i, line in enumerate(f):\n                if i>200: break\n                parts = line.strip().split()\n                if len(parts)>=2 and parts[1].isupper(): \n                    lm_is_upper=True; break\n    except: pass\n    \n    d_vocab = [c.upper() for c in vocab_list] if lm_is_upper else list(vocab_list)\n    print(f\"🔧 Decoder Vocab Upper: {lm_is_upper}\")\n    \n    # 🔥 优化点 4: 强力解码参数 (解决 uprice/countin)\n    # alpha: 2.2 -> 2.8 (更听 LM 的话，修复拼写)\n    # beta: -1.0 -> -2.5 (惩罚短词和乱码插入)\n    # beam_width: 100 -> 200 (搜索更深，反正只推一次)\n    print(\"🔥 Using Optimized Decoder: alpha=2.8, beta=-2.5, beam=200\")\n    decoder = build_ctcdecoder(d_vocab, LM_PATH, alpha=2.8, beta=-2.5)\n    \n    infer_loader = DataLoader(tta_ds, batch_size=64, shuffle=False, collate_fn=PseudoLabelDataset.collate_fn) \n    preds = [\"\"] * len(tta_ds)\n    \n    print(\"🚀 Final Decoding with Adapted Model...\")\n    with torch.no_grad():\n        for feats, _, in_len, _, sess_ids, idxs in tqdm(infer_loader):\n            feats, sess_ids = feats.to(DEVICE), sess_ids.to(DEVICE)\n            with torch.cuda.amp.autocast():\n                logits = model(feats, in_len, sess_ids)\n            \n            logits_np = logits.float().cpu().numpy()\n            for i, original_idx in enumerate(idxs):\n                text = decoder.decode(logits_np[i, :in_len[i], :], beam_width=200)\n                preds[original_idx.item()] = text.lower()\n                \n    preds = [p if p.strip() else \"the\" for p in preds]\n    sub = pd.DataFrame({\"id\": range(len(preds)), \"text\": preds})\n    out_path = os.path.join(OUTPUT_DIR, \"submission_tta_optimized.csv\")\n    sub.to_csv(out_path, index=False)\n    print(f\"🏆 Optimized TTA Submission Saved: {out_path}\")\n    print(sub.head())\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T09:05:24.492017Z","iopub.execute_input":"2025-12-31T09:05:24.492310Z","iopub.status.idle":"2025-12-31T09:14:15.860849Z","shell.execute_reply.started":"2025-12-31T09:05:24.492290Z","shell.execute_reply":"2025-12-31T09:14:15.859957Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#直接使用强大的llm去润色，我最终是使用了的deepseekv3去修改提升了0.02分\n**In the final stage, DeepSeek V3 was used for post-editing, resulting in an additional +0.02 improvement in score.**\n#!/usr/bin/env python3\n\"\"\"\nBrain-to-Text CSV Optimizer using DeepSeek API\nOptimizes decoded brain signal text with high-confidence corrections only.\n\"\"\"\n\nimport csv\nimport json\nimport time\nimport requests\nfrom datetime import datetime\nfrom pathlib import Path\n\n# Configuration\nDEEPSEEK_API_KEY = \"\"#fill your api in or you can install this model in local（i dont konw if it can use in kaggle ） \nDEEPSEEK_API_URL = \"https://api.deepseek.com/v1/chat/completions\"\nMODEL = \"deepseek-chat\"\n\n#这些都是我在本地使用的 these csv all in local\nINPUT_FILE = Path(__file__).parent / \"submission_improved_v2.csv\"#change\nOUTPUT_FILE = Path(__file__).parent / \"submission_deepseek_optimized.csv\"#change\nLOG_FILE = Path(__file__).parent / \"deepseek_correction_log.txt\"\n\n# Rate limiting\nREQUEST_DELAY = 0.3  # seconds between requests\nMAX_RETRIES = 3\nBATCH_SIZE = 5  # Process in batches for logging\n\nSYSTEM_PROMPT = \"\"\"You are a text correction specialist for brain-to-text decoding. \nYour task is to correct ONLY obvious errors while preserving the original meaning.\n\nSTRICT RULES:\n1. Fix clear spelling errors (e.g., \"epecion\" → \"education\", \"beautifull\" → \"beautiful\")\n2. Fix obvious missing/wrong words (e.g., \"I to go\" → \"I want to go\")\n3. Fix grammar issues that clearly don't match intended meaning\n4. PRESERVE the sentence if it's too corrupted to understand with HIGH confidence\n5. DO NOT add substantial new content\n6. DO NOT change the meaning of the sentence\n7. If uncertain, return the original sentence unchanged\n\nReturn ONLY a JSON object (no markdown, no explanation):\n{\"corrected\": \"the corrected sentence\", \"confidence\": \"high\" or \"low\", \"changes\": \"brief description or none\"}\n\nIf you cannot confidently correct it, return:\n{\"corrected\": \"original sentence here\", \"confidence\": \"low\", \"changes\": \"none\"}\"\"\"\n\n\ndef call_deepseek_api(sentence: str) -> dict:\n    \"\"\"Call DeepSeek API to correct a sentence.\"\"\"\n    headers = {\n        \"Authorization\": f\"Bearer {DEEPSEEK_API_KEY}\",\n        \"Content-Type\": \"application/json\"\n    }\n    \n    payload = {\n        \"model\": MODEL,\n        \"messages\": [\n            {\"role\": \"system\", \"content\": SYSTEM_PROMPT},\n            {\"role\": \"user\", \"content\": f\"Correct this sentence: {sentence}\"}\n        ],\n        \"temperature\": 0.1,  # Low temperature for consistent corrections\n        \"max_tokens\": 200\n    }\n    \n    for attempt in range(MAX_RETRIES):\n        try:\n            response = requests.post(DEEPSEEK_API_URL, headers=headers, json=payload, timeout=30)\n            response.raise_for_status()\n            \n            result = response.json()\n            content = result[\"choices\"][0][\"message\"][\"content\"].strip()\n            \n            # Parse JSON response\n            # Handle potential markdown code blocks\n            if content.startswith(\"```\"):\n                content = content.split(\"```\")[1]\n                if content.startswith(\"json\"):\n                    content = content[4:]\n                content = content.strip()\n            \n            parsed = json.loads(content)\n            return parsed\n            \n        except json.JSONDecodeError as e:\n            print(f\"  JSON parse error: {e}, raw: {content[:100]}\")\n            return {\"corrected\": sentence, \"confidence\": \"low\", \"changes\": \"parse_error\"}\n        except requests.exceptions.RequestException as e:\n            print(f\"  API error (attempt {attempt + 1}): {e}\")\n            if attempt < MAX_RETRIES - 1:\n                time.sleep(2 ** attempt)  # Exponential backoff\n            else:\n                return {\"corrected\": sentence, \"confidence\": \"low\", \"changes\": \"api_error\"}\n    \n    return {\"corrected\": sentence, \"confidence\": \"low\", \"changes\": \"max_retries\"}\n\n\ndef process_csv():\n    \"\"\"Process the CSV file and apply high-confidence corrections.\"\"\"\n    print(f\"Starting optimization at {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\")\n    print(f\"Input: {INPUT_FILE}\")\n    print(f\"Output: {OUTPUT_FILE}\")\n    print(\"-\" * 60)\n    \n    # Read input CSV\n    rows = []\n    with open(INPUT_FILE, 'r', encoding='utf-8') as f:\n        reader = csv.DictReader(f)\n        rows = list(reader)\n    \n    total = len(rows)\n    print(f\"Total sentences to process: {total}\")\n    \n    # Open log file\n    log_entries = []\n    corrections_made = 0\n    high_confidence_count = 0\n    \n    # Process each sentence\n    for i, row in enumerate(rows):\n        row_id = row['id']\n        original_text = row['text']\n        \n        # Skip empty sentences\n        if not original_text or original_text.strip() == \"\":\n            continue\n        \n        # Call API\n        result = call_deepseek_api(original_text)\n        \n        corrected = result.get(\"corrected\", original_text)\n        confidence = result.get(\"confidence\", \"low\")\n        changes = result.get(\"changes\", \"none\")\n        \n        # Only apply high-confidence corrections\n        if confidence == \"high\" and corrected != original_text and changes != \"none\":\n            row['text'] = corrected\n            corrections_made += 1\n            high_confidence_count += 1\n            \n            log_entry = f\"[ID {row_id}] HIGH CONFIDENCE\\n  Original: {original_text}\\n  Corrected: {corrected}\\n  Changes: {changes}\\n\"\n            log_entries.append(log_entry)\n            print(f\"[{i+1}/{total}] ID {row_id}: CORRECTED - {changes}\")\n        else:\n            if confidence == \"high\":\n                high_confidence_count += 1\n            if (i + 1) % 50 == 0:\n                print(f\"[{i+1}/{total}] Processing... ({corrections_made} corrections so far)\")\n        \n        # Rate limiting\n        time.sleep(REQUEST_DELAY)\n    \n    # Write output CSV\n    with open(OUTPUT_FILE, 'w', encoding='utf-8', newline='') as f:\n        writer = csv.DictWriter(f, fieldnames=['id', 'text'])\n        writer.writeheader()\n        writer.writerows(rows)\n    \n    # Write log file\n    with open(LOG_FILE, 'w', encoding='utf-8') as f:\n        f.write(f\"DeepSeek Optimization Log\\n\")\n        f.write(f\"Generated: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\\n\")\n        f.write(f\"Total sentences: {total}\\n\")\n        f.write(f\"High confidence responses: {high_confidence_count}\\n\")\n        f.write(f\"Corrections applied: {corrections_made}\\n\")\n        f.write(\"=\" * 60 + \"\\n\\n\")\n        f.write(\"\\n\".join(log_entries))\n    \n    print(\"-\" * 60)\n    print(f\"COMPLETED!\")\n    print(f\"Total sentences: {total}\")\n    print(f\"High confidence responses: {high_confidence_count}\")\n    print(f\"Corrections applied: {corrections_made}\")\n    print(f\"Output saved to: {OUTPUT_FILE}\")\n    print(f\"Log saved to: {LOG_FILE}\")\n\n\nif __name__ == \"__main__\":\n    process_csv()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}