{"metadata":{"kernelspec":{"display_name":"Python 3 (ipykernel)","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.12"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":15231210,"isSourceIdPinned":false},{"sourceType":"datasetVersion","sourceId":14604295,"datasetId":9328538,"databundleVersionId":15440074},{"sourceType":"datasetVersion","sourceId":14962460,"datasetId":9577079,"databundleVersionId":15833819},{"sourceType":"datasetVersion","sourceId":14822484,"datasetId":9479395,"databundleVersionId":15679639},{"sourceType":"datasetVersion","sourceId":11118830,"datasetId":6933267,"databundleVersionId":11511771},{"sourceType":"datasetVersion","sourceId":14874339,"datasetId":9502242,"databundleVersionId":15736806},{"sourceType":"datasetVersion","sourceId":14962495,"datasetId":9577097,"databundleVersionId":15833858},{"sourceType":"datasetVersion","sourceId":11922858,"datasetId":7495841,"databundleVersionId":12430654},{"sourceType":"datasetVersion","sourceId":14519720,"datasetId":9271415,"databundleVersionId":15347344},{"sourceType":"datasetVersion","sourceId":14534919,"datasetId":9283271,"databundleVersionId":15363994},{"sourceType":"datasetVersion","sourceId":14805765,"datasetId":9467172,"databundleVersionId":15661298},{"sourceType":"datasetVersion","sourceId":10855324,"datasetId":6742586,"databundleVersionId":11219268},{"sourceType":"kernelVersion","sourceId":300796201,"isSourceIdPinned":false}],"dockerImageVersionId":31260,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"a1558147","cell_type":"markdown","source":"Reference\n\nhttps://www.kaggle.com/code/qiweiyin/protenix-v1-inference-2026\n\nhttps://www.kaggle.com/code/nihilisticneuralnet/0-409-stanford-rna-folding-2-protenix-template\n\nhttps://www.kaggle.com/code/alexxanderlarko/protenix-v1","metadata":{}},{"id":"649c084f","cell_type":"markdown","source":"dependency\n\npip install biotite\npip install rdkit\npip install biopython","metadata":{}},{"id":"b2bab9c6","cell_type":"code","source":"!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\n# copy the prepared binaries into the PATH of the offline kernel\nimport os, shutil, stat\n\nsrc = \"/kaggle/input/notebooks/llkh0a/kalign-hmmer\"            # adjust to your dataset name\nfor fn in [\"kalign\", \"hmmsearch\"]:                 # add other hmmer tools if needed\n    srcpath = os.path.join(src, fn)\n    if os.path.exists(srcpath):\n        dst = f\"/usr/local/bin/{fn}\"\n        shutil.copy(srcpath, dst)\n        os.chmod(dst, stat.S_IRWXU | stat.S_IRGRP | stat.S_IXGRP | stat.S_IROTH | stat.S_IXOTH)\n        print(\"installed\", dst)\n    else:\n        print(\"missing\", srcpath)\n\n# if you included .deb files instead, you can install them in-place:\n# !dpkg -i /kaggle/input/notebooks/llkh0a/kalign-hmmer/kalign_*.deb\n# !dpkg -i /kaggle/input/notebooks/llkh0a/kalign-hmmer/hmmer_*.deb","metadata":{},"outputs":[],"execution_count":null},{"id":"ce29af9a","cell_type":"code","source":"!which kalign    # e.g. prints /usr/bin/kalign\n!which hmmsearch # hmmer installs a suite of tool","metadata":{},"outputs":[],"execution_count":null},{"id":"bc94f100","cell_type":"code","source":"import os\nimport sys\nimport pandas as pd\n\n# ── Local vs Kaggle mode ─────────────────────────────────────────────────────\n# On Kaggle competition rerun, KAGGLE_IS_COMPETITION_RERUN is set to a truthy value.\n# When running locally we do NOT exit — instead we cap the test set to a small\n# number of samples so the notebook finishes quickly.\n\nIS_KAGGLE = bool(os.environ.get(\"KAGGLE_IS_COMPETITION_RERUN\", \"\"))\n\n# How many test samples to use when running locally\nLOCAL_N_SAMPLES = 1\n\nif IS_KAGGLE:\n    print(\"Running in KAGGLE COMPETITION mode — all test targets will be processed.\")\nelse:\n    print(f\"Running in LOCAL mode — only the first {LOCAL_N_SAMPLES} test targets \"\n          f\"will be processed to save time.\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"dac15ed6","cell_type":"code","source":"import gc\nimport json\nimport os\nimport time\nimport shutil\nimport subprocess\nimport csv\n\nos.environ[\"LAYERNORM_TYPE\"] = \"torch\"\nos.environ.setdefault(\"RNA_MSA_DEPTH_LIMIT\", \"512\")\n\nimport sys\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom Bio.Align import PairwiseAligner\nfrom tqdm import tqdm\n","metadata":{},"outputs":[],"execution_count":null},{"id":"f167b4de","cell_type":"code","source":"def get_c1_mask(data: dict, atom_array) -> torch.Tensor:\n    # 1. Try atom_array attributes first\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            \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\n    # 2. Fallback to feature dict\n    f = data[\"input_feature_dict\"]\n    \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        \n    # Heuristic fallback: check which index gives us roughly N_token atoms\n    n_tokens = data.get(\"N_token\", torch.tensor(0)).item()\n    mask11 = (f[\"atom_to_tokatom_idx\"] == 11).bool()\n    mask12 = (f[\"atom_to_tokatom_idx\"] == 12).bool()\n    \n    c11 = mask11.sum().item()\n    c12 = mask12.sum().item()\n    \n    # Return the one closer to N_tokens (likely one per residue)\n    if abs(c11 - n_tokens) < abs(c12 - n_tokens):\n        return mask11\n    else:\n        return mask12\n","metadata":{},"outputs":[],"execution_count":null},{"id":"4ea87e94","cell_type":"markdown","source":"# Config","metadata":{}},{"id":"8cd7477d","cell_type":"code","source":"\n# ─────────────── Paths & Constants ───────────────────────────────────────────\nDATA_BASE              = \"/kaggle/input/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\n\nMODEL_NAME    = \"protenix_base_20250630_v1.0.0\"\nPROTENIX_PREDICTION_SAMPLES = 3  # Generate more samples than needed to pick the best ones\nN_SAMPLE      = 5\nSEED          = 42\nMAX_SEQ_LEN   = int(os.environ.get(\"MAX_SEQ_LEN\",   \"512\"))\nCHUNK_OVERLAP = int(os.environ.get(\"CHUNK_OVERLAP\",  \"64\"))\n\n# TBM quality thresholds — sequences below these get routed to Protenix\nMIN_SIMILARITY = float(os.environ.get(\"MIN_SIMILARITY\", \"-999.0\"))\nMIN_PERCENT_IDENTITY = float(os.environ.get(\"MIN_PERCENT_IDENTITY\", \"0.0\"))\n\n# Set False to skip Protenix and use de-novo fallback instead\nUSE_PROTENIX = True\nUSE_RNAPRO = True  # Enable RNAPro predictions\n\n# ─────────────── Prediction Slot Allocation ──────────────────────────────────\n# Total predictions per target = 5 (N_SAMPLE)\n# TBM = 2 slots, PROTENIX = 1 slot, RNAPro = 2 slots\nTBM = 2\nPROTENIX = 2\nRNAPro = 1\n\n# Verify allocation sums to N_SAMPLE\nassert TBM + PROTENIX + RNAPro == N_SAMPLE, f\"Allocation mismatch: {TBM}+{PROTENIX}+{RNAPro} != {N_SAMPLE}\"\n\nPRESERVE_FOR_PROTENIX = PROTENIX\n\n# RNAPro settings\nRNAPRO_MAX_LEN = 1000  # Sequences longer than this will skip RNAPro\nMIN_TBM_TARGET_LENGTH = 300  # Minimum sequence length for TBM processing\n\ndef parse_bool(value: str, default: bool = False) -> str:\n    v = str(value).strip().lower()\n    if v in {\"1\", \"true\", \"t\", \"yes\", \"y\", \"on\"}:\n        return \"true\"\n    if v in {\"0\", \"false\", \"f\", \"no\", \"n\", \"off\"}:\n        return \"false\"\n    return \"true\" if default else \"false\"\n\n\nUSE_MSA      = parse_bool(os.environ.get(\"USE_MSA\",      \"false\"))\nUSE_TEMPLATE = parse_bool(os.environ.get(\"USE_TEMPLATE\", \"false\"))\nUSE_RNA_MSA  = parse_bool(os.environ.get(\"USE_RNA_MSA\",  \"false\"))\nN_CYCLE      = os.environ.get(\"N_CYCLE\",12)\nN_STEP       = os.environ.get(\"N_STEP\",200)\nMODEL_N_SAMPLE = int(os.environ.get(\"MODEL_N_SAMPLE\", str(N_SAMPLE)))\n\n","metadata":{},"outputs":[],"execution_count":null},{"id":"f6e03d7d","cell_type":"markdown","source":"# Utiities","metadata":{}},{"id":"e7546d31","cell_type":"code","source":"\n# ─────────────── General Utilities ───────────────────────────────────────────\ndef seed_everything(seed: int) -> None:\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    os.environ[\"CUBLAS_WORKSPACE_CONFIG\"] = \":4096:8\"\n    torch.manual_seed(seed)\n    torch.cuda.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\n\ndef resolve_paths():\n    test_csv   = os.environ.get(\"TEST_CSV\",           DEFAULT_TEST_CSV)\n    output_csv = os.environ.get(\"SUBMISSION_CSV\",     DEFAULT_OUTPUT)\n    code_dir   = os.environ.get(\"PROTENIX_CODE_DIR\",  DEFAULT_CODE_DIR)\n    \n    # FIX: Check if we created a writable root in previous cells\n    if os.path.isdir(\"/kaggle/working/protenix_data\"):\n        root_dir = \"/kaggle/working/protenix_data\"\n    elif os.path.isdir(\"/kaggle/working/protenix_root\"):\n        root_dir = \"/kaggle/working/protenix_root\"\n    else:\n        root_dir = os.environ.get(\"PROTENIX_ROOT_DIR\",  DEFAULT_ROOT_DIR)\n        \n    return test_csv, output_csv, code_dir, root_dir\n\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\n\ndef split_and_save_msa(tid: str,\n                       seq: str,\n                       segs: list,\n                       msa_path: str,\n                       out_dir: str,\n                       debug: bool = True) -> list:\n    \"\"\"Splits the Kaggle full-length MSA into chain-specific MSAs for Protenix.\"\"\"\n    if not os.path.exists(msa_path):\n        if debug:\n            print(f\"[DEBUG] {tid}: no MSA file at {msa_path}\")\n        return [None] * len(segs)\n\n    with open(msa_path, \"r\") as f:\n        lines = f.read().splitlines()\n\n    headers, seqs, cur_seq = [], [], []\n    for line in lines:\n        if line.startswith(\">\"):\n            if cur_seq:\n                seqs.append(\"\".join(cur_seq))\n                cur_seq = []\n            headers.append(line)\n        else:\n            cur_seq.append(line.strip())\n    if cur_seq:\n        seqs.append(\"\".join(cur_seq))\n\n    if not seqs or len(seqs[0]) != len(seq):\n        if debug:\n            print(f\"[DEBUG] {tid}: length mismatch \"\n                  f\"(msa {len(seqs[0]) if seqs else 'None'} vs seq {len(seq)})\")\n        return [None] * len(segs)\n\n    if debug:\n        print(f\"[DEBUG] {tid}: read {len(seqs)} sequences, \"\n              f\"full-length={len(seqs[0])}, segments={segs}\")\n\n    os.makedirs(out_dir, exist_ok=True)\n    split_paths = []\n    for i, (s, e) in enumerate(segs):\n        chain_msa_path = os.path.join(out_dir, f\"{tid}_chain_{i}.fasta\")\n        valid_seqs = 0\n        with open(chain_msa_path, \"w\") as f:\n            for h, sq in zip(headers, seqs):\n                sub_sq = sq[s:e]\n                if set(sub_sq) != {\"-\"}:\n                    f.write(f\"{h}\\n{sub_sq}\\n\")\n                    valid_seqs += 1\n        if debug:\n            print(f\"[DEBUG] {tid} chain {i}: span=({s},{e}) \"\n                  f\"len={e-s} valid={valid_seqs}\")\n        split_paths.append(chain_msa_path if valid_seqs > 0 else None)\n\n    return split_paths\n\ndef build_input_json(df: pd.DataFrame, json_path: str, msa_out_dir: str,debug: bool = True) -> None:\n    os.makedirs(msa_out_dir, exist_ok=True)\n    data = []\n    for _, row in df.iterrows():\n        tid = row[\"target_id\"]\n        full_seq = row[\"sequence\"]\n        msa_path = f\"/kaggle/input/stanford-rna-3d-folding-2/MSA/{tid}.MSA.fasta\"\n        \n        segs = get_chain_segments(row)\n        split_msa_paths = split_and_save_msa(tid, full_seq, segs, msa_path, msa_out_dir)\n        \n        sequences_list = []\n        for i, (s, e) in enumerate(segs):\n            chain_seq = full_seq[s:e]\n            seq_entry = {\"sequence\": chain_seq, \"count\": 1}\n            if split_msa_paths[i] is not None:\n                seq_entry[\"unpairedMsaPath\"] = split_msa_paths[i]\n            sequences_list.append({\"rnaSequence\": seq_entry})\n            \n        data.append({\n            \"name\": tid,\n            \"covalent_bonds\": [],\n            \"sequences\": sequences_list,\n        })\n    with open(json_path, \"w\", encoding=\"utf-8\") as f:\n        json.dump(data, f)\n    if debug:\n        with open(json_path) as f:\n            txt = f.read()\n        print(f\"[DEBUG] wrote {len(data)} samples to {json_path}:\\n{txt}\")\n\ndef build_configs(input_json_path: str, dump_dir: str, model_name: str):\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\"--model.N_cycle {N_CYCLE}\",\n        f\"--sample_diffusion.N_step {N_STEP}\",\n        f\"--use_template {USE_TEMPLATE}\",\n        f\"--use_rna_msa {USE_RNA_MSA}\",\n        f\"--sample_diffusion.N_sample {MODEL_N_SAMPLE}\",\n        f\"--seeds {SEED}\",\n    ])\n    return parse_configs(configs=base, arg_str=arg_str, fill_required_with_null=True)\n\n\ndef get_c1_mask(data: dict, atom_array) -> torch.Tensor:\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            \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\n    f = data[\"input_feature_dict\"]\n    \n    if \"center_atom_mask\" in f:\n        return (f[\"center_atom_mask\"] == 1).bool()\n    if \"centre_atom_mask\" in f:\n        return (f[\"centre_atom_mask\"] == 1).bool()\n        \n    return (f[\"atom_to_tokatom_idx\"] == 11).bool()\n\n\ndef get_feature_c1_mask(data: dict) -> torch.Tensor:\n    f = data[\"input_feature_dict\"]\n    if \"centre_atom_mask\" in f:\n        return f[\"centre_atom_mask\"].long() == 1\n    return f[\"atom_to_tokatom_idx\"].long() == 12\n\n\ndef coords_to_rows(target_id: str, seq: str, coords: np.ndarray) -> list:\n    \"\"\"coords shape: (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            if s < coords.shape[0] and i < coords.shape[1]:\n                x, y, z = coords[s, i]\n            else:\n                x, y, z = 0.0, 0.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 pad_samples(coords: np.ndarray, n: int) -> np.ndarray:\n    if coords.shape[0] >= n:\n        return coords[:n]\n    if coords.shape[0] == 0:\n        return np.zeros((n, coords.shape[1], 3), dtype=coords.dtype)\n    extra = np.repeat(coords[:1], n - coords.shape[0], axis=0)\n    return np.concatenate([coords, extra], axis=0)\n\n\ndef split_into_chunks(seq_len: int, max_len: int, overlap: int) -> list:\n    \"\"\"Split a sequence into overlapping (start, end) chunks.\"\"\"\n    if seq_len <= max_len:\n        return [(0, seq_len)]\n    chunks = []\n    step = max_len - overlap\n    pos = 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    \"\"\"Compute optimal rotation R and translation t so that  R @ P + t ≈ Q.\"\"\"\n    centroid_P = P.mean(axis=0)\n    centroid_Q = Q.mean(axis=0)\n    Pc = P - centroid_P\n    Qc = Q - centroid_Q\n    H = Pc.T @ Qc\n    U, _, Vt = np.linalg.svd(H)\n    d = np.linalg.det(Vt.T @ U.T)\n    S = np.eye(3)\n    if d < 0:\n        S[2, 2] = -1\n    R = Vt.T @ S @ U.T\n    t = centroid_Q - R @ centroid_P\n    return R, t\n\n\ndef stitch_chunk_coords(chunk_coords_list: list,\n                        chunk_ranges: list,\n                        seq_len: int) -> np.ndarray:\n    \"\"\"\n    Merge overlapping chunk coordinates into a full sequence geometry.\n    Applies Kabsch alignment on overlapping residues, and smoothly\n    blends the coordinates using a linear weight ramp.\n    \"\"\"\n    if len(chunk_coords_list) == 1:\n        coords = chunk_coords_list[0]\n        if coords.shape[0] >= seq_len:\n            return coords[:seq_len]\n        out = np.zeros((seq_len, 3), dtype=coords.dtype)\n        out[:coords.shape[0]] = coords\n        return out\n\n    aligned = [chunk_coords_list[0].copy()]\n\n    for i in range(1, len(chunk_coords_list)):\n        prev_start, prev_end = chunk_ranges[i - 1]\n        cur_start, cur_end = chunk_ranges[i]\n\n        ov_start = cur_start\n        ov_end = min(prev_end, cur_end)\n        ov_len = ov_end - ov_start\n\n        if ov_len < 3:\n            aligned.append(chunk_coords_list[i].copy())\n            continue\n\n        prev_ov = aligned[i - 1][ov_start - prev_start: ov_end - prev_start]\n        cur_ov = chunk_coords_list[i][ov_start - cur_start: ov_end - cur_start]\n\n        valid = ~(np.isnan(prev_ov).any(axis=1) | np.isnan(cur_ov).any(axis=1))\n        if valid.sum() < 3:\n            aligned.append(chunk_coords_list[i].copy())\n            continue\n\n        R, t = kabsch_align(cur_ov[valid], prev_ov[valid])\n        transformed = (chunk_coords_list[i] @ R.T) + t\n        aligned.append(transformed)\n\n    full = np.zeros((seq_len, 3), dtype=np.float64)\n    weights = np.zeros(seq_len, dtype=np.float64)\n\n    for i, ((s, e), coords) in enumerate(zip(chunk_ranges, aligned)):\n        chunk_len = coords.shape[0]\n        actual_end = min(s + chunk_len, seq_len)\n        used_len = actual_end - s\n\n        w = np.ones(used_len, dtype=np.float64)\n\n        if i > 0:\n            ov_start = s\n            ov_end = min(chunk_ranges[i - 1][1], e)\n            ramp_len = ov_end - ov_start\n            if ramp_len > 0:\n                w[:ramp_len] = np.linspace(0.0, 1.0, ramp_len)\n\n        if i < len(chunk_ranges) - 1:\n            next_s = chunk_ranges[i + 1][0]\n            ramp_start = next_s - s\n            ramp_len = actual_end - next_s\n            if ramp_len > 0 and ramp_start < used_len:\n                w[ramp_start:used_len] = np.linspace(1.0, 0.0, ramp_len)\n\n        full[s:actual_end] += coords[:used_len] * w[:, None]\n        weights[s:actual_end] += w\n\n    mask = weights > 0\n    full[mask] /= weights[mask, None]\n\n    return full\n\n\n# ─────────────── TBM Core Functions ──────────────────────────────────────────\ndef _make_aligner(seq_len: int = None) -> PairwiseAligner:\n    al = PairwiseAligner()\n    al.mode                           = \"global\"\n    al.match_score                    = 2\n    al.mismatch_score                 = -1.5\n    open_gap = -8.0\n    extend_gap = -0.4\n\n    if seq_len is not None:\n        length_factor = max(0.5, min(2.0, 1.0 + (seq_len - 100) / 1000.0))\n        open_gap *= length_factor\n        extend_gap *= length_factor\n\n    al.open_gap_score                 = open_gap\n    al.extend_gap_score               = extend_gap\n    al.query_left_open_gap_score      = open_gap\n    al.query_left_extend_gap_score    = extend_gap\n    al.query_right_open_gap_score     = open_gap\n    al.query_right_extend_gap_score   = extend_gap\n    al.target_left_open_gap_score     = open_gap\n    al.target_left_extend_gap_score   = extend_gap\n    al.target_right_open_gap_score    = open_gap\n    al.target_right_extend_gap_score  = extend_gap\n    return al\n\n\n_aligner = _make_aligner()\n\n\ndef parse_stoichiometry(stoich: str) -> list:\n    if pd.isna(stoich) or str(stoich).strip() == \"\":\n        return []\n    return [(ch.strip(), int(cnt)) for part in str(stoich).split(\";\")\n            for ch, cnt in [part.split(\":\")]]\n\n\ndef parse_fasta(fasta_content: str) -> dict:\n    out, cur, parts = {}, None, []\n    for line in str(fasta_content).splitlines():\n        line = line.strip()\n        if not line:\n            continue\n        if line.startswith(\">\"):\n            if cur is not None:\n                out[cur] = \"\".join(parts)\n            cur = line[1:].split()[0]\n            parts = []\n        else:\n            parts.append(line.replace(\" \", \"\"))\n    if cur is not None:\n        out[cur] = \"\".join(parts)\n    return out\n\n\ndef get_chain_segments(row) -> list:\n    seq    = row[\"sequence\"]\n    stoich = row.get(\"stoichiometry\", \"\")\n    all_sq = row.get(\"all_sequences\", \"\")\n    if (pd.isna(stoich) or pd.isna(all_sq)\n            or str(stoich).strip() == \"\" or str(all_sq).strip() == \"\"):\n        return [(0, len(seq))]\n    try:\n        chain_dict = parse_fasta(all_sq)\n        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:\n                return [(0, len(seq))]\n            for _ in range(cnt):\n                segs.append((pos, pos + len(base)))\n                pos += len(base)\n        return segs if pos == len(seq) else [(0, len(seq))]\n    except Exception:\n        return [(0, len(seq))]\n\n\ndef build_segments_map(df: pd.DataFrame) -> tuple:\n    seg_map, stoich_map = {}, {}\n    for _, r in df.iterrows():\n        tid               = r[\"target_id\"]\n        seg_map[tid]      = get_chain_segments(r)\n        raw_s             = r.get(\"stoichiometry\", \"\")\n        stoich_map[tid]   = \"\" if pd.isna(raw_s) else str(raw_s)\n    return seg_map, stoich_map\n\n\ndef process_labels(labels_df: pd.DataFrame) -> dict:\n    coords = {}\n    prefixes = labels_df[\"ID\"].str.rsplit(\"_\", n=1).str[0]\n    for prefix, grp in labels_df.groupby(prefixes):\n        coords[prefix] = grp.sort_values(\"resid\")[[\"x_1\", \"y_1\", \"z_1\"]].values\n    return coords\n\n\ndef _build_aligned_strings(query_seq, template_seq, alignment):\n    q_segs, t_segs = alignment.aligned\n    aq, at, qi, ti = [], [], 0, 0\n    for (qs, qe), (ts, te) in zip(q_segs, t_segs):\n        while qi < qs: aq.append(query_seq[qi]);    at.append(\"-\");              qi += 1\n        while ti < ts: aq.append(\"-\");              at.append(template_seq[ti]); ti += 1\n        for qp, tp in zip(range(qs, qe), range(ts, te)):\n            aq.append(query_seq[qp]); at.append(template_seq[tp])\n        qi, ti = qe, te\n    while qi < len(query_seq):    aq.append(query_seq[qi]);    at.append(\"-\");              qi += 1\n    while ti < len(template_seq): aq.append(\"-\");              at.append(template_seq[ti]); ti += 1\n    return \"\".join(aq), \"\".join(at)\n\n\ndef find_similar_sequences_detailed(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:\n            continue\n        if abs(len(tseq) - len(query_seq)) / max(len(tseq), len(query_seq)) > 0.8:\n            continue\n        local_aligner = _make_aligner(min(len(query_seq), len(tseq)))\n        aln       = next(iter(local_aligner.align(query_seq, tseq)))\n        norm_s    = aln.score / (2 * min(len(query_seq), len(tseq)))\n        identical = sum(\n            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        )\n        pct_id = 100 * identical / len(query_seq)\n        aq, at = _build_aligned_strings(query_seq, tseq, aln)\n        results.append((tid, tseq, norm_s, train_coords_dict[tid], pct_id, aq, at))\n    results.sort(key=lambda x: x[2], reverse=True)\n    return results[:top_n]\n\n\ndef adapt_template_to_query(query_seq, template_seq, template_coords) -> np.ndarray:\n    local_aligner = _make_aligner(min(len(query_seq), len(template_seq)))\n    aln        = next(iter(local_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, 0, 0]\n            elif nv >= 0:\n                new_coords[i] = new_coords[nv] + [3, 0, 0]\n            else:\n                new_coords[i] = [i * 3, 0, 0]\n    return np.nan_to_num(new_coords)\n\n\ndef adaptive_rna_constraints(coords, target_id, segments_map, confidence=1.0, passes=2) -> np.ndarray:\n    X        = coords.copy()\n    segments = segments_map.get(target_id, [(0, len(X))])\n    strength = max(0.75 * (1.0 - min(confidence, 0.97)), 0.02)\n    for _ in range(passes):\n        for s, e in segments:\n            C = X[s:e]; L = e - s\n            if L < 3:\n                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(axis=1)\n                    C[idx] += (0.015 * strength) * vec\n            X[s:e] = C\n    return X\n\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\n\ndef apply_hinge(coords, seg, rng, deg=22, confidence=1.0):\n    s, e = seg; L = e - s\n    if L < 30: return coords\n    deg_scaled = deg * max(0.2, min(2.0, 1.5 - confidence))\n    pivot = s + int(rng.integers(10, L - 10))\n    R = _rotmat(rng.normal(size=3), np.deg2rad(float(rng.uniform(-deg_scaled, deg_scaled))))\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\n\ndef jitter_chains(coords, segs, rng, deg=12, trans=1.5, confidence=1.0):\n    X = coords.copy(); gc_ = X.mean(0, keepdims=True)\n    scale = max(0.2, min(2.0, 1.5 - confidence))\n    deg_scaled = deg * scale\n    trans_scaled = trans * scale\n    for s, e in segs:\n        R     = _rotmat(rng.normal(size=3), np.deg2rad(float(rng.uniform(-deg_scaled, deg_scaled))))\n        shift = rng.normal(size=3); shift = shift / (np.linalg.norm(shift) + 1e-12) * float(rng.uniform(0, trans_scaled))\n        c     = X[s:e].mean(0, keepdims=True)\n        X[s:e] = (X[s:e] - c) @ R.T + c + shift\n    X -= X.mean(0, keepdims=True) - gc_\n    return X\n\n\ndef smooth_wiggle(coords, segs, rng, amp=0.8, confidence=1.0):\n    X = coords.copy()\n    scale = max(0.2, min(2.0, 1.5 - confidence))\n    amp_scaled = amp * scale\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_scaled, (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\n\ndef generate_rna_structure(sequence: str, seed=None) -> np.ndarray:\n    \"\"\"Idealized A-form RNA helix — last-resort de-novo fallback.\"\"\"\n    if seed is not None:\n        np.random.seed(seed)\n    n = len(sequence); coords = np.zeros((n, 3))\n    for i in range(n):\n        ang = i * 0.6\n        coords[i] = [10.0 * np.cos(ang), 10.0 * np.sin(ang), i * 2.5]\n    return coords\n\n\n# ─────────────── Global tracking dictionaries ────────────────────────────────\nTBM_COUNTS = {}\nPROTENIX_COUNTS = {}\nRNAPRO_COUNTS = {}\nSOURCES_BY_TARGET = {} \n\n# ─────────────── TBM Phase ───────────────────────────────────────────────────\ndef tbm_phase(test_df, train_seqs_df, train_coords_dict, segments_map):\n    \"\"\"\n    Phase 1 — Template-Based Modeling.\n\n    Returns\n    -------\n    template_predictions : {target_id: [np.ndarray(seq_len, 3), ...]}\n        0 to TBM predictions per target, from real templates.\n    protenix_queue : {target_id: (n_needed, full_sequence)}\n        Targets that still need Protenix predictions.\n    \"\"\"\n    global TBM_COUNTS\n    print(f\"\\n{'='*60}\")\n    print(f\"PHASE 1: Template-Based Modeling (TBM={TBM} slots)\")\n    print(f\"  MIN_SIMILARITY = {MIN_SIMILARITY}  |  MIN_PCT_IDENTITY = {MIN_PERCENT_IDENTITY}\")\n    print(f\"{'='*60}\")\n    t0 = time.time()\n\n    template_predictions: dict = {}\n    protenix_queue:       dict = {}\n\n    for _, row in test_df.iterrows():\n        tid = row[\"target_id\"]\n        seq = row[\"sequence\"]\n        segs = segments_map.get(tid, [(0, len(seq))])\n\n        similar = find_similar_sequences_detailed(seq, train_seqs_df, train_coords_dict, top_n=30)\n        preds   = []\n        used    = set()\n\n        # Enforce TBM slot limit\n        # FIX: Generate enough TBM candidates to fill all slots if other models fail\n        max_tbm = N_SAMPLE \n\n        for i, (tmpl_id, tmpl_seq, sim, tmpl_coords, pct_id, _, _) in enumerate(similar):\n            if len(preds) >= max_tbm:\n                break\n            if sim < MIN_SIMILARITY or pct_id < MIN_PERCENT_IDENTITY:\n                break\n            if tmpl_id in used:\n                continue\n\n            rng     = np.random.default_rng((abs(hash(tid)) + i * 10007) % (2**32))\n            adapted = adapt_template_to_query(seq, tmpl_seq, tmpl_coords)\n\n            slot = len(preds)\n            if slot == 0:\n                X = adapted\n            elif slot == 1:\n                X = adapted + rng.normal(0, max(0.01, (0.40 - sim) * 0.06), adapted.shape)\n            elif slot == 2:\n                longest = max(segs, key=lambda se: se[1] - se[0])\n                X = apply_hinge(adapted, longest, rng, confidence=sim)\n            elif slot == 3:\n                X = jitter_chains(adapted, segs, rng, confidence=sim)\n            else:\n                X = smooth_wiggle(adapted, segs, rng, confidence=sim)\n\n            refined = adaptive_rna_constraints(X, tid, segments_map, confidence=sim)\n            preds.append(refined)\n            used.add(tmpl_id)\n\n        template_predictions[tid] = preds\n        TBM_COUNTS[tid] = len(preds)\n        \n        # Always request PROTENIX predictions (Protenix slot is after TBM)\n        if USE_PROTENIX:\n            n_needed = PROTENIX\n            protenix_queue[tid] = (n_needed, seq)\n            print(f\"  {tid} ({len(seq)} nt): {len(preds)} TBM → need {n_needed} from Protenix\")\n        else:\n            print(f\"  {tid} ({len(seq)} nt): {len(preds)} TBM ✓\")\n\n    elapsed = time.time() - t0\n    print(f\"\\nPhase 1 done in {elapsed:.1f}s\")\n    print(f\"  Total TBM predictions generated\")\n    return template_predictions, protenix_queue\n\n\n# ─────────────── RNAPro Phase ────────────────────────────────────────────────\ndef rnapro_phase(test_df, template_preds, segments_map):\n    \"\"\"\n    Phase 2.5 — RNAPro Predictions.\n    \n    Uses TBM predictions as templates to generate RNAPro structure predictions.\n    \n    Returns\n    -------\n    rnapro_preds : {target_id: np.ndarray (RNAPro, seq_len, 3) or None}\n    \"\"\"\n    global RNAPRO_COUNTS\n    \n    if not USE_RNAPRO:\n        print(\"\\nRNAPro Phase skipped (USE_RNAPRO=False)\")\n        return {}\n    \n    print(f\"\\n{'='*60}\")\n    print(f\"PHASE 2.5: RNAPro Predictions (RNAPro={RNAPro} slots)\")\n    print(f\"{'='*60}\")\n    \n    import csv\n    import subprocess\n    \n    rnapro_preds: dict = {}\n    rnapro_queue = []\n    \n    # Determine which targets qualify for RNAPro (sequence length < RNAPRO_MAX_LEN)\n    for _, row in test_df.iterrows():\n        tid = row[\"target_id\"]\n        seq = row[\"sequence\"]\n        \n        if len(seq) >= MIN_TBM_TARGET_LENGTH and len(seq) <= RNAPRO_MAX_LEN:\n            rnapro_queue.append((tid, seq))\n            print(f\"  {tid} ({len(seq)} nt): queued for RNAPro\")\n        else:\n            print(f\"  {tid} ({len(seq)} nt): skipped (length outside range {MIN_TBM_TARGET_LENGTH}-{RNAPRO_MAX_LEN})\")\n            rnapro_preds[tid] = None\n            RNAPRO_COUNTS[tid] = 0\n    \n    if not rnapro_queue:\n        print(\"\\nNo targets qualified for RNAPro\")\n        return rnapro_preds\n    \n    # Step 1: Generate TBM submission CSV for RNAPro templates\n    print(\"\\n  Step 1: Generating TBM template file for RNAPro...\")\n    tbm_rows = []\n    for tid, seq in rnapro_queue:\n        preds = template_preds.get(tid, [])\n        # Pad to 5 predictions (RNAPro expects 5 template slots)\n        while len(preds) < 5:\n            if len(preds) > 0:\n                preds.append(preds[-1])  # Repeat last prediction\n            else:\n                # Generate de-novo structure\n                seed_val = hash(tid) % 10000\n                dn = generate_rna_structure(seq, seed=seed_val)\n                preds.append(dn)\n        \n        for i in range(len(seq)):\n            row_data = {\n                \"ID\": f\"{tid}_{i+1}\",\n                \"resname\": seq[i],\n                \"resid\": i + 1\n            }\n            for s in range(5):\n                if i < len(preds[s]):\n                    x, y, z = preds[s][i]\n                else:\n                    x, y, z = 0.0, 0.0, 0.0\n                row_data[f\"x_{s+1}\"] = float(x)\n                row_data[f\"y_{s+1}\"] = float(y)\n                row_data[f\"z_{s+1}\"] = float(z)\n            tbm_rows.append(row_data)\n    \n    # Write TBM submission CSV\n    tbm_csv_path = \"/kaggle/working/submission_tbm.csv\"\n    cols = [\"ID\", \"resname\", \"resid\"] + [f\"{c}_{i}\" for i in range(1, 6) for c in [\"x\", \"y\", \"z\"]]\n    pd.DataFrame(tbm_rows)[cols].to_csv(tbm_csv_path, index=False)\n    print(f\"  Written TBM templates to {tbm_csv_path}\")\n    \n    # Step 2: Convert templates to .pt file for RNAPro\n    print(\"\\n  Step 2: Converting templates to .pt format...\")\n    rnapro_dir = \"/kaggle/working/RNAPro\"\n    template_pt_path = f\"{rnapro_dir}/release_data/kaggle/templates.pt\"\n    os.makedirs(os.path.dirname(template_pt_path), exist_ok=True)\n    \n    # Run conversion script\n    convert_cmd = [\n        \"python\", f\"{rnapro_dir}/preprocess/convert_templates_to_pt_files.py\",\n        \"--input_csv\", tbm_csv_path,\n        \"--output_name\", template_pt_path\n    ]\n    result = subprocess.run(convert_cmd, capture_output=True, text=True, cwd=rnapro_dir)\n    if result.returncode != 0:\n        print(f\"  WARNING: Template conversion failed: {result.stderr}\")\n    else:\n        print(f\"  Templates converted to {template_pt_path}\")\n    \n    # Step 3: Create sequence CSV for RNAPro\n    print(\"\\n  Step 3: Creating RNAPro input sequences...\")\n    rnapro_seq_csv = \"/kaggle/working/rnapro_sequences.csv\"\n    rnapro_df = pd.DataFrame([{\"target_id\": tid, \"sequence\": seq} for tid, seq in rnapro_queue])\n    rnapro_df.to_csv(rnapro_seq_csv, index=False)\n    print(f\"  Written {len(rnapro_queue)} sequences to {rnapro_seq_csv}\")\n    \n    # Step 4: Run RNAPro inference\n    print(\"\\n  Step 4: Running RNAPro inference...\")\n    \n    # Update the shell script with correct paths\n    rnapro_sh = f\"{rnapro_dir}/rnapro_inference_kaggle.sh\"\n    \n    # Change to RNAPro directory and run inference\n    old_cwd = os.getcwd()\n    try:\n        os.chdir(rnapro_dir)\n        # Set PYTHONPATH explicitly so subprocess finds RNAPro modules\n        import copy\n        rnapro_env = copy.copy(os.environ)\n        rnapro_env[\"PYTHONPATH\"] = rnapro_dir + \":\" + rnapro_env.get(\"PYTHONPATH\", \"\")\n        result = subprocess.run([\"bash\", rnapro_sh], capture_output=True, text=True, env=rnapro_env)\n        if result.returncode != 0:\n            print(f\"  RNAPro inference error:\\n{result.stderr}\")\n        else:\n            print(\"  RNAPro inference completed\")\n    finally:\n        os.chdir(old_cwd)\n    \n    # Move output file\n    rnapro_output = f\"{rnapro_dir}/submission.csv\"\n    rnapro_final = \"/kaggle/working/submission_rnapro.csv\"\n    if os.path.exists(rnapro_output):\n        shutil.move(rnapro_output, rnapro_final)\n        print(f\"  RNAPro output saved to {rnapro_final}\")\n    \n    # Step 5: Parse RNAPro output and extract coordinates\n    print(\"\\n  Step 5: Parsing RNAPro predictions...\")\n    if os.path.exists(rnapro_final):\n        rnapro_df = pd.read_csv(rnapro_final)\n        rnapro_df[\"target_id\"] = rnapro_df[\"ID\"].apply(lambda x: \"_\".join(str(x).split(\"_\")[:-1]))\n        \n        for tid, seq in rnapro_queue:\n            group = rnapro_df[rnapro_df[\"target_id\"] == tid].sort_values(\"resid\")\n            if len(group) == 0:\n                rnapro_preds[tid] = None\n                RNAPRO_COUNTS[tid] = 0\n                continue\n            \n            # Extract coordinates for RNAPro predictions (up to RNAPro slots)\n            coords_list = []\n            for slot in range(1, min(RNAPro + 1, 6)):\n                coords = group[[f\"x_{slot}\", f\"y_{slot}\", f\"z_{slot}\"]].values.astype(np.float32)\n                # Check if valid (not all zeros)\n                if np.abs(coords).sum() > 1e-6:\n                    coords_list.append(coords)\n            \n            if coords_list:\n                rnapro_preds[tid] = np.stack(coords_list, axis=0)\n                RNAPRO_COUNTS[tid] = len(coords_list)\n                print(f\"  {tid}: {len(coords_list)} RNAPro predictions extracted\")\n            else:\n                rnapro_preds[tid] = None\n                RNAPRO_COUNTS[tid] = 0\n                print(f\"  {tid}: No valid RNAPro predictions\")\n    else:\n        print(f\"  WARNING: RNAPro output file not found at {rnapro_final}\")\n        for tid, seq in rnapro_queue:\n            rnapro_preds[tid] = None\n            RNAPRO_COUNTS[tid] = 0\n    \n    return rnapro_preds\n\n\n# ─────────────── Main ────────────────────────────────────────────────────────\ndef main() -> None:\n    global PROTENIX_COUNTS\n    \n    test_csv, output_csv, code_dir, root_dir = resolve_paths()\n\n    if not os.path.isdir(code_dir):\n        raise FileNotFoundError(\n            f\"Missing PROTENIX_CODE_DIR: {code_dir}. \"\n            \"Set PROTENIX_CODE_DIR to the repo path.\"\n        )\n\n    os.environ[\"PROTENIX_ROOT_DIR\"] = root_dir\n    if code_dir not in sys.path:\n        sys.path.insert(0, code_dir)\n    seed_everything(SEED)\n    sys.path = [p for p in sys.path if not p.startswith(\"/kaggle/working/RNAPro\")]\n\n    ensure_required_files(root_dir)\n    seed_everything(SEED)\n    # ── Load test data ──────────────────────────────────────────────────────\n    test_df_full = pd.read_csv(test_csv)\n    if not IS_KAGGLE:\n        if TARGET is not None:\n            test_df_full = test_df_full[test_df_full[\"target_id\"].isin(TARGET)]\n        else:\n            test_df_full = test_df_full.head(LOCAL_N_SAMPLES)\n\n    \n    # test_df      = (test_df_full.head(LOCAL_N_SAMPLES) if not IS_KAGGLE\n    #                 else test_df_full).reset_index(drop=True)\n    test_df = test_df_full\n    print(f\"Test targets : {len(test_df)}\"\n          + (\" (LOCAL MODE)\" if not IS_KAGGLE else \"\"))\n    \n    print(f\"\\n=== Prediction Allocation ===\")\n    print(f\"  TBM:      {TBM} slots\")\n    print(f\"  Protenix: {PROTENIX} slot\")\n    print(f\"  RNAPro:   {RNAPro} slots\")\n    print(f\"  Total:    {N_SAMPLE} slots\")\n\n    seq_by_id = dict(zip(test_df[\"target_id\"], test_df[\"sequence\"]))\n\n    # Truncated copy for Protenix (Protenix has token limits)\n    test_df_trunc = test_df.copy()\n    test_df_trunc[\"sequence\"] = test_df_trunc[\"sequence\"].str[:MAX_SEQ_LEN]\n\n    # ── Load training data for TBM ──────────────────────────────────────────\n    print(\"\\nLoading training data for TBM...\")\n    train_seqs   = pd.read_csv(DEFAULT_TRAIN_CSV)\n    train_labels = pd.read_csv(DEFAULT_TRAIN_LBLS)\n\n    if not IS_KAGGLE:\n        combined_seqs   = train_seqs.copy()\n        combined_labels = train_labels.copy()\n        print(\"Using training set only as template pool (LOCAL mode).\")\n    else:\n        val_seqs     = pd.read_csv(DEFAULT_VAL_CSV)\n        val_labels   = pd.read_csv(DEFAULT_VAL_LBLS)\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        print(f\"Using train+val templates ({len(val_seqs)} extra sequences).\")\n    train_coords    = process_labels(combined_labels)\n    segments_map, _ = build_segments_map(test_df)\n\n    print(f\"Template pool: {len(combined_seqs)} sequences, {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\n    )\n\n    # ─── PHASE 2: Protenix ─────────────────────────────────────────────────\n    protenix_preds: dict = {}\n\n    if protenix_queue and USE_PROTENIX:\n        print(f\"\\n{'='*60}\")\n        print(f\"PHASE 2: Protenix for {len(protenix_queue)} targets (PROTENIX={PROTENIX} slot)\")\n        print(f\"{'='*60}\")\n\n        work_dir = Path(\"/kaggle/working\")\n        work_dir.mkdir(parents=True, exist_ok=True)\n\n        tasks = []\n        chunk_info = {}\n        \n        for target_id, (n_needed, full_seq) in protenix_queue.items():\n            seq_len = len(full_seq)\n            if seq_len <= MAX_SEQ_LEN:\n                tasks.append({\"target_id\": target_id, \"sequence\": full_seq[:MAX_SEQ_LEN]})\n                chunk_info[target_id] = [{\"name\": target_id, \"range\": (0, seq_len)}]\n                print(f\"  {target_id} ({seq_len} nt): single pass queued\")\n            else:\n                chunks = split_into_chunks(seq_len, MAX_SEQ_LEN, CHUNK_OVERLAP)\n                print(f\"  {target_id} ({seq_len} nt): {len(chunks)} chunks queued \"\n                      f\"{[(s, e) for s, e in chunks]}\")\n                chunk_info[target_id] = []\n                for ci, (cs, ce) in enumerate(chunks):\n                    chunk_name = f\"{target_id}_chunk{ci}\"\n                    sub_seq = full_seq[cs:ce]\n                    tasks.append({\"target_id\": chunk_name, \"sequence\": sub_seq})\n                    chunk_info[target_id].append({\"name\": chunk_name, \"range\": (cs, ce)})\n\n        tasks_df = pd.DataFrame(tasks)\n        input_json_path = str(work_dir / \"protenix_queue_input.json\")\n        msa_out_dir = str(work_dir / \"split_msas\")\n        build_input_json(tasks_df, input_json_path, msa_out_dir)\n\n        from protenix.data.inference.infer_dataloader import InferenceDataset\n        from runner.inference import (\n            InferenceRunner,\n            update_gpu_compatible_configs,\n            update_inference_configs,\n        )\n\n        configs = build_configs(input_json_path, str(work_dir / \"outputs\"), MODEL_NAME)\n        configs = update_gpu_compatible_configs(configs)\n        runner  = InferenceRunner(configs)\n        dataset = InferenceDataset(configs)\n\n        raw_predictions = {}\n\n        def _extract_c1_coords(data, atom_array, chunk_seq_len, raw_coords):\n            if \"input_feature_dict\" in data:\n                feat = data[\"input_feature_dict\"]\n                if \"profile\" not in feat and \"aatype\" in feat:\n                    aatype = feat[\"aatype\"].long()\n                    feat[\"profile\"] = torch.nn.functional.one_hot(aatype, num_classes=32).float()\n                if \"deletion_mean\" not in feat and \"aatype\" in feat:\n                    feat[\"deletion_mean\"] = torch.zeros((feat[\"aatype\"].shape[0], 1), dtype=torch.float32)\n\n            mask = get_c1_mask(data, atom_array)\n            coords = raw_coords[:, mask, :]\n            if isinstance(coords, torch.Tensor):\n                coords = coords.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                    print(\"    WARNING: Collapsed coordinates detected\")\n                    return None\n\n            if coords.shape[1] != chunk_seq_len:\n                if coords.shape[1] == 1 and chunk_seq_len > 1:\n                    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)\n                padded[:, :ml, :] = coords[:, :ml, :]\n                coords = padded\n            return coords\n\n        for i in tqdm(range(len(dataset)), desc=\"Protenix\"):\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}: data error — {error_message}\")\n                raw_predictions[sample_name] = None\n                del data, atom_array, error_message\n                gc.collect(); torch.cuda.empty_cache(); gc.collect()\n                continue\n\n            target_id = sample_name.split(\"_chunk\")[0] if \"_chunk\" in sample_name else sample_name\n            n_needed = protenix_queue.get(target_id, (PROTENIX, \"\"))[0]\n            sub_seq_len = data[\"N_token\"].item()\n\n            try:\n                new_cfg = update_inference_configs(configs, sub_seq_len)\n                new_cfg.sample_diffusion.N_sample = n_needed\n                runner.update_model_configs(new_cfg)\n\n                prediction = runner.predict(data)\n                raw_coords = prediction[\"coordinate\"]\n                coords = _extract_c1_coords(data, atom_array, sub_seq_len, raw_coords)\n                raw_predictions[sample_name] = coords\n\n            except Exception as exc:\n                print(f\"  {sample_name}: Protenix FAILED — {exc}\")\n                raw_predictions[sample_name] = None\n\n            finally:\n                try: del prediction, raw_coords, data, atom_array\n                except: pass\n                gc.collect(); torch.cuda.empty_cache(); gc.collect()\n\n        for target_id, (n_needed, full_seq) in protenix_queue.items():\n            seq_len = len(full_seq)\n            chunks = chunk_info.get(target_id, [])\n            if not chunks:\n                protenix_preds[target_id] = None\n                PROTENIX_COUNTS[target_id] = 0\n                continue\n\n            if len(chunks) == 1:\n                coords = raw_predictions.get(target_id)\n                if coords is not None:\n                    coords = coords[:PROTENIX]\n                protenix_preds[target_id] = coords\n                PROTENIX_COUNTS[target_id] = 0 if coords is None else coords.shape[0]\n                if coords is not None:\n                    print(f\"  {target_id}: {coords.shape[0]} Protenix predictions generated\")\n                else:\n                    print(f\"  {target_id}: FAILED\")\n            else:\n                chunk_results_per_sample = {s: [] for s in range(n_needed)}\n                all_ok = True\n\n                for cinfo in chunks:\n                    cname = cinfo[\"name\"]\n                    crange = cinfo[\"range\"]\n                    ccoords = raw_predictions.get(cname)\n                    if ccoords is None:\n                        all_ok = False\n                        break\n                    for s_idx in range(n_needed):\n                        if s_idx < ccoords.shape[0]:\n                            chunk_results_per_sample[s_idx].append((ccoords[s_idx], crange))\n                        else:\n                            chunk_results_per_sample[s_idx].append((ccoords[-1], crange))\n\n                if not all_ok:\n                    print(f\"  {target_id}: chunked inference incomplete, using fallback\")\n                    protenix_preds[target_id] = None\n                    PROTENIX_COUNTS[target_id] = 0\n                    continue\n\n                stitched_samples = []\n                for s_idx in range(n_needed):\n                    items = chunk_results_per_sample[s_idx]\n                    coords_list = [c for c, _ in items]\n                    ranges_list = [r for _, r in items]\n                    full_coords = stitch_chunk_coords(coords_list, ranges_list, seq_len)\n                    stitched_samples.append(full_coords)\n\n                result = np.stack(stitched_samples, axis=0)[:PROTENIX]\n                protenix_preds[target_id] = result\n                PROTENIX_COUNTS[target_id] = result.shape[0]\n                print(f\"  {target_id}: {result.shape[0]} stitched Protenix predictions generated\")\n\n    elif not USE_PROTENIX:\n        print(f\"\\nPHASE 2 skipped (USE_PROTENIX=False)\")\n        for tid in protenix_queue:\n            PROTENIX_COUNTS[tid] = 0\n\n    # ─── PHASE 2.5: RNAPro ─────────────────────────────────────────────────\n    rnapro_preds = rnapro_phase(test_df, template_preds, segments_map)\n    global SOURCES_BY_TARGET\n    SOURCES_BY_TARGET.clear()\n    # ─── PHASE 3: Combine TBM + Protenix + RNAPro + de-novo ────────────────\n    print(f\"\\n{'='*60}\")\n    print(\"PHASE 3: Combine TBM + Protenix + RNAPro + de-novo fallback\")\n    print(f\"  Slot allocation: {TBM} TBM + {PROTENIX} Protenix + {RNAPro} RNAPro = {N_SAMPLE}\")\n    print(f\"{'='*60}\")\n\n    all_rows = []\n    for _, row in test_df.iterrows():\n        tid = row[\"target_id\"]\n        seq = row[\"sequence\"]\n\n        tbm_all = template_preds.get(tid, [])\n        combined: list = list(tbm_all)[:TBM]\n        combined_src = [\"TBM\"] * len(combined)          # track origins\n\n        # pad primary TBM slots …\n        while len(combined) < TBM:\n            seed_val = hash(tid) % 10000 + len(combined) * 1000\n            dn = generate_rna_structure(seq, seed=seed_val)\n            combined.append(adaptive_rna_constraints(dn, tid, segments_map, confidence=0.2))\n            combined_src.append(\"de-novo\")\n\n        # add Protenix\n        ptx = protenix_preds.get(tid)\n        if ptx is not None and ptx.ndim == 3:\n            for j in range(min(ptx.shape[0], PROTENIX)):\n                combined.append(ptx[j])\n                combined_src.append(\"Protenix\")\n\n        # fill missing Protenix slots with unused TBMs\n        while len(combined) < TBM + PROTENIX:\n            needed_idx = len(combined)\n            if len(tbm_all) > needed_idx:\n                combined.append(tbm_all[needed_idx])\n                combined_src.append(\"TBM\")\n            else:\n                seed_val = hash(tid) % 10000 + len(combined) * 1000\n                dn = generate_rna_structure(seq, seed=seed_val)\n                combined.append(adaptive_rna_constraints(dn, tid, segments_map, confidence=0.2))\n                combined_src.append(\"de-novo\")\n\n        # add RNAPro\n        rnapro = rnapro_preds.get(tid)\n        if rnapro is not None and rnapro.ndim == 3:\n            for j in range(min(rnapro.shape[0], RNAPro)):\n                combined.append(rnapro[j])\n                combined_src.append(\"RNAPro\")\n\n        # fill missing RNAPro slots with unused TBMs\n        while len(combined) < N_SAMPLE:\n            needed_idx = len(combined)\n            if len(tbm_all) > needed_idx:\n                combined.append(tbm_all[needed_idx])\n                combined_src.append(\"TBM\")\n            else:\n                seed_val = hash(tid) % 10000 + len(combined) * 1000\n                dn = generate_rna_structure(seq, seed=seed_val)\n                combined.append(adaptive_rna_constraints(dn, tid, segments_map, confidence=0.2))\n                combined_src.append(\"de-novo\")\n\n        # save the source list for this target\n        SOURCES_BY_TARGET[tid] = combined_src[:N_SAMPLE]\n\n        # stack & emit rows as before…\n        stacked = np.stack(combined[:N_SAMPLE], axis=0)\n        all_rows.extend(coords_to_rows(tid, seq, stacked))\n\n    # Write 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    ]\n    coord_cols = [c for c in cols if c.startswith((\"x_\", \"y_\", \"z_\"))]\n\n    sub[coord_cols] = sub[coord_cols].clip(-999.999, 9999.999)\n    sub[cols].to_csv(output_csv, index=False)\n\n    print(f\"\\n✓ Saved submission to {output_csv}  ({len(sub):,} rows)\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"f83c2c14","cell_type":"code","source":"# ─────────────────────────────────────────────────────────────────────────────\n# FIX: Setup writable directory for Templates\n# ─────────────────────────────────────────────────────────────────────────────\n# The competition PDB_RNA folder contains .cif files we can use as templates.\n# We create a writable 'root' directory and symlink the read-only data into it.\n\n# 1. New writable root for Protenix data\nNEW_ROOT = \"/kaggle/working/protenix_root\"\nos.makedirs(NEW_ROOT, exist_ok=True)\nos.environ[\"PROTENIX_ROOT_DIR\"] = NEW_ROOT\n\n# 2. Symlink the competition's PDB_RNA to 'mmcif' (where Protenix looks for templates)\n#    Source: /kaggle/input/stanford-rna-3d-folding-2/PDB_RNA ({pdb_id}.cif)\n#    Target: /kaggle/working/protenix_root/mmcif\nsrc_pdb = \"/kaggle/input/stanford-rna-3d-folding-2/PDB_RNA\"\ndst_pdb = f\"{NEW_ROOT}/mmcif\"\n\nif not os.path.exists(dst_pdb):\n    try:\n        os.symlink(src_pdb, dst_pdb)\n        print(f\"Created symlink: {dst_pdb} -> {src_pdb}\")\n    except OSError:\n        # Fallback if symlink fails (copying is slow but safe)\n        import shutil\n        print(\"Symlink failed, copying PDB_RNA folder (this may take a minute)...\")\n        shutil.copytree(src_pdb, dst_pdb)\n\n# 3. Symlink 'common' and 'checkpoint' from the original dataset\n#    The original dataset path (adjust if your dataset path is different):\nORIGINAL_DATA_DIR = \"/kaggle/input/datasets/qiweiyin/protenix-v1-adjusted/Protenix-v1-adjust-v2/Protenix-v1-adjust-v2/Protenix-v1\"\n\nfor folder in [\"common\", \"checkpoint\"]:\n    src = f\"{ORIGINAL_DATA_DIR}/{folder}\"\n    dst = f\"{NEW_ROOT}/{folder}\"\n    if not os.path.exists(dst):\n        if os.path.exists(src):\n            os.symlink(src, dst)\n            print(f\"Created symlink: {dst} -> {src}\")\n        else:\n            print(f\"WARNING: Could not find original {folder} at {src}\")\n\n# 4. Update the path resolution function to use our new root\ndef resolve_paths():\n    test_csv   = os.environ.get(\"TEST_CSV\",           DEFAULT_TEST_CSV)\n    output_csv = os.environ.get(\"SUBMISSION_CSV\",     DEFAULT_OUTPUT)\n    code_dir   = os.environ.get(\"PROTENIX_CODE_DIR\",  DEFAULT_CODE_DIR)\n    # FORCE the root dir to be our new writable dir\n    root_dir   = NEW_ROOT \n    return test_csv, output_csv, code_dir, root_dir # Returns new root\n","metadata":{},"outputs":[],"execution_count":null},{"id":"6aac8848","cell_type":"code","source":"# /kaggle/input/datasets/qiweiyin/protenix-v1-adjusted/Protenix-v1-adjust-v2/Protenix-v1-adjust-v2/Protenix-v1/checkpoint","metadata":{},"outputs":[],"execution_count":null},{"id":"169b8cfa","cell_type":"code","source":"# ─── FIX for USE_TEMPLATE ────────────────────────────────────────────────────\n# Protenix expects a specific folder structure for templates.\n# We create a writable root in /kaggle/working and symlink the necessary data.\n\n# 1. Define writable root\nWRITABLE_ROOT = \"/kaggle/working/protenix_data\"\nos.makedirs(WRITABLE_ROOT, exist_ok=True)\nos.environ[\"PROTENIX_ROOT_DIR\"] = WRITABLE_ROOT\n\n# 2. Link the competition's PDB_RNA folder to 'mmcif' (where Protenix looks)\n#    Note: Protenix expects $ROOT/mmcif to contain the .cif files\nmmcif_target = f\"{WRITABLE_ROOT}/mmcif\"\nif not os.path.exists(mmcif_target):\n    # Symlink the directory directly if possible, or create folder and link files\n    # The competition data is at: /kaggle/input/stanford-rna-3d-folding-2/PDB_RNA\n    source_pdb = \"/kaggle/input/stanford-rna-3d-folding-2/PDB_RNA\"\n    \n    # We create a symlink to the folder. \n    # Valid validation: Protenix checks if os.path.exists(template_mmcif_dir)\n    os.symlink(source_pdb, mmcif_target)\n    print(f\"Symlinked {source_pdb} -> {mmcif_target}\")\n\n# 3. Link other static assets from the original read-only code dataset\n#    The original code expected data in DEFAULT_CODE_DIR/../.. or similar.\n#    We need 'common' folder (components.cif etc) and 'checkpoint'.\n\n# Original read-only data assets path\n# Based on your previous configs, the assets seem to be here:\nREADONLY_ASSETS = \"/kaggle/input/datasets/qiweiyin/protenix-v1-adjusted/Protenix-v1-adjust-v2/Protenix-v1-adjust-v2\"\n\nfor folder in [\"common\", \"checkpoint\"]:\n    src = f\"{READONLY_ASSETS}/{folder}\"\n    dst = f\"{WRITABLE_ROOT}/{folder}\"\n    if not os.path.exists(dst) and os.path.exists(src):\n        os.symlink(src, dst)\n        print(f\"Symlinked {src} -> {dst}\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"44e333e4","cell_type":"markdown","source":"# RNAPro Setup\n\nSetup RNAPro model and dependencies. RNAPro is an RNA structure prediction model that uses TBM templates as input.\n","metadata":{}},{"id":"bde29d16","cell_type":"code","source":"# ─────────────── RNAPro Setup ─────────────────────────────────────────────────\n# Copy RNAPro source and model checkpoint\n\nimport shutil\n\nRNAPRO_SRC = \"/kaggle/input/datasets/theoviel/rnapro-src/RNAPro\"\nRNAPRO_CKPT = \"/kaggle/input/datasets/theoviel/rnapro-src/rnapro-private-best-500m.ckpt\"\nRNAPRO_CCD_CACHE = \"/kaggle/input/datasets/jaejohn/rnapro-ccd-cache\"\n\n# Copy RNAPro to working directory\nif not os.path.exists(\"/kaggle/working/RNAPro\"):\n    shutil.copytree(RNAPRO_SRC, \"/kaggle/working/RNAPro\")\n    print(\"Copied RNAPro source to /kaggle/working/RNAPro\")\n\n# Copy checkpoint\nif not os.path.exists(\"/kaggle/working/rnapro-private-best-500m.ckpt\"):\n    shutil.copy(RNAPRO_CKPT, \"/kaggle/working/rnapro-private-best-500m.ckpt\")\n    print(\"Copied RNAPro checkpoint\")\n\nprint(\"RNAPro setup complete\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"65cab521","cell_type":"code","source":"# Setup CCD cache for RNAPro (no pip install - RNAPro runs as subprocess only)\nRNAPRO_CCD_DEST = \"/kaggle/working/RNAPro/release_data/ccd_cache\"\nos.makedirs(RNAPRO_CCD_DEST, exist_ok=True)\n\n# Copy CCD cache files\nfor fname in [\"components.cif\", \"components.cif.rdkit_mol.pkl\"]:\n    src = f\"{RNAPRO_CCD_CACHE}/ccd_cache/{fname}\"\n    dst = f\"{RNAPRO_CCD_DEST}/{fname}\"\n    if os.path.exists(src) and not os.path.exists(dst):\n        shutil.copy(src, dst)\n        print(f\"Copied {fname} to RNAPro CCD cache\")\n\nprint(\"RNAPro CCD cache setup complete\")\n\n# ─── Extract RibonanzaNet2 weights from the 500M checkpoint ─────────────────\n# The checkpoint embeds RibonanzaNet2 weights directly (ribonanza_net.* keys).\n# We extract them so RNAPro can use use_RibonanzaNet2=true without needing the\n# separate pytorch_model_fsdp.bin from the Kaggle ribonanzanet2 model.\nRNAPRO_CKPT_PATH = \"/kaggle/working/rnapro-private-best-500m.ckpt\"\nRNET2_DIR = \"/kaggle/working/ribonanzanet2_weights\"\nos.makedirs(RNET2_DIR, exist_ok=True)\n\nfsdp_bin = f\"{RNET2_DIR}/pytorch_model_fsdp.bin\"\nyaml_out = f\"{RNET2_DIR}/pairwise.yaml\"\n\nif not os.path.exists(fsdp_bin):\n    print(\"Extracting ribonanza_net weights from 500M checkpoint...\")\n    import torch\n    ckpt = torch.load(RNAPRO_CKPT_PATH, map_location=\"cpu\")\n    model_state = ckpt.get(\"model\", ckpt)\n    if any(k.startswith(\"module.\") for k in model_state.keys()):\n        model_state = {k[len(\"module.\"):]: v for k, v in model_state.items()}\n    rnet_state = {k[len(\"ribonanza_net.\"):]: v for k, v in model_state.items()\n                  if k.startswith(\"ribonanza_net.\")}\n    torch.save(rnet_state, fsdp_bin)\n    del ckpt, model_state, rnet_state\n    print(f\"Saved ribonanza_net weights to {fsdp_bin}\")\nelse:\n    print(f\"RibonanzaNet2 weights already exist at {fsdp_bin}\")\n\nif not os.path.exists(yaml_out):\n    # pairwise.yaml: ninp=384, nlayers=48, dim_msa=32, pairwise_dimension=128\n    pairwise_yaml = \"\"\"# Inferred from rnapro-private-best-500m.ckpt\nlearning_rate: 0.002\nbatch_size: 3\ntest_batch_size: 8\nepochs: 1\ndropout: 0.1\nweight_decay: 0.0001\nk: 5\nninp: 384\nnlayers: 48\nnclass: 10\nntoken: 6\nnhead: 12\nuse_flip_aug: false\ngradient_accumulation_steps: 1\nuse_triangular_attention: false\npairwise_dimension: 128\ndim_msa: 32\nclip_grad_norm: 1\nmax_len: 1000\nlog_interval: 2000\nuse_noise_aug: false\nuse_data_percentage: 1\nuse_dirty_data: true\nfold: 0\nnfolds: 6\ninput_dir: \"../../input/\"\ngpu_id: \"0\"\n\"\"\"\n    with open(yaml_out, \"w\") as f:\n        f.write(pairwise_yaml)\n    print(f\"Written pairwise.yaml to {yaml_out}\")\nelse:\n    print(f\"pairwise.yaml already exists at {yaml_out}\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"a152ada6","cell_type":"code","source":"# ─────────────── RNAPro Inference Runner ─────────────────────────────────────\n# This creates a custom inference runner script for RNAPro\n\nrnapro_inference_code = '''\nimport os\nimport shutil\nimport logging\nimport traceback\nimport warnings\nimport argparse\nfrom contextlib import nullcontext\nfrom os.path import join as opjoin\nfrom typing import Any, Mapping\n\nimport json\nimport torch\nimport pandas as pd\nimport numpy as np\nfrom biotite.structure.io import pdbx\n\nfrom configs.configs_base import configs as configs_base\nfrom configs.configs_data import data_configs\nfrom configs.configs_inference import inference_configs\nfrom runner.dumper import DataDumper\n\nfrom rnapro.config import parse_sys_args\nfrom rnapro.config.config import ConfigManager, ArgumentNotSet\nfrom rnapro.data.infer_data_pipeline import get_inference_dataloader\nfrom rnapro.model.RNAPro import RNAPro\nfrom rnapro.utils.distributed import DIST_WRAPPER\nfrom rnapro.utils.seed import seed_everything\nfrom rnapro.utils.torch_utils import to_device\n\nwarnings.filterwarnings(\"ignore\", category=UserWarning)\nwarnings.filterwarnings(\"ignore\", category=FutureWarning)\nwarnings.filterwarnings(\"ignore\", category=DeprecationWarning)\n\nlogger = logging.getLogger(__name__)\nlogging.basicConfig(level=logging.WARNING)\nlogging.getLogger(\"rnapro.data\").setLevel(logging.WARNING)\nlogging.getLogger(\"rnapro\").setLevel(logging.WARNING)\n\n\ndef parse_configs(configs, arg_str=None, fill_required_with_null=False):\n    manager = ConfigManager(configs, fill_required_with_null=fill_required_with_null)\n    parser = argparse.ArgumentParser()\n    parser.add_argument(\"--max_len\", type=int, default=10000, required=False)\n    \n    for key, (dtype, default_value, allow_none, required) in manager.config_infos.items():\n        parser.add_argument(\"--\" + key, type=str, default=ArgumentNotSet(), required=required)\n    \n    merged_configs = manager.merge_configs(\n        vars(parser.parse_args(arg_str.split())) if arg_str else {}\n    )\n    max_len = parser.parse_args(arg_str.split()).max_len\n    merged_configs.max_len = max_len\n    return merged_configs\n\n\nclass dotdict(dict):\n    __setattr__ = dict.__setitem__\n    __delattr__ = dict.__delitem__\n    def __getattr__(self, name):\n        try:\n            return self[name]\n        except KeyError:\n            raise AttributeError(name)\n\n\nclass InferenceRunner(object):\n    def __init__(self, configs):\n        self.configs = configs\n        self.init_env()\n        self.init_basics()\n        self.init_model()\n        self.load_checkpoint()\n        self.init_dumper(\n            need_atom_confidence=configs.need_atom_confidence,\n            sorted_by_ranking_score=configs.sorted_by_ranking_score,\n        )\n\n    def init_env(self):\n        self.use_cuda = torch.cuda.device_count() > 0\n        if self.use_cuda:\n            self.device = torch.device(\"cuda:{}\".format(DIST_WRAPPER.local_rank))\n            torch.cuda.set_device(self.device)\n        else:\n            self.device = torch.device(\"cpu\")\n\n    def init_basics(self):\n        self.dump_dir = self.configs.dump_dir\n        self.error_dir = opjoin(self.dump_dir, \"ERR\")\n        os.makedirs(self.dump_dir, exist_ok=True)\n        os.makedirs(self.error_dir, exist_ok=True)\n\n    def init_model(self):\n        self.model = RNAPro(self.configs).to(self.device)\n        num_params = sum(p.numel() for p in self.model.parameters())\n        print(f\"RNAPro model loaded: {num_params:,} parameters\")\n\n    def load_checkpoint(self):\n        checkpoint_path = self.configs.load_checkpoint_path\n        if not os.path.exists(checkpoint_path):\n            raise Exception(f\"Checkpoint not found: {checkpoint_path}\")\n        checkpoint = torch.load(checkpoint_path, self.device)\n        sample_key = list(checkpoint[\"model\"].keys())[0]\n        if sample_key.startswith(\"module.\"):\n            checkpoint[\"model\"] = {k[len(\"module.\"):]: v for k, v in checkpoint[\"model\"].items()}\n        self.model.load_state_dict(checkpoint[\"model\"], strict=True)\n        self.model.eval()\n\n    def init_dumper(self, need_atom_confidence=False, sorted_by_ranking_score=True):\n        self.dumper = DataDumper(\n            base_dir=self.dump_dir,\n            need_atom_confidence=need_atom_confidence,\n            sorted_by_ranking_score=sorted_by_ranking_score,\n        )\n\n    @torch.no_grad()\n    def predict(self, data):\n        eval_precision = {\"fp32\": torch.float32, \"bf16\": torch.bfloat16, \"fp16\": torch.float16}[self.configs.dtype]\n        enable_amp = torch.autocast(device_type=\"cuda\", dtype=eval_precision) if torch.cuda.is_available() else nullcontext()\n        data = to_device(data, self.device)\n        with enable_amp:\n            prediction, _, _ = self.model(\n                input_feature_dict=data[\"input_feature_dict\"],\n                label_full_dict=None,\n                label_dict=None,\n                mode=\"inference\",\n            )\n        return prediction\n\n    def update_model_configs(self, new_configs):\n        self.model.configs = new_configs\n\n\ndef update_inference_configs(configs, N_token):\n    if N_token > 3840:\n        configs.skip_amp.confidence_head = False\n        configs.skip_amp.sample_diffusion = False\n    elif N_token > 2560:\n        configs.skip_amp.confidence_head = False\n        configs.skip_amp.sample_diffusion = True\n    else:\n        configs.skip_amp.confidence_head = True\n        configs.skip_amp.sample_diffusion = True\n    return configs\n\n\ndef infer_predict(runner, configs):\n    try:\n        dataloader = get_inference_dataloader(configs=configs)\n    except Exception as e:\n        error_message = f\"{e}:\\\\n{traceback.format_exc()}\"\n        logger.info(error_message)\n        with open(opjoin(runner.error_dir, \"error.txt\"), \"a\") as f:\n            f.write(error_message)\n        return\n\n    num_data = len(dataloader.dataset)\n    for seed in configs.seeds:\n        seed_everything(seed=seed, deterministic=configs.deterministic)\n        for batch in dataloader:\n            try:\n                data, atom_array, data_error_message = batch[0]\n                sample_name = data[\"sample_name\"]\n                if len(data_error_message) > 0:\n                    continue\n\n                new_configs = update_inference_configs(configs, data[\"N_token\"].item())\n                runner.update_model_configs(new_configs)\n                prediction = runner.predict(data)\n                runner.dumper.dump(\n                    dataset_name=\"\",\n                    pdb_id=sample_name,\n                    seed=seed,\n                    pred_dict=prediction,\n                    atom_array=atom_array,\n                    entity_poly_type=data[\"entity_poly_type\"],\n                )\n                torch.cuda.empty_cache()\n            except Exception as e:\n                error_message = f\"[Rank {DIST_WRAPPER.rank}]{data['sample_name']} {e}:\\\\n{traceback.format_exc()}\"\n                if hasattr(torch.cuda, \"empty_cache\"):\n                    torch.cuda.empty_cache()\n\n\ndef make_dummy_solution(valid_df):\n    solution = dotdict()\n    for i, row in valid_df.iterrows():\n        target_id = row.target_id\n        sequence = row.sequence\n        solution[target_id] = dotdict(target_id=target_id, sequence=sequence, coord=[])\n    return solution\n\n\ndef solution_to_submit_df(solution):\n    submit_df = []\n    for k, s in solution.items():\n        df = coord_to_df(s.sequence, s.coord, s.target_id)\n        submit_df.append(df)\n    submit_df = pd.concat(submit_df)\n    return submit_df\n\n\ndef coord_to_df(sequence, coord, target_id):\n    L = len(sequence)\n    df = pd.DataFrame()\n    df[\"ID\"] = [f\"{target_id}_{i+1}\" for i in range(L)]\n    df[\"resname\"] = [s for s in sequence]\n    df[\"resid\"] = [i+1 for i in range(L)]\n    num_coord = len(coord)\n    for j in range(num_coord):\n        df[f\"x_{j+1}\"] = coord[j][:, 0]\n        df[f\"y_{j+1}\"] = coord[j][:, 1]\n        df[f\"z_{j+1}\"] = coord[j][:, 2]\n    return df\n\n\ndef create_input_json(sequence, target_id):\n    return [{\n        \"sequences\": [{\"rnaSequence\": {\"sequence\": sequence, \"count\": 1}}],\n        \"name\": target_id,\n    }]\n\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        atom_names_clean = np.char.strip(atom_array.atom_name.astype(str))\n        mask_c1 = atom_names_clean == \"C1\\\\'\"\n        c1_atoms = atom_array[mask_c1]\n        if len(c1_atoms) == 0:\n            return None\n        sort_indices = np.argsort(c1_atoms.res_id)\n        c1_atoms_sorted = c1_atoms[sort_indices]\n        return c1_atoms_sorted.coord\n    except Exception as e:\n        print(f\"Error extracting C1\\\\' coordinates: {e}\")\n        return None\n\n\ndef process_sequence(sequence, target_id, temp_dir):\n    input_json = create_input_json(sequence, target_id)\n    os.makedirs(temp_dir, exist_ok=True)\n    input_json_path = os.path.join(temp_dir, f\"{target_id}_input.json\")\n    with open(input_json_path, \"w\") as f:\n        json.dump(input_json, f, indent=4)\n\n\ndef run_ptx(target_id, sequence, configs, solution, template_idx, runner):\n    temp_dir = f\"./{configs.dump_dir}/input\"\n    output_dir = f\"./{configs.dump_dir}/output\"\n    os.makedirs(temp_dir, exist_ok=True)\n    os.makedirs(output_dir, exist_ok=True)\n\n    process_sequence(sequence=sequence, target_id=target_id, temp_dir=temp_dir)\n    configs.input_json_path = os.path.join(temp_dir, f\"{target_id}_input.json\")\n    configs.template_idx = int(template_idx)\n\n    infer_predict(runner, configs)\n\n    cif_file_path = f\"{configs.dump_dir}/{target_id}/seed_42/predictions/{target_id}_sample_0.cif\"\n    coord = extract_c1_coordinates(cif_file_path)\n    if coord is None:\n        coord = np.zeros((len(sequence), 3), dtype=np.float32)\n    elif coord.shape[0] < len(sequence):\n        pad_len = len(sequence) - coord.shape[0]\n        pad = np.zeros((pad_len, 3), dtype=np.float32)\n        coord = np.concatenate([coord, pad], axis=0)\n    solution[target_id].coord.append(coord)\n\n\ndef run():\n    LOG_FORMAT = \"%(asctime)s,%(msecs)-3d %(levelname)-8s [%(filename)s:%(lineno)s %(funcName)s] %(message)s\"\n    logging.basicConfig(format=LOG_FORMAT, level=logging.WARNING, datefmt=\"%Y-%m-%d %H:%M:%S\", filemode=\"w\")\n    logging.getLogger(\"rnapro.data\").setLevel(logging.WARNING)\n    logging.getLogger(\"rnapro\").setLevel(logging.WARNING)\n    \n    configs_base[\"use_deepspeed_evo_attention\"] = os.environ.get(\"USE_DEEPSPEED_EVO_ATTENTION\", False) == \"true\"\n    configs = {**configs_base, **{\"data\": data_configs}, **inference_configs}\n    configs = parse_configs(configs=configs, arg_str=parse_sys_args(), fill_required_with_null=True)\n\n    valid_df = pd.read_csv(configs.sequences_csv)\n    print(f\"\\\\n -> Loaded {len(valid_df)} sequence(s)\")\n\n    print(\"\\\\n -> Building model and loading checkpoint\")\n    runner = InferenceRunner(configs)\n    print(\"\\\\n -> Done, starting inference...\")\n\n    solution = make_dummy_solution(valid_df)\n    for idx, row in valid_df.iterrows():\n        print(f\"\\\\n -> Sequence {row.target_id}: {row.sequence[:50]}...\")\n\n        if len(row.sequence) > configs.max_len:\n            print(f\"Sequence too long ({len(row.sequence)} > {configs.max_len}), skipping\")\n            for template_idx in range(5):\n                coord = np.zeros((len(row.sequence), 3), dtype=np.float32)\n                solution[row.target_id].coord.append(coord)\n            continue\n\n        try:\n            for template_idx in range(5):\n                run_ptx(\n                    target_id=row.target_id,\n                    sequence=row.sequence,\n                    configs=configs,\n                    solution=solution,\n                    template_idx=template_idx,\n                    runner=runner,\n                )\n        except Exception as e:\n            print(f\"Error processing {row.target_id}: {e}\")\n            continue\n\n    print(\"\\\\n\\\\n -> Inference done! Saving to submission.csv\")\n    submit_df = solution_to_submit_df(solution)\n    submit_df = submit_df.fillna(0.0)\n    submit_df.to_csv(\"./submission.csv\", index=False)\n\n\nif __name__ == \"__main__\":\n    run()\n'''\n\n# Write the inference script to RNAPro directory\nwith open(\"/kaggle/working/RNAPro/runner/inference.py\", \"w\") as f:\n    f.write(rnapro_inference_code)\nprint(\"Written RNAPro inference.py script\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"8a993d01","cell_type":"code","source":"# Create RNAPro inference shell script\n# RibonanzaNet2 weights are extracted from the 500M checkpoint into /kaggle/working/ribonanzanet2_weights\nRNET2_DIR = \"/kaggle/working/ribonanzanet2_weights\"\n\nrnapro_sh_script = f'''#!/bin/bash\nexport PYTHONPATH=\"/kaggle/working/RNAPro:$PYTHONPATH\"\nexport LAYERNORM_TYPE=torch\n\n# Inference parameters\nSEED=42\nN_SAMPLE=1\nN_STEP=200\nN_CYCLE=10\n\n# Paths\nDUMP_DIR=\"../output\"\nCHECKPOINT_PATH=\"../rnapro-private-best-500m.ckpt\"\n\n# Template/MSA settings\nTEMPLATE_DATA=\"./release_data/kaggle/templates.pt\"\nTEMPLATE_IDX=0\nRNA_MSA_DIR=\"/kaggle/input/stanford-rna-3d-folding-2/MSA\"\n\nSEQUENCES_CSV=\"/kaggle/working/rnapro_sequences.csv\"\nMODEL_NAME=\"rnapro_base\"\n\n# RibonanzaNet2 weights extracted from the main checkpoint\nRIBONANZA_PATH=\"{RNET2_DIR}\"\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 true \\\\\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 ${{TEMPLATE_IDX}} \\\\\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 1000\n'''\n\nwith open(\"/kaggle/working/RNAPro/rnapro_inference_kaggle.sh\", \"w\") as f:\n    f.write(rnapro_sh_script)\nprint(\"Written RNAPro inference shell script\")\nprint(f\"  RIBONANZA_PATH = {RNET2_DIR}\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"b34b950e","cell_type":"markdown","source":"# Main","metadata":{}},{"id":"2cd78edb","cell_type":"code","source":"TARGET = None\n# TARGET = [\"9MME\"]\nif __name__ == \"__main__\":\n    main()\n","metadata":{},"outputs":[],"execution_count":null},{"id":"228cc4f4","cell_type":"code","source":"#read submission.csv\nsubmission_path = \"/kaggle/working/submission.csv\"\nsubmission_df = pd.read_csv(submission_path)\nprint(submission_df.head(20))","metadata":{},"outputs":[],"execution_count":null},{"id":"658f6dca","cell_type":"markdown","source":"# Evaluation","metadata":{}},{"id":"6472352e","cell_type":"code","source":"!rm -f /kaggle/working/metric.py","metadata":{},"outputs":[],"execution_count":null},{"id":"6a22ebdc","cell_type":"code","source":"%%writefile /kaggle/working/metric.py\n#!/usr/bin/env python\n# coding: utf-8\n\n# In[ ]:\n\nimport numpy as np\nimport os\nimport re\nimport math\nimport pandas as pd\nfrom pathlib import Path\nimport shutil\nimport sys\nimport csv\n\n# ---------------------\n# Helper: parse USalign output\n# ---------------------\ndef parse_tmscore_output(output: str) -> float:\n    matches = re.findall(r'TM-score=\\s+([\\d.]+)', output)\n    if len(matches) < 2:\n        raise ValueError('No TM score found in USalign output')\n    return float(matches[1])\n\n# ---------------------\n# PDB writers\n# ---------------------\n\ndef sanitize(xyz):\n    MIN_COORD=-999.999\n    MAX_COORD=9999.999\n    return min(max(xyz,MIN_COORD),MAX_COORD)\n\ndef write_target_line(atom_name, atom_serial, residue_name, chain_id, residue_num,\n                      x_coord, y_coord, z_coord, occupancy=1.0, b_factor=0.0, atom_type='P') -> str:\n    return f'ATOM  {atom_serial:>5d}  {atom_name:4s}{residue_name:>3s} {chain_id:1s}{residue_num:>4d}    {sanitize(x_coord):>8.3f}{sanitize(y_coord):>8.3f}{sanitize(z_coord):>8.3f}{occupancy:>6.2f}{b_factor:>6.2f}           {atom_type}\\n'\n    \ndef write2pdb(df: pd.DataFrame, xyz_id: int, target_path: str) -> int:\n    \"\"\"\n    Write single-chain PDB (chain 'A') using row['resid'] as residue_num.\n    Raises exceptions on invalid data.\n    \"\"\"\n    resolved_cnt = 0\n    with open(target_path, 'w') as fh:\n        for _, row in df.iterrows():\n            x = row[f'x_{xyz_id}']\n            y = row[f'y_{xyz_id}']\n            z = row[f'z_{xyz_id}']\n            if x > -1e6 and y > -1e6 and z > -1e6:\n                resolved_cnt += 1\n                resid_num = int(row['resid'])\n                fh.write(write_target_line(\"C1'\", resid_num, row['resname'], 'A', resid_num, x, y, z, atom_type='C'))\n    return resolved_cnt\n\ndef write2pdb_singlechain_native(df_native: pd.DataFrame, xyz_id: int, target_path: str) -> int:\n    \"\"\"\n    Write native single-chain using row['resid'] as residue numbers.\n    Assumes all required columns exist and are valid.\n    \"\"\"\n    df_sorted = df_native.copy()\n    df_sorted['__resid_int'] = df_sorted['resid'].astype(int)\n    df_sorted = df_sorted.sort_values('__resid_int').reset_index(drop=True)\n\n    resolved_cnt = 0\n    with open(target_path, 'w') as fh:\n        for _, row in df_sorted.iterrows():\n            x = row[f'x_{xyz_id}']\n            y = row[f'y_{xyz_id}']\n            z = row[f'z_{xyz_id}']\n            if x > -1e6 and y > -1e6 and z > -1e6:\n                resolved_cnt += 1\n                resid_num = int(row['resid'])\n                fh.write(write_target_line(\"C1'\", resid_num, row['resname'], 'A', resid_num, x, y, z, atom_type='C'))\n    return resolved_cnt\n\ndef write2pdb_multichain_from_solution(df_solution: pd.DataFrame, xyz_id: int, target_path: str) -> int:\n    \"\"\"\n    Write multi-chain PDB for native solution using columns 'chain' and 'copy' to assign chain letters.\n    Expects 'resid' convertible to int and chain/copy present. No fallbacks.\n    \"\"\"\n    df_sorted = df_solution.copy()\n    df_sorted['__resid_int'] = df_sorted['resid'].astype(int)\n    df_sorted = df_sorted.sort_values('__resid_int')\n\n    chain_map = {}\n    next_ord = ord('A')\n    written = 0\n    with open(target_path, 'w') as fh:\n        for _, row in df_sorted.iterrows():\n            x = row[f'x_{xyz_id}']\n            y = row[f'y_{xyz_id}']\n            z = row[f'z_{xyz_id}']\n            if not (x > -1e6 and y > -1e6 and z > -1e6):\n                continue\n            chain_val = row['chain']\n            copy_key = int(row['copy'])\n            g = (str(chain_val), copy_key)\n            if g not in chain_map:\n                if next_ord <= ord('Z'):\n                    ch = chr(next_ord)\n                else:\n                    ov = next_ord - ord('Z') - 1\n                    if ov < 26:\n                        ch = chr(ord('a') + ov)\n                    else:\n                        ch = chr(ord('0') + (ov - 26) % 10)\n                chain_map[g] = ch\n                next_ord += 1\n            chain_id = chain_map[g]\n            written += 1\n            resid_num = int(row['resid'])\n            fh.write(write_target_line(\"C1'\", resid_num, row['resname'], chain_id, resid_num, x, y, z, atom_type='C'))\n    return written\n\ndef write2pdb_multichain_from_groups(df_pred: pd.DataFrame, xyz_id: int, target_path: str, groups_list) -> (int, list):\n    \"\"\"\n    Write predicted multichain PDB based on a positional groups_list (tuple per residue: (chain, copy)).\n    Requires groups_list length == number of residues in df_pred (after sorting).\n    Returns (written_count, chain_letters_per_res).\n    \"\"\"\n    df_sorted = df_pred.copy()\n    df_sorted['__resid_int'] = df_sorted['resid'].astype(int)\n    df_sorted = df_sorted.sort_values('__resid_int').reset_index(drop=True)\n\n    if groups_list is None or len(groups_list) != len(df_sorted):\n        raise ValueError(\"groups_list must be provided and match number of residues in predicted df\")\n\n    chain_map = {}\n    next_ord = ord('A')\n    chain_letters = []\n    written = 0\n    with open(target_path, 'w') as fh:\n        for idx, row in df_sorted.iterrows():\n            g = groups_list[idx]\n            if isinstance(g, tuple):\n                gkey = (str(g[0]), int(g[1]))\n            else:\n                gkey = (str(g), None)\n            if gkey not in chain_map:\n                if next_ord <= ord('Z'):\n                    ch = chr(next_ord)\n                else:\n                    ov = next_ord - ord('Z') - 1\n                    if ov < 26:\n                        ch = chr(ord('a') + ov)\n                    else:\n                        ch = chr(ord('0') + (ov - 26) % 10)\n                chain_map[gkey] = ch\n                next_ord += 1\n            chain_id = chain_map[gkey]\n            chain_letters.append(chain_id)\n            x = row[f'x_{xyz_id}']\n            y = row[f'y_{xyz_id}']\n            z = row[f'z_{xyz_id}']\n            if x > -1e6 and y > -1e6 and z > -1e6:\n                written += 1\n                resid_num = int(row['resid'])\n                fh.write(write_target_line(\"C1'\", resid_num, row['resname'], chain_id, resid_num, x, y, z, atom_type='C'))\n    return written, chain_letters\n\ndef write2pdb_singlechain_permuted_pred(df_pred: pd.DataFrame, xyz_id: int, permuted_indices: list, target_path: str) -> int:\n    \"\"\"\n    Create single-chain PDB by concatenating predicted residues in permuted_indices order.\n    Output residue numbers are sequential starting at 1 and increase for every permuted position.\n    Raises exception if indices out of range.\n    \"\"\"\n    df_sorted = df_pred.copy()\n    df_sorted['__resid_int'] = df_sorted['resid'].astype(int)\n    df_sorted = df_sorted.sort_values('__resid_int').reset_index(drop=True)\n\n    written = 0\n    next_res = 1\n    with open(target_path, 'w') as fh:\n        for idx in permuted_indices:\n            if idx < 0 or idx >= len(df_sorted):\n                # strict behavior: raise error for invalid index\n                raise IndexError(f\"permuted index {idx} out of range for predicted residues\")\n            row = df_sorted.iloc[idx]\n            x = row[f'x_{xyz_id}']\n            y = row[f'y_{xyz_id}']\n            z = row[f'z_{xyz_id}']\n            out_resnum = next_res\n            if x > -1e6 and y > -1e6 and z > -1e6:\n                written += 1\n                fh.write(write_target_line(\"C1'\", out_resnum, row['resname'], 'A', out_resnum, x, y, z, atom_type='C'))\n            next_res += 1\n    return written\n\n# ---------------------\n# USalign wrappers\n# ---------------------\ndef run_usalign_raw(predicted_pdb: str, native_pdb: str, usalign_bin='USalign', align_sequence=False, tmscore=None) -> str:\n    cmd = f'{usalign_bin} {predicted_pdb} {native_pdb} -atom \" C1\\'\"'\n    if tmscore is not None:\n        cmd += f' -TMscore {tmscore}'\n        if int(tmscore) == 0:\n            cmd += ' -mm 1 -ter 0'\n    elif not align_sequence:\n        cmd += ' -TMscore 1'\n    return os.popen(cmd).read()\n\ndef parse_usalign_chain_orders(output: str):\n    \"\"\"\n    Parse USalign output for both Structure_1 and Structure_2 chain lists.\n    Returns (chain_list_structure1, chain_list_structure2).\n    Raises if parsing fails to find either line.\n    \"\"\"\n    chain1 = None\n    chain2 = None\n    for line in output.splitlines():\n        line = line.strip()\n        if line.startswith('Name of Structure_1:'):\n            parts = line.split(':')\n            clist = []\n            for part in parts[2:]:\n                token = part.strip()\n                if token == '':\n                    continue\n                token0 = token.split()[0]\n                last = token0.split(',')[-1]\n                ch = re.sub(r'[^A-Za-z0-9]', '', last)\n                if ch:\n                    clist.append(ch)\n            chain1 = clist\n        elif line.startswith('Name of Structure_2:'):\n            parts = line.split(':')\n            clist = []\n            for part in parts[2:]:\n                token = part.strip()\n                if token == '':\n                    continue\n                token0 = token.split()[0]\n                last = token0.split(',')[-1]\n                ch = re.sub(r'[^A-Za-z0-9]', '', last)\n                if ch:\n                    clist.append(ch)\n            chain2 = clist\n    if chain1 is None or chain2 is None:\n        raise ValueError(\"Failed to parse chain orders from USalign output\")\n    return chain1, chain2\n\n# ---------------------\n# Main scoring function (no try/except, no fallbacks)\n# ---------------------\n# ---------------------\n# Main scoring function (no try/except, no fallbacks)\n# ---------------------\ndef score(solution: pd.DataFrame, submission: pd.DataFrame, row_id_column_name: str, usalign_bin_hint: str = None) -> float:\n    \"\"\"\n    Enhanced scoring with chain-permutation handling for multicopy targets.\n    This version contains no try/except blocks and will raise on any error.\n    \n    Returns:\n    mean_score: float\n    details_by_target: dict { target_id: (best_idx, [score_slot_1, score_slot_2, ...]) }\n    \"\"\"\n    # determine usalign binary\n    if usalign_bin_hint:\n        usalign_bin = usalign_bin_hint\n    else:\n        if os.path.exists('/kaggle/input/datasets/metric/usalign/USalign') and not os.path.exists('/kaggle/working/USalign'):\n            shutil.copy2('/kaggle/input/datasets/metric/usalign/USalign', '/kaggle/working/USalign')\n            os.chmod('/kaggle/working/USalign', 0o755)\n        usalign_bin = '/kaggle/working/USalign' if os.path.exists('/kaggle/working/USalign') else 'USalign'\n\n    sol = solution.copy()\n    sub = submission.copy()\n    sol['target_id'] = sol['ID'].apply(lambda x: '_'.join(str(x).split('_')[:-1]))\n    sub['target_id'] = sub['ID'].apply(lambda x: '_'.join(str(x).split('_')[:-1]))\n\n    results = []\n    details_by_target = {}\n    \n    for target_id, group_native in sol.groupby('target_id'):\n        group_predicted = sub[sub['target_id'] == target_id]\n        if group_predicted.empty:\n            zeros = [0.0] * 5\n            results.append(0.0)\n            details_by_target[target_id] = (0, zeros)\n            continue\n        has_chain_copy = ('chain' in group_native.columns) and ('copy' in group_native.columns)\n        is_multicopy = has_chain_copy and (group_native['copy'].astype(float).max() > 1)\n\n        # precompute native models that have coords\n        native_with_coords = []\n        for native_cnt in range(1, 41):\n            native_pdb = f'native_{target_id}_{native_cnt}.pdb'\n            resolved_native = write2pdb(group_native, native_cnt, native_pdb)\n            if resolved_native > 0:\n                native_with_coords.append(native_cnt)\n            else:\n                if os.path.exists(native_pdb):\n                    os.remove(native_pdb)\n\n        if not native_with_coords:\n            raise ValueError(f\"No native models with coordinates for target {target_id}\")\n\n        best_per_pred = []\n        for pred_cnt in range(1, 6):\n            if not is_multicopy:\n                predicted_pdb = f'predicted_{target_id}_{pred_cnt}.pdb'\n                resolved_pred = write2pdb(group_predicted, pred_cnt, predicted_pdb)\n                if resolved_pred <= 2:\n                    #print(f\"Predicted model {pred_cnt} for target {target_id} has insufficient coordinates\")\n                    best_per_pred.append( 0.0 )\n                    continue\n                \n                scores = []\n                for native_cnt in native_with_coords:\n                    native_pdb = f'native_{target_id}_{native_cnt}.pdb'\n                    out = run_usalign_raw(predicted_pdb, native_pdb, usalign_bin=usalign_bin, align_sequence=False, tmscore=1)\n                    s = parse_tmscore_output(out)\n                    scores.append(s)\n                best_per_pred.append(max(scores))\n\n            else:\n                # multicopy\n                # strict: require chain and copy columns convertible\n                gn_sorted = group_native.copy()\n                gn_sorted['__resid_int'] = gn_sorted['resid'].astype(int)\n                gn_sorted = gn_sorted.sort_values('__resid_int').reset_index(drop=True)\n                groups_list = []\n                for _, r in gn_sorted.iterrows():\n                    chain_val = r['chain']\n                    copy_i = int(r['copy'])\n                    groups_list.append((chain_val, copy_i))\n\n                # predicted multichain - groups_list must match predicted residue count or error\n                dfp_sorted = group_predicted.copy()\n                dfp_sorted['__resid_int'] = dfp_sorted['resid'].astype(int)\n                dfp_sorted = dfp_sorted.sort_values('__resid_int').reset_index(drop=True)\n                if len(groups_list) != len(dfp_sorted):\n                    raise ValueError(f\"groups_list length ({len(groups_list)}) does not match predicted residue count ({len(dfp_sorted)}) for target {target_id}\")\n\n                predicted_multi_pdb = f'pred_multi_{target_id}_{pred_cnt}.pdb'\n                resolved_pred_multi, pred_chain_letters = write2pdb_multichain_from_groups(group_predicted, pred_cnt, predicted_multi_pdb, groups_list)\n                if resolved_pred_multi == 0:\n                    #print(f\"Predicted multi model {pred_cnt} for target {target_id} has no coordinates\")\n                    best_per_pred.append( 0.0 )\n                    continue\n\n                scores = []\n                for native_cnt in native_with_coords:\n                    native_multi_pdb = f'native_multi_{target_id}_{native_cnt}.pdb'\n                    resolved_native_multi = write2pdb_multichain_from_solution(group_native, native_cnt, native_multi_pdb)\n                    if resolved_native_multi == 0:\n                        continue\n\n                    raw_out = run_usalign_raw(predicted_multi_pdb, native_multi_pdb, usalign_bin=usalign_bin, align_sequence=True, tmscore=0)\n                    chain1, chain2 = parse_usalign_chain_orders(raw_out)  # will raise if parsing fails\n\n                    # build native->pred mapping chain2[i] -> chain1[i]\n                    native_to_pred = {n_ch: p_ch for n_ch, p_ch in zip(chain2, chain1)}\n\n                    # canonical native order = chain2 unique in order seen\n                    #native_chain_order = []\n                    #for ch in chain2:\n                    #    if ch not in native_chain_order:\n                    #        native_chain_order.append(ch)\n                    native_chain_order = list(native_to_pred.keys())\n                    native_chain_order.sort() # this is critical...\n\n                    # predicted chain order by following native chain A,B,...\n                    pred_chain_order = [native_to_pred[n_ch] for n_ch in native_chain_order if native_to_pred.get(n_ch) is not None]\n\n                    # construct pred_positions_by_chain\n                    pred_positions_by_chain = {}\n                    for idx, ch in enumerate(pred_chain_letters):\n                        if ch is None:\n                            continue\n                        pred_positions_by_chain.setdefault(ch, []).append(idx)\n\n                    # require that each chain in pred_chain_order exists in pred_positions_by_chain\n                    pred_chain_order = [p for p in pred_chain_order if p in pred_positions_by_chain]\n\n                    # form permuted indices by concatenation\n                    permuted_indices = []\n                    for ch in pred_chain_order:\n                        permuted_indices.extend(pred_positions_by_chain[ch])\n                    # append any remaining\n                    for idx in range(len(pred_chain_letters)):\n                        if idx not in permuted_indices:\n                            permuted_indices.append(idx)\n\n                    # write permuted single-chain predicted and native single-chain\n                    pred_single_perm = f'pred_permuted_{target_id}_{pred_cnt}_{native_cnt}.pdb'\n                    written_pred_single = write2pdb_singlechain_permuted_pred(group_predicted, pred_cnt, permuted_indices, pred_single_perm)\n                    native_single = f'native_single_{target_id}_{native_cnt}.pdb'\n                    written_native = write2pdb_singlechain_native(group_native, native_cnt, native_single)\n\n                    if written_pred_single <= 2 or written_native <= 2:\n                        raise ValueError(f\"Insufficient residues after permutation for target {target_id}, pred {pred_cnt}, native {native_cnt}\")\n\n                    out = run_usalign_raw(pred_single_perm, native_single, usalign_bin=usalign_bin, align_sequence=False, tmscore=1)\n                    score_final = parse_tmscore_output(out)\n                    scores.append(score_final)\n\n                best_per_pred.append(max(scores))\n\n        target_best_score = max(best_per_pred)\n        results.append(target_best_score)\n        \n        # Store index of best score + all scores\n        best_idx = int(np.argmax(best_per_pred))\n        details_by_target[target_id] = (best_idx, best_per_pred)\n        \n    mean_score = float(sum(results) / len(results)) if results else 0.0\n    return mean_score, details_by_target","metadata":{},"outputs":[],"execution_count":null},{"id":"0d431d67","cell_type":"code","source":"import runpy\nmodule_globals = runpy.run_path(\"/kaggle/working/metric.py\")\nscore = module_globals['score']","metadata":{},"outputs":[],"execution_count":null},{"id":"13b87561","cell_type":"code","source":"# ─────────────── Score Calculation with Source Attribution ───────────────────\n# This cell calculates scores and shows which source (TBM/Protenix/RNAPro) \n# contributed each prediction slot.\n\nif not IS_KAGGLE:\n    import pandas as pd\n    sub = pd.read_csv('/kaggle/working/submission.csv')\n    sol = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/validation_labels.csv')\n\n    sub['target_id'] = sub['ID'].apply(lambda x: '_'.join(str(x).split('_')[:-1]))\n    sol['target_id'] = sol['ID'].apply(lambda x: '_'.join(str(x).split('_')[:-1]))\n\n    sub_targets = sub['target_id'].unique()\n\n    # only score the targets we actually produced in local mode\n    sol = sol[sol['target_id'].isin(sub_targets)]\n\n    # run the scoring once\n    mean_score, details = score(sol, sub, \"ID\")\n\n    print(f\"\\n{'='*80}\")\n    print(f\"Score Results: TBM={TBM}, Protenix={PROTENIX}, RNAPro={RNAPro}\")\n    print(f\"{'='*80}\\n\")\n\n    results = []\n    per_target_scores = []          # new list of best‑score tuples\n    best_counts = np.zeros(N_SAMPLE, dtype=int)\n    tbm_wins = protenix_wins = rnapro_wins = 0\n\n    for target_id in sub_targets:\n        if target_id not in details:\n            continue\n        best_idx, all_scores = details[target_id]\n\n        # obtain the true source list (filled earlier in main)\n        sources_list = SOURCES_BY_TARGET.get(\n            target_id,\n            [\"TBM\"] * TBM + [\"Protenix\"] * PROTENIX +\n            [\"RNAPro\"] * RNAPro +\n            [\"de-novo\"] * (N_SAMPLE - TBM - PROTENIX - RNAPro)\n        )\n\n        score_val = all_scores[best_idx]\n        results.append(score_val)\n        per_target_scores.append((target_id, score_val))   # remember it\n        best_counts[best_idx] += 1\n        src = sources_list[best_idx]\n        if src == \"TBM\":\n            tbm_wins += 1\n        elif src == \"Protenix\":\n            protenix_wins += 1\n        elif src == \"RNAPro\":\n            rnapro_wins += 1\n\n        slot_strs = [f\"{s:.4f}({sources_list[i]})\" for i, s in enumerate(all_scores)]\n        print(f\"{target_id}: {'-'.join(slot_strs)}  best slot {best_idx+1} ({src})\")\n\n    # print best score for each target\n    print(f\"\\nBest score per target:\")\n    for tid, val in per_target_scores:\n        print(f\"  {tid}: {val:.4f}\")\n\n    # compute an average from the results list as a sanity check\n    avg_score = sum(results) / len(results) if results else 0.0\n\n    print(f\"\\n{'='*80}\")\n    print(f\"Summary Statistics\")\n    print(f\"{'='*80}\")\n    print(f\"Average score: {avg_score:.4f}  (mean_score returned {mean_score:.4f})\")\n    print(f\"\\nBest slot distribution:\")\n    for i in range(N_SAMPLE):\n        if i < TBM:\n            label = \"TBM\"\n        elif i < TBM + PROTENIX:\n            label = \"Protenix\"\n        elif i < TBM + PROTENIX + RNAPro:\n            label = \"RNAPro\"\n        else:\n            label = \"de-novo\"\n        print(f\"  Slot {i+1} ({label}): {best_counts[i]} targets\")\n\n    print(f\"\\nWins by source:\")\n    print(f\"  TBM:      {tbm_wins} ({100*tbm_wins/len(results):.1f}%)\")\n    print(f\"  Protenix: {protenix_wins} ({100*protenix_wins/len(results):.1f}%)\")\n    print(f\"  RNAPro:   {rnapro_wins} ({100*rnapro_wins/len(results):.1f}%)\")","metadata":{},"outputs":[],"execution_count":null}]}