{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":15231210,"isSourceIdPinned":false},{"sourceType":"datasetVersion","sourceId":11519231,"datasetId":7224412,"databundleVersionId":11970024},{"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":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":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,"isSourceIdPinned":false}],"dockerImageVersionId":31259,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"ONLY_INFER = True # True # False\n\nif ONLY_INFER:\n    # --- Early-exit guard for local runs (put this as the very first cell) ---\n    import os\n    import sys\n    import pandas as pd\n    \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        # Try to derive the correct row count / IDs from an available sequences file.\n        # (adjust paths if your competition uses different names)\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","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2026-02-15T00:26:45.93218Z","iopub.execute_input":"2026-02-15T00:26:45.932512Z","iopub.status.idle":"2026-02-15T00:26:45.937515Z","shell.execute_reply.started":"2026-02-15T00:26:45.932477Z","shell.execute_reply":"2026-02-15T00:26:45.936924Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# trRosettaRNA -> ×","metadata":{}},{"cell_type":"markdown","source":"# Proteinix","metadata":{}},{"cell_type":"code","source":"import os\n\nIS_SCORING_RUN = os.environ.get('KAGGLE_IS_COMPETITION_RERUN')\nprint(IS_SCORING_RUN)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T12:32:51.174346Z","iopub.execute_input":"2026-03-23T12:32:51.175349Z","iopub.status.idle":"2026-03-23T12:32:51.179565Z","shell.execute_reply.started":"2026-03-23T12:32:51.175301Z","shell.execute_reply":"2026-03-23T12:32:51.178925Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install --no-deps /kaggle/input/datasets/tobimichigan/biotite-1-2/biotite-1.2.0-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T12:32:51.180652Z","iopub.execute_input":"2026-03-23T12:32:51.181502Z","iopub.status.idle":"2026-03-23T12:32:54.498631Z","shell.execute_reply.started":"2026-03-23T12:32:51.181474Z","shell.execute_reply":"2026-03-23T12:32:54.497674Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !cp -r /kaggle/input/datasets/qiweiyin/protenix-v1-adjusted/Protenix-v1 /kaggle/working/Protenix-v1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-22T00:08:53.086575Z","iopub.execute_input":"2026-03-22T00:08:53.087015Z","iopub.status.idle":"2026-03-22T00:08:53.090806Z","shell.execute_reply.started":"2026-03-22T00:08:53.086981Z","shell.execute_reply":"2026-03-22T00:08:53.090067Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !cp -r /kaggle/input/stanford-rna-3d-folding-2/PDB_RNA /kaggle/working/Protenix-v1/PDB_RNA","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-22T00:08:53.091863Z","iopub.execute_input":"2026-03-22T00:08:53.092146Z","iopub.status.idle":"2026-03-22T00:08:53.1038Z","shell.execute_reply.started":"2026-03-22T00:08:53.092116Z","shell.execute_reply":"2026-03-22T00:08:53.102921Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import gc\nimport json\nimport os\n\nos.environ[\"LAYERNORM_TYPE\"] = \"torch\"\nos.environ.setdefault(\"RNA_MSA_DEPTH_LIMIT\", \"512\")\n\nimport sys\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom tqdm import tqdm\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\"# \"/kaggle/input/protenix-v1/Protenix-v1\"\nDEFAULT_ROOT_DIR = \"/kaggle/input/datasets/qiweiyin/protenix-v1-adjusted/Protenix-v1-adjust-v2/Protenix-v1-adjust-v2/Protenix-v1\"# \"/kaggle/input/protenix-v1/Protenix-v1\"\n# MODEL_NAME = \"protenix_base_default_v1.0.0\"\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\n\n#如果 TEST_TARGETS 为空，就按全量运行。\n\n# TEST_TARGETS = os.environ.get(\"TEST_TARGETS\", \"9MME,9ZCC\").strip()\n# TEST_TARGETS = os.environ.get(\"TEST_TARGETS\", \"8ZNQ\").strip()\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\n# os.environ[\"USE_MSA\"] = \"true\" # コメントアウトした方が良い可能性あり\n# os.environ[\"USE_TEMPLATE\"] = \"true\"\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                            # \"msa\": {\n                            #     \"precomputed_msa_dir\": f\"{input_path}/MSA/{target_id}.MSA.fasta\",\n                            #     \"pairing_db\": \"rnacentral\"\n                            # }\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                        # \"msa\": {\n                        #     \"precomputed_msa_dir\": f\"{input_path}/MSA/{target_id}.MSA.fasta\",\n                        #     \"pairing_db\": \"rnacentral\"\n                        # }\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\ndef extract_c1_coords_aligned(\n    coords_all,\n    res_ids,\n    atom_names,\n    seq_len: int,\n) -> np.ndarray:\n    \"\"\"\n    coords_all: torch.Tensor or np.ndarray of shape (S, N_atom, 3)\n    res_ids: atom_array.res_id\n    atom_names: atom_array.atom_name\n    seq_len: number of residues in the sequence\n\n    Returns:\n        np.ndarray of shape (S, seq_len, 3), aligned by residue id.\n        Missing residues are filled by linear interpolation, and edges by nearest fill.\n    \"\"\"\n    if hasattr(coords_all, \"detach\"):\n        coords_all = coords_all.detach().cpu().numpy()\n    else:\n        coords_all = np.asarray(coords_all)\n\n    S = coords_all.shape[0]\n    out = np.full((S, seq_len, 3), np.nan, dtype=np.float32)\n\n    # scatter C1' coordinates to the correct residue slot\n    for r in range(1, seq_len + 1):\n        idx = np.where((res_ids == r) & (atom_names == \"C1'\"))[0]\n        if len(idx) > 0:\n            out[:, r - 1, :] = coords_all[:, idx[0], :]\n\n    # fill missing positions sample-wise, axis-wise\n    x = np.arange(seq_len)\n    for s in range(S):\n        for k in range(3):\n            y = out[s, :, k]\n            valid = np.isfinite(y)\n\n            if valid.sum() == 0:\n                out[s, :, k] = 0.0\n            elif valid.sum() == 1:\n                out[s, :, k] = y[valid][0]\n            elif not valid.all():\n                out[s, :, k] = np.interp(x, x[valid], y[valid])\n\n    return out\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() -> 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]\n    test_df = test_df[(test_df['sequence_len'] < 250) | (test_df[\"sequence_len\"] >= 1000)]\n\n    # print(\"MAX_SEQ_LEN =\", MAX_SEQ_LEN)\n    # print(\"Sequence length =\", test_df[\"sequence_len\"].max())\n    \n    if not IS_SCORING_RUN:\n    #     test_df = test_df.head(5)\n    \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    # dataset = InferenceDataset(configs)\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 i in tqdm(range(len(dataset)), total=len(dataset)):\n    #     data, atom_array, error_message = dataset[i]\n        # print(atom_array.res_id[:50])\n        # print(atom_array.atom_name[:50])\n        # target_id = data.get(\"sample_name\", f\"sample_{i}\")\n        # seq = seq_by_id.get(target_id, \"\")\n    for _, row in tqdm(test_df.iterrows(), total=len(test_df)):\n        target_id = row[\"target_id\"]\n        seq = row[\"sequence\"]\n        # data, atom_array, error_message = dataset[i]\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 = coords[:, mask, :].detach().cpu().numpy()\n                coords_all = prediction[\"coordinate\"]  # (S, N_atom, 3)\n                # print(\"coords_all shape:\", coords_all.shape)\n                # print(\"sample0 atom0-10 x:\", coords_all[0, :10, 0])\n                # print(\"sample0 atom20-30 x:\", coords_all[0, 20:30, 0])\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                res_ids = chunk_atom_array.res_id\n                atom_names = chunk_atom_array.atom_name\n                \n                coords = extract_c1_coords_aligned(\n                    coords_all=coords_all,\n                    res_ids=res_ids,\n                    atom_names=atom_names,\n                    seq_len=len(subseq),\n                )\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            # print(f\"{target_id} mask sum:\", mask.sum().item(), \"seq_len:\", len(seq))\n            \n            # coords = coords[:, mask, :].detach().cpu().numpy()\n            coords_all = prediction[\"coordinate\"]  # (S, N_atom, 3)\n            # print(\"coords_all shape:\", coords_all.shape)\n            # print(\"sample0 atom0-10 x:\", coords_all[0, :10, 0])\n            # print(\"sample0 atom20-30 x:\", coords_all[0, 20:30, 0])\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            res_ids = atom_array.res_id\n            atom_names = atom_array.atom_name\n            \n            coords = extract_c1_coords_aligned(\n                coords_all=coords_all,\n                res_ids=res_ids,\n                atom_names=atom_names,\n                seq_len=len(seq),\n            )\n    \n            coords = pad_samples(coords, N_SAMPLE)\n            # print(coords.shape)\n            # print(np.abs(coords[0] - coords[1]).max())\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()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T12:43:46.335847Z","iopub.execute_input":"2026-03-23T12:43:46.336177Z","iopub.status.idle":"2026-03-23T13:00:07.853591Z","shell.execute_reply.started":"2026-03-23T12:43:46.336149Z","shell.execute_reply":"2026-03-23T13:00:07.852789Z"},"_kg_hide-output":true,"_kg_hide-input":false},"outputs":[],"execution_count":null},{"cell_type":"code","source":"protenix_pred = pd.read_csv('/kaggle/working/protenix_submission.csv')\nprotenix_pred","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-23T13:00:07.854997Z","iopub.execute_input":"2026-03-23T13:00:07.855316Z","iopub.status.idle":"2026-03-23T13:00:07.892412Z","shell.execute_reply.started":"2026-03-23T13:00:07.855294Z","shell.execute_reply":"2026-03-23T13:00:07.891784Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # import os\n# # import random\n# # import pandas as pd\n# # import numpy as np\n# # import torch\n# # from tqdm import tqdm\n# # import warnings\n# # warnings.filterwarnings(\"ignore\")  \n\n# # #time0=time.time()\n\n# # print('IMPORT OK !!!!')\n\n# if os.path.exists('/kaggle/input/stanford-rna-3d-folding-2'):\n#     # run on kaggle\n#     TEST_SEQUENCES_PATH = '/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv'\n#     DATA_KAGGLE_DIR = '/kaggle/input/stanford-rna-3d-folding-2'\n#     # CHECKPOINT_PATH = '/kaggle/input/datasets/geraseva/protenix-checkpoints/model_v0.2.0.pt'\n#     CHECKPOINT_PATH = '/kaggle/input/datasets/qiweiyin/protenix-v1-adjusted/Protenix-v1/checkpoint/protenix_base_default_v1.0.0.pt'\n#     # os.environ[\"PROTENIX_DATA_ROOT_DIR\"] = '/kaggle/input/datasets/geraseva/protenix-checkpoints'\n#     os.environ[\"PROTENIX_DATA_ROOT_DIR\"] = '/kaggle/input/datasets/qiweiyin/protenix-v1-adjusted/Protenix-v1/checkpoint'\n# else:\n#     # run locally\n#     DATA_KAGGLE_DIR = '.' \n#     CHECKPOINT_PATH = '/protenix/release_data/checkpoint/model_v0.2.0.pt'\n#     os.environ[\"PROTENIX_DATA_ROOT_DIR\"] = '/protenix/release_data/ccd_cache'\n\n# # from runner.inference import update_inference_configs, InferenceRunner   \n# # from protenix.data.infer_data_pipeline import InferenceDataset\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 protenix.config.config import parse_configs\n\n\n# def parse_output_to_df(output, seq, target_id):\n#     df = []\n#     chain_data = []\n#     for i, res in enumerate(seq):\n#         d={\"ID\": target_id, \"resname\": res, \"resid\": i + 1}\n#         for n in range(len(output)):\n#             d={**d, f'x_{n+1}': round(output[n,i,0].item(),3),\n#                      f'y_{n+1}': round(output[n,i,1].item(),3),\n#                      f'z_{n+1}': round(output[n,i,2].item(),3)}\n#         chain_data.append(d)\n\n#     if len(chain_data)!=0:\n#         chain_df = pd.DataFrame(chain_data)\n#         df.append(chain_df)\n#         ##print(chain_df)\n#     return df\n\n# def seed_everything(seed=0):\n#     np.random.seed(seed)\n#     torch.random.manual_seed(seed)\n#     torch.cuda.manual_seed_all(seed)\n#     os.environ[\"PYTHONHASHSEED\"] = str(seed)\n#     random.seed(seed)\n#     torch.backends.cudnn.deterministic = True\n#     torch.backends.cudnn.benchmark = False\n#     print(f\"Random seed set to {seed}\")\n\n# seed_everything()\n\n# class DictDataset(InferenceDataset):\n#     def __init__(\n#         self,\n#         seq_list: list,\n#         dump_dir: str,\n#         id_list: list = None,\n#         use_msa: bool = False,\n#     ) -> None:\n\n#         self.dump_dir = dump_dir\n#         self.use_msa = use_msa\n#         if isinstance(id_list,type(None)):\n#             self.inputs = [{\"sequences\": \n#                             [{\"rnaSequence\": \n#                                 {\"sequence\": seq, \n#                                 \"count\": 1}}],\n#                             \"name\": \"query\"} for seq in seq_list]\n#         else:\n#             self.inputs = [{\"sequences\": \n#                             [{\"rnaSequence\": \n#                                 {\"sequence\": seq, \n#                                 \"count\": 1}}],\n#                             \"name\": i} for i, seq in zip(id_list,seq_list)]\n\n\n\n# configs_base[\"use_deepspeed_evo_attention\"] = (\n# os.environ.get(\"USE_DEEPSPEED_EVO_ATTENTION\", False) == \"true\")\n# configs_base[\"model\"][\"N_cycle\"] = 10 #10\n# configs_base[\"sample_diffusion\"][\"N_sample\"] = 5\n# configs_base[\"sample_diffusion\"][\"N_step\"] = 200\n# inference_configs['load_checkpoint_path']=CHECKPOINT_PATH\n# configs = {**configs_base, **{\"data\": data_configs}, **inference_configs}\n\n# configs = parse_configs(\n#         configs=configs,\n#         fill_required_with_null=True,\n#     )\n# #print('CONFIGS:')\n# #print(configs)\n# #print('DATA CONFIGS:')\n# #print(data_configs)\n# runner=InferenceRunner(configs)\n\n\n# test_df=pd.read_csv(TEST_SEQUENCES_PATH)\n# test_df[\"sequence_len\"] = test_df[\"sequence\"].str.len()\n# test_df = test_df[(test_df['sequence_len'] >= 200) & (test_df['sequence_len'] < 600)]\n# test_df.reset_index(drop=True, inplace=True)\n# print(test_df.shape)\n\n# dataset = DictDataset(test_df.sequence, dump_dir='output', id_list=test_df.target_id, use_msa=False)\n# num_data = len(dataset)\n# # results = pd.DataFrame()\n# for i, seq in tqdm(enumerate(test_df.sequence),total=num_data):\n#     try:\n#         data, atom_array, data_error_message=dataset[i]\n#         target_id = data[\"sample_name\"]\n#         assert target_id==test_df.target_id[i]\n#         assert data_error_message==''\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#         prediction=prediction['coordinate'][:,data['input_feature_dict']['atom_to_tokatom_idx']==12]\n\n#         result = parse_output_to_df(prediction, seq, target_id)[0]\n#     except Exception as exc:\n#         print(exc)\n#         print(f\"Error message: {data_error_message}\")        \n#         print('Failed to predict', target_id)\n#         result=pd.DataFrame(columns=['ID', 'resname', 'resid', \n#                                         '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#                                         data=[[target_id, x, j+1] + [0.0]*15 for j, x in enumerate(seq)])\n        \n#     result['ID'] = result.apply(lambda x: x.ID + '_' + str(x.resid), axis=1)\n#     if i == 0:\n#         results = result.copy()\n#     if i != 0:\n#         results = pd.concat([results, result], axis=0)\n#     torch.cuda.empty_cache()\n\n# results.to_csv('protenix_submission.csv', index=False, mode='a', header=(i==0))\n\n# # print(pd.read_csv('protenix_submission.csv'))","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2026-02-19T12:58:38.08567Z","iopub.execute_input":"2026-02-19T12:58:38.085917Z","iopub.status.idle":"2026-02-19T12:58:38.092129Z","shell.execute_reply.started":"2026-02-19T12:58:38.085894Z","shell.execute_reply":"2026-02-19T12:58:38.091408Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# protenix_pred = results # pd.read_csv('protenix_submission.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T12:58:38.092857Z","iopub.execute_input":"2026-02-19T12:58:38.093132Z","iopub.status.idle":"2026-02-19T12:58:38.106651Z","shell.execute_reply.started":"2026-02-19T12:58:38.093105Z","shell.execute_reply":"2026-02-19T12:58:38.106063Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# protenix_pred","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T12:58:38.107532Z","iopub.execute_input":"2026-02-19T12:58:38.107791Z","iopub.status.idle":"2026-02-19T12:58:38.118875Z","shell.execute_reply.started":"2026-02-19T12:58:38.107764Z","shell.execute_reply":"2026-02-19T12:58:38.118383Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# DRFold2","metadata":{}},{"cell_type":"code","source":"import os\n\nIS_SCORING_RUN = os.environ.get('KAGGLE_IS_COMPETITION_RERUN')\nprint(IS_SCORING_RUN)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T13:52:50.328169Z","iopub.execute_input":"2026-02-20T13:52:50.328894Z","iopub.status.idle":"2026-02-20T13:52:50.332903Z","shell.execute_reply.started":"2026-02-20T13:52:50.328867Z","shell.execute_reply":"2026-02-20T13:52:50.33227Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import time\nimport sys\nimport torch\nimport random\nimport shutil\nimport numpy as np\nimport pandas as pd\n\nfrom Bio.Seq import Seq\nfrom Bio import pairwise2\n\nfrom tqdm import tqdm\nfrom scipy.spatial import distance_matrix\nfrom scipy.spatial.transform import Rotation as R\n\nimport warnings\nwarnings.filterwarnings('ignore')\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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T13:52:50.623385Z","iopub.execute_input":"2026-02-20T13:52:50.624009Z","iopub.status.idle":"2026-02-20T13:52:54.060651Z","shell.execute_reply.started":"2026-02-20T13:52:50.623984Z","shell.execute_reply":"2026-02-20T13:52:54.060007Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"os.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'\nCFG_MERGE=False\nDR_SCORE=False\nNO_SORT=False\nGET_CENTER=True\n\nFULL_ENERGY=False\n\nOPTIM_LENGTH=0\n\nDEVICE=device #'cuda' #'cpu'#\nPREC=sys_dtype ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T13:52:54.061576Z","iopub.execute_input":"2026-02-20T13:52:54.061802Z","iopub.status.idle":"2026-02-20T13:52:54.114453Z","shell.execute_reply.started":"2026-02-20T13:52:54.06178Z","shell.execute_reply":"2026-02-20T13:52:54.113676Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if PREC=='fp16':\n    torch.set_default_dtype(torch.float16)\nif PREC=='bf16':\n    torch.set_default_dtype(torch.bfloat16)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T13:52:54.423011Z","iopub.execute_input":"2026-02-20T13:52:54.423692Z","iopub.status.idle":"2026-02-20T13:52:54.427418Z","shell.execute_reply.started":"2026-02-20T13:52:54.423664Z","shell.execute_reply":"2026-02-20T13:52:54.42666Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from 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\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\ndef set_aspect_equal(ax):\n\tx_limits = ax.get_xlim()\n\ty_limits = ax.get_ylim()\n\tz_limits = ax.get_zlim()\n\n\t# Compute the mean of each axis\n\tx_middle = np.mean(x_limits)\n\ty_middle = np.mean(y_limits)\n\tz_middle = np.mean(z_limits)\n\n\t# Compute the max range across all axes\n\tmax_range = max(x_limits[1] - x_limits[0],\n\t\t\t\t\ty_limits[1] - y_limits[0],\n\t\t\t\t\tz_limits[1] - z_limits[0]) / 2.0\n\n\t# Set the new limits to ensure equal scaling\n\tax.set_xlim(x_middle - max_range, x_middle + max_range)\n\tax.set_ylim(y_middle - max_range, y_middle + max_range)\n\tax.set_zlim(z_middle - max_range, z_middle + max_range)\n\n\nprint('torch',torch.__version__)\nprint('torch.cuda',torch.version.cuda)\n\nprint('IMPORT OK!!!')\nMODE = 'submit' #'local' # submit\n\nDATA_KAGGLE_DIR = '/kaggle/input/stanford-rna-3d-folding-2'\n\nif MODE == 'local':\n    valid_df = pd.read_csv(f'{DATA_KAGGLE_DIR}/validation_sequences.csv')\n    label_df = pd.read_csv(f'{DATA_KAGGLE_DIR}/validation_labels.csv')\n    label_df['target_id'] = label_df['ID'].apply(lambda x: '_'.join(x.split('_')[:-1]))\n\nif MODE == 'submit':\n    valid_df = pd.read_csv(f'{DATA_KAGGLE_DIR}/test_sequences.csv')\n    # Sort test sequences by length to process shorter ones with DRfold2\n    valid_df[\"sequence_len\"] = valid_df[\"sequence\"].str.len()\n    # test_sequences = test_sequences[(test_sequences['sequence_len'] >= 200) & (test_sequences['sequence_len'] < 600)]\n    # valid_df = valid_df[valid_df['sequence_len'] < 100]\n    valid_df = valid_df[valid_df['sequence_len'] < 250]\n    # valid_df = valid_df[valid_df['sequence_len'] < 350]\n    valid_df.reset_index(drop=True, inplace=True)\n    print(valid_df.shape)\n \n    if 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('MODE:', MODE)\nprint('SETTING OK!!!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T13:52:54.770358Z","iopub.execute_input":"2026-02-20T13:52:54.770931Z","iopub.status.idle":"2026-02-20T13:52:54.819585Z","shell.execute_reply.started":"2026-02-20T13:52:54.770903Z","shell.execute_reply":"2026-02-20T13:52:54.818858Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sys.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()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T13:52:58.996252Z","iopub.execute_input":"2026-02-20T13:52:58.996798Z","iopub.status.idle":"2026-02-20T13:52:59.065848Z","shell.execute_reply.started":"2026-02-20T13:52:58.99677Z","shell.execute_reply":"2026-02-20T13:52:59.065314Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def 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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T13:53:01.149015Z","iopub.execute_input":"2026-02-20T13:53:01.149516Z","iopub.status.idle":"2026-02-20T13:53:01.154634Z","shell.execute_reply.started":"2026-02-20T13:53:01.149466Z","shell.execute_reply":"2026-02-20T13:53:01.153891Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import 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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T13:53:01.274865Z","iopub.execute_input":"2026-02-20T13:53:01.275377Z","iopub.status.idle":"2026-02-20T13:53:01.282602Z","shell.execute_reply.started":"2026-02-20T13:53:01.275351Z","shell.execute_reply":"2026-02-20T13:53:01.281859Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\n\ndef score_energy_one_full(seq, target_id, out, paired=None,\n                          # bond parameters\n                          d0_P_C4=1.60, d0_C4_N=1.47, k_bond=100.0,\n                          # angle parameters\n                          theta0=np.deg2rad(109.5), k_angle=20.0,\n                          # stacking parameters\n                          d0_stack=3.4, k_stack=5.0,\n                          # dihedral parameters\n                          phi0=np.deg2rad(180.0), k_dihedral=5.0,\n                          # hydrogen-bond parameters\n                          d0_hb=2.9, k_hb=10.0,\n                          # Lennard-Jones parameters\n                          sigma=4.0, epsilon=0.1,\n                          # Debye-Hückel electrostatics\n                          q_P=-1.0, epsilon_r=80.0, kappa=10.0,\n                          k_e=332.0637):\n    coor = out['coor']\n    L = len(seq)\n    if paired is None:\n        paired = out.get('paired', [])\n\n    E_bond = 0.0\n    for i in range(L):\n        P  = coor[i,0]; C4 = coor[i,1]; N = coor[i,2]\n        d_PC4 = np.linalg.norm(P - C4)\n        d_C4N = np.linalg.norm(C4 - N)\n        E_bond += k_bond * (d_PC4 - d0_P_C4)**2\n        E_bond += k_bond * (d_C4N - d0_C4_N)**2\n    E_angle = 0.0\n    for i in range(L):\n        P, C4, N = coor[i]\n        v1 = P - C4; v2 = N - C4\n        cost = np.dot(v1, v2) / (np.linalg.norm(v1)*np.linalg.norm(v2) + 1e-8)\n        theta = np.arccos(np.clip(cost, -1.0, 1.0))\n        E_angle += k_angle * (theta - theta0)**2\n\n    E_stack = 0.0\n    for i in range(L-1):\n        d = np.linalg.norm(coor[i,1] - coor[i+1,1])\n        E_stack += k_stack * (d - d0_stack)**2\n\n    def torsion_angle(a, b, c, d):\n        b1, b2, b3 = b-a, c-b, d-c\n        n1 = np.cross(b1, b2); n2 = np.cross(b2, b3)\n        n1 /= np.linalg.norm(n1) + 1e-8; n2 /= np.linalg.norm(n2) + 1e-8\n        cos_phi = np.dot(n1, n2)\n        return np.arccos(np.clip(cos_phi, -1, 1))\n\n    E_dihedral = 0.0\n    for i in range(L-1):\n        a = coor[i,0]; b = coor[i,1]; c = coor[i,2]; d = coor[i+1,0]\n        phi = torsion_angle(a, b, c, d)\n        E_dihedral += k_dihedral * (phi - phi0)**2\n\n    E_hb = 0.0\n    for i, j in paired:\n        d = np.linalg.norm(coor[i,2] - coor[j,2])\n        E_hb += k_hb * (d - d0_hb)**2\n\n    E_LJ = 0.0\n    for i in range(L):\n        for j in range(i+2, L): \n            r = np.linalg.norm(coor[i,1] - coor[j,1])\n            sr6 = (sigma / (r + 1e-8))**6\n            sr12 = sr6 * sr6\n            E_LJ += 4 * epsilon * (sr12 - sr6)\n\n    E_elec = 0.0\n    for i in range(L):\n        for j in range(i+1, L):\n            r = np.linalg.norm(coor[i,0] - coor[j,0])\n            prefac = k_e * q_P * q_P / epsilon_r\n            E_elec += prefac * np.exp(-r / kappa) / (r + 1e-8)\n\n    total_energy = (E_bond + E_angle + E_stack +\n                    E_dihedral + E_hb + E_LJ + E_elec)\n\n    # print(f\"[{target_id}] bond={E_bond:.2f}, angle={E_angle:.2f}, stack={E_stack:.2f}, \\\n#          dihedral={E_dihedral:.2f}, hb={E_hb:.2f}, LJ={E_LJ:.2f}, elec={E_elec:.2f} -> total={total_energy:.2f}\")\n    return total_energy","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T13:53:03.857082Z","iopub.execute_input":"2026-02-20T13:53:03.857705Z","iopub.status.idle":"2026-02-20T13:53:03.869326Z","shell.execute_reply.started":"2026-02-20T13:53:03.857667Z","shell.execute_reply":"2026-02-20T13:53:03.868726Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import 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\ndef score_energy_one(seq, target_id, out):\n    with tempfile.TemporaryDirectory() as tmpdirname:\n        fastafile = os.path.join(tmpdirname, 'tmp.fasta')\n        with open(fastafile, 'w') as f:\n            f.write(f'>{target_id}\\n{seq}\\n')\n        retfile = os.path.join(tmpdirname, 'tmp.ret')\n        with open(retfile, 'wb') as f:\n            f.write(pickle.dumps(out))\n        # foldconfig = '/kaggle/input/drfold-model/DRfold2/cfg_for_selection.json'\n        # foldconfig = '/kaggle/input/drfold-model/DRfold2/cfg_for_folding.json'\n        foldconfig = '/kaggle/input/drfold-model-bf16/cfg_for_folding.json'\n        # foldconfig = 'cfg_for_folding.json'\n        save_prefix = os.path.join(tmpdirname, 'tmp.json')\n        stru=Structure(fastafile,[retfile],save_prefix,0,foldconfig)\n        rama=stru.init_quat(0).data.numpy()\n        energy=stru.obj_func_np(rama)\n        return energy\n\n\n\ndef optimize_coor(seq, target_id, out):\n    print('Optimizing structure for ', target_id)\n    if len(out['coor'])>len(out['plddt']):\n        if len(out['plddt'])>0:\n            mean_plddt = np.mean(out['plddt'])\n            out['plddt'] = np.concatenate([out['plddt'], np.full((len(out['coor'])-len(out['plddt'])), mean_plddt)])\n        else:\n            out['plddt'] = np.full((len(out['coor'])), 0.0)\n            \n    if len(out['coor'])<len(out['plddt']):\n        mean_plddt = np.mean(out['plddt'])\n        out['plddt'] = np.full((len(out['coor'])), mean_plddt)\n    \n    with tempfile.TemporaryDirectory() as tmpdirname:\n        fastafile = os.path.join(tmpdirname, 'tmp.fasta')\n        with open(fastafile, 'w') as f:\n            f.write(f'>{target_id}\\n{seq}\\n')\n\n        retfile = os.path.join(tmpdirname, 'tmp.ret')\n        with open(retfile, 'wb') as f:\n            pickle.dump(out, f)\n\n        # foldconfig = '/kaggle/input/drfold-model/DRfold2/cfg_for_folding.json'\n        foldconfig = '/kaggle/input/drfold-model-bf16/cfg_for_folding.json'\n        # foldconfig = 'cfg_for_folding.json'\n        save_prefix = os.path.join(tmpdirname, 'tmp')\n        stru = Structure(fastafile, [retfile], save_prefix, 0, foldconfig)\n        stru.foldning()\n\n        pdb_file = save_prefix + '.pdb'\n        parser = PDBParser(QUIET=True)\n        structure = parser.get_structure(target_id, pdb_file)\n\n        residues = [\n            res for res in structure.get_residues()\n            if res.id[0] == ' '\n        ]\n        L = len(seq)\n        if len(residues) != L:\n            raise ValueError(f\"PDB 中残基数 ({len(residues)}) 与序列长度 ({L}) 不一致\")\n\n        atom_order = ['P', \"C4'\", 'N1/N9']\n        coor = np.zeros((L, 3, 3), dtype=float)\n\n        for i, res in enumerate(residues):\n            coor[i, :, :] = np.nan  \n            \n            if 'P' in res:\n                coord = res['P'].get_vector().get_array()\n                coor[i, 0, :] = coord\n            if \"C4'\" in res:\n                coord = res[\"C4'\"].get_vector().get_array()\n                coor[i, 1, :] = coord\n            if 'N1' in res:\n                coord = res['N1'].get_vector().get_array()\n                coor[i, 2, :] = coord\n            elif 'N9' in res:\n                coord = res['N9'].get_vector().get_array()\n                coor[i, 2, :] = coord\n\n        return coor","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T13:53:05.783944Z","iopub.execute_input":"2026-02-20T13:53:05.784234Z","iopub.status.idle":"2026-02-20T13:53:05.934813Z","shell.execute_reply.started":"2026-02-20T13:53:05.784209Z","shell.execute_reply":"2026-02-20T13:53:05.934193Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sys.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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T13:53:08.054098Z","iopub.execute_input":"2026-02-20T13:53:08.05493Z","iopub.status.idle":"2026-02-20T13:53:11.778314Z","shell.execute_reply.started":"2026-02-20T13:53:08.0549Z","shell.execute_reply":"2026-02-20T13:53:11.777741Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 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 solution_to_submit_df(solution):\n    submit_df = []\n    for k,s in solution.items():\n        df = coord_to_df(s.sequence, s.coord, s.target_id)\n        submit_df.append(df)\n    \n    submit_df = pd.concat(submit_df)\n    return submit_df\n \n\ndef coord_to_df(sequence, coord, target_id):\n    L = len(sequence)\n    df = pd.DataFrame()\n    df['ID'] = [f'{target_id}_{i + 1}' for i in range(L)]\n    df['resname'] = [s for s in sequence]\n    df['resid'] = [i + 1 for i in range(L)]\n\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\nout_dir = '/kaggle/working/model-output'\nos.makedirs(out_dir, exist_ok=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T13:53:11.77955Z","iopub.execute_input":"2026-02-20T13:53:11.77985Z","iopub.status.idle":"2026-02-20T13:53:11.788155Z","shell.execute_reply.started":"2026-02-20T13:53:11.779818Z","shell.execute_reply":"2026-02-20T13:53:11.787358Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import 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    if CFG_MERGE:\n        msa2xyz_file = [\n            f'/kaggle/input/datasets/z1493916656/drfold-model-bf16/model_hub/cfg_97/model_{i}'\n            for i in range(20)\n        ] + [\n            f'/kaggle/input/datasets/z1493916656/drfold-model-bf16/model_hub/cfg_95/model_{i}'\n            for i in range(20)\n        ] + [\n            f'/kaggle/input/datasets/z1493916656/drfold-model-bf16/model_hub/cfg_96/model_{i}'\n            for i in range(20)\n        ] + [\n            f'/kaggle/input/datasets/z1493916656/drfold-model-bf16/model_hub/cfg_99/model_{i}'\n            for i in range(20)\n        ]\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=True))\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    # start here !!!!!!!!!!!!!!!!!!!!!!\n    #valid_df = valid_df.iloc[[0,1]].reset_index(drop=True)\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        # else:\n        #     continue\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 = [0,1,2,8,9]\n            model_to_try = [13, 6, 14, 5, 3]\n        elif  len(seq)>100:\n            # model_to_try = [0,2,4,6,8,10,12,14,16,18]#list(range(min(num_msa2xyz,10)))\n            model_to_try = [13, 6, 14, 12, 7, 2, 5, 19, 10, 9]\n            if CFG_MERGE:\n                model_to_try = [24, 34, 20, 13, 6, 37, 28, 25, 14, 39]\n        else:\n            # model_to_try = list(range(num_msa2xyz))\n            \n            # model_to_try = list(range(20))\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            # if CFG_MERGE:\n            #     model_to_try = list(range(20)) + list(range(40,60)) + list(range(60, 80))\n            \n        if NO_SORT:\n            model_to_try=model_to_try[:5]\n\n\n        # 分段预测\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            if NO_SORT:\n                e=0\n            elif len(model_to_try)>5 and DR_SCORE:\n                e = score_energy_one(seq, target_id, out_cat)\n            elif FULL_ENERGY:\n                e = score_energy_one_full(seq, target_id, out_cat)\n            else:\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            #print('time_taken:', time_to_str(time_taken, mode='sec'))\n\n            gpu_mem_used = gpu_memory_use()\n            max_gpu_mem_used = max(max_gpu_mem_used,gpu_mem_used)\n            #print('gpu_mem_used:', gpu_mem_used, 'GB')\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        if GET_CENTER:\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        if L <= OPTIM_LENGTH:\n            out_opt= outputs[argsort[0]]\n            out_opt['coor'] = optimize_coor(seq, target_id, out_opt)\n            coordinate[argsort[0]] = frame_coor_to_C1(out_opt['coor'], sequence, BASE_COOR, OTHER_COOR)\n            torch.cuda.empty_cache()\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!!!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T14:45:04.074038Z","iopub.execute_input":"2026-02-20T14:45:04.074343Z","iopub.status.idle":"2026-02-20T14:47:37.477869Z","shell.execute_reply.started":"2026-02-20T14:45:04.074316Z","shell.execute_reply":"2026-02-20T14:47:37.477036Z"},"_kg_hide-output":true,"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"drfold_pred = pd.read_csv('/kaggle/working/drfold_submission.csv')\ndrfold_pred","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T14:47:37.486888Z","iopub.execute_input":"2026-02-20T14:47:37.48715Z","iopub.status.idle":"2026-02-20T14:47:37.508798Z","shell.execute_reply.started":"2026-02-20T14:47:37.487127Z","shell.execute_reply":"2026-02-20T14:47:37.508146Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# test_sequences = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv\")\n\n# # Set up directories\n# predictions_dir = \"/kaggle/working/predictions\"\n# os.makedirs(predictions_dir, exist_ok=True)\n# fasta_dir = \"/kaggle/working/fasta_files\"\n# os.makedirs(fasta_dir, exist_ok=True)\n\n# # Set time limit for DRfold2 (in seconds)\n# DRFOLD_TIME_LIMIT = 7 * 60 * 60  # 7 hours\n# start_time_global = time.time()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T13:02:03.098985Z","iopub.execute_input":"2026-02-19T13:02:03.099263Z","iopub.status.idle":"2026-02-19T13:02:03.102477Z","shell.execute_reply.started":"2026-02-19T13:02:03.099242Z","shell.execute_reply":"2026-02-19T13:02:03.101764Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !cp -r /kaggle/input/datasets/jaejohn/drfold2-repo/DRfold2 /kaggle/working/\n# %cd DRfold2\n# %cd Arena\n# !make Arena\n# %cd ..","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T13:02:03.103314Z","iopub.execute_input":"2026-02-19T13:02:03.103671Z","iopub.status.idle":"2026-02-19T13:02:03.113587Z","shell.execute_reply.started":"2026-02-19T13:02:03.103642Z","shell.execute_reply":"2026-02-19T13:02:03.112879Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !cp -r /kaggle/input/datasets/jaejohn/drfold2/model_hub /kaggle/working/DRfold2/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T13:02:03.114454Z","iopub.execute_input":"2026-02-19T13:02:03.114711Z","iopub.status.idle":"2026-02-19T13:02:03.12468Z","shell.execute_reply.started":"2026-02-19T13:02:03.11469Z","shell.execute_reply":"2026-02-19T13:02:03.124008Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %%writefile /kaggle/working/DRfold2/DRfold_infer.py\n# import os,sys\n# import torch\n# import numpy as np\n# from subprocess import Popen, PIPE, STDOUT\n\n# # Get the directory where the script is located\n# exp_dir = os.path.dirname(os.path.abspath(__file__))\n\n# device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n# # dlexps = ['cfg_95','cfg_96','cfg_97','cfg_99']\n# dlexps = ['cfg_97']\n\n# print(f\"[DRfold2] Starting prediction pipeline on {device} device\")\n\n# # Get input FASTA file and output directory from command line arguments\n# fastafile =  os.path.realpath(sys.argv[1])\n# outdir = os.path.realpath(sys.argv[2])\n\n# print(f\"[DRfold2] Input: {fastafile}\")\n# print(f\"[DRfold2] Output: {outdir}\")\n\n# # Initialize clustering flag\n# pclu = False\n\n# # If third argument is '1', enable clustering\n# if len(sys.argv) == 4 and sys.argv[3] == '1': \n#     print('[DRfold2] Clustering enabled - will generate multiple models')\n#     pclu = True\n# else:\n#     print('[DRfold2] Clustering disabled - will generate single model')\n\n# # Create output directory if it doesn't exist\n# if not os.path.isdir(outdir):\n#     os.makedirs(outdir)\n#     print(f\"[DRfold2] Created output directory: {outdir}\")\n\n# # Create subdirectories for different outputs\n# ret_dir = os.path.join(outdir,'rets_dir')  # For return files\n# if not os.path.isdir(ret_dir):\n#     os.makedirs(ret_dir)\n#     print(f\"[DRfold2] Created returns directory: {ret_dir}\")\n\n# folddir = os.path.join(outdir,'folds')     # For folded structures\n# if not os.path.isdir(folddir):\n#     os.makedirs(folddir)\n#     print(f\"[DRfold2] Created folds directory: {folddir}\")\n\n# refdir = os.path.join(outdir,'relax')      # For relaxed structures\n# if not os.path.isdir(refdir):\n#     os.makedirs(refdir)\n#     print(f\"[DRfold2] Created relaxation directory: {refdir}\")\n\n# # Helper function to run commands and capture output\n# def run_cmd(cmd, description):\n#     print(f\"[DRfold2] {description}\")\n#     print(f\"[DRfold2] Command: {cmd}\")\n    \n#     # Execute the command and capture output in real-time\n#     process = Popen(cmd, shell=True, stdout=PIPE, stderr=STDOUT, universal_newlines=True, bufsize=1)\n    \n#     # Print output line by line as it becomes available\n#     for line in iter(process.stdout.readline, ''):\n#         line = line.strip()\n#         if line:\n#             print(f\"[DRfold2 subprocess] {line}\")\n    \n#     # Get return code\n#     return_code = process.wait()\n#     if return_code == 0:\n#         print(f\"[DRfold2] {description} completed successfully\")\n#     else:\n#         print(f\"[DRfold2] {description} failed with return code {return_code}\")\n#     return return_code\n\n# # Create paths for model directories and test scripts\n# dlmains = [os.path.join(exp_dir, one_exp, 'test_modeldir.py') for one_exp in dlexps]\n# dirs = [os.path.join(exp_dir, 'model_hub', one_exp) for one_exp in dlexps]\n\n# # Check if processing has been done before\n# if not os.path.isfile(ret_dir + '/done'): \n#     print(\"[DRfold2] Step 1/4: GENERATING INITIAL PREDICTIONS\")\n#     print(f\"[DRfold2] No previous predictions found, will generate e2e and geo files\")\n    \n#     # Run each model configuration\n#     for idx, (dlmain, one_exp, mdir) in enumerate(zip(dlmains, dlexps, dirs)):\n#         # Construct command to run the model\n#         cmd = f'python {dlmain} {device} {fastafile} {ret_dir}/{one_exp}_ {mdir}'\n#         description = f\"Running model {idx+1}/{len(dlexps)}: {one_exp}\"\n#         run_cmd(cmd, description)\n\n#     # Mark processing as complete\n#     wfile = open(ret_dir+'/done','w')\n#     wfile.write('1')\n#     wfile.close()\n#     print(\"[DRfold2] Initial predictions generation completed\")\n# else:\n#     print(\"[DRfold2] Step 1/4: USING EXISTING PREDICTIONS\")\n#     print(f\"[DRfold2] Found previous predictions in {ret_dir}, using existing e2e and geo files\")\n\n# # Helper function to get model PDB file\n# def get_model_pdb(tdir,opt):\n#     files = os.listdir(tdir)\n#     files = [afile for afile in files if afile.startswith(opt)][0]\n#     return files\n\n# # Set up directory paths and configuration files\n# cso_dir = folddir                                                    # Directory for coarse-grained structures\n# clufile = os.path.join(folddir,'clu.txt')                            # Clustering results file\n# config_sel = os.path.join(exp_dir,'cfg_for_selection.json')          # Selection configuration\n# foldconfig = os.path.join(exp_dir,'cfg_for_folding.json')            # Folding configuration\n# selpython = os.path.join(exp_dir,'PotentialFold','Selection.py')     # Selection script\n# optpython = os.path.join(exp_dir,'PotentialFold','Optimization.py')  # Optimization script\n# clupy = os.path.join(exp_dir,'PotentialFold','Clust.py')             # Clustering script\n# arena = os.path.join(exp_dir,'Arena','Arena')                        # Arena executable for structure refinement\n\n# # Set up initial save prefixes for optimization and selection\n# optsaveprefix = os.path.join(cso_dir, f'opt_0')\n# save_prefix = os.path.join(cso_dir, f'sel_0')\n\n# # Get all .ret files from the return directory\n# rets = os.listdir(ret_dir)\n# rets = [afile for afile in rets if afile.endswith('.ret')]\n# rets = [os.path.join(ret_dir,aret) for aret in rets ]\n# ret_str = ' '.join(rets)\n\n# print(\"[DRfold2] Step 2/4: SELECTION PROCESS\")\n# print(f\"[DRfold2] Found {len(rets)} return files for selection\")\n# print(f\"[DRfold2] Using selection config: {config_sel}\")\n# print(f\"[DRfold2] Output prefix: {save_prefix}\")\n\n# # Run selection process\n# cmd = f'python {selpython} {fastafile} {config_sel} {save_prefix} {ret_str}'\n# run_cmd(cmd, \"Running selection process\")\n\n# print(\"[DRfold2] Step 3/4: OPTIMIZATION PROCESS\")\n# print(f\"[DRfold2] Using fold config: {foldconfig}\")\n# print(f\"[DRfold2] Optimization output prefix: {optsaveprefix}\")\n\n# # Run optimization process\n# cmd = f'python {optpython} {fastafile} {optsaveprefix} {ret_dir} {save_prefix} {foldconfig}'\n# run_cmd(cmd, \"Running optimization process\")\n\n# # Get the coarse-grained PDB and save refined structure\n# cgpdb = os.path.join(folddir,get_model_pdb(folddir,'opt_0'))\n# savepdb = os.path.join(refdir,'model_1.pdb')\n\n# print(\"[DRfold2] Step 4/4: STRUCTURE REFINEMENT\")\n# print(f\"[DRfold2] Found optimized structure: {cgpdb}\")\n# print(f\"[DRfold2] Final output will be saved to: {savepdb}\")\n\n# cmd = f'{arena} {cgpdb} {savepdb} 7'\n# run_cmd(cmd, \"Running structure refinement\")\n\n# # If clustering is enabled (pclu=True)\n# if pclu:\n#     print(\"[DRfold2] ADDITIONAL STEP: CLUSTERING\")\n#     print(f\"[DRfold2] Running clustering process, output: {clufile}\")\n    \n#     # Run clustering process\n#     cmd = f'python {clupy} {ret_dir} {clufile}'\n#     run_cmd(cmd, \"Running clustering\")\n\n#     # Read clustering results\n#     lines = open(clufile).readlines()\n#     lines = [aline.strip() for aline in lines]\n#     lines = [aline for aline in lines if aline]\n    \n#     cluster_count = len(lines) - 1\n#     print(f\"[DRfold2] Found {cluster_count} additional clusters to process\")\n\n#     # Process each cluster\n#     for i in range(1,len(lines)):\n#         print(f\"[DRfold2] PROCESSING CLUSTER {i}/{cluster_count}\")\n        \n#         # Get return files for this cluster\n#         rets = lines[i].split()\n#         rets = [os.path.join(ret_dir,aret.replace('.pdb','.ret')) for aret in rets ]\n#         ret_str = ' '.join(rets)\n\n#         # Set up save prefixes for this cluster\n#         optsaveprefix =  os.path.join(cso_dir,f'opt_{str(i+1)}')\n#         save_prefix = os.path.join(cso_dir,f'sel_{str(i+1)}')\n        \n#         print(f\"[DRfold2] Cluster {i} Selection Process\")\n#         print(f\"[DRfold2] Found {len(rets)} return files for selection\")\n#         print(f\"[DRfold2] Selection output prefix: {save_prefix}\")\n\n#         # Run selection process for this cluster\n#         cmd = f'python {selpython} {fastafile} {config_sel} {save_prefix} {ret_str}'\n#         run_cmd(cmd, f\"Running selection for cluster {i}\")\n        \n#         print(f\"[DRfold2] Cluster {i} Optimization Process\")\n#         print(f\"[DRfold2] Optimization output prefix: {optsaveprefix}\")\n\n#         # Run optimization process for this cluster\n#         cmd = f'python {optpython} {fastafile} {optsaveprefix} {ret_dir} {save_prefix} {foldconfig}'\n#         run_cmd(cmd, f\"Running optimization for cluster {i}\")\n\n#         # Get the coarse-grained PDB and save refined structure for this cluster\n#         cgpdb = os.path.join(folddir,get_model_pdb(folddir,f'opt_{str(i+1)}'))\n#         savepdb = os.path.join(refdir,f'model_{str(i+1)}.pdb')\n        \n#         print(f\"[DRfold2] Cluster {i} Refinement Process\")\n#         print(f\"[DRfold2] Found optimized structure: {cgpdb}\")\n#         print(f\"[DRfold2] Final output will be saved to: {savepdb}\")\n\n#         cmd = f'{arena} {cgpdb} {savepdb} 7'\n#         run_cmd(cmd, f\"Running refinement for cluster {i}\")\n\n# print(\"[DRfold2] PREDICTION PIPELINE COMPLETED SUCCESSFULLY\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T13:02:03.125484Z","iopub.execute_input":"2026-02-19T13:02:03.125735Z","iopub.status.idle":"2026-02-19T13:02:03.137235Z","shell.execute_reply.started":"2026-02-19T13:02:03.125708Z","shell.execute_reply":"2026-02-19T13:02:03.136711Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %%writefile /kaggle/working/DRfold2/PotentialFold/operations.py\n# \"\"\"\n# operations.py: Core Mathematical Operations for RNA Structure Analysis\n\n# This module provides essential mathematical operations for manipulating and analyzing\n# RNA 3D structures, organized into four main categories:\n\n# 1. Basic Vector Operations:\n#    Functions for selecting coordinates and calculating distances between points,\n#    which form the foundation for all structural calculations.\n\n# 2. Angle Calculations:\n#    Functions for computing bond angles and dihedral (torsion) angles between atoms,\n#    with differentiable implementations suitable for gradient-based optimization.\n\n# 3. Rigid Body Transformations:\n#    Functions for determining optimal rotations and translations between sets of\n#    coordinates, enabling structure alignment and manipulation.\n\n# 4. Sequence Utilities:\n#    Functions for converting RNA sequence data into standard 3D coordinate templates,\n#    allowing sequence-structure mapping.\n\n# These operations support the core functionality of RNA structure prediction, analysis,\n# and optimization throughout the codebase.\n# \"\"\"\n\n# import os\n# import torch\n# import torch.nn as nn\n# import numpy as np \n# import math, sys, math\n# from io import BytesIO\n# import torch.nn.functional as F\n# from torch.autograd import Function\n# from torch.nn.parameter import Parameter\n# from subprocess import Popen, PIPE, STDOUT\n\n# # Use consistent epsilon value across all functions\n# EPS = 1e-8\n\n\n# # === Basic Vector Operations ===\n# def coor_selection(coor,mask):\n#     #[L,n,3],[L,n],byte\n#     return torch.masked_select(coor,mask.bool()).view(-1,3)\n\n# def pair_distance(x1, x2, eps=1e-6, p=2):\n#     # Use torch.cdist for p=2 (Euclidean) which is highly optimized\n#     if p == 2:\n#         return torch.cdist(x1, x2, p=2)\n    \n#     # For other p-norms, avoid memory expansion with broadcasting\n#     x1_ = x1.unsqueeze(1)  # [n1, 1, dim]\n#     x2_ = x2.unsqueeze(0)  # [1, n2, dim]\n#     diff = torch.abs(x1_ - x2_)\n#     out = torch.pow(diff + eps, p).sum(dim=2)\n#     return torch.pow(out, 1. / p)\n\n\n# # === Angle Calculations ===\n# def angle(p0, p1, p2):\n#     # [b 3] \n#     b0 = p0-p1\n#     b1 = p2-p1\n\n#     b0 = b0 / (torch.norm(b0, dim =-1, keepdim=True) + EPS)\n#     b1 = b1 / (torch.norm(b1, dim =-1, keepdim=True) + EPS)\n    \n#     recos = torch.sum(b0*b1, -1)\n#     recos = torch.clamp(recos, -0.9999, 0.9999)\n#     return torch.acos(recos)\n\n# class torsion(Function):\n#     #PyTorch class to calculate differentiable torsion angle\n#     #https://stackoverflow.com/questions/20305272/dihedral-torsion-angle-from-four-points-in-cartesian-coordinates-in-python\n#     #https://salilab.org/modeller/manual/node492.html\n#     @staticmethod\n#     def forward(ctx, p0, p1, p2, p3):\n#         # Save input points for backward pass\n#         ctx.save_for_backward(p0, p1, p2, p3)\n\n#         # Calculate bond vectors\n#         b0 = p0 - p1\n#         b1 = p2 - p1\n#         b2 = p3 - p2\n\n#         # Normalize the middle bond vector\n#         b1_norm = torch.norm(b1, dim=-1, keepdim=True) + 1e-8\n#         b1_unit = b1 / b1_norm\n\n#         # Project the other bonds onto the plane perpendicular to middle bond\n#         v = b0 - torch.sum(b0 * b1_unit, dim=-1, keepdim=True) * b1_unit\n#         w = b2 - torch.sum(b2 * b1_unit, dim=-1, keepdim=True) * b1_unit\n\n#         # Calculate torsion using the arctan2 formula (more stable than arccos)\n#         x = torch.sum(v * w, dim=-1)                                # cosine component\n#         y = torch.sum(torch.cross(b1_unit, v, dim=-1) * w, dim=-1)  # sine component\n\n#         return torch.atan2(y, x)\n\n    \n#     @staticmethod\n#     def backward(ctx, grad_output):\n#         # Retrieve saved tensors from forward pass\n#         p0, p1, p2, p3 = ctx.saved_tensors\n\n#         # Calculate bond vectors\n#         r01 = p0 - p1\n#         r12 = p2 - p1\n#         r23 = p3 - p2\n\n#         # Calculate bond lengths with numerical stability\n#         d01 = torch.norm(r01, dim=-1, keepdim=True) + 1e-8\n#         d12 = torch.norm(r12, dim=-1, keepdim=True) + 1e-8\n#         d23 = torch.norm(r23, dim=-1, keepdim=True) + 1e-8\n\n#         # Normalize bond vectors\n#         e01 = r01 / d01\n#         e12 = r12 / d12\n#         e23 = r23 / d23\n\n#         # Calculate normal vectors to the two planes\n#         n1 = torch.cross(e01, e12, dim=-1)\n#         n2 = torch.cross(e12, e23, dim=-1)\n\n#         # Normalize normal vectors\n#         n1_norm = torch.norm(n1, dim=-1, keepdim=True) + 1e-8\n#         n2_norm = torch.norm(n2, dim=-1, keepdim=True) + 1e-8\n#         n1 = n1 / n1_norm\n#         n2 = n2 / n2_norm\n\n#         # Calculate gradients for each atom\n#         # These are based on the analytical derivatives of dihedral angles\n#         g0 = torch.cross(e01, n1, dim=-1) / d01\n#         g1 = -g0 - torch.cross(e12, n1, dim=-1) / d12\n#         g2 = torch.cross(e12, n2, dim=-1) / d12 - torch.cross(e23, n2, dim=-1) / d23\n#         g3 = torch.cross(e23, n2, dim=-1) / d23\n\n#         # Apply chain rule with incoming gradient\n#         g0 = g0 * grad_output.unsqueeze(-1)\n#         g1 = g1 * grad_output.unsqueeze(-1)\n#         g2 = g2 * grad_output.unsqueeze(-1)\n#         g3 = g3 * grad_output.unsqueeze(-1)\n\n#         return g0, g1, g2, g3\n\n\n# def dihedral(input1, input2, input3, input4):\n#     return torsion.apply(input1, input2, input3, input4)\n\n\n\n# # === Rigid Body Transformations ===\n# def rigidFrom3Points(x):    \n#     x1, x2, x3 = x[:, 0], x[:, 1], x[:, 2]\n#     v1 = x3 - x2\n#     v2 = x1 - x2\n    \n#     # Normalize v1 to get e1\n#     e1 = F.normalize(v1, p=2, dim=-1)\n    \n#     # Project v2 onto e1 and subtract to get the component orthogonal to e1\n#     u2 = v2 - e1 * (torch.einsum('bn,bn->b', e1, v2)[:, None])\n    \n#     # Normalize u2 to get e2\n#     e2 = F.normalize(u2, p=2, dim=-1)\n    \n#     # Cross product to get e3\n#     e3 = torch.cross(e1, e2, dim=-1)\n    \n#     return torch.stack([e1, e2, e3], dim=1)\n\n\n# # return the direction from to_q to from_p\n# def Kabsch_rigid(bases,x1,x2,x3):\n#     # Early return for empty input\n#     if x1.shape[0] == 0:\n#         return torch.empty(0, 3, 3), torch.empty(0, 3)\n    \n#     the_dim=1\n#     to_q = torch.stack([x1,x2,x3],dim=the_dim)\n#     biasq=torch.mean(to_q,dim=the_dim,keepdim=True)\n#     q=to_q-biasq\n#     m = torch.einsum('bnz,bny->bzy',bases,q)\n#     u, s, v = torch.svd(m)\n#     vt = torch.transpose(v, 1, 2)\n#     det = torch.det(torch.matmul(u, vt))\n#     det = det.view(-1, 1, 1)\n#     vt = torch.cat((vt[:, :2, :], vt[:, -1:, :] * det), 1)\n#     r = torch.matmul(u, vt)\n#     return r,biasq.squeeze()\n\n\n\n# # === Sequence Utilities ===\n# def Get_base(seq,basenpy_standard):\n#     base_num = basenpy_standard.shape[1]\n#     basenpy = np.zeros([len(seq),base_num,3])\n#     seqnpy = np.array(list(seq))\n#     basenpy[seqnpy=='A']=basenpy_standard[0]\n#     basenpy[seqnpy=='a']=basenpy_standard[0]\n\n#     basenpy[seqnpy=='G']=basenpy_standard[1]\n#     basenpy[seqnpy=='g']=basenpy_standard[1]\n\n#     basenpy[seqnpy=='C']=basenpy_standard[2]\n#     basenpy[seqnpy=='c']=basenpy_standard[2]\n\n#     basenpy[seqnpy=='U']=basenpy_standard[3]\n#     basenpy[seqnpy=='u']=basenpy_standard[3]\n\n#     basenpy[seqnpy=='T']=basenpy_standard[3]\n#     basenpy[seqnpy=='t']=basenpy_standard[3]\n    \n#     return torch.from_numpy(basenpy).double()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T13:02:03.138267Z","iopub.execute_input":"2026-02-19T13:02:03.138793Z","iopub.status.idle":"2026-02-19T13:02:03.154462Z","shell.execute_reply.started":"2026-02-19T13:02:03.138772Z","shell.execute_reply":"2026-02-19T13:02:03.153792Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %%writefile /kaggle/working/DRfold2/PotentialFold/Optimization.py\n# #! /nfs/amino-home/liyangum/miniconda3/bin/python\n# import torch\n# import random\n# import numpy as np \n# import os, json, sys\n\n# import Cubic, Potential\n# import operations\n# import a2b, rigid\n# import torch.optim as opt\n# from scipy.optimize import minimize\n# import pickle\n\n# torch.manual_seed(6)\n# np.random.seed(9)\n# random.seed(9)\n\n\n# Scale_factor = 1.0\n# USEGEO = False\n\n# device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n\n# def readconfig(configfile=''):\n#     config=[]\n#     expdir=os.path.dirname(os.path.abspath(__file__))\n#     if configfile=='':\n#         configfile=os.path.join(expdir,'lib','ddf.json')\n#     config=json.load(open(configfile,'r'))\n#     return config \n\n    \n# class Structure:\n#     def __init__(self,fastafile,geofiles,saveprefix,initial_ret,foldconfig):\n#         self.config=readconfig(foldconfig)\n#         self.seqfile=fastafile\n#         self.init_ret = initial_ret\n#         self.foldconfig = foldconfig\n#         self.geofiles = geofiles\n#         self.rets = [pickle.load(open(refile,'rb')) for refile  in geofiles]\n#         self.txs=[]\n#         for ret in self.rets:\n#             self.txs.append( torch.from_numpy(ret['coor']).double().to(device))\n#         self.handle_geo()\n#         self.pair = []\n#         for ret in self.rets:\n#             self.pair.append(torch.from_numpy(ret['plddt']).double().to(device))\n#         self.saveprefix=saveprefix\n#         self.seq=open(fastafile).readlines()[1].strip()\n#         self.L=len(self.seq)\n#         basenpy = np.load(os.path.join(os.path.dirname(os.path.abspath(__file__)),'lib','base.npy'))\n#         self.basex = operations.Get_base(self.seq,basenpy).to(device)\n#         othernpy = np.load(os.path.join(os.path.dirname(os.path.abspath(__file__)),'lib','other2.npy'))\n#         self.otherx = operations.Get_base(self.seq,othernpy).to(device)\n#         sidenpy = np.load(os.path.join(os.path.dirname(os.path.abspath(__file__)),'lib','side.npy'))\n#         self.sidex = operations.Get_base(self.seq,sidenpy).to(device)\n        \n#         self.init_mask()\n#         self.init_paras()\n#         self._init_fape()\n#         self.tx2ds = [td.to(device) for td in self.tx2ds]\n#         self.local_weight = torch.ones(self.L,self.L).to(device)\n        \n#         for i in range(self.L):\n#             for j in range(i+1,min(self.L,i+2)):\n#                 self.local_weight[i,j] = self.local_weight[j,i] = 4\n#             for j in range(i+2,min(self.L,i+3)):\n#                 self.local_weight[i,j] = self.local_weight[j,i] = 3\n#             for j in range(i+3,min(self.L,i+4)):\n#                 self.local_weight[i,j] = self.local_weight[j,i] = 2\n\n#     def _init_fape(self):\n#         self.tx2ds = []\n#         for tx in self.txs:\n#             true_rot,true_trans = operations.Kabsch_rigid(self.basex,tx[:,0],tx[:,1],tx[:,2])\n#             true_x2 = tx[:,None,:,:] - true_trans[None,:,None,:]\n#             true_x2 = torch.einsum('ijnd,jde->ijne',true_x2,true_rot.transpose(-1,-2))\n#             self.tx2ds.append(true_x2)\n    \n#     def handle_geo(self):\n#         oldkeys=['dist_p','dist_c','dist_n']\n#         newkeys=['pp','cc','nn']\n#         self.geos=[]\n#         for ret in self.rets:\n#             geo = {}\n#             for nk,ok in zip(newkeys,oldkeys):\n#                 geo[nk] = torch.from_numpy(ret[ok].astype(np.float64)).to(device) + 0\n#             self.geos.append(geo)\n\n\n#     def init_mask(self):\n#         halfmask=np.zeros([self.L,self.L])\n#         fullmask=np.zeros([self.L,self.L])\n#         for i in range(self.L):\n#             for j in range(i+1,self.L):\n#                 halfmask[i,j]=1\n#                 fullmask[i,j]=1\n#                 fullmask[j,i]=1\n#         self.halfmask=(torch.DoubleTensor(halfmask) > 0.5).to(device)\n#         self.fullmask=(torch.DoubleTensor(fullmask) > 0.5).to(device)\n#         self.clash_mask = torch.zeros([self.L,self.L,22,22], device=device)\n#         for i in range(self.L):\n#             for j in range(i+1,self.L):\n#                 self.clash_mask[i,j]=1\n\n#         for i in range(self.L):\n#              self.clash_mask[i,i,:6,7:]=1\n\n#         for i in range(self.L-1):\n#             self.clash_mask[i,i+1,:,0]=0\n#             self.clash_mask[i,i+1,0,:]=0\n#             self.clash_mask[i,i+1,:,5]=0\n#             self.clash_mask[i,i+1,5,:]=0\n\n#         self.side_mask = rigid.side_mask(self.seq).to(device)\n#         self.side_mask = (self.side_mask[:,None,:,None] * self.side_mask[None,:,None,:]).to(device)\n#         self.clash_mask = ((self.clash_mask > 0.5) * (self.side_mask > 0.5)).to(device)\n\n#         self.geo_confimask_cc = []\n#         self.geo_confimask_pp = []\n#         self.geo_confimask_nn = []\n#         for geo in self.geos:\n#             confimask_cc = geo['cc'][:,:,-1] < 0.5\n#             confimask_pp = geo['pp'][:,:,-1] < 0.5\n#             confimask_nn = geo['nn'][:,:,-1] < 0.5\n#             self.geo_confimask_cc.append(confimask_cc)\n#             self.geo_confimask_pp.append(confimask_pp)\n#             self.geo_confimask_nn.append(confimask_nn)\n\n\n#     def init_paras(self):\n#         self.geo_cc = []\n#         self.geo_pp = []\n#         self.geo_nn = []\n#         self.cs_coefs = {'cc': [], 'pp': [], 'nn': []}\n#         self.cs_knots = {'cc': [], 'pp': [], 'nn': []}\n#         for geo in self.geos:\n#             cc_cs,cc_decs=Cubic.dis_cubic(geo['cc'],2,40,36)\n#             pp_cs,pp_decs=Cubic.dis_cubic(geo['pp'],2,40,36)\n#             nn_cs,nn_decs=Cubic.dis_cubic(geo['nn'],2,40,36)\n#             self.geo_cc.append([cc_cs,cc_decs])\n#             self.geo_pp.append([pp_cs,pp_decs])\n#             self.geo_nn.append([nn_cs,nn_decs])\n            \n#             L = self.L\n#             cc_coefs_np  = np.stack([[cc_cs[i,j].c for j in range(L)] for i in range(L)], axis=0)\n#             cc_knots_np  = np.stack([[cc_cs[i,j].x for j in range(L)] for i in range(L)], axis=0)\n#             self.cs_coefs['cc'].append(torch.from_numpy(cc_coefs_np).to(device))\n#             self.cs_knots['cc'].append(torch.from_numpy(cc_knots_np).to(device))\n            \n#             pp_coefs_np  = np.stack([[pp_cs[i,j].c for j in range(L)] for i in range(L)], axis=0)\n#             pp_knots_np  = np.stack([[pp_cs[i,j].x for j in range(L)] for i in range(L)], axis=0)\n#             self.cs_coefs['pp'].append(torch.from_numpy(pp_coefs_np).to(device))\n#             self.cs_knots['pp'].append(torch.from_numpy(pp_knots_np).to(device))\n            \n#             nn_coefs_np  = np.stack([[nn_cs[i,j].c for j in range(L)] for i in range(L)], axis=0)\n#             nn_knots_np  = np.stack([[nn_cs[i,j].x for j in range(L)] for i in range(L)], axis=0)\n#             self.cs_coefs['nn'].append(torch.from_numpy(nn_coefs_np).to(device))\n#             self.cs_knots['nn'].append(torch.from_numpy(nn_knots_np).to(device))\n\n\n#     def compute_bb_clash(self,coor,other_coor):\n#         com_coor = torch.cat([coor,other_coor],dim=1)\n#         com_dis  = (com_coor[:,None,:,None,:] - com_coor[None,:,None,:,:]).norm(dim=-1)\n#         dynamicmask2_vdw= (com_dis <= 3.15) * (self.clash_mask)\n#         vdw_dynamic = Potential.LJpotential(com_dis[dynamicmask2_vdw],3.15)\n#         return vdw_dynamic.sum()*self.config['weight_vdw']\n\n#     def compute_full_clash(self,coor,other_coor,side_coor):\n#         com_coor = torch.cat([coor[:,:2],other_coor,side_coor],dim=1)\n#         com_dis  = (com_coor[:,None,:,None,:] - com_coor[None,:,None,:,:]).norm(dim=-1)\n#         dynamicmask2_vdw= (com_dis <= 2.5) * (self.clash_mask)\n#         vdw_dynamic = Potential.LJpotential(com_dis[dynamicmask2_vdw],2.5)\n#         return vdw_dynamic.sum()*self.config['weight_vdw']\n\n\n#     def _cubic_pair_energy(self, atom_map, geo_cs, geo_confimask, weight_key):\n#         \"\"\"General cubic-spline energy for CC/PP/NN pairs.\"\"\"\n#         min_dis, max_dis, bin_num = 2, 40, 36\n#         dev = atom_map.device\n#         upper_th = max_dis - ((max_dis - min_dis) / bin_num) * 0.5\n#         lower_th = 2.5\n#         total = torch.zeros((), device=dev, dtype=torch.double)\n#         spline_key   = weight_key.split('_')[1]  # 'cc', 'pp', or 'nn'\n#         coeffs_list  = self.cs_coefs[spline_key]\n#         knots_list   = self.cs_knots[spline_key]\n#         for block_idx, mask_block in enumerate(geo_confimask):\n#             mask = (atom_map <= upper_th) & mask_block & self.fullmask & (atom_map >= lower_th)\n#             idx = mask.nonzero(as_tuple=True)\n#             if idx[0].numel() > 1:\n#                 coef  = coeffs_list[block_idx][idx]\n#                 knots = knots_list[block_idx][idx]\n#                 part1 = Potential.cubic_distance(atom_map[mask], coef, knots, min_dis, max_dis, bin_num).sum() * self.config[weight_key] * 0.5\n#             else:\n#                 part1 = torch.zeros((), device=dev)\n#             part2 = ((atom_map <= lower_th) & mask_block & self.fullmask).sum() * self.config[weight_key]\n#             total = total + part1 + part2\n#         return total\n\n#     def compute_cc_energy(self, coor):\n#         atom_map = operations.pair_distance(coor[:,1], coor[:,1])\n#         return self._cubic_pair_energy(atom_map, self.geo_cc, self.geo_confimask_cc, 'weight_cc')\n    \n#     def compute_pp_energy(self, coor):\n#         atom_map = operations.pair_distance(coor[:,0], coor[:,0])\n#         return self._cubic_pair_energy(atom_map, self.geo_pp, self.geo_confimask_pp, 'weight_pp')\n    \n#     def compute_nn_energy(self, coor):\n#         atom_map = operations.pair_distance(coor[:,-1], coor[:,-1])\n#         return self._cubic_pair_energy(atom_map, self.geo_nn, self.geo_confimask_nn, 'weight_nn')\n\n#     def compute_pccp_energy(self,coor):\n#         p_atoms=coor[:,0]\n#         c_atoms=coor[:,1]\n#         pccpmap=operations.dihedral( p_atoms[self.pccpi], c_atoms[self.pccpi], c_atoms[self.pccpj] ,p_atoms[self.pccpj]                  )\n#         neg_log = Potential.cubic_torsion(pccpmap,self.pccp_coe,self.pccp_x,36)\n#         return neg_log.sum()*self.config['weight_pccp']\n\n#     def compute_cnnc_energy(self,coor):\n#         n_atoms=coor[:,-1]\n#         c_atoms=coor[:,1]\n#         pccpmap=operations.dihedral( c_atoms[self.cnnci], n_atoms[self.cnnci], n_atoms[self.cnncj] ,c_atoms[self.cnncj]                  )\n#         neg_log = Potential.cubic_torsion(pccpmap,self.cnnc_coe,self.cnnc_x,36)\n#         return neg_log.sum()*self.config['weight_cnnc']\n\n#     def compute_pnnp_energy(self,coor):\n#         n_atoms=coor[:,-1]\n#         p_atoms=coor[:,0]\n#         pccpmap=operations.dihedral( p_atoms[self.pnnpi], n_atoms[self.pnnpi], n_atoms[self.pnnpj] ,p_atoms[self.pnnpj]                  )\n#         neg_log = Potential.cubic_torsion(pccpmap,self.pnnp_coe,self.pnnp_x,36)\n#         return neg_log.sum()*self.config['weight_pnnp']\n\n#     def compute_pcc_energy(self,coor):\n#         p_atoms=coor[:,1]\n#         c_atoms=coor[:,2]\n#         pccmap=operations.angle( p_atoms[self.pcci], c_atoms[self.pcci], c_atoms[self.pccj]                   )\n#         neg_log = Potential.cubic_angle(pccmap,self.pcc_coe,self.pcc_x,12)\n#         return neg_log.sum()*self.config['weight_pcc']\n\n#     def compute_fape_energy(self,coor,ep=1e-3,epmax=20):\n#         energy= 0\n#         for tx in self.tx2ds:\n#             px_mean = coor[:,[1]]\n#             p_rot   = operations.rigidFrom3Points(coor)\n#             p_tran  = px_mean[:,0]\n#             pred_x2 = coor[:,None,:,:] - p_tran[None,:,None,:] # Lx Lrot N , 3\n#             pred_x2 = torch.einsum('ijnd,jde->ijne',pred_x2,p_rot.transpose(-1,-2)) # transpose should be equal to inverse\n#             errmap=torch.sqrt( ((pred_x2 - tx)**2).sum(dim=-1) + ep )\n#             energy = energy + torch.sum(  torch.clamp(errmap,max=epmax)        )\n#         return energy * self.config['weight_fape']\n\n#     def compute_bond_energy(self,coor,other_coor):\n#         # 3.87\n#         o3 = other_coor[:-1,-2]\n#         p  = coor[1:,0]\n#         dis = (o3-p).norm(dim=-1)\n#         energy = ((dis-1.607)**2).sum()\n#         return energy * self.config['weight_bond']\n\n#     def tooth_func(self,errmap, ep = 0.05):\n#         return -1/(errmap/10+ep) + (1/ep)\n\n#     def reweight_func(self,ww):\n#         reweighting = torch.pow(ww,self.config['pair_weight_power'])\n#         reweighting[ww < self.config['pair_weight_min']] = 0\n#         return reweighting\n\n#     def compute_fape_energy_fromquat(self,x,coor,ep=1e-6,epmax=100):\n#         energy= 0\n#         p_rot,px_mean = a2b.Non2rot(x[:,:9],x.shape[0]),x[:,9:]\n#         pred_x2 = coor[:,None,:,:] - px_mean[None,:,None,:] # Lx Lrot N , 3\n#         pred_x2 = torch.einsum('ijnd,jde->ijne',pred_x2,p_rot.transpose(-1,-2)) # transpose should be equal to inverse\n#         for tx,weightplddt in zip(self.tx2ds,self.pair):\n\n#             tamplate_dist_map = torch.min( tx.norm(dim=-1), dim=2   )[0]\n#             errmap=torch.sqrt( ((pred_x2 - tx)**2).sum(dim=-1) + ep ) \n#             energy = energy + torch.sum( ( (torch.clamp(errmap,max=self.config['FAPE_max'])**self.config['pair_error_power'])  * self.reweight_func(weightplddt[...,None]) * self.local_weight[...,None] )[tamplate_dist_map>self.config['pair_rest_min_dist']]    )\n\n#         return energy * self.config['weight_fape']\n\n\n#     def energy(self,rama):\n#         coor=a2b.quat2b(self.basex,rama[:,9:])\n#         other_coor = a2b.quat2b(self.otherx,rama[:,9:])\n#         side_coor = a2b.quat2b(self.sidex,torch.cat([rama[:,:9],coor[:,-1]],dim=-1))\n        \n#         if self.config['weight_cc']>0:\n#             E_cc= self.compute_cc_energy(coor) / len(self.rets)\n#         else:\n#             E_cc=0\n#         if self.config['weight_pp']>0:\n#             E_pp= self.compute_pp_energy(coor) / len(self.rets)\n#         else:\n#             E_pp=0\n#         if self.config['weight_nn']>0:\n#             E_nn= self.compute_nn_energy(coor) / len(self.rets)\n#         else:\n#             E_nn=0\n\n#         if self.config['weight_pccp']>0:\n#             E_pccp= self.compute_pccp_energy(coor) / len(self.rets)\n#         else:\n#             E_pccp=0\n\n#         if self.config['weight_cnnc']>0:\n#             E_cnnc= self.compute_cnnc_energy(coor)  / len(self.rets)\n#         else:\n#             E_cnnc=0\n\n#         if self.config['weight_pnnp']>0:\n#             E_pnnp= self.compute_pnnp_energy(coor) / len(self.rets)\n#         else:\n#             E_pnnp=0\n\n#         if self.config['weight_vdw']>0:\n#             E_vdw= self.compute_full_clash(coor,other_coor,side_coor)\n#         else:\n#             E_vdw=0\n\n#         if self.config['weight_fape']>0:\n#             E_fape= self.compute_fape_energy_fromquat(rama[:,9:],coor) / len(self.rets)\n#         else:\n#             E_fape=0\n#         if self.config['weight_bond']>0:\n#             E_bond= self.compute_bond_energy(coor,other_coor)\n#         else:\n#             E_bond=0\n#         return  E_vdw + E_fape + E_bond + E_pp + E_cc + E_nn + E_pccp + E_cnnc + E_pnnp\n\n\n#     def obj_func_grad_np(self,rama_):\n#         rama=torch.DoubleTensor(rama_)\n#         rama.requires_grad=True\n#         if rama.grad:\n#             rama.grad.zero_()\n#         f=self.energy(rama.view(self.L,21))*Scale_factor\n#         grad_value=autograd.grad(f,rama)[0]\n#         return grad_value.data.numpy().astype(np.float64)\n    \n#     def obj_func_np(self,rama_):\n#         rama=torch.DoubleTensor(rama_)\n#         rama=rama.view(self.L,21)\n#         with torch.no_grad():\n#             f=self.energy(rama)*Scale_factor\n#             return f.item()\n\n\n#     def foldning(self):\n#         ilter = self.init_ret\n#         # 1) get initial quaternions (double precision)\n#         try:\n#             init_q = self.init_quat(ilter).double()\n#         except:\n#             init_q = self.init_quat_safe(ilter).double()\n\n#         # 2) move to target device (GPU if available), enable grad\n#         param = init_q.to(device).clone().detach().requires_grad_(True)\n\n#         # 3) set up PyTorch LBFGS optimizer over `param`\n#         optimizer = opt.LBFGS(\n#             [param],\n#             max_iter=self.config.get('max_iter', 300),\n#             tolerance_grad=1e-6,\n#             tolerance_change=1e-9,\n#             history_size=10,\n#             line_search_fn='strong_wolfe'\n#         )\n\n#         # 4) define the “closure” that LBFGS will call to reevaluate loss + gradients\n#         def closure():\n#             optimizer.zero_grad()                                 # clear old grads\n#             E = self.energy(param.view(self.L,21)) * Scale_factor # compute ∂E/∂param\n#             E.backward()\n#             return E\n\n#         # 5) run LBFGS until convergence (it calls closure repeatedly)\n#         optimizer.step(closure)\n\n#         # 6) write out final PDB\n#         final_energy = self.energy(param.view(self.L,21)).item()\n#         self.outpdb(param, self.saveprefix + '.pdb', energystr=str(final_energy))\n\n\n#     def outpdb(self,rama,savefile,start=0,end=10000,energystr=''):\n#         # bring baseframes and quaternion data onto CPU to prevent device mismatch\n#         basex_cpu = self.basex.detach().cpu()\n#         otherx_cpu = self.otherx.detach().cpu()\n#         sidex_cpu = self.sidex.detach().cpu()\n#         shaped_rama = rama.view(self.L,21).detach().cpu()\n#         # compute backbone and other coords\n#         coor_np = a2b.quat2b(basex_cpu, shaped_rama[:,9:]).detach().cpu().numpy()\n#         other_np = a2b.quat2b(otherx_cpu, shaped_rama[:,9:]).detach().cpu().numpy()\n#         coor = torch.FloatTensor(coor_np)\n#         # compute side atom coords\n#         side_coor_NP = a2b.quat2b(sidex_cpu, torch.cat([shaped_rama[:,:9], coor[:,-1]], dim=-1)).detach().cpu().numpy()\n        \n#         Atom_name=[' P  ',\" C4'\",' N1 ']\n#         Other_Atom_name = [\" O5'\",\" C5'\",\" C3'\",\" O3'\",\" C1'\"]\n#         other_last_name = ['O',\"C\",\"C\",\"O\",\"C\"]\n\n#         side_atoms=         [' N1 ',' C2 ',' O2 ',' N2 ',' N3 ',' N4 ',' C4 ',' O4 ',' C5 ',' C6 ',' O6 ',' N6 ',' N7 ',' N8 ',' N9 ']\n#         side_last_name =    ['N',      \"C\",   \"O\",   \"N\",   \"N\",   'N',   'C',   'O',   'C',   'C',   'O',   'N',    'N', 'N','N']\n\n#         base_dict = rigid.base_table()\n#         last_name=['P','C','N']\n#         wstr=[f'REMARK {str(energystr)}']\n#         templet='%6s%5d %4s %3s %1s%4d    %8.3f%8.3f%8.3f%6.2f%6.2f          %2s%2s'\n#         count=1\n#         for i in range(self.L):\n#             if self.seq[i] in ['a','g','A','G']:\n#                 Atom_name = [' P  ',\" C4'\",' N9 ']\n#                 #atoms = ['P','C4']\n\n#             elif self.seq[i] in ['c','u','C','U']:\n#                 Atom_name = [' P  ',\" C4'\",' N1 ']\n#             for j in range(coor_np.shape[1]):\n#                 outs=('ATOM  ',count,Atom_name[j],self.seq[i],'A',i+1,coor_np[i][j][0],coor_np[i][j][1],coor_np[i][j][2],0,0,last_name[j],'')\n#                 if i>=start-1 and i < end:\n#                     wstr.append(templet % outs)\n#                     count+=1\n\n#             for j in range(other_np.shape[1]):\n#                 outs=('ATOM  ',count,Other_Atom_name[j],self.seq[i],'A',i+1,other_np[i][j][0],other_np[i][j][1],other_np[i][j][2],0,0,other_last_name[j],'')\n#                 if i>=start-1 and i < end:\n#                     wstr.append(templet % outs)\n#                     count+=1\n            \n#         wstr='\\n'.join(wstr)\n#         wfile=open(savefile,'w')\n#         wfile.write(wstr)\n#         wfile.close()\n    \n#     def outpdb_coor(self,coor_np,savefile,start=0,end=1000,energystr=''):\n#         Atom_name=[' P  ',\" C4'\",' N1 ']\n#         last_name=['P','C','N']\n#         wstr=[f'REMARK {str(energystr)}']\n#         templet='%6s%5d %4s %3s %1s%4d    %8.3f%8.3f%8.3f%6.2f%6.2f          %2s%2s'\n#         count=1\n#         for i in range(self.L):\n#             if self.seq[i] in ['a','g','A','G']:\n#                 Atom_name = [' P  ',\" C4'\",' N9 ']\n\n#             elif self.seq[i] in ['c','u','C','U']:\n#                 Atom_name = [' P  ',\" C4'\",' N1 ']\n#             for j in range(coor_np.shape[1]):\n#                 outs=('ATOM  ',count,Atom_name[j],self.seq[i],'A',i+1,coor_np[i][j][0],coor_np[i][j][1],coor_np[i][j][2],0,0,last_name[j],'')\n#                 if i>=start-1 and i < end:\n#                     wstr.append(templet % outs)\n#                 count+=1\n            \n#         wstr='\\n'.join(wstr)\n#         wfile=open(savefile,'w')\n#         wfile.write(wstr)\n#         wfile.close()\n\n\n#     def init_quat(self,ii):\n#         x = torch.rand([self.L,21])\n#         x[:,18:] = self.txs[ii].mean(dim=1)\n#         init_coor = self.txs[ii]\n#         biasq = torch.mean(init_coor,dim=1,keepdim=True)\n#         q = init_coor - biasq\n#         m = torch.einsum('bnz,bny->bzy',self.basex,q).reshape([self.L,-1])\n#         x[:,:9] = x[:,9:18] = m\n#         x.requires_grad_()\n#         return x\n\n#     def init_quat_safe(self,ii):\n#         x = torch.rand([self.L,21])\n#         x[:,18:] = self.txs[ii].mean(dim=1)\n#         init_coor = self.txs[ii]\n#         biasq = torch.mean(init_coor,dim=1,keepdim=True)\n#         q = init_coor - biasq + torch.rand([self.L,3,3])\n#         m = (torch.einsum('bnz,bny->bzy',self.basex,q) + torch.eye(3)[None,:,:]).reshape([self.L,-1])\n#         x[:,:9] = x[:,9:18] = m\n#         x.requires_grad_()\n#         return x\n\n\n# if __name__ == '__main__': \n\n#     fastafile=sys.argv[1]\n#     saveprefix=sys.argv[2]\n#     retdirs  =sys.argv[3]\n#     ret_score = sys.argv[4]\n#     foldconfig = sys.argv[5]\n\n#     savepare = os.path.dirname(saveprefix)\n#     if not os.path.isdir(savepare):\n#         os.makedirs(savepare)\n\n#     num_of_models = readconfig(foldconfig)['num_of_models']\n\n#     score_dict = readconfig(ret_score)\n#     sorted_items = sorted(score_dict.items(), key=lambda x: x[1])\n#     lowest_n_keys = [item[0] for item in sorted_items][:num_of_models]\n#     bestkey = lowest_n_keys[0] + ''\n#     print(\"Before sort:\", lowest_n_keys)\n#     lowest_n_keys.sort()\n#     print(\"After sort:\", lowest_n_keys)\n#     bestindex = lowest_n_keys.index(bestkey)\n\n#     current_ret = bestkey\n#     retfiles = [os.path.join(retdirs, afile) for afile in lowest_n_keys]\n#     stru = Structure(fastafile, retfiles, saveprefix + '_from_' + current_ret, bestindex, foldconfig)\n#     stru.foldning()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T13:02:03.159207Z","iopub.execute_input":"2026-02-19T13:02:03.15949Z","iopub.status.idle":"2026-02-19T13:02:03.175109Z","shell.execute_reply.started":"2026-02-19T13:02:03.159471Z","shell.execute_reply":"2026-02-19T13:02:03.174388Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %%writefile /kaggle/working/DRfold2/PotentialFold/Selection.py\n# #! /nfs/amino-home/liyangum/miniconda3/bin/python\n# import numpy\n# import torch\n# import torch.autograd as autograd\n# import numpy as np \n\n# import random\n# import Cubic, Potential\n# import operations\n# import os, json, sys\n\n# import a2b, rigid\n# import torch.optim as opt\n# from torch.nn.parameter import Parameter\n# import torch.nn as nn\n# import math\n# from scipy.optimize import fmin_l_bfgs_b,fmin_cg,fmin_bfgs\n# from scipy.optimize import minimize\n# import lbfgs_rosetta\n# import pickle\n# import shutil\n\n# torch.manual_seed(6)\n# torch.set_num_threads(4)\n# np.random.seed(9)\n# random.seed(9)\n\n# Scale_factor = 1.0\n# USEGEO = False\n\n# device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# def readconfig(configfile=''):\n#     config=[]\n#     expdir=os.path.dirname(os.path.abspath(__file__))\n#     if configfile=='':\n#         configfile=os.path.join(expdir,'lib','ddf.json')\n#     config=json.load(open(configfile,'r'))\n#     return config \n\n    \n# class Structure:\n#     def __init__(self, fastafile, geofiles, foldconfig, saveprefix):\n#         # Load Configuration and Inputs\n#         self.config = readconfig(foldconfig)\n#         self.seqfile = fastafile\n#         self.foldconfig = foldconfig\n#         self.geofiles = geofiles\n\n#         # Load Model Results\n#         self.rets = [pickle.load(open(refile, 'rb')) for refile  in geofiles]\n        \n#         # Extract Coordinates\n#         self.txs = []\n#         for ret in self.rets:\n#             self.txs.append(torch.from_numpy(ret['coor']).double().to(device))\n        \n#         # Handle Geometrical Data\n#         self.handle_geo()\n\n#         # Extract pLDDT Scores\n#         self.pair = []\n#         for ret in self.rets:\n#             self.pair.append( torch.from_numpy(ret['plddt']).double().to(device))\n        \n#         # Store Output and Sequence Info\n#         self.saveprefix = saveprefix\n#         self.seq = open(fastafile).readlines()[1].strip()\n#         self.L = len(self.seq)\n        \n#         # Load Reference Arrays for Structure Construction\n#         basenpy = np.load(os.path.join(os.path.dirname(os.path.abspath(__file__)), 'lib', 'base.npy'))\n#         self.basex = operations.Get_base(self.seq, basenpy).double().to(device)\n        \n#         othernpy = np.load(os.path.join(os.path.dirname(os.path.abspath(__file__)), 'lib', 'other2.npy'))\n#         self.otherx = operations.Get_base(self.seq, othernpy).double().to(device)\n        \n#         sidenpy = np.load(os.path.join(os.path.dirname(os.path.abspath(__file__)), 'lib', 'side.npy'))\n#         self.sidex = operations.Get_base(self.seq, sidenpy).double().to(device)        \n        \n#         # Initialize Masks, Parameters, and FAPE\n#         self.init_mask()\n#         self.init_paras()\n#         self._init_fape()\n    \n\n#     def _init_fape(self):\n#         self.tx2ds = []\n#         for tx in self.txs:\n#             true_rot, true_trans = operations.Kabsch_rigid(self.basex, tx[:, 0], tx[:, 1], tx[:, 2])\n#             true_x2 = tx[:, None, :, :] - true_trans[None, :, None, :]\n#             true_x2 = torch.einsum('ijnd,jde->ijne', true_x2, true_rot.transpose(-1,-2))\n#             self.tx2ds.append(true_x2)\n    \n\n#     def handle_geo(self):\n#         oldkeys = ['dist_p', 'dist_c', 'dist_n']\n#         newkeys = ['pp', 'cc', 'nn']\n#         self.geos = []\n#         geo = {'pp':0, 'cc':0, 'nn':0}\n        \n#         for ret in self.rets:    \n#             for nk, ok in zip(newkeys, oldkeys):\n#                 geo[nk] = geo[nk] + (ret[ok].astype(np.float64) /(len(self.rets)))\n#         self.geos.append(geo)\n\n\n#     def init_mask(self):\n#         halfmask=np.zeros([self.L,self.L])\n#         fullmask=np.zeros([self.L,self.L])\n#         for i in range(self.L):\n#             for j in range(i+1,self.L):\n#                 halfmask[i,j]=1\n#                 fullmask[i,j]=1\n#                 fullmask[j,i]=1\n#         self.halfmask=torch.DoubleTensor(halfmask) > 0.5\n#         self.fullmask=torch.DoubleTensor(fullmask) > 0.5\n#         self.clash_mask = torch.zeros([self.L,self.L,22,22])\n#         for i in range(self.L):\n#             for j in range(i+1,self.L):\n#                 self.clash_mask[i,j]=1\n\n#         for i in range(self.L):\n#              self.clash_mask[i,i,:6,7:]=1\n\n#         for i in range(self.L-1):\n#             self.clash_mask[i,i+1,:,0]=0\n#             self.clash_mask[i,i+1,0,:]=0\n#             self.clash_mask[i,i+1,:,5]=0\n#             self.clash_mask[i,i+1,5,:]=0\n\n#         self.side_mask = rigid.side_mask(self.seq)\n#         self.side_mask = self.side_mask[:,None,:,None] * self.side_mask[None,:,None,:]\n#         self.clash_mask = (self.clash_mask > 0.5) * (self.side_mask > 0.5)\n\n#         self.geo_confimask_cc = []\n#         self.geo_confimask_pp = []\n#         self.geo_confimask_nn = []\n#         for geo in self.geos:\n#             confimask_cc = torch.DoubleTensor(geo['cc'][:,:,-1]) < 0.5\n#             confimask_pp = torch.DoubleTensor(geo['pp'][:,:,-1]) < 0.5\n#             confimask_nn = torch.DoubleTensor(geo['nn'][:,:,-1]) < 0.5\n#             self.geo_confimask_cc.append(confimask_cc)\n#             self.geo_confimask_pp.append(confimask_pp)\n#             self.geo_confimask_nn.append(confimask_nn)\n\n#         # Move masks and confimasks to the GPU/CPU device\n#         self.halfmask = self.halfmask.to(device)\n#         self.fullmask = self.fullmask.to(device)\n#         self.clash_mask = self.clash_mask.to(device)\n#         self.side_mask = self.side_mask.to(device)\n#         # geo_confimasks are lists\n#         self.geo_confimask_cc = [m.to(device) for m in self.geo_confimask_cc]\n#         self.geo_confimask_pp = [m.to(device) for m in self.geo_confimask_pp]\n#         self.geo_confimask_nn = [m.to(device) for m in self.geo_confimask_nn]\n\n\n#     def init_paras(self):\n#         self.geo_cc = []\n#         self.geo_pp = []\n#         self.geo_nn = []\n#         self.cs_coefs = {'cc': [], 'pp': [], 'nn': []}\n#         self.cs_knots = {'cc': [], 'pp': [], 'nn': []}\n#         for geo in self.geos:\n#             cc_cs, cc_decs = Cubic.dis_cubic(geo['cc'], 2, 40, 36)\n#             pp_cs, pp_decs = Cubic.dis_cubic(geo['pp'], 2, 40, 36)\n#             nn_cs, nn_decs = Cubic.dis_cubic(geo['nn'], 2, 40, 36)\n#             self.geo_cc.append([cc_cs, cc_decs])\n#             self.geo_pp.append([pp_cs, pp_decs])\n#             self.geo_nn.append([nn_cs, nn_decs])\n#             L = self.L\n#             cc_coefs_np = np.stack([[cc_cs[i,j].c for j in range(L)] for i in range(L)], axis=0)\n#             cc_knots_np = np.stack([[cc_cs[i,j].x for j in range(L)] for i in range(L)], axis=0)\n#             self.cs_coefs['cc'].append(torch.from_numpy(cc_coefs_np).to(device))\n#             self.cs_knots['cc'].append(torch.from_numpy(cc_knots_np).to(device))\n#             pp_coefs_np = np.stack([[pp_cs[i,j].c for j in range(L)] for i in range(L)], axis=0)\n#             pp_knots_np = np.stack([[pp_cs[i,j].x for j in range(L)] for i in range(L)], axis=0)\n#             self.cs_coefs['pp'].append(torch.from_numpy(pp_coefs_np).to(device))\n#             self.cs_knots['pp'].append(torch.from_numpy(pp_knots_np).to(device))\n#             nn_coefs_np = np.stack([[nn_cs[i,j].c for j in range(L)] for i in range(L)], axis=0)\n#             nn_knots_np = np.stack([[nn_cs[i,j].x for j in range(L)] for i in range(L)], axis=0)\n#             self.cs_coefs['nn'].append(torch.from_numpy(nn_coefs_np).to(device))\n#             self.cs_knots['nn'].append(torch.from_numpy(nn_knots_np).to(device))\n     \n\n#     def _cubic_pair_energy(self, atom_map, geo_cs, geo_confimask, weight_key):\n#         \"\"\"General cubic-spline energy for CC/PP/NN pairs.\"\"\"\n#         min_dis, max_dis, bin_num = 2, 40, 36\n#         dev = atom_map.device\n#         upper_th = max_dis - ((max_dis - min_dis) / bin_num) * 0.5\n#         lower_th = 2.5\n#         total = torch.zeros((), device=dev, dtype=torch.double)\n#         spline_key = weight_key.split('_')[1]\n#         coeffs_list = self.cs_coefs[spline_key]\n#         knots_list = self.cs_knots[spline_key]\n#         for block_idx, mask_block in enumerate(geo_confimask):\n#             mask = (atom_map <= upper_th) & mask_block & self.fullmask & (atom_map >= lower_th)\n#             idx = mask.nonzero(as_tuple=True)\n#             if idx[0].numel() > 1:\n#                 coef = coeffs_list[block_idx][idx]\n#                 knots = knots_list[block_idx][idx]\n#                 part1 = Potential.cubic_distance(atom_map[mask], coef, knots, min_dis, max_dis, bin_num).sum() * self.config[weight_key] * 0.5\n#             else:\n#                 part1 = torch.zeros((), device=dev, dtype=torch.double)\n#             part2 = ((atom_map <= lower_th) & mask_block & self.fullmask).sum() * self.config[weight_key]\n#             total = total + part1 + part2\n#         return total\n\n#     # GPU-friendly torsion and angle energy helpers\n#     def _cubic_torsion_energy(self, atom_map, coef, x_vals, weight_key, num_bin):\n#         energy = Potential.cubic_torsion(atom_map, coef, x_vals, num_bin)\n#         return energy.sum() * self.config[weight_key]\n\n#     def _cubic_angle_energy(self, atom_map, coef, x_vals, weight_key, num_bin):\n#         energy = Potential.cubic_angle(atom_map, coef, x_vals, num_bin)\n#         return energy.sum() * self.config[weight_key]\n\n#     def compute_cc_energy(self, coor):\n#         atom_map = operations.pair_distance(coor[:,1], coor[:,1])\n#         return self._cubic_pair_energy(atom_map, self.geo_cc, self.geo_confimask_cc, 'weight_cc')\n\n#     def compute_pp_energy(self, coor):\n#         atom_map = operations.pair_distance(coor[:,0], coor[:,0])\n#         return self._cubic_pair_energy(atom_map, self.geo_pp, self.geo_confimask_pp, 'weight_pp')\n\n#     def compute_nn_energy(self, coor):\n#         atom_map = operations.pair_distance(coor[:,-1], coor[:,-1])\n#         return self._cubic_pair_energy(atom_map, self.geo_nn, self.geo_confimask_nn, 'weight_nn')\n\n#     def compute_pccp_energy(self, coor):\n#         # P-C-C-P dihedral energy on GPU\n#         p = coor[:, 0]\n#         c = coor[:, 1]\n#         dia = operations.dihedral(\n#             p[self.pccpi], c[self.pccpi], c[self.pccpj], p[self.pccpj]\n#         )\n#         return self._cubic_torsion_energy(dia, self.pccp_coe, self.pccp_x, 'weight_pccp', 36)\n\n#     def compute_cnnc_energy(self, coor):\n#         # C-N-N-C dihedral energy on GPU\n#         n = coor[:, -1]\n#         c = coor[:, 1]\n#         dia = operations.dihedral(\n#             c[self.cnnci], n[self.cnnci], n[self.cnncj], c[self.cnncj]\n#         )\n#         return self._cubic_torsion_energy(dia, self.cnnc_coe, self.cnnc_x, 'weight_cnnc', 36)\n\n#     def compute_pnnp_energy(self, coor):\n#         # P-N-N-P dihedral energy on GPU\n#         n = coor[:, -1]\n#         p = coor[:, 0]\n#         dia = operations.dihedral(\n#             p[self.pnnpi], n[self.pnnpi], n[self.pnnpj], p[self.pnnpj]\n#         )\n#         return self._cubic_torsion_energy(dia, self.pnnp_coe, self.pnnp_x, 'weight_pnnp', 36)\n\n#     def compute_pcc_energy(self, coor):\n#         # P-C-C angle energy on GPU\n#         p = coor[:, 1]\n#         c = coor[:, 2]\n#         ang = operations.angle(\n#             p[self.pcci], c[self.pcci], c[self.pccj]\n#         )\n#         return self._cubic_angle_energy(ang, self.pcc_coe, self.pcc_x, 'weight_pcc', 12)\n\n#     def compute_fape_energy(self,coor,ep=1e-3,epmax=20):\n#         energy= 0\n#         for tx in self.tx2ds:\n#             px_mean = coor[:,[1]]\n#             p_rot   = operations.rigidFrom3Points(coor)\n#             p_tran  = px_mean[:,0]\n#             pred_x2 = coor[:,None,:,:] - p_tran[None,:,None,:] # Lx Lrot N , 3\n#             pred_x2 = torch.einsum('ijnd,jde->ijne',pred_x2,p_rot.transpose(-1,-2)) # transpose should be equal to inverse\n#             errmap=torch.sqrt( ((pred_x2 - tx)**2).sum(dim=-1) + ep )\n#             energy = energy + torch.sum(  torch.clamp(errmap,max=epmax)        )\n#         return energy * self.config['weight_fape']\n\n#     def compute_bond_energy(self,coor,other_coor):\n#         # 3.87\n#         o3 = other_coor[:-1,-2]\n#         p  = coor[1:,0]\n#         dis = (o3-p).norm(dim=-1)\n#         energy = ((dis-1.607)**2).sum()\n#         return energy * self.config['weight_bond']\n\n#     def tooth_func(self,errmap, ep = 0.05):\n#         return -1/(errmap/10+ep) + (1/ep)\n    \n#     def reweight_func(self,ww):\n#         reweighting = torch.pow(ww,self.config['pair_weight_power'])\n#         reweighting[ww < self.config['pair_weight_min']] = 0\n#         return reweighting\n    \n#     def compute_fape_energy_fromquat(self,x,coor,ep=1e-6,epmax=100):\n#         energy= 0\n#         p_rot,px_mean = a2b.Non2rot(x[:,:9],x.shape[0]),x[:,9:]\n#         pred_x2 = coor[:,None,:,:] - px_mean[None,:,None,:] # Lx Lrot N , 3\n#         pred_x2 = torch.einsum('ijnd,jde->ijne',pred_x2,p_rot.transpose(-1,-2)) # transpose should be equal to inverse\n\n#         for tx,weightplddt in zip(self.tx2ds,self.pair):\n#             tamplate_dist_map = torch.min( tx.norm(dim=-1), dim=2   )[0]\n#             errmap=torch.sqrt( ((pred_x2 - tx)**2).sum(dim=-1) + ep ) \n#             energy = energy + torch.sum( ( (torch.clamp(errmap,max=self.config['FAPE_max'])**self.config['pair_error_power'])  * self.reweight_func(weightplddt[...,None]) )[tamplate_dist_map>self.config['pair_rest_min_dist']]    )\n\n#         return energy * self.config['weight_fape']\n    \n#     def compute_fape_energy_fromcoor(self,coor,ep=1e-6,epmax=100):\n#         energy= 0\n        \n#         p_rot,px_mean = operations.Kabsch_rigid(self.basex,coor[:,0],coor[:,1],coor[:,2])\n#         pred_x2 = coor[:,None,:,:] - px_mean[None,:,None,:] # Lx Lrot N , 3\n#         pred_x2 = torch.einsum('ijnd,jde->ijne',pred_x2,p_rot.transpose(-1,-2)) # transpose should be equal to inverse\n        \n#         for tx,weightplddt in zip(self.tx2ds,self.pair):\n#             tamplate_dist_map = torch.min( tx.norm(dim=-1), dim=2   )[0]\n#             errmap=torch.sqrt( ((pred_x2 - tx)**2).sum(dim=-1) + ep ) \n#             energy = energy + torch.sum( ( (torch.clamp(errmap,max=self.config['FAPE_max'])**self.config['pair_error_power'])  * self.reweight_func(weightplddt[...,None]) )[tamplate_dist_map>self.config['pair_rest_min_dist']]    )\n\n#         return energy * self.config['weight_fape']\n    \n    \n#     def energy(self, rama):\n#         coor = a2b.quat2b(self.basex, rama[:, 9:])\n#         other_coor = a2b.quat2b(self.otherx, rama[:, 9:])\n#         side_coor = a2b.quat2b(self.sidex, torch.cat([rama[:, :9], coor[:, -1]], dim=-1))\n\n#         E_cc = self.compute_cc_energy(coor) / len(self.geofiles) if self.config['weight_cc'] > 0 else 0\n#         E_pp = self.compute_pp_energy(coor) / len(self.geofiles) if self.config['weight_pp'] > 0 else 0\n#         E_nn = self.compute_nn_energy(coor) / len(self.geofiles) if self.config['weight_nn'] > 0 else 0\n#         E_pccp = self.compute_pccp_energy(coor) / len(self.geofiles) if self.config['weight_pccp'] > 0 else 0\n#         E_cnnc = self.compute_cnnc_energy(coor) / len(self.geofiles) if self.config['weight_cnnc'] > 0 else 0\n#         E_pnnp = self.compute_pnnp_energy(coor) / len(self.geofiles) if self.config['weight_pnnp'] > 0 else 0\n#         E_vdw = self.compute_full_clash(coor, other_coor, side_coor) if self.config['weight_vdw'] > 0 else 0\n#         E_fape = self.compute_fape_energy_fromquat(rama[:, 9:], coor) / len(self.geofiles) if self.config['weight_fape'] > 0 else 0\n#         E_bond = self.compute_bond_energy(coor, other_coor) if self.config['weight_bond'] > 0 else 0\n\n#         return E_vdw + E_fape + E_bond + E_pp + E_cc + E_nn + E_pccp + E_cnnc + E_pnnp\n\n\n#     def energy_from_coor(self, coor):\n#         E_cc = self.compute_cc_energy(coor) if self.config['weight_cc'] > 0 else 0\n#         E_pp = self.compute_pp_energy(coor) if self.config['weight_pp'] > 0 else 0\n#         E_nn = self.compute_nn_energy(coor) if self.config['weight_nn'] > 0 else 0\n#         E_fape = (self.compute_fape_energy_fromcoor(coor) / len(self.geofiles)) if self.config['weight_fape'] > 0 else 0\n#         print(E_fape, E_pp, E_cc, E_nn)\n#         return E_fape + E_pp + E_cc + E_nn \n\n#     def obj_func_grad_np(self,rama_):\n#         rama=torch.DoubleTensor(rama_)\n#         rama.requires_grad=True\n#         if rama.grad:\n#             rama.grad.zero_()\n#         f=self.energy(rama.view(self.L,21))*Scale_factor\n#         grad_value=autograd.grad(f,rama)[0]\n#         return grad_value.data.numpy().astype(np.float64)\n    \n#     def obj_func_np(self,rama_):\n#         rama=torch.DoubleTensor(rama_)\n#         rama=rama.view(self.L,21)\n#         with torch.no_grad():\n#             f = self.energy(rama)*Scale_factor\n#             return f.item()\n\n#     def saveconfig(self,dict,confile):\n#         json_object = json.dumps(dict, indent = 4)\n#         wfile = open(confile,'w')\n#         wfile.write(json_object)\n#         wfile.close()\n    \n#     def scoring(self):\n#         geoscale = self.config['geo_scale']\n#         self.config['weight_pp'] = geoscale * self.config['weight_pp']\n#         self.config['weight_cc'] = geoscale * self.config['weight_cc']\n#         self.config['weight_nn'] = geoscale * self.config['weight_nn']\n#         self.config['weight_pccp'] = geoscale * self.config['weight_pccp']\n#         self.config['weight_cnnc'] = geoscale * self.config['weight_cnnc']\n#         self.config['weight_pnnp'] = geoscale * self.config['weight_pnnp']  \n        \n#         energy_dict = {}\n#         saveenergy_dict  = {}\n        \n#         with torch.no_grad():\n#             for retfile, tx in zip(self.geofiles, self.txs):\n#                 one = self.energy_from_coor(tx)\n#                 aaretfile = os.path.basename(retfile) \n#                 energy_dict[aaretfile] = one.item()\n#                 saveenergy_dict[retfile] = one.item()\n#             self.saveconfig(energy_dict, self.saveprefix)\n\n\n#     def foldning(self):\n#         minenergy=1e16\n#         count=0\n#         for tx in self.txs:\n#             count+=1\n        \n#         minirama=None\n\n#         ilter = self.init_ret\n#         selected_ret = self.geofiles[ilter]\n#         try:\n#             rama=self.init_quat(ilter).data.numpy()\n#             self.config=readconfig(os.path.join(os.path.dirname(os.path.abspath(__file__)),'lib','vdw.json'))\n#             rama = fmin_l_bfgs_b(func=self.obj_func_np, x0=rama,  fprime=self.obj_func_grad_np,iprint=10,maxfun=100)[0]\n#             rama = rama.flatten()\n#         except:\n#             rama=self.init_quat_safe(ilter).data.numpy()\n#             self.config=readconfig(os.path.join(os.path.dirname(os.path.abspath(__file__)),'lib','vdw.json'))\n#             rama = fmin_l_bfgs_b(func=self.obj_func_np, x0=rama,  fprime=self.obj_func_grad_np,iprint=10,maxfun=100)[0]\n#             rama = rama.flatten()\n            \n#         self.config=readconfig(self.foldconfig)\n#         geoscale = self.config['geo_scale']\n#         self.config['weight_pp'] =geoscale * self.config['weight_pp']\n#         self.config['weight_cc'] =geoscale * self.config['weight_cc']\n#         self.config['weight_nn'] =geoscale * self.config['weight_nn']\n#         self.config['weight_pccp'] =geoscale * self.config['weight_pccp']\n#         self.config['weight_cnnc'] =geoscale * self.config['weight_cnnc']\n#         self.config['weight_pnnp'] =geoscale * self.config['weight_pnnp']\n#         for i in range(3):\n#             line_min = lbfgs_rosetta.ArmijoLineMinimization(self.obj_func_np,self.obj_func_grad_np,True,len(rama),120)\n#             lbfgs_opt = lbfgs_rosetta.lbfgs(self.obj_func_np,self.obj_func_grad_np)\n#             rama=lbfgs_opt.run(rama,256,lbfgs_rosetta.absolute_converge_test,line_min,8000,self.obj_func_np,self.obj_func_grad_np,1e-9)\n#         newrama=rama+0.0\n#         newrama=torch.DoubleTensor(newrama) \n#         current_energy =self.obj_func_np(rama)\n\n#         if current_energy < minenergy:\n#             print(current_energy,minenergy)\n#             minenergy=current_energy\n#             self.outpdb(newrama,self.saveprefix+'.pdb',energystr=str(current_energy))\n\n\n#     def outpdb(self,rama,savefile,start=0,end=10000,energystr=''):\n#         coor_np=a2b.quat2b(self.basex,rama.view(self.L,21)[:,9:]).data.numpy()\n#         other_np=a2b.quat2b(self.otherx,rama.view(self.L,21)[:,9:]).data.numpy()\n#         shaped_rama=rama.view(self.L,21)\n#         coor = torch.FloatTensor(coor_np)\n#         side_coor_NP = a2b.quat2b(self.sidex,torch.cat([shaped_rama[:,:9],coor[:,-1]],dim=-1)).data.numpy()\n        \n#         Atom_name=[' P  ',\" C4'\",' N1 ']\n#         Other_Atom_name = [\" O5'\",\" C5'\",\" C3'\",\" O3'\",\" C1'\"]\n#         other_last_name = ['O',\"C\",\"C\",\"O\",\"C\"]\n\n#         side_atoms=         [' N1 ',' C2 ',' O2 ',' N2 ',' N3 ',' N4 ',' C4 ',' O4 ',' C5 ',' C6 ',' O6 ',' N6 ',' N7 ',' N8 ',' N9 ']\n#         side_last_name =    ['N',      \"C\",   \"O\",   \"N\",   \"N\",   'N',   'C',   'O',   'C',   'C',   'O',   'N',    'N', 'N','N']\n\n#         base_dict = rigid.base_table()\n        \n#         last_name=['P','C','N']\n#         wstr=[f'REMARK {str(energystr)}']\n#         templet='%6s%5d %4s %3s %1s%4d    %8.3f%8.3f%8.3f%6.2f%6.2f          %2s%2s'\n#         count=1\n#         for i in range(self.L):\n#             if self.seq[i] in ['a','g','A','G']:\n#                 Atom_name = [' P  ',\" C4'\",' N9 ']\n\n#             elif self.seq[i] in ['c','u','C','U']:\n#                 Atom_name = [' P  ',\" C4'\",' N1 ']\n#             for j in range(coor_np.shape[1]):\n#                 outs=('ATOM  ',count,Atom_name[j],self.seq[i],'A',i+1,coor_np[i][j][0],coor_np[i][j][1],coor_np[i][j][2],0,0,last_name[j],'')\n#                 if i>=start-1 and i < end:\n#                     wstr.append(templet % outs)\n#                     count+=1\n\n#             for j in range(other_np.shape[1]):\n#                 outs=('ATOM  ',count,Other_Atom_name[j],self.seq[i],'A',i+1,other_np[i][j][0],other_np[i][j][1],other_np[i][j][2],0,0,other_last_name[j],'')\n#                 if i>=start-1 and i < end:\n#                     wstr.append(templet % outs)\n#                     count+=1\n            \n#         wstr='\\n'.join(wstr)\n#         wfile=open(savefile,'w')\n#         wfile.write(wstr)\n#         wfile.close()\n    \n    \n#     def outpdb_coor(self,coor_np,savefile,start=0,end=1000,energystr=''):\n#         Atom_name=[' P  ',\" C4'\",' N1 ']\n#         last_name=['P','C','N']\n#         wstr=[f'REMARK {str(energystr)}']\n#         templet='%6s%5d %4s %3s %1s%4d    %8.3f%8.3f%8.3f%6.2f%6.2f          %2s%2s'\n#         count=1\n#         for i in range(self.L):\n#             if self.seq[i] in ['a','g','A','G']:\n#                 Atom_name = [' P  ',\" C4'\",' N9 ']\n\n#             elif self.seq[i] in ['c','u','C','U']:\n#                 Atom_name = [' P  ',\" C4'\",' N1 ']\n            \n#             for j in range(coor_np.shape[1]):\n#                 outs=('ATOM  ',count,Atom_name[j],self.seq[i],'A',i+1,coor_np[i][j][0],coor_np[i][j][1],coor_np[i][j][2],0,0,last_name[j],'')\n#                 if i>=start-1 and i < end:\n#                     wstr.append(templet % outs)\n#                 count+=1\n            \n#         wstr='\\n'.join(wstr)\n#         wfile=open(savefile,'w')\n#         wfile.write(wstr)\n#         wfile.close()\n\n\n# if __name__ == '__main__': \n\n#     fastafile = sys.argv[1]\n#     foldconfig = sys.argv[2]\n#     save_prefix = sys.argv[3]\n#     retfiles = sys.argv[4:]\n\n#     save_parent_dir = os.path.dirname(save_prefix)\n#     if not os.path.isdir(save_parent_dir):\n#         os.makedirs(save_parent_dir)\n\n#     retfiles.sort()\n#     print(retfiles)\n\n#     stru = Structure(fastafile, retfiles, foldconfig, save_prefix)    \n#     stru.scoring()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T13:02:03.176138Z","iopub.execute_input":"2026-02-19T13:02:03.176489Z","iopub.status.idle":"2026-02-19T13:02:03.195056Z","shell.execute_reply.started":"2026-02-19T13:02:03.176468Z","shell.execute_reply":"2026-02-19T13:02:03.194486Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %%writefile /kaggle/working/DRfold2/PotentialFold/Cubic.py\n# import numpy as np \n# from scipy.interpolate import CubicSpline,UnivariateSpline\n# import os\n# from torch.autograd import Function\n# import torch\n# import math\n\n# def fit_dis_cubic(dis_matrix,min_dis,max_dis,num_bin):\n#     # convert torch Tensor on GPU to numpy array for SciPy\n#     if isinstance(dis_matrix, torch.Tensor):\n#         dis_matrix = dis_matrix.detach().cpu().numpy()\n#     dis_region=np.zeros(num_bin)\n#     for i in range(num_bin):\n#         dis_region[i]=min_dis+(i+0.5)*(max_dis-min_dis)*1.0/num_bin\n#     L=dis_matrix.shape[0]\n#     csnp=[]\n#     decsnp=[]\n#     for i in range(L):\n#         css=[]\n#         decss=[]\n#         for j in range(L):\n#             y=-np.log(      (dis_matrix[i,j,1:-1]+1e-8) / (dis_matrix[i,j,[-2]]+1e-8)              )\n#             x=dis_region\n#             x[0]=-0.0001\n#             y[0]= max(10,y[1]+4)\n#             cs= CubicSpline(x,y)\n#             decs=cs.derivative()\n#             css.append(cs)\n#             decss.append(decs)\n#         csnp.append(css)\n#         decsnp.append(decss)\n#     return np.array(csnp),np.array(decsnp)\n\n# def dis_cubic(out,min_dis,max_dis,num_bin):\n#     print('fitting cubic distance')\n#     cs,decs=fit_dis_cubic(out,min_dis,max_dis,num_bin)\n#     return cs,decs\n\n\n\n# def cubic_matrix_torsion(dis_matrix,min_dis,max_dis,num_bin):\n#     dis_region=np.zeros(num_bin)\n#     bin_size=(max_dis-min_dis)/num_bin\n#     for i in range(num_bin):\n#         dis_region[i]=min_dis+(i+0.5)*(max_dis-min_dis)*1.0/num_bin\n#     L=dis_matrix.shape[0]\n#     csnp=[]\n#     decsnp=[]\n#     for i in range(L):\n#         css=[]\n#         decss=[]\n#         for j in range(L):\n#             y=-np.log(      dis_matrix[i,j,:-1]+1e-8             )\n#             x=dis_region\n#             x=np.append(x,x[-1]+bin_size)\n#             y=np.append(y,y[0])\n#             cs= CubicSpline(x,y,bc_type='periodic')\n#             decs=cs.derivative()\n#             css.append(cs)\n#             decss.append(decs)\n#         csnp.append(css)\n#         decsnp.append(decss)\n#     return np.array(csnp),np.array(decsnp)\n# def torsion_cubic(out,min_dis,max_dis,num_bin):\n#     print('fitting cubic')\n#     cs,decs=cubic_matrix_torsion(out,min_dis,max_dis,num_bin)\n#     return cs,decs\n\n# def cubic_matrix_angle(dis_matrix,min_dis,max_dis,num_bin): # 0 - np.pi 12\n#     dis_region=np.zeros(num_bin)\n#     bin_size=(max_dis-min_dis)/num_bin\n#     for i in range(num_bin):\n#         dis_region[i]=min_dis+(i+0.5)*(max_dis-min_dis)*1.0/num_bin\n#     L=dis_matrix.shape[0]\n#     csnp=[]\n#     decsnp=[]\n#     for i in range(L):\n#         css=[]\n#         decss=[]\n#         for j in range(L):\n#             y=-np.log(      dis_matrix[i,j,:-1]+1e-8             )\n#             x=dis_region\n\n#             x=np.concatenate([[x[0]-bin_size*3,x[0]-bin_size*2,x[0]-bin_size], x,[x[-1]+bin_size,x[-1]+bin_size*2,x[-1]+bin_size*3]               ])\n#             y=np.concatenate([ [y[2],y[1],y[0]],y,[y[-1],y[-2],y[-3]]                                                                                                                    ])\n\n#             cs= CubicSpline(x,y)\n#             decs=cs.derivative()\n\n#             css.append(cs)\n#             decss.append(decs)\n#         csnp.append(css)\n#         decsnp.append(decss)\n\n#     return np.array(csnp),np.array(decsnp)\n# def angle_cubic(out,min_dis,max_dis,num_bin):\n\n#     print('fitting angle cubic')\n#     cs,decs=cubic_matrix_angle(out,min_dis,max_dis,num_bin)\n\n#     return cs,decs","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T13:02:03.196136Z","iopub.execute_input":"2026-02-19T13:02:03.196404Z","iopub.status.idle":"2026-02-19T13:02:03.208474Z","shell.execute_reply.started":"2026-02-19T13:02:03.196379Z","shell.execute_reply":"2026-02-19T13:02:03.207905Z"},"jupyter":{"source_hidden":true},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # Define an improved function to run DRfold2 that captures output\n# def predict_rna_structures_drfold2(sequence, target_id):\n#     \"\"\"Use DRfold2 to predict RNA structures with proper output capture\"\"\"\n#     import subprocess\n#     from subprocess import PIPE, STDOUT\n    \n#     # Create FASTA file for this sequence\n#     fasta_path = os.path.join(fasta_dir, f\"{target_id}.fasta\")\n#     with open(fasta_path, \"w\") as f:\n#         f.write(f\">{target_id}\\n{sequence}\\n\")\n    \n#     # Run DRfold2 with proper output capture\n#     output_dir = os.path.join(predictions_dir, target_id)\n#     cmd = f\"python /kaggle/working/DRfold2/DRfold_infer.py {fasta_path} {output_dir} 1\"\n    \n#     print(f\"Running command: {cmd}\")\n#     process = subprocess.Popen(\n#         cmd, \n#         shell=True, \n#         stdout=PIPE, \n#         stderr=STDOUT,\n#         universal_newlines=True,\n#         bufsize=1\n#     )\n    \n#     # Print output in real-time\n#     for line in iter(process.stdout.readline, ''):\n#         line = line.strip()\n#         if line:\n#             print(line)\n    \n#     # Get return code and check success\n#     return_code = process.wait()\n#     if return_code != 0:\n#         print(f\"DRfold2 failed with return code {return_code}\")\n#         return None\n    \n#     # Clean up FASTA file to save space\n#     os.remove(fasta_path)\n    \n#     # Extract coordinates\n#     relax_dir = os.path.join(output_dir, \"relax\")\n#     if not os.path.isdir(relax_dir):\n#         print(f\"Warning: No relax directory found for {target_id}\")\n#         relax_dir = output_dir\n    \n#     # Get up to 5 PDB files\n#     pdb_files = sorted([f for f in os.listdir(relax_dir) if f.endswith(\".pdb\")])[:5]\n    \n#     if not pdb_files:\n#         print(f\"Warning: No PDB files found for {target_id}\")\n#         # Return None to indicate failure\n#         return None\n    \n#     # Parse PDB files to extract C1' coordinates\n#     predictions = []\n#     for pdb_file in pdb_files:\n#         file_path = os.path.join(relax_dir, pdb_file)\n        \n#         # Read PDB file\n#         coords = []\n#         with open(file_path, \"r\") as f:\n#             residue_map = {}\n#             for line in f:\n#                 if line.startswith(\"ATOM\") and \" C1' \" in line:\n#                     parts = line.split()\n#                     resid = int(parts[5])  # Residue ID as integer\n#                     x, y, z = float(parts[6]), float(parts[7]), float(parts[8])\n#                     residue_map[resid] = (x, y, z)\n            \n#             # Ensure we have coordinates for all residues\n#             for j in range(1, len(sequence) + 1):\n#                 if j in residue_map:\n#                     coords.append(residue_map[j])\n#                 else:\n#                     # If residue not found, use zeros\n#                     print(f\"Warning: Residue {j} not found in {pdb_file} for {target_id}\")\n#                     coords.append((0.0, 0.0, 0.0))\n        \n#         predictions.append(coords)\n    \n#     # Clean up PDB files to save space\n#     if is_submission_mode:\n#         shutil.rmtree(output_dir)\n    \n#     # If we have fewer than 5 predictions, duplicate the last one\n#     while len(predictions) < 5:\n#         predictions.append(predictions[-1] if predictions else [(0.0, 0.0, 0.0) for _ in range(len(sequence))])\n    \n#     return predictions[:5]  # Return exactly 5 predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T13:02:03.209314Z","iopub.execute_input":"2026-02-19T13:02:03.209672Z","iopub.status.idle":"2026-02-19T13:02:03.222822Z","shell.execute_reply.started":"2026-02-19T13:02:03.209646Z","shell.execute_reply":"2026-02-19T13:02:03.222099Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # Vectorized version of process_labels function\n# def process_labels_vectorized(labels_df):\n#     # Extract target_id from ID column (remove last part after underscore)\n#     labels_df = labels_df.copy()\n#     labels_df['target_id'] = labels_df['ID'].str.rsplit('_', n=1).str[0]\n    \n#     # Sort by target_id and resid for proper ordering\n#     labels_df = labels_df.sort_values(['target_id', 'resid'])\n    \n#     # Group by target_id and convert coordinates to arrays\n#     coords_dict = {}\n#     for target_id, group in labels_df.groupby('target_id'):\n#         # Extract coordinates as numpy array in one operation\n#         coords_dict[target_id] = group[['x_1', 'y_1', 'z_1']].values\n    \n#     return coords_dict\n\n# def find_similar_sequences(query_seq, train_seqs_df, train_coords_dict, top_n=5):\n#     similar_seqs = []\n#     query_seq_obj = Seq(query_seq)\n\n#     for _, row in train_seqs_df.iterrows():\n#         target_id = row['target_id']\n#         train_seq = row['sequence']\n\n#         # Skip if coordinates not available\n#         if target_id not in train_coords_dict:\n#             continue\n\n#         # Skip if sequence is too different in length (more than 40% difference)\n#         if abs(len(train_seq) - len(query_seq)) / max(len(train_seq), len(query_seq)) > 0.4:\n#             continue\n\n#         # Perform sequence alignment\n#         alignments = pairwise2.align.globalms(query_seq_obj, train_seq, 2.9, -1, -10, -0.5, one_alignment_only=True)\n\n#         if alignments:\n#             alignment = alignments[0]\n#             similarity_score = alignment.score / (2 * min(len(query_seq), len(train_seq)))\n#             similar_seqs.append((target_id, train_seq, similarity_score, train_coords_dict[target_id]))\n\n#     # Sort by similarity score (higher is better) and return top N\n#     similar_seqs.sort(key=lambda x: x[2], reverse=True)\n#     return similar_seqs[:top_n]\n\n\n# def adaptive_rna_constraints(coordinates, sequence, confidence=1.0):\n#     # Make a copy of coordinates to refine\n#     refined_coords = coordinates.copy()\n#     n_residues = len(sequence)\n    \n#     # Calculate constraint strength (inverse of confidence)\n#     # High confidence templates receive gentler constraints\n#     constraint_strength = 0.8 * (1.0 - min(confidence, 0.8))\n    \n#     # 1. Sequential distance constraints (consecutive nucleotides)\n#     # More flexible distance range (statistical distribution from PDB)\n#     seq_min_dist = 5.5  # Minimum sequential distance\n#     seq_max_dist = 6.5  # Maximum sequential distance\n    \n#     for i in range(n_residues - 1):\n#         current_pos = refined_coords[i]\n#         next_pos = refined_coords[i+1]\n        \n#         # Calculate current distance\n#         current_dist = np.linalg.norm(next_pos - current_pos)\n        \n#         # Only adjust if significantly outside expected range\n#         if current_dist < seq_min_dist or current_dist > seq_max_dist:\n#             # Calculate target distance (midpoint of range)\n#             target_dist = (seq_min_dist + seq_max_dist) / 2\n            \n#             # Get direction vector\n#             direction = next_pos - current_pos\n#             direction = direction / (np.linalg.norm(direction) + 1e-10)\n            \n#             # Apply partial adjustment based on constraint strength\n#             adjustment = (target_dist - current_dist) * constraint_strength\n            \n#             # Only adjust the next position to preserve the overall fold\n#             refined_coords[i+1] = current_pos + direction * (current_dist + adjustment)\n    \n#     # 2. Steric clash prevention (more conservative)\n#     min_allowed_distance = 3.8  # Minimum distance between non-consecutive C1' atoms\n    \n#     # Calculate all pairwise distances\n#     dist_matrix = distance_matrix(refined_coords, refined_coords)\n    \n#     # Find severe clashes (atoms too close)\n#     severe_clashes = np.where((dist_matrix < min_allowed_distance) & (dist_matrix > 0))\n    \n#     # Fix severe clashes\n#     for idx in range(len(severe_clashes[0])):\n#         i, j = severe_clashes[0][idx], severe_clashes[1][idx]\n        \n#         # Skip consecutive nucleotides and previously processed pairs\n#         if abs(i - j) <= 1 or i >= j:\n#             continue\n            \n#         # Get current positions and distance\n#         pos_i = refined_coords[i]\n#         pos_j = refined_coords[j]\n#         current_dist = dist_matrix[i, j]\n        \n#         # Calculate necessary adjustment but scale by constraint strength\n#         direction = pos_j - pos_i\n#         direction = direction / (np.linalg.norm(direction) + 1e-10)\n        \n#         # Calculate partial adjustment\n#         adjustment = (min_allowed_distance - current_dist) * constraint_strength\n        \n#         # Move points apart\n#         refined_coords[i] = pos_i - direction * (adjustment / 2)\n#         refined_coords[j] = pos_j + direction * (adjustment / 2)\n    \n#     # 3. Very light base-pair constraining (if confidence is low)\n#     if constraint_strength > 0.3:  # Only apply if template confidence is low\n#         # Simple Watson-Crick base pairs\n#         pairs = {'A': 'U', 'U': 'A', 'G': 'C', 'C': 'G'}\n        \n#         # Scan for potential base pairs\n#         for i in range(n_residues):\n#             base_i = sequence[i]\n#             complement = pairs.get(base_i)\n            \n#             if not complement:\n#                 continue\n                \n#             # Look for complementary bases within a reasonable range\n#             for j in range(i + 3, min(i + 20, n_residues)):\n#                 if sequence[j] == complement:\n#                     # Calculate current distance\n#                     current_dist = np.linalg.norm(refined_coords[i] - refined_coords[j])\n                    \n#                     # Only consider if distance suggests potential pairing\n#                     if 8.0 < current_dist < 14.0:\n#                         # Target 10.5Å as generic base-pair C1'-C1' distance\n#                         target_dist = 10.5\n                        \n#                         # Calculate very gentle adjustment (scaled by constraint_strength)\n#                         adjustment = (target_dist - current_dist) * (constraint_strength * 0.3)\n                        \n#                         # Get direction vector\n#                         direction = refined_coords[j] - refined_coords[i]\n#                         direction = direction / (np.linalg.norm(direction) + 1e-10)\n                        \n#                         # Apply very gentle adjustment to both positions\n#                         refined_coords[i] = refined_coords[i] - direction * (adjustment / 2)\n#                         refined_coords[j] = refined_coords[j] + direction * (adjustment / 2)\n                        \n#                         # Only consider one potential pair per base (closest match)\n#                         break\n    \n#     return refined_coords\n\n# def adapt_template_to_query(query_seq, template_seq, template_coords, alignment=None):\n#     if alignment is None:\n#         from Bio.Seq import Seq\n#         from Bio import pairwise2\n        \n#         query_seq_obj = Seq(query_seq)\n#         template_seq_obj = Seq(template_seq)\n#         alignments = pairwise2.align.globalms(query_seq_obj, template_seq_obj, 2.9, -1, -10, -0.5, one_alignment_only=True)\n        \n#         if not alignments:\n#             return generate_improved_rna_structure(query_seq)\n            \n#         alignment = alignments[0]\n    \n#     aligned_query = alignment.seqA\n#     aligned_template = alignment.seqB\n    \n#     query_coords = np.zeros((len(query_seq), 3))\n#     query_coords.fill(np.nan)\n    \n#     # Map template coordinates to query\n#     query_idx = 0\n#     template_idx = 0\n    \n#     for i in range(len(aligned_query)):\n#         query_char = aligned_query[i]\n#         template_char = aligned_template[i]\n        \n#         if query_char != '-' and template_char != '-':\n#             if template_idx < len(template_coords):\n#                 query_coords[query_idx] = template_coords[template_idx]\n#             template_idx += 1\n#             query_idx += 1\n#         elif query_char != '-' and template_char == '-':\n#             query_idx += 1\n#         elif query_char == '-' and template_char != '-':\n#             template_idx += 1\n    \n#     # IMPROVED GAP FILLING - maintains RNA backbone geometry\n#     backbone_distance = 5.9  # Typical C1'-C1' distance\n    \n#     # Fill gaps by maintaining realistic backbone connectivity\n#     for i in range(len(query_coords)):\n#         if np.isnan(query_coords[i, 0]):\n#             # Find nearest valid neighbors\n#             prev_valid = next_valid = None\n            \n#             for j in range(i-1, -1, -1):\n#                 if not np.isnan(query_coords[j, 0]):\n#                     prev_valid = j\n#                     break\n                    \n#             for j in range(i+1, len(query_coords)):\n#                 if not np.isnan(query_coords[j, 0]):\n#                     next_valid = j\n#                     break\n            \n#             if prev_valid is not None and next_valid is not None:\n#                 # Interpolate along realistic RNA backbone path\n#                 gap_size = next_valid - prev_valid\n#                 total_distance = np.linalg.norm(query_coords[next_valid] - query_coords[prev_valid])\n#                 expected_distance = gap_size * backbone_distance\n                \n#                 # If gap is compressed, extend it realistically\n#                 if total_distance < expected_distance * 0.7:\n#                     direction = query_coords[next_valid] - query_coords[prev_valid]\n#                     direction = direction / (np.linalg.norm(direction) + 1e-10)\n                    \n#                     # Place intermediate points along extended path\n#                     for k, idx in enumerate(range(prev_valid + 1, next_valid)):\n#                         progress = (k + 1) / gap_size\n#                         base_pos = query_coords[prev_valid] + direction * expected_distance * progress\n                        \n#                         # Add slight curvature for realism\n#                         perpendicular = np.cross(direction, [0, 0, 1])\n#                         if np.linalg.norm(perpendicular) < 1e-6:\n#                             perpendicular = np.cross(direction, [1, 0, 0])\n#                         perpendicular = perpendicular / (np.linalg.norm(perpendicular) + 1e-10)\n                        \n#                         curve_amplitude = 2.0 * np.sin(progress * np.pi)\n#                         query_coords[idx] = base_pos + perpendicular * curve_amplitude\n#                 else:\n#                     # Linear interpolation for normal gaps\n#                     for k, idx in enumerate(range(prev_valid + 1, next_valid)):\n#                         weight = (k + 1) / gap_size\n#                         query_coords[idx] = (1 - weight) * query_coords[prev_valid] + weight * query_coords[next_valid]\n            \n#             elif prev_valid is not None:\n#                 # Extend from previous position\n#                 if prev_valid > 0 and not np.isnan(query_coords[prev_valid-1, 0]):\n#                     direction = query_coords[prev_valid] - query_coords[prev_valid-1]\n#                     direction = direction / (np.linalg.norm(direction) + 1e-10)\n#                 else:\n#                     direction = np.array([1.0, 0.0, 0.0])\n                \n#                 steps_needed = i - prev_valid\n#                 for step in range(1, steps_needed + 1):\n#                     pos_idx = prev_valid + step\n#                     if pos_idx < len(query_coords):\n#                         query_coords[pos_idx] = query_coords[prev_valid] + direction * backbone_distance * step\n            \n#             elif next_valid is not None:\n#                 # Work backwards from next position\n#                 direction = np.array([-1.0, 0.0, 0.0])  # Default backward direction\n#                 steps_needed = next_valid - i\n#                 for step in range(steps_needed, 0, -1):\n#                     pos_idx = next_valid - step\n#                     if pos_idx >= 0:\n#                         query_coords[pos_idx] = query_coords[next_valid] - direction * backbone_distance * step\n    \n#     # Final cleanup\n#     query_coords = np.nan_to_num(query_coords)\n#     return query_coords\n\n\n\n\n# def generate_improved_rna_structure(sequence):\n#     \"\"\"\n#     Generate a more realistic RNA structure fallback based on sequence patterns\n#     and basic RNA structure principles.\n    \n#     Args:\n#         sequence: RNA sequence string\n        \n#     Returns:\n#         Array of 3D coordinates\n#     \"\"\"\n#     n_residues = len(sequence)\n#     coordinates = np.zeros((n_residues, 3))\n    \n#     # Analyze sequence to predict structural elements\n#     # Look for complementary regions that could form base pairs\n#     potential_stems = identify_potential_stems(sequence)\n    \n#     # Default parameters\n#     radius_helix = 10.0\n#     radius_loop = 15.0\n#     rise_per_residue_helix = 2.5\n#     rise_per_residue_loop = 1.5\n#     angle_per_residue_helix = 0.6\n#     angle_per_residue_loop = 0.3\n    \n#     # Assign structural classifications\n#     structure_types = assign_structure_types(sequence, potential_stems)\n    \n#     # Generate coordinates based on predicted structure\n#     current_pos = np.array([0.0, 0.0, 0.0])\n#     current_direction = np.array([0.0, 0.0, 1.0])\n#     current_angle = 0.0\n    \n#     for i in range(n_residues):\n#         if structure_types[i] == 'stem':\n#             # Part of a helical stem\n#             current_angle += angle_per_residue_helix\n#             coordinates[i] = [\n#                 radius_helix * np.cos(current_angle), \n#                 radius_helix * np.sin(current_angle), \n#                 current_pos[2] + rise_per_residue_helix\n#             ]\n#             current_pos = coordinates[i]\n#         elif structure_types[i] == 'loop':\n#             # Part of a loop\n#             current_angle += angle_per_residue_loop\n#             z_shift = rise_per_residue_loop * np.sin(current_angle * 0.5)\n#             coordinates[i] = [\n#                 radius_loop * np.cos(current_angle), \n#                 radius_loop * np.sin(current_angle), \n#                 current_pos[2] + z_shift\n#             ]\n#             current_pos = coordinates[i]\n#         else:\n#             # Single-stranded region\n#             # Add some randomness to make it look more realistic\n#             jitter = np.random.normal(0, 1, 3) * 2.0\n#             coordinates[i] = current_pos + jitter\n#             current_pos = coordinates[i]\n            \n#     return coordinates\n\n# def identify_potential_stems(sequence):\n#     \"\"\"\n#     Identify potential stem regions by looking for self-complementary segments.\n    \n#     Args:\n#         sequence: RNA sequence string\n        \n#     Returns:\n#         List of tuples (start1, end1, start2, end2) representing potentially paired regions\n#     \"\"\"\n#     complementary_bases = {'A': 'U', 'U': 'A', 'G': 'C', 'C': 'G'}\n#     min_stem_length = 3\n#     potential_stems = []\n    \n#     # Simple stem identification\n#     for i in range(len(sequence) - min_stem_length):\n#         for j in range(i + min_stem_length + 3, len(sequence) - min_stem_length + 1):\n#             # Check if regions could form a stem\n#             potential_stem_len = min(min_stem_length, len(sequence) - j)\n#             is_stem = True\n            \n#             for k in range(potential_stem_len):\n#                 if sequence[i+k] not in complementary_bases or \\\n#                    complementary_bases[sequence[i+k]] != sequence[j+potential_stem_len-k-1]:\n#                     is_stem = False\n#                     break\n            \n#             if is_stem:\n#                 potential_stems.append((i, i+potential_stem_len-1, j, j+potential_stem_len-1))\n    \n#     return potential_stems\n\n# def assign_structure_types(sequence, potential_stems):\n#     \"\"\"\n#     Assign each nucleotide to a structural element type.\n    \n#     Args:\n#         sequence: RNA sequence string\n#         potential_stems: List of tuples representing stem regions\n        \n#     Returns:\n#         List of structure types ('stem', 'loop', 'single')\n#     \"\"\"\n#     structure_types = ['single'] * len(sequence)\n    \n#     # Mark stem regions\n#     for stem in potential_stems:\n#         start1, end1, start2, end2 = stem\n#         for i in range(end1 - start1 + 1):\n#             structure_types[start1 + i] = 'stem'\n#             structure_types[end2 - i] = 'stem'\n    \n#     # Mark loop regions (regions between paired regions)\n#     for i in range(len(potential_stems) - 1):\n#         _, end1, start2, _ = potential_stems[i]\n#         next_start1, _, _, _ = potential_stems[i+1]\n        \n#         if next_start1 > end1 + 1 and start2 > next_start1:\n#             for j in range(end1 + 1, next_start1):\n#                 structure_types[j] = 'loop'\n    \n#     return structure_types\n\n\n\n# # Function to create a more realistic RNA structure when no good templates are found\n# def generate_rna_structure(sequence, seed=None):\n#     if seed is not None:\n#         np.random.seed(seed)\n#         random.seed(seed)\n    \n#     n_residues = len(sequence)\n#     coordinates = np.zeros((n_residues, 3))\n    \n#     # Initialize the first few residues in a helix\n#     for i in range(min(3, n_residues)):\n#         angle = i * 0.6\n#         coordinates[i] = [10.0 * np.cos(angle), 10.0 * np.sin(angle), i * 2.5]\n    \n#     # Add more complex folding patterns\n#     current_direction = np.array([0.0, 0.0, 1.0])  # Start moving along z-axis\n    \n#     # Define base-pairing tendencies (G-C and A-U pairs)\n#     for i in range(3, n_residues):\n#         # Check for potential base-pairing in the sequence\n#         has_pair = False\n#         pair_idx = -1\n        \n#         # Simple detection of complementary bases (G-C, A-U)\n#         complementary = {'G': 'C', 'C': 'G', 'A': 'U', 'U': 'A'}\n#         current_base = sequence[i]\n        \n#         # Look for potential base-pairing within a window before the current position\n#         window_size = min(i, 15)  # Look back up to 15 bases\n#         for j in range(i-window_size, i):\n#             if j >= 0 and sequence[j] == complementary.get(current_base, 'X'):\n#                 # Found a potential pair\n#                 has_pair = True\n#                 pair_idx = j\n#                 break\n        \n#         if has_pair and i - pair_idx <= 10 and random.random() < 0.7:\n#             # Try to create a base-pair by positioning this nucleotide near its pair\n#             pair_pos = coordinates[pair_idx]\n            \n#             # Create a position that's roughly opposite to the pair\n#             random_offset = np.random.normal(0, 1, 3) * 2.0\n#             base_pair_distance = 10.0 + random.uniform(-1.0, 1.0)\n            \n#             # Calculate a vector from base-pair toward center of structure\n#             center = np.mean(coordinates[:i], axis=0)\n#             direction = center - pair_pos\n#             direction = direction / (np.linalg.norm(direction) + 1e-10)\n            \n#             # Position new nucleotide in the general direction of the \"center\"\n#             coordinates[i] = pair_pos + direction * base_pair_distance + random_offset\n            \n#             # Update direction for next nucleotide\n#             current_direction = np.random.normal(0, 0.3, 3)\n#             current_direction = current_direction / (np.linalg.norm(current_direction) + 1e-10)\n            \n#         else:\n#             # No base-pairing detected, continue with the current fold direction\n#             # Randomly rotate current direction to simulate RNA flexibility\n#             if random.random() < 0.3:\n#                 # More significant direction change\n#                 angle = random.uniform(0.2, 0.6)\n#                 axis = np.random.normal(0, 1, 3)\n#                 axis = axis / (np.linalg.norm(axis) + 1e-10)\n#                 rotation = R.from_rotvec(angle * axis)\n#                 current_direction = rotation.apply(current_direction)\n#             else:\n#                 # Small random changes in direction\n#                 current_direction += np.random.normal(0, 0.15, 3)\n#                 current_direction = current_direction / (np.linalg.norm(current_direction) + 1e-10)\n            \n#             # Distance between consecutive nucleotides (3.5-4.5Å is typical)\n#             step_size = random.uniform(3.5, 4.5)\n            \n#             # Update position\n#             coordinates[i] = coordinates[i-1] + step_size * current_direction\n    \n#     return coordinates\n\n\n# def predict_rna_structures(sequence, target_id, train_seqs_df, train_coords_dict, n_predictions=5):\n#     predictions = []\n    \n#     # Find similar sequences in the training data\n#     similar_seqs = find_similar_sequences(sequence, train_seqs_df, train_coords_dict, top_n=n_predictions)\n    \n#     # If we found any similar sequences, use them as templates\n#     if similar_seqs:\n#         for i, (template_id, template_seq, similarity_score, template_coords) in enumerate(similar_seqs):\n#             # Adapt template coordinates to the query sequence\n#             adapted_coords = adapt_template_to_query(sequence, template_seq, template_coords)\n            \n#             if adapted_coords is not None:\n#                 # Apply adaptive constraints based on template similarity\n#                 # For high similarity templates, apply very gentle constraints\n#                 refined_coords = adaptive_rna_constraints(adapted_coords, sequence, confidence=similarity_score)\n                \n#                 # Add some randomness (less for better templates)\n#                 random_scale = max(0.05, 0.8 - similarity_score)  # Reduced randomness\n#                 randomized_coords = refined_coords.copy()\n#                 randomized_coords += np.random.normal(0, random_scale, randomized_coords.shape)\n                \n#                 predictions.append(randomized_coords)\n                \n#                 if len(predictions) >= n_predictions:\n#                     break\n    \n#     # If we don't have enough predictions from templates, generate de novo structures\n#     while len(predictions) < n_predictions:\n#         seed_value = hash(target_id) % 10000 + len(predictions) * 1000\n#         de_novo_coords = generate_rna_structure(sequence, seed=seed_value)\n        \n#         # Apply stronger constraints to de novo structures (lower confidence)\n#         refined_de_novo = adaptive_rna_constraints(de_novo_coords, sequence, confidence=0.2)\n        \n#         predictions.append(refined_de_novo)\n    \n#     return predictions[:n_predictions]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T13:02:03.223692Z","iopub.execute_input":"2026-02-19T13:02:03.223942Z","iopub.status.idle":"2026-02-19T13:02:03.243062Z","shell.execute_reply.started":"2026-02-19T13:02:03.223921Z","shell.execute_reply":"2026-02-19T13:02:03.242381Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # Initialize counters and range settings\n# if is_submission_mode:\n#     DRFOLD_START_IDX = 14\n#     DRFOLD_END_IDX = len(test_sequences) - 1\n# else:\n#     DRFOLD_START_IDX = 0\n#     DRFOLD_END_IDX = 0\n\n# drfold_processed = 0\n# template_processed = 0\n\n# # train_coords_dict = process_labels_vectorized(train_labels_final)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T13:02:03.243957Z","iopub.execute_input":"2026-02-19T13:02:03.244309Z","iopub.status.idle":"2026-02-19T13:02:03.256748Z","shell.execute_reply.started":"2026-02-19T13:02:03.244281Z","shell.execute_reply":"2026-02-19T13:02:03.256033Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # Sort test sequences by length to process shorter ones with DRfold2\n# test_sequences[\"sequence_len\"] = test_sequences[\"sequence\"].str.len()\n# # test_sequences = test_sequences[(test_sequences['sequence_len'] >= 200) & (test_sequences['sequence_len'] < 600)]\n# test_sequences = test_sequences[test_sequences['sequence_len'] < 100]\n# test_sequences.reset_index(drop=True, inplace=True)\n# print(test_sequences.shape)\n\n# if not IS_SCORING_RUN:\n#     test_sequences = test_sequences.head(5)\n\n# # List to store all prediction records\n# all_predictions = []\n\n# # For each sequence in the test set\n# for idx, row in test_sequences.iterrows():\n#     target_id = row['target_id']\n#     sequence = row['sequence']\n    \n#     # Generate 5 different structure predictions\n#     print(f\"Using DRfold2 for target {target_id} (index {idx})\")\n#     predictions = predict_rna_structures_drfold2(sequence, target_id)\n    \n#     # For each residue in the sequence\n#     for j in range(len(sequence)):\n#         pred_row = {\n#             'ID': f\"{target_id}_{j+1}\",\n#             'resname': sequence[j],\n#             'resid': j + 1\n#         }\n        \n#         # Add coordinates from all 5 predictions\n#         for i in range(5):\n#             pred_row[f'x_{i+1}'] = predictions[i][j][0]\n#             pred_row[f'y_{i+1}'] = predictions[i][j][1]\n#             pred_row[f'z_{i+1}'] = predictions[i][j][2]\n        \n#         all_predictions.append(pred_row)\n    \n#     # Free up memory\n#     if torch.cuda.is_available():\n#         torch.cuda.empty_cache()\n\n# # Create DataFrame with predictions\n# drfold_pred = pd.DataFrame(all_predictions)\n\n# # Ensure the submission file has the correct format\n# column_order = ['ID', 'resname', 'resid']\n# for i in range(1, 6):\n#     for coord in ['x', 'y', 'z']:\n#         column_order.append(f'{coord}_{i}')\n        \n# drfold_pred = drfold_pred[column_order]\n\n# # Clean the working directory before saving\n# print(\"Cleaning working directory...\")\n# for item in os.listdir(\"/kaggle/working/\"):\n#     item_path = os.path.join(\"/kaggle/working/\", item)\n#     if os.path.isfile(item_path) and item != \"submission.csv\":\n#         os.remove(item_path)\n#     elif os.path.isdir(item_path) and item != \"predictions\" and item != \"fasta_files\" and item != \"DRfold2\":\n#         shutil.rmtree(item_path)\n\n# # Save the submission\n# drfold_pred.to_csv('/kaggle/working/drfold_submission.csv', index=False)\n# print(f\"Submission file saved to /kaggle/working/drfold_submission.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T13:02:03.257535Z","iopub.execute_input":"2026-02-19T13:02:03.257778Z","iopub.status.idle":"2026-02-19T13:02:03.267851Z","shell.execute_reply.started":"2026-02-19T13:02:03.257753Z","shell.execute_reply":"2026-02-19T13:02:03.267204Z"},"_kg_hide-output":true,"jupyter":{"source_hidden":true},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# drfold_pred","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T13:02:03.268698Z","iopub.execute_input":"2026-02-19T13:02:03.268967Z","iopub.status.idle":"2026-02-19T13:02:03.27973Z","shell.execute_reply.started":"2026-02-19T13:02:03.268937Z","shell.execute_reply":"2026-02-19T13:02:03.279022Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# TBM","metadata":{}},{"cell_type":"code","source":"!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","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2026-03-22T00:27:20.752781Z","iopub.execute_input":"2026-03-22T00:27:20.753253Z","iopub.status.idle":"2026-03-22T00:27:24.050352Z","shell.execute_reply.started":"2026-03-22T00:27:20.753224Z","shell.execute_reply":"2026-03-22T00:27:24.049535Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import 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\n# sub[cols].to_csv('submission.csv', index=False)\n# print(\"submission.csv! saved\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-22T00:27:24.052304Z","iopub.execute_input":"2026-03-22T00:27:24.052632Z","iopub.status.idle":"2026-03-22T00:29:54.374488Z","shell.execute_reply.started":"2026-03-22T00:27:24.052603Z","shell.execute_reply":"2026-03-22T00:29:54.373797Z"},"_kg_hide-input":false},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pred_tbm = sub","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-22T00:29:54.376088Z","iopub.execute_input":"2026-03-22T00:29:54.376365Z","iopub.status.idle":"2026-03-22T00:29:54.380623Z","shell.execute_reply.started":"2026-03-22T00:29:54.376342Z","shell.execute_reply":"2026-03-22T00:29:54.379893Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# print('Submission.csv shape:', rosetta_pred.shape)\npred_tbm.to_csv('/kaggle/working/pred_tbm.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-22T00:29:54.381441Z","iopub.execute_input":"2026-03-22T00:29:54.38177Z","iopub.status.idle":"2026-03-22T00:29:54.608427Z","shell.execute_reply.started":"2026-03-22T00:29:54.381721Z","shell.execute_reply":"2026-03-22T00:29:54.607474Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pred_tbm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-22T00:29:54.61013Z","iopub.execute_input":"2026-03-22T00:29:54.610477Z","iopub.status.idle":"2026-03-22T00:29:54.630985Z","shell.execute_reply.started":"2026-03-22T00:29:54.610456Z","shell.execute_reply":"2026-03-22T00:29:54.630184Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# RNAPro","metadata":{}},{"cell_type":"code","source":"!cp -r /kaggle/input/datasets/theoviel/rnapro-src/RNAPro .\n!cp /kaggle/input/datasets/theoviel/rnapro-src/rnapro-private-best-500m.ckpt .","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T07:41:24.59322Z","iopub.execute_input":"2026-03-01T07:41:24.593525Z","iopub.status.idle":"2026-03-01T07:42:02.516937Z","shell.execute_reply.started":"2026-03-01T07:41:24.593503Z","shell.execute_reply":"2026-03-01T07:42:02.515977Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cd RNAPro","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T07:42:02.518336Z","iopub.execute_input":"2026-03-01T07:42:02.518579Z","iopub.status.idle":"2026-03-01T07:42:02.524059Z","shell.execute_reply.started":"2026-03-01T07:42:02.518553Z","shell.execute_reply":"2026-03-01T07:42:02.523416Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -e . --no-deps","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T07:42:02.525055Z","iopub.execute_input":"2026-03-01T07:42:02.525305Z","iopub.status.idle":"2026-03-01T07:42:11.04391Z","shell.execute_reply.started":"2026-03-01T07:42:02.52528Z","shell.execute_reply":"2026-03-01T07:42:11.043231Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DIST = \"/kaggle/working/RNAPro/release_data/ccd_cache/\"\n!mkdir -p $DIST","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T07:42:11.045025Z","iopub.execute_input":"2026-03-01T07:42:11.045292Z","iopub.status.idle":"2026-03-01T07:42:11.210529Z","shell.execute_reply.started":"2026-03-01T07:42:11.045266Z","shell.execute_reply":"2026-03-01T07:42:11.209823Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T07:42:11.211702Z","iopub.execute_input":"2026-03-01T07:42:11.212337Z","iopub.status.idle":"2026-03-01T07:42:15.484946Z","shell.execute_reply.started":"2026-03-01T07:42:11.212309Z","shell.execute_reply":"2026-03-01T07:42:15.484019Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T07:42:15.486282Z","iopub.execute_input":"2026-03-01T07:42:15.486512Z","iopub.status.idle":"2026-03-01T07:42:18.822225Z","shell.execute_reply.started":"2026-03-01T07:42:15.486487Z","shell.execute_reply":"2026-03-01T07:42:18.821575Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python preprocess/convert_templates_to_pt_files.py --input_csv /kaggle/working/pred_tbm.csv --output_name templates.pt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T07:42:18.823453Z","iopub.execute_input":"2026-03-01T07:42:18.823696Z","iopub.status.idle":"2026-03-01T07:42:28.241811Z","shell.execute_reply.started":"2026-03-01T07:42:18.82367Z","shell.execute_reply":"2026-03-01T07:42:28.241132Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"IS_SCORING_RUN = os.environ.get('KAGGLE_IS_COMPETITION_RERUN')\nprint(IS_SCORING_RUN)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T07:42:28.24316Z","iopub.execute_input":"2026-03-01T07:42:28.243718Z","iopub.status.idle":"2026-03-01T07:42:28.247896Z","shell.execute_reply.started":"2026-03-01T07:42:28.243689Z","shell.execute_reply":"2026-03-01T07:42:28.24736Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %%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'] >= 250) & (df['sequence_len'] < 1000)]\n# df = df[df['sequence_len'] < 600]\ndf = df[df['sequence_len'] < 1000] # 1000 2000\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\")\ndf.to_csv('/kaggle/working/sample_sequences.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T08:43:50.305029Z","iopub.execute_input":"2026-03-01T08:43:50.305619Z","iopub.status.idle":"2026-03-01T08:43:50.317648Z","shell.execute_reply.started":"2026-03-01T08:43:50.305594Z","shell.execute_reply":"2026-03-01T08:43:50.316961Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile runner/inference.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    \"\"\" Overlap領域で各チャンクの座標を結合し、線形補間でブレンドする \"\"\"\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:]  # 前のチャンクの後方部分\n        B = curr_coords[:overlap]    # 今のチャンクの前方部分\n\n        # Kabschで回転行列と並進ベクトルを計算\n        R, t = kabsch_umeyama(A, B)\n\n        # チャンク全体に変換を適用\n        curr_coords_transformed = (curr_coords @ R.T) + t\n\n        # Overlap部分を線形補間（スムージング）\n        weights = np.linspace(1, 0, overlap).reshape(-1, 1)\n        blended_overlap = A * weights + curr_coords_transformed[:overlap] * (1 - weights)\n\n        # 最終座標を更新（Overlap部分を上書きし、残りを追加）\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 # オーバーラップさせる塩基数\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            # 短い配列はそのまま予測\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            # 長い配列はチャンクに分割\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                    # 各チャンクを推論                  \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                # 重なり合う部分（overlap）のサイズを計算して結合\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,"execution":{"iopub.status.busy":"2026-03-01T09:00:37.071897Z","iopub.execute_input":"2026-03-01T09:00:37.072482Z","iopub.status.idle":"2026-03-01T09:00:37.084078Z","shell.execute_reply.started":"2026-03-01T09:00:37.072447Z","shell.execute_reply":"2026-03-01T09:00:37.083374Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile rnapro_inference_kaggle.sh\n\nexport LAYERNORM_TYPE=torch # fast_layernorm, torch\n\n\n# Inference parameters (RNAPro)\nSEED=42\nN_SAMPLE=1\nN_STEP=250 # 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.py \\\n    --model_name \"${MODEL_NAME}\" \\\n    --seeds ${SEED} \\\n    --dump_dir \"${DUMP_DIR}\" \\\n    --load_checkpoint_path \"${CHECKPOINT_PATH}\" \\\n    --use_msa true \\\n    --use_template \"ca_precomputed\" \\\n    --model.use_template \"ca_precomputed\" \\\n    --model.use_RibonanzaNet2 true \\\n    --model.template_embedder.n_blocks 2 \\\n    --model.ribonanza_net_path \"${RIBONANZA_PATH}\" \\\n    --template_data \"${TEMPLATE_DATA}\" \\\n    --template_idx ${TEMPLATE_IDX} \\\n    --rna_msa_dir \"${RNA_MSA_DIR}\" \\\n    --model.N_cycle ${N_CYCLE} \\\n    --sample_diffusion.N_sample ${N_SAMPLE} \\\n    --sample_diffusion.N_step ${N_STEP} \\\n    --load_strict true \\\n    --num_workers 0 \\\n    --triangle_attention \"torch\" \\\n    --triangle_multiplicative \"torch\" \\\n    --sequences_csv \"${SEQUENCES_CSV}\" \\\n    --max_len 1000\n\n\n# --triangle_attention supports 'triattention', 'cuequivariance', 'deepspeed', 'torch'\n# --triangle_multiplicative supports 'cuequivariance', 'torch'\n# --max_len 1000: Sequences longer than max_len will be skipped to avoid oom","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T09:00:37.085305Z","iopub.execute_input":"2026-03-01T09:00:37.085516Z","iopub.status.idle":"2026-03-01T09:00:37.104419Z","shell.execute_reply.started":"2026-03-01T09:00:37.08549Z","shell.execute_reply":"2026-03-01T09:00:37.103924Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!bash ./rnapro_inference_kaggle.sh","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T09:00:37.167289Z","iopub.execute_input":"2026-03-01T09:00:37.167496Z","iopub.status.idle":"2026-03-01T09:16:36.937365Z","shell.execute_reply.started":"2026-03-01T09:00:37.167479Z","shell.execute_reply":"2026-03-01T09:16:36.936537Z"},"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!mv submission.csv ..","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T08:39:03.199408Z","iopub.execute_input":"2026-03-01T08:39:03.199674Z","iopub.status.idle":"2026-03-01T08:39:03.37544Z","shell.execute_reply.started":"2026-03-01T08:39:03.199647Z","shell.execute_reply":"2026-03-01T08:39:03.374432Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pred_rnapro = pd.read_csv(\"/kaggle/working/rnapro_submission.csv\")\npred_rnapro","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T09:16:47.992584Z","iopub.execute_input":"2026-03-01T09:16:47.993015Z","iopub.status.idle":"2026-03-01T09:16:48.018965Z","shell.execute_reply.started":"2026-03-01T09:16:47.992981Z","shell.execute_reply":"2026-03-01T09:16:48.018414Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Boltz2","metadata":{}},{"cell_type":"code","source":"!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","metadata":{"trusted":true,"_kg_hide-input":false,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2026-02-20T23:48:46.627665Z","iopub.execute_input":"2026-02-20T23:48:46.628392Z","iopub.status.idle":"2026-02-20T23:48:49.985799Z","shell.execute_reply.started":"2026-02-20T23:48:46.628344Z","shell.execute_reply":"2026-02-20T23:48:49.984986Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!cp -r /kaggle/input/datasets/lbugnon/boltz-src-minimal ./\n# Install boltz from src (some weird bug in kaggle notebooks with fairscale so it was removed)\n!pip install --no-index --no-build-isolation -e ./boltz-src-minimal","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T23:48:50.282574Z","iopub.execute_input":"2026-02-20T23:48:50.283581Z","iopub.status.idle":"2026-02-20T23:48:58.205032Z","shell.execute_reply.started":"2026-02-20T23:48:50.283537Z","shell.execute_reply":"2026-02-20T23:48:58.204177Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install --no-index /kaggle/input/datasets/youhanlee/boltz-dependencies/mashumaro-3.14-py3-none-any.whl --no-deps\n!pip install --no-index /kaggle/input/datasets/youhanlee/boltz-dependencies/ihm-2.2-py3-none-any.whl --no-deps\n!pip install --no-index /kaggle/input/datasets/youhanlee/boltz-dependencies/modelcif-1.3-py3-none-any.whl --no-deps","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T23:48:58.206716Z","iopub.execute_input":"2026-02-20T23:48:58.206985Z","iopub.status.idle":"2026-02-20T23:49:02.607338Z","shell.execute_reply.started":"2026-02-20T23:48:58.206954Z","shell.execute_reply":"2026-02-20T23:49:02.606618Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# need to create tar so botlz do not try to download, also clean the trash dataset creation does\n!mkdir boltz_cache\n!cp -r /kaggle/input/datasets/lbugnon/boltz2 boltz_cache\n!mv boltz_cache/boltz2/mols/mols/* boltz_cache/boltz2/mols/\n!rm -r boltz_cache/boltz2/mols/mols/\n!tar -cf boltz_cache/boltz2/mols.tar boltz_cache/boltz2/mols","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T23:49:02.60858Z","iopub.execute_input":"2026-02-20T23:49:02.60884Z","iopub.status.idle":"2026-02-20T23:54:23.405749Z","shell.execute_reply.started":"2026-02-20T23:49:02.608812Z","shell.execute_reply":"2026-02-20T23:54:23.40467Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\nMAX_LENGTH = 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\n# mask = sequences[\"sequence_len\"] >= MAX_LENGTH\n# sequences.loc[mask, \"sequence\"] = (\n#     sequences.loc[mask, \"sequence\"].str.slice(0, MAX_LENGTH)\n# )\n\n# # 長さを再計算\n# sequences[\"sequence_len\"] = sequences[\"sequence\"].str.len()\n\n# # sequences = sequences[sequences['sequence_len'] < 100]\n# sequences = sequences[sequences['sequence_len'] < 1000]\nprint(sequences.shape)\n\nprint(max(sequences[\"sequence_len\"].values))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T23:54:23.408192Z","iopub.execute_input":"2026-02-20T23:54:23.40856Z","iopub.status.idle":"2026-02-20T23:54:23.734621Z","shell.execute_reply.started":"2026-02-20T23:54:23.408528Z","shell.execute_reply":"2026-02-20T23:54:23.733942Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# boltz２用のyaml作成\nimport yaml\nimport os\nfrom pathlib import Path\n\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    # --- リガンド処理 ---\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    # --- 書き出し ---\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# 実行例\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\")","metadata":{"trusted":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2026-02-20T23:54:23.735484Z","iopub.execute_input":"2026-02-20T23:54:23.735747Z","iopub.status.idle":"2026-02-20T23:54:23.778169Z","shell.execute_reply.started":"2026-02-20T23:54:23.735724Z","shell.execute_reply":"2026-02-20T23:54:23.777566Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# with open('/kaggle/working/inputs/9E74.yaml', encoding='utf-8')as f:\n#     document = yaml.safe_load(f)\n# document","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T13:28:34.710024Z","iopub.execute_input":"2026-02-19T13:28:34.710265Z","iopub.status.idle":"2026-02-19T13:28:34.713479Z","shell.execute_reply.started":"2026-02-19T13:28:34.710246Z","shell.execute_reply":"2026-02-19T13:28:34.712794Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!rm -r /kaggle/working/boltz_repeat_0*","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T23:54:23.779008Z","iopub.execute_input":"2026-02-20T23:54:23.779482Z","iopub.status.idle":"2026-02-20T23:54:23.897167Z","shell.execute_reply.started":"2026-02-20T23:54:23.779455Z","shell.execute_reply":"2026-02-20T23:54:23.896325Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import time\nfrom random import shuffle\nimport gc\nimport torch\n        \n!mkdir inputs\n\nuse_all_chains = False\nmax_tokens = 1000\nmax_repeats = 999\npred_repeats = 1\n\nnucleotides = {\"A\", \"G\", \"C\", \"U\"}\naminoacids = {\"A\", \"R\", \"N\", \"D\", \"C\", \"E\", \"Q\", \"G\", \"H\", \"I\", \"L\", \"K\", \"M\", \"F\", \"P\", \"S\", \"T\", \"W\", \"Y\", \"V\"}\n\nimport pytorch_lightning as pl\nfrom functools import partial\n\n# 1. Trainerの初期化を「改造」する\n# これにより、Boltzが内部でTrainerを作るときに、強制的に logger=False にします\noriginal_trainer_init = pl.Trainer.__init__\n\ndef patched_trainer_init(self, *args, **kwargs):\n    # ロガーを無効化し、エラーの元になるロギングを物理的に遮断\n    kwargs[\"logger\"] = False\n    # ついでに、進捗バー以外のログ出力も最小限にする\n    kwargs[\"enable_checkpointing\"] = False\n    return original_trainer_init(self, *args, **kwargs)\n\n# Trainerのコンストラクタを差し替え\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. Firs load all chains. This may fail if all_sequences is empty or different format\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                elif use_all_chains:\n                    entry_type = \"rna\" if set(entry)<=nucleotides else \"protein\"  if set(entry) <= aminoacids else None\n                    if entry_type is None:\n                        continue\n                    chains[chain_ind] = entry, entry_type\n                    chain_ind += 1\n                if use_all_chains:\n                    for i in range(1, repeats):\n                        chains[chain_ind] = entry, entry_type\n                        chain_ind += 1\n    \n        # Now keep the chain 0 plus all chains randomly picked to fit max_tokens (asuming chain 0 is ok under tokens, otherwise will fail)\n        filtered_chains, ntokens = {}, 0\n        other_chains = [c for c in chains if c!=0]\n        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 fasta, chain 0 first\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（存在する場合のみ）---\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    # --- setting ---\n    boltz_args = {\n        \"diffusion_samples\": 10,\n        \"diffusion_samples_affinity\": 5, # 10,\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    !cat /kaggle/working/inputs/{name}.yaml\n    print()\n    for repeat in range(pred_repeats):\n        cmd = f\"\"\"\n        boltz predict /kaggle/working/inputs/{name}.yaml \\\n        --diffusion_samples {boltz_args['diffusion_samples']} \\\n        --diffusion_samples_affinity {boltz_args['diffusion_samples_affinity']} \\\n        --devices {boltz_args['devices']} \\\n        --num_workers {boltz_args['num_workers']} \\\n        --max_parallel_samples {boltz_args['max_parallel_samples']} \\\n        --output_format {boltz_args['output_format']} \\\n        --cache {boltz_args['cache']} \\\n        --out_dir boltz_repeat_{repeat}\n        \"\"\"\n        !{cmd}\n        # !boltz predict /kaggle/working/inputs/{name}.yaml --diffusion_samples 10 --diffusion_samples_affinity 10 --devices 2 --num_workers 1\t--max_parallel_samples 1 --output_format pdb --cache boltz_cache/boltz2/ --out_dir boltz_repeat_{repeat} \n        gc.collect()\n        torch.cuda.empty_cache()\n        \n    #     break\n    # break","metadata":{"trusted":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2026-02-20T23:54:23.898506Z","iopub.execute_input":"2026-02-20T23:54:23.898752Z","iopub.status.idle":"2026-02-20T23:56:07.458677Z","shell.execute_reply.started":"2026-02-20T23:54:23.898726Z","shell.execute_reply":"2026-02-20T23:56:07.457971Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# {'confidence_score': 0.659386396408081,\n#  'ptm': 0.4263673722743988,\n#  'iptm': 0.0,\n#  'ligand_iptm': 0.0,\n#  'protein_iptm': 0.0,\n#  'complex_plddt': 0.7176411747932434,\n#  'complex_iplddt': 0.7176411747932434,\n#  'complex_pde': 0.4417988955974579,\n#  'complex_ipde': 0.0,\n#  'chains_ptm': {'0': 0.4263673722743988},\n#  'pair_chains_iptm': {'0': {'0': 0.4263673722743988}}}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T14:21:29.164115Z","iopub.execute_input":"2026-02-19T14:21:29.164975Z","iopub.status.idle":"2026-02-19T14:21:29.168274Z","shell.execute_reply.started":"2026-02-19T14:21:29.164944Z","shell.execute_reply":"2026-02-19T14:21:29.167599Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==========================================\n# 1. 座標抽出の関数\n# ==========================================\nimport json\nimport gemmi\nimport numpy as np\nfrom Bio.PDB.MMCIF2Dict import MMCIF2Dict\n\ndef extract_coords(pdb_file, cif_file):\n    \"\"\"mmCIFファイルからC1'原子の座標を抽出\"\"\"\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. モデル選択のロジック\n# ==========================================\n\ndef get_top_n_predictions(diffusion_results_dir: Path, file_prefix: str, n: int = 3):\n    \"\"\"10個の結果からスコア上位n個の情報を取得\"\"\"\n    results = []\n    for i in range(10):\n        json_path = diffusion_results_dir / file_prefix / f\"confidence_{file_prefix}_model_{i}.json\"\n        # print(json_path)\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    # スコア順にソートして上位n個を昇順で返す\n    results.sort(key=lambda x: (x[\"plddt\"], x[\"conf\"]), reverse=True)\n    # top_n = results[:n]\n    # top_n.sort(key=lambda x: (x[\"plddt\"], x[\"conf\"]), reverse=False)\n    return results[:n] # top_n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T23:56:16.315824Z","iopub.execute_input":"2026-02-20T23:56:16.316379Z","iopub.status.idle":"2026-02-20T23:56:16.408726Z","shell.execute_reply.started":"2026-02-20T23:56:16.316342Z","shell.execute_reply":"2026-02-20T23:56:16.408095Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==========================================\n# 3. メイン処理\n# ==========================================\n\ndef main():\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    # すべての座標カラムを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/RNAPro/boltz_repeat_0/boltz_results_{target_id}/predictions/\")\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        # ターゲットに該当する行のマスク\n        target_mask = submission['ID'].apply(lambda x: x.startswith(f\"{target_id}_\"))\n        expected_count = target_mask.sum()\n\n        # 各モデル（1〜5番目）の座標を格納\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                # 座標数の調整（足りなければ0埋め、多ければカット）\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                # 対応する列（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        # モデルが3つ未満の場合や、1000塩基超の場合の処理は、\n        # 初期値（0.0）のままになるか、必要に応じてここでゼロ埋めを明示します。\n\n    # 保存\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()","metadata":{"trusted":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2026-02-20T23:57:28.869263Z","iopub.execute_input":"2026-02-20T23:57:28.870231Z","iopub.status.idle":"2026-02-20T23:57:29.156032Z","shell.execute_reply.started":"2026-02-20T23:57:28.870192Z","shell.execute_reply":"2026-02-20T23:57:29.155299Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"boltz_pred = pd.read_csv(\"/kaggle/working/boltz_submission.csv\")\nboltz_pred\n# 信頼度高い順","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T23:57:39.146071Z","iopub.execute_input":"2026-02-20T23:57:39.146744Z","iopub.status.idle":"2026-02-20T23:57:39.191321Z","shell.execute_reply.started":"2026-02-20T23:57:39.146712Z","shell.execute_reply":"2026-02-20T23:57:39.190763Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"boltz_pred.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T23:57:43.774371Z","iopub.execute_input":"2026-02-20T23:57:43.774745Z","iopub.status.idle":"2026-02-20T23:57:43.786844Z","shell.execute_reply.started":"2026-02-20T23:57:43.77471Z","shell.execute_reply":"2026-02-20T23:57:43.786091Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 結合","metadata":{}},{"cell_type":"code","source":"# TBM、Boltzは信頼度が高いものを1, 低いものを5","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# test_seq[[\"target_id\", \"sequence_len\"]]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T14:21:38.045916Z","iopub.execute_input":"2026-02-19T14:21:38.046167Z","iopub.status.idle":"2026-02-19T14:21:38.054928Z","shell.execute_reply.started":"2026-02-19T14:21:38.046144Z","shell.execute_reply":"2026-02-19T14:21:38.054235Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# sequences boltz_pred pred_tbm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T14:21:38.056139Z","iopub.execute_input":"2026-02-19T14:21:38.056388Z","iopub.status.idle":"2026-02-19T14:21:38.064735Z","shell.execute_reply.started":"2026-02-19T14:21:38.05633Z","shell.execute_reply":"2026-02-19T14:21:38.064065Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# boltz_pred = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/sample_submission.csv')\n# pred_tbm = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/sample_submission.csv')\n# test_seqs = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv')\n# sequences = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv')\n# pred_rnapro = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/sample_submission.csv')\n# drfold_pred = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/sample_submission.csv')\n# protenix_pred = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/sample_submission.csv')\n# boltz_pred.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T08:59:03.172407Z","iopub.execute_input":"2026-02-28T08:59:03.173141Z","iopub.status.idle":"2026-02-28T08:59:03.241159Z","shell.execute_reply.started":"2026-02-28T08:59:03.173113Z","shell.execute_reply":"2026-02-28T08:59:03.240454Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pred_tbm_ = pred_tbm.copy()\nboltz_pred_ = boltz_pred.copy()\npred_rnapro_ = pred_rnapro.copy()\ndrfold_pred_ = drfold_pred.copy()\nprotenix_pred_ = protenix_pred.copy()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T08:59:20.429596Z","iopub.execute_input":"2026-02-28T08:59:20.429915Z","iopub.status.idle":"2026-02-28T08:59:20.438682Z","shell.execute_reply.started":"2026-02-28T08:59:20.429863Z","shell.execute_reply":"2026-02-28T08:59:20.438088Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ① pred_tbm側に target_id 列を作る（_より前を抽出）\npred_tbm_[\"target_id\"] = pred_tbm_[\"ID\"].str.split(\"_\").str[0]\n\n# ② test_seq から必要な列だけ抽出\ntest_seqs[\"sequence_len\"] = test_seqs[\"sequence\"].apply(lambda x : len(x))\nseq_len_df = test_seqs[[\"target_id\", \"sequence_len\"]]\n\n# ③ merge\npred_tbm_ = pred_tbm_.merge(seq_len_df, on=\"target_id\", how=\"left\")\n\n# 4 boltz_pred側にtarget_id 列を作る（_より前を抽出）\nboltz_pred_[\"target_id\"] = boltz_pred_[\"ID\"].str.split(\"_\").str[0]\n\n# 5 seaquences から必要な列だけ抽出\n# sequences[\"sequence_len\"] = sequences[\"sequence\"].apply(lambda x : len(x))\nseq_len_df = sequences[[\"target_id\", \"sequence_len\"]]\n\n# 6 boltz_predにmerge\nboltz_pred_ = boltz_pred_.merge(seq_len_df, on=\"target_id\", how=\"left\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T08:59:21.522189Z","iopub.execute_input":"2026-02-28T08:59:21.522491Z","iopub.status.idle":"2026-02-28T08:59:21.563674Z","shell.execute_reply.started":"2026-02-28T08:59:21.522464Z","shell.execute_reply":"2026-02-28T08:59:21.563119Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# merge用にprotenix_predの準備\ntest_df = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv\")\ntest_df[\"sequence_len\"] = test_df[\"sequence\"].str.len()\nseq_len_df = test_df[[\"target_id\", \"sequence_len\"]]\nprotenix_pred_[\"target_id\"] = protenix_pred_[\"ID\"].str.split(\"_\").str[0]\nprotenix_pred_ = protenix_pred_.merge(seq_len_df, on=\"target_id\", how=\"left\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T08:59:24.550468Z","iopub.execute_input":"2026-02-28T08:59:24.551028Z","iopub.status.idle":"2026-02-28T08:59:24.562896Z","shell.execute_reply.started":"2026-02-28T08:59:24.551Z","shell.execute_reply":"2026-02-28T08:59:24.562319Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# merge用にpred_rnaproの準備\n# test_df = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv\")\n# test_df[\"sequence_len\"] = test_df[\"sequence\"].str.len()\n# seq_len_df = test_df[[\"target_id\", \"sequence_len\"]]\n# pred_rnapro_[\"target_id\"] = pred_rnapro_[\"ID\"].str.split(\"_\").str[0]\n# pred_rnapro_ = pred_rnapro_.merge(seq_len_df, on=\"target_id\", how=\"left\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T08:59:26.138063Z","iopub.execute_input":"2026-02-28T08:59:26.138626Z","iopub.status.idle":"2026-02-28T08:59:26.142147Z","shell.execute_reply.started":"2026-02-28T08:59:26.138596Z","shell.execute_reply":"2026-02-28T08:59:26.141427Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pred_tbm_","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T08:59:26.328625Z","iopub.execute_input":"2026-02-28T08:59:26.329441Z","iopub.status.idle":"2026-02-28T08:59:26.341625Z","shell.execute_reply.started":"2026-02-28T08:59:26.329411Z","shell.execute_reply":"2026-02-28T08:59:26.340944Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# filterd_pred_tbm = pred_tbm_[pred_tbm_['sequence_len'] >= 100]\nfilterd_pred_tbm = pred_tbm_[pred_tbm_['sequence_len'] >= 250]\n# filterd_pred_tbm = pred_tbm_[pred_tbm_['sequence_len'] >= 350]\n# filterd_boltz_pred = boltz_pred_[boltz_pred_['sequence_len'] < 100]\nfilterd_boltz_pred = boltz_pred_[boltz_pred_['sequence_len'] < 250]\n# filterd_boltz_pred = boltz_pred_[boltz_pred_['sequence_len'] < 350]\n\nmerged_pred = pd.concat([filterd_pred_tbm, filterd_boltz_pred], axis=0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T08:59:31.006848Z","iopub.execute_input":"2026-02-28T08:59:31.007629Z","iopub.status.idle":"2026-02-28T08:59:31.016037Z","shell.execute_reply.started":"2026-02-28T08:59:31.007598Z","shell.execute_reply":"2026-02-28T08:59:31.01534Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# RNAProを結合\n# 将来的にはRNAProの信頼度を高いものから左に並べたい\n\n# target_cols = [\"x_3\", \"y_3\", \"z_3\", \"x_4\", \"y_4\", \"z_4\", \"x_5\", \"y_5\", \"z_5\"]\n# target_cols = [\"x_3\", \"y_3\", \"z_3\", \"x_4\", \"y_4\", \"z_4\"]\ntarget_cols = [\"x_1\", \"y_1\", \"z_1\", \"x_2\", \"y_2\", \"z_2\"]\nreplace_cols = [\"x_3\", \"y_3\", \"z_3\", \"x_4\", \"y_4\", \"z_4\"]\n# target_cols = [\"x_1\", \"y_1\", \"z_1\", \"x_2\", \"y_2\", \"z_2\"]\n# replace_cols = [\"x_4\", \"y_4\", \"z_4\", \"x_5\", \"y_5\", \"z_5\"]\n\n# 1. 一旦 ID 列をインデックスに変更（既存の ID 列名を 'ID_col' と仮定）\n#    ※ 既に ID がインデックスならこの工程は不要です\nmerged_pred = merged_pred.set_index(\"ID\") \n\n# 2. loc を使って置換（右辺に .values を忘れずに）\n# filterd_pred_rnapro = pred_rnapro_[(pred_rnapro_['sequence_len'] >= 400) &  (pred_rnapro_['sequence_len'] < 1000)]\nmerged_pred.loc[pred_rnapro_[\"ID\"], replace_cols] = pred_rnapro_[target_cols].values\n# merged_pred.loc[pred_rnapro_[\"ID\"], replace_cols] = filterd_pred_rnapro[target_cols].values\n\nmerged_pred.reset_index(drop=False, inplace=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T08:59:33.164567Z","iopub.execute_input":"2026-02-28T08:59:33.164861Z","iopub.status.idle":"2026-02-28T08:59:33.180202Z","shell.execute_reply.started":"2026-02-28T08:59:33.164837Z","shell.execute_reply":"2026-02-28T08:59:33.179555Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# DRfoldを結合\n\n# target_cols = [\"x_5\", \"y_5\", \"z_5\"]\ntarget_cols = [\"x_1\", \"y_1\", \"z_1\"]\nreplace_cols = [\"x_5\", \"y_5\", \"z_5\"]\n\n# 1. 一旦 ID 列をインデックスに変更（既存の ID 列名を 'ID_col' と仮定）\n#    ※ 既に ID がインデックスならこの工程は不要です\nmerged_pred = merged_pred.set_index(\"ID\") \n\n# 2. loc を使って置換（右辺に .values を忘れずに）\nmerged_pred.loc[drfold_pred_[\"ID\"], replace_cols] = drfold_pred_[target_cols].values\n\nmerged_pred.reset_index(drop=False, inplace=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T08:59:34.950129Z","iopub.execute_input":"2026-02-28T08:59:34.950425Z","iopub.status.idle":"2026-02-28T08:59:34.965788Z","shell.execute_reply.started":"2026-02-28T08:59:34.950394Z","shell.execute_reply":"2026-02-28T08:59:34.965105Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Boltzを結合\n# target_cols = [\"x_2\", \"y_2\", \"z_2\", \"x_5\", \"y_5\", \"z_5\"]\ntarget_cols = [\"x_1\", \"y_1\", \"z_1\", \"x_2\", \"y_2\", \"z_2\"]\nreplace_cols = [\"x_2\", \"y_2\", \"z_2\", \"x_5\", \"y_5\", \"z_5\"]\n\n# 1. 一旦 ID 列をインデックスに変更（既存の ID 列名を 'ID_col' と仮定）\n#    ※ 既に ID がインデックスならこの工程は不要です\nmerged_pred = merged_pred.set_index(\"ID\") \n\n# 2. loc を使って置換（右辺に .values を忘れずに）\n# filterd_boltz_pred = boltz_pred_[(boltz_pred_['sequence_len'] >= 350) & (boltz_pred_['sequence_len'] < 1000)]\nfilterd_boltz_pred = boltz_pred_[(boltz_pred_['sequence_len'] >= 250) & (boltz_pred_['sequence_len'] < 1000)]\nmerged_pred.loc[filterd_boltz_pred[\"ID\"], replace_cols] = filterd_boltz_pred[target_cols].values\n\nmerged_pred.reset_index(drop=False, inplace=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T08:59:36.756247Z","iopub.execute_input":"2026-02-28T08:59:36.756998Z","iopub.status.idle":"2026-02-28T08:59:36.770969Z","shell.execute_reply.started":"2026-02-28T08:59:36.75697Z","shell.execute_reply":"2026-02-28T08:59:36.770371Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Protenixを結合\n# target_cols = [\"x_4\", \"y_4\", \"z_4\"]\ntarget_cols = [\"x_1\", \"y_1\", \"z_1\"]\nreplace_cols = [\"x_4\", \"y_4\", \"z_4\"]\n# target_cols = [\"x_1\", \"y_1\", \"z_1\", \"x_2\", \"y_2\", \"z_2\"]\n# replace_cols = [\"x_3\", \"y_3\", \"z_3\", \"x_4\", \"y_4\", \"z_4\"]\n\n# 1. 一旦 ID 列をインデックスに変更（既存の ID 列名を 'ID_col' と仮定）\n#    ※ 既に ID がインデックスならこの工程は不要です\nmerged_pred = merged_pred.set_index(\"ID\") \n\n# filterd_protenix_pred = protenix_pred_[(protenix_pred_['sequence_len'] >= 350) & (protenix_pred_['sequence_len'] < 1000)]\n# filterd_protenix_pred = protenix_pred_[protenix_pred_['sequence_len'] < 350]\nfilterd_protenix_pred = protenix_pred_[protenix_pred_['sequence_len'] < 250]\n# filterd_protenix_pred = protenix_pred_[protenix_pred_['sequence_len'] < 400]\n# filterd_protenix_pred = protenix_pred_[protenix_pred_['sequence_len'] < 1000]\nmerged_pred.loc[filterd_protenix_pred[\"ID\"], replace_cols] = filterd_protenix_pred[target_cols].values\n# merged_pred.loc[protenix_pred_[\"ID\"], replace_cols] = protenix_pred_[target_cols].values\n\nmerged_pred.reset_index(drop=False, inplace=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Protenixを結合\ntarget_cols = [\"x_1\", \"y_1\", \"z_1\", \"x_2\", \"y_2\", \"z_2\"]\nreplace_cols = [\"x_4\", \"y_4\", \"z_4\", \"x_5\", \"y_5\", \"z_5\"]\n\n# 1. 一旦 ID 列をインデックスに変更（既存の ID 列名を 'ID_col' と仮定）\n#    ※ 既に ID がインデックスならこの工程は不要です\nmerged_pred = merged_pred.set_index(\"ID\") \n\nfilterd_protenix_pred = protenix_pred_[protenix_pred_['sequence_len'] >= 1000]\nmerged_pred.loc[filterd_protenix_pred[\"ID\"], replace_cols] = filterd_protenix_pred[target_cols].values\n\nmerged_pred.reset_index(drop=False, inplace=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# RNAProを結合\n\n# target_cols = [\"x_3\", \"y_3\", \"z_3\", \"x_4\", \"y_4\", \"z_4\", \"x_5\", \"y_5\", \"z_5\"]\n# target_cols = [\"x_3\", \"y_3\", \"z_3\", \"x_4\", \"y_4\", \"z_4\"]\n# target_cols = [\"x_1\", \"y_1\", \"z_1\", \"x_2\", \"y_2\", \"z_2\"]\n# replace_cols = [\"x_3\", \"y_3\", \"z_3\", \"x_4\", \"y_4\", \"z_4\"]\n# target_cols = [\"x_1\", \"y_1\", \"z_1\"]\n# replace_cols = [\"x_3\", \"y_3\", \"z_3\"]\n\n# 1. 一旦 ID 列をインデックスに変更（既存の ID 列名を 'ID_col' と仮定）\n#    ※ 既に ID がインデックスならこの工程は不要です\n# merged_pred = merged_pred.set_index(\"ID\") \n\n# 2. loc を使って置換（右辺に .values を忘れずに）\n# filterd_pred_rnapro = pred_rnapro_[(pred_rnapro_['sequence_len'] >= 250) &  (pred_rnapro_['sequence_len'] < 400)]\n# merged_pred.loc[pred_rnapro_[\"ID\"], replace_cols] = pred_rnapro_[target_cols].values\n# merged_pred.loc[pred_rnapro_[\"ID\"], replace_cols] = filterd_pred_rnapro[target_cols].values\n\n# merged_pred.reset_index(drop=False, inplace=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# merged_pred.query('target_id == \"9LEL\"')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T09:02:38.746171Z","iopub.execute_input":"2026-02-28T09:02:38.746757Z","iopub.status.idle":"2026-02-28T09:02:38.750132Z","shell.execute_reply.started":"2026-02-28T09:02:38.746726Z","shell.execute_reply":"2026-02-28T09:02:38.749143Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# merged_pred.iloc[protenix_pred_[\"ID\"]]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T14:21:38.251753Z","iopub.execute_input":"2026-02-19T14:21:38.25202Z","iopub.status.idle":"2026-02-19T14:21:38.260912Z","shell.execute_reply.started":"2026-02-19T14:21:38.251997Z","shell.execute_reply":"2026-02-19T14:21:38.260187Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub = 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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T09:01:54.423505Z","iopub.execute_input":"2026-02-28T09:01:54.42381Z","iopub.status.idle":"2026-02-28T09:01:54.461095Z","shell.execute_reply.started":"2026-02-28T09:01:54.423786Z","shell.execute_reply":"2026-02-28T09:01:54.460501Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"output_path = '/kaggle/working/submission.csv'\nsub.to_csv(output_path, index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T09:02:01.22748Z","iopub.execute_input":"2026-02-28T09:02:01.228112Z","iopub.status.idle":"2026-02-28T09:02:01.281931Z","shell.execute_reply.started":"2026-02-28T09:02:01.228082Z","shell.execute_reply":"2026-02-28T09:02:01.281345Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Evaluation","metadata":{}},{"cell_type":"code","source":"PATH = \"/kaggle/input/stanford-rna-3d-folding-2/\"\n\nsolution = pd.read_csv(os.path.join(PATH, 'validation_labels.csv'))\nsolution.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T01:14:23.604112Z","iopub.execute_input":"2026-02-23T01:14:23.604422Z","iopub.status.idle":"2026-02-23T01:14:23.874334Z","shell.execute_reply.started":"2026-02-23T01:14:23.604398Z","shell.execute_reply":"2026-02-23T01:14:23.87345Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# solution = pd.merge(pred_rnapro[\"ID\"], solution, how='left', on='ID')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T14:21:38.708367Z","iopub.execute_input":"2026-02-19T14:21:38.708748Z","iopub.status.idle":"2026-02-19T14:21:38.712564Z","shell.execute_reply.started":"2026-02-19T14:21:38.708712Z","shell.execute_reply":"2026-02-19T14:21:38.711919Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# def parse_tmscore_output(output):\n#     # Extract TM-score based on length of reference structure (second)\n#     tm_score_match = re.findall(r'TM-score=\\s+([\\d.]+)', output)[1]\n#     if not tm_score_match:\n#         raise ValueError('No TM score found')\n#     return float(tm_score_match)\n\n\n# def write_target_line(\n#     atom_name, atom_serial, residue_name, chain_id, residue_num, x_coord, y_coord, z_coord, occupancy=1.0, b_factor=0.0, atom_type='P'\n# ) -> str:\n#     \"\"\"\n#     Writes a single line of PDB format based on provided atom information.\n\n#     Args:\n#         atom_name (str): Name of the atom (e.g., \"N\", \"CA\").\n#         atom_serial (int): Atom serial number.\n#         residue_name (str): Residue name (e.g., \"ALA\").\n#         chain_id (str): Chain identifier.\n#         residue_num (int): Residue number.\n#         x_coord (float): X coordinate.\n#         y_coord (float): Y coordinate.\n#         z_coord (float): Z coordinate.\n#         occupancy (float, optional): Occupancy value (default: 1.0).\n#         b_factor (float, optional): B-factor value (default: 0.0).\n\n#     Returns:\n#         str: A single line of PDB string.\n#     \"\"\"\n#     return f'ATOM  {atom_serial:>5d}  {atom_name:<5s} {residue_name:<3s} {residue_num:>3d}    {x_coord:>8.3f}{y_coord:>8.3f}{z_coord:>8.3f}{occupancy:>6.2f}{b_factor:>6.2f}           {atom_type}\\n'\n\n\n# def write2pdb(df: pd.DataFrame, xyz_id: str, target_path: str) -> int:\n#     resolved_cnt = 0\n#     with open(target_path, 'w') as target_file:\n#         for _, row in df.iterrows():\n#             x_coord = row[f'x_{xyz_id}']\n#             y_coord = row[f'y_{xyz_id}']\n#             z_coord = row[f'z_{xyz_id}']\n\n#             if x_coord > -1e17 and y_coord > -1e17 and z_coord > -1e17:\n#                 # if True:\n#                 resolved_cnt += 1\n#                 target_line = write_target_line(\n#                     atom_name=\"C1'\",\n#                     atom_serial=int(row['resid']),\n#                     residue_name=row['resname'],\n#                     chain_id='0',\n#                     residue_num=int(row['resid']),\n#                     x_coord=x_coord,\n#                     y_coord=y_coord,\n#                     z_coord=z_coord,\n#                     atom_type='C',\n#                 )\n#                 target_file.write(target_line)\n#     return resolved_cnt\n\n\n# def score(solution: pd.DataFrame, submission: pd.DataFrame, row_id_column_name: str) -> float:\n#     \"\"\"\n#     Computes the TM-score between predicted and native RNA structures using USalign.\n\n#     This function evaluates the structural similarity of RNA predictions to native structures\n#     by computing the TM-score. It uses USalign, a structural alignment tool, to compare\n#     the predicted structures with the native structures.\n\n#     Workflow:\n#     1. Copies the USalign binary to the working directory and grants execution permissions.\n#     2. Extracts the `target_id` from the `ID` column of both the solution and submission DataFrames.\n#     3. Iterates over each unique `target_id`, grouping the native and predicted structures.\n#     4. Writes PDB files for native and predicted structures.\n#     5. Runs USalign on each predicted-native pair and extracts the TM-score.\n#     6. Computes the highest TM-score per target and returns aggregated results.\n\n#     Args:\n#         solution (pd.DataFrame): A DataFrame containing the native RNA structures.\n#         submission (pd.DataFrame): A DataFrame containing the predicted RNA structures.\n#         row_id_column_name (str): The name of the column containing unique row identifiers.\n\n#     Returns:\n#         float: the average highest TM-scores.\n#     \"\"\"\n    \n#     os.system('cp /kaggle/input/datasets/metric/usalign/USalign /kaggle/working')\n#     os.system('sudo chmod u+x /kaggle/working//USalign')\n\n#     # Extract target_id from ID (target_resid)\n#     solution['target_id'] = solution['ID'].apply(lambda x: x.split('_')[0])\n#     submission['target_id'] = submission['ID'].apply(lambda x: x.split('_')[0])\n\n#     results = {}\n#     # Iterate through each target_id and generate PDB files for both clean and corrupted data\n#     for target_id, group_native in tqdm(solution.groupby('target_id'), desc='TM-scoring'):\n#         group_predicted = submission[submission['target_id'] == target_id]\n#         native_pdb = 'native.pdb'\n#         predicted_pdb = 'predicted.pdb'\n\n#         target_id_scores = []\n#         for pred_cnt in range(1, 6):\n#             prediction_scores = []\n#             for native_cnt in range(1, 41):\n#                 # Write solution PDB\n#                 resolved_cnt = write2pdb(group_native, native_cnt, native_pdb)\n\n#                 # Write predicted PDB\n#                 _ = write2pdb(group_predicted, pred_cnt, predicted_pdb)\n\n#                 if resolved_cnt > 0:\n#                     command = f'/kaggle/working/USalign {predicted_pdb} {native_pdb} -atom \" C1\\'\"'\n#                     usalign_output = os.popen(command).read()\n#                     prediction_scores.append(parse_tmscore_output(usalign_output))\n\n#             target_id_scores.append(max(prediction_scores))\n#         results[target_id] = max(target_id_scores)\n\n#     return results","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T14:21:38.713535Z","iopub.execute_input":"2026-02-19T14:21:38.713759Z","iopub.status.idle":"2026-02-19T14:21:38.724757Z","shell.execute_reply.started":"2026-02-19T14:21:38.713736Z","shell.execute_reply":"2026-02-19T14:21:38.723992Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import os\n# import re\n# import subprocess\n# from multiprocessing import Pool, cpu_count\n# import numpy as np\n# import pandas as pd\n# from tqdm import tqdm\n\n\n# def parse_tmscore_output(output):\n#     m = re.findall(r'TM-score=\\s+([\\d.]+)', output)\n#     if len(m) < 2:\n#         return 0.0\n#     return float(m[1])\n\n\n# def write2pdb_fast(df, xyz_id, target_path):\n#     x = df[f'x_{xyz_id}'].to_numpy()\n#     y = df[f'y_{xyz_id}'].to_numpy()\n#     z = df[f'z_{xyz_id}'].to_numpy()\n#     resid = df['resid'].to_numpy()\n#     resname = df['resname'].to_numpy()\n\n#     mask = (x > -1e17) & (y > -1e17) & (z > -1e17)\n\n#     if not np.any(mask):\n#         return 0\n\n#     x = x[mask]\n#     y = y[mask]\n#     z = z[mask]\n#     resid = resid[mask]\n#     resname = resname[mask]\n\n#     with open(target_path, 'w') as f:\n#         for i in range(len(x)):\n#             line = (\n#                 f\"ATOM  {int(resid[i]):>5d}  C1'  {resname[i]:<3s} \"\n#                 f\"{int(resid[i]):>3d}    \"\n#                 f\"{x[i]:>8.3f}{y[i]:>8.3f}{z[i]:>8.3f}\"\n#                 f\"{1.0:>6.2f}{0.0:>6.2f}           C\\n\"\n#             )\n#             f.write(line)\n\n#     return len(x)\n\n\n# def run_usalign(args):\n#     pred_pdb, native_pdb = args\n#     cmd = [\n#         \"/kaggle/working/USalign\",\n#         pred_pdb,\n#         native_pdb,\n#         \"-atom\",\n#         \" C1'\"\n#     ]\n#     try:\n#         result = subprocess.run(\n#             cmd,\n#             capture_output=True,\n#             text=True,\n#             timeout=10  # ← 重要\n#         )\n#         return parse_tmscore_output(result.stdout)\n#     except subprocess.TimeoutExpired:\n#         print(f\"Timeout: {pred_pdb} vs {native_pdb}\")\n#         return 0.0\n\n\n# def score(solution: pd.DataFrame, submission: pd.DataFrame):\n#     os.system('cp /kaggle/input/datasets/metric/usalign/USalign /kaggle/working/')\n#     os.system('chmod u+x /kaggle/working/USalign')\n\n#     solution['target_id'] = solution['ID'].str.split('_').str[0]\n#     submission['target_id'] = submission['ID'].str.split('_').str[0]\n\n#     results = {}\n#     # n_workers = cpu_count()\n\n#     n_workers = min(4, cpu_count())\n\n#     for target_id, group_native in tqdm(solution.groupby('target_id')):\n    \n#         print(f\"Processing {target_id}\")\n    \n#         group_pred = submission[submission['target_id'] == target_id]\n    \n#         native_paths = []\n#         for n in range(1, 41):\n#             path = f\"native_{target_id}_{n}.pdb\"\n#             resolved = write2pdb_fast(group_native, n, path)\n#             if resolved > 5:   # ← 重要：最低原子数制限\n#                 native_paths.append(path)\n    \n#         pred_paths = []\n#         for p in range(1, 6):\n#             path = f\"pred_{target_id}_{p}.pdb\"\n#             write2pdb_fast(group_pred, p, path)\n#             pred_paths.append(path)\n    \n#         tasks = [(pred, native) for pred in pred_paths for native in native_paths]\n\n#         # print(f\"Target: {target_id}, tasks={len(tasks)}\")\n    \n#         if len(tasks) == 0:\n#             results[target_id] = 0.0\n#             continue\n    \n#         with Pool(n_workers, maxtasksperchild=20) as pool:\n#             scores = pool.map(run_usalign, tasks)\n    \n#         results[target_id] = max(scores)\n\n#     return results\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T14:21:38.725693Z","iopub.execute_input":"2026-02-19T14:21:38.726423Z","iopub.status.idle":"2026-02-19T14:21:38.740046Z","shell.execute_reply.started":"2026-02-19T14:21:38.726373Z","shell.execute_reply":"2026-02-19T14:21:38.739089Z"},"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport re\nimport subprocess\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nfrom concurrent.futures import ThreadPoolExecutor, as_completed\nimport multiprocessing\n\n\n# -----------------------------\n# TM-score parser\n# -----------------------------\ndef parse_tmscore_output(output):\n    m = re.findall(r\"TM-score=\\s+([\\d.]+)\", output)\n    if len(m) < 2:\n        return 0.0\n    return float(m[1])\n\n\n# -----------------------------\n# Fast PDB writer\n# -----------------------------\ndef write2pdb_fast(df, xyz_id, target_path):\n    x = df[f\"x_{xyz_id}\"].to_numpy()\n    y = df[f\"y_{xyz_id}\"].to_numpy()\n    z = df[f\"z_{xyz_id}\"].to_numpy()\n    resid = df[\"resid\"].to_numpy()\n    resname = df[\"resname\"].to_numpy()\n\n    mask = (x > -1e17) & (y > -1e17) & (z > -1e17)\n    if not np.any(mask):\n        return 0\n\n    x, y, z = x[mask], y[mask], z[mask]\n    resid, resname = resid[mask], resname[mask]\n\n    with open(target_path, \"w\") as f:\n        for i in range(len(x)):\n            line = (\n                f\"ATOM  {i+1:5d}  C1' {resname[i]:>3s} A\"\n                f\"{int(resid[i]):4d}    \"\n                f\"{x[i]:8.3f}{y[i]:8.3f}{z[i]:8.3f}\"\n                f\"{1.00:6.2f}{0.00:6.2f}           C\\n\"\n            )\n            f.write(line)\n\n    return len(x)\n\n\n# -----------------------------\n# USalign runner (fast mode)\n# -----------------------------\ndef run_usalign(pred_pdb, native_pdb):\n    cmd = [\n        \"/kaggle/working/USalign\",\n        pred_pdb,\n        native_pdb,\n        \"-atom\",\n        \" C1'\",\n        \"-fast\",\n        \"-TMscore\",\n        \"5\",\n    ]\n\n    try:\n        result = subprocess.run(\n            cmd,\n            capture_output=True,\n            text=True,\n            timeout=30,   # ← 長鎖用\n        )\n        return parse_tmscore_output(result.stdout)\n    except subprocess.TimeoutExpired:\n        return 0.0\n\n\n# -----------------------------\n# Main score function\n# -----------------------------\ndef score(solution: pd.DataFrame, submission: pd.DataFrame):\n\n    os.system(\"cp /kaggle/input/datasets/metric/usalign/USalign /kaggle/working/\")\n    os.system(\"chmod u+x /kaggle/working/USalign\")\n\n    solution[\"target_id\"] = solution[\"ID\"].str.split(\"_\").str[0]\n    submission[\"target_id\"] = submission[\"ID\"].str.split(\"_\").str[0]\n\n    results = {}\n    max_workers = min(8, multiprocessing.cpu_count())\n\n    for target_id, group_native in tqdm(solution.groupby(\"target_id\")):\n\n        group_pred = submission[submission[\"target_id\"] == target_id]\n\n        # ------------------\n        # write PDBs once\n        # ------------------\n        native_paths = []\n        for n in range(1, 41):\n            path = f\"native_{target_id}_{n}.pdb\"\n            resolved = write2pdb_fast(group_native, n, path)\n            if resolved > 10:\n                native_paths.append(path)\n\n        pred_paths = []\n        for p in range(1, 6):\n            path = f\"pred_{target_id}_{p}.pdb\"\n            write2pdb_fast(group_pred, p, path)\n            pred_paths.append(path)\n\n        if len(native_paths) == 0:\n            results[target_id] = 0.0\n            continue\n\n        tasks = [(pred, native) for pred in pred_paths for native in native_paths]\n\n        scores = []\n\n        # ------------------\n        # ThreadPool (重要)\n        # ------------------\n        with ThreadPoolExecutor(max_workers=max_workers) as executor:\n            futures = [\n                executor.submit(run_usalign, pred, native)\n                for pred, native in tasks\n            ]\n\n            for f in as_completed(futures):\n                scores.append(f.result())\n        # print(scores)\n        results[target_id] = max(scores) if scores else 0.0\n\n    return results\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T01:14:28.75728Z","iopub.execute_input":"2026-02-23T01:14:28.757687Z","iopub.status.idle":"2026-02-23T01:14:28.78991Z","shell.execute_reply.started":"2026-02-23T01:14:28.757601Z","shell.execute_reply":"2026-02-23T01:14:28.788294Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# pred_tbm_ = pred_tbm.copy()\n# pred_rnapro_ = pred_rnapro.copy()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T14:21:38.757091Z","iopub.execute_input":"2026-02-19T14:21:38.757371Z","iopub.status.idle":"2026-02-19T14:21:38.77042Z","shell.execute_reply.started":"2026-02-19T14:21:38.757325Z","shell.execute_reply":"2026-02-19T14:21:38.769733Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# pred_tbm_[[\"x_1\", \"y_1\", \"z_1\", \"x_3\", \"y_3\", \"z_3\"]] = 0\n# pred_tbm_ = pred_tbm_.drop(\"target_id\", axis=1)\n# pred_tbm_","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T14:21:38.771282Z","iopub.execute_input":"2026-02-19T14:21:38.771665Z","iopub.status.idle":"2026-02-19T14:21:38.780848Z","shell.execute_reply.started":"2026-02-19T14:21:38.77163Z","shell.execute_reply":"2026-02-19T14:21:38.780172Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# pred_rnapro_","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T14:21:38.781734Z","iopub.execute_input":"2026-02-19T14:21:38.782022Z","iopub.status.idle":"2026-02-19T14:21:38.790124Z","shell.execute_reply.started":"2026-02-19T14:21:38.781989Z","shell.execute_reply":"2026-02-19T14:21:38.789583Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# pred_rnapro_[\"target_id\"] = pred_rnapro_[\"ID\"].str.split(\"_\").str[0]\n# pred_rnapro_[[\"x_1\", \"y_1\", \"z_1\", \"x_2\", \"y_2\", \"z_2\"]] = 0\n# pred_rnapro_ = pred_rnapro_.drop(\"target_id\", axis=1)\n# pred_rnapro_\n# pred_rnapro_[[\"x_3\", \"y_3\", \"z_3\", \"x_4\", \"y_4\", \"z_4\", \"x_5\", \"y_5\", \"z_5\"]]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T14:21:38.790853Z","iopub.execute_input":"2026-02-19T14:21:38.791093Z","iopub.status.idle":"2026-02-19T14:21:38.800371Z","shell.execute_reply.started":"2026-02-19T14:21:38.791068Z","shell.execute_reply":"2026-02-19T14:21:38.799667Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# pred_tbm_.iloc[:, 9:] = 0\n# pred_tbm_","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T14:21:38.801265Z","iopub.execute_input":"2026-02-19T14:21:38.80163Z","iopub.status.idle":"2026-02-19T14:21:38.813423Z","shell.execute_reply.started":"2026-02-19T14:21:38.801596Z","shell.execute_reply":"2026-02-19T14:21:38.812695Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\nimport re\n\nscores = score(solution, sub) # sub pred_tbm_ submission pred_rnapro_ pred_tbm\nresult = {'target_id': [],\n          'length': [],\n          'tm_score': []}","metadata":{"trusted":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2026-02-23T03:15:33.446372Z","iopub.execute_input":"2026-02-23T03:15:33.447478Z","iopub.status.idle":"2026-02-23T03:16:38.088322Z","shell.execute_reply.started":"2026-02-23T03:15:33.447442Z","shell.execute_reply":"2026-02-23T03:16:38.086466Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for idx, rna in test_seqs.iterrows(): # test_seqs sequences test_df test_df.iloc[:2] df test_sequences valid_df\n    \n    target_id = rna.target_id\n    length = len(rna.sequence)\n    tm_score = scores[target_id]\n    \n    result['target_id'].append(target_id)\n    result['length'].append(length)\n    result['tm_score'].append(tm_score)\n\nresult_df = pd.DataFrame(result)\nresult_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T03:16:38.089999Z","iopub.execute_input":"2026-02-23T03:16:38.09037Z","iopub.status.idle":"2026-02-23T03:16:38.108365Z","shell.execute_reply.started":"2026-02-23T03:16:38.09034Z","shell.execute_reply":"2026-02-23T03:16:38.106994Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\n\ndef plot_results(df, desc, axs=None):\n    if axs is None:\n        fig, axs = plt.subplots(nrows=1, ncols=2, figsize=(15, 5), gridspec_kw={'width_ratios': [1, 2]})\n    else:\n        fig = None\n\n    df['color'] = df['length'].apply(lambda x: colors[0] if x > 400 else colors[-1])\n\n    sns.scatterplot(data=df, x='length', y='tm_score', hue='color', palette={colors[0]: colors[0], colors[-1]: colors[-1]}, ax=axs[0], legend=False)\n    axs[0].set_title('TM-score per length')\n    axs[0].set_xlabel('RNA length')\n    axs[0].set_ylabel('TM-score')\n    axs[0].set_ylim(0., 1.)\n\n    mean_tm = df['tm_score'].mean()\n    axs[0].axhline(mean_tm, color=colors[1], linestyle='--', label=f'Mean = {mean_tm:.2f}')\n    axs[0].legend()\n\n    sns.barplot(data=df, x='target_id', y='tm_score', palette=df['color'].to_list(), ax=axs[1])\n    axs[1].set_title('TM-score per target')\n    axs[1].set_xlabel('Target id')\n    axs[1].set_ylabel('TM-score')\n    axs[1].set_ylim(0., 1.)\n    axs[1].tick_params(axis='x', labelrotation=45)\n    axs[1].axhline(mean_tm, color=colors[1], linestyle='--', label=f'Mean = {mean_tm:.2f}')\n    axs[1].legend()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T01:24:30.636857Z","iopub.execute_input":"2026-02-23T01:24:30.637253Z","iopub.status.idle":"2026-02-23T01:24:30.661326Z","shell.execute_reply.started":"2026-02-23T01:24:30.637213Z","shell.execute_reply":"2026-02-23T01:24:30.659963Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"colors = ['#c1121f', '#ffb703', '#003049']\nprint('\\n----- Color -----\\n')\nsns.palplot(sns.color_palette(colors))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T01:24:30.664026Z","iopub.execute_input":"2026-02-23T01:24:30.665542Z","iopub.status.idle":"2026-02-23T01:24:30.729911Z","shell.execute_reply.started":"2026-02-23T01:24:30.665506Z","shell.execute_reply":"2026-02-23T01:24:30.728763Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_results(result_df, desc='Results')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T01:24:30.731123Z","iopub.execute_input":"2026-02-23T01:24:30.731431Z","iopub.status.idle":"2026-02-23T01:24:31.444471Z","shell.execute_reply.started":"2026-02-23T01:24:30.731402Z","shell.execute_reply":"2026-02-23T01:24:31.443662Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def log_resluts(result_df):\n    long_score = result_df.loc[result_df.length > 400, 'tm_score'].mean()\n    short_score = result_df.loc[result_df.length <= 400, 'tm_score'].mean()\n    all_score = result_df.tm_score.mean()\n\n    print(f'[TM-score] > 400: {long_score:.4f}')\n    print(f'[TM-score] <= 400: {short_score:.4f}')\n    print(f'[TM-score] all: {all_score:.4f}')\n\nlog_resluts(result_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-23T01:24:31.446165Z","iopub.execute_input":"2026-02-23T01:24:31.446543Z","iopub.status.idle":"2026-02-23T01:24:31.456159Z","shell.execute_reply.started":"2026-02-23T01:24:31.446517Z","shell.execute_reply":"2026-02-23T01:24:31.454568Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# [TM-score] > 400: 0.0611\n# [TM-score] <= 400: 0.4631\n# [TM-score] all: 0.4200","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T14:22:38.301178Z","iopub.execute_input":"2026-02-19T14:22:38.301498Z","iopub.status.idle":"2026-02-19T14:22:38.312368Z","shell.execute_reply.started":"2026-02-19T14:22:38.301464Z","shell.execute_reply":"2026-02-19T14:22:38.311665Z"}},"outputs":[],"execution_count":null}]}