{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":106809,"databundleVersionId":13056355,"sourceType":"competition"},{"sourceId":14126434,"sourceType":"datasetVersion","datasetId":9000582},{"sourceId":681234,"sourceType":"modelInstanceVersion","modelInstanceId":516905,"modelId":531567}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"file_path = \"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final/t15.2023.08.11/data_train.hdf5\"\n\nimport h5py, numpy as np\n\nwith h5py.File(file_path, \"r\") as f:\n    k = list(f.keys())[0]\n    g = f[k]\n    print(\"key:\", k)\n    print(\"input_features shape:\", g[\"input_features\"].shape)\n    print(\"n_time_steps:\", g.attrs[\"n_time_steps\"])\n    print(\"seq_class_ids shape:\", g[\"seq_class_ids\"][:].shape)\n    print(\"seq_len:\", g.attrs[\"seq_len\"])\n    print(\"label min/max:\", g[\"seq_class_ids\"][:].min(), g[\"seq_class_ids\"][:].max())\n    print(\"session:\", g.attrs[\"session\"])\n    print(\"block_num:\", g.attrs[\"block_num\"], \"trial_num:\", g.attrs[\"trial_num\"])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:21:50.507247Z","iopub.execute_input":"2025-12-13T10:21:50.507532Z","iopub.status.idle":"2025-12-13T10:21:50.699540Z","shell.execute_reply.started":"2025-12-13T10:21:50.507506Z","shell.execute_reply":"2025-12-13T10:21:50.698732Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport glob\nimport h5py\nimport numpy as np\nfrom collections import Counter\nfrom tqdm import tqdm\n\n# ======= 1. 设置根目录（这里按你当前环境改一下） =======\n# Kaggle:\nROOT_DIR = \"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData\"\n# AutoDL:\n# ROOT_DIR = \"/root/autodl-tmp/t15_copyTask_neuralData\"\n\n# ======= 2. 找到所有 hdf5 文件 =======\nh5_files = glob.glob(os.path.join(ROOT_DIR, \"**\", \"*.hdf5\"), recursive=True)\nprint(f\"📂 找到 HDF5 文件数: {len(h5_files)}\")\nfor p in h5_files[:5]:\n    print(\"  示例文件:\", p)\n\n# ======= 3. 统计器初始化 =======\ntotal_trials = 0\ntrials_with_labels = 0\n\nsplit_counter = Counter()   # train / val / test\nsession_counter = Counter()\nblock_counter = Counter()\n\ntime_steps_list = []        # 每条 trial 的 T\nfeat_dim_set = set()        # 特征维度集合\n\nn_time_mismatch = 0         # n_time_steps attr 和 实际 T 不一致计数\nmismatch_examples = []\n\nseq_len_list = []           # attr 'seq_len'\nraw_label_len_list = []     # seq_class_ids 的原始长度（通常 500）\nlabel_min_list = []\nlabel_max_list = []\n\nzero_in_firstL = 0          # 0 出现在前 seq_len 里的 trial 数\nzero_only_in_padding = 0    # 0 只在 padding 段（L:）出现\ntrials_check_zero = 0\n\n# ======= 4. 遍历所有文件 / trial 做统计 =======\nfor f_path in tqdm(h5_files, desc=\"Scanning HDF5\"):\n    # 粗略判断一下 split\n    if \"data_train\" in f_path:\n        split = \"train\"\n    elif \"data_val\" in f_path:\n        split = \"val\"\n    elif \"data_test\" in f_path:\n        split = \"test\"\n    else:\n        split = \"unknown\"\n\n    try:\n        with h5py.File(f_path, \"r\") as f:\n            keys = list(f.keys())\n            for k in keys:\n                g = f[k]\n                total_trials += 1\n                split_counter[split] += 1\n\n                # ---------- 4.1 读取特征 ----------\n                raw_feats = g[\"input_features\"][:]  # numpy array\n                feats = np.asarray(raw_feats)\n                shape = feats.shape  # (T, 512) 或 (512, T)\n                \n                # 智能判断，把所有东西统一为 (T, 512)\n                if shape[1] == 512:\n                    T = shape[0]\n                    D = shape[1]\n                elif shape[0] == 512:\n                    T = shape[1]\n                    D = shape[0]\n                else:\n                    # 非预期形状，记录一下\n                    mismatch_examples.append((f_path, k, shape))\n                    continue\n\n                time_steps_list.append(T)\n                feat_dim_set.add(D)\n\n                # n_time_steps attr 检查\n                n_attr = g.attrs.get(\"n_time_steps\", None)\n                if n_attr is not None and int(n_attr) != int(T):\n                    n_time_mismatch += 1\n                    if len(mismatch_examples) < 10:\n                        mismatch_examples.append(\n                            (f_path, k, shape, int(n_attr))\n                        )\n\n                # session / block 信息\n                sess = g.attrs.get(\"session\", None)\n                if isinstance(sess, bytes):\n                    sess = sess.decode(\"utf-8\")\n                if sess is not None:\n                    session_counter[sess] += 1\n\n                blk = g.attrs.get(\"block_num\", None)\n                if blk is not None:\n                    block_counter[int(blk)] += 1\n\n                # ---------- 4.2 label 相关（只在 train/val 有） ----------\n                if \"seq_class_ids\" in g:\n                    trials_with_labels += 1\n                    raw_labels = g[\"seq_class_ids\"][:]  # (500,)\n                    raw_label_len_list.append(len(raw_labels))\n\n                    if len(raw_labels) > 0:\n                        label_min_list.append(int(raw_labels.min()))\n                        label_max_list.append(int(raw_labels.max()))\n\n                    # 用 seq_len attr 截取真实长度\n                    L = g.attrs.get(\"seq_len\", None)\n                    if L is not None:\n                        L = int(L)\n                    else:\n                        L = len(raw_labels)\n\n                    seq_len_list.append(L)\n\n                    # 检查 0 在不在前 L 段\n                    if len(raw_labels) > 0 and L > 0:\n                        trials_check_zero += 1\n                        firstL = raw_labels[:L]\n                        tail = raw_labels[L:]\n\n                        has_zero_first = np.any(firstL == 0)\n                        has_zero_tail = np.any(tail == 0)\n\n                        if has_zero_first:\n                            zero_in_firstL += 1\n                        if (not has_zero_first) and has_zero_tail:\n                            zero_only_in_padding += 1\n\n    except Exception as e:\n        print(f\"⚠️ 读取文件失败: {f_path}, error={e}\")\n\n# ======= 5. 汇总统计结果 =======\nprint(\"\\n================ 数据整体情况 ================\")\nprint(f\"总 trial 数: {total_trials}\")\nprint(f\"带标签的 trial 数: {trials_with_labels}\")\nprint(\"按 split 统计:\", dict(split_counter))\nprint(f\"特征维度集合 feat_dim_set: {feat_dim_set}\")\n\nif time_steps_list:\n    ts = np.array(time_steps_list)\n    print(\"\\n—— 时间步 T 分布 ——\")\n    print(f\"  min / mean / median / max = {ts.min()} / {ts.mean():.1f} / {np.median(ts)} / {ts.max()}\")\n    print(\"  25/75 分位数:\", np.percentile(ts, [25, 75]))\n\nprint(f\"\\n n_time_steps attr 与真实长度不一致的 trial 数: {n_time_mismatch}\")\nif n_time_mismatch > 0:\n    print(\"  示例几条 mismatch:\")\n    for item in mismatch_examples[:5]:\n        print(\"   \", item)\n\nif seq_len_list:\n    sl = np.array(seq_len_list)\n    print(\"\\n—— seq_len (真实 phoneme 序列长度) 分布 ——\")\n    print(f\"  min / mean / median / max = {sl.min()} / {sl.mean():.1f} / {np.median(sl)} / {sl.max()}\")\n    print(\"  25/75 分位数:\", np.percentile(sl, [25, 75]))\n\nif raw_label_len_list:\n    rll = np.array(raw_label_len_list)\n    print(\"\\n—— seq_class_ids 原始长度 (通常应该是 padding 后长度) ——\")\n    print(f\"  唯一长度集合: {set(raw_label_len_list)}\")\n    print(f\"  min / mean / max = {rll.min()} / {rll.mean():.1f} / {rll.max()}\")\n\nif label_min_list and label_max_list:\n    print(\"\\n—— label 取值范围 ——\")\n    print(f\"  全局 label min: {min(label_min_list)}, max: {max(label_max_list)}\")\n    print(\"  按 trial 统计的 min 分布 (前几个):\", sorted(label_min_list)[:10])\n    print(\"  按 trial 统计的 max 分布 (前几个):\", sorted(label_max_list)[-10:])\n\nif trials_check_zero > 0:\n    print(\"\\n—— label=0 是否只出现在 padding ——\")\n    print(f\"  检查了 {trials_check_zero} 条带 label 的 trial\")\n    print(f\"  有 0 出现在前 seq_len 部分的 trial 数: {zero_in_firstL}\")\n    print(f\"  只在 padding 段出现 0 的 trial 数: {zero_only_in_padding}\")\n\nprint(\"\\n—— Session 分布 (前 10 个按 trial 数排序) ——\")\nfor sess, cnt in session_counter.most_common(10):\n    print(f\"  {sess}: {cnt} trials\")\n\nprint(\"\\n—— Block_Num 分布 (前 10 个按 trial 数排序) ——\")\nfor blk, cnt in sorted(block_counter.items())[:10]:\n    print(f\"  block {blk}: {cnt} trials\")\n\nprint(\"\\n✅ EDA 完成。\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T08:08:37.884239Z","iopub.execute_input":"2025-12-13T08:08:37.884483Z","iopub.status.idle":"2025-12-13T08:12:47.814428Z","shell.execute_reply.started":"2025-12-13T08:08:37.884463Z","shell.execute_reply":"2025-12-13T08:12:47.813796Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====== ONE CELL: train GRU-CTC + 3gram+lexicon decode + write submission ======\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\n\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\n\n# --------------------------\n# 0) Config (Kaggle paths)\n# --------------------------\nDATA_DIR   = \"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final\"\nOUTPUT_DIR = \"/kaggle/working\"\nos.makedirs(OUTPUT_DIR, exist_ok=True)\n\nLEXICON_PATH = \"/kaggle/input/lexicon/librispeech-lexicon.txt\"\nLM_PATH      = \"/kaggle/input/3gm/pytorch/default/1/3-gram.pruned.1e-7.arpa\"\n\nBATCH_SIZE  = 64\nNUM_EPOCHS  = 30\nNUM_WORKERS = 2\n\nBEAM_SIZE      = 150\nBEAM_THRESHOLD = 15.0\nLM_WEIGHT      = 1.0\nWORD_SCORE     = 1.0\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"🚀 Device: {DEVICE}  (gpus={torch.cuda.device_count()})\")\n\ntorch.backends.cudnn.benchmark = True\n\n# --------------------------\n# 1) Vocab (char CTC)\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'\n]\nCHAR2IDX = {c: i+1 for i, c in enumerate(VOCAB)}  # 0 is blank\nIDX2CHAR = {i+1: c for i, c in enumerate(VOCAB)}\nIDX2CHAR[0] = \"\"  # blank\n\n# normalize output text to: lowercase words separated by single spaces\n_re_multi_space = re.compile(r\"\\s+\")\n_re_keep = re.compile(r\"[^a-z0-9'\\s]\")\n\ndef normalize_text(s: str) -> str:\n    s = (s or \"\").strip().lower()\n    s = _re_keep.sub(\" \", s)\n    s = _re_multi_space.sub(\" \", s).strip()\n    return s\n\n# --------------------------\n# 2) Dataset\n# --------------------------\nclass BrainDataset(Dataset):\n    def __init__(self, data_dir, split=\"train\", session_map=None):\n        search_pattern = os.path.join(data_dir, \"**\", f\"data_{split}.hdf5\")\n        if split == \"test\":\n            search_pattern = os.path.join(data_dir, \"**\", \"data_test.hdf5\")\n\n        files = sorted(glob.glob(search_pattern, recursive=True))\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        # Build session map (lightweight scan: first group of each file)\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:\n                    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        for fp in tqdm(files):\n            try:\n                with h5py.File(fp, \"r\") as f:\n                    for k in f.keys():\n                        g = f[k]\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)  # (T,512)\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                            # char indices (keep only in vocab)\n                            idxs = [CHAR2IDX[c] for c in sent if c in CHAR2IDX]\n                            label = torch.tensor(idxs, dtype=torch.long)\n\n                        self.data.append({\"key\": k, \"feats\": feats, \"label\": label, \"sess_id\": sess_id})\n            except:\n                pass\n\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        # sort by length desc for pack_padded_sequence(enforce_sorted=True)\n        batch.sort(key=lambda x: x[\"feats\"].shape[0], reverse=True)\n\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)  # (B,T,512)\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# 3) Model: GRU + Day Adapter + Conv (your structure, but vectorized adapter)\n# --------------------------\nclass OfficialRNN(nn.Module):\n    def __init__(self, n_days, n_chars=len(VOCAB)+1):\n        super().__init__()\n        # day-specific linear: x @ W_day + b_day\n        W = torch.stack([torch.eye(512) for _ in range(n_days)], dim=0)   # (D,512,512)\n        b = torch.zeros(n_days, 512)                                      # (D,512)\n        self.day_W = nn.Parameter(W)\n        self.day_b = nn.Parameter(b)\n\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        self.gru = nn.GRU(\n            512, 512, num_layers=5, bidirectional=True,\n            batch_first=True, dropout=0.4\n        )\n        self.fc = nn.Linear(1024, n_chars)\n\n    def forward(self, x, lengths, day_ids):\n        # x: (B,T,512)\n        day_ids = day_ids.clamp(min=0, max=self.day_W.shape[0]-1)\n        W = self.day_W[day_ids]                 # (B,512,512)\n        b = self.day_b[day_ids]                 # (B,512)\n        x = torch.einsum(\"btc,bcd->btd\", x, W) + b[:, None, :]\n\n        x = x.transpose(1, 2)                   # (B,512,T)\n        x = self.conv(x)\n        x = x.transpose(1, 2)                   # (B,T,512)\n\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)\n        return self.fc(out)                     # (B,T,C)\n\n# --------------------------\n# 4) Decoder: build char-lexicon from librispeech word->phones + 3gram word LM\n# --------------------------\n# torchaudio CTC decoder needs flashlight-text\nif importlib.util.find_spec(\"flashlight\") is None:\n    print(\"Installing flashlight-text ...\")\n    subprocess.check_call([sys.executable, \"-m\", \"pip\", \"-q\", \"install\", \"flashlight-text\"])\n\nfrom torchaudio.models.decoder import ctc_decoder\n\nTOKENS = [\"BLANK\"] + VOCAB  # must match model class order\n\ndef build_char_lexicon(src_lex: str, out_lex: str):\n    seen = set()\n    kept, total = 0, 0\n    with open(src_lex, \"r\", encoding=\"utf-8\", errors=\"ignore\") as f, \\\n         open(out_lex, \"w\", encoding=\"utf-8\") as o:\n        for line in f:\n            line = line.strip()\n            if not line:\n                continue\n            parts = line.split()\n            if len(parts) < 2:\n                continue\n            total += 1\n\n            w = parts[0].lower()\n            w = re.sub(r\"\\(\\d+\\)$\", \"\", w)  # remove pronunciation index like WORD(2)\n\n            if not w or w in seen:\n                continue\n\n            letters = []\n            ok = True\n            for ch in w:\n                if ch in [\"'\", *list(\"abcdefghijklmnopqrstuvwxyz\")]:\n                    letters.append(ch)\n                else:\n                    ok = False\n                    break\n            if not ok or len(letters) == 0:\n                continue\n\n            o.write(w + \" \" + \" \".join(letters) + \"\\n\")\n            kept += 1\n            seen.add(w)\n\n    print(f\"🔤 char-lexicon: kept={kept}/{total} -> {out_lex}\")\n\ndef build_decoder(work_dir: str):\n    lex_out = os.path.join(work_dir, \"lexicon_char.txt\")\n    build_char_lexicon(LEXICON_PATH, lex_out)\n    dec = ctc_decoder(\n        lexicon=lex_out,\n        tokens=TOKENS,\n        lm=LM_PATH,\n        nbest=1,\n        beam_size=BEAM_SIZE,\n        beam_threshold=BEAM_THRESHOLD,\n        lm_weight=LM_WEIGHT,\n        word_score=WORD_SCORE,\n        sil_score=0.0,\n        blank_token=\"BLANK\",\n        sil_token=\" \",      # space is in VOCAB\n        unk_word=\"<unk>\",\n    )\n    return dec\n\ndef hyp_to_text(h):\n    if h is None:\n        return \"\"\n    if hasattr(h, \"words\") and h.words is not None:\n        return \" \".join(h.words)\n    if hasattr(h, \"text\"):\n        return str(h.text)\n    return str(h)\n\n# --------------------------\n# 5) Train + Decode + Submission\n# --------------------------\ndef main():\n    # Train\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,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    n_days = len(train_ds.session_map) if train_ds.session_map else 45\n    model = OfficialRNN(n_days).to(DEVICE)\n\n    optimizer = optim.AdamW(model.parameters(), lr=1e-3, weight_decay=1e-2)\n    scheduler = optim.lr_scheduler.OneCycleLR(\n        optimizer, max_lr=1e-3,\n        steps_per_epoch=len(train_loader), epochs=NUM_EPOCHS\n    )\n    criterion = nn.CTCLoss(blank=0, zero_infinity=True)\n    scaler = GradScaler(enabled=(DEVICE.type == \"cuda\"))\n\n    best_loss = float(\"inf\")\n    best_path = os.path.join(OUTPUT_DIR, \"best_gru.pth\")\n\n    print(\"🚀 Training GRU-CTC...\")\n    for epoch in range(NUM_EPOCHS):\n        model.train()\n        total = 0.0\n        pbar = tqdm(train_loader, desc=f\"Ep {epoch+1}/{NUM_EPOCHS}\")\n        for feats, targets, in_len, tgt_len, sess_ids, _ in pbar:\n            feats = feats.to(DEVICE, non_blocking=True)\n            targets = targets.to(DEVICE, non_blocking=True)\n            sess_ids = sess_ids.to(DEVICE, non_blocking=True)\n\n            optimizer.zero_grad(set_to_none=True)\n            with autocast(enabled=(DEVICE.type == \"cuda\")):\n                logits = model(feats, in_len, sess_ids)  # (B,T,C)\n                log_probs = logits.log_softmax(2).permute(1, 0, 2)  # (T,B,C)\n                loss = criterion(log_probs, targets, in_len, tgt_len)\n\n            scaler.scale(loss).backward()\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            scaler.step(optimizer)\n            scaler.update()\n            scheduler.step()\n\n            total += float(loss.item())\n            pbar.set_postfix(loss=f\"{loss.item():.4f}\")\n\n        avg = total / max(1, len(train_loader))\n        print(f\"✅ Epoch {epoch+1} avg_loss={avg:.4f}\")\n\n        if avg < best_loss:\n            best_loss = avg\n            torch.save(model.state_dict(), best_path)\n            print(f\"💾 saved best -> {best_path} (loss={best_loss:.4f})\")\n\n    # Load best for inference\n    model.load_state_dict(torch.load(best_path, map_location=\"cpu\"))\n    model.to(DEVICE).eval()\n    print(\"✅ Loaded best ckpt:\", best_path)\n\n    # Build decoder once\n    decoder = build_decoder(OUTPUT_DIR)\n\n    # Test decode -> submission\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,\n        num_workers=NUM_WORKERS, pin_memory=True\n    )\n\n    print(\"🚀 Decoding TEST with 3-gram LM + lexicon ...\")\n    results = []\n    with torch.no_grad():\n        for feats, _, in_len, _, sess_ids, keys in tqdm(test_loader, desc=\"TEST\"):\n            feats = feats.to(DEVICE, non_blocking=True)\n            sess_ids = sess_ids.to(DEVICE, non_blocking=True)\n\n            logits = model(feats, in_len, sess_ids)          # (B,T,C)\n            logp = logits.log_softmax(2).cpu()               # decoder wants CPU\n            hypos = decoder(logp, in_len.cpu())              # list[B][nbest]\n\n            for i in range(len(keys)):\n                hyp = hypos[i][0] if (len(hypos[i]) > 0) else None\n                text = normalize_text(hyp_to_text(hyp))\n                results.append({\"key\": keys[i], \"text\": text})\n\n    # Kaggle usually wants id=0..N-1\n    df = pd.DataFrame(results)\n    df[\"id\"] = range(len(df))\n    sub_path = os.path.join(OUTPUT_DIR, \"submission.csv\")\n    df[[\"id\", \"text\"]].to_csv(sub_path, index=False)\n    print(\"✅ submission saved ->\", sub_path)\n\n    # quick preview (print a few decoded sentences)\n    print(\"\\n🧾 Preview decoded (first 30):\")\n    for i in range(min(30, len(df))):\n        print(f\"[{i:04d}] key={df.loc[i,'key']}  text='{df.loc[i,'text']}'\")\n\nmain()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T14:22:21.527736Z","iopub.execute_input":"2025-12-15T14:22:21.528291Z","iopub.status.idle":"2025-12-15T16:10:21.995327Z","shell.execute_reply.started":"2025-12-15T14:22:21.528264Z","shell.execute_reply":"2025-12-15T16:10:21.994299Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install pyctcdecode\n!pip install https://github.com/kpu/kenlm/archive/master.zip","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T14:27:18.030595Z","iopub.execute_input":"2025-12-16T14:27:18.031174Z","iopub.status.idle":"2025-12-16T14:28:23.664556Z","shell.execute_reply.started":"2025-12-16T14:27:18.031149Z","shell.execute_reply":"2025-12-16T14:28:23.663791Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport sys\nimport glob\nimport re\nimport json\nimport csv\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.nn.utils.rnn import pad_sequence, pack_padded_sequence, pad_packed_sequence\nfrom tqdm import tqdm\nimport pandas as pd\nimport h5py\n\n# --- 引入新库 ---\nfrom pyctcdecode import build_ctcdecoder\n\n# ==========================================\n# 0. 配置\n# ==========================================\nDATA_DIR = \"/kaggle/input/brain-to-text-25/t15_copyTask_neuralData/hdf5_data_final\"\nMODEL_PATH = \"/kaggle/working/best_gru.pth\" # 确保路径对\nLM_PATH = \"/kaggle/input/3gm/pytorch/default/1/3-gram.pruned.1e-7.arpa\"\n# pyctcdecode 不需要单独的 lexicon 文件，它直接读 LM\nOUTPUT_CSV = \"submission.csv\"\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"🚀 Inference Device: {DEVICE}\")\n\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'\n]\nTOKENS = [\"\"] + VOCAB  # ⚠️ 注意: pyctcdecode 中 BLANK 通常用空字符串 \"\" 表示\n\n# ==========================================\n# 1. 模型架构 (保持不变)\n# ==========================================\nclass OfficialRNN(nn.Module):\n    def __init__(self, n_days, n_chars=len(VOCAB)+1):\n        super().__init__()\n        W = torch.stack([torch.eye(512) for _ in range(n_days)], dim=0)\n        b = torch.zeros(n_days, 512)\n        self.day_W = nn.Parameter(W)\n        self.day_b = nn.Parameter(b)\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        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        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        x = x.transpose(1, 2)\n        x = self.conv(x)\n        x = x.transpose(1, 2)\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)\n        return self.fc(out)\n\n# ==========================================\n# 2. 测试集加载器 (保持不变)\n# ==========================================\nclass TestDataset(Dataset):\n    def __init__(self, data_dir):\n        print(\"🔍 Scanning sessions...\")\n        all_files = sorted(glob.glob(os.path.join(data_dir, \"**\", \"*.hdf5\"), recursive=True))\n        sessions = set()\n        for fp in all_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        self.session_map = {s: i for i, s in enumerate(sorted(list(sessions)))}\n        \n        self.data = []\n        test_files = sorted(glob.glob(os.path.join(data_dir, \"**\", \"data_test.hdf5\"), recursive=True))\n        print(f\"📥 Loading {len(test_files)} test files...\")\n        for fp in tqdm(test_files):\n            try:\n                with h5py.File(fp, 'r') as f:\n                    for k in f.keys():\n                        g = f[k]\n                        feats = torch.from_numpy(g['input_features'][:]).float()\n                        if feats.shape[0] == 512: feats = feats.transpose(0, 1)\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                        self.data.append({\"key\": k, \"feats\": feats, \"sess_id\": sess_id})\n            except: pass\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        sess_ids = torch.tensor([x['sess_id'] for x in batch], dtype=torch.long)\n        keys = [x['key'] for x in batch]\n        lens = torch.tensor([f.shape[0] for f in feats], dtype=torch.long)\n        padded_feats = pad_sequence(feats, batch_first=True)\n        return padded_feats, lens, sess_ids, keys\n\n# ==========================================\n# 3. 核心：PyCTCDecode 构建函数\n# ==========================================\ndef get_decoder():\n    print(\"🔧 Building PyCTCDecode Decoder...\")\n    \n    # ⚠️ 关键点：\n    # alpha: LM 权重 (0.5 ~ 0.8 是 pyctcdecode 的推荐值，和 torchaudio 不同)\n    # beta: 单词奖励 (设为正值鼓励长词，设为负值惩罚单词数。为了避免碎片化，可以设高一点)\n    \n    decoder = build_ctcdecoder(\n        labels=TOKENS,          # 包含了 \"\" 和 \" \", a-z\n        kenlm_model_path=LM_PATH,\n        alpha=0.6,              # Language Model Weight\n        beta=1.5,               # Word Bonus (越高越倾向于生成完整的词)\n    )\n    return decoder\n\n# ==========================================\n# 4. 主程序\n# ==========================================\ndef main():\n    ds = TestDataset(DATA_DIR)\n    dl = DataLoader(ds, batch_size=32, collate_fn=TestDataset.collate_fn, num_workers=2)\n    \n    n_days = 45\n    model = OfficialRNN(n_days=n_days).to(DEVICE)\n    \n    print(f\"📥 Loading weights from {MODEL_PATH}...\")\n    checkpoint = torch.load(MODEL_PATH, map_location=DEVICE)\n    state_dict = {}\n    for k, v in checkpoint.items():\n        if k.startswith('module.'): state_dict[k[7:]] = v\n        else: state_dict[k] = v\n    model.load_state_dict(state_dict)\n    model.eval()\n    \n    decoder = get_decoder()\n    results = []\n    \n    print(\"🚀 Starting Decoding (PyCTCDecode)...\")\n    with torch.no_grad():\n        for feats, lens, sess_ids, keys in tqdm(dl):\n            feats, sess_ids = feats.to(DEVICE), sess_ids.to(DEVICE)\n            \n            # 1. Forward\n            logits = model(feats, lens, sess_ids) # (B, T, 42)\n            \n            # 2. PyCTCDecode 需要 numpy 格式的 logits (CPU)\n            # 形状需要是 (Time, Vocabulary) 或者是 Batch List\n            # 我们逐个样本解码\n            logits_np = logits.cpu().numpy()\n            \n            for i in range(len(keys)):\n                # 截取有效长度\n                valid_len = lens[i].item()\n                sample_logits = logits_np[i, :valid_len, :] # (T, 42)\n                \n                # 3. Decode\n                text = decoder.decode(sample_logits)\n                \n                # 4. 后处理 (防止空值)\n                if not text or len(text.strip()) == 0:\n                    text = \" \"\n                \n                results.append({\"id\": keys[i], \"text\": text})\n\n    # 5. 保存\n    results.sort(key=lambda x: x['id'])\n    df = pd.DataFrame(results)\n    df['id'] = range(len(df))\n    df.to_csv(OUTPUT_CSV, index=False)\n    \n    print(f\"\\n✅ Submission generated: {OUTPUT_CSV}\")\n    print(\"👀 Preview:\")\n    print(df.head(10))\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T14:30:29.537588Z","iopub.execute_input":"2025-12-16T14:30:29.537888Z","iopub.status.idle":"2025-12-16T14:34:21.011012Z","shell.execute_reply.started":"2025-12-16T14:30:29.537867Z","shell.execute_reply":"2025-12-16T14:34:21.010308Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}