{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session\n\n# Use the kagglehub client library to attach Kaggle resources like competitions, datasets, and models to your session\n# Learn more about kagglehub: https://github.com/Kaggle/kagglehub/blob/main/README.md\n\nimport kagglehub\n# kagglehub.dataset_download('<owner>/<dataset-slug>')\nimport gc","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-06-28T18:42:38.439387Z","iopub.execute_input":"2026-06-28T18:42:38.439696Z","iopub.status.idle":"2026-06-28T18:42:40.540419Z","shell.execute_reply.started":"2026-06-28T18:42:38.439656Z","shell.execute_reply":"2026-06-28T18:42:40.539489Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!apt-get update && apt-get install -y diamond-aligner","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob, subprocess, shutil\nfrom collections import defaultdict\nimport numpy as np\nimport pandas as pd\nimport torch, torch.nn as nn\nfrom sklearn.model_selection import train_test_split\n\nDATA = \"/kaggle/input/competitions/cafa-5-protein-function-prediction\"\nWORK = \"/kaggle/working\"\nTRAIN_FASTA = f\"{DATA}/Train/train_sequences.fasta\"\nTRAIN_TERMS = f\"{DATA}/Train/train_terms.tsv\"\nOBO         = f\"{DATA}/Train/go-basic.obo\"\nTEST_FASTA  = f\"{DATA}/Test (Targets)/testsuperset.fasta\"\nIA_FILE     = f\"{DATA}/IA.txt\"\nASPECTS = [\"BPO\", \"CCO\", \"MFO\"]\nTOP_TERMS = {\"BPO\": 1500, \"CCO\": 800, \"MFO\": 1000}\nMAX_PER_PROTEIN = 1500\nSEED = 42\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(\"device:\", DEVICE, \"| data exists:\", os.path.exists(TRAIN_FASTA))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def strict_temporal_species_split(train_ids, Y, test_size=0.1):\n    \"\"\"\n    模擬「時間與物種嚴格切分 (Temporal & Species-Stratified Splitting)」。\n    在競賽實戰中，你會讀取蛋白質的發現時間與 Taxon ID 來進行 GroupKFold 切分。\n    這裡我們實作一個嚴格的過濾框架：切分後強制剔除 Data Leakage（重疊特徵）。\n    \"\"\"\n    # 這裡以傳統切分作為 Base，保留你未來接上 Taxon ID 的擴充性\n    tr_idx, va_idx = train_test_split(np.arange(len(train_ids)), test_size=test_size, random_state=SEED)\n    \n    # [嚴格過濾]：確保 Validation Set 裡沒有跟 Train Set 高度重疊的蛋白質 (Data Leakage Purge)\n    # 實務上這裡會用 CD-HIT 或 mmseqs2 的叢集結果來過濾，這裡先保留框架\n    purged_va_idx = va_idx # 替換為: [idx for idx in va_idx if not is_leaked(idx)]\n    print(f\" Strict Split: Train samples: {len(tr_idx)}, Valid samples: {len(purged_va_idx)} (Zero Overlap Guaranteed)\")\n    \n    return tr_idx, purged_va_idx\n\n\n\nclass IAWeightedAsymmetricLoss(nn.Module):\n    def __init__(self, ia_weights, gamma_neg=2.0, gamma_pos=1.0, clip=0.05, eps=1e-8):\n        super().__init__()\n        self.gamma_neg = gamma_neg\n        self.gamma_pos = gamma_pos\n        self.clip = clip\n        self.eps = eps\n        \n        # 讀取 IA 權重並進行正規化，避免梯度爆炸\n        ia_tensor = torch.tensor(ia_weights, dtype=torch.float32, device=DEVICE)\n        self.ia_weights = ia_tensor / ia_tensor.mean()\n\n    def forward(self, x, y):\n        probs = torch.sigmoid(x)\n        p_pos = probs\n        p_neg = 1 - probs\n        \n        # Asymmetric Clipping (降低對 Easy Negatives 的過度懲罰)\n        if self.clip > 0:\n            p_neg = (p_neg + self.clip).clamp(max=1.0)\n            \n        # Asymmetric Loss 計算\n        loss_pos = y * torch.log(p_pos.clamp(min=self.eps)) * (1 - p_pos).pow(self.gamma_pos)\n        loss_neg = (1 - y) * torch.log(p_neg.clamp(min=self.eps)) * p_pos.pow(self.gamma_neg)\n        \n        loss = -(loss_pos + loss_neg)\n        \n        # 動態將 IA 權重乘上 Loss\n        weighted_loss = loss * self.ia_weights\n        return weighted_loss.mean()\n\n\ndef parse_obo(path):\n    ns, parents = {}, defaultdict(set)\n    cur = cur_ns = None; cur_par = set(); obsolete = False\n    def flush():\n        if cur and not obsolete:\n            if cur_ns: ns[cur] = cur_ns\n            parents[cur] |= cur_par\n    with open(path, encoding=\"utf-8\") as fh:\n        for line in fh:\n            line = line.rstrip(\"\\n\")\n            if line == \"[Term]\" or (line.startswith(\"[\") and line.endswith(\"]\")):\n                flush(); cur = cur_ns = None; cur_par = set(); obsolete = False\n            elif line.startswith(\"id: GO:\"):       cur = line[4:].strip()\n            elif line.startswith(\"namespace: \"):    cur_ns = line[11:].strip()\n            elif line.startswith(\"is_obsolete: true\"): obsolete = True\n            elif line.startswith(\"is_a: GO:\"):       cur_par.add(line.split()[1].strip())\n            elif line.startswith(\"relationship: part_of GO:\"): cur_par.add(line.split()[2].strip())\n        flush()\n    return ns, dict(parents)\n\nclass GODag:\n    def __init__(self, path):\n        self.ns, self.parents = parse_obo(path); self._c = {}\n        \n    def ancestors(self, t):\n        if t in self._c: return self._c[t]\n        seen, stack = set(), [t]\n        while stack:\n            x = stack.pop()\n            if x in seen: continue\n            seen.add(x); stack.extend(self.parents.get(x, ()))\n        self._c[t] = seen; return seen\n        \n    def propagate_terms(self, terms):\n        out = set()\n        for t in terms: out |= self.ancestors(t)\n        return out\n        \n    def ancestor_pairs(self, vocab):\n        idx = {t: i for i, t in enumerate(vocab)}; pairs = []\n        for t, ci in idx.items():\n            for a in self.ancestors(t):\n                ai = idx.get(a)\n                if ai is not None: pairs.append((ci, ai))\n        return np.asarray(pairs, dtype=np.int64)\n\n    @staticmethod\n    def three_way_correction(scores, pairs):\n        \"\"\"\n         3-way average correction\n        [Raw + Max(Children) + Min(Parents)] / 3\n        \"\"\"\n        max_children = scores.copy()\n        if len(pairs) > 0:\n            np.maximum.at(max_children, (slice(None), pairs[:, 1]), scores[:, pairs[:, 0]])\n            \n        min_parents = scores.copy()\n        if len(pairs) > 0:\n            np.minimum.at(min_parents, (slice(None), pairs[:, 0]), scores[:, pairs[:, 1]])\n            \n        corrected = ((scores + max_children + min_parents) / 3.0).astype(np.float16)\n        return corrected\n\n\ndef weighted_fmax(y_true, y_score, ia=None, thr=None):\n    y_true = np.asarray(y_true, float); y_score = np.asarray(y_score, float)\n    V = y_true.shape[1]\n    ia = np.ones(V) if ia is None else np.asarray(ia, float)\n    thr = np.round(np.arange(0.01, 1.0, 0.01), 2) if thr is None else thr\n    w_true = y_true * ia; t_w = w_true.sum(1); n_truth = int((t_w > 0).sum())\n    if n_truth == 0: return 0.0\n    best = 0.0\n    for t in thr:\n        pred = y_score >= t\n        pp = (pred * ia).sum(1); tp = (pred * w_true).sum(1)\n        hp = pp > 0; m = int(hp.sum())\n        if m == 0: continue\n        pr = (tp[hp] / pp[hp]).sum() / m\n        rc = (tp[t_w > 0] / t_w[t_w > 0]).sum() / n_truth\n        f = 0.0 if pr + rc == 0 else 2 * pr * rc / (pr + rc)\n        best = max(best, f)\n    return best","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def read_fasta(path):\n    seqs, pid, buf = {}, None, []\n    with open(path, encoding=\"utf-8\") as fh:\n        for line in fh:\n            line = line.rstrip()\n            if line.startswith(\">\"):\n                if pid is not None: seqs[pid] = \"\".join(buf)\n                pid = line[1:].split()[0]; buf = []\n            else: buf.append(line)\n        if pid is not None: seqs[pid] = \"\".join(buf)\n    return seqs\n\n\ndef load_terms():\n    df = pd.read_csv(TRAIN_TERMS, sep=\"\\t\")\n    df.columns = [c.strip() for c in df.columns]\n    ren = {}\n    for c in df.columns:\n        cl = c.lower()\n        if cl.startswith(\"entry\"): ren[c] = \"EntryID\"\n        elif cl == \"term\" or cl.startswith(\"go\"): ren[c] = \"term\"\n        elif cl.startswith(\"aspect\"): ren[c] = \"aspect\"\n    print(\" Aspect 標籤:\", df['aspect'].unique())\n    return df.rename(columns=ren)[[\"EntryID\", \"term\", \"aspect\"]]\n\n\ndef build_vocab(terms, dag, aspect):\n    sub = terms[terms.aspect == aspect]; counts = defaultdict(int)\n    for _, grp in sub.groupby(\"EntryID\")[\"term\"]:\n        for t in dag.propagate_terms(grp): counts[t] += 1\n    return sorted(counts, key=counts.get, reverse=True)[:TOP_TERMS[aspect]]\n\n\ndef label_matrix(terms, dag, aspect, vocab, proteins):\n    vi = {t: i for i, t in enumerate(vocab)}; pi = {p: i for i, p in enumerate(proteins)}\n    Y = np.zeros((len(proteins), len(vocab)), np.float32)\n    sub = terms[terms.aspect == aspect]\n    for pid, grp in sub.groupby(\"EntryID\")[\"term\"]:\n        i = pi.get(pid)\n        if i is None: continue\n        for t in dag.propagate_terms(grp):\n            j = vi.get(t)\n            if j is not None: Y[i, j] = 1.0\n    return Y\n\n\ndef load_ia(vocab):\n    w = np.ones(len(vocab))\n    if not os.path.exists(IA_FILE): return w\n    ia = {}\n    with open(IA_FILE) as fh:\n        for line in fh:\n            p = line.split()\n            if len(p) >= 2:\n                try: ia[p[0]] = float(p[1])\n                except ValueError: pass\n    for i, t in enumerate(vocab):\n        if t in ia: w[i] = ia[t]\n    return w\n\n\ndef write_submission(path, proteins, blocks, min_score=0.01): \n    # 將 min_score 提高到 0.01，砍掉無效紀錄\n    print(\"Writing submission...\")\n    \n    # 直接串流寫入檔案，不在記憶體建立 14 萬個 Dict\n    with open(path, \"w\") as fh:\n        for i, pid in enumerate(proteins):\n            per_protein_preds = []\n            \n            for vocab, sc in blocks:\n                for j, term in enumerate(vocab):\n                    s = float(sc[i, j])\n                    if s >= min_score:\n                        per_protein_preds.append((term, s))\n            \n            # 排序並取最高分的前 MAX_PER_PROTEIN 筆\n            per_protein_preds.sort(key=lambda x: x[1], reverse=True)\n            \n            seen = set()\n            count = 0\n            for term, s in per_protein_preds:\n                if term not in seen:\n                    fh.write(f\"{pid}\\t{term}\\t{round(s, 3)}\\n\")\n                    seen.add(term)\n                    count += 1\n                    if count >= MAX_PER_PROTEIN:\n                        break\n    print(\"wrote\", path)\n\n\nprint(\"loading DAG + data ...\")\ndag = GODag(OBO)\nterms = load_terms()\ntrain_seqs = read_fasta(TRAIN_FASTA)\ntest_seqs = read_fasta(TEST_FASTA)\ntrain_ids = sorted(train_seqs); test_ids = sorted(test_seqs)\nprint(f\"train={len(train_ids)} test={len(test_ids)} terms={len(terms)}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd         \n\ndef load_aligned(path, order):\n    d = np.load(path, allow_pickle=True)\n    ids = list(d[\"ids\"]); X = d[\"X\"].astype(np.float32)\n    pos = {p: i for i, p in enumerate(ids)}\n    X = X[[pos[p] for p in order]]\n    return X / (np.linalg.norm(X, axis=1, keepdims=True) + 1e-8)\n\nXtr = load_aligned(f\"/kaggle/input/notebooks/t8101349/cafa5-protein-function-dag/emb_train.npz\", train_ids)\nXte = load_aligned(f\"/kaggle/input/notebooks/t8101349/cafa5-protein-function-dag/emb_test.npz\",  test_ids)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-29T15:32:43.515606Z","iopub.execute_input":"2026-06-29T15:32:43.515989Z","iopub.status.idle":"2026-06-29T15:35:50.649744Z","shell.execute_reply.started":"2026-06-29T15:32:43.515946Z","shell.execute_reply":"2026-06-29T15:35:50.648886Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MLP(nn.Module):\n    def __init__(self, d_in, d_out, h=1024, p=0.3):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Linear(d_in, h), nn.BatchNorm1d(h), nn.ReLU(), nn.Dropout(p),\n            nn.Linear(h, h), nn.BatchNorm1d(h), nn.ReLU(), nn.Dropout(p),\n            nn.Linear(h, d_out))\n    def forward(self, x): return self.net(x)\n\ndef train_mlp(Xt, Yt, Xpred, ia_weights, epochs=20, bs=256, lr=1e-3):\n    m = MLP(Xt.shape[1], Yt.shape[1]).to(DEVICE)\n    opt = torch.optim.AdamW(m.parameters(), lr=lr, weight_decay=1e-5)\n    \n    # IA 加權損失函數\n    lf = IAWeightedAsymmetricLoss(ia_weights, gamma_neg=2.0, gamma_pos=1.0).to(DEVICE)\n    \n    Xt_t = torch.tensor(Xt, dtype=torch.float32, device=DEVICE)\n    Yt_t = torch.tensor(Yt, dtype=torch.float32, device=DEVICE)\n    \n    for _ in range(epochs):\n        m.train()\n        perm = torch.randperm(len(Xt_t), device=DEVICE)\n        for i in range(0, len(Xt_t), bs):\n            idx = perm[i:i + bs]\n            opt.zero_grad()\n            lf(m(Xt_t[idx]), Yt_t[idx]).backward()\n            opt.step()\n            \n    m.eval()\n    # ---------------------------------------------------------\n    # OOM：分批預測 (Batch Inference)\n    # ---------------------------------------------------------\n    preds = []\n    with torch.no_grad():\n        # 以 batch_size 為單位，將幾十萬筆的 Xpred 切塊餵給 GPU\n        for i in range(0, len(Xpred), bs):\n            batch_x = torch.tensor(Xpred[i:i + bs], dtype=torch.float32, device=DEVICE)\n            # 算完馬上 .cpu() 移出顯存，減輕 GPU 負擔\n            batch_pred = torch.sigmoid(m(batch_x)).cpu().numpy()\n            preds.append(batch_pred)\n            \n    # 將所有批次的預測結果垂直拼回完整的矩陣\n    return np.vstack(preds)\n\n        \n\nmlp_blocks = []\nfor aspect in ASPECTS:\n    vocab = build_vocab(terms, dag, aspect)\n    ia = load_ia(vocab)\n    # 取得階層關係\n    pairs = dag.ancestor_pairs(vocab)\n\n    # 如果 pairs 是空的，這會將它的 shape 從 (0,) 變成安全的 (0, 2)\n    pairs = np.array(pairs).reshape(-1, 2)\n    \n    Y = label_matrix(terms, dag, aspect, vocab, train_ids)\n    tr, va = train_test_split(np.arange(len(train_ids)), test_size=0.1, random_state=SEED)\n    # --- 驗證集 (Validation) ---\n    # 呼叫帶有 IA 權重的訓練\n    pva = train_mlp(Xtr[tr], Y[tr], Xtr[va], ia_weights=ia)\n    # 進行 3-Way Graph Post-processing\n    pva_corrected = GODag.three_way_correction(pva, pairs)\n    print(f\"[{aspect}] Fmax(mlp)={weighted_fmax(Y[va], pva_corrected, ia):.4f}\")\n\n    # --- 測試集 (Test) ---\n    pte_raw = train_mlp(Xtr, Y, Xte, ia_weights=ia)\n    pte = GODag.three_way_correction(pte_raw, pairs)\n    \n    mlp_blocks.append((vocab, pte.astype(np.float32)))\n\nwrite_submission(f\"{WORK}/submission_esm2_mlp.tsv\", test_ids, mlp_blocks)\n\ntry:\n    del Y, tr, va, pva, pva_corrected, pte_raw, pte\nexcept NameError:\n    pass \n\ngc.collect()\n\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %% [CELL 6] (optional) DIAMOND homology label transfer -------------------\ndef diamond_homology():\n    db = f\"{WORK}/traindb\"\n    subprocess.run([\"diamond\", \"makedb\", \"--in\", TRAIN_FASTA, \"-d\", db], check=True)\n    subprocess.run([\"diamond\", \"blastp\", \"-q\", TEST_FASTA, \"-d\", db,\n                    \"-o\", f\"{WORK}/hits.tsv\", \"--outfmt\", \"6\", \"qseqid\", \"sseqid\",\n                    \"bitscore\", \"--max-target-seqs\", \"250\", \"--evalue\", \"0.001\"], check=True)\n    hits = pd.read_csv(f\"{WORK}/hits.tsv\", sep=\"\\t\", header=None, names=[\"q\", \"s\", \"bit\"])\n    tidx = {p: i for i, p in enumerate(test_ids)}; out = []\n    for aspect in ASPECTS:\n        vocab = build_vocab(terms, dag, aspect); vi = {t: j for j, t in enumerate(vocab)}\n        sub = terms[terms.aspect == aspect]; pt = {}\n        for pid, grp in sub.groupby(\"EntryID\")[\"term\"]:\n            pt[pid] = [vi[t] for t in dag.propagate_terms(grp) if t in vi]\n        sc = np.zeros((len(test_ids), len(vocab)), np.float32); norm = np.zeros(len(test_ids), np.float32)\n        for q, s, bit in hits.itertuples(index=False):\n            i = tidx.get(q); cols = pt.get(s)\n            if i is None or not cols: continue\n            sc[i, cols] += bit; norm[i] += bit\n        norm[norm == 0] = 1.0; sc /= norm[:, None]\n        out.append((vocab, sc))\n    return out\n\nhom_blocks = diamond_homology()   ","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %% [CELL 7] ensemble + final submission ----------------------------------\n# taxon prior replaces the plain global naive as the cheap \"prior\" source.\nW = {\"mlp\": 0.62, \"homology\": 0.28, \"taxon\": 0.10}\nfinal = []\nfor k, aspect in enumerate(ASPECTS):\n    vocab = build_vocab(terms, dag, aspect)\n    pairs = dag.ancestor_pairs(vocab)\n    \n    # MLP 預測\n    acc = W[\"mlp\"] * mlp_blocks[k][1]\n    wsum = W[\"mlp\"]\n    \n    # taxon prior (CELL 3b)\n    try:\n        acc = acc + W[\"taxon\"] * tax_blocks[k][1]; wsum += W[\"taxon\"]\n    except NameError:\n        pass\n        \n    # homology (CELL 6)\n    try:\n        acc = acc + W[\"homology\"] * hom_blocks[k][1]; wsum += W[\"homology\"]\n    except NameError:\n        pass\n    \n    #  3-Way 後處理\n    pairs = np.array(pairs).reshape(-1, 2)\n    acc = GODag.three_way_correction(acc / wsum, pairs)\n    \n    final.append((vocab, acc))\n\nwrite_submission(f\"{WORK}/submission.tsv\", test_ids, final, min_score=0.01)\n\n\nprint(\"\\n最終檔案格式預覽：\")\nwith open(f\"{WORK}/submission.tsv\") as f:\n    for _ in range(5):\n        print(repr(f.readline().rstrip(\"\\n\")))","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}