{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":15231210},{"sourceType":"datasetVersion","sourceId":14604295,"datasetId":9328538,"databundleVersionId":15440074},{"sourceType":"datasetVersion","sourceId":14962460,"datasetId":9577079,"databundleVersionId":15833819},{"sourceType":"datasetVersion","sourceId":14874339,"datasetId":9502242,"databundleVersionId":15736806},{"sourceType":"datasetVersion","sourceId":14962495,"datasetId":9577097,"databundleVersionId":15833858},{"sourceType":"datasetVersion","sourceId":10855324,"datasetId":6742586,"databundleVersionId":11219268},{"sourceType":"kernelVersion","sourceId":290004465}],"dockerImageVersionId":31328,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Stanford_RNA_P2_TBM_Protenix_FastRoute\n\n## Timeout Fix Strategy\n\nThe previous version timed out because:\n- `9MME` (4640 nt, octamer A:8) required 12 Protenix chunks × ~4 min = ~48 min alone\n- `9ZCC` (1460 nt) required 4 chunks × ~4 min = ~16 min\n- Total Protenix time exceeded the 9-hour GPU notebook limit\n\n**Three key optimizations in this version:**\n\n1. **Symmetric chain expansion** — For homo-oligomers (stoichiometry `X:N`), predict\n   the *single chain* with Protenix (or TBM), then tile+superimpose to build the full complex.\n   `9MME` drops from 12 chunks → 1 chunk (580 nt), saving ~44 min.\n\n2. **Hard time budget** — Each target gets at most `MAX_PROTENIX_SECONDS` seconds.\n   If a target exceeds the budget, it falls back to TBM/de-novo instead of timing out.\n\n3. **Protenix chunk limit** — Sequences longer than `PROTENIX_HARD_MAX` nt skip\n   Protenix entirely and use TBM + de-novo, which are fast.\n\n## Required Kaggle Datasets\n| Dataset | Kaggle path |\n|---|---|\n| biopython-cp312 | `kami1976/biopython-cp312` |\n| biotite | `amirrezaaleyasin/biotite` |\n| rdkit-2025-9-5 | `amirrezaaleyasin/rdkit-2025-9-5` |\n| **Protenix v1** | `qiweiyin/protenix-v1-adjusted` |\n| TM-score metric | `rhijudas/tm-score-permutechains` (kernel) |\n| USalign | `metric/usalign` (dataset) |","metadata":{}},{"cell_type":"code","source":"# Install offline wheel packages\n!pip install --no-index --no-deps /kaggle/input/datasets/kami1976/biopython-cp312/biopython-1.86-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl\n!pip install --no-index --no-deps /kaggle/input/datasets/amirrezaaleyasin/biotite/biotite-1.6.0-cp312-cp312-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl\n!pip install --no-index --no-deps /kaggle/input/datasets/amirrezaaleyasin/rdkit-2025-9-5/rdkit-2025.9.5-cp312-cp312-manylinux_2_28_x86_64.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-22T04:27:43.933745Z","iopub.execute_input":"2026-03-22T04:27:43.933943Z","iopub.status.idle":"2026-03-22T04:27:54.530683Z","shell.execute_reply.started":"2026-03-22T04:27:43.933922Z","shell.execute_reply":"2026-03-22T04:27:54.529701Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import gc, json, os, time, sys, math, glob\nfrom pathlib import Path\n\nos.environ[\"LAYERNORM_TYPE\"] = \"torch\"\nos.environ.setdefault(\"RNA_MSA_DEPTH_LIMIT\", \"512\")\n\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom Bio.Align import PairwiseAligner\nfrom tqdm import tqdm\n\nprint(\"Imports OK\")\nprint(f\"CUDA available: {torch.cuda.is_available()}\")\nif torch.cuda.is_available():\n    print(f\"GPU: {torch.cuda.get_device_name(0)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-22T04:29:55.257318Z","iopub.execute_input":"2026-03-22T04:29:55.258162Z","iopub.status.idle":"2026-03-22T04:29:59.493649Z","shell.execute_reply.started":"2026-03-22T04:29:55.258122Z","shell.execute_reply":"2026-03-22T04:29:59.49291Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ══════════════════════════════════════════════════════════════\n# TUNING PARAMETERS\n# ══════════════════════════════════════════════════════════════\n\nDATA_BASE          = \"/kaggle/input/competitions/stanford-rna-3d-folding-2\"\nDEFAULT_TEST_CSV   = f\"{DATA_BASE}/test_sequences.csv\"\nDEFAULT_TRAIN_CSV  = f\"{DATA_BASE}/train_sequences.csv\"\nDEFAULT_TRAIN_LBLS = f\"{DATA_BASE}/train_labels.csv\"\nDEFAULT_VAL_CSV    = f\"{DATA_BASE}/validation_sequences.csv\"\nDEFAULT_VAL_LBLS   = f\"{DATA_BASE}/validation_labels.csv\"\nDEFAULT_OUTPUT     = \"/kaggle/working/submission.csv\"\n\nDEFAULT_CODE_DIR = (\n    \"/kaggle/input/datasets/qiweiyin/protenix-v1-adjusted\"\n    \"/Protenix-v1-adjust-v2/Protenix-v1-adjust-v2/Protenix-v1\"\n)\nDEFAULT_ROOT_DIR = DEFAULT_CODE_DIR\nMODEL_NAME       = \"protenix_base_20250630_v1.0.0\"\n\nN_SAMPLE  = 5\nSEED      = 42\nIS_KAGGLE = True\n\n# TBM threshold: 30 gives wider coverage than the original 50\nMIN_PERCENT_IDENTITY = float(os.environ.get(\"MIN_PERCENT_IDENTITY\", \"30.0\"))\nMIN_SIMILARITY       = 0.0\n\n# Protenix: skip sequences longer than this (feed to de-novo instead)\n# 9MME monomer = 580 nt, fits well within 900\nPROTENIX_HARD_MAX    = int(os.environ.get(\"PROTENIX_HARD_MAX\", \"900\"))\nMAX_SEQ_LEN          = int(os.environ.get(\"MAX_SEQ_LEN\",       \"512\"))\nCHUNK_OVERLAP        = int(os.environ.get(\"CHUNK_OVERLAP\",     \"128\"))\n\n# With GPU T4 a 512-nt chunk takes ~3-4 min; 600s = safe budget for 1-2 chunks\nMAX_PROTENIX_SECONDS = int(os.environ.get(\"MAX_PROTENIX_SECONDS\", \"600\"))\n\n# Total wall-clock budget (Kaggle GPU limit ~9 h = 32400 s)\n# Leave 30 min for TBM/startup → 30000 s for Protenix\nTOTAL_TIME_BUDGET_S  = int(os.environ.get(\"TOTAL_TIME_BUDGET_S\", \"30000\"))\n\nUSE_PROTENIX = True\n\ndef _pb(v, default=False):\n    return \"true\" if str(v).lower() in {\"1\",\"true\",\"t\",\"yes\"} else (\"true\" if default else \"false\")\n\nUSE_MSA      = _pb(os.environ.get(\"USE_MSA\",      \"false\"))\nUSE_TEMPLATE = _pb(os.environ.get(\"USE_TEMPLATE\", \"false\"))\nUSE_RNA_MSA  = _pb(os.environ.get(\"USE_RNA_MSA\",  \"true\"))\n\nWALL_CLOCK_START = time.time()\n\nprint(f\"GPU enabled      : {torch.cuda.is_available()}\")\nprint(f\"MIN_PCT_IDENTITY : {MIN_PERCENT_IDENTITY}\")\nprint(f\"PROTENIX_HARD_MAX: {PROTENIX_HARD_MAX} nt\")\nprint(f\"MAX_SEQ_LEN      : {MAX_SEQ_LEN} nt\")\nprint(f\"MAX_PROTENIX_S   : {MAX_PROTENIX_SECONDS} s\")\nprint(f\"USE_RNA_MSA      : {USE_RNA_MSA}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-22T04:30:08.015293Z","iopub.execute_input":"2026-03-22T04:30:08.015692Z","iopub.status.idle":"2026-03-22T04:30:08.024297Z","shell.execute_reply.started":"2026-03-22T04:30:08.015664Z","shell.execute_reply":"2026-03-22T04:30:08.023425Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ══════════════════════════════════════════════════════════════\n# UTILITY FUNCTIONS\n# ══════════════════════════════════════════════════════════════\n\ndef seed_everything(seed):\n    os.environ[\"CUBLAS_WORKSPACE_CONFIG\"] = \":4096:8\"\n    torch.manual_seed(seed); torch.cuda.manual_seed_all(seed); np.random.seed(seed)\n    torch.backends.cudnn.benchmark = False\n    torch.backends.cudnn.deterministic = True\n    torch.use_deterministic_algorithms(True)\n\ndef ensure_required_files(root_dir):\n    for p, name in [\n        (Path(root_dir)/\"checkpoint\"/f\"{MODEL_NAME}.pt\",         \"checkpoint\"),\n        (Path(root_dir)/\"common\"/\"components.cif\",               \"CCD file\"),\n        (Path(root_dir)/\"common\"/\"components.cif.rdkit_mol.pkl\", \"CCD cache\"),\n    ]:\n        if not p.exists():\n            raise FileNotFoundError(f\"Missing {name}: {p}\")\n\ndef build_input_json(df, json_path):\n    data = [{\"name\": r[\"target_id\"], \"covalent_bonds\": [],\n             \"sequences\": [{\"rnaSequence\": {\"sequence\": r[\"sequence\"], \"count\": 1}}]}\n            for _, r in df.iterrows()]\n    with open(json_path, \"w\") as f:\n        json.dump(data, f)\n\ndef build_configs(input_json_path, dump_dir, model_name, n_sample):\n    from configs.configs_base import configs as configs_base\n    from configs.configs_data import data_configs\n    from configs.configs_inference import inference_configs\n    from configs.configs_model_type import model_configs\n    from protenix.config.config import parse_configs\n    base = {**configs_base, **{\"data\": data_configs}, **inference_configs}\n    def _upd(t, p):\n        for k, v in p.items():\n            if isinstance(v, dict) and k in t and isinstance(t[k], dict): _upd(t[k], v)\n            else: t[k] = v\n    _upd(base, model_configs[model_name])\n    arg_str = (f\"--model_name {model_name} --input_json_path {input_json_path} \"\n               f\"--dump_dir {dump_dir} --use_msa {USE_MSA} --use_template {USE_TEMPLATE} \"\n               f\"--use_rna_msa {USE_RNA_MSA} --sample_diffusion.N_sample {n_sample} \"\n               f\"--seeds {SEED}\")\n    return parse_configs(configs=base, arg_str=arg_str, fill_required_with_null=True)\n\ndef parse_stoichiometry(stoich):\n    if pd.isna(stoich) or str(stoich).strip() == \"\": return []\n    return [(ch.strip(), int(cnt)) for part in str(stoich).split(\";\")\n            for ch, cnt in [part.split(\":\")]]\n\ndef parse_fasta(fasta_content):\n    out, cur, parts = {}, None, []\n    for line in str(fasta_content).splitlines():\n        line = line.strip()\n        if not line: continue\n        if line.startswith(\">\"):\n            if cur is not None: out[cur] = \"\".join(parts)\n            cur = line[1:].split()[0]; parts = []\n        else: parts.append(line.replace(\" \", \"\"))\n    if cur is not None: out[cur] = \"\".join(parts)\n    return out\n\ndef get_chain_segments(row):\n    seq = row[\"sequence\"]; stoich = row.get(\"stoichiometry\",\"\"); all_sq = row.get(\"all_sequences\",\"\")\n    if pd.isna(stoich) or pd.isna(all_sq) or str(stoich).strip() == \"\": return [(0,len(seq))]\n    try:\n        chain_dict = parse_fasta(all_sq); order = parse_stoichiometry(stoich)\n        segs, pos = [], 0\n        for ch, cnt in order:\n            base = chain_dict.get(ch)\n            if base is None: return [(0,len(seq))]\n            for _ in range(cnt): segs.append((pos, pos+len(base))); pos += len(base)\n        return segs if pos == len(seq) else [(0,len(seq))]\n    except: return [(0,len(seq))]\n\ndef build_segments_map(df):\n    return {r[\"target_id\"]: get_chain_segments(r) for _, r in df.iterrows()}\n\ndef process_labels(labels_df):\n    coords = {}\n    prefixes = labels_df[\"ID\"].str.rsplit(\"_\", n=1).str[0]\n    for prefix, grp in labels_df.groupby(prefixes):\n        arr = grp.sort_values(\"resid\")[[\"x_1\",\"y_1\",\"z_1\"]].values\n        valid = arr[:,0] > -1e17\n        if valid.sum() > 0: coords[prefix] = arr[valid]\n    return coords\n\ndef coords_to_rows(target_id, seq, coords):\n    \"\"\"coords: (N_SAMPLE, seq_len, 3)\"\"\"\n    rows = []\n    for i in range(len(seq)):\n        row = {\"ID\": f\"{target_id}_{i+1}\", \"resname\": seq[i], \"resid\": i+1}\n        for s in range(N_SAMPLE):\n            x,y,z = (coords[s,i] if s < coords.shape[0] and i < coords.shape[1] else (0,0,0))\n            row[f\"x_{s+1}\"] = float(x); row[f\"y_{s+1}\"] = float(y); row[f\"z_{s+1}\"] = float(z)\n        rows.append(row)\n    return rows\n\ndef split_into_chunks(seq_len, max_len, overlap):\n    if seq_len <= max_len: return [(0, seq_len)]\n    chunks, step, pos = [], max_len-overlap, 0\n    while pos < seq_len:\n        end = min(pos+max_len, seq_len); chunks.append((pos,end))\n        if end == seq_len: break\n        pos += step\n    return chunks\n\ndef kabsch_align(P, Q):\n    cp, cq = P.mean(0), Q.mean(0)\n    Pc, Qc = P-cp, Q-cq\n    U, _, Vt = np.linalg.svd(Pc.T @ Qc)\n    d = np.linalg.det(Vt.T @ U.T)\n    S = np.diag([1,1,d])\n    R = Vt.T @ S @ U.T\n    return R, cq - R @ cp\n\ndef stitch_chunk_coords(chunk_coords_list, chunk_ranges, seq_len):\n    if len(chunk_coords_list) == 1:\n        c = chunk_coords_list[0]\n        out = np.zeros((seq_len,3), dtype=c.dtype)\n        out[:min(c.shape[0],seq_len)] = c[:seq_len]\n        return out\n    aligned = [chunk_coords_list[0].copy()]\n    for i in range(1, len(chunk_coords_list)):\n        ps,pe = chunk_ranges[i-1]; cs,ce = chunk_ranges[i]\n        ov_s,ov_e = cs,min(pe,ce); ov_len = ov_e-ov_s\n        if ov_len < 3: aligned.append(chunk_coords_list[i].copy()); continue\n        pov = aligned[i-1][ov_s-ps:ov_e-ps]; cov = chunk_coords_list[i][ov_s-cs:ov_e-cs]\n        valid = ~(np.isnan(pov).any(1) | np.isnan(cov).any(1))\n        if valid.sum() < 3: aligned.append(chunk_coords_list[i].copy()); continue\n        R,t = kabsch_align(cov[valid], pov[valid])\n        aligned.append((chunk_coords_list[i] @ R.T) + t)\n    full = np.zeros((seq_len,3),dtype=np.float64); weights = np.zeros(seq_len)\n    for i,((s,e),coords) in enumerate(zip(chunk_ranges,aligned)):\n        ae = min(s+coords.shape[0],seq_len); ul = ae-s\n        w = np.ones(ul)\n        if i > 0:\n            rl = min(chunk_ranges[i-1][1],e)-s\n            if rl > 0: w[:rl] = np.linspace(0,1,rl)\n        if i < len(chunk_ranges)-1:\n            ns = chunk_ranges[i+1][0]; rs=ns-s; rl=ae-ns\n            if rl > 0 and rs < ul: w[rs:ul] = np.linspace(1,0,rl)\n        full[s:ae] += coords[:ul]*w[:,None]; weights[s:ae] += w\n    mask = weights > 0; full[mask] /= weights[mask,None]\n    return full\n\nprint(\"Utilities ready.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-22T04:30:13.484555Z","iopub.execute_input":"2026-03-22T04:30:13.485422Z","iopub.status.idle":"2026-03-22T04:30:13.512403Z","shell.execute_reply.started":"2026-03-22T04:30:13.485392Z","shell.execute_reply":"2026-03-22T04:30:13.511682Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ══════════════════════════════════════════════════════════════\n# HOMO-OLIGOMER EXPANSION\n# For targets like 9MME (octamer A:8), predict the monomer\n# with Protenix (580 nt) instead of the full 4640 nt complex.\n# ══════════════════════════════════════════════════════════════\n\ndef get_monomer_info(row):\n    stoich = row.get(\"stoichiometry\",\"\"); all_sq = row.get(\"all_sequences\",\"\")\n    seq    = row[\"sequence\"]\n    if pd.isna(stoich) or str(stoich).strip() == \"\": return seq, 1\n    parsed = parse_stoichiometry(stoich)\n    if len(parsed) == 1:\n        chain, n_copies = parsed[0]\n        if n_copies > 1 and not pd.isna(all_sq):\n            chain_dict = parse_fasta(all_sq)\n            monomer    = chain_dict.get(chain)\n            if monomer and len(monomer)*n_copies == len(seq):\n                return monomer, n_copies\n    return seq, 1\n\ndef tile_monomer_coords(monomer_coords, n_copies, monomer_len):\n    \"\"\"Tile monomer coordinates into a full complex.\n    Each copy is offset along X by bounding-box width + 5 Å gap.\n    \"\"\"\n    full = np.zeros((n_copies*monomer_len, 3))\n    span = monomer_coords.max(0) - monomer_coords.min(0)\n    for i in range(n_copies):\n        offset = np.array([(span[0]+5.0)*i, 0.0, 0.0])\n        full[i*monomer_len:(i+1)*monomer_len] = monomer_coords + offset\n    return full\n\nprint(\"Homo-oligomer expansion ready.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-22T04:30:19.697781Z","iopub.execute_input":"2026-03-22T04:30:19.698392Z","iopub.status.idle":"2026-03-22T04:30:19.705138Z","shell.execute_reply.started":"2026-03-22T04:30:19.698365Z","shell.execute_reply":"2026-03-22T04:30:19.70438Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ══════════════════════════════════════════════════════════════\n# TBM CORE\n# ══════════════════════════════════════════════════════════════\n\ndef _make_aligner():\n    al = PairwiseAligner()\n    al.mode = \"global\"\n    al.match_score = 2; al.mismatch_score = -1.5\n    al.open_gap_score = -8; al.extend_gap_score = -0.4\n    for attr in [\"query_left_open_gap_score\",\"query_left_extend_gap_score\",\n                 \"query_right_open_gap_score\",\"query_right_extend_gap_score\",\n                 \"target_left_open_gap_score\",\"target_left_extend_gap_score\",\n                 \"target_right_open_gap_score\",\"target_right_extend_gap_score\"]:\n        try: setattr(al, attr, -8 if \"open\" in attr else -0.4)\n        except: pass\n    return al\n\n_aligner = _make_aligner()\n\ndef find_similar_sequences(query_seq, train_seqs_df, train_coords_dict, top_n=30):\n    results = []\n    for _, row in train_seqs_df.iterrows():\n        tid, tseq = row[\"target_id\"], row[\"sequence\"]\n        if tid not in train_coords_dict: continue\n        if abs(len(tseq)-len(query_seq))/max(len(tseq),len(query_seq)) > 0.3: continue\n        aln      = next(iter(_aligner.align(query_seq, tseq)))\n        norm_s   = aln.score / (2*min(len(query_seq),len(tseq)))\n        identical = sum(1 for (qs,qe),(ts,te) in zip(*aln.aligned)\n                        for qp,tp in zip(range(qs,qe),range(ts,te))\n                        if query_seq[qp]==tseq[tp])\n        pct_id = 100*identical/len(query_seq)\n        results.append((tid, tseq, norm_s, train_coords_dict[tid], pct_id))\n    results.sort(key=lambda x: x[2], reverse=True)\n    return results[:top_n]\n\ndef adapt_template_to_query(query_seq, template_seq, template_coords):\n    aln = next(iter(_aligner.align(query_seq, template_seq)))\n    new_coords = np.full((len(query_seq),3), np.nan)\n    for (qs,qe),(ts,te) in zip(*aln.aligned):\n        chunk = template_coords[ts:te]\n        if len(chunk) == (qe-qs): new_coords[qs:qe] = chunk\n    for i in range(len(new_coords)):\n        if np.isnan(new_coords[i,0]):\n            pv = next((j for j in range(i-1,-1,-1) if not np.isnan(new_coords[j,0])),-1)\n            nv = next((j for j in range(i+1,len(new_coords)) if not np.isnan(new_coords[j,0])),-1)\n            if pv>=0 and nv>=0: w=(i-pv)/(nv-pv); new_coords[i]=(1-w)*new_coords[pv]+w*new_coords[nv]\n            elif pv>=0: new_coords[i]=new_coords[pv]+[3,0,0]\n            elif nv>=0: new_coords[i]=new_coords[nv]+[3,0,0]\n            else: new_coords[i]=[i*3,0,0]\n    return np.nan_to_num(new_coords)\n\ndef adaptive_rna_constraints(coords, segs, confidence=1.0, passes=2):\n    X = coords.copy()\n    strength = max(0.75*(1.0-min(confidence,0.97)), 0.02)\n    for _ in range(passes):\n        for s,e in segs:\n            C=X[s:e]; L=e-s\n            if L<3: continue\n            d=C[1:]-C[:-1]; dist=np.linalg.norm(d,axis=1)+1e-6\n            adj=d*((5.95-dist)/dist)[:,None]*(0.22*strength)\n            C[:-1]-=adj; C[1:]+=adj\n            d2=C[2:]-C[:-2]; d2n=np.linalg.norm(d2,axis=1)+1e-6\n            adj2=d2*((10.2-d2n)/d2n)[:,None]*(0.10*strength)\n            C[:-2]-=adj2; C[2:]+=adj2\n            C[1:-1]+=(0.06*strength)*(0.5*(C[:-2]+C[2:])-C[1:-1])\n            if L>=25:\n                idx=np.linspace(0,L-1,min(L,160)).astype(int) if L>220 else np.arange(L)\n                P=C[idx]; diff=P[:,None,:]-P[None,:,:]\n                dm=np.linalg.norm(diff,axis=2)+1e-6\n                sep=np.abs(idx[:,None]-idx[None,:])\n                mask=(sep>2)&(dm<3.2)\n                if np.any(mask):\n                    vec=(diff*((3.2-dm)/dm)[:,:,None]*mask[:,:,None]).sum(1)\n                    C[idx]+=(0.015*strength)*vec\n            X[s:e]=C\n    return X\n\ndef _rotmat(axis,ang):\n    a=np.asarray(axis,float); a/=np.linalg.norm(a)+1e-12\n    x,y,z=a; c,s=np.cos(ang),np.sin(ang); CC=1-c\n    return np.array([[c+x*x*CC,x*y*CC-z*s,x*z*CC+y*s],\n                     [y*x*CC+z*s,c+y*y*CC,y*z*CC-x*s],\n                     [z*x*CC-y*s,z*y*CC+x*s,c+z*z*CC]])\n\ndef apply_hinge(coords,seg,rng,deg=22):\n    s,e=seg; L=e-s\n    if L<30: return coords\n    pivot=s+int(rng.integers(10,L-10))\n    R=_rotmat(rng.normal(size=3),np.deg2rad(float(rng.uniform(-deg,deg))))\n    X=coords.copy(); p0=X[pivot].copy()\n    X[pivot+1:e]=(X[pivot+1:e]-p0)@R.T+p0\n    return X\n\ndef jitter_chains(coords,segs,rng,deg=12,trans=1.5):\n    X=coords.copy(); gc_=X.mean(0,keepdims=True)\n    for s,e in segs:\n        R=_rotmat(rng.normal(size=3),np.deg2rad(float(rng.uniform(-deg,deg))))\n        sh=rng.normal(size=3); sh=sh/(np.linalg.norm(sh)+1e-12)*float(rng.uniform(0,trans))\n        c=X[s:e].mean(0,keepdims=True); X[s:e]=(X[s:e]-c)@R.T+c+sh\n    X-=X.mean(0,keepdims=True)-gc_; return X\n\ndef smooth_wiggle(coords,segs,rng,amp=0.8):\n    X=coords.copy()\n    for s,e in segs:\n        L=e-s\n        if L<20: continue\n        ctrl=np.linspace(0,L-1,6); disp=rng.normal(0,amp,(6,3)); t=np.arange(L)\n        X[s:e]+=np.vstack([np.interp(t,ctrl,disp[:,k]) for k in range(3)]).T\n    return X\n\ndef generate_rna_structure(sequence, seed=None):\n    if seed is not None: np.random.seed(seed)\n    n=len(sequence); coords=np.zeros((n,3))\n    for i in range(n): ang=i*0.6; coords[i]=[10.0*np.cos(ang),10.0*np.sin(ang),i*2.5]\n    return coords\n\nprint(\"TBM core ready.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-21T11:58:36.835836Z","iopub.execute_input":"2026-03-21T11:58:36.836091Z","iopub.status.idle":"2026-03-21T11:58:36.868077Z","shell.execute_reply.started":"2026-03-21T11:58:36.836061Z","shell.execute_reply":"2026-03-21T11:58:36.866947Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ══════════════════════════════════════════════════════════════\n# PHASE 1: TBM\n# ══════════════════════════════════════════════════════════════\n\ndef tbm_phase(test_df, combined_seqs, train_coords, segments_map, monomer_info_map):\n    print(f\"\\n{'='*60}\")\n    print(f\"PHASE 1: TBM  (MIN_PCT_IDENTITY={MIN_PERCENT_IDENTITY})\")\n    print(f\"{'='*60}\")\n    t0 = time.time()\n    template_preds = {}; protenix_queue = {}\n\n    for _, row in test_df.iterrows():\n        tid  = row[\"target_id\"]; seq = row[\"sequence\"]\n        segs = segments_map.get(tid, [(0,len(seq))])\n        mono_seq, n_copies = monomer_info_map[tid]\n        similar = find_similar_sequences(mono_seq, combined_seqs, train_coords, top_n=30)\n        preds, used = [], set()\n\n        for i,(tmpl_id,tmpl_seq,sim,tmpl_coords,pct_id) in enumerate(similar):\n            if len(preds) >= N_SAMPLE: break\n            if pct_id < MIN_PERCENT_IDENTITY or sim < MIN_SIMILARITY: break\n            if tmpl_id in used: continue\n            rng     = np.random.default_rng((row.name*10000000000+i*10007)%(2**32))\n            adapted = adapt_template_to_query(mono_seq, tmpl_seq, tmpl_coords)\n            if n_copies > 1:\n                adapted = tile_monomer_coords(adapted, n_copies, len(mono_seq))\n            slot = len(preds)\n            if   slot==0: X=adapted\n            elif slot==1: X=adapted+rng.normal(0,max(0.01,(0.40-sim)*0.06),adapted.shape)\n            elif slot==2: longest=max(segs,key=lambda se:se[1]-se[0]); X=apply_hinge(adapted,longest,rng)\n            elif slot==3: X=jitter_chains(adapted,segs,rng)\n            else:         X=smooth_wiggle(adapted,segs,rng)\n            refined=adaptive_rna_constraints(X,segs,confidence=sim)\n            preds.append(refined); used.add(tmpl_id)\n\n        template_preds[tid] = preds\n        n_needed = N_SAMPLE-len(preds)\n        if n_needed > 0:\n            protenix_seq = mono_seq if (n_copies>1 and len(mono_seq)<=PROTENIX_HARD_MAX) else seq\n            is_mono      = (protenix_seq==mono_seq and n_copies>1)\n            route = \"mono\" if is_mono else (\"skip\" if len(protenix_seq)>PROTENIX_HARD_MAX else \"full\")\n            protenix_queue[tid] = (n_needed, protenix_seq, n_copies if is_mono else 1, route)\n            print(f\"  {tid} ({len(seq)} nt): {len(preds)} TBM → {n_needed} Protenix [{route}]\")\n        else:\n            print(f\"  {tid} ({len(seq)} nt): all {N_SAMPLE} from TBM ✓\")\n\n    elapsed=time.time()-t0\n    skipped=sum(1 for v in protenix_queue.values() if v[3]==\"skip\")\n    print(f\"\\nPhase 1: {elapsed:.1f}s | TBM-complete={len(test_df)-len(protenix_queue)} | queue={len(protenix_queue)} | skip={skipped}\")\n    return template_preds, protenix_queue\n\nprint(\"TBM phase ready.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-21T11:58:36.869216Z","iopub.execute_input":"2026-03-21T11:58:36.869473Z","iopub.status.idle":"2026-03-21T11:58:36.897112Z","shell.execute_reply.started":"2026-03-21T11:58:36.869449Z","shell.execute_reply":"2026-03-21T11:58:36.895629Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ══════════════════════════════════════════════════════════════\n# PHASE 2: PROTENIX\n# ══════════════════════════════════════════════════════════════\n\ndef _extract_c1_coords(pred, feat, chunk_seq_len, raw_coords):\n    if \"centre_atom_mask\" in feat:\n        mask = (feat[\"centre_atom_mask\"]==1).to(raw_coords.device)\n    elif \"atom_to_tokatom_idx\" in feat:\n        m11=(feat[\"atom_to_tokatom_idx\"]==11).to(raw_coords.device)\n        m12=(feat[\"atom_to_tokatom_idx\"]==12).to(raw_coords.device)\n        c11,c12=m11.sum(),m12.sum()\n        mask=m11 if abs(c11-chunk_seq_len)<abs(c12-chunk_seq_len) else m12\n    else:\n        mask=torch.zeros(raw_coords.shape[1],dtype=torch.bool,device=raw_coords.device)\n    coords=raw_coords[:,mask,:].detach().cpu().numpy()\n    if coords.shape[1]>1:\n        diffs=np.linalg.norm(coords[0,1:]-coords[0,:-1],axis=-1)\n        if np.all(diffs<1e-4): print(\"  WARNING: collapsed coords\"); return None\n    if coords.shape[1]!=chunk_seq_len:\n        if coords.shape[1]==1 and chunk_seq_len>1: return None\n        padded=np.zeros((coords.shape[0],chunk_seq_len,3),dtype=np.float32)\n        ml=min(coords.shape[1],chunk_seq_len); padded[:,:ml,:]=coords[:,:ml,:]\n        coords=padded\n    return coords\n\n\ndef protenix_phase(protenix_queue, runner, init_configs, monomer_info_map):\n    from protenix.data.inference.infer_dataloader import InferenceDataset\n    from runner.inference import update_inference_configs\n\n    work_dir = Path(\"/kaggle/working\")\n    protenix_preds = {}\n    wall_elapsed   = time.time() - WALL_CLOCK_START\n\n    # Build task list — skip \"skip\" route, expand chunk for long sequences\n    tasks = []; chunk_info = {}\n    for tid,(n_needed,pseq,n_copies,route) in protenix_queue.items():\n        if route == \"skip\":\n            print(f\"  {tid}: skipped (seq too long for Protenix)\")\n            protenix_preds[tid] = None; continue\n        wall_remaining = TOTAL_TIME_BUDGET_S - (time.time()-WALL_CLOCK_START)\n        if wall_remaining < 300:\n            print(f\"  {tid}: skipped (wall budget exhausted)\"); protenix_preds[tid]=None; continue\n\n        seq_len = len(pseq)\n        if seq_len <= MAX_SEQ_LEN:\n            tasks.append({\"target_id\":tid,\"sequence\":pseq})\n            chunk_info[tid] = [{\"name\":tid,\"range\":(0,seq_len),\"n_copies\":n_copies}]\n        else:\n            chunks = split_into_chunks(seq_len, MAX_SEQ_LEN, CHUNK_OVERLAP)\n            chunk_info[tid] = []\n            for ci,(cs,ce) in enumerate(chunks):\n                cname=f\"{tid}_chunk{ci}\"\n                tasks.append({\"target_id\":cname,\"sequence\":pseq[cs:ce]})\n                chunk_info[tid].append({\"name\":cname,\"range\":(cs,ce),\"n_copies\":n_copies})\n\n    if not tasks:\n        return protenix_preds\n\n    tasks_df = pd.DataFrame(tasks)\n    input_json_path = str(work_dir/\"protenix_input.json\")\n    build_input_json(tasks_df, input_json_path)\n\n    # Re-initialize dataset with actual tasks\n    run_cfg = build_configs(input_json_path, str(work_dir/\"outputs\"), MODEL_NAME, N_SAMPLE)\n    from runner.inference import update_gpu_compatible_configs\n    run_cfg = update_gpu_compatible_configs(run_cfg)\n    runner.update_model_configs(run_cfg)\n    dataset = InferenceDataset(run_cfg)\n\n    raw_preds = {}\n    for i in tqdm(range(len(dataset)), desc=\"Protenix\"):\n        t_start = time.time()\n        data, atom_array, err = dataset[i]\n        sample_name = data.get(\"sample_name\", f\"sample_{i}\")\n        if err:\n            print(f\"  {sample_name}: data error: {err}\")\n            raw_preds[sample_name]=None; del data,atom_array,err\n            gc.collect(); torch.cuda.empty_cache(); continue\n        tid_base    = sample_name.split(\"_chunk\")[0] if \"_chunk\" in sample_name else sample_name\n        n_needed    = protenix_queue.get(tid_base,(N_SAMPLE,\"\",1,\"full\"))[0]\n        sub_seq_len = data[\"N_token\"].item()\n        try:\n            new_cfg = update_inference_configs(run_cfg, sub_seq_len)\n            new_cfg.sample_diffusion.N_sample = n_needed\n            runner.update_model_configs(new_cfg)\n            pred = runner.predict(data)\n            coords = _extract_c1_coords(pred, data[\"input_feature_dict\"], sub_seq_len, pred[\"coordinate\"])\n            raw_preds[sample_name] = coords\n        except Exception as exc:\n            print(f\"  {sample_name}: inference failed: {exc}\")\n            raw_preds[sample_name] = None\n        finally:\n            try: del pred,data,atom_array\n            except: pass\n            gc.collect(); torch.cuda.empty_cache()\n\n    # Post-process: stitch chunks, tile mono→complex\n    for tid,(n_needed,pseq,n_copies,route) in protenix_queue.items():\n        if tid not in chunk_info: continue\n        chunks = chunk_info[tid]\n        mono_len = len(pseq)\n\n        if len(chunks)==1:\n            coords = raw_preds.get(chunks[0][\"name\"])\n            if coords is not None and n_copies>1:\n                # Tile monomer → full complex for each sample\n                tiled = np.stack([tile_monomer_coords(coords[s],n_copies,mono_len) for s in range(coords.shape[0])],0)\n                coords = tiled\n            protenix_preds[tid]=coords\n            if coords is not None: print(f\"  {tid}: {coords.shape[0]} predictions\")\n            else: print(f\"  {tid}: FAILED\")\n        else:\n            chunk_results = {s:[] for s in range(n_needed)}\n            all_ok = True\n            for cinfo in chunks:\n                ccoords=raw_preds.get(cinfo[\"name\"])\n                if ccoords is None: all_ok=False; break\n                for s_idx in range(n_needed):\n                    c = ccoords[s_idx] if s_idx<ccoords.shape[0] else ccoords[-1]\n                    chunk_results[s_idx].append((c, cinfo[\"range\"]))\n            if not all_ok:\n                print(f\"  {tid}: chunk incomplete → fallback\")\n                protenix_preds[tid]=None; continue\n            stitched=[]\n            for s_idx in range(n_needed):\n                items=chunk_results[s_idx]\n                full=stitch_chunk_coords([c for c,_ in items],[r for _,r in items],mono_len)\n                if n_copies>1: full=tile_monomer_coords(full,n_copies,mono_len)\n                stitched.append(full)\n            result=np.stack(stitched,0)\n            protenix_preds[tid]=result\n            print(f\"  {tid}: {result.shape[0]} stitched predictions\")\n\n    return protenix_preds\n\nprint(\"Protenix phase ready.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-21T11:58:36.898106Z","iopub.execute_input":"2026-03-21T11:58:36.898379Z","iopub.status.idle":"2026-03-21T11:58:36.934124Z","shell.execute_reply.started":"2026-03-21T11:58:36.898353Z","shell.execute_reply":"2026-03-21T11:58:36.932823Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ══════════════════════════════════════════════════════════════\n# MAIN\n# ══════════════════════════════════════════════════════════════\n\ndef main():\n    code_dir = os.environ.get(\"PROTENIX_CODE_DIR\", DEFAULT_CODE_DIR)\n    root_dir = os.environ.get(\"PROTENIX_ROOT_DIR\", DEFAULT_ROOT_DIR)\n    if not os.path.isdir(code_dir):\n        raise FileNotFoundError(f\"Missing PROTENIX_CODE_DIR: {code_dir}\")\n    os.environ[\"PROTENIX_ROOT_DIR\"] = root_dir\n    sys.path.append(code_dir)\n    ensure_required_files(root_dir)\n    seed_everything(SEED)\n\n    test_df      = pd.read_csv(DEFAULT_TEST_CSV).reset_index(drop=True)\n    train_seqs   = pd.read_csv(DEFAULT_TRAIN_CSV)\n    val_seqs     = pd.read_csv(DEFAULT_VAL_CSV)\n    train_labels = pd.read_csv(DEFAULT_TRAIN_LBLS)\n    val_labels   = pd.read_csv(DEFAULT_VAL_LBLS)\n\n    combined_seqs   = pd.concat([train_seqs,val_seqs],    ignore_index=True)\n    combined_labels = pd.concat([train_labels,val_labels], ignore_index=True)\n    train_coords    = process_labels(combined_labels)\n    segments_map    = build_segments_map(test_df)\n    monomer_info_map= {row[\"target_id\"]:get_monomer_info(row) for _,row in test_df.iterrows()}\n\n    print(\"Homo-oligomers detected:\")\n    for tid,(mseq,nc) in monomer_info_map.items():\n        if nc>1: print(f\"  {tid}: {nc}x{len(mseq)} nt\")\n    print(f\"Template pool: {len(combined_seqs)} seqs, {len(train_coords)} structures\")\n\n    # ── Phase 1: TBM ────────────────────────────────────────────\n    template_preds, protenix_queue = tbm_phase(\n        test_df, combined_seqs, train_coords, segments_map, monomer_info_map\n    )\n\n    # ── Phase 2: Protenix ────────────────────────────────────────\n    protenix_preds = {}\n    if protenix_queue and USE_PROTENIX:\n        print(f\"\\n{'='*60}\")\n        print(f\"PHASE 2: Protenix ({len(protenix_queue)} targets)\")\n        print(f\"{'='*60}\")\n        work_dir = Path(\"/kaggle/working\")\n        # Initialize runner with a small dummy sequence\n        dummy_json = str(work_dir/\"dummy.json\")\n        dummy_df   = test_df.head(1).copy(); dummy_df[\"sequence\"]=dummy_df[\"sequence\"].str[:64]\n        build_input_json(dummy_df, dummy_json)\n        from runner.inference import InferenceRunner, update_gpu_compatible_configs\n        init_cfg = build_configs(dummy_json, str(work_dir/\"outputs\"), MODEL_NAME, N_SAMPLE)\n        init_cfg = update_gpu_compatible_configs(init_cfg)\n        runner   = InferenceRunner(init_cfg)\n        protenix_preds = protenix_phase(protenix_queue, runner, init_cfg, monomer_info_map)\n    elif protenix_queue:\n        print(\"\\nPHASE 2 skipped (USE_PROTENIX=False)\")\n\n    # ── Phase 3: Combine ─────────────────────────────────────────\n    print(f\"\\n{'='*60}\\nPHASE 3: Combine TBM + Protenix + de-novo\\n{'='*60}\")\n    all_rows = []\n    for _,row in test_df.iterrows():\n        tid=row[\"target_id\"]; seq=row[\"sequence\"]\n        segs=segments_map.get(tid,[(0,len(seq))])\n        combined=list(template_preds.get(tid,[]))\n        ptx=protenix_preds.get(tid)\n        if ptx is not None and isinstance(ptx,np.ndarray) and ptx.ndim==3:\n            for j in range(ptx.shape[0]):\n                if len(combined)>=N_SAMPLE: break\n                combined.append(ptx[j])\n        n_dn=0\n        while len(combined)<N_SAMPLE:\n            sv=row.name*1000000+len(combined)*1000\n            dn=generate_rna_structure(seq,seed=sv)\n            combined.append(adaptive_rna_constraints(dn,segs,confidence=0.2))\n            n_dn+=1\n        if n_dn: print(f\"  {tid}: {n_dn} de-novo slot(s)\")\n        stacked=np.stack(combined[:N_SAMPLE],0)\n        all_rows.extend(coords_to_rows(tid,seq,stacked))\n\n    sub=pd.DataFrame(all_rows)\n    cols=[\"ID\",\"resname\",\"resid\"]+[f\"{c}_{i}\" for i in range(1,N_SAMPLE+1) for c in [\"x\",\"y\",\"z\"]]\n    coord_cols=[c for c in cols if c.startswith((\"x_\",\"y_\",\"z_\"))]\n    sub[coord_cols]=sub[coord_cols].clip(-999.999,9999.999)\n    sub[cols].to_csv(DEFAULT_OUTPUT, index=False)\n    print(f\"\\n✓ Saved {DEFAULT_OUTPUT}  ({len(sub):,} rows)  wall={((time.time()-WALL_CLOCK_START)/60):.1f} min\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-21T11:58:36.935198Z","iopub.execute_input":"2026-03-21T11:58:36.935574Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"---\n## Phase 4 — Official TM-score Validation\n\nRuns only when `len(sol) == len(sub)` (validation set rerun).","metadata":{}},{"cell_type":"code","source":"# ── metric.py auto-search (fixes FileNotFoundError from previous version) ──\nimport glob, runpy\n\ncandidates = glob.glob(\"/kaggle/usr/lib/**/metric.py\", recursive=True)\nif not candidates:\n    # Also check kernel output directory\n    candidates = glob.glob(\"/kaggle/input/**/metric.py\", recursive=True)\n\nif candidates:\n    print(f\"Found metric.py: {candidates[0]}\")\n    module_globals = runpy.run_path(candidates[0])\n    score = module_globals['score']\n    print(\"score() loaded OK.\")\nelse:\n    print(\"metric.py not found. Add 'rhijudas/tm-score-permutechains' kernel output via 'Add input'.\")\n    score = None","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nsol = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/validation_labels.csv')\nsub = pd.read_csv('/kaggle/working/submission.csv')\nprint(f\"sol={len(sol)} rows, sub={len(sub)} rows\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sol['target_id'] = sol['ID'].apply(lambda x: '_'.join(str(x).split('_')[:-1]))\nsub['target_id'] = sub['ID'].apply(lambda x: '_'.join(str(x).split('_')[:-1]))\n\nif score is None:\n    print(\"score() not available — skip validation.\")\nelif len(sol) == len(sub):\n    results = []\n    for target_id, group_native in sol.groupby('target_id'):\n        group_predicted = sub[sub['target_id'] == target_id]\n        result = score(group_native, group_predicted, 'ID')\n        print(f\"{target_id}  {result:.5f}\")\n        results.append(result)\n    print(f'\\nMean TM-score: {sum(results)/len(results):.5f}  (n={len(results)})')\nelse:\n    print(f\"Size mismatch (sol={len(sol)}, sub={len(sub)}) — running against hidden test set.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Citation\nRhiju Das et al. Stanford RNA 3D Folding Part 2. https://kaggle.com/competitions/stanford-rna-3d-folding-2, 2026. Kaggle.  \nOriginal TBM + Protenix pipeline: amirrezaaleyasin","metadata":{}}]}