{"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":"gpu","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":11899194,"datasetId":7479946,"databundleVersionId":12404228},{"sourceType":"datasetVersion","sourceId":14874339,"datasetId":9502242,"databundleVersionId":15736806},{"sourceType":"datasetVersion","sourceId":14962495,"datasetId":9577097,"databundleVersionId":15833858},{"sourceType":"kernelVersion","sourceId":242152007}],"dockerImageVersionId":31287,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ── Install offline wheels (same as baseline) ─────────────────────────────\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\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\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\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-14T01:28:32.363349Z","iopub.execute_input":"2026-03-14T01:28:32.363853Z","iopub.status.idle":"2026-03-14T01:28:43.262189Z","shell.execute_reply.started":"2026-03-14T01:28:32.363818Z","shell.execute_reply":"2026-03-14T01:28:43.261428Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nRNA 3D Folding — New Architecture\n==================================\nStrategy that targets 0.55+ TM-score:\n\n  Phase 0 — Load jaejohn PDB-wide TBM templates (pre-computed, offline)\n             jaejohn's algorithm searches the *full PDB* RNA structure\n             database (~20k structures) with MMseqs2, vs the baseline\n             which only searches ~5744 competition training structures.\n             This is the single biggest gap to the 0.432 baseline.\n\n  Phase 1 — Run fine-tuned Protenix with TBM templates injected via\n             use_template=true and the jaejohn template CSV.\n             Checkpoint: zoushuxian/protenix-finetuned-rna3db-all-1599\n             (fine-tuned on RNA3DB — 1599 diverse RNA structures)\n\n  Phase 2 — For diversity (best-of-5 benefits from spread):\n             Run fine-tuned Protenix with 5 independent seeds, no\n             templates, use_rna_msa=true to get structurally diverse\n             diffusion samples.\n\n  Phase 3 — Merge: use template-guided predictions for slots where\n             we have good TBM coverage, diffusion samples otherwise.\n             Always output exactly 5 per target.\n\nKey datasets to add in Kaggle UI:\n  1. jaejohn/rna-3d-folds-tbm-only-approach        (kernelVersion — TBM templates CSV)\n  2. zoushuxian/protenix-finetuned-rna3db-all-1599 (dataset — fine-tuned .pt checkpoint)\n  3. qiweiyin/protenix-v1-adjusted                 (dataset — Protenix code base, same as baseline)\n  4. kami1976/biopython-cp312                      (dataset — offline wheel)\n  5. amirrezaaleyasin/biotite                      (dataset — offline wheel)\n  6. amirrezaaleyasin/rdkit-2025-9-5               (dataset — offline wheel)\n\"\"\"\n\nimport gc\nimport json\nimport os\nimport sys\nimport time\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom tqdm import tqdm\n\nos.environ[\"LAYERNORM_TYPE\"] = \"torch\"\nos.environ.setdefault(\"RNA_MSA_DEPTH_LIMIT\", \"512\")\n\n# ─────────────── Paths ───────────────────────────────────────────────────────\nDATA_BASE     = \"/kaggle/input/competitions/stanford-rna-3d-folding-2\"\nTEST_CSV      = f\"{DATA_BASE}/test_sequences.csv\"\nTRAIN_CSV     = f\"{DATA_BASE}/train_sequences.csv\"\nTRAIN_LBLS    = f\"{DATA_BASE}/train_labels.csv\"\nVAL_CSV       = f\"{DATA_BASE}/validation_sequences.csv\"\nVAL_LBLS      = f\"{DATA_BASE}/validation_labels.csv\"\nOUTPUT_CSV    = \"/kaggle/working/submission.csv\"\n\n# Protenix code (same as baseline — qiweiyin/protenix-v1-adjusted)\nPROTENIX_CODE_DIR = (\n    \"/kaggle/input/datasets/qiweiyin/protenix-v1-adjusted\"\n    \"/Protenix-v1-adjust-v2/Protenix-v1-adjust-v2/Protenix-v1\"\n)\nPROTENIX_ROOT_DIR = PROTENIX_CODE_DIR\n\n# ── Fine-tuned checkpoint ────────────────────────────────────────────────────\n# zoushuxian/protenix-finetuned-rna3db-all-1599\n# RNA3DB-fine-tuned Protenix — substantially better RNA geometry than the\n# base checkpoint used in the 0.432 baseline.\n# Add dataset \"zoushuxian/protenix-finetuned-rna3db-all-1599\" in Kaggle UI.\n# The exact .pt filename may vary; we search for it automatically below.\nFINETUNED_CKPT_DIR = \"/kaggle/input/datasets/zoushuxian/protenix-finetuned-rna3db-all-1599\"\n\n# ── jaejohn TBM templates ────────────────────────────────────────────────────\n# Output of jaejohn/rna-3d-folds-tbm-only-approach kernel.\n# Add as a kernelVersion data source in Kaggle UI.\n# The kernel outputs a submission.csv at /kaggle/input/<kernel-slug>/submission.csv\nJAEJOHN_TEMPLATE_CSV = (\n    \"/kaggle/input/notebooks/jaejohn/rna-3d-folds-tbm-only-approach/submission.csv\"\n)\n\n# ── Model name in the base Protenix repo ─────────────────────────────────────\nBASE_MODEL_NAME = \"protenix_base_20250630_v1.0.0\"\n\n# ── Run settings ─────────────────────────────────────────────────────────────\nN_SAMPLE   = 5        # always output 5 predictions\nSEED       = 42\nMAX_SEQ_LEN  = 512    # Protenix token limit per chunk\nCHUNK_OVERLAP = 128\n\nIS_KAGGLE = True\n\n\n# ─────────────── Utilities ───────────────────────────────────────────────────\ndef seed_everything(seed: int) -> None:\n    os.environ[\"CUBLAS_WORKSPACE_CONFIG\"] = \":4096:8\"\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    np.random.seed(seed)\n    torch.backends.cudnn.benchmark = False\n    torch.backends.cudnn.deterministic = True\n    torch.use_deterministic_algorithms(True)\n\n\ndef find_checkpoint(ckpt_dir: str) -> str:\n    \"\"\"Find the .pt file in the fine-tuned checkpoint directory.\"\"\"\n    for p in Path(ckpt_dir).rglob(\"*.pt\"):\n        return str(p)\n    raise FileNotFoundError(\n        f\"No .pt checkpoint found under {ckpt_dir}. \"\n        \"Make sure zoushuxian/protenix-finetuned-rna3db-all-1599 is added as a dataset.\"\n    )\n\n\ndef split_into_chunks(seq_len: int, max_len: int, overlap: int) -> list:\n    if seq_len <= max_len:\n        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)\n        chunks.append((pos, end))\n        if end == seq_len:\n            break\n        pos += step\n    return chunks\n\n\ndef kabsch_align(P: np.ndarray, Q: np.ndarray):\n    cP, cQ = P.mean(0), Q.mean(0)\n    H = (P - cP).T @ (Q - cQ)\n    U, _, Vt = np.linalg.svd(H)\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\n\ndef stitch_chunks(chunk_coords_list: list, chunk_ranges: list, seq_len: int) -> np.ndarray:\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(len(c), seq_len)] = c[:seq_len]\n        return out\n\n    aligned = [chunk_coords_list[0].copy()]\n    for i in range(1, len(chunk_coords_list)):\n        ps, pe = chunk_ranges[i - 1]\n        cs, ce = chunk_ranges[i]\n        ov_s, ov_e = cs, min(pe, ce)\n        ov_len = ov_e - ov_s\n        if ov_len < 3:\n            aligned.append(chunk_coords_list[i].copy())\n            continue\n        prev_ov = aligned[i - 1][ov_s - ps: ov_e - ps]\n        cur_ov  = chunk_coords_list[i][ov_s - cs: ov_e - cs]\n        valid = ~(np.isnan(prev_ov).any(1) | np.isnan(cur_ov).any(1))\n        if valid.sum() < 3:\n            aligned.append(chunk_coords_list[i].copy()); continue\n        R, t = kabsch_align(cur_ov[valid], prev_ov[valid])\n        aligned.append((chunk_coords_list[i] @ R.T) + t)\n\n    full = np.zeros((seq_len, 3), dtype=np.float64)\n    wts  = np.zeros(seq_len, dtype=np.float64)\n    for i, ((s, e), coords) in enumerate(zip(chunk_ranges, aligned)):\n        ae = min(s + len(coords), seq_len)\n        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            rs = chunk_ranges[i+1][0] - s\n            rl = ae - chunk_ranges[i+1][0]\n            if rl > 0 and rs < ul: w[rs:ul] = np.linspace(1, 0, rl)\n        full[s:ae] += coords[:ul] * w[:, None]\n        wts[s:ae]  += w\n    mask = wts > 0\n    full[mask] /= wts[mask, None]\n    return full\n\n\ndef coords_to_rows(target_id: str, seq: str, coords: np.ndarray) -> list:\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]\n                       else (0., 0., 0.))\n            row[f\"x_{s+1}\"] = float(x)\n            row[f\"y_{s+1}\"] = float(y)\n            row[f\"z_{s+1}\"] = float(z)\n        rows.append(row)\n    return rows\n\n\ndef generate_aform_helix(seq: str, seed=None) -> np.ndarray:\n    if seed is not None:\n        np.random.seed(seed)\n    n = len(seq)\n    coords = np.zeros((n, 3))\n    for i in range(n):\n        a = i * 0.6\n        coords[i] = [10.0 * np.cos(a), 10.0 * np.sin(a), i * 2.5]\n    return coords\n\n\n# ─────────────── Phase 0: Load jaejohn TBM templates ─────────────────────────\ndef load_jaejohn_templates(template_csv: str) -> dict:\n    \"\"\"\n    Load the precomputed jaejohn TBM template submission.\n    Returns {target_id: np.ndarray (N_preds, seq_len, 3)}\n\n    The jaejohn submission CSV has the same format as the competition\n    submission: ID, resname, resid, x_1..z_5\n    \"\"\"\n    print(f\"\\n{'='*60}\")\n    print(\"PHASE 0: Loading jaejohn PDB-wide TBM templates\")\n    print(f\"  from: {template_csv}\")\n    print(f\"{'='*60}\")\n\n    if not os.path.exists(template_csv):\n        print(f\"  WARNING: jaejohn template CSV not found at {template_csv}\")\n        print(\"  Will fall back to fine-tuned Protenix only (no TBM templates)\")\n        return {}\n\n    df = pd.read_csv(template_csv)\n    templates = {}\n    id_prefix = df[\"ID\"].str.rsplit(\"_\", n=1).str[0]\n\n    for tid, grp in df.groupby(id_prefix):\n        grp = grp.sort_values(\"resid\")\n        preds = []\n        for k in range(1, 6):\n            if f\"x_{k}\" in grp.columns:\n                coords = grp[[f\"x_{k}\", f\"y_{k}\", f\"z_{k}\"]].values\n                preds.append(coords)\n        if preds:\n            templates[tid] = np.stack(preds, axis=0)  # (n_preds, seq_len, 3)\n\n    n_targets = len(templates)\n    print(f\"  Loaded templates for {n_targets} targets\")\n    for tid, t in list(templates.items())[:5]:\n        print(f\"    {tid}: {t.shape[0]} template predictions, {t.shape[1]} residues\")\n    return templates\n\n\n# ─────────────── Protenix inference helpers ───────────────────────────────────\ndef build_input_json(df: pd.DataFrame, json_path: str) -> None:\n    data = [\n        {\n            \"name\": row[\"target_id\"],\n            \"covalent_bonds\": [],\n            \"sequences\": [{\"rnaSequence\": {\"sequence\": row[\"sequence\"], \"count\": 1}}],\n        }\n        for _, row in df.iterrows()\n    ]\n    with open(json_path, \"w\") as f:\n        json.dump(data, f)\n\n\ndef build_configs(input_json_path, dump_dir, model_name,\n                  use_msa=\"false\", use_template=\"false\", use_rna_msa=\"true\",\n                  n_sample=5, seed=42, checkpoint_path=None):\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\n    base = {**configs_base, **{\"data\": data_configs}, **inference_configs}\n\n    def deep_update(t, p):\n        for k, v in p.items():\n            if isinstance(v, dict) and k in t and isinstance(t[k], dict):\n                deep_update(t[k], v)\n            else:\n                t[k] = v\n\n    deep_update(base, model_configs[model_name])\n\n    arg_parts = [\n        f\"--model_name {model_name}\",\n        f\"--input_json_path {input_json_path}\",\n        f\"--dump_dir {dump_dir}\",\n        f\"--use_msa {use_msa}\",\n        f\"--use_template {use_template}\",\n        f\"--use_rna_msa {use_rna_msa}\",\n        f\"--sample_diffusion.N_sample {n_sample}\",\n        f\"--seeds {seed}\",\n    ]\n    if checkpoint_path:\n        arg_parts.append(f\"--load_checkpoint_path {checkpoint_path}\")\n\n    return parse_configs(configs=base, arg_str=\" \".join(arg_parts),\n                         fill_required_with_null=True)\n\n\ndef extract_c1_coords(prediction, feat, seq_len, raw_coords):\n    \"\"\"Extract C1' atom coordinates from Protenix prediction.\"\"\"\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        mask = m11 if abs(m11.sum() - seq_len) < abs(m12.sum() - seq_len) else m12\n    else:\n        return None\n\n    coords = raw_coords[:, mask, :].detach().cpu().numpy()\n\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):\n            return None   # collapsed\n\n    if coords.shape[1] != seq_len:\n        padded = np.zeros((coords.shape[0], seq_len, 3), dtype=np.float32)\n        ml = min(coords.shape[1], seq_len)\n        padded[:, :ml] = coords[:, :ml]\n        coords = padded\n    return coords\n\n\ndef run_protenix_batch(tasks_df, work_dir, model_name, checkpoint_path,\n                       use_rna_msa=\"true\", n_sample=5, seed=42,\n                       seq_len_map=None, label=\"Protenix\"):\n    \"\"\"\n    Run Protenix on a batch of sequences.\n    Returns {sample_name: np.ndarray (n_sample, seq_len, 3) or None}\n    \"\"\"\n    from protenix.data.inference.infer_dataloader import InferenceDataset\n    from runner.inference import (InferenceRunner,\n                                  update_gpu_compatible_configs,\n                                  update_inference_configs)\n\n    input_json = str(work_dir / f\"{label}_input.json\")\n    dump_dir   = str(work_dir / f\"{label}_outputs\")\n    build_input_json(tasks_df, input_json)\n\n    configs = build_configs(\n        input_json, dump_dir, model_name,\n        use_rna_msa=use_rna_msa,\n        n_sample=n_sample, seed=seed,\n        checkpoint_path=checkpoint_path,\n    )\n    configs = update_gpu_compatible_configs(configs)\n    runner  = InferenceRunner(configs)\n    dataset = InferenceDataset(configs)\n\n    results = {}\n    for i in tqdm(range(len(dataset)), desc=label):\n        data, atom_array, err = dataset[i]\n        name = data.get(\"sample_name\", f\"sample_{i}\")\n        if err:\n            print(f\"  {name} data error: {err}\")\n            results[name] = None\n            gc.collect(); torch.cuda.empty_cache()\n            continue\n\n        sl = data[\"N_token\"].item()\n        try:\n            new_cfg = update_inference_configs(configs, sl)\n            new_cfg.sample_diffusion.N_sample = n_sample\n            runner.update_model_configs(new_cfg)\n            pred = runner.predict(data)\n            coords = extract_c1_coords(pred, data[\"input_feature_dict\"],\n                                       sl, pred[\"coordinate\"])\n            results[name] = coords\n        except Exception as exc:\n            print(f\"  {name} failed: {exc}\")\n            import traceback; traceback.print_exc()\n            results[name] = None\n        finally:\n            try: del pred, data, atom_array\n            except: pass\n            gc.collect(); torch.cuda.empty_cache()\n\n    return results\n\n\ndef collect_protenix_predictions(raw_preds, target_id, full_seq_len, chunk_info):\n    \"\"\"Stitch chunked Protenix predictions back into full-length coordinates.\"\"\"\n    chunks = chunk_info.get(target_id, [])\n    if not chunks:\n        return None\n\n    if len(chunks) == 1:\n        return raw_preds.get(target_id)\n\n    # Multi-chunk: stitch per-sample\n    n_samples = None\n    chunk_coords_per_sample = {}\n\n    for cinfo in chunks:\n        cname  = cinfo[\"name\"]\n        crange = cinfo[\"range\"]\n        cc     = raw_preds.get(cname)\n        if cc is None:\n            return None  # any missing chunk → fallback\n        if n_samples is None:\n            n_samples = cc.shape[0]\n            for s in range(n_samples):\n                chunk_coords_per_sample[s] = []\n        for s in range(n_samples):\n            idx = s if s < cc.shape[0] else cc.shape[0] - 1\n            chunk_coords_per_sample[s].append((cc[idx], crange))\n\n    stitched = []\n    for s in range(n_samples):\n        items = chunk_coords_per_sample[s]\n        stitched.append(\n            stitch_chunks([c for c, _ in items],\n                          [r for _, r in items],\n                          full_seq_len)\n        )\n    return np.stack(stitched, axis=0)\n\n\n# ─────────────── Main ────────────────────────────────────────────────────────\ndef main():\n    print(f\"\\n{'='*60}\")\n    print(\"RNA 3D Folding — New Architecture\")\n    print(\"  Targets: 0.55+ TM-score\")\n    print(\"  Strategy: jaejohn PDB TBM + fine-tuned Protenix (RNA3DB)\")\n    print(f\"{'='*60}\")\n\n    # ── Setup ──────────────────────────────────────────────────────────────\n    if not os.path.isdir(PROTENIX_CODE_DIR):\n        raise FileNotFoundError(f\"Protenix code dir missing: {PROTENIX_CODE_DIR}\")\n\n    os.environ[\"PROTENIX_ROOT_DIR\"] = PROTENIX_ROOT_DIR\n    sys.path.append(PROTENIX_CODE_DIR)\n    seed_everything(SEED)\n\n    work_dir = Path(\"/kaggle/working\")\n    work_dir.mkdir(parents=True, exist_ok=True)\n\n    # ── Find fine-tuned checkpoint ─────────────────────────────────────────\n    try:\n        ckpt_path = find_checkpoint(FINETUNED_CKPT_DIR)\n        print(f\"\\nFine-tuned checkpoint: {ckpt_path}\")\n    except FileNotFoundError as e:\n        print(f\"\\nWARNING: {e}\")\n        print(\"Falling back to base Protenix checkpoint.\")\n        ckpt_path = None\n\n    # ── Load test sequences ────────────────────────────────────────────────\n    test_df = pd.read_csv(TEST_CSV).reset_index(drop=True)\n    print(f\"\\nTest targets: {len(test_df)}\")\n\n    # ── Phase 0: Load jaejohn templates ───────────────────────────────────\n    tbm_templates = load_jaejohn_templates(JAEJOHN_TEMPLATE_CSV)\n    # tbm_templates: {target_id: (N_preds, seq_len, 3)}\n    # Each prediction is a full-length coordinate set from a PDB template.\n\n    # ── Phase 1: Fine-tuned Protenix — ONLY for targets with weak/missing TBM ──\n    #\n    # jaejohn's TBM already gives 5 valid predictions for every target.\n    # Running Protenix on ALL 28 targets takes ~3h and causes timeout on P100.\n    #\n    # Fix: only run Protenix on targets where TBM quality is low\n    # (bond-distance std > 5 Å in majority of predictions = exploded template gaps).\n    # For all other targets, jaejohn's 5 predictions are used directly as-is.\n    # This cuts runtime from ~3h to ~20-40 min.\n\n    print(f\"\\n{'='*60}\")\n    print(\"PHASE 1: Fine-tuned Protenix (RNA3DB checkpoint)\")\n    print(\"  Only for targets where TBM quality is insufficient\")\n    print(f\"{'='*60}\")\n\n    def tbm_quality_ok(coords_NxSx3: np.ndarray, max_bond_std: float = 5.0) -> bool:\n        \"\"\"True if at least 3 of the 5 TBM predictions look structurally valid.\"\"\"\n        if coords_NxSx3 is None or coords_NxSx3.shape[1] < 2:\n            return False\n        good = 0\n        for k in range(coords_NxSx3.shape[0]):\n            diffs = np.linalg.norm(np.diff(coords_NxSx3[k], axis=0), axis=1)\n            if diffs.std() < max_bond_std:\n                good += 1\n        return good >= 3\n\n    protenix_rows = []\n    for _, row in test_df.iterrows():\n        tid = row[\"target_id\"]\n        tbm = tbm_templates.get(tid)\n        if tbm is None or not tbm_quality_ok(tbm):\n            protenix_rows.append(row)\n            reason = \"missing\" if tbm is None else \"low quality\"\n            print(f\"  {tid} ({len(row['sequence'])} nt): TBM {reason} → Protenix queued\")\n        else:\n            print(f\"  {tid} ({len(row['sequence'])} nt): TBM ok → skipping Protenix\")\n\n    print(f\"\\n  Protenix queue: {len(protenix_rows)}/{len(test_df)} targets\")\n\n    tasks, chunk_info = [], {}\n    for row in protenix_rows:\n        tid, seq = row[\"target_id\"], row[\"sequence\"]\n        slen = len(seq)\n        if slen <= MAX_SEQ_LEN:\n            tasks.append({\"target_id\": tid, \"sequence\": seq})\n            chunk_info[tid] = [{\"name\": tid, \"range\": (0, slen)}]\n        else:\n            chunks = split_into_chunks(slen, 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\": seq[cs:ce]})\n                chunk_info[tid].append({\"name\": cname, \"range\": (cs, ce)})\n\n    if tasks:\n        tasks_df = pd.DataFrame(tasks)\n        raw_preds_p1 = run_protenix_batch(\n            tasks_df, work_dir,\n            model_name=BASE_MODEL_NAME,\n            checkpoint_path=ckpt_path,\n            use_rna_msa=\"true\",\n            n_sample=N_SAMPLE,\n            seed=42,\n            label=\"Phase1_finetuned\",\n        )\n    else:\n        print(\"  All targets covered by TBM — skipping Protenix entirely.\")\n        raw_preds_p1 = {}\n    # Phase 2 removed — one Protenix run is sufficient.\n    # Slot budget: 3 TBM (jaejohn) + 2 Protenix(MSA) = 5.\n    # Running Protenix twice would double runtime and risk timeout on Kaggle P100.\n    raw_preds_p2 = {}\n\n    # ── Phase 3: Assemble final 5 predictions per target ──────────────────\n    print(f\"\\n{'='*60}\")\n    print(\"PHASE 3: Assemble final 5 predictions\")\n    print(\"  Priority: TBM(jaejohn)[slots 1-3] > Protenix/MSA[slots 4-5] > helix fallback\")\n    print(f\"{'='*60}\")\n\n    all_rows = []\n\n    for _, row in test_df.iterrows():\n        tid  = row[\"target_id\"]\n        seq  = row[\"sequence\"]\n        slen = len(seq)\n\n        combined = []\n\n        # ── Slot filling strategy:\n        #    jaejohn's TBM gives the best structural predictions for targets\n        #    that have PDB homologs. We use them in the first slots.\n        #    For targets with no PDB template, Protenix fills all 5 slots.\n        #    The second Protenix run (no MSA, different seed) gives structural\n        #    diversity to maximize best-of-5 TM-score.\n\n        # 1. jaejohn TBM templates (up to 3 slots — they're high quality but\n        #    correlated; save 2 slots for Protenix diversity)\n        tbm = tbm_templates.get(tid)\n        if tbm is not None:\n            for j in range(min(3, tbm.shape[0])):\n                c = tbm[j]\n                if len(c) >= slen:\n                    combined.append(c[:slen])\n                else:\n                    # pad with helix extension if template is shorter\n                    padded = np.zeros((slen, 3))\n                    padded[:len(c)] = c\n                    for k in range(len(c), slen):\n                        a = k * 0.6\n                        padded[k] = [10.0*np.cos(a), 10.0*np.sin(a), k*2.5]\n                    combined.append(padded)\n\n        # 2. Fine-tuned Protenix with RNA-MSA (up to N_SAMPLE - len(combined))\n        p1 = collect_protenix_predictions(raw_preds_p1, tid, slen, chunk_info)\n        if p1 is not None:\n            for j in range(p1.shape[0]):\n                if len(combined) >= N_SAMPLE:\n                    break\n                combined.append(p1[j])\n\n        # 3. Protenix no-MSA run for diversity\n        p2 = collect_protenix_predictions(raw_preds_p2, tid, slen, chunk_info)\n        if p2 is not None:\n            for j in range(p2.shape[0]):\n                if len(combined) >= N_SAMPLE:\n                    break\n                combined.append(p2[j])\n\n        # 4. A-form helix fallback (should rarely be needed)\n        n_denovo = 0\n        while len(combined) < N_SAMPLE:\n            combined.append(generate_aform_helix(seq, seed=row.name * 1000 + len(combined)))\n            n_denovo += 1\n\n        if n_denovo:\n            print(f\"  {tid}: {n_denovo} slot(s) filled with A-form helix fallback\")\n\n        # Report source mix for this target\n        n_tbm = min(3, len(tbm_templates.get(tid, [])))\n        n_p1  = min(N_SAMPLE - n_tbm, p1.shape[0] if p1 is not None else 0)\n        print(f\"  {tid} ({slen} nt): TBM={n_tbm} Protenix={n_p1} helix={n_denovo}\")\n\n        stacked = np.stack(combined[:N_SAMPLE], axis=0)  # (5, seq_len, 3)\n        all_rows.extend(coords_to_rows(tid, seq, stacked))\n\n    # ── Save submission ────────────────────────────────────────────────────\n    sub = pd.DataFrame(all_rows)\n    cols = ([\"ID\", \"resname\", \"resid\"] +\n            [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[0] in \"xyz\" and \"_\" in c]\n    sub[coord_cols] = sub[coord_cols].clip(-999.999, 9999.999)\n    sub[cols].to_csv(OUTPUT_CSV, index=False)\n    print(f\"\\n✓ Saved submission to {OUTPUT_CSV}  ({len(sub):,} rows)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T01:28:43.26392Z","iopub.execute_input":"2026-03-14T01:28:43.264201Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    main()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}