{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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"},"accelerator":"GPU"},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# CAFA 5 — self-contained Kaggle pipeline.\n# Paste each \"# %%\" block into its own Kaggle notebook cell.\n# Settings: Accelerator = GPU, Internet = ON (needed to download ESM-2 weights).\n# Add the competition data to the notebook (it mounts at the path below).","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\ns = '/kaggle/input/datasets/t8101349/submit/submission.tsv'\n\n# Read the TSV file (specify the tab separator)\narr = pd.read_csv(s, sep='\\t')\n\n# Save to a new TSV file\narr.to_csv(\"submission.tsv\", sep='\\t', index=False)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"stop","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %% [CELL 1] paths, imports, GO DAG, metric -------------------------------\nimport 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))\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\n\nclass GODag:\n    def __init__(self, path):\n        self.ns, self.parents = parse_obo(path); self._c = {}\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    def propagate_terms(self, terms):\n        out = set()\n        for t in terms: out |= self.ancestors(t)\n        return out\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    @staticmethod\n    def propagate_scores(scores, pairs):\n        out = scores.copy()\n        np.maximum.at(out, (slice(None), pairs[:, 1]), scores[:, pairs[:, 0]])\n        return out\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,"execution":{"iopub.status.busy":"2026-06-27T19:38:48.716455Z","iopub.execute_input":"2026-06-27T19:38:48.716729Z","iopub.status.idle":"2026-06-27T19:38:52.416619Z","shell.execute_reply.started":"2026-06-27T19:38:48.716706Z","shell.execute_reply":"2026-06-27T19:38:52.415882Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %% [CELL 2] data loaders -------------------------------------------------\ndef 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    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=1e-3):\n    per = [dict() for _ in proteins]\n    for vocab, sc in blocks:\n        for j, term in enumerate(vocab):\n            col = sc[:, j]\n            for i in np.where(col >= min_score)[0]:\n                s = float(col[i])\n                if s > per[i].get(term, 0.0): per[i][term] = s\n    with open(path, \"w\") as fh:\n        for i, pid in enumerate(proteins):\n            for term, s in sorted(per[i].items(), key=lambda kv: -kv[1])[:MAX_PER_PROTEIN]:\n                fh.write(f\"{pid}\\t{term}\\t{round(s, 3)}\\n\")\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":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %% [CELL 3] naive baseline (CPU, ~minutes) -> first submission -----------\nblocks = []\nfor aspect in ASPECTS:\n    vocab = build_vocab(terms, dag, aspect); ia = load_ia(vocab)\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    f = weighted_fmax(Y[va], np.tile(Y[tr].mean(0), (len(va), 1)), ia)\n    print(f\"[{aspect}] Fmax(naive)={f:.4f}\")\n    blocks.append((vocab, np.tile(Y.mean(0), (len(test_ids), 1)).astype(np.float32)))\nwrite_submission(f\"{WORK}/submission_naive.tsv\", test_ids, blocks)\ndel blocks","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %% [CELL 3b] per-taxon naive prior (CPU) ---------------------------------\n# Term frequency conditioned on species, smoothed toward the global frequency.\n# Beats the global naive because many GO terms are strongly species-biased.\nTRAIN_TAX = f\"{DATA}/Train/train_taxonomy.tsv\"\nTEST_TAX  = f\"{DATA}/Test (Targets)/testsuperset-taxon-list.tsv\"\nALPHA = 50.0   # backoff strength: higher => trust the global freq more for rare taxa\n\n\ndef load_taxonomy(path):\n    df = pd.read_csv(path, sep=\"\\t\", header=None, dtype=str, encoding=\"latin1\")\n    if not str(df.iloc[0, 1]).strip().isdigit():   # drop a header row if present\n        df = df.iloc[1:]\n    return dict(zip(df.iloc[:, 0].str.strip(), df.iloc[:, 1].str.strip()))\n\ntrain_tax = load_taxonomy(TRAIN_TAX)\ntest_tax  = load_taxonomy(TEST_TAX)\nprint(f\"taxa: train={len(set(train_tax.values()))} test={len(set(test_tax.values()))}\")\n\ndef taxon_prior(Y, row_tax, query_tax, alpha=ALPHA):\n    g = Y.mean(0)                                  # global frequency [V]\n    sums, cnts = {}, {}\n    for i, tx in enumerate(row_tax):\n        if tx not in sums:\n            sums[tx] = np.zeros(Y.shape[1], np.float64); cnts[tx] = 0\n        sums[tx] += Y[i]; cnts[tx] += 1\n    prior = {tx: (sums[tx] + alpha * g) / (cnts[tx] + alpha) for tx in sums}\n    out = np.empty((len(query_tax), Y.shape[1]), np.float32)\n    for r, tx in enumerate(query_tax):\n        out[r] = prior.get(tx, g)                  # unseen taxon -> global\n    return out\n\ntax_blocks = []\nfor aspect in ASPECTS:\n    vocab = build_vocab(terms, dag, aspect); ia = load_ia(vocab)\n    Y = label_matrix(terms, dag, aspect, vocab, train_ids)\n    row_tax = [train_tax.get(p, \"NA\") for p in train_ids]\n    tr, va = train_test_split(np.arange(len(train_ids)), test_size=0.1, random_state=SEED)\n    f_tax = weighted_fmax(Y[va], taxon_prior(Y[tr], [row_tax[i] for i in tr],\n                                             [row_tax[i] for i in va]), ia)\n    f_glob = weighted_fmax(Y[va], np.tile(Y[tr].mean(0), (len(va), 1)), ia)\n    print(f\"[{aspect}] Fmax taxon-prior={f_tax:.4f}  vs global-naive={f_glob:.4f}\")\n    q_tax = [test_tax.get(p, \"NA\") for p in test_ids]\n    tax_blocks.append((vocab, taxon_prior(Y, row_tax, q_tax)))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-27T06:04:51.418241Z","iopub.execute_input":"2026-06-27T06:04:51.418541Z","iopub.status.idle":"2026-06-27T06:04:51.587254Z","shell.execute_reply.started":"2026-06-27T06:04:51.418501Z","shell.execute_reply":"2026-06-27T06:04:51.586174Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %% [CELL 4] ESM-2 embeddings (GPU) ---------------------------------------\n!pip install -q transformers sentencepiece  \n\nfrom transformers import AutoTokenizer, AutoModel\n\nESM_NAME = \"facebook/esm2_t33_650M_UR50D\"   # -> esm2_t36_3B_UR50D if VRAM allows\n\n@torch.no_grad()\ndef embed_esm(seqs, batch_size=8, max_len=1022):\n    tok = AutoTokenizer.from_pretrained(ESM_NAME)\n    model = AutoModel.from_pretrained(ESM_NAME).to(DEVICE).eval().half()\n    ids = list(seqs); out = []\n    for i in range(0, len(ids), batch_size):\n        chunk = ids[i:i + batch_size]\n        enc = tok([seqs[p][:max_len] for p in chunk], return_tensors=\"pt\",\n                  padding=True, truncation=True, max_length=max_len).to(DEVICE)\n        rep = model(**enc).last_hidden_state\n        mask = enc[\"attention_mask\"].unsqueeze(-1)\n        pooled = (rep * mask).sum(1) / mask.sum(1).clamp(min=1)\n        out.append(pooled.float().cpu().numpy())\n        if i % (batch_size * 50) == 0: print(f\"  {i}/{len(ids)}\", end=\"\\r\")\n    return ids, np.concatenate(out)\n\nids_tr, Xtr = embed_esm(train_seqs); np.savez_compressed(f\"{WORK}/emb_train.npz\", ids=np.array(ids_tr), X=Xtr)\nids_te, Xte = embed_esm(test_seqs);  np.savez_compressed(f\"{WORK}/emb_test.npz\",  ids=np.array(ids_te), X=Xte)\nprint(\"\\nembeddings:\", Xtr.shape, Xte.shape)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %% [CELL 5] per-aspect MLP -> predictions --------------------------------\ndef l2(X): return X / (np.linalg.norm(X, axis=1, keepdims=True) + 1e-8)\nXtr = l2(np.load(f\"{WORK}/emb_train.npz\")[\"X\"].astype(np.float32))\nXte = l2(np.load(f\"{WORK}/emb_test.npz\")[\"X\"].astype(np.float32))\n\nclass 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, 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    lf = nn.BCEWithLogitsLoss()\n    Xt_t = torch.tensor(Xt, device=DEVICE); Yt_t = torch.tensor(Yt, device=DEVICE)\n    for _ in range(epochs):\n        m.train(); perm = torch.randperm(len(Xt_t), device=DEVICE)\n        for i in range(0, len(Xt_t), bs):\n            idx = perm[i:i + bs]; opt.zero_grad()\n            lf(m(Xt_t[idx]), Yt_t[idx]).backward(); opt.step()\n    m.eval()\n    with torch.no_grad():\n        return torch.sigmoid(m(torch.tensor(Xpred, device=DEVICE))).cpu().numpy()\n\nmlp_blocks = []\nfor aspect in ASPECTS:\n    vocab = build_vocab(terms, dag, aspect); ia = load_ia(vocab)\n    pairs = dag.ancestor_pairs(vocab)\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    pva = GODag.propagate_scores(train_mlp(Xtr[tr], Y[tr], Xtr[va]), pairs)\n    print(f\"[{aspect}] Fmax(mlp)={weighted_fmax(Y[va], pva, ia):.4f}\")\n    pte = GODag.propagate_scores(train_mlp(Xtr, Y, Xte), pairs)\n    mlp_blocks.append((vocab, pte.astype(np.float32)))\nwrite_submission(f\"{WORK}/submission_mlp.tsv\", test_ids, mlp_blocks)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!apt-get update && apt-get install -y diamond-aligner","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-27T19:37:52.632079Z","iopub.execute_input":"2026-06-27T19:37:52.63235Z","iopub.status.idle":"2026-06-27T19:38:07.572601Z","shell.execute_reply.started":"2026-06-27T19:37:52.632322Z","shell.execute_reply":"2026-06-27T19:38:07.571957Z"}},"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,"execution":{"iopub.status.busy":"2026-06-27T19:38:58.387365Z","iopub.execute_input":"2026-06-27T19:38:58.38777Z","iopub.status.idle":"2026-06-27T19:40:56.13295Z","shell.execute_reply.started":"2026-06-27T19:38:58.387745Z","shell.execute_reply":"2026-06-27T19:40:56.131321Z"}},"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); pairs = dag.ancestor_pairs(vocab)\n    acc = W[\"mlp\"] * mlp_blocks[k][1]\n    wsum = W[\"mlp\"]\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    # homology (CELL 6)\n    try:\n        acc = acc + W[\"homology\"] * hom_blocks[k][1]; wsum += W[\"homology\"]\n    except NameError:\n        pass\n    acc = GODag.propagate_scores(acc / wsum, pairs)\n    final.append((vocab, acc))\nwrite_submission(f\"{WORK}/submission.csv\", test_ids, final)\nprint(\"done -> /kaggle/working/submission.csv\")","metadata":{},"outputs":[],"execution_count":null}]}