{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":16320058},{"sourceType":"datasetVersion","sourceId":14604295,"datasetId":9328538,"databundleVersionId":15440074},{"sourceType":"datasetVersion","sourceId":11472091,"datasetId":7189531,"databundleVersionId":11916011},{"sourceType":"datasetVersion","sourceId":11695366,"datasetId":6791615,"databundleVersionId":12170579},{"sourceType":"datasetVersion","sourceId":14796373,"datasetId":9460350,"databundleVersionId":15650962},{"sourceType":"datasetVersion","sourceId":14787388,"datasetId":9453383,"databundleVersionId":15641141},{"sourceType":"datasetVersion","sourceId":14786962,"datasetId":9447634,"databundleVersionId":15640661},{"sourceType":"datasetVersion","sourceId":11837219,"datasetId":7436926,"databundleVersionId":12332619},{"sourceType":"datasetVersion","sourceId":11056335,"datasetId":6888333,"databundleVersionId":11442412},{"sourceType":"datasetVersion","sourceId":11056379,"datasetId":6888367,"databundleVersionId":11442459},{"sourceType":"datasetVersion","sourceId":5123458,"datasetId":2975803,"databundleVersionId":5194879},{"sourceType":"datasetVersion","sourceId":10880374,"datasetId":6760482,"databundleVersionId":11247092},{"sourceType":"datasetVersion","sourceId":11118830,"datasetId":6933267,"databundleVersionId":11511771},{"sourceType":"datasetVersion","sourceId":11451236,"datasetId":7174725,"databundleVersionId":11892125},{"sourceType":"datasetVersion","sourceId":11230242,"datasetId":7014687,"databundleVersionId":11640305},{"sourceType":"datasetVersion","sourceId":13282339,"datasetId":7162026,"databundleVersionId":13982667},{"sourceType":"datasetVersion","sourceId":14874339,"datasetId":9502242,"databundleVersionId":15736806},{"sourceType":"datasetVersion","sourceId":14764131,"datasetId":9437051,"databundleVersionId":15615594},{"sourceType":"datasetVersion","sourceId":10923077,"datasetId":6785143,"databundleVersionId":11294302},{"sourceType":"datasetVersion","sourceId":14519720,"datasetId":9271415,"databundleVersionId":15347344},{"sourceType":"datasetVersion","sourceId":14534919,"datasetId":9283271,"databundleVersionId":15363994},{"sourceType":"datasetVersion","sourceId":10855324,"datasetId":6742586,"databundleVersionId":11219268},{"sourceType":"modelInstanceVersion","sourceId":311741,"databundleVersionId":11641144,"modelInstanceId":264400,"modelId":285488}],"dockerImageVersionId":31260,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"#!/usr/bin/env python\n# coding: utf-8\n# Stanford RNA 3D Folding 2 - Winning Solution\n# Models: Protenix, DRFold2, TBM, RNApro, Boltz2\n# Ensemble strategy: model selection based on sequence length","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Dependency Installation (Offline Wheels)\n# ============================================================\n# pip install \n!pip install /kaggle/input/datasets/kami1976/biopython-cp312/biopython-1.86-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl\n!pip install /kaggle/input/datasets/yutaroito/boltz-env-depend/boltz_wheels_cuda_v2/gemmi-0.6.5-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n!pip install /kaggle/input/datasets/yutaroito/rdkit-312/rdkit-2025.9.4-cp312-cp312-manylinux_2_28_x86_64.whl\n!pip install /kaggle/input/datasets/yutaroito/boltz-env-depend/boltz_wheels_cuda_v2/chembl_structure_pipeline-1.2.2-py3-none-any.whl \n!pip install /kaggle/input/datasets/yutaroito/boltz-env-depend/boltz_wheels_cuda_v2/einx-0.3.0-py3-none-any.whl    \n\n# ============================================================\n# Standard library imports\n# ============================================================\nimport builtins\nimport gc\nimport json\nimport logging\nimport os\nimport pickle\nimport random\nimport shutil\nimport subprocess\nimport sys\nimport tempfile\nimport time\nimport traceback\nimport typing\nimport warnings\nfrom contextlib import nullcontext\nfrom datetime import datetime\nfrom functools import partial\nfrom pathlib import Path\nfrom timeit import default_timer as timer\n\nwarnings.filterwarnings('ignore')\n# ============================================================\n# Third-party imports\n# ============================================================\nimport numpy as np\nimport pandas as pd\nimport pytz\nimport torch\nimport torch.nn.functional as F\nimport yaml\nfrom Bio import pairwise2\nfrom Bio.Align import PairwiseAligner\nfrom Bio.PDB import PDBParser\nfrom Bio.PDB.MMCIF2Dict import MMCIF2Dict\nfrom Bio.Seq import Seq\nfrom scipy.spatial import distance_matrix\nfrom scipy.spatial.transform import Rotation as R\nfrom tqdm import tqdm\nimport gemmi\nimport pytorch_lightning as pl","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# SECTION 0: Early Exit Guard (Kaggle submission mode check)\n# ============================================================\n\nONLY_INFER = False  # True # False\n\nif ONLY_INFER:\n    # --- Early-exit guard for local runs (put this as the very first cell) ---\n    flag = os.environ.get(\"KAGGLE_IS_COMPETITION_RERUN\")\n\n    # Kaggle sets this to a truthy value during the submission rerun environment.\n    # If not in that environment, write an all-zeros submission and stop.\n    if not flag:\n        candidate_seq_paths = [\n            \"/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv\",\n            \"/kaggle/input/stanford-rna-3d-folding-2/validation_sequences.csv\",\n            \"/kaggle/input/stanford-rna-3d-folding-2/train_sequences.csv\",\n        ]\n\n        seq_df = None\n        for p in candidate_seq_paths:\n            if os.path.exists(p):\n                seq_df = pd.read_csv(p)\n                break\n\n        if seq_df is None or \"target_id\" not in seq_df.columns or \"sequence\" not in seq_df.columns:\n            raise RuntimeError(\n                \"Not in competition rerun, and could not find a sequences file to build a valid zero submission.\\n\"\n                \"Make sure test_sequences.csv is available under /kaggle/input/... or adjust candidate_seq_paths.\"\n            )\n\n        rows = []\n        for row in seq_df.itertuples(index=False):\n            tid = row.target_id\n            seq = row.sequence\n            for i, base in enumerate(seq, start=1):\n                rec = {\n                    \"ID\": f\"{tid}_{i}\",\n                    \"resname\": base,\n                    \"resid\": i,\n                }\n                for k in range(1, 6):\n                    rec[f\"x_{k}\"] = 0.0\n                    rec[f\"y_{k}\"] = 0.0\n                    rec[f\"z_{k}\"] = 0.0\n                rows.append(rec)\n\n        out = pd.DataFrame(rows)\n\n        # Enforce official column order\n        cols = (\n            [\"ID\", \"resname\", \"resid\"] +\n            [f\"{ax}_{k}\" for k in range(1, 6) for ax in (\"x\", \"y\", \"z\")]\n        )\n        out = out[cols]\n        out.to_csv(\"submission.csv\", index=False)\n        print(\"Not in Kaggle competition rerun. Wrote zero submission.csv and exiting early.\")\n        raise SystemExit(0)\n\n    # If flag is truthy: do nothing, notebook continues normally.\n    print(\"KAGGLE_IS_COMPETITION_RERUN detected -> continuing normally.\")\n\n\nIS_SCORING_RUN = os.environ.get('KAGGLE_IS_COMPETITION_RERUN')\nprint(IS_SCORING_RUN)\n\n\nsubprocess.run(\n    [\"pip\", \"install\", \"--no-deps\",\n     \"/kaggle/input/datasets/tobimichigan/biotite-1-2/biotite-1.2.0-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\"],\n    check=False\n)\n\nos.environ[\"LAYERNORM_TYPE\"] = \"torch\"\nos.environ.setdefault(\"RNA_MSA_DEPTH_LIMIT\", \"512\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# SECTION 1: Protenix Inference (seq_len < 250 or >= 1000)\n# ============================================================\n\n# ---- user-configurable paths ----\nDEFAULT_TEST_CSV = \"/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv\"\nDEFAULT_OUTPUT = \"/kaggle/working/protenix_submission.csv\"\nDEFAULT_CODE_DIR = \"/kaggle/input/datasets/qiweiyin/protenix-v1-adjusted/Protenix-v1-adjust-v2/Protenix-v1-adjust-v2/Protenix-v1\"\nDEFAULT_ROOT_DIR = \"/kaggle/input/datasets/qiweiyin/protenix-v1-adjusted/Protenix-v1-adjust-v2/Protenix-v1-adjust-v2/Protenix-v1\"\nMODEL_NAME = \"protenix_base_20250630_v1.0.0\"\ninput_path = \"/kaggle/input/stanford-rna-3d-folding-2\"\nN_SAMPLE = 5\nSEED = 42\n\nMAX_SEQ_LEN = int(os.environ.get(\"MAX_SEQ_LEN\", \"512\"))\nCHUNK_OVERLAP = int(os.environ.get(\"CHUNK_OVERLAP\", \"128\"))  # 64\nMODEL_N_SAMPLE = int(os.environ.get(\"MODEL_N_SAMPLE\", str(N_SAMPLE)))\n\nTEST_TARGETS = False\nif not IS_SCORING_RUN:\n    TEST_TARGETS = os.environ.get(\"TEST_TARGETS\", \"9ZCC\").strip()\n\ndef parse_bool(value: str, default: bool = False) -> str:\n    normalized = str(value).strip().lower()\n    if normalized in {\"1\", \"true\", \"t\", \"yes\", \"y\", \"on\"}:\n        return \"true\"\n    if normalized in {\"0\", \"false\", \"f\", \"no\", \"n\", \"off\"}:\n        return \"false\"\n    return \"true\" if default else \"false\"\n\nos.environ[\"USE_RNA_MSA\"] = \"true\"\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\"))\nprint(f\"USE_MSA: {USE_MSA}\")\nprint(f\"USE_TEMPLATE: {USE_TEMPLATE}\")\nprint(f\"USE_RNA_MSA: {USE_RNA_MSA}\")\n\nTRIANGLE_ATTENTION = os.environ.get(\"TRIANGLE_ATTENTION\", \"torch\").strip().lower()\nTRIANGLE_MULTIPLICATIVE = os.environ.get(\"TRIANGLE_MULTIPLICATIVE\", \"torch\").strip().lower()\n\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() -> tuple[str, str, str]:\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    root_dir = os.environ.get(\"PROTENIX_ROOT_DIR\", DEFAULT_ROOT_DIR)\n    return test_csv, output_csv, code_dir, root_dir\n\n\ndef ensure_required_files(root_dir: str) -> None:\n    checkpoint_path = Path(root_dir) / \"checkpoint\" / f\"{MODEL_NAME}.pt\"\n    components_path = Path(root_dir) / \"common\" / \"components.cif\"\n    components_pkl_path = Path(root_dir) / \"common\" / \"components.cif.rdkit_mol.pkl\"\n    if not checkpoint_path.exists():\n        raise FileNotFoundError(\n            f\"Missing checkpoint: {checkpoint_path}.\\n\"\n            \"Please include checkpoint/protenix_base_default_v1.0.0.pt in PROTENIX_ROOT_DIR.\"\n        )\n    if not components_path.exists():\n        raise FileNotFoundError(\n            f\"Missing CCD file: {components_path}.\\n\"\n            \"Please include common/components.cif in PROTENIX_ROOT_DIR.\"\n        )\n    if not components_pkl_path.exists():\n        raise FileNotFoundError(\n            f\"Missing CCD cache: {components_pkl_path}.\\n\"\n            \"Please include common/components.cif.rdkit_mol.pkl in PROTENIX_ROOT_DIR.\"\n        )\n\n\ndef build_input_json(test_df: pd.DataFrame, json_path: str) -> None:\n    data = []\n    for _, row in test_df.iterrows():\n        seq = row[\"sequence\"]\n        target_id = row[\"target_id\"]\n        data.append(\n            {\n                \"name\": target_id,\n                \"covalent_bonds\": [],\n                \"sequences\": [\n                    {\n                        \"rnaSequence\": {\n                            \"sequence\": seq,\n                            \"count\": 1,\n                        }\n                    }\n                ],\n            }\n        )\n    with open(json_path, \"w\", encoding=\"utf-8\") as f:\n        json.dump(data, f)\n\n\ndef build_single_input_json(target_id: str, seq: str, json_path: str) -> None:\n    data = [\n        {\n            \"name\": target_id,\n            \"covalent_bonds\": [],\n            \"sequences\": [\n                {\n                    \"rnaSequence\": {\n                        \"sequence\": seq,\n                        \"count\": 1,\n                    }\n                }\n            ],\n        }\n    ]\n    with open(json_path, \"w\", encoding=\"utf-8\") as f:\n        json.dump(data, f)\n\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 = {**configs_base, **{\"data\": data_configs}, **inference_configs}\n\n    def deep_update(target, patch):\n        for key, value in patch.items():\n            if isinstance(value, dict) and key in target and isinstance(target[key], dict):\n                deep_update(target[key], value)\n            else:\n                target[key] = value\n\n    deep_update(base_configs, model_configs[model_name])\n\n    arg_str = \" \".join(\n        [\n            f\"--model_name {model_name}\",\n            f\"--input_json_path {input_json_path}\",\n            f\"--dump_dir {dump_dir}\",\n            f\"--use_msa {USE_MSA}\",\n            f\"--use_template {USE_TEMPLATE}\",\n            f\"--use_rna_msa {USE_RNA_MSA}\",\n            f\"--sample_diffusion.N_sample {MODEL_N_SAMPLE}\",\n            f\"--seeds {SEED}\",\n        ]\n    )\n    configs = parse_configs(\n        configs=base_configs,\n        arg_str=arg_str,\n        fill_required_with_null=True,\n    )\n    return configs\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                centre_mask = atom_array.centre_atom_mask == 1\n                if hasattr(atom_array, \"is_rna\"):\n                    centre_mask = centre_mask & (atom_array.is_rna)\n                return torch.from_numpy(centre_mask)\n            if hasattr(atom_array, \"atom_name\"):\n                if hasattr(atom_array, \"is_rna\"):\n                    return torch.from_numpy(\n                        (atom_array.atom_name == \"C1'\") & (atom_array.is_rna)\n                    )\n                return torch.from_numpy(atom_array.atom_name == \"C1'\")\n        except Exception:\n            pass\n\n    features = data[\"input_feature_dict\"]\n    if \"centre_atom_mask\" in features:\n        return features[\"centre_atom_mask\"].long() == 1\n    return features[\"atom_to_tokatom_idx\"].long() == 12\n\n\ndef get_feature_c1_mask(data: dict) -> torch.Tensor:\n    features = data[\"input_feature_dict\"]\n    if \"centre_atom_mask\" in features:\n        return features[\"centre_atom_mask\"].long() == 1\n    return features[\"atom_to_tokatom_idx\"].long() == 12\n\n\ndef coords_to_rows(target_id: str, seq: str, coords: np.ndarray) -> list[dict]:\n    rows = []\n    n_res = len(seq)\n    for i in range(n_res):\n        row = {\n            \"ID\": f\"{target_id}_{i + 1}\",\n            \"resname\": seq[i],\n            \"resid\": i + 1,\n        }\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, target_samples: int) -> np.ndarray:\n    if coords.shape[0] >= target_samples:\n        return coords\n    if coords.shape[0] == 0:\n        return np.zeros((target_samples, coords.shape[1], 3), dtype=coords.dtype)\n    repeat = target_samples - coords.shape[0]\n    extra = np.repeat(coords[:1], repeat, axis=0)\n    return np.concatenate([coords, extra], axis=0)\n\n\ndef split_sequence(seq: str, max_len: int, overlap: int) -> list[tuple[int, int, str]]:\n    if max_len <= 0:\n        raise ValueError(\"max_len must be > 0\")\n    if overlap >= max_len:\n        raise ValueError(\"overlap must be smaller than max_len\")\n    if len(seq) <= max_len:\n        return [(0, len(seq), seq)]\n    step = max_len - overlap\n    chunks = []\n    start = 0\n    while start < len(seq):\n        end = min(start + max_len, len(seq))\n        chunks.append((start, end, seq[start:end]))\n        if end == len(seq):\n            break\n        start += step\n    return chunks\n\ndef main_protenix() -> None:\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}. Set PROTENIX_CODE_DIR to the repo path.\"\n        )\n\n    os.environ[\"PROTENIX_ROOT_DIR\"] = root_dir\n    sys.path.append(code_dir)\n\n    ensure_required_files(root_dir)\n    seed_everything(SEED)\n\n    test_df = pd.read_csv(test_csv)\n    test_df[\"sequence_len\"] = test_df[\"sequence\"].str.len()\n    test_df = test_df[(test_df['sequence_len'] < 250) | (test_df[\"sequence_len\"] >= 1000)]\n\n    if not IS_SCORING_RUN:\n        if TEST_TARGETS:\n            keep = {t.strip() for t in TEST_TARGETS.split(\",\") if t.strip()}\n            test_df = test_df[test_df[\"target_id\"].isin(keep)].reset_index(drop=True)\n            if test_df.empty:\n                raise ValueError(\"TEST_TARGETS did not match any target_id values\")\n\n    work_dir = Path(\"/kaggle/working\")\n\n    work_dir.mkdir(parents=True, exist_ok=True)\n    input_json_path = str(work_dir / \"protenix_input.json\")\n    build_input_json(test_df, input_json_path)\n\n    chunks_dir = work_dir / \"chunks\"\n    chunks_dir.mkdir(parents=True, exist_ok=True)\n\n    from protenix.data.inference.infer_dataloader import InferenceDataset\n    from runner.inference import InferenceRunner, update_gpu_compatible_configs, update_inference_configs\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\n    seq_by_id = dict(zip(test_df.target_id.tolist(), test_df.sequence.tolist()))\n\n    all_predictions = []\n\n    debug_path = Path(output_csv).with_suffix(\".debug.csv\")\n    debug_rows = []\n\n    for _, row in tqdm(test_df.iterrows(), total=len(test_df)):\n        target_id = row[\"target_id\"]\n        seq = row[\"sequence\"]\n\n        if len(seq) > MAX_SEQ_LEN:\n            chunks = split_sequence(seq, MAX_SEQ_LEN, CHUNK_OVERLAP)\n            print(\n                f\"[{target_id}] long sequence: {len(seq)} -> {len(chunks)} chunks \"\n                f\"(max_len={MAX_SEQ_LEN}, overlap={CHUNK_OVERLAP})\"\n            )\n            full_coords = np.zeros((N_SAMPLE, len(seq), 3), dtype=np.float32)\n            for chunk_idx, (start, end, subseq) in enumerate(chunks):\n                chunk_id = f\"{target_id}__chunk{chunk_idx:03d}\"\n                chunk_json_path = str(chunks_dir / f\"{chunk_id}.json\")\n                build_single_input_json(chunk_id, subseq, chunk_json_path)\n                chunk_configs = build_configs(\n                    chunk_json_path, str(work_dir / \"outputs\"), MODEL_NAME\n                )\n                chunk_configs = update_gpu_compatible_configs(chunk_configs)\n                chunk_dataset = InferenceDataset(chunk_configs)\n                chunk_data, chunk_atom_array, chunk_error = chunk_dataset[0]\n\n                if chunk_error:\n                    print(f\"[{chunk_id}] data_error: {chunk_error.splitlines()[0]}\")\n                    debug_rows.append(\n                        {\n                            \"target_id\": target_id,\n                            \"seq_len\": len(seq),\n                            \"c1_count\": 0,\n                            \"is_rna_atoms\": -1,\n                            \"coord_abs_max\": 0.0,\n                            \"error\": chunk_error,\n                            \"split_mode\": \"chunked\",\n                            \"chunk_idx\": chunk_idx,\n                            \"chunk_start\": start,\n                            \"chunk_end\": end,\n                            \"chunk_len\": end - start,\n                        }\n                    )\n                    continue\n\n                new_configs = update_inference_configs(\n                    configs, chunk_data[\"N_token\"].item()\n                )\n                runner.update_model_configs(new_configs)\n\n                prediction = runner.predict(chunk_data)\n                coords = prediction[\"coordinate\"]\n                mask = get_c1_mask(chunk_data, chunk_atom_array)\n                c1_count = int(mask.sum().item()) if hasattr(mask, \"sum\") else 0\n                if c1_count == 0:\n                    mask = get_feature_c1_mask(chunk_data)\n                    c1_count = int(mask.sum().item()) if hasattr(mask, \"sum\") else 0\n                    print(\n                        f\"[{chunk_id}] C1' atoms: {c1_count} (seq_len={len(subseq)}) fallback\"\n                    )\n                else:\n                    print(\n                        f\"[{chunk_id}] C1' atoms: {c1_count} (seq_len={len(subseq)})\"\n                    )\n                is_rna_count = (\n                    int(chunk_atom_array.is_rna.sum())\n                    if hasattr(chunk_atom_array, \"is_rna\")\n                    else -1\n                )\n                coord_abs_max = float(coords.abs().max().item())\n                debug_rows.append(\n                    {\n                        \"target_id\": target_id,\n                        \"seq_len\": len(seq),\n                        \"c1_count\": c1_count,\n                        \"is_rna_atoms\": is_rna_count,\n                        \"coord_abs_max\": coord_abs_max,\n                        \"error\": \"\",\n                        \"split_mode\": \"chunked\",\n                        \"chunk_idx\": chunk_idx,\n                        \"chunk_start\": start,\n                        \"chunk_end\": end,\n                        \"chunk_len\": end - start,\n                    }\n                )\n                coords_all = prediction[\"coordinate\"]  # (S, N_atom, 3)\n\n                res_ids = chunk_atom_array.res_id\n                atom_names = chunk_atom_array.atom_name\n\n                c1_indices = []\n                for r in range(1, len(subseq) + 1):\n                    idx = np.where(\n                        (res_ids == r) &\n                        (atom_names == \"C1'\")\n                    )[0]\n                    if len(idx) == 0:\n                        continue\n                    c1_indices.append(idx[0])\n\n                coords = coords_all[:, c1_indices, :]\n                coords = coords.detach().cpu().numpy()\n\n                if coords.shape[1] != len(subseq):\n                    padded = np.zeros((coords.shape[0], len(subseq), 3), dtype=np.float32)\n                    min_len = min(coords.shape[1], len(subseq))\n                    if min_len > 0:\n                        padded[:, :min_len, :] = coords[:, :min_len, :]\n                    coords = padded\n\n                coords = pad_samples(coords, N_SAMPLE)\n\n                full_coords[:, start:end, :] = coords[:, : end - start, :]\n                del prediction, coords, mask, chunk_data, chunk_atom_array, chunk_dataset\n                torch.cuda.empty_cache()\n                gc.collect()\n\n            rows = coords_to_rows(target_id, seq, full_coords)\n            all_predictions.extend(rows)\n            continue\n\n        elif len(seq) <= MAX_SEQ_LEN:\n            single_json = str(work_dir / f\"{target_id}.json\")\n            build_single_input_json(target_id, seq, single_json)\n\n            single_configs = build_configs(\n                single_json,\n                str(work_dir / \"outputs\"),\n                MODEL_NAME,\n            )\n            single_configs = update_gpu_compatible_configs(single_configs)\n\n            dataset = InferenceDataset(single_configs)\n            data, atom_array, error_message = dataset[0]\n            if error_message:\n                print(f\"[{target_id}] data_error: {error_message.splitlines()[0]}\")\n                debug_rows.append(\n                    {\n                        \"target_id\": target_id,\n                        \"seq_len\": len(seq),\n                        \"c1_count\": 0,\n                        \"is_rna_atoms\": -1,\n                        \"coord_abs_max\": 0.0,\n                        \"error\": error_message,\n                        \"split_mode\": \"full\",\n                        \"chunk_idx\": -1,\n                        \"chunk_start\": 0,\n                        \"chunk_end\": len(seq),\n                        \"chunk_len\": len(seq),\n                    }\n                )\n                rows = coords_to_rows(\n                    target_id,\n                    seq,\n                    np.zeros((0, 0, 3), dtype=np.float32),\n                )\n                all_predictions.extend(rows)\n                continue\n\n            new_configs = update_inference_configs(configs, data[\"N_token\"].item())\n            runner.update_model_configs(new_configs)\n\n            prediction = runner.predict(data)\n            coords = prediction[\"coordinate\"]\n            mask = get_c1_mask(data, atom_array)\n            c1_count = int(mask.sum().item()) if hasattr(mask, \"sum\") else 0\n            if c1_count == 0:\n                mask = get_feature_c1_mask(data)\n                c1_count = int(mask.sum().item()) if hasattr(mask, \"sum\") else 0\n                print(\n                    f\"[{target_id}] C1' atoms: {c1_count} (seq_len={len(seq)}) fallback\"\n                )\n            else:\n                print(f\"[{target_id}] C1' atoms: {c1_count} (seq_len={len(seq)})\")\n            is_rna_count = (\n                int(atom_array.is_rna.sum()) if hasattr(atom_array, \"is_rna\") else -1\n            )\n            coord_abs_max = float(coords.abs().max().item())\n            debug_rows.append(\n                {\n                    \"target_id\": target_id,\n                    \"seq_len\": len(seq),\n                    \"c1_count\": c1_count,\n                    \"is_rna_atoms\": is_rna_count,\n                    \"coord_abs_max\": coord_abs_max,\n                    \"error\": \"\",\n                    \"split_mode\": \"full\",\n                    \"chunk_idx\": -1,\n                    \"chunk_start\": 0,\n                    \"chunk_end\": len(seq),\n                    \"chunk_len\": len(seq),\n                }\n            )\n\n            coords_all = prediction[\"coordinate\"]  # (S, N_atom, 3)\n\n            res_ids = atom_array.res_id\n            atom_names = atom_array.atom_name\n\n            c1_indices = []\n            for r in range(1, len(seq) + 1):\n                idx = np.where(\n                    (res_ids == r) &\n                    (atom_names == \"C1'\")\n                )[0]\n                if len(idx) == 0:\n                    continue\n                c1_indices.append(idx[0])\n\n            coords = coords_all[:, c1_indices, :]\n            coords = coords.detach().cpu().numpy()\n\n            if coords.shape[1] != len(seq):\n                padded = np.zeros((coords.shape[0], len(seq), 3), dtype=np.float32)\n                min_len = min(coords.shape[1], len(seq))\n                if min_len > 0:\n                    padded[:, :min_len, :] = coords[:, :min_len, :]\n                coords = padded\n\n            coords = pad_samples(coords, N_SAMPLE)\n\n            rows = coords_to_rows(target_id, seq, coords)\n            all_predictions.extend(rows)\n            del prediction, coords, mask, data, atom_array\n            torch.cuda.empty_cache()\n            gc.collect()\n\n    if debug_rows:\n        pd.DataFrame(debug_rows).to_csv(debug_path, index=False)\n\n    sub = pd.DataFrame(all_predictions)\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    sub[coord_cols] = sub[coord_cols].clip(-999.999, 9999.999)\n    sub[cols].to_csv(output_csv, index=False)\n\n    print(f\"Saved submission to {output_csv}\")\n\n\nif __name__ == \"__main__\":\n    main_protenix()\n\n\nprotenix_pred = pd.read_csv('/kaggle/working/protenix_submission.csv')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# SECTION 2: DRFold2 Inference (seq_len < 250)\n# ============================================================\n\ntest_sequences = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv\")\n\nis_submission_mode = len(test_sequences) != 12\n\n\nos.environ['CUDA_VISIBLE_DEVICES'] = '0'\n\nsys.argv = ['notebook', 'cuda' if torch.cuda.is_available() else 'cpu', 'fp32'] #'fp16' #'fp32'# 'bf16' #\ndevice = sys.argv[1]\nsys_dtype = sys.argv[2] if len(sys.argv) > 2 else 'fp32'\n\nprint('Using device:', device)\nprint('Using dtype:', sys_dtype)\n\n# dr settings\nNUM_CONF=5\nMAX_LENGTH=480\nMAX_CAT_LENGTH=2400\n\nCFG_DIR='cfg_97'\n\nDEVICE=device #'cuda' #'cpu'#\nPREC=sys_dtype \n\n\nif PREC=='fp16':\n    torch.set_default_dtype(torch.float16)\nif PREC=='bf16':\n    torch.set_default_dtype(torch.bfloat16)\n\n\nfrom datetime import datetime\nimport pytz\nprint('LOGGING TIME OF START:',  datetime.strftime(datetime.now(pytz.timezone('Asia/Singapore')), \"%Y-%m-%d %H:%M:%S\"))\n\n\nprint('PIP INSTALL OK !!!!')\nimport os,sys\n\nimport pandas as pd\npd.set_option('display.max_columns', 20)\npd.set_option('display.expand_frame_repr', False)\n\nimport numpy as np\nimport torch\nimport torch.nn.functional as F\nfrom timeit import default_timer as timer\n\n\n# helper--\nclass dotdict(dict):\n\t__setattr__ = dict.__setitem__\n\t__delattr__ = dict.__delitem__\n\n\tdef __getattr__(self, name):\n\t\ttry:\n\t\t\treturn self[name]\n\t\texcept KeyError:\n\t\t\traise AttributeError(name)\n\ndef time_to_str(t, mode='min'):\n\tif mode=='min':\n\t\tt  = int(t)/60\n\t\thr = t//60\n\t\tmin = t%60\n\t\treturn '%2d hr %02d min'%(hr,min) \n\telif mode=='sec':\n\t\tt   = int(t)\n\t\tmin = t//60\n\t\tsec = t%60\n\t\treturn '%2d min %02d sec'%(min,sec)\n\n\telse:\n\t\traise NotImplementedError\n\ndef gpu_memory_use():\n    if torch.cuda.is_available():\n        device = torch.device(0)\n        free, total = torch.cuda.mem_get_info(device)\n        used= (total - free) / 1024 ** 3\n        return round(used,2)\n    else:\n        return 0\n\nprint('torch',torch.__version__)\nprint('torch.cuda',torch.version.cuda)\n\nprint('IMPORT OK!!!')\n\nDATA_KAGGLE_DIR = '/kaggle/input/stanford-rna-3d-folding-2'\n\nvalid_df = pd.read_csv(f'{DATA_KAGGLE_DIR}/test_sequences.csv')\n# Sort test sequences by length to process shorter ones with DRfold2\nvalid_df[\"sequence_len\"] = valid_df[\"sequence\"].str.len()\nvalid_df = valid_df[valid_df['sequence_len'] < 250]\nvalid_df.reset_index(drop=True, inplace=True)\nprint(valid_df.shape)\n \nif not IS_SCORING_RUN:\n    valid_df = valid_df.head(5)\n\nprint('len(valid_df)',len(valid_df))\nprint(valid_df.iloc[0])\nprint('')\n\n\nprint('SETTING OK!!!')\n\n\nsys.path.append('/kaggle/input/datasets/z1493916656/drfold-model-bf16/PotentialFold')\nimport a2b\ndef frame_coor_to_C1(coor, seq, BASE_COOR, OTHER_COOR):\n    tx = torch.as_tensor(coor, dtype=torch.float32)\n    basex = torch.from_numpy(Get_base(seq, BASE_COOR)).to(tx.dtype)\n    otherx = torch.from_numpy(Get_base(seq, OTHER_COOR)).to(tx.dtype)\n    L = len(seq)\n\n    x = torch.rand((L, 21), dtype=tx.dtype, device=tx.device)\n\n    biasq = tx.mean(dim=1, keepdim=True)            # (L, 1, 3)\n    q = tx - biasq                                  # (L, 3, 3)\n\n    m = torch.einsum('bnz,bny->bzy', basex, q).reshape(L, -1)\n    x[:, :9] = m\n    x[:, 9:18] = m\n\n    x[:, 18:] = biasq.squeeze(1)\n    rama = x.double()  \n\n    #    otherx: (L, 5, 3) -> other_xyz: (L, 5, 3)\n    other_xyz = a2b.quat2b(otherx.double(), rama[:, 9:]).float().cpu().numpy()\n\n    c1_xyz = other_xyz[:, 4, :]\n    return c1_xyz\n\n\ndef concat_coor(out1: dict, out2: dict) -> np.ndarray:\n    coor1 = torch.as_tensor(out1['coor'], dtype=torch.float64)   # (L1,3,3)\n    coor2 = torch.as_tensor(out2['coor'], dtype=torch.float64)   # (L2,3,3)\n\n    f1 = coor1[-1]   # (3,3)\n    f2 = coor2[0]    # (3,3)\n\n    bias1 = f1.mean(dim=0)   # (3,)\n    bias2 = f2.mean(dim=0)   # (3,)\n    basex = f1 - bias1       # (3,3)\n    q     = f2 - bias2       # (3,3)\n\n    #    R_{ij} = sum_z basex_{iz} * q_{jz}\n    R = torch.einsum('iz,jz->ij', basex, q)   # (3,3)\n\n    t = bias1 - (R @ bias2)                  # (3,)\n\n    L2 = coor2.shape[0]\n    rama = torch.empty((L2, 12), dtype=torch.float64, device=coor2.device)\n    R_flat = R.reshape(1, 9).repeat(L2, 1)    # (L2,9)\n    t_rep  = t.reshape(1, 3).repeat(L2, 1)    # (L2,3)\n    rama[:, :9] = R_flat\n    rama[:, 9:] = t_rep\n\n    coor2_aligned = a2b.quat2b(coor2, rama)   # torch.Tensor (L2,3,3)\n\n    coor_cat = torch.cat([coor1, coor2_aligned[1:]], dim=0)  # (L1+L2-1,3,3)\n\n    return coor_cat.cpu().numpy()\n\n\ndef Get_base(seq, basenpy_standard):\n    n_atoms = basenpy_standard.shape[1]\n    basenpy = np.zeros([len(seq), n_atoms, 3])\n    seqnpy = np.array(list(seq))\n    basenpy[seqnpy=='A'] = basenpy_standard[0]\n    basenpy[seqnpy=='a'] = basenpy_standard[0]\n    basenpy[seqnpy=='G'] = basenpy_standard[1]\n    basenpy[seqnpy=='g'] = basenpy_standard[1]\n    basenpy[seqnpy=='C'] = basenpy_standard[2]\n    basenpy[seqnpy=='c'] = basenpy_standard[2]\n    basenpy[seqnpy=='U'] = basenpy_standard[3]\n    basenpy[seqnpy=='u'] = basenpy_standard[3]\n    basenpy[seqnpy=='T'] = basenpy_standard[3]\n    basenpy[seqnpy=='t'] = basenpy_standard[3]\n    return basenpy\n\n\nimport numpy as np\n\ndef score_energy_one_simple(seq, target_id, out):\n\n    coor = out['coor']  # (L, 3, 3)\n    L = len(seq)\n\n    d0_P_C4  = 1.60\n    d0_C4_N  = 1.47\n    k_bond   = 100.0 \n\n    energy_bond = 0.0\n    for i in range(L):\n        P  = coor[i,0]\n        C4 = coor[i,1]\n        N  = coor[i,2]\n        d_PC4 = np.linalg.norm(P - C4)\n        d_C4N = np.linalg.norm(C4 - N)\n        energy_bond += k_bond * (d_PC4 - d0_P_C4)**2\n        energy_bond += k_bond * (d_C4N - d0_C4_N)**2\n\n    theta0   = np.deg2rad(109.5)\n    k_angle  = 20.0   \n\n    energy_angle = 0.0\n    for i in range(L):\n        P  = coor[i,0]\n        C4 = coor[i,1]\n        N  = coor[i,2]\n        v1 = P  - C4\n        v2 = N  - C4\n        cos_theta = np.dot(v1, v2) / (np.linalg.norm(v1)*np.linalg.norm(v2) + 1e-8)\n        theta = np.arccos(np.clip(cos_theta, -1.0, 1.0))\n        energy_angle += k_angle * (theta - theta0)**2\n\n    d0_stack = 3.4\n    k_stack  = 5.0   # (kcal/mol/Å²)\n\n    energy_stack = 0.0\n    for i in range(L-1):\n        C4_i   = coor[i  ,1]\n        C4_ip1 = coor[i+1,1]\n        d = np.linalg.norm(C4_i - C4_ip1)\n        energy_stack += k_stack * (d - d0_stack)**2\n\n    total_energy = energy_bond + energy_angle + energy_stack\n\n    # print(f\"[{target_id}] bond={energy_bond:.2f}, angle={energy_angle:.2f}, stack={energy_stack:.2f} → total={total_energy:.2f}\")\n\n    return total_energy\n\n\nimport numpy as np\nimport os\nimport json\nimport pickle\nimport numpy as np\nimport tempfile\nfrom Bio.PDB import PDBParser\nsys.path.append('/kaggle/input/drfold-model/DRfold2/PotentialFold')\nfrom Optimization import Structure\n\nsys.path.append('/kaggle/input/datasets/z1493916656/drfold-model-bf16')\nsys.path.append('/kaggle/input/datasets/z1493916656/drfold-model-bf16/PotentialFold')\nsys.path.append(f'/kaggle/input/datasets/z1493916656/drfold-model-bf16/{CFG_DIR}')\nsys.path.append(f'/kaggle/input/datasets/z1493916656/drfold-model-bf16/{CFG_DIR}/RNALM2')\n\nBASE_COOR = np.load('/kaggle/input/datasets/z1493916656/drfold-model-bf16/PotentialFold/lib/base.npy')\nOTHER_COOR = np.load('/kaggle/input/datasets/z1493916656/drfold-model-bf16/PotentialFold/lib/other2.npy')\nSIDE_COOR = np.load('/kaggle/input/datasets/z1493916656/drfold-model-bf16/PotentialFold/lib/side.npy')\n\nfrom EvoMSA2XYZ import MSA2XYZ\nfrom RNALM2.Model import RNA2nd\nfrom data import parse_seq\n\n\n# data helper\ndef make_data(seq, device):\n    aa_type = parse_seq(seq)\n    base = Get_base(seq, BASE_COOR)\n    seq_idx = np.arange(len(seq)) + 1\n\n    msa = aa_type[None, :]\n    msa = torch.from_numpy(msa)\n    msa = torch.cat([msa, msa], 0)  # ???\n    msa = F.one_hot(msa.long(), 6).float()\n\n    base_x = torch.from_numpy(base).float()\n    seq_idx = torch.from_numpy(seq_idx).long()\n\n    msa, base_x, seq_idx = msa.to(device), base_x.to(device), seq_idx.to(device)\n    return msa, base_x, seq_idx\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\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\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\ndef run_submit(valid_df):\n    \n    #load model (these are moified versions, not the same from their github repo)\n    rnalm = RNA2nd(dict(\n        s_in_dim=5,\n        z_in_dim=2,\n        s_dim= 512,\n        z_dim= 128,\n        N_elayers=18,\n    ))\n    rnalm_file = '/kaggle/input/datasets/z1493916656/drfold-model-bf16/model_hub/RCLM/epoch_67000'\n    print(rnalm_file)\n    print(\n        rnalm.load_state_dict(torch.load(rnalm_file, map_location='cpu', weights_only=True), strict=False)\n        #Unexpected key(s) in state_dict: \"ss_head.linear.weight\", \"ss_head.linear.bias\".\n    )\n    rnalm = rnalm.to(DEVICE)\n    if(PREC=='fp16'):\n        rnalm=rnalm.half()\n        \n    if PREC=='bf16':\n        rnalm = rnalm.bfloat16()\n        \n    rnalm = rnalm.eval()\n\n    #---\n    msa2xyz = MSA2XYZ(\n        seq_dim=6,\n        msa_dim=7,\n        N_ensemble=1,\n        N_cycle=8,  # 8\n        m_dim=64,\n        s_dim=64,\n        z_dim=64,\n    )\n    msa2xyz_file = [f'/kaggle/input/datasets/z1493916656/drfold-model-bf16/model_hub/{CFG_DIR}/model_{i}' for i in range(20)]\n    num_msa2xyz = len(msa2xyz_file) \n    msa2xyz_state_dict = []\n    for c in range(num_msa2xyz):\n        if c==0: print(msa2xyz_file[c])\n        m = torch.load(msa2xyz_file[c], map_location='cpu', weights_only=True)\n        msa2xyz_state_dict.append(m)\n        \n    print(msa2xyz.load_state_dict(msa2xyz_state_dict[0], strict=False))\n    msa2xyz = msa2xyz.to(DEVICE)\n    if(PREC=='fp16'):\n        msa2xyz=msa2xyz.half()\n\n    if PREC=='bf16':\n        msa2xyz = msa2xyz.bfloat16()\n    msa2xyz = msa2xyz.eval()\n    \n    msa2xyz.msaxyzone.premsa.rnalm = rnalm\n\n    #---\n\n\n    drfold_pred = [] \n    total_time_taken = 0\n    max_gpu_mem_used = 0\n\n    for i, row in valid_df.iterrows():\n        start_timer = timer()\n        target_id = row.target_id  # 'R1116' #casp15 R1116: len(157)\n        sequence = row.sequence\n        seq = row.sequence  \n        L = len(seq)\n        if L > MAX_CAT_LENGTH:\n            seq = seq[:MAX_CAT_LENGTH]\n        print(i, target_id, L, len(seq), seq[:75] + '...')\n\n        \n        if len(seq)>480:\n            model_to_try=[16, 9, 1, 2, 0]\n        elif len(seq)>200:\n            model_to_try = [13, 6, 14, 5, 3]\n        elif  len(seq)>100:\n            model_to_try = [13, 6, 14, 12, 7, 2, 5, 19, 10, 9]\n        else:\n\n            model_to_try = [1, 2, 0, 8, 7, 5, 6, 14, 10, 18, 4, 13, 3, 17, 19, 11, 12, 15, 16, 9]\n\n\n        # Segment prediction for long sequences\n        def predict_segment(seq):\n            msa, base_x, seq_idx = make_data(seq, DEVICE)\n            with torch.no_grad():\n                if PREC=='fp16':\n                    msa, base_x = msa.half(), base_x.half()\n                if PREC=='bf16':\n                    msa, base_x = msa.bfloat16(), base_x.bfloat16()\n                return msa2xyz.pred(msa, seq_idx, None, base_x, np.array(list(seq)))\n\n                \n        energy = []\n        coordinate=[]\n        outputs=[]\n        for c in model_to_try:\n            msa2xyz.load_state_dict(msa2xyz_state_dict[c], strict=False)\n\n            if len(seq) <= MAX_LENGTH:\n                outs = [ predict_segment(seq) ]\n            else:\n                step = MAX_LENGTH - 1\n                outs = []\n                for s in range(0, len(seq), step):\n                    seg = seq[s : min(s+MAX_LENGTH, len(seq))]\n                    outs.append(predict_segment(seg))\n                    \n            out_cat = outs[0]\n            for out_seg in outs[1:]:\n                out_cat = {'coor': concat_coor(out_cat, out_seg)}\n                    \n            e = score_energy_one_simple(seq, target_id, out_cat)\n            energy.append(e) #tranucated sequence\n            \n            if L != len(seq):\n                out_cat['coor'] = np.pad(out_cat['coor'], ((0, L - len(seq)), (0, 0), (0, 0)), 'constant', constant_values=0)\n                \n            outputs.append(out_cat)\n            \n            \n            xyz = frame_coor_to_C1(out_cat['coor'], sequence, BASE_COOR, OTHER_COOR)\n            \n            coordinate.append(xyz)\n            \n\n            time_taken = timer() - start_timer\n            total_time_taken += time_taken\n    \n            gpu_mem_used = gpu_memory_use()\n            max_gpu_mem_used = max(max_gpu_mem_used,gpu_mem_used)\n    \n            print(f'{c:02d}   energy:{e:10.0f}   out_cat{str(out_cat[\"coor\"].shape)}  time:{time_to_str(time_taken, mode=\"sec\")}   gpu={gpu_mem_used} gb')\n\n            \n        #------- \n        torch.cuda.empty_cache()\n        \n        energy = np.array(energy)\n        energy_mean = np.mean(energy)\n        energy = np.abs(energy - energy_mean)\n        #select top5\n        argsort = np.array(energy).argsort()\n        argsort = argsort[:5]\n        \n        df = coord_to_df(row.sequence, [coordinate[k] for k in argsort], row.target_id)\n        drfold_pred.append(df)\n    \n    print('----------------------------------------')\n    print('MAX_LENGTH', MAX_LENGTH)\n    print('### total_time_taken:', time_to_str(total_time_taken, mode='min'))\n    print('### max_gpu_mem_used:', max_gpu_mem_used, 'GB')\n    print('')\n\n    drfold_pred = pd.concat(drfold_pred)\n    drfold_pred.to_csv(f'/kaggle/working/drfold_submission.csv', index=False)\n    print(drfold_pred)\n    return drfold_pred\n\nrun_submit(valid_df)\n\nprint('SUBMIT OK!!!')\n\n\ndrfold_pred = pd.read_csv('/kaggle/working/drfold_submission.csv')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# SECTION 3: TBM (Template-Based Modeling, all sequences)\n# ============================================================\n\n!pip install --no-index /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\nimport pandas as pd\nimport numpy as np\nimport random\nimport time\nimport warnings\nimport os, sys\n\nwarnings.filterwarnings('ignore')\n\nDATA_PATH = '/kaggle/input/stanford-rna-3d-folding-2/'\ntrain_seqs = pd.read_csv(DATA_PATH + 'train_sequences.csv')\ntest_seqs = pd.read_csv(DATA_PATH + 'test_sequences.csv')\ntrain_labels = pd.read_csv(DATA_PATH + 'train_labels.csv')\n\nsys.path.append(os.path.join(DATA_PATH, \"extra\"))\n\n# --- Robust import for Kaggle's extra/parse_fasta_py.py (it may miss typing imports) ---\ntry:\n    import typing as _typing\n    import builtins as _builtins\n\n    # Make these names available during module import-time annotation evaluation\n    _builtins.Dict  = getattr(_typing, \"Dict\")\n    _builtins.Tuple = getattr(_typing, \"Tuple\")\n    _builtins.List  = getattr(_typing, \"List\")\n\n    from parse_fasta_py import parse_fasta as _parse_fasta_raw\n\n    # Normalize output to: {chain_id: sequence_string}\n    def parse_fasta(fasta_content: str):\n        d = _parse_fasta_raw(fasta_content)\n        out = {}\n        for k, v in d.items():\n            # some variants return (sequence, headers/lines) or similar\n            out[k] = v[0] if isinstance(v, tuple) else v\n        return out\n\nexcept Exception:\n    # Fallback FASTA parser: {chain_id: sequence_string}\n    def parse_fasta(fasta_content: str):\n        out = {}\n        cur = None\n        seq_parts = []\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(seq_parts)\n                header = line[1:]\n                # First token is usually chain id in this dataset\n                cur = header.split()[0]\n                seq_parts = []\n            else:\n                seq_parts.append(line.replace(\" \", \"\"))\n        if cur is not None:\n            out[cur] = \"\".join(seq_parts)\n        return out\n\ndef parse_stoichiometry(stoich: str):\n    if pd.isna(stoich) or str(stoich).strip() == \"\":\n        return []\n    out = []\n    for part in str(stoich).split(';'):\n        ch, cnt = part.split(':')\n        out.append((ch.strip(), int(cnt)))\n    return out\n\ndef get_chain_segments(row):\n    \"\"\"\n    Returns list of (start,end) segments in row['sequence'] corresponding to chain copies in stoichiometry order.\n    Falls back to single segment if parsing fails.\n    \"\"\"\n    seq = row['sequence']\n    stoich = row.get('stoichiometry', '')\n    all_seq = row.get('all_sequences', '')\n\n    if pd.isna(stoich) or pd.isna(all_seq) or str(stoich).strip()==\"\" or str(all_seq).strip()==\"\":\n        return [(0, len(seq))]\n\n    try:\n        chain_dict = parse_fasta(all_seq)  # dict: chain_id -> sequence\n        order = parse_stoichiometry(stoich)\n        segs = []\n        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                L = len(base)\n                segs.append((pos, pos + L))\n                pos += L\n        if pos != len(seq):\n            return [(0, len(seq))]\n        return segs\n    except Exception:\n        return [(0, len(seq))]\n\ndef build_segments_map(df):\n    seg_map = {}\n    stoich_map = {}\n    for _, r in df.iterrows():\n        tid = r['target_id']\n        seg_map[tid] = get_chain_segments(r)\n        stoich_map[tid] = str(r.get('stoichiometry', '') if not pd.isna(r.get('stoichiometry', '')) else '')\n    return seg_map, stoich_map\n\ntrain_segs_map, train_stoich_map = build_segments_map(train_seqs)\ntest_segs_map,  test_stoich_map  = build_segments_map(test_seqs)\n\ndef process_labels(labels_df):\n    coords_dict = {}\n    # Faster + safer prefix extraction\n    prefixes = labels_df['ID'].str.rsplit('_', n=1).str[0]\n    for id_prefix, group in labels_df.groupby(prefixes):\n        coords_dict[id_prefix] = group.sort_values('resid')[['x_1', 'y_1', 'z_1']].values\n    return coords_dict\n\ntrain_coords_dict = process_labels(train_labels)\n\nfrom Bio.Align import PairwiseAligner\n\naligner = PairwiseAligner()\naligner.mode = 'global'\naligner.match_score = 2\naligner.mismatch_score = -1.5\n\n# Stronger gap penalties discourage \"sliding\" (critical: residue numbering must match)\naligner.open_gap_score   = -8\naligner.extend_gap_score = -0.4\n\n# Also penalize terminal gaps (prevents end-gap semi-global behavior)\naligner.query_left_open_gap_score  = -8\naligner.query_left_extend_gap_score = -0.4\naligner.query_right_open_gap_score = -8\naligner.query_right_extend_gap_score = -0.4\naligner.target_left_open_gap_score = -8\naligner.target_left_extend_gap_score = -0.4\naligner.target_right_open_gap_score = -8\naligner.target_right_extend_gap_score = -0.4\n\ndef find_similar_sequences(query_seq, train_seqs_df, train_coords_dict, top_n=5):\n    similar_seqs = []\n    \n    # Pre-filter: Iterate only valid targets\n    # Note: aligner.score is much faster than generating full alignments\n    for _, row in train_seqs_df.iterrows():\n        target_id, train_seq = row['target_id'], row['sequence']\n        if target_id not in train_coords_dict: continue\n        \n        # Length filter (keep your original logic)\n        if abs(len(train_seq) - len(query_seq)) / max(len(train_seq), len(query_seq)) > 0.3: continue\n        \n        # FAST SCORE: Calculates score without traceback overhead\n        raw_score = aligner.score(query_seq, train_seq)\n        \n        normalized_score = raw_score / (2 * min(len(query_seq), len(train_seq)))\n        similar_seqs.append((target_id, train_seq, normalized_score, train_coords_dict[target_id]))\n    \n    similar_seqs.sort(key=lambda x: x[2], reverse=True)\n    return similar_seqs[:top_n]\n\ndef adapt_template_to_query(query_seq, template_seq, template_coords):\n    # Generate the alignment object\n    # aligner.align returns an iterator; we take the first optimal alignment\n    alignment = next(iter(aligner.align(query_seq, template_seq)))\n    \n    new_coords = np.full((len(query_seq), 3), np.nan)\n    \n    # VECTORIZED MAPPING:\n    # alignment.aligned returns lists of (start, end) tuples for matched segments.\n    # This avoids the slow python loop \"for char_q, char_t in zip...\"\n    for (q_start, q_end), (t_start, t_end) in zip(*alignment.aligned):\n        # Map the coordinate chunk directly\n        t_chunk = template_coords[t_start:t_end]\n        \n        # Safety check to ensure shapes match (handles edge cases)\n        if len(t_chunk) == (q_end - q_start):\n            new_coords[q_start:q_end] = t_chunk\n\n    # --- Interpolation Logic (Unchanged) ---\n    for i in range(len(new_coords)):\n        if np.isnan(new_coords[i, 0]):\n            prev_v = next((j for j in range(i-1, -1, -1) if not np.isnan(new_coords[j, 0])), -1)\n            next_v = next((j for j in range(i+1, len(new_coords)) if not np.isnan(new_coords[j, 0])), -1)\n            if prev_v >= 0 and next_v >= 0:\n                w = (i - prev_v) / (next_v - prev_v)\n                new_coords[i] = (1-w)*new_coords[prev_v] + w*new_coords[next_v]\n            elif prev_v >= 0: new_coords[i] = new_coords[prev_v] + [3, 0, 0]\n            elif next_v >= 0: new_coords[i] = new_coords[next_v] + [3, 0, 0]\n            else: new_coords[i] = [i*3, 0, 0]\n            \n    return np.nan_to_num(new_coords)\n\ndef adaptive_rna_constraints(coordinates, target_id, confidence=1.0, passes=2):\n    \"\"\"\n    Evaluation-driven constraints:\n    - US-align is show-only rigid body => internal geometry errors are fatal\n    - apply within each chain segment (no fake bond across chain breaks)\n    \"\"\"\n    coords = coordinates.copy()\n    segments = test_segs_map.get(target_id, [(0, len(coords))])\n\n    # stronger corrections when confidence is low\n    strength = 0.75 * (1.0 - min(confidence, 0.97))\n    strength = max(strength, 0.02)\n\n    for _ in range(passes):\n        for (s, e) in segments:\n            X = coords[s:e]\n            L = e - s\n            if L < 3:\n                coords[s:e] = X\n                continue\n\n            # (1) bond i,i+1 to ~5.95Å (vectorized, symmetric)\n            d = X[1:] - X[:-1]\n            dist = np.linalg.norm(d, axis=1) + 1e-6\n            target = 5.95\n            scale = (target - dist) / dist\n            adj = (d * scale[:, None]) * (0.22 * strength)\n            X[:-1] -= adj\n            X[1:]  += adj\n\n            # (2) soft i,i+2 to ~10.2Å (vectorized, symmetric)\n            d2 = X[2:] - X[:-2]\n            dist2 = np.linalg.norm(d2, axis=1) + 1e-6\n            target2 = 10.2\n            scale2 = (target2 - dist2) / dist2\n            adj2 = (d2 * scale2[:, None]) * (0.10 * strength)\n            X[:-2] -= adj2\n            X[2:]  += adj2\n\n            # (3) Laplacian smoothing (removes kinks US-align cannot fix)\n            lap = 0.5 * (X[:-2] + X[2:]) - X[1:-1]\n            X[1:-1] += (0.06 * strength) * lap\n\n            # (4) light self-avoidance (prevents steric collapse)\n            if L >= 25:\n                k = min(L, 160) if L > 220 else L\n                if k < L:\n                    idx = np.linspace(0, L - 1, k).astype(int)\n                else:\n                    idx = np.arange(L)\n\n                P = X[idx]\n                diff = P[:, None, :] - P[None, :, :]\n                distm = np.linalg.norm(diff, axis=2) + 1e-6\n                sep = np.abs(idx[:, None] - idx[None, :])\n\n                mask = (sep > 2) & (distm < 3.2)\n                if np.any(mask):\n                    force = (3.2 - distm) / distm\n                    vec = (diff * force[:, :, None] * mask[:, :, None]).sum(axis=1)\n                    X[idx] += (0.015 * strength) * vec\n\n            coords[s:e] = X\n\n    return coords\n\ndef _rotmat(axis, ang):\n    axis = np.asarray(axis, float)\n    axis = axis / (np.linalg.norm(axis) + 1e-12)\n    x, y, z = axis\n    c, s = np.cos(ang), np.sin(ang)\n    C = 1.0 - c\n    return np.array([\n        [c + x*x*C,     x*y*C - z*s, x*z*C + y*s],\n        [y*x*C + z*s,   c + y*y*C,   y*z*C - x*s],\n        [z*x*C - y*s,   z*y*C + x*s, c + z*z*C]\n    ], dtype=float)\n\ndef apply_hinge(coords, seg, rng, max_angle_deg=25):\n    s, e = seg\n    L = e - s\n    if L < 30:\n        return coords\n    pivot = s + int(rng.integers(10, L - 10))\n    axis = rng.normal(size=3)\n    ang = np.deg2rad(float(rng.uniform(-max_angle_deg, max_angle_deg)))\n    R = _rotmat(axis, ang)\n    X = coords.copy()\n    p0 = X[pivot].copy()\n    X[pivot+1:e] = (X[pivot+1:e] - p0) @ R.T + p0\n    return X\n\ndef jitter_chains(coords, segments, rng, max_angle_deg=12, max_trans=1.5):\n    X = coords.copy()\n    global_center = X.mean(axis=0, keepdims=True)\n    for (s, e) in segments:\n        axis = rng.normal(size=3)\n        ang = np.deg2rad(float(rng.uniform(-max_angle_deg, max_angle_deg)))\n        R = _rotmat(axis, ang)\n        shift = rng.normal(size=3)\n        shift = shift / (np.linalg.norm(shift) + 1e-12) * float(rng.uniform(0.0, max_trans))\n        c = X[s:e].mean(axis=0, keepdims=True)\n        X[s:e] = (X[s:e] - c) @ R.T + c + shift\n    # recenter\n    X -= X.mean(axis=0, keepdims=True) - global_center\n    return X\n\ndef smooth_wiggle(coords, segments, rng, amp=0.8):\n    X = coords.copy()\n    for (s, e) in segments:\n        L = e - s\n        if L < 20:\n            continue\n        n_ctrl = 6\n        ctrl_x = np.linspace(0, L - 1, n_ctrl)\n        ctrl_disp = rng.normal(0, amp, size=(n_ctrl, 3))\n        t = np.arange(L)\n        disp = np.vstack([np.interp(t, ctrl_x, ctrl_disp[:, k]) for k in range(3)]).T\n        X[s:e] += disp\n    return X\n\ndef predict_rna_structures(row, train_seqs_df, train_coords_dict, n_predictions=5):\n    tid = row['target_id']\n    seq = row['sequence']\n\n    # Data constraint: should already be canonical A/C/G/U\n    assert set(seq).issubset(set(\"ACGU\")), f\"Non-ACGU in {tid}; do not remap here.\"\n\n    segments = test_segs_map.get(tid, [(0, len(seq))])\n\n    # Grab a larger candidate pool, then sample for diversity (best-of-5)\n    cands = find_similar_sequences(query_seq=seq, train_seqs_df=train_seqs_df, train_coords_dict=train_coords_dict, top_n=30)\n    assert all(len(c[3]) == len(c[1]) for c in cands), \"Template coords/seq length mismatch\"\n    predictions = []\n    used = set()\n\n    for i in range(n_predictions):\n        seed = (abs(hash(tid)) + i * 10007) % (2**32)\n        rng = np.random.default_rng(seed)\n\n        if not cands:\n            # hard fallback (straight line per chain)\n            coords = np.zeros((len(seq), 3), dtype=float)\n            for (s, e) in segments:\n                for j in range(s+1, e):\n                    coords[j] = coords[j-1] + [5.95, 0, 0]\n            predictions.append(coords)\n            continue\n\n        # Choose template:\n        # i=0 => best template; others => sample among top-K with weights, avoid duplicates\n        if i == 0:\n            t_id, t_seq, sim, t_coords = cands[0]\n        else:\n            K = min(12, len(cands))\n            sims = np.array([cands[k][2] for k in range(K)], float)\n            w = np.exp((sims - sims.max()) / 0.08)\n            # penalize already used templates\n            for k in range(K):\n                if cands[k][0] in used:\n                    w[k] *= 0.10\n            w = w / (w.sum() + 1e-12)\n            k = int(rng.choice(np.arange(K), p=w))\n            t_id, t_seq, sim, t_coords = cands[k]\n\n        used.add(t_id)\n\n        # Transfer coords with diagonal-guard mapping (no sliding)\n        adapted = adapt_template_to_query(query_seq=seq, template_seq=t_seq, template_coords=t_coords)\n\n        # Diversity transforms (then re-refine constraints)\n        if i == 0:\n            X = adapted\n        elif i == 1:\n            # mild noise\n            X = adapted + rng.normal(0, max(0.01, (0.40 - sim) * 0.06), adapted.shape)\n        elif i == 2:\n            # hinge within the longest chain\n            longest = max(segments, key=lambda se: se[1] - se[0])\n            X = apply_hinge(adapted, longest, rng, max_angle_deg=22)\n        elif i == 3:\n            # inter-chain jitter (small, safe)\n            X = jitter_chains(adapted, segments, rng, max_angle_deg=10, max_trans=1.0)\n        else:\n            # smooth low-frequency deformation\n            X = smooth_wiggle(adapted, segments, rng, amp=0.7)\n\n        refined = adaptive_rna_constraints(X, tid, confidence=sim, passes=2)\n        predictions.append(refined)\n\n    return predictions\n\nall_predictions = []\nstart_time = time.time()\nfor idx, row in test_seqs.iterrows():\n    if idx % 10 == 0: print(f\"Processing {idx} | {time.time()-start_time:.1f}s\")\n    tid, seq = row['target_id'], row['sequence']\n    preds = predict_rna_structures(row, train_seqs, train_coords_dict)\n    for j in range(len(seq)):\n        res = {'ID': f\"{tid}_{j+1}\", 'resname': seq[j], 'resid': j+1}\n        for i in range(5):\n            res[f'x_{i+1}'], res[f'y_{i+1}'], res[f'z_{i+1}'] = preds[i][j]\n        all_predictions.append(res)\n\nsub = pd.DataFrame(all_predictions)\ncols = ['ID', 'resname', 'resid'] + [f'{c}_{i}' for i in range(1,6) for c in ['x','y','z']]\n\n# Safety: competition clips coords; do it explicitly to avoid out-of-range explosions\ncoord_cols = [c for c in cols if c.startswith(('x_','y_','z_'))]\nsub[coord_cols] = sub[coord_cols].clip(-999.999, 9999.999)\n\npred_tbm = sub\n\n\n# print('Submission.csv shape:', rosetta_pred.shape)\npred_tbm.to_csv('/kaggle/working/pred_tbm.csv', index=False)\ndisplay(pred_tbm.head())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# SECTION 4: RNApro Inference (seq_len < 1000, uses TBM templates)\n# ============================================================\n\n!cp -r /kaggle/input/datasets/theoviel/rnapro-src/RNAPro .\n!cp /kaggle/input/datasets/theoviel/rnapro-src/rnapro-private-best-500m.ckpt .\n\n\n%cd /kaggle/working/RNAPro\n\n\n!pip install -e . --no-deps\n\n\nDIST = \"/kaggle/working/RNAPro/release_data/ccd_cache/\"\n!mkdir -p $DIST\n\n\n# updated file paths\n!cp /kaggle/input/datasets/jaejohn/rnapro-ccd-cache/ccd_cache/components.cif $DIST\n!cp /kaggle/input/datasets/jaejohn/rnapro-ccd-cache/ccd_cache/components.cif.rdkit_mol.pkl $DIST\n\n\n!pip install /kaggle/input/datasets/tobimichigan/biotite-1-2/biotite-1.2.0-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl --no-deps\n\n\n!python preprocess/convert_templates_to_pt_files.py --input_csv /kaggle/working/pred_tbm.csv --output_name templates.pt\n\n\nIS_SCORING_RUN = os.environ.get('KAGGLE_IS_COMPETITION_RERUN')\nprint(IS_SCORING_RUN)\n\n\n# %%python\ndf = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv\")\ndf[\"sequence_len\"] = df[\"sequence\"].str.len()\n# df = df[(df['sequence_len'] >= 200) & (df['sequence_len'] < 600)]\n# df = df[df['sequence_len'] < 600]\ndf = df[df['sequence_len'] < 1000]\ndf.reset_index(drop=True, inplace=True)\nprint(df.shape)\nif not IS_SCORING_RUN:\n    # df = df.head(5)\n    TEST_TARGETS = \"8ZNQ,9ZCC\"\n    keep = {t.strip() for t in TEST_TARGETS.split(\",\") if t.strip()}\n    df = df[df[\"target_id\"].isin(keep)].reset_index(drop=True)\n    print(df.shape)\n    if df.empty:\n        raise ValueError(\"TEST_TARGETS did not match any target_id values\")\n\ndf.to_csv('/kaggle/working/sample_sequences.csv', index=False)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile rnapro_inference_kaggle.sh\n\n%export LAYERNORM_TYPE=torch # fast_layernorm, torch\n\n\n# Inference parameters (RNAPro)\nSEED=42\nN_SAMPLE=1\nN_STEP=200\nN_CYCLE=10\n\n# Paths\nDUMP_DIR=\"../output\"\n# Set a valid checkpoint file path below\nCHECKPOINT_PATH=\"../rnapro-private-best-500m.ckpt\"\n\n# Template/MSA settings\nTEMPLATE_DATA=\"./release_data/kaggle/templates.pt\"\n# Note: template_idx supports 5 choices and maps to top-k:\n# 0->top1, 1->top2, 2->top3, 3->top4, 4->top5\nTEMPLATE_IDX=0\nRNA_MSA_DIR=\"/kaggle/input/stanford-rna-3d-folding-2/MSA\"\n\n# SEQUENCES_CSV=\"/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv\"\nSEQUENCES_CSV=\"/kaggle/working/sample_sequences.csv\"\n\n# RibonanzaNet2 path (keep as-is per request)\nRIBONANZA_PATH=\"/kaggle/input/models/shujun717/ribonanzanet2/pytorch/alpha/1\"\n\n# Model selection: keep to an existing key to align defaults (N_step=200, N_cycle=10)\nMODEL_NAME=\"rnapro_base\"\nmkdir -p \"${DUMP_DIR}\"\n\npython3 runner/inference_kaggle.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","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile runner/inference_kaggle.py\n# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.\n# SPDX-License-Identifier: Apache-2.0\n\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\nimport gc\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__)\n\nlogging.basicConfig(level=logging.INFO)\nlogging.getLogger(\"rnapro.data\").setLevel(logging.INFO)\nlogging.getLogger(\"rnapro\").setLevel(logging.INFO)\n\n\ndef parse_configs(configs: dict, arg_str: str = None, fill_required_with_null: bool = 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=1000, 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    merged_configs = manager.merge_configs(vars(parser.parse_args(arg_str.split())) if arg_str else {})\n    merged_configs.max_len = parser.parse_args(arg_str.split()).max_len\n    return merged_configs\n\nclass dotdict(dict):\n    __setattr__ = dict.__setitem__\n    __delattr__ = dict.__delitem__\n    def __getattr__(self, name):\n        try: return self[name]\n        except KeyError: raise AttributeError(name)\n\nclass InferenceRunner(object):\n    def __init__(self, configs: Any) -> None:\n        self.configs = configs\n        self.init_env()\n        self.init_basics()\n        self.init_model()\n        self.load_checkpoint()\n        self.init_dumper(need_atom_confidence=configs.need_atom_confidence, sorted_by_ranking_score=configs.sorted_by_ranking_score)\n\n    def init_env(self) -> None:\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) -> None:\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) -> None:\n        self.model = RNAPro(self.configs).to(self.device)\n\n    def load_checkpoint(self) -> None:\n        checkpoint_path = self.configs.load_checkpoint_path\n        checkpoint = torch.load(checkpoint_path, self.device)\n        sample_key = [k for k in 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(state_dict=checkpoint[\"model\"], strict=True)\n        self.model.eval()\n\n    def init_dumper(self, need_atom_confidence: bool = False, sorted_by_ranking_score: bool = True):\n        self.dumper = DataDumper(base_dir=self.dump_dir, need_atom_confidence=need_atom_confidence, sorted_by_ranking_score=sorted_by_ranking_score)\n\n    @torch.no_grad()\n    def predict(self, data: Mapping[str, Mapping[str, Any]]) -> dict[str, torch.Tensor]:\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(input_feature_dict=data[\"input_feature_dict\"], label_full_dict=None, label_dict=None, mode=\"inference\")\n        return prediction\n\n    def update_model_configs(self, new_configs: Any) -> None:\n        self.model.configs = new_configs\n\n\ndef update_inference_configs(configs: Any, N_token: int):\n    # Force enable AMP (skip_amp=False) to save memory\n    configs.skip_amp.confidence_head = False\n    configs.skip_amp.sample_diffusion = False\n    return configs\n\n\ndef infer_predict(runner: InferenceRunner, configs: Any) -> None:\n    try:\n        dataloader = get_inference_dataloader(configs=configs)\n    except Exception as e:\n        logger.error(f\"Dataloader error: {e}\")\n        traceback.print_exc()\n        return\n\n    # --- Patch for chunked inference ---\n    if getattr(configs, \"chunk_target_id\", None) and getattr(configs, \"chunk_span\", None):\n        chunk_id = configs.chunk_target_id\n        original_id = chunk_id.split(\"_chk\")[0]\n        start, end = configs.chunk_span\n        dataset = dataloader.dataset\n        if hasattr(dataset, \"template_features\"):\n            if original_id in dataset.template_features:\n                orig_feats = dataset.template_features[original_id]\n                # Slice template features\n                chunk_feats = {}\n                logger.info(f\"Processing template features for {chunk_id} (span: {start}-{end})\")\n                \n                for k, v in orig_feats.items():\n                    if isinstance(v, (torch.Tensor, np.ndarray)):\n                        try:\n                            shape = v.shape\n                            logger.info(f\"  Feature '{k}': shape={shape}\")\n                            \n                            # Case 1: (N, L, L, ...)\n                            if len(shape) >= 3 and shape[1] == shape[2] and shape[1] >= end:\n                                chunk_feats[k] = v[:, start:end, start:end]\n                                logger.info(f\"    -> Sliced dim 1,2: {chunk_feats[k].shape}\")\n                            # Case 2: (L, L, ...)\n                            elif len(shape) >= 2 and shape[0] == shape[1] and shape[0] >= end:\n                                chunk_feats[k] = v[start:end, start:end]\n                                logger.info(f\"    -> Sliced dim 0,1: {chunk_feats[k].shape}\")\n                            # Case 3: (N, L, ...)\n                            elif len(shape) >= 2 and shape[1] >= end:\n                                chunk_feats[k] = v[:, start:end]\n                                logger.info(f\"    -> Sliced dim 1: {chunk_feats[k].shape}\")\n                            # Case 4: (L, ...)\n                            elif len(shape) >= 1 and shape[0] >= end:\n                                chunk_feats[k] = v[start:end]\n                                logger.info(f\"    -> Sliced dim 0: {chunk_feats[k].shape}\")\n                            else:\n                                chunk_feats[k] = v\n                                logger.info(f\"    -> Kept original\")\n                        except IndexError as e:\n                            logger.warning(f\"IndexError slicing template feature {k} for {chunk_id}: {e}\")\n                            chunk_feats[k] = v\n                    else:\n                        logger.info(f\"  Feature '{k}': type={type(v)} (not sliced)\")\n                        chunk_feats[k] = v\n                \n                dataset.template_features[chunk_id] = chunk_feats\n                logger.info(f\"Mapped and sliced template for {chunk_id} from {original_id} ({start}:{end})\")\n            else:\n                logger.warning(f\"Original template {original_id} not found for chunk {chunk_id}\")\n\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: 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=\"\", pdb_id=sample_name, seed=seed,\n                    pred_dict=prediction, atom_array=atom_array, entity_poly_type=data[\"entity_poly_type\"]\n                )\n                \n                # Explicitly cleanup memory\n                del prediction\n                del data\n                torch.cuda.empty_cache()\n                gc.collect()\n                \n            except Exception as e:\n                logger.error(f\"Inference error for {sample_name}: {e}\")\n                traceback.print_exc()\n                if hasattr(torch.cuda, \"empty_cache\"): torch.cuda.empty_cache()\n                gc.collect()\n\n\ndef make_dummy_solution(valid_df):\n    solution = dotdict()\n    for i, row in valid_df.iterrows():\n        solution[row.target_id] = dotdict(target_id=row.target_id, sequence=row.sequence, coord=[])\n    return solution\n\ndef solution_to_submit_df(solution):\n    submit_df = []\n    for k, s in solution.items():\n        L = len(s.sequence)\n        df = pd.DataFrame()\n        df[\"ID\"] = [f\"{s.target_id}_{i + 1}\" for i in range(L)]\n        df[\"resname\"] = list(s.sequence)\n        df[\"resid\"] = [i + 1 for i in range(L)]\n        for j in range(len(s.coord)):\n            df[f\"x_{j+1}\"] = s.coord[j][:, 0]\n            df[f\"y_{j+1}\"] = s.coord[j][:, 1]\n            df[f\"z_{j+1}\"] = s.coord[j][:, 2]\n        submit_df.append(df)\n    return pd.concat(submit_df)\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        mask_c1 = np.char.strip(atom_array.atom_name.astype(str)) == \"C1'\"\n        c1_atoms = atom_array[mask_c1]\n        sort_indices = np.argsort(c1_atoms.res_id)\n        return c1_atoms[sort_indices].coord\n    except Exception as e:\n        return None\n\n# ==========================================\n# 3D Assembly (Kabsch & Stitching) Functions\n# ==========================================\ndef kabsch_umeyama(A, B):\n    \"\"\" Kabsch algorithm to align chunk B onto chunk A \"\"\"\n    centroid_A = np.mean(A, axis=0)\n    centroid_B = np.mean(B, axis=0)\n    AA = A - centroid_A\n    BB = B - centroid_B\n    H = BB.T @ AA\n    U, S, Vt = np.linalg.svd(H)\n    R = Vt.T @ U.T\n    if np.linalg.det(R) < 0:\n        Vt[-1, :] *= -1\n        R = Vt.T @ U.T\n    t = centroid_A - R @ centroid_B\n    return R, t\n\ndef stitch_chunks(chunks_coords, overlaps):\n    \"\"\"Stitch chunk coordinates using Kabsch alignment and linear interpolation over overlap regions.\"\"\"\n    final_coords = chunks_coords[0].copy()\n\n    for i in range(1, len(chunks_coords)):\n        curr_coords = chunks_coords[i]\n        overlap = overlaps[i-1]\n\n        A = final_coords[-overlap:]  # Tail portion of the previous chunk\n        B = curr_coords[:overlap]    # Head portion of the current chunk\n\n        # Compute rotation matrix and translation vector via Kabsch alignment\n        R, t = kabsch_umeyama(A, B)\n\n        # Apply the transformation to the entire chunk\n        curr_coords_transformed = (curr_coords @ R.T) + t\n\n        # Linearly interpolate over the overlap region (smoothing)\n        weights = np.linspace(1, 0, overlap).reshape(-1, 1)\n        blended_overlap = A * weights + curr_coords_transformed[:overlap] * (1 - weights)\n\n        # Update final coordinates: overwrite overlap region and append the rest\n        final_coords[-overlap:] = blended_overlap\n        final_coords = np.vstack([final_coords, curr_coords_transformed[overlap:]])\n\n    return final_coords\n\n# ==========================================\n# Modified runner processing\n# ==========================================\ndef run_ptx(target_id, sequence, configs, template_idx, runner, chunk_span=None):\n    temp_dir = f\"./{configs.dump_dir}/input\"\n    os.makedirs(temp_dir, exist_ok=True)\n\n    input_json = [{\"sequences\": [{\"rnaSequence\": {\"sequence\": sequence, \"count\": 1}}], \"name\": target_id}]\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    configs.input_json_path = input_json_path\n    configs.template_idx = int(template_idx)\n    \n    # If target_id contains \"_chk\", pass it as chunk_target_id to trigger mapping in infer_predict\n    if \"_chk\" in target_id:\n        configs.chunk_target_id = target_id\n        configs.chunk_span = chunk_span # Tuple (start, end)\n    else:\n        configs.chunk_target_id = None\n        configs.chunk_span = None\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\n    if coord is None:\n        coord = np.zeros((len(sequence), 3), dtype=np.float32)\n    elif coord.shape[0] < len(sequence):\n        pad = np.zeros((len(sequence) - coord.shape[0], 3), dtype=np.float32)\n        coord = np.concatenate([coord, pad], axis=0)\n\n    return coord\n\n\ndef run() -> None:\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    valid_df = pd.read_csv(configs.sequences_csv)\n    runner = InferenceRunner(configs)\n    solution = make_dummy_solution(valid_df)\n\n    chunk_size = 500 # configs.max_len\n    overlap_size = 100 # number of overlapping bases between chunks\n\n    for idx, row in valid_df.iterrows():\n        target_id = row.target_id\n        sequence = row.sequence\n        L = len(sequence)\n        print(f\"\\n -> Processing {target_id} (Length: {L})\")\n\n        if L < 1000:\n            # Short sequences: run inference directly without chunking\n            for template_idx in range(2):\n                coord = run_ptx(target_id, sequence, configs, template_idx, runner, chunk_span=(0, L))\n                solution[target_id].coord.append(coord)\n        else:\n            # Long sequences: split into overlapping chunks for inference\n            print(f\"    Sequence exceeds max_len ({chunk_size}), triggering chunked prediction...\")\n            starts, ends = [], []\n            start = 0\n            while start < L:\n                end = min(start + chunk_size, L)\n                starts.append(start)\n                ends.append(end)\n                if end == L: break\n                start = end - overlap_size\n\n            for template_idx in range(1):\n                chunk_coords = []\n                for i in range(len(starts)):\n                    c_seq = sequence[starts[i]:ends[i]]\n                    c_id = f\"{target_id}_chk{i}\"\n                    print(f\"    -> Predicting chunk {i+1}/{len(starts)} ({starts[i]} to {ends[i]})\")\n\n                    # Run inference for each chunk\n                    coord = run_ptx(c_id, c_seq, configs, template_idx, runner, chunk_span=(starts[i], ends[i]))\n                    chunk_coords.append(coord)               \n\n                # Compute overlap sizes and stitch chunks together\n                overlaps = [ends[i-1] - starts[i] for i in range(1, len(starts))]\n                final_coord = stitch_chunks(chunk_coords, overlaps)\n                solution[target_id].coord.append(final_coord)\n\n    print('\\n\\n -> Inference done! Saving to rnapro_submission.csv')\n    submit_df = solution_to_submit_df(solution)\n    submit_df = submit_df.fillna(0.0)\n    submit_df.to_csv(\"/kaggle/working/rnapro_submission.csv\", index=False)\n\nif __name__ == \"__main__\":\n    run()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# run\n!bash ./rnapro_inference_kaggle.sh\n\npred_rnapro = pd.read_csv(\"/kaggle/working/rnapro_submission.csv\")\ndisplay(pred_rnapro.head())\n\n%cd /kaggle/working","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# SECTION 5: Boltz2 Inference (seq_len < 1000)\n# ============================================================\n\nsubprocess.run([\"cp\", \"-r\", \"/kaggle/input/datasets/lbugnon/boltz-src-minimal\", \"./\"], check=False)\n# Install boltz from src\nsubprocess.run([\"pip\", \"install\", \"--no-index\", \"--no-build-isolation\", \"-e\", \"./boltz-src-minimal\"], check=False)\n\nsubprocess.run([\"pip\", \"install\", \"--no-index\", \"/kaggle/input/datasets/youhanlee/boltz-dependencies/mashumaro-3.14-py3-none-any.whl\", \"--no-deps\"], check=False)\nsubprocess.run([\"pip\", \"install\", \"--no-index\", \"/kaggle/input/datasets/youhanlee/boltz-dependencies/ihm-2.2-py3-none-any.whl\", \"--no-deps\"], check=False)\nsubprocess.run([\"pip\", \"install\", \"--no-index\", \"/kaggle/input/datasets/youhanlee/boltz-dependencies/modelcif-1.3-py3-none-any.whl\", \"--no-deps\"], check=False)\n\n# need to create tar so boltz does not try to download\nos.makedirs(\"boltz_cache\", exist_ok=True)\nsubprocess.run([\"cp\", \"-r\", \"/kaggle/input/datasets/lbugnon/boltz2\", \"boltz_cache\"], check=False)\n!mv boltz_cache/boltz2/mols/mols boltz_cache/boltz2/mols_new\n!rm -rf boltz_cache/boltz2/mols\n!mv boltz_cache/boltz2/mols_new boltz_cache/boltz2/mols\nshutil.rmtree(\"boltz_cache/boltz2/mols/mols\", ignore_errors=True)\nsubprocess.run([\"tar\", \"-cf\", \"boltz_cache/boltz2/mols.tar\", \"boltz_cache/boltz2/mols\"], check=False)\n\nMAX_LENGTH_BOLTZ = 1000\n\nsequences = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv\")\nprint(sequences.shape)\n\nsequences[\"sequence_len\"] = sequences[\"sequence\"].str.len()\n\nsequences = sequences[sequences['sequence_len'] < 1000]\nsequences.reset_index(drop=True, inplace=True)\n\nif not IS_SCORING_RUN:\n  sequences = sequences.head(1)\n\nprint(sequences.shape)\n\nprint(max(sequences[\"sequence_len\"].values))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile /kaggle/working/patch_torch.py\nimport torch\nimport functools\n\n# Monkey-patch torch.load to override safety settings\noriginal_load = torch.load\n\n@functools.wraps(original_load)\ndef patched_load(*args, **kwargs):\n    # Force weights_only=False even when the library explicitly passes True\n    kwargs['weights_only'] = False\n    return original_load(*args, **kwargs)\n\n# Apply the patch\ntorch.load = patched_load\n\nprint(\"✅ PyTorch Safety Check FORCE DISABLED (Overriding explicit arguments)\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Build YAML input files for Boltz2\ndef create_boltz_yaml(row, output_dir):\n\n    data = {\n        \"id\": str(row[\"target_id\"]),\n        \"sequences\": [\n            {\n                \"rna\": {\n                    \"id\": str(row[\"stoichiometry\"]).split(\":\")[0],\n                    \"sequence\": str(row[\"sequence\"])\n                }\n            }\n        ]\n    }\n\n    # --- Ligand processing ---\n    if pd.notna(row[\"ligand_SMILES\"]):\n\n        smiles_list = str(row[\"ligand_SMILES\"]).split(\";\")\n        ids_list = str(row[\"ligand_ids\"]).split(\";\")\n\n        ligands = []\n        for s, i in zip(smiles_list, ids_list):\n            ligands.append({\n                \"id\": i.strip(),\n                \"smiles\": s.strip()\n            })\n\n        if ligands:\n            data[\"ligands\"] = ligands\n\n    # --- Write output ---\n    file_path = Path(output_dir) / f\"{row['target_id']}.yaml\"\n\n    with open(file_path, \"w\", encoding=\"utf-8\") as f:\n        yaml.safe_dump(\n            data,\n            f,\n            sort_keys=False,\n            allow_unicode=True\n        )\n\n    return file_path\n\n# Create YAML inputs for all test sequences\noutput_path = \"/kaggle/working/inputs\"\nos.makedirs(output_path, exist_ok=True)\n\nfor _, row in sequences.iterrows():\n    yaml_file = create_boltz_yaml(row, output_path)\nprint(f\"Created\")\n\n\nfor p in Path(\"/kaggle/working\").glob(\"boltz_repeat_0*\"):\n    shutil.rmtree(p, ignore_errors=True)\n\nos.makedirs(\"inputs\", exist_ok=True)\n\nmax_tokens = 1000\nmax_repeats = 999\npred_repeats = 1\n\n\n# Monkey-patch PyTorch Lightning Trainer to suppress logger errors in Kaggle\noriginal_trainer_init = pl.Trainer.__init__\n\ndef patched_trainer_init(self, *args, **kwargs):\n    # Disable logger to prevent errors in Kaggle environment\n    kwargs[\"logger\"] = False\n    # Minimize log output beyond the progress bar\n    kwargs[\"enable_checkpointing\"] = False\n    return original_trainer_init(self, *args, **kwargs)\n\n# Apply the patch\npl.Trainer.__init__ = patched_trainer_init\n\nfor k in range(len(sequences)):\n\n    t0 = time.time()\n    name = sequences.iloc[k].target_id\n    print(f\"preparing {name} ({k+1} of {len(sequences)})\")\n\n    pred_seq = sequences.iloc[k].sequence\n\n    # ignore super large seqs\n    if len(sequences.iloc[k].sequence)>max_tokens:\n        print(f\"skipped {name}\")\n        continue\n\n    # A for the chain to predict. First load all chains.\n    try:\n        chains, chain_ind = {}, 1\n        for entry in sequences.iloc[k].all_sequences.split(\"\\n\"):\n            if entry.startswith(\">\"):\n                # get number of repeats\n                repeats = min(len(entry.split(\"|\")[1].split(\",\")), max_repeats)\n            else:\n                if entry == pred_seq:\n                    chains[0] = entry, \"rna\"\n\n        # Now keep the chain 0 plus all chains randomly picked to fit max_tokens\n        filtered_chains, ntokens = {}, 0\n        other_chains = [c for c in chains if c!=0]\n        random.shuffle(other_chains)\n        for c in [0] + other_chains:\n            ntokens += len(chains[c][0])\n            if ntokens >= max_tokens:\n                break\n            filtered_chains[c] = chains[c]\n        chains = filtered_chains\n    except:\n        # use only target sequence\n        chains = {0: (pred_seq, \"rna\")}\n\n    # save yaml\n    yaml_data = {\n        \"id\": name,\n        \"sequences\": []\n    }\n\n    # --- sequences ---\n    for chain in sorted(chains.keys()):\n        seq, seq_type = chains[chain]\n        yaml_data[\"sequences\"].append({\n            seq_type: {\n                \"id\": chr(ord(\"A\") + chain),\n                \"sequence\": seq\n            }\n        })\n\n    # --- Ligands (only if present) ---\n    row = sequences.iloc[k]\n\n    if pd.notna(row[\"ligand_SMILES\"]):\n        smiles_list = str(row[\"ligand_SMILES\"]).split(\";\")\n        ids_list = str(row[\"ligand_ids\"]).split(\";\")\n\n        ligands = []\n        for s, i in zip(smiles_list, ids_list):\n            s = s.strip()\n            i = i.strip()\n            if s:\n                ligands.append({\n                    \"id\": i,\n                    \"smiles\": s\n                })\n\n        if ligands:\n            yaml_data[\"ligands\"] = ligands\n\n    # --- save ---\n    with open(f\"/kaggle/working/inputs/{name}.yaml\", \"w\") as fout:\n        yaml.safe_dump(yaml_data, fout, sort_keys=False)\n\n    # print yaml content\n    with open(f\"/kaggle/working/inputs/{name}.yaml\") as _f:\n        print(_f.read())\n\n    # --- setting ---\n    boltz_args = {\n        \"diffusion_samples\": 10,\n        \"diffusion_samples_affinity\": 5,\n        \"devices\": 2,\n        \"num_workers\": 4,\n        \"max_parallel_samples\": 2,\n        \"output_format\": \"pdb\",\n        \"cache\": \"boltz_cache/boltz2/\",\n    }\n\n    # --- inferring ---\n    print(name, \":\")\n    print()\n    for repeat in range(pred_repeats):\n        cmd = (\n            f\"export PYTHONPATH=$PYTHONPATH:/kaggle/working && \"\n            f\"python3 -c 'import patch_torch; from boltz.main import cli; cli()' predict /kaggle/working/inputs/{name}.yaml\"\n            f\" --diffusion_samples {boltz_args['diffusion_samples']}\"\n            f\" --diffusion_samples_affinity {boltz_args['diffusion_samples_affinity']}\"\n            f\" --devices {boltz_args['devices']}\"\n            f\" --num_workers {boltz_args['num_workers']}\"\n            f\" --max_parallel_samples {boltz_args['max_parallel_samples']}\"\n            f\" --output_format {boltz_args['output_format']}\"\n            f\" --cache {boltz_args['cache']}\"\n            f\" --out_dir boltz_repeat_{repeat}\"\n        )\n        \n        # Log prefix to distinguish runs\n        print(f\"Running with PyTorch 2.6 patch...\")\n        subprocess.run(cmd, shell=True, check=False)\n        gc.collect()\n        torch.cuda.empty_cache()\n\n\n# ==========================================\n# 1. Coordinate extraction\n# ==========================================\ndef extract_coords(pdb_file, cif_file):\n    \"\"\"Extract C1' atom coordinates from mmCIF file.\"\"\"\n    try:\n        structure = gemmi.read_structure(f\"{pdb_file}\")\n        structure.make_mmcif_document().write_file(f\"{cif_file}\")\n        mmcif_dict = MMCIF2Dict(cif_file)\n        x_coords = mmcif_dict[\"_atom_site.Cartn_x\"]\n        y_coords = mmcif_dict[\"_atom_site.Cartn_y\"]\n        z_coords = mmcif_dict[\"_atom_site.Cartn_z\"]\n        atom_names = mmcif_dict[\"_atom_site.label_atom_id\"]\n\n        c1_coords = []\n        for i, atom in enumerate(atom_names):\n            if atom == \"C1'\":\n                c1_coords.append([float(x_coords[i]), float(y_coords[i]), float(z_coords[i])])\n        return np.array(c1_coords)\n    except Exception as e:\n        print(f\"Error parsing {cif_file}: {e}\")\n        return None\n\n# ==========================================\n# 2. Model selection logic\n# ==========================================\n\ndef get_top_n_predictions(diffusion_results_dir: Path, file_prefix: str, n: int = 3):\n    \"\"\"Return top-n predictions sorted by pLDDT and confidence score.\"\"\"\n    results = []\n    target_subdir = diffusion_results_dir / file_prefix\n    \n    for i in range(10):\n        json_path = target_subdir / f\"confidence_{file_prefix}_model_{i}.json\"\n        if not json_path.exists():\n            continue\n\n        with open(json_path, 'r') as f:\n            data = json.load(f)\n\n        results.append({\n            \"idx\": i,\n            \"plddt\": data.get(\"complex_plddt\", -float(\"inf\")),\n            \"conf\": data.get(\"confidence_score\", -float(\"inf\"))\n        })\n\n    # Sort by score (descending) and return top-n\n    results.sort(key=lambda x: (x[\"plddt\"], x[\"conf\"]), reverse=True)\n    return results[:n]\n\n\n# ==========================================\n# 3. Main processing\n# ==========================================\n\ndef main_boltz():\n    df_test = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv\")\n    submission = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/sample_submission.csv')\n\n    # Initialize all coordinate columns as float\n    coord_cols = ['x_1', 'y_1', 'z_1',\n                  'x_2', 'y_2', 'z_2',\n                  'x_3', 'y_3', 'z_3',\n                  'x_4', 'y_4', 'z_4',\n                  'x_5', 'y_5', 'z_5',]\n    submission[coord_cols] = submission[coord_cols].astype(float)\n\n    for target_id, sequence in zip(df_test[\"target_id\"], df_test[\"sequence\"]):\n        print(f'--- Processing {target_id} ---')\n        diffusion_results_dir = Path(f\"/kaggle/working/boltz_repeat_0/boltz_results_{target_id}/predictions\")\n\n        seq_len = len(sequence)\n        top_models = get_top_n_predictions(diffusion_results_dir, target_id, n=5)\n\n        # Mask rows corresponding to this target\n        target_mask = submission['ID'].apply(lambda x: x.startswith(f\"{target_id}_\"))\n        expected_count = target_mask.sum()\n\n        # Store coordinates for each of the 5 predictions\n        for rank, model_info in enumerate(top_models, start=1):\n            m_idx = model_info[\"idx\"]\n            cif_file = diffusion_results_dir / target_id / f\"{target_id}_model_{m_idx}.cif\"\n            pdb_file = diffusion_results_dir / target_id / f\"{target_id}_model_{m_idx}.pdb\"\n            coords = extract_coords(pdb_file, cif_file)\n\n            if coords is not None:\n                # Pad with zeros if too few residues; truncate if too many\n                if len(coords) < expected_count:\n                    coords = np.vstack([coords, np.zeros((expected_count - len(coords), 3))])\n                else:\n                    coords = coords[:expected_count]\n\n                # Assign to coordinate columns (x_N, y_N, z_N)\n                cols = [f'x_{rank}', f'y_{rank}', f'z_{rank}']\n                submission.loc[target_mask, cols] = coords\n                print(f\"  Rank {rank} (Model {m_idx}) loaded.\")\n            else:\n                print(f\"  Warning: Rank {rank} coordinates not found.\")\n\n    # Save output to CSV\n    output_path = '/kaggle/working/boltz_submission.csv'\n    submission.to_csv(output_path, index=False)\n    print(f\"\\nFinal submission saved to {output_path}\")\n\nif __name__ == \"__main__\":\n    main_boltz()\n\n\nboltz_pred = pd.read_csv(\"/kaggle/working/boltz_submission.csv\")\ndisplay(boltz_pred.head())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# SECTION 6: Ensemble & Final Submission\n#\n# Prediction slot assignment by sequence length:\n#   seq_len < 250        : [Boltz2_1, Boltz2_2, RNApro_1, Protenix_1, DRFold2_1]\n#   250 <= seq_len < 1000: [TBM_1,    Boltz2_1, RNApro_1, RNApro_2,   Boltz2_2 ]\n#   seq_len >= 1000      : [TBM_1,    TBM_2,    TBM_3,    Protenix_1, Protenix_2]\n# ============================================================\n\ndef overlay(merged, src_df, src_cols, dst_cols):\n    \"\"\"Write src_cols from src_df into dst_cols of merged (aligned by ID).\"\"\"\n    merged = merged.set_index(\"ID\")\n    merged.loc[src_df[\"ID\"], dst_cols] = src_df[src_cols].values\n    return merged.reset_index(drop=False)\n\n\n# --- Attach sequence_len to each prediction dataframe ---\n# test_seqs already has sequence_len from Section 3\ntest_seqs[\"sequence_len\"] = [len(i) for i in test_seqs[\"sequence\"]]\nseq_len_df = test_seqs[[\"target_id\", \"sequence_len\"]]\n\ndef add_seq_len(df):\n    df = df.copy()\n    df[\"target_id\"] = df[\"ID\"].str.split(\"_\").str[0]\n    return df.merge(seq_len_df, on=\"target_id\", how=\"left\")\n\npred_tbm_      = add_seq_len(pred_tbm)\nboltz_pred_    = add_seq_len(boltz_pred)\nprotenix_pred_ = add_seq_len(protenix_pred)\npred_rnapro_   = pred_rnapro.copy()  # seq_len < 1000 by construction (Section 4 filter)\ndrfold_pred_   = drfold_pred.copy()  # seq_len < 250  by construction (Section 2 filter)\n\nSHORT = 250\n\n# --- Step 1: Base layer (TBM for long seqs, Boltz2 for short seqs) ---\nmerged_pred = pd.concat([\n    pred_tbm_  [pred_tbm_  [\"sequence_len\"] >= SHORT],\n    boltz_pred_[boltz_pred_[\"sequence_len\"] <  SHORT],\n], axis=0)\n\n# --- Step 2: RNApro → slot3 & slot4 (seq_len < 1000 only, by construction) ---\nmerged_pred = overlay(merged_pred, pred_rnapro_,\n                      src_cols=[\"x_1\", \"y_1\", \"z_1\", \"x_2\", \"y_2\", \"z_2\"],\n                      dst_cols=[\"x_3\", \"y_3\", \"z_3\", \"x_4\", \"y_4\", \"z_4\"])\n\n# --- Step 3: DRFold2 → slot5 (seq_len < 250 only, by construction) ---\nmerged_pred = overlay(merged_pred, drfold_pred_,\n                      src_cols=[\"x_1\", \"y_1\", \"z_1\"],\n                      dst_cols=[\"x_5\", \"y_5\", \"z_5\"])\n\n# --- Step 4: Boltz2 → slot2 & slot5 (250 <= seq_len < 1000) ---\nboltz_medium = boltz_pred_[\n    (boltz_pred_[\"sequence_len\"] >= SHORT) & (boltz_pred_[\"sequence_len\"] < 1000)\n]\nmerged_pred = overlay(merged_pred, boltz_medium,\n                      src_cols=[\"x_1\", \"y_1\", \"z_1\", \"x_2\", \"y_2\", \"z_2\"],\n                      dst_cols=[\"x_2\", \"y_2\", \"z_2\", \"x_5\", \"y_5\", \"z_5\"])\n\n# --- Step 5a: Protenix → slot4 (seq_len < 250) ---\nprotenix_short = protenix_pred_[protenix_pred_[\"sequence_len\"] < SHORT]\nmerged_pred = overlay(merged_pred, protenix_short,\n                      src_cols=[\"x_1\", \"y_1\", \"z_1\"],\n                      dst_cols=[\"x_4\", \"y_4\", \"z_4\"])\n\n# --- Step 5b: Protenix → slot4 & slot5 (seq_len >= 1000) ---\nprotenix_long = protenix_pred_[protenix_pred_[\"sequence_len\"] >= 1000]\nmerged_pred = overlay(merged_pred, protenix_long,\n                      src_cols=[\"x_1\", \"y_1\", \"z_1\", \"x_2\", \"y_2\", \"z_2\"],\n                      dst_cols=[\"x_4\", \"y_4\", \"z_4\", \"x_5\", \"y_5\", \"z_5\"])\n\n# --- Final: align to sample submission order and save ---\nsub = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding-2/sample_submission.csv\")[\"ID\"]\nsub = pd.merge(sub, merged_pred, how=\"left\", on=\"ID\").iloc[:, :18]\nsub[\"resid\"] = sub[\"resid\"].astype(int)\nsub.to_csv('/kaggle/working/submission.csv', index=False)\ndisplay(sub.head())","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}