{"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":5123458,"datasetId":2975803,"databundleVersionId":5194879},{"sourceType":"datasetVersion","sourceId":11118830,"datasetId":6933267,"databundleVersionId":11511771},{"sourceType":"datasetVersion","sourceId":14874339,"datasetId":9502242,"databundleVersionId":15736806},{"sourceType":"datasetVersion","sourceId":11469248,"datasetId":7187409,"databundleVersionId":11912757},{"sourceType":"datasetVersion","sourceId":14519720,"datasetId":9271415,"databundleVersionId":15347344},{"sourceType":"datasetVersion","sourceId":14534919,"datasetId":9283271,"databundleVersionId":15363994},{"sourceType":"modelInstanceVersion","sourceId":311741,"databundleVersionId":11641144,"modelInstanceId":264400,"modelId":285488}],"dockerImageVersionId":31287,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport sys\n\nIS_KAGGLE = bool(os.environ.get(\"KAGGLE_IS_COMPETITION_RERUN\", \"\"))\nLOCAL_N_SAMPLES = 2\n\nif IS_KAGGLE:\n    print(\"Running in KAGGLE COMPETITION mode — all targets processed.\")\nelse:\n    print(f\"Running in LOCAL mode — only {LOCAL_N_SAMPLES} targets.\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-09T17:07:36.692105Z","iopub.execute_input":"2026-03-09T17:07:36.692393Z","iopub.status.idle":"2026-03-09T17:07:36.702689Z","shell.execute_reply.started":"2026-03-09T17:07:36.692367Z","shell.execute_reply":"2026-03-09T17:07:36.701917Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import gc\nimport json\nimport time\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom Bio.Align import PairwiseAligner\nfrom tqdm import tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T17:08:27.698235Z","iopub.execute_input":"2026-03-09T17:08:27.699102Z","iopub.status.idle":"2026-03-09T17:08:32.036303Z","shell.execute_reply.started":"2026-03-09T17:08:27.699068Z","shell.execute_reply":"2026-03-09T17:08:32.035668Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── Paths & Constants ───\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\"\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\n\nos.environ[\"LAYERNORM_TYPE\"] = \"torch\"\nos.environ.setdefault(\"RNA_MSA_DEPTH_LIMIT\", \"512\")\n\nMODEL_NAME    = \"protenix_base_20250630_v1.0.0\"\nPTX_N_SAMPLE  = 5          # Only 1 sample — we just need templates changed to 5\nSEED          = 42\nMAX_SEQ_LEN   = 512\nCHUNK_OVERLAP = 128\n\nUSE_MSA      = \"false\"\nUSE_TEMPLATE = \"false\"\nUSE_RNA_MSA  = \"true\"\n\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.backends.cudnn.enabled = True\n    torch.use_deterministic_algorithms(True)\n\ndef ensure_required_files(root_dir: str) -> None:\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: 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\", encoding=\"utf-8\") as f:\n        json.dump(data, f)\n\ndef build_configs(input_json_path, dump_dir, model_name):\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    arg_str = \" \".join([\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 {PTX_N_SAMPLE}\",\n        f\"--seeds {SEED}\",\n    ])\n    return parse_configs(configs=base, arg_str=arg_str, fill_required_with_null=True)\n\ndef get_c1_mask(data, atom_array):\n    if atom_array is not None:\n        try:\n            if hasattr(atom_array, \"centre_atom_mask\"):\n                m = atom_array.centre_atom_mask == 1\n                if hasattr(atom_array, \"is_rna\"):\n                    m = m & atom_array.is_rna\n                return torch.from_numpy(m).bool()\n            if hasattr(atom_array, \"atom_name\"):\n                base = atom_array.atom_name == \"C1'\"\n                if hasattr(atom_array, \"is_rna\"):\n                    base = base & atom_array.is_rna\n                return torch.from_numpy(base).bool()\n        except Exception:\n            pass\n    f = data[\"input_feature_dict\"]\n    if \"centre_atom_mask\" in f:\n        return (f[\"centre_atom_mask\"] == 1).bool()\n    if \"center_atom_mask\" in f:\n        return (f[\"center_atom_mask\"] == 1).bool()\n    return (f[\"atom_to_tokatom_idx\"] == 11).bool()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T17:09:44.122853Z","iopub.execute_input":"2026-03-09T17:09:44.123302Z","iopub.status.idle":"2026-03-09T17:09:44.140401Z","shell.execute_reply.started":"2026-03-09T17:09:44.123251Z","shell.execute_reply":"2026-03-09T17:09:44.139515Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── TBM (for filling template slots 2-5) ───\ndef _make_aligner():\n    al = PairwiseAligner()\n    al.mode = \"global\"\n    al.match_score = 2\n    al.mismatch_score = -1.5\n    al.open_gap_score = -8\n    al.extend_gap_score = -0.4\n    return al\n\n_aligner = _make_aligner()\n\ndef find_best_templates(query_seq, seq_lookup, coords_dict, top_k=5):\n    results = []\n    for tid, tseq in seq_lookup.items():\n        if tid not in coords_dict:\n            continue\n        if abs(len(tseq) - len(query_seq)) / max(len(tseq), len(query_seq)) > 0.3:\n            continue\n        aln = next(iter(_aligner.align(query_seq, tseq)))\n        norm_score = aln.score / (2 * min(len(query_seq), len(tseq)))\n        results.append((tid, norm_score, coords_dict[tid], tseq))\n    results.sort(key=lambda x: x[1], reverse=True)\n    return results[:top_k]\n\ndef adapt_template_aligned(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):\n            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:\n                w = (i - pv) / (nv - pv)\n                new_coords[i] = (1 - w) * new_coords[pv] + w * new_coords[nv]\n            elif pv >= 0:\n                new_coords[i] = new_coords[pv] + [3.8, 0, 0]\n            elif nv >= 0:\n                new_coords[i] = new_coords[nv] - [3.8, 0, 0]\n            else:\n                new_coords[i] = [i * 3.8, 0, 0]\n    return np.nan_to_num(new_coords)\n\ndef process_labels_chunked(labels_path, chunksize=500000):\n    coords_dict = {}\n    temp_data = {}\n    for chunk in tqdm(pd.read_csv(labels_path, chunksize=chunksize), desc=\"Reading chunks\"):\n        chunk['target_id'] = chunk['ID'].str.rsplit('_', n=1).str[0]\n        for _, row in chunk.iterrows():\n            tid = row['target_id']\n            if tid not in temp_data:\n                temp_data[tid] = []\n            temp_data[tid].append((row['resid'], row['x_1'], row['y_1'], row['z_1']))\n        del chunk\n        gc.collect()\n    for tid, data in temp_data.items():\n        data.sort(key=lambda x: x[0])\n        coords_dict[tid] = np.array([[d[1], d[2], d[3]] for d in data])\n    del temp_data\n    gc.collect()\n    return coords_dict\n\n# ─── Chunking ───\ndef plan_chunks(seq_len, max_len=MAX_SEQ_LEN, overlap=CHUNK_OVERLAP):\n    if seq_len <= max_len:\n        return [(0, seq_len)]\n    chunks = []\n    stride = max_len - overlap\n    start = 0\n    while start < seq_len:\n        end = min(start + max_len, seq_len)\n        chunks.append((start, end))\n        if end == seq_len:\n            break\n        start += stride\n    if len(chunks) > 1 and chunks[-1][1] - chunks[-1][0] < overlap:\n        chunks.pop()\n    if len(chunks) > 1:\n        last_s, last_e = chunks[-1]\n        if last_e - last_s > max_len:\n            chunks[-1] = (last_e - max_len, last_e)\n    return chunks\n\ndef stitch_chunks(chunk_coords_list, chunk_ranges, full_len):\n    full_coords = np.full((full_len, 3), np.nan)\n    if len(chunk_coords_list) == 1:\n        s, e = chunk_ranges[0]\n        clen = min(len(chunk_coords_list[0]), e - s)\n        full_coords[s:s+clen] = chunk_coords_list[0][:clen]\n        return np.nan_to_num(full_coords)\n    s0, e0 = chunk_ranges[0]\n    c0 = chunk_coords_list[0]\n    clen0 = min(len(c0), e0 - s0)\n    full_coords[s0:s0+clen0] = c0[:clen0]\n    placed_end = s0 + clen0\n    for k in range(1, len(chunk_coords_list)):\n        sk, ek = chunk_ranges[k]\n        ck = chunk_coords_list[k]\n        clenk = min(len(ck), ek - sk)\n        overlap_start = sk\n        overlap_end = min(placed_end, sk + clenk)\n        overlap_len = overlap_end - overlap_start\n        ck_aligned = ck[:clenk].copy()\n        if overlap_len >= 3:\n            ref_overlap = full_coords[overlap_start:overlap_end]\n            local_overlap = ck[:overlap_len]\n            valid = ~np.any(np.isnan(ref_overlap), axis=1)\n            if valid.sum() >= 3:\n                ref_pts = ref_overlap[valid]\n                chk_pts = local_overlap[valid]\n                ref_c = ref_pts.mean(axis=0)\n                chk_c = chk_pts.mean(axis=0)\n                H = (chk_pts - chk_c).T @ (ref_pts - ref_c)\n                U, S, Vt = np.linalg.svd(H)\n                d = np.linalg.det(Vt.T @ U.T)\n                R = Vt.T @ np.diag([1, 1, np.sign(d)]) @ U.T\n                ck_aligned = (ck[:clenk] - chk_c) @ R.T + ref_c\n        for pos in range(clenk):\n            gpos = sk + pos\n            if gpos >= full_len:\n                break\n            if np.isnan(full_coords[gpos, 0]):\n                full_coords[gpos] = ck_aligned[pos]\n            else:\n                t_linear = pos / max(overlap_len - 1, 1) if overlap_len > 1 else 0.5\n                t = 0.5 * (1 - np.cos(np.pi * t_linear))\n                full_coords[gpos] = (1 - t) * full_coords[gpos] + t * ck_aligned[pos]\n        placed_end = max(placed_end, sk + clenk)\n    return np.nan_to_num(full_coords)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T17:09:49.03459Z","iopub.execute_input":"2026-03-09T17:09:49.034927Z","iopub.status.idle":"2026-03-09T17:09:49.05912Z","shell.execute_reply.started":"2026-03-09T17:09:49.0349Z","shell.execute_reply":"2026-03-09T17:09:49.05838Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── Load test data ───\ntest_df_full = pd.read_csv(DEFAULT_TEST_CSV)\ntest_df = (test_df_full.head(LOCAL_N_SAMPLES) if not IS_KAGGLE else test_df_full).reset_index(drop=True)\nprint(f\"Test targets: {len(test_df)}\")\n\n# Skip TBM — all 5 template slots will be Protenix predictions\ntbm_templates = {}  # Empty — fallback handled in Cell 6\nprint(\"✓ Skipping TBM (all slots from Protenix)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T17:09:52.907623Z","iopub.execute_input":"2026-03-09T17:09:52.907961Z","iopub.status.idle":"2026-03-09T17:16:28.73406Z","shell.execute_reply.started":"2026-03-09T17:09:52.907935Z","shell.execute_reply":"2026-03-09T17:16:28.733341Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── Phase A: Protenix generates slot 1 templates ───\nprint(f\"\\n{'='*60}\")\nprint(\"PHASE A: Protenix inference (N_SAMPLE=1 for template generation)\")\nprint(f\"{'='*60}\")\n\nseed_everything(SEED)\ncode_dir = DEFAULT_CODE_DIR\nroot_dir = DEFAULT_ROOT_DIR\n\nos.environ[\"PROTENIX_ROOT_DIR\"] = root_dir\nsys.path.append(code_dir)\nensure_required_files(root_dir)\n\nfrom protenix.data.inference.infer_dataloader import InferenceDataset\nfrom runner.inference import (InferenceRunner, update_gpu_compatible_configs, update_inference_configs)\n\n# Plan chunks for all targets\nchunk_map = {}\nchunk_entries = []\n\nfor _, row in test_df.iterrows():\n    tid = row['target_id']\n    seq = row['sequence']\n    chunks = plan_chunks(len(seq))\n    chunk_map[tid] = chunks\n    if len(chunks) == 1:\n        chunk_entries.append((tid, tid, seq[:MAX_SEQ_LEN], chunks[0]))\n        print(f\"  {tid} ({len(seq)} nt): 1 chunk\")\n    else:\n        for ci, (cs, ce) in enumerate(chunks):\n            chunk_id = f\"{tid}__chunk{ci}\"\n            chunk_entries.append((chunk_id, tid, seq[cs:ce], (cs, ce)))\n        print(f\"  {tid} ({len(seq)} nt): {len(chunks)} chunks {chunks}\")\n\n# Build input JSON\nchunk_data = []\nfor chunk_id, tid, chunk_seq, _ in chunk_entries:\n    chunk_data.append({\n        \"name\": chunk_id,\n        \"covalent_bonds\": [],\n        \"sequences\": [{\"rnaSequence\": {\"sequence\": chunk_seq, \"count\": 1}}],\n    })\n\nwork_dir = Path(\"/kaggle/working\")\ninput_json_path = str(work_dir / \"protenix_template_input.json\")\nwith open(input_json_path, \"w\") as f:\n    json.dump(chunk_data, f)\n\nconfigs = build_configs(input_json_path, str(work_dir / \"ptx_outputs\"), MODEL_NAME)\nconfigs = update_gpu_compatible_configs(configs)\nrunner = InferenceRunner(configs)\ndataset = InferenceDataset(configs)\n\n# Run Protenix inference\nprotenix_results = {}  # target_id -> (1, seq_len, 3) array\n\nchunk_results = {}\nptx_start = time.time()\n\nfor i in tqdm(range(len(dataset)), desc=\"Protenix (template gen)\"):\n    data, atom_array, error_message = dataset[i]\n    sample_name = data.get(\"sample_name\", f\"sample_{i}\")\n    \n    if error_message:\n        print(f\"  {sample_name}: error — {error_message}\")\n        del data, atom_array, error_message\n        gc.collect(); torch.cuda.empty_cache()\n        continue\n    \n    matching = [e for e in chunk_entries if e[0] == sample_name]\n    if not matching:\n        del data, atom_array, error_message\n        gc.collect(); torch.cuda.empty_cache()\n        continue\n    \n    chunk_id, parent_tid, chunk_seq, chunk_range = matching[0]\n    \n    try:\n        new_cfg = update_inference_configs(configs, data[\"N_token\"].item())\n        new_cfg.sample_diffusion.N_sample = 5\n        runner.update_model_configs(new_cfg)\n        \n        prediction = runner.predict(data)\n        raw_coords = prediction[\"coordinate\"]\n        feat = data[\"input_feature_dict\"]\n        \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            target_len = len(chunk_seq)\n            mask = m11 if abs(m11.sum() - target_len) < abs(m12.sum() - target_len) else m12\n        else:\n            mask = torch.zeros(raw_coords.shape[1], dtype=torch.bool, device=raw_coords.device)\n        \n        coords = raw_coords[:, mask, :].detach().cpu().numpy()\n        chunk_results[chunk_id] = coords\n        print(f\"  {chunk_id}: shape {coords.shape}\")\n        \n    except Exception as exc:\n        print(f\"  {chunk_id}: FAILED — {exc}\")\n        chunk_results[chunk_id] = None\n    finally:\n        del prediction, raw_coords, mask, data, atom_array\n        gc.collect(); torch.cuda.empty_cache()\n\n# Stitch chunks per target\n# Stitch chunks per target\nfor tid in test_df['target_id']:\n    seq = test_df[test_df['target_id'] == tid]['sequence'].iloc[0]\n    full_len = len(seq)\n    chunks = chunk_map[tid]\n    \n    if len(chunks) == 1:\n        cr = chunk_results.get(tid)\n        if cr is not None:\n            coords = cr\n            if coords.shape[1] != full_len:\n                padded = np.zeros((coords.shape[0], full_len, 3), dtype=np.float32)\n                ml = min(coords.shape[1], full_len)\n                padded[:, :ml, :] = coords[:, :ml, :]\n                coords = padded\n            protenix_results[tid] = coords  # Keep all samples\n        else:\n            protenix_results[tid] = None\n    else:\n        chunk_coords = []\n        all_ok = True\n        for ci in range(len(chunks)):\n            cr = chunk_results.get(f\"{tid}__chunk{ci}\")\n            if cr is None:\n                all_ok = False\n                break\n            chunk_coords.append(cr)\n        \n        if all_ok and chunk_coords:\n            n_samples = chunk_coords[0].shape[0]\n            stitched_all = np.zeros((n_samples, full_len, 3), dtype=np.float32)\n            for si in range(n_samples):\n                sample_chunks = [cc[min(si, cc.shape[0]-1)] for cc in chunk_coords]\n                stitched_all[si] = stitch_chunks(sample_chunks, chunks, full_len)\n            protenix_results[tid] = stitched_all\n            print(f\"  {tid}: stitched {len(chunks)} chunks × {n_samples} samples ✓\")\n        else:\n            protenix_results[tid] = None\n\nptx_elapsed = time.time() - ptx_start\nprint(f\"\\n✓ Protenix done in {ptx_elapsed/60:.1f}min\")\nprint(f\"  Successful: {sum(1 for v in protenix_results.values() if v is not None)}/{len(protenix_results)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T17:16:53.927717Z","iopub.execute_input":"2026-03-09T17:16:53.928059Z","iopub.status.idle":"2026-03-09T17:19:52.737112Z","shell.execute_reply.started":"2026-03-09T17:16:53.928031Z","shell.execute_reply":"2026-03-09T17:19:52.736368Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── Build submission_protenix.csv (Protenix templates for RNAPro) ───\nprint(\"\\nBuilding Protenix template CSV...\")\nall_rows = []\n\nfor idx, row in test_df.iterrows():\n    tid = row['target_id']\n    seq = row['sequence']\n    slen = len(seq)\n    \n    # All 5 slots from Protenix\n    ptx_full = protenix_results.get(tid)  # shape: (5, seq_len, 3) or None\n    slots = []\n    \n    if ptx_full is not None and hasattr(ptx_full, 'ndim') and ptx_full.ndim == 2:\n        # Single sample (shouldn't happen with N_SAMPLE=5, but safety)\n        slots.append(ptx_full if len(ptx_full) >= slen else np.vstack([ptx_full, np.zeros((slen - len(ptx_full), 3))]))\n    \n    if ptx_full is not None and hasattr(ptx_full, 'ndim') and ptx_full.ndim == 3:\n        # Multiple samples — use all of them\n        for si in range(ptx_full.shape[0]):\n            sample = ptx_full[si]\n            if len(sample) < slen:\n                sample = np.vstack([sample, np.zeros((slen - len(sample), 3))])\n            elif len(sample) > slen:\n                sample = sample[:slen]\n            slots.append(sample)\n    \n    # Fill remaining from TBM if Protenix didn't produce 5\n    tbm_list = tbm_templates.get(tid, [])\n    for i in range(len(slots), 5):\n        if i - len(slots) < len(tbm_list):\n            slots.append(tbm_list[i - len(slots)])\n        elif slots:\n            slots.append(slots[0] + np.random.normal(0, 0.3, slots[0].shape))\n        else:\n            slots.append(np.array([[10*np.cos(k*0.6), 10*np.sin(k*0.6), k*2.5] for k in range(slen)]))\n    \n    slots = slots[:5]\n    \n    for ri, nt in enumerate(seq):\n        rd = {'ID': f\"{tid}_{ri+1}\", 'resname': nt, 'resid': ri+1}\n        for s in range(5):\n            c = slots[s][ri] if ri < len(slots[s]) else np.array([0.0, 0.0, 0.0])\n            rd[f'x_{s+1}'] = float(np.nan_to_num(c[0]))\n            rd[f'y_{s+1}'] = float(np.nan_to_num(c[1]))\n            rd[f'z_{s+1}'] = float(np.nan_to_num(c[2]))\n        all_rows.append(rd)\n\nsub_ptx = pd.DataFrame(all_rows).sort_values(['ID']).reset_index(drop=True)\ncols = ['ID', 'resname', 'resid'] + [f'{c}_{i}' for i in range(1,6) for c in ['x','y','z']]\nsub_ptx = sub_ptx[cols].fillna(0.0)\nsub_ptx.to_csv('/kaggle/working/submission_protenix.csv', index=False)\nprint(f\"✓ Protenix template CSV: {sub_ptx.shape}\")\n\n# ─── ALSO save as a complete backup submission ───\n# If RNAPro times out, this is our fallback\nsub_ptx.to_csv('/kaggle/working/submission.csv', index=False)\nprint(\"✓ Backup submission.csv saved (Protenix-only fallback)\")\n\n# ─── Clear Protenix from GPU ───\nprint(\"\\nClearing Protenix from GPU...\")\ndel runner, dataset, configs\ndel protenix_results, chunk_results\ngc.collect()\ntorch.cuda.empty_cache()\ngc.collect()\nprint(f\"✓ GPU freed: {torch.cuda.memory_allocated()/1024**2:.0f} MB allocated\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T17:26:42.023298Z","iopub.execute_input":"2026-03-09T17:26:42.024409Z","iopub.status.idle":"2026-03-09T17:26:42.617704Z","shell.execute_reply.started":"2026-03-09T17:26:42.024374Z","shell.execute_reply":"2026-03-09T17:26:42.616925Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 2: Copy RNAPro source code and checkpoint\n!cp -r /kaggle/input/datasets/theoviel/rnapro-src/RNAPro /kaggle/working/\n!cp /kaggle/input/datasets/theoviel/rnapro-src/rnapro-private-best-500m.ckpt /kaggle/working/\n!ls -la /kaggle/working/\nprint(\"✓ Copied RNAPro code and checkpoint\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T17:26:50.5021Z","iopub.execute_input":"2026-03-09T17:26:50.502809Z","iopub.status.idle":"2026-03-09T17:27:32.852468Z","shell.execute_reply.started":"2026-03-09T17:26:50.502779Z","shell.execute_reply":"2026-03-09T17:27:32.851045Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 3: Setup RNAPro (Offline Compatible)\nimport os\nimport sys\n\n# Add RNAPro to Python path\nsys.path.insert(0, '/kaggle/working/RNAPro')\n\n# Change to RNAPro directory for installation\nos.chdir('/kaggle/working/RNAPro')\n!pip install -e . --no-deps -q\nos.chdir('/kaggle/working')\n\nprint(\"✓ RNAPro setup complete\")\n\n# Verify imports work\ntry:\n    from biotite.structure.io import pdbx\n    from rdkit import Chem\n    from Bio.Align import PairwiseAligner\n    print(\"✓ All dependencies imported successfully\")\nexcept ImportError as e:\n    print(f\"✗ Import error: {e}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T17:27:39.227688Z","iopub.execute_input":"2026-03-09T17:27:39.228084Z","iopub.status.idle":"2026-03-09T17:27:45.200003Z","shell.execute_reply.started":"2026-03-09T17:27:39.228049Z","shell.execute_reply":"2026-03-09T17:27:45.198859Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 4: Setup CCD cache (required for structure generation)\n!mkdir -p /kaggle/working/RNAPro/release_data/ccd_cache/\n!cp /kaggle/input/datasets/jaejohn/rnapro-ccd-cache/ccd_cache/components.cif /kaggle/working/RNAPro/release_data/ccd_cache/\n!cp /kaggle/input/datasets/jaejohn/rnapro-ccd-cache/ccd_cache/components.cif.rdkit_mol.pkl /kaggle/working/RNAPro/release_data/ccd_cache/\nprint(\"✓ CCD cache ready\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T17:27:58.22786Z","iopub.execute_input":"2026-03-09T17:27:58.228239Z","iopub.status.idle":"2026-03-09T17:28:02.925222Z","shell.execute_reply.started":"2026-03-09T17:27:58.228205Z","shell.execute_reply":"2026-03-09T17:28:02.924222Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nos.chdir('/kaggle/working/RNAPro')\nos.makedirs('release_data/kaggle', exist_ok=True)\n\nos.system('python preprocess/convert_templates_to_pt_files.py '\n          '--input_csv /kaggle/working/submission_protenix.csv '\n          '--output_name templates.pt')\n\nos.chdir('/kaggle/working')\n\ntemplate_path = '/kaggle/working/RNAPro/release_data/kaggle/templates.pt'\nif os.path.exists(template_path):\n    import torch as _torch\n    tdata = _torch.load(template_path, map_location='cpu', weights_only=False)\n    print(f\"✓ Templates: {os.path.getsize(template_path)/1024:.1f} KB, entries: {len(tdata) if isinstance(tdata, dict) else '?'}\")\n    del tdata\nelse:\n    print(f\"✗ CRITICAL: Template file not found!\")\n\ntest_df_rnapro = pd.read_csv(DEFAULT_TEST_CSV)\nif not IS_KAGGLE:\n    test_df_rnapro = test_df_rnapro.head(LOCAL_N_SAMPLES)\ntest_df_rnapro.to_csv('/kaggle/working/test_sequences_input.csv', index=False)\nprint(f\"✓ Test sequences: {len(test_df_rnapro)}\")\n\n# ─── OPTIMIZED RNAPro script ───\nscript_content = \"\"\"export LAYERNORM_TYPE=torch\nSEED=42\nN_SAMPLE=5\nN_STEP=150\nN_CYCLE=4\nMAX_LEN=700\nDUMP_DIR=\"../output\"\nCHECKPOINT_PATH=\"../rnapro-private-best-500m.ckpt\"\nTEMPLATE_DATA=\"./release_data/kaggle/templates.pt\"\nRNA_MSA_DIR=\"/kaggle/input/competitions/stanford-rna-3d-folding-2/MSA\"\nSEQUENCES_CSV=\"/kaggle/working/test_sequences_input.csv\"\nRIBONANZA_PATH=\"/kaggle/input/models/shujun717/ribonanzanet2/pytorch/alpha/1\"\nMODEL_NAME=\"rnapro_base\"\n\nmkdir -p \"${DUMP_DIR}\"\n\npython3 runner/inference.py \\\\\n    --model_name \"${MODEL_NAME}\" \\\\\n    --seeds ${SEED} \\\\\n    --dump_dir \"${DUMP_DIR}\" \\\\\n    --load_checkpoint_path \"${CHECKPOINT_PATH}\" \\\\\n    --use_msa false \\\\\n    --use_template \"ca_precomputed\" \\\\\n    --model.use_template \"ca_precomputed\" \\\\\n    --model.use_RibonanzaNet2 true \\\\\n    --model.template_embedder.n_blocks 2 \\\\\n    --model.ribonanza_net_path \"${RIBONANZA_PATH}\" \\\\\n    --template_data \"${TEMPLATE_DATA}\" \\\\\n    --template_idx 0 \\\\\n    --rna_msa_dir \"${RNA_MSA_DIR}\" \\\\\n    --model.N_cycle ${N_CYCLE} \\\\\n    --sample_diffusion.N_sample ${N_SAMPLE} \\\\\n    --sample_diffusion.N_step ${N_STEP} \\\\\n    --load_strict true \\\\\n    --num_workers 0 \\\\\n    --triangle_attention \"torch\" \\\\\n    --triangle_multiplicative \"torch\" \\\\\n    --sequences_csv \"${SEQUENCES_CSV}\" \\\\\n    --max_len ${MAX_LEN}\n\necho \"CIF files: $(find ../output -name '*.cif' | wc -l)\"\n\"\"\"\n\nwith open('/kaggle/working/RNAPro/rnapro_inference.sh', 'w') as f:\n    f.write(script_content)\n\nprint(\"✓ RNAPro script ready (N_SAMPLE=3, N_STEP=100, N_CYCLE=4, MAX_LEN=500)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T17:28:48.292939Z","iopub.execute_input":"2026-03-09T17:28:48.293338Z","iopub.status.idle":"2026-03-09T17:28:51.879022Z","shell.execute_reply.started":"2026-03-09T17:28:48.293303Z","shell.execute_reply":"2026-03-09T17:28:51.878196Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob, time\n\nos.chdir('/kaggle/working/RNAPro')\nprint(\"Running RNAPro with Protenix templates...\")\nrnapro_start = time.time()\nos.system('bash rnapro_inference.sh')\nos.chdir('/kaggle/working')\n\nrnapro_elapsed = time.time() - rnapro_start\ncif_files = glob.glob('/kaggle/working/output/**/*.cif', recursive=True)\nprint(f\"\\n✓ RNAPro done in {rnapro_elapsed/60:.1f}min\")\nprint(f\"  CIF files: {len(cif_files)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T17:29:01.777532Z","iopub.execute_input":"2026-03-09T17:29:01.777869Z","iopub.status.idle":"2026-03-09T17:49:50.884446Z","shell.execute_reply.started":"2026-03-09T17:29:01.777843Z","shell.execute_reply":"2026-03-09T17:49:50.883734Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, glob\nimport numpy as np\nimport pandas as pd\nfrom biotite.structure.io import pdbx\n\ndef extract_c1_coordinates(cif_file_path):\n    try:\n        with open(cif_file_path, \"r\") as f:\n            cif_data = pdbx.CIFFile.read(f)\n        atom_array = pdbx.get_structure(cif_data, model=1)\n        names = np.char.strip(atom_array.atom_name.astype(str))\n        c1 = atom_array[names == \"C1'\"]\n        if len(c1) == 0:\n            return None\n        return c1[np.argsort(c1.res_id)].coord\n    except Exception as e:\n        return None\n\ntest_df = pd.read_csv('/kaggle/working/test_sequences_input.csv')\noutput_dir = \"/kaggle/working/output\"\nptx_df = pd.read_csv('/kaggle/working/submission_protenix.csv')\n\nall_rows = []\nstats = {'rnapro': 0, 'protenix_fallback': 0}\n\nfor idx, row in test_df.iterrows():\n    tid = row['target_id']\n    seq = row['sequence']\n    slen = len(seq)\n    \n    # Try RNAPro CIF predictions first\n    clist = []\n    tdir = f\"{output_dir}/{tid}/seed_42/predictions\"\n    if os.path.isdir(tdir):\n        for cp in sorted(glob.glob(f\"{tdir}/{tid}_sample_*.cif\")):\n            co = extract_c1_coordinates(cp)\n            if co is not None:\n                if len(co) < slen:\n                    co = np.vstack([co, np.zeros((slen - len(co), 3))])\n                elif len(co) > slen:\n                    co = co[:slen]\n                clist.append(co)\n    \n    n_rnapro = len(clist)\n    \n    # Fill remaining slots from Protenix backup\n    if len(clist) < 5:\n        trows = ptx_df[ptx_df['ID'].str.startswith(tid + '_')].sort_values('resid')\n        if len(trows) > 0:\n            for si in range(1, 6):\n                if len(clist) >= 5:\n                    break\n                xc, yc, zc = f'x_{si}', f'y_{si}', f'z_{si}'\n                if xc in trows.columns:\n                    tc = trows[[xc, yc, zc]].values.astype(float)\n                    if not (np.all(tc == 0) or np.any(np.isnan(tc))):\n                        if len(tc) < slen:\n                            tc = np.vstack([tc, np.zeros((slen - len(tc), 3))])\n                        elif len(tc) > slen:\n                            tc = tc[:slen]\n                        clist.append(tc)\n    \n    # Last resort\n    if len(clist) == 0:\n        c = np.array([[10*np.cos(k*0.6), 10*np.sin(k*0.6), k*2.5] for k in range(slen)])\n        clist.append(c)\n    \n    while len(clist) < 5:\n        clist.append(clist[-1] + np.random.normal(0, 0.3, clist[-1].shape))\n    \n    if n_rnapro > 0:\n        stats['rnapro'] += 1\n    else:\n        stats['protenix_fallback'] += 1\n    \n    clist = clist[:5]\n    for ri, nt in enumerate(seq):\n        rd = {'ID': f\"{tid}_{ri+1}\", 'resname': nt, 'resid': ri+1}\n        for s in range(5):\n            rd[f'x_{s+1}'] = float(clist[s][ri, 0])\n            rd[f'y_{s+1}'] = float(clist[s][ri, 1])\n            rd[f'z_{s+1}'] = float(clist[s][ri, 2])\n        all_rows.append(rd)\n\nsub = pd.DataFrame(all_rows)\ncols = ['ID','resname','resid'] + [f'{c}_{i}' for i in range(1,6) for c in ['x','y','z']]\nsub = sub[cols].fillna(0.0)\n\n# OVERWRITE the backup with the improved version\nsub.to_csv('/kaggle/working/submission.csv', index=False)\n\nprint(f\"\\nFinal submission: {sub.shape}\")\nprint(f\"  RNAPro refined: {stats['rnapro']}\")\nprint(f\"  Protenix only: {stats['protenix_fallback']}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T17:50:54.658458Z","iopub.execute_input":"2026-03-09T17:50:54.659256Z","iopub.status.idle":"2026-03-09T17:50:55.055898Z","shell.execute_reply.started":"2026-03-09T17:50:54.65922Z","shell.execute_reply":"2026-03-09T17:50:55.054994Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}