{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.12"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":15231210},{"sourceType":"datasetVersion","sourceId":14604295,"datasetId":9328538,"databundleVersionId":15440074},{"sourceType":"datasetVersion","sourceId":14962460,"datasetId":9577079,"databundleVersionId":15833819},{"sourceType":"datasetVersion","sourceId":14834274,"datasetId":9487358,"databundleVersionId":15692465},{"sourceType":"datasetVersion","sourceId":14874339,"datasetId":9502242,"databundleVersionId":15736806},{"sourceType":"datasetVersion","sourceId":14962495,"datasetId":9577097,"databundleVersionId":15833858},{"sourceType":"datasetVersion","sourceId":10855324,"datasetId":6742586,"databundleVersionId":11219268},{"sourceType":"datasetVersion","sourceId":15114187,"datasetId":9677457,"databundleVersionId":16001204}],"dockerImageVersionId":31260,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":4701.131185,"end_time":"2026-03-09T11:40:27.376497","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2026-03-09T10:22:06.245312","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"| version | description                                            | LB    |\n|---------|-------------------------------------------------------|-------|\n| 23      | v4 + pLDDT sort                                        | 0.436 |\n| 32      | v23 + GUDHI H1 PH Rerank + Noise                     | 0.447 |\n| 33      | v32 No noise (counterproductive)                                | 0.418 |\n| 34      | v23 + Noise only (counterproductive)                              | 0.426 |\n| 35      | v32 + Official gudhi installation                            | TBC   |\n| 36      | v35 + Auto-draw H1 persistent diagram                     | TBC   |\n","metadata":{"execution":{"iopub.execute_input":"2026-02-28T12:18:56.684691Z","iopub.status.busy":"2026-02-28T12:18:56.68442Z","iopub.status.idle":"2026-02-28T12:19:01.878617Z","shell.execute_reply":"2026-02-28T12:19:01.877688Z"},"papermill":{"duration":0.003311,"end_time":"2026-03-09T10:22:08.989737","exception":false,"start_time":"2026-03-09T10:22:08.986426","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# GUDHI H1 Persistent Homology Reranking\n\n## Overview\n\nApplies Persistent Homology (PH) to the C1' coordinate arrays of Protenix samples\nand reranks them so that the topologically most representative structure is placed first.\nTBM sample ordering is never modified.\n\n---\n\n## Pipeline\n```\nPhase 1 : TBM          → up to N_SAMPLE predictions (by similarity)\nPhase 2 : Protenix     → fills remaining slots (sorted by pLDDT descending)\n          ↓\n          PH rerank    → applied only within the Protenix block\nPhase 3 : de-novo      → last-resort fallback\n```\n\n---\n\n## Formulation\n\n### Step 1 — H1 Feature Extraction\n\nFor candidate $C_i \\in \\mathbb{R}^{L \\times 3}$, compute the pairwise distance matrix $D_{jk} = \\|C_i^{(j)} - C_i^{(k)}\\|_2$, build a Rips filtration, and extract H1 birth–death pairs:\n\n$$\\text{pairs} = \\{(b_k,\\, d_k) \\mid d_k - b_k \\geq \\epsilon_{\\min}\\}$$\n\nAfter scale normalization, vectorize into a feature vector:\n\n$$\\phi(C_i) = \\bigl[\\underbrace{p_1,\\ldots,p_{16}}_{\\text{top-16 persistence}},\\ \\underbrace{\\text{8 statistics}}_{\\text{mean, std, etc.}},\\ \\underbrace{h_b}_{\\text{birth histogram}},\\ \\underbrace{h_p}_{\\text{persistence histogram}}\\bigr] \\in \\mathbb{R}^{64}$$\n\n### Step 2 — Reference Feature (Median Ensemble)\n\n$$\\phi_{\\text{ref}} = \\operatorname{median}\\bigl(\\phi(C_0),\\ldots,\\phi(C_{n-1})\\bigr)$$\n\n### Step 3 — Scoring\n\n$$s_i^{\\text{PH}} = -\\|\\phi(C_i) - \\phi_{\\text{ref}}\\|_2 - \\beta\\, B(C_i) - \\gamma\\, K(C_i)$$\n\n$$B(C_i) = \\frac{1}{L-1}\\sum_{j=1}^{L-1}\\bigl(\\|C_i^{(j+1)}-C_i^{(j)}\\|_2 - 6.0\\bigr)^2 \\quad\\text{(bond length penalty,}\\ \\beta=0.25\\text{)}$$\n\n$$K(C_i) = \\frac{1}{|M|}\\sum_{(j,k)\\in M}(2.2 - d_{jk})^2 \\quad M=\\{d_{jk}<2.2,\\ |j-k|>2\\} \\quad\\text{(clash penalty,}\\ \\gamma=0.20\\text{)}$$\n\n### Step 4 — Final Score and Reranking\n\nNormalize the PH scores and blend with the original rank order:\n\n$$\\tilde{s}_i^{\\text{PH}} = \\frac{s_i^{\\text{PH}} - \\mu}{\\sigma}$$\n\n$$\\text{score}_i = \\underbrace{-0.035 \\cdot i}_{\\text{original rank penalty}} + \\underbrace{0.010 \\cdot \\tilde{s}_i^{\\text{PH}}}_{\\text{PH signal}}$$\n\nCandidates are sorted by $\\text{score}_i$ in descending order.\n\n---\n\n## Intuition\n\n> **\"Place the candidate whose topology is closest to the median of all candidates at rank 1.\"**\n\nThe PH signal weight $0.010$ is roughly 28% of the rank penalty $0.035$,\nso reranking is a mild correction that rarely causes large rank swaps.\n\n---\n\n## Parameters\n\n| Parameter | Value | Description |\n|---|---|---|\n| `PH_BLEND_LAMBDA` | 0.010 | Blend weight for PH signal |\n| `PH_BLEND_BASE_GAP` | 0.035 | Per-rank penalty coefficient |\n| `PH_ALPHA` | 1.0 | PH score weight |\n| `PH_BETA` | 0.25 | Bond length penalty weight |\n| `PH_GAMMA` | 0.20 | Clash penalty weight |\n| `PH_MAX_POINTS` | 96 | Subsampling cap |\n| `PH_H1_TOPK` | 16 | Number of top persistence values |\n| `PH_BETTI_BINS` | 24 | Histogram bin count |\n| `PH_MAX_EDGE_MULT` | 3.5 | Max edge length multiplier for Rips |\n| `PH_MIN_PERSISTENCE` | 1e-3 | Minimum persistence threshold |\n| `EXPECTED_C1_STEP` | 6.0 Å | Ideal C1'–C1' distance |\n| `CLASH_DIST` | 2.2 Å | Clash detection distance |","metadata":{}},{"cell_type":"code","source":"# =============\n# 1) Install\n# =============\n# Assumes Kaggle offline wheels (adjust the path to match your wheel directory)\n!pip install --no-index --find-links=/kaggle/input/datasets/takahiro3110/gungun-cp312/gudhi_wheels_cp312/wheels gudhi","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T10:49:28.08608Z","iopub.execute_input":"2026-03-15T10:49:28.086332Z","iopub.status.idle":"2026-03-15T10:49:32.937232Z","shell.execute_reply.started":"2026-03-15T10:49:28.086311Z","shell.execute_reply":"2026-03-15T10:49:32.936346Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Reference\n\nhttps://www.kaggle.com/code/llkh0a/stanford-rna-3d-folding-part-2-protenix-tbm","metadata":{}},{"cell_type":"code","source":"!pip install --no-index --no-deps /kaggle/input/datasets/kami1976/biopython-cp312/biopython-1.86-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl\n\n!pip install --no-index --no-deps /kaggle/input/datasets/amirrezaaleyasin/biotite/biotite-1.6.0-cp312-cp312-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl\n\n!pip install --no-index --no-deps /kaggle/input/datasets/amirrezaaleyasin/rdkit-2025-9-5/rdkit-2025.9.5-cp312-cp312-manylinux_2_28_x86_64.whl\n\n!pip install --no-index --find-links=/kaggle/input/datasets/takahiro3110/gungun-cp312/gudhi_wheels_cp312/wheels gudhi\n","metadata":{"execution":{"iopub.status.busy":"2026-03-15T10:49:32.939231Z","iopub.execute_input":"2026-03-15T10:49:32.93955Z","iopub.status.idle":"2026-03-15T10:49:46.653042Z","shell.execute_reply.started":"2026-03-15T10:49:32.93952Z","shell.execute_reply":"2026-03-15T10:49:46.652166Z"},"papermill":{"duration":10.203616,"end_time":"2026-03-09T10:22:19.196535","exception":false,"start_time":"2026-03-09T10:22:08.992919","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport sys\nimport pandas as pd\n\n# ── Local vs Kaggle mode ─────────────────────────────────────────────────────\n# On Kaggle competition rerun, KAGGLE_IS_COMPETITION_RERUN is set to a truthy value.\n# When running locally we do NOT exit — instead we cap the test set to a small\n# number of samples so the notebook finishes quickly.\n\nIS_KAGGLE = True #bool(os.environ.get(\"KAGGLE_IS_COMPETITION_RERUN\", \"\"))\n\n# How many test samples to use when running locally\nLOCAL_N_SAMPLES = None\n\nif IS_KAGGLE:\n    print(\"Running in KAGGLE COMPETITION mode — all test targets will be processed.\")\nelse:\n    print(f\"Running in LOCAL mode — only the first {LOCAL_N_SAMPLES} test targets \"\n          f\"will be processed to save time.\")","metadata":{"_cell_guid":"fccfef83-959a-48a1-a35d-47f796ad39a2","_uuid":"68a01f46-f15d-42e4-9432-34d7f5bd253c","collapsed":false,"execution":{"iopub.status.busy":"2026-03-15T10:49:46.654304Z","iopub.execute_input":"2026-03-15T10:49:46.654569Z","iopub.status.idle":"2026-03-15T10:49:46.95709Z","shell.execute_reply.started":"2026-03-15T10:49:46.654532Z","shell.execute_reply":"2026-03-15T10:49:46.956464Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.926443,"end_time":"2026-03-09T10:22:20.126975","exception":false,"start_time":"2026-03-09T10:22:19.200532","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import gc\nimport json\nimport os\nimport time\n\nos.environ[\"LAYERNORM_TYPE\"] = \"torch\"\nos.environ.setdefault(\"RNA_MSA_DEPTH_LIMIT\", \"512\")\n\nimport sys\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom Bio.Align import PairwiseAligner\nfrom tqdm import tqdm","metadata":{"_cell_guid":"9bcebd97-b977-437d-a0bd-65786f1ca5b7","_uuid":"153fdfa6-3982-4e92-ba15-550f5fd68472","collapsed":false,"execution":{"iopub.status.busy":"2026-03-15T10:49:46.958637Z","iopub.execute_input":"2026-03-15T10:49:46.959071Z","iopub.status.idle":"2026-03-15T10:49:50.480691Z","shell.execute_reply.started":"2026-03-15T10:49:46.959047Z","shell.execute_reply":"2026-03-15T10:49:50.48013Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":4.029671,"end_time":"2026-03-09T10:22:24.160342","exception":false,"start_time":"2026-03-09T10:22:20.130671","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_c1_mask(data: dict, atom_array) -> torch.Tensor:\n    # 1. Try atom_array attributes first\n    if atom_array is not None:\n        try:\n            if hasattr(atom_array, \"centre_atom_mask\"):\n                m = atom_array.centre_atom_mask == 1\n                if hasattr(atom_array, \"is_rna\"):\n                    m = m & atom_array.is_rna\n                return torch.from_numpy(m).bool()\n            \n            if hasattr(atom_array, \"atom_name\"):\n                base = atom_array.atom_name == \"C1'\"\n                if hasattr(atom_array, \"is_rna\"):\n                    base = base & atom_array.is_rna\n                return torch.from_numpy(base).bool()\n        except Exception:\n            pass\n\n    # 2. Fallback to feature dict\n    f = data[\"input_feature_dict\"]\n    \n    if \"centre_atom_mask\" in f:\n        return (f[\"centre_atom_mask\"] == 1).bool()\n    if \"center_atom_mask\" in f:\n        return (f[\"center_atom_mask\"] == 1).bool()\n        \n    # Heuristic fallback: check which index gives us roughly N_token atoms\n    n_tokens = data.get(\"N_token\", torch.tensor(0)).item()\n    mask11 = (f[\"atom_to_tokatom_idx\"] == 11).bool()\n    mask12 = (f[\"atom_to_tokatom_idx\"] == 12).bool()\n    \n    c11 = mask11.sum().item()\n    c12 = mask12.sum().item()\n    \n    # Return the one closer to N_tokens (likely one per residue)\n    if abs(c11 - n_tokens) < abs(c12 - n_tokens):\n        return mask11\n    else:\n        return mask12","metadata":{"_cell_guid":"da0eeb13-1bb3-4dca-8af4-dbe018709908","_uuid":"f06a4b36-7bce-40ff-be22-fbe776b64e7d","collapsed":false,"execution":{"iopub.status.busy":"2026-03-15T10:49:50.481771Z","iopub.execute_input":"2026-03-15T10:49:50.482249Z","iopub.status.idle":"2026-03-15T10:49:50.489222Z","shell.execute_reply.started":"2026-03-15T10:49:50.482208Z","shell.execute_reply":"2026-03-15T10:49:50.488538Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.013173,"end_time":"2026-03-09T10:22:24.177416","exception":false,"start_time":"2026-03-09T10:22:24.164243","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ─────────────── Paths & Constants ───────────────────────────────────────────\nDATA_BASE              = \"/kaggle/input/stanford-rna-3d-folding-2\"\nDEFAULT_TEST_CSV       = f\"{DATA_BASE}/test_sequences.csv\"\nDEFAULT_TRAIN_CSV      = f\"{DATA_BASE}/train_sequences.csv\"\nDEFAULT_TRAIN_LBLS     = f\"{DATA_BASE}/train_labels.csv\"\nDEFAULT_VAL_CSV        = f\"{DATA_BASE}/validation_sequences.csv\"\nDEFAULT_VAL_LBLS       = f\"{DATA_BASE}/validation_labels.csv\"\nDEFAULT_OUTPUT         = \"/kaggle/working/submission.csv\"\n\nDEFAULT_CODE_DIR = (\n    \"/kaggle/input/datasets/qiweiyin/protenix-v1-adjusted\"\n    \"/Protenix-v1-adjust-v2/Protenix-v1-adjust-v2/Protenix-v1\"\n)\nDEFAULT_ROOT_DIR = DEFAULT_CODE_DIR\n\nMODEL_NAME    = \"protenix_base_20250630_v1.0.0\"\nN_SAMPLE      = 5\nSEED          = 42\nMAX_SEQ_LEN   = int(os.environ.get(\"MAX_SEQ_LEN\",   \"512\"))\nCHUNK_OVERLAP = int(os.environ.get(\"CHUNK_OVERLAP\",  \"128\"))\n\n# TBM quality thresholds — sequences below these get routed to Protenix\nMIN_SIMILARITY       = float(os.environ.get(\"MIN_SIMILARITY\",       \"0.0\"))\nMIN_PERCENT_IDENTITY = float(os.environ.get(\"MIN_PERCENT_IDENTITY\", \"50.0\"))\n\n# Set False to skip Protenix and use de-novo fallback instead\nUSE_PROTENIX = True\n\n\ndef parse_bool(value: str, default: bool = False) -> str:\n    v = str(value).strip().lower()\n    if v in {\"1\", \"true\", \"t\", \"yes\", \"y\", \"on\"}:\n        return \"true\"\n    if v in {\"0\", \"false\", \"f\", \"no\", \"n\", \"off\"}:\n        return \"false\"\n    return \"true\" if default else \"false\"\n\n\nUSE_MSA      = parse_bool(os.environ.get(\"USE_MSA\",      \"false\"))\nUSE_TEMPLATE = parse_bool(os.environ.get(\"USE_TEMPLATE\", \"false\"))\nUSE_RNA_MSA  = parse_bool(os.environ.get(\"USE_RNA_MSA\",  \"true\"))\n\nMODEL_N_SAMPLE = int(os.environ.get(\"MODEL_N_SAMPLE\", str(N_SAMPLE)))\n\n\n# ─────────────── General Utilities ───────────────────────────────────────────\ndef seed_everything(seed: int) -> None:\n    os.environ[\"CUBLAS_WORKSPACE_CONFIG\"] = \":4096:8\"\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    np.random.seed(seed)\n    torch.backends.cudnn.benchmark = False\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.enabled = True\n    torch.use_deterministic_algorithms(True)\n\n\ndef resolve_paths():\n    test_csv   = os.environ.get(\"TEST_CSV\",           DEFAULT_TEST_CSV)\n    output_csv = os.environ.get(\"SUBMISSION_CSV\",     DEFAULT_OUTPUT)\n    code_dir   = os.environ.get(\"PROTENIX_CODE_DIR\",  DEFAULT_CODE_DIR)\n    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    for p, name in [\n        (Path(root_dir) / \"checkpoint\" / f\"{MODEL_NAME}.pt\",          \"checkpoint\"),\n        (Path(root_dir) / \"common\" / \"components.cif\",                \"CCD file\"),\n        (Path(root_dir) / \"common\" / \"components.cif.rdkit_mol.pkl\",  \"CCD cache\"),\n    ]:\n        if not p.exists():\n            raise FileNotFoundError(f\"Missing {name}: {p}\")\n\n\n# ─────────────── Protenix Input / Config Helpers ─────────────────────────────\ndef build_input_json(df: pd.DataFrame, json_path: str) -> None:\n    data = [\n        {\n            \"name\": row[\"target_id\"],\n            \"covalent_bonds\": [],\n            \"sequences\": [{\"rnaSequence\": {\"sequence\": row[\"sequence\"], \"count\": 1}}],\n        }\n        for _, row in df.iterrows()\n    ]\n    with open(json_path, \"w\", 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_base, **{\"data\": data_configs}, **inference_configs}\n\n    def deep_update(t, p):\n        for k, v in p.items():\n            if isinstance(v, dict) and k in t and isinstance(t[k], dict):\n                deep_update(t[k], v)\n            else:\n                t[k] = v\n\n    deep_update(base, model_configs[model_name])\n    arg_str = \" \".join([\n        f\"--model_name {model_name}\",\n        f\"--input_json_path {input_json_path}\",\n        f\"--dump_dir {dump_dir}\",\n        f\"--use_msa {USE_MSA}\",\n        f\"--use_template {USE_TEMPLATE}\",\n        f\"--use_rna_msa {USE_RNA_MSA}\",\n        f\"--sample_diffusion.N_sample {MODEL_N_SAMPLE}\",\n        f\"--seeds {SEED}\",\n    ])\n    return parse_configs(configs=base, arg_str=arg_str, fill_required_with_null=True)\n\n\ndef get_c1_mask(data: dict, atom_array) -> torch.Tensor:\n    # 1. Try atom_array attributes first\n    if atom_array is not None:\n        try:\n            if hasattr(atom_array, \"centre_atom_mask\"):\n                m = atom_array.centre_atom_mask == 1\n                if hasattr(atom_array, \"is_rna\"):\n                    m = m & atom_array.is_rna\n                return torch.from_numpy(m).bool()\n            \n            if hasattr(atom_array, \"atom_name\"):\n                base = atom_array.atom_name == \"C1'\"\n                if hasattr(atom_array, \"is_rna\"):\n                    base = base & atom_array.is_rna\n                return torch.from_numpy(base).bool()\n        except Exception:\n            pass\n\n    # 2. Fallback to feature dict\n    f = data[\"input_feature_dict\"]\n    \n    # CASE A: center_atom_mask exists\n    if \"center_atom_mask\" in f:\n        return (f[\"center_atom_mask\"] == 1).bool()\n    if \"centre_atom_mask\" in f:\n        return (f[\"centre_atom_mask\"] == 1).bool()\n        \n    # CASE B: Use atom_name\n    if \"atom_name\" in f:\n        # Check against \"C1'\" (byte encoded or string?)\n        # For now assume typical behavior is center_atom_mask is present.\n        pass\n\n    # CASE C: atom_to_tokatom_idx fallback\n    # The index for C1' is typically 11 or 12 depending on featurizer.\n    # Let's try to match exactly C1' if possible.\n    # But usually 'centre_atom_mask' should be there.\n    \n    # If we fall through, assume standard mask\n    return (f[\"atom_to_tokatom_idx\"] == 11).bool()\n\n\ndef get_feature_c1_mask(data: dict) -> torch.Tensor:\n    f = data[\"input_feature_dict\"]\n    if \"centre_atom_mask\" in f:\n        return f[\"centre_atom_mask\"].long() == 1\n    return f[\"atom_to_tokatom_idx\"].long() == 12\n\n\ndef coords_to_rows(target_id: str, seq: str, coords: np.ndarray) -> list:\n    \"\"\"coords shape: (N_SAMPLE, seq_len, 3)\"\"\"\n    rows = []\n    for i in range(len(seq)):\n        row = {\"ID\": f\"{target_id}_{i + 1}\", \"resname\": seq[i], \"resid\": i + 1}\n        for s in range(N_SAMPLE):\n            if s < coords.shape[0] and i < coords.shape[1]:\n                x, y, z = coords[s, i]\n            else:\n                x, y, z = 0.0, 0.0, 0.0\n            row[f\"x_{s + 1}\"] = float(x)\n            row[f\"y_{s + 1}\"] = float(y)\n            row[f\"z_{s + 1}\"] = float(z)\n        rows.append(row)\n    return rows\n\n\ndef pad_samples(coords: np.ndarray, n: int) -> np.ndarray:\n    if coords.shape[0] >= n:\n        return coords[:n]\n    if coords.shape[0] == 0:\n        return np.zeros((n, coords.shape[1], 3), dtype=coords.dtype)\n    extra = np.repeat(coords[:1], n - coords.shape[0], axis=0)\n    return np.concatenate([coords, extra], axis=0)\n\n\ndef split_into_chunks(seq_len: int, max_len: int, overlap: int) -> list:\n    \"\"\"Split a sequence into overlapping (start, end) chunks.\"\"\"\n    if seq_len <= max_len:\n        return [(0, seq_len)]\n    chunks = []\n    step = max_len - overlap\n    pos = 0\n    while pos < seq_len:\n        end = min(pos + max_len, seq_len)\n        chunks.append((pos, end))\n        if end == seq_len:\n            break\n        pos += step\n    return chunks\n\n\ndef kabsch_align(P: np.ndarray, Q: np.ndarray):\n    \"\"\"Compute optimal rotation R and translation t so that  R @ P + t ≈ Q.\"\"\"\n    centroid_P = P.mean(axis=0)\n    centroid_Q = Q.mean(axis=0)\n    Pc = P - centroid_P\n    Qc = Q - centroid_Q\n    H = Pc.T @ Qc\n    U, _, Vt = np.linalg.svd(H)\n    d = np.linalg.det(Vt.T @ U.T)\n    S = np.eye(3)\n    if d < 0:\n        S[2, 2] = -1\n    R = Vt.T @ S @ U.T\n    t = centroid_Q - R @ centroid_P\n    return R, t\n\n\ndef stitch_chunk_coords(chunk_coords_list: list,\n                        chunk_ranges: list,\n                        seq_len: int) -> np.ndarray:\n    \"\"\"\n    Merge overlapping chunk coordinates into a full sequence geometry.\n    Applies Kabsch alignment on overlapping residues, and smoothly\n    blends the coordinates using a linear weight ramp.\n    \"\"\"\n    if len(chunk_coords_list) == 1:\n        coords = chunk_coords_list[0]\n        if coords.shape[0] >= seq_len:\n            return coords[:seq_len]\n        out = np.zeros((seq_len, 3), dtype=coords.dtype)\n        out[:coords.shape[0]] = coords\n        return out\n\n    # Start with the first chunk aligned to itself (identity)\n    aligned = [chunk_coords_list[0].copy()]\n\n    for i in range(1, len(chunk_coords_list)):\n        prev_start, prev_end = chunk_ranges[i - 1]\n        cur_start, cur_end = chunk_ranges[i]\n\n        ov_start = cur_start\n        ov_end = min(prev_end, cur_end)\n        ov_len = ov_end - ov_start\n\n        if ov_len < 3:\n            # Cannot align reliably, just trust the coordinates as-is\n            aligned.append(chunk_coords_list[i].copy())\n            continue\n\n        prev_ov = aligned[i - 1][ov_start - prev_start: ov_end - prev_start]\n        cur_ov = chunk_coords_list[i][ov_start - cur_start: ov_end - cur_start]\n\n        # Ignore invalid residues (e.g. padding/blank)\n        valid = ~(np.isnan(prev_ov).any(axis=1) | np.isnan(cur_ov).any(axis=1))\n        if valid.sum() < 3:\n            aligned.append(chunk_coords_list[i].copy())\n            continue\n\n        # Align current chunk to previous chunk using only the overlap region\n        R, t = kabsch_align(cur_ov[valid], prev_ov[valid])\n        transformed = (chunk_coords_list[i] @ R.T) + t\n        aligned.append(transformed)\n\n    # Blend them together\n    full = np.zeros((seq_len, 3), dtype=np.float64)\n    weights = np.zeros(seq_len, dtype=np.float64)\n\n    for i, ((s, e), coords) in enumerate(zip(chunk_ranges, aligned)):\n        chunk_len = coords.shape[0]\n        actual_end = min(s + chunk_len, seq_len)\n        used_len = actual_end - s\n\n        w = np.ones(used_len, dtype=np.float64)\n\n        if i > 0:\n            ov_start = s\n            ov_end = min(chunk_ranges[i - 1][1], e)\n            ramp_len = ov_end - ov_start\n            if ramp_len > 0:\n                w[:ramp_len] = np.linspace(0.0, 1.0, ramp_len)\n\n        if i < len(chunk_ranges) - 1:\n            next_s = chunk_ranges[i + 1][0]\n            ramp_start = next_s - s\n            ramp_len = actual_end - next_s\n            if ramp_len > 0 and ramp_start < used_len:\n                w[ramp_start:used_len] = np.linspace(1.0, 0.0, ramp_len)\n\n        full[s:actual_end] += coords[:used_len] * w[:, None]\n        weights[s:actual_end] += w\n\n    mask = weights > 0\n    full[mask] /= weights[mask, None]\n\n    return full\n\n\n# ─────────────── TBM Core Functions ──────────────────────────────────────────\ndef _make_aligner() -> PairwiseAligner:\n    al = PairwiseAligner()\n    al.mode                           = \"global\"\n    al.match_score                    = 2\n    al.mismatch_score                 = -1.5\n    al.open_gap_score                 = -8\n    al.extend_gap_score               = -0.4\n    al.query_left_open_gap_score      = -8\n    al.query_left_extend_gap_score    = -0.4\n    al.query_right_open_gap_score     = -8\n    al.query_right_extend_gap_score   = -0.4\n    al.target_left_open_gap_score     = -8\n    al.target_left_extend_gap_score   = -0.4\n    al.target_right_open_gap_score    = -8\n    al.target_right_extend_gap_score  = -0.4\n    return al\n\n\n_aligner = _make_aligner()\n\n\ndef parse_stoichiometry(stoich: str) -> list:\n    if pd.isna(stoich) or str(stoich).strip() == \"\":\n        return []\n    return [(ch.strip(), int(cnt)) for part in str(stoich).split(\";\")\n            for ch, cnt in [part.split(\":\")]]\n\n\ndef parse_fasta(fasta_content: str) -> dict:\n    out, cur, parts = {}, None, []\n    for line in str(fasta_content).splitlines():\n        line = line.strip()\n        if not line:\n            continue\n        if line.startswith(\">\"):\n            if cur is not None:\n                out[cur] = \"\".join(parts)\n            cur = line[1:].split()[0]\n            parts = []\n        else:\n            parts.append(line.replace(\" \", \"\"))\n    if cur is not None:\n        out[cur] = \"\".join(parts)\n    return out\n\n\ndef get_chain_segments(row) -> list:\n    seq    = row[\"sequence\"]\n    stoich = row.get(\"stoichiometry\", \"\")\n    all_sq = row.get(\"all_sequences\", \"\")\n    if (pd.isna(stoich) or pd.isna(all_sq)\n            or str(stoich).strip() == \"\" or str(all_sq).strip() == \"\"):\n        return [(0, len(seq))]\n    try:\n        chain_dict = parse_fasta(all_sq)\n        order = parse_stoichiometry(stoich)\n        segs, pos = [], 0\n        for ch, cnt in order:\n            base = chain_dict.get(ch)\n            if base is None:\n                return [(0, len(seq))]\n            for _ in range(cnt):\n                segs.append((pos, pos + len(base)))\n                pos += len(base)\n        return segs if pos == len(seq) else [(0, len(seq))]\n    except Exception:\n        return [(0, len(seq))]\n\n\ndef build_segments_map(df: pd.DataFrame) -> tuple:\n    seg_map, stoich_map = {}, {}\n    for _, r in df.iterrows():\n        tid               = r[\"target_id\"]\n        seg_map[tid]      = get_chain_segments(r)\n        raw_s             = r.get(\"stoichiometry\", \"\")\n        stoich_map[tid]   = \"\" if pd.isna(raw_s) else str(raw_s)\n    return seg_map, stoich_map\n\n\ndef process_labels(labels_df: pd.DataFrame) -> dict:\n    coords = {}\n    prefixes = labels_df[\"ID\"].str.rsplit(\"_\", n=1).str[0]\n    for prefix, grp in labels_df.groupby(prefixes):\n        coords[prefix] = grp.sort_values(\"resid\")[[\"x_1\", \"y_1\", \"z_1\"]].values\n    return coords\n\n\ndef _build_aligned_strings(query_seq, template_seq, alignment):\n    q_segs, t_segs = alignment.aligned\n    aq, at, qi, ti = [], [], 0, 0\n    for (qs, qe), (ts, te) in zip(q_segs, t_segs):\n        while qi < qs: aq.append(query_seq[qi]);    at.append(\"-\");              qi += 1\n        while ti < ts: aq.append(\"-\");              at.append(template_seq[ti]); ti += 1\n        for qp, tp in zip(range(qs, qe), range(ts, te)):\n            aq.append(query_seq[qp]); at.append(template_seq[tp])\n        qi, ti = qe, te\n    while qi < len(query_seq):    aq.append(query_seq[qi]);    at.append(\"-\");              qi += 1\n    while ti < len(template_seq): aq.append(\"-\");              at.append(template_seq[ti]); ti += 1\n    return \"\".join(aq), \"\".join(at)\n\n\ndef find_similar_sequences_detailed(query_seq, train_seqs_df, train_coords_dict, top_n=30):\n    results = []\n    for _, row in train_seqs_df.iterrows():\n        tid, tseq = row[\"target_id\"], row[\"sequence\"]\n        if tid not in train_coords_dict:\n            continue\n        if abs(len(tseq) - len(query_seq)) / max(len(tseq), len(query_seq)) > 0.3:\n            continue\n        aln       = next(iter(_aligner.align(query_seq, tseq)))\n        norm_s    = aln.score / (2 * min(len(query_seq), len(tseq)))\n        identical = sum(\n            1 for (qs, qe), (ts, te) in zip(*aln.aligned)\n            for qp, tp in zip(range(qs, qe), range(ts, te))\n            if query_seq[qp] == tseq[tp]\n        )\n        pct_id = 100 * identical / len(query_seq)\n        aq, at = _build_aligned_strings(query_seq, tseq, aln)\n        results.append((tid, tseq, norm_s, train_coords_dict[tid], pct_id, aq, at))\n    results.sort(key=lambda x: x[2], reverse=True)\n    return results[:top_n]\n\n\ndef adapt_template_to_query(query_seq, template_seq, template_coords) -> np.ndarray:\n    aln        = next(iter(_aligner.align(query_seq, template_seq)))\n    new_coords = np.full((len(query_seq), 3), np.nan)\n    for (qs, qe), (ts, te) in zip(*aln.aligned):\n        chunk = template_coords[ts:te]\n        if len(chunk) == (qe - qs):\n            new_coords[qs:qe] = chunk\n    for i in range(len(new_coords)):\n        if np.isnan(new_coords[i, 0]):\n            pv = next((j for j in range(i - 1, -1, -1) if not np.isnan(new_coords[j, 0])), -1)\n            nv = next((j for j in range(i + 1, len(new_coords)) if not np.isnan(new_coords[j, 0])), -1)\n            if pv >= 0 and nv >= 0:\n                w = (i - pv) / (nv - pv)\n                new_coords[i] = (1 - w) * new_coords[pv] + w * new_coords[nv]\n            elif pv >= 0:\n                new_coords[i] = new_coords[pv] + [3, 0, 0]\n            elif nv >= 0:\n                new_coords[i] = new_coords[nv] + [3, 0, 0]\n            else:\n                new_coords[i] = [i * 3, 0, 0]\n    return np.nan_to_num(new_coords)\n\n\ndef adaptive_rna_constraints(coords, target_id, segments_map, confidence=1.0, passes=2) -> np.ndarray:\n    X        = coords.copy()\n    segments = segments_map.get(target_id, [(0, len(X))])\n    strength = max(0.75 * (1.0 - min(confidence, 0.97)), 0.02)\n    for _ in range(passes):\n        for s, e in segments:\n            C = X[s:e]; L = e - s\n            if L < 3:\n                continue\n            # bond i–i+1  ~5.95 Å\n            d    = C[1:] - C[:-1]; dist = np.linalg.norm(d, axis=1) + 1e-6\n            adj  = d * ((5.95 - dist) / dist)[:, None] * (0.22 * strength)\n            C[:-1] -= adj; C[1:] += adj\n            # soft i–i+2  ~10.2 Å\n            d2   = C[2:] - C[:-2]; d2n = np.linalg.norm(d2, axis=1) + 1e-6\n            adj2 = d2 * ((10.2 - d2n) / d2n)[:, None] * (0.10 * strength)\n            C[:-2] -= adj2; C[2:] += adj2\n            # Laplacian smoothing\n            C[1:-1] += (0.06 * strength) * (0.5 * (C[:-2] + C[2:]) - C[1:-1])\n            # self-avoidance\n            if L >= 25:\n                idx  = np.linspace(0, L - 1, min(L, 160)).astype(int) if L > 220 else np.arange(L)\n                P    = C[idx]; diff = P[:, None, :] - P[None, :, :]\n                dm   = np.linalg.norm(diff, axis=2) + 1e-6\n                sep  = np.abs(idx[:, None] - idx[None, :])\n                mask = (sep > 2) & (dm < 3.2)\n                if np.any(mask):\n                    vec = (diff * ((3.2 - dm) / dm)[:, :, None] * mask[:, :, None]).sum(axis=1)\n                    C[idx] += (0.015 * strength) * vec\n            X[s:e] = C\n    return X\n\n\ndef _rotmat(axis, ang):\n    a = np.asarray(axis, float); a /= np.linalg.norm(a) + 1e-12\n    x, y, z = a; c, s = np.cos(ang), np.sin(ang); CC = 1 - c\n    return np.array([[c+x*x*CC, x*y*CC-z*s, x*z*CC+y*s],\n                     [y*x*CC+z*s, c+y*y*CC, y*z*CC-x*s],\n                     [z*x*CC-y*s, z*y*CC+x*s, c+z*z*CC]])\n\n\ndef apply_hinge(coords, seg, rng, deg=22):\n    s, e = seg; L = e - s\n    if L < 30: return coords\n    pivot = s + int(rng.integers(10, L - 10))\n    R = _rotmat(rng.normal(size=3), np.deg2rad(float(rng.uniform(-deg, deg))))\n    X = coords.copy(); p0 = X[pivot].copy()\n    X[pivot+1:e] = (X[pivot+1:e] - p0) @ R.T + p0\n    return X\n\n\ndef jitter_chains(coords, segs, rng, deg=12, trans=1.5):\n    X = coords.copy(); gc_ = X.mean(0, keepdims=True)\n    for s, e in segs:\n        R     = _rotmat(rng.normal(size=3), np.deg2rad(float(rng.uniform(-deg, deg))))\n        shift = rng.normal(size=3); shift = shift / (np.linalg.norm(shift) + 1e-12) * float(rng.uniform(0, trans))\n        c     = X[s:e].mean(0, keepdims=True)\n        X[s:e] = (X[s:e] - c) @ R.T + c + shift\n    X -= X.mean(0, keepdims=True) - gc_\n    return X\n\n\ndef smooth_wiggle(coords, segs, rng, amp=0.8):\n    X = coords.copy()\n    for s, e in segs:\n        L = e - s\n        if L < 20: continue\n        ctrl = np.linspace(0, L - 1, 6); disp = rng.normal(0, amp, (6, 3)); t = np.arange(L)\n        X[s:e] += np.vstack([np.interp(t, ctrl, disp[:, k]) for k in range(3)]).T\n    return X\n\n\ndef generate_rna_structure(sequence: str, seed=None) -> np.ndarray:\n    \"\"\"Idealized A-form RNA helix — last-resort de-novo fallback.\"\"\"\n    if seed is not None:\n        np.random.seed(seed)\n    n = len(sequence); coords = np.zeros((n, 3))\n    for i in range(n):\n        ang = i * 0.6\n        coords[i] = [10.0 * np.cos(ang), 10.0 * np.sin(ang), i * 2.5]\n    return coords\n\n\n# ─────────────── TBM Phase ───────────────────────────────────────────────────\ndef tbm_phase(test_df, train_seqs_df, train_coords_dict, segments_map):\n    \"\"\"\n    Phase 1 — Template-Based Modeling.\n\n    Returns\n    -------\n    template_predictions : {target_id: [np.ndarray(seq_len, 3), ...]}\n        0 to N_SAMPLE predictions per target, from real templates.\n    protenix_queue : {target_id: (n_needed, full_sequence)}\n        Targets that still need more predictions.\n    \"\"\"\n    print(f\"\\n{'='*60}\")\n    print(f\"PHASE 1: Template-Based Modeling\")\n    print(f\"  MIN_SIMILARITY = {MIN_SIMILARITY}  |  MIN_PCT_IDENTITY = {MIN_PERCENT_IDENTITY}\")\n    print(f\"{'='*60}\")\n    t0 = time.time()\n\n    template_predictions: dict = {}\n    protenix_queue:       dict = {}\n\n    for _, row in test_df.iterrows():\n        tid = row[\"target_id\"]\n        seq = row[\"sequence\"]\n        segs = segments_map.get(tid, [(0, len(seq))])\n\n        similar = find_similar_sequences_detailed(seq, train_seqs_df, train_coords_dict, top_n=30)\n        preds   = []\n        used    = set()\n\n        for i, (tmpl_id, tmpl_seq, sim, tmpl_coords, pct_id, _, _) in enumerate(similar):\n            if len(preds) >= N_SAMPLE:\n                break\n            if sim < MIN_SIMILARITY or pct_id < MIN_PERCENT_IDENTITY:\n                break           # list is sorted by sim, so no point continuing\n            if tmpl_id in used:\n                continue\n\n            rng     = np.random.default_rng((row.name * 10000000000 + i * 10007) % (2**32))\n            adapted = adapt_template_to_query(seq, tmpl_seq, tmpl_coords)\n\n            # Diversity transforms (same strategy as the 0-409 TBM notebook)\n            slot = len(preds)\n            if slot == 0:\n                X = adapted\n            elif slot == 1:\n                X = adapted + rng.normal(0, max(0.01, (0.40 - sim) * 0.06), adapted.shape)\n            elif slot == 2:\n                longest = max(segs, key=lambda se: se[1] - se[0])\n                X = apply_hinge(adapted, longest, rng)\n            elif slot == 3:\n                X = jitter_chains(adapted, segs, rng)\n            else:\n                X = smooth_wiggle(adapted, segs, rng)\n\n            refined = adaptive_rna_constraints(X, tid, segments_map, confidence=sim)\n            preds.append(refined)\n            used.add(tmpl_id)\n\n        template_predictions[tid] = preds\n        n_needed = N_SAMPLE - len(preds)\n        if n_needed > 0:\n            protenix_queue[tid] = (n_needed, seq)\n            print(f\"  {tid} ({len(seq)} nt): {len(preds)} TBM → need {n_needed} from Protenix\")\n        else:\n            print(f\"  {tid} ({len(seq)} nt): all {N_SAMPLE} from TBM ✓\")\n\n    elapsed = time.time() - t0\n    n_full  = len(test_df) - len(protenix_queue)\n    print(f\"\\nPhase 1 done in {elapsed:.1f}s\")\n    print(f\"  Fully covered by TBM : {n_full}\")\n    print(f\"  Need Protenix        : {len(protenix_queue)}\")\n    return template_predictions, protenix_queue\n\n\n\n# ============================================================\n# GUDHI-based H1 PH rerank for RNA 3D candidates\n# ============================================================\n# Notes:\n# - If gudhi is unavailable, reranking is skipped safely.\n# - To enable online install in environments that allow it, uncomment:\n#     !pip -q install gudhi\n# ============================================================\n\nPH_RERANK_ENABLED = True\nPH_MODE = \"gudhi_h1\"\n\n# PH is used only as a small auxiliary signal on top of the original candidate order.\n# final_score = base_order_score + PH_BLEND_LAMBDA * normalized_ph_score\nPH_BLEND_LAMBDA   = 0.080\nPH_BLEND_BASE_GAP = 0.035\n\nPH_ALPHA = 1.0\nPH_BETA  = 0.25   # bond penalty weight\nPH_GAMMA = 0.20   # clash penalty weight\n\nPH_MAX_POINTS = 96\nPH_H1_TOPK = 16\nPH_BETTI_BINS = 24\nPH_MAX_EDGE_MULT = 3.5\nPH_MIN_PERSISTENCE = 1e-3\nPH_VERBOSE = False\n\nEXPECTED_C1_STEP = 6.0\nCLASH_DIST = 2.2\n\ntry:\n    import gudhi as gd\n    GUDHI_AVAILABLE = True\nexcept Exception as _gudhi_exc:\n    gd = None\n    GUDHI_AVAILABLE = False\n    print(f\"[WARN] gudhi is not available. PH rerank will be skipped. ({_gudhi_exc})\")\n\ndef _safe_valid_coords(coords: np.ndarray) -> np.ndarray:\n    x = np.asarray(coords, dtype=np.float32)\n    if x.ndim != 2 or x.shape[1] != 3:\n        return np.zeros((0, 3), dtype=np.float32)\n    mask = np.isfinite(x).all(axis=1)\n    return x[mask]\n\ndef _pairwise_distances(x: np.ndarray) -> np.ndarray:\n    diff = x[:, None, :] - x[None, :, :]\n    D = np.sqrt(np.sum(diff * diff, axis=-1))\n    return D.astype(np.float32)\n\ndef _nearest_neighbor_scale_from_D(D: np.ndarray) -> float:\n    n = D.shape[0]\n    if n <= 1:\n        return 1.0\n    A = D.copy()\n    np.fill_diagonal(A, np.inf)\n    nn = np.min(A, axis=1)\n    s = float(np.median(nn[np.isfinite(nn)])) if np.isfinite(nn).any() else 1.0\n    return max(s, 1e-6)\n\ndef _subsample_coords(x: np.ndarray, max_points: int = PH_MAX_POINTS) -> np.ndarray:\n    n = len(x)\n    if n <= max_points:\n        return x\n    idx = np.linspace(0, n - 1, max_points).round().astype(int)\n    idx = np.unique(idx)\n    return x[idx]\n\ndef bond_length_penalty(coords: np.ndarray, expected_step: float = EXPECTED_C1_STEP) -> float:\n    x = _safe_valid_coords(coords)\n    if len(x) <= 1:\n        return 0.0\n    d = np.sqrt(np.sum((x[1:] - x[:-1]) ** 2, axis=1))\n    return float(np.mean((d - expected_step) ** 2))\n\ndef clash_penalty(coords: np.ndarray, clash_dist: float = CLASH_DIST) -> float:\n    x = _safe_valid_coords(coords)\n    n = len(x)\n    if n <= 3:\n        return 0.0\n    D = _pairwise_distances(x)\n    pen = 0.0\n    cnt = 0\n    for i in range(n):\n        for j in range(i + 2, n):\n            dij = float(D[i, j])\n            if dij < clash_dist:\n                pen += (clash_dist - dij) ** 2\n                cnt += 1\n    return pen / cnt if cnt > 0 else 0.0\n\ndef _extract_h1_pairs_from_simplex_tree(st, min_persistence: float = PH_MIN_PERSISTENCE):\n    pairs = []\n    pers = st.persistence()\n    for dim, bd in pers:\n        if dim != 1:\n            continue\n        b, d = float(bd[0]), float(bd[1])\n        if not np.isfinite(d):\n            continue\n        p = d - b\n        if p < min_persistence:\n            continue\n        pairs.append((b, d, p))\n    return pairs\n\ndef _h1_feature_from_pairs(\n    h1_pairs: list[tuple[float, float, float]],\n    topk: int = PH_H1_TOPK,\n    bins: int = PH_BETTI_BINS,\n    scale: float = 1.0,\n) -> np.ndarray:\n    if len(h1_pairs) == 0:\n        return np.zeros(topk + 8 + bins + bins, dtype=np.float32)\n\n    arr = np.asarray(h1_pairs, dtype=np.float32)\n    births = arr[:, 0] / max(scale, 1e-6)\n    deaths = arr[:, 1] / max(scale, 1e-6)\n    pers   = arr[:, 2] / max(scale, 1e-6)\n\n    pers_sorted = np.sort(pers)[::-1]\n    top = np.zeros(topk, dtype=np.float32)\n    m = min(topk, len(pers_sorted))\n    top[:m] = pers_sorted[:m]\n\n    stats = np.array([\n        float(len(pers)),\n        float(np.sum(pers)),\n        float(np.mean(pers)),\n        float(np.std(pers)),\n        float(np.max(pers)),\n        float(np.mean(births)),\n        float(np.std(births)),\n        float(np.mean(deaths)),\n    ], dtype=np.float32)\n\n    max_birth = max(float(np.max(births)), 1e-6)\n    max_pers  = max(float(np.max(pers)), 1e-6)\n\n    birth_hist, _ = np.histogram(\n        births, bins=bins, range=(0.0, max_birth), density=False\n    )\n    pers_hist, _ = np.histogram(\n        pers, bins=bins, range=(0.0, max_pers), density=False\n    )\n\n    birth_hist = birth_hist.astype(np.float32)\n    pers_hist  = pers_hist.astype(np.float32)\n\n    if birth_hist.sum() > 0:\n        birth_hist /= birth_hist.sum()\n    if pers_hist.sum() > 0:\n        pers_hist /= pers_hist.sum()\n\n    return np.concatenate([top, stats, birth_hist, pers_hist], axis=0).astype(np.float32)\n\ndef gudhi_h1_feature(coords: np.ndarray) -> np.ndarray:\n    if not GUDHI_AVAILABLE:\n        raise RuntimeError(\"gudhi is not installed\")\n\n    x = _safe_valid_coords(coords)\n    if len(x) < 4:\n        return np.zeros(PH_H1_TOPK + 8 + PH_BETTI_BINS + PH_BETTI_BINS, dtype=np.float32)\n\n    x = _subsample_coords(x, PH_MAX_POINTS)\n    if len(x) < 4:\n        return np.zeros(PH_H1_TOPK + 8 + PH_BETTI_BINS + PH_BETTI_BINS, dtype=np.float32)\n\n    D = _pairwise_distances(x)\n    scale = _nearest_neighbor_scale_from_D(D)\n    max_edge = max(PH_MAX_EDGE_MULT * scale, 1e-3)\n\n    rips = gd.RipsComplex(distance_matrix=D, max_edge_length=max_edge)\n    st = rips.create_simplex_tree(max_dimension=2)\n    h1_pairs = _extract_h1_pairs_from_simplex_tree(st, min_persistence=PH_MIN_PERSISTENCE)\n\n    return _h1_feature_from_pairs(\n        h1_pairs,\n        topk=PH_H1_TOPK,\n        bins=PH_BETTI_BINS,\n        scale=scale,\n    )\n\ndef _make_reference_ph_feature(candidates: list[np.ndarray]) -> np.ndarray:\n    feats = []\n    for c in candidates:\n        try:\n            feats.append(gudhi_h1_feature(c))\n        except Exception as e:\n            if PH_VERBOSE:\n                print(\"[PH ref feature error]\", e)\n\n    if len(feats) == 0:\n        return np.zeros(PH_H1_TOPK + 8 + PH_BETTI_BINS + PH_BETTI_BINS, dtype=np.float32)\n\n    F = np.stack(feats, axis=0)\n    return np.median(F, axis=0).astype(np.float32)\n\ndef ph_score_candidate(coords: np.ndarray, ref_feat: np.ndarray) -> tuple[float, dict]:\n    feat = gudhi_h1_feature(coords)\n    ph_dist = float(np.linalg.norm(feat - ref_feat))\n    p_bond  = bond_length_penalty(coords)\n    p_clash = clash_penalty(coords)\n\n    score = (\n        PH_ALPHA * (-ph_dist)\n        - PH_BETA  * p_bond\n        - PH_GAMMA * p_clash\n    )\n    aux = {\n        \"ph_dist\": ph_dist,\n        \"bond_pen\": p_bond,\n        \"clash_pen\": p_clash,\n        \"score\": score,\n    }\n    return score, aux\n\n\ndef plot_persistence_diagram(h1_pairs: list, target_id: str, sample_idx: int,\n                              ref_feat: np.ndarray = None) -> None:\n    \"\"\"\n    H1パーシステント図（birth-death plot）を描画する。\n    h1_pairs: [(birth, death, persistence), ...]\n    \"\"\"\n    try:\n        import matplotlib\n        matplotlib.use(\"Agg\")\n        import matplotlib.pyplot as plt\n        import matplotlib.gridspec as gridspec\n\n        fig = plt.figure(figsize=(12, 5))\n        fig.suptitle(f\"{target_id}  sample_{sample_idx}  —  H1 Persistence Diagram\",\n                     fontsize=12, fontweight=\"bold\")\n        gs = gridspec.GridSpec(1, 2, figure=fig, wspace=0.35)\n\n        # ── Left: birth-death plot ────────────────────────────────────────\n        ax1 = fig.add_subplot(gs[0])\n        if len(h1_pairs) > 0:\n            arr    = np.array(h1_pairs)\n            births = arr[:, 0]\n            deaths = arr[:, 1]\n            pers   = arr[:, 2]\n\n            # 寿命でサイズを変える（長寿命ほど大きい点）\n            sizes  = 20 + 200 * (pers / (pers.max() + 1e-9))\n            sc = ax1.scatter(births, deaths, c=pers, s=sizes,\n                             cmap=\"plasma\", alpha=0.8, edgecolors=\"k\", linewidths=0.4)\n            plt.colorbar(sc, ax=ax1, label=\"Persistence\")\n\n            # 対角線（birth = death）\n            lim = max(deaths.max(), births.max()) * 1.05\n            ax1.plot([0, lim], [0, lim], \"k--\", lw=0.8, alpha=0.5, label=\"birth=death\")\n            ax1.set_xlim(-0.02 * lim, lim)\n            ax1.set_ylim(-0.02 * lim, lim)\n        else:\n            ax1.text(0.5, 0.5, \"No H1 pairs\", ha=\"center\", va=\"center\",\n                     transform=ax1.transAxes, fontsize=11, color=\"gray\")\n\n        ax1.set_xlabel(\"Birth (Å)\")\n        ax1.set_ylabel(\"Death (Å)\")\n        ax1.set_title(\"H1  Birth–Death Plot\")\n        ax1.set_aspect(\"equal\")\n        ax1.grid(True, alpha=0.3)\n\n        # ── Right: persistence barcode (寿命順) ──────────────────────────\n        ax2 = fig.add_subplot(gs[1])\n        if len(h1_pairs) > 0:\n            arr_s  = sorted(h1_pairs, key=lambda x: x[2], reverse=True)\n            colors = plt.cm.plasma(\n                np.linspace(0.9, 0.2, len(arr_s))\n            )\n            for k, (b, d, p) in enumerate(arr_s):\n                ax2.plot([b, d], [k, k], lw=2.5, color=colors[k], alpha=0.85)\n            ax2.set_yticks(range(len(arr_s)))\n            ax2.set_yticklabels([f\"p={p:.2f}\" for _, _, p in arr_s], fontsize=7)\n            ax2.set_xlabel(\"Filtration value (Å)\")\n        else:\n            ax2.text(0.5, 0.5, \"No H1 pairs\", ha=\"center\", va=\"center\",\n                     transform=ax2.transAxes, fontsize=11, color=\"gray\")\n\n        ax2.set_title(\"H1  Barcode (寿命順)\")\n        ax2.grid(True, alpha=0.3, axis=\"x\")\n\n        plt.tight_layout()\n        save_path = f\"/kaggle/working/ph_diagram_{target_id}_s{sample_idx}.png\"\n        plt.savefig(save_path, dpi=100, bbox_inches=\"tight\")\n        plt.show()\n        plt.close(fig)\n        print(f\"  [PH plot] saved → {save_path}\")\n    except Exception as e:\n        print(f\"  [PH plot] failed: {e}\")\n\n\ndef _get_h1_pairs_for_coords(coords: np.ndarray) -> list:\n    \"\"\"coords (L,3) から H1 birth-death ペアを返す（可視化用）\"\"\"\n    if not GUDHI_AVAILABLE:\n        return []\n    x = _safe_valid_coords(coords)\n    if len(x) < 4:\n        return []\n    x   = _subsample_coords(x, PH_MAX_POINTS)\n    D   = _pairwise_distances(x)\n    sc  = _nearest_neighbor_scale_from_D(D)\n    me  = max(PH_MAX_EDGE_MULT * sc, 1e-3)\n    rips = gd.RipsComplex(distance_matrix=D, max_edge_length=me)\n    st   = rips.create_simplex_tree(max_dimension=2)\n    return _extract_h1_pairs_from_simplex_tree(st, min_persistence=PH_MIN_PERSISTENCE)\n\ndef ph_rerank_candidates(\n    candidates: list[np.ndarray],\n    target_id: str | None = None,\n    verbose: bool = False,\n) -> tuple[list[np.ndarray], list[dict]]:\n    \"\"\"\n    Original candidate order remains the primary signal.\n    This function is intended to rerank only a local candidate block\n    (now used for the Protenix block only).\n\n    final_score_i = base_score_i + PH_BLEND_LAMBDA * ph_norm_i\n    where\n        base_score_i = -PH_BLEND_BASE_GAP * original_rank_i\n    \"\"\"\n    if candidates is None or len(candidates) <= 1:\n        dbg = [{\"rank\": 1, \"old_idx\": 0, \"score\": 0.0}] if candidates else []\n        return candidates, dbg\n\n    n = len(candidates)\n\n    # If GUDHI is unavailable, keep the original order.\n    if not GUDHI_AVAILABLE:\n        dbg = [{\n            \"rank\": i + 1,\n            \"old_idx\": i,\n            \"base_score\": -PH_BLEND_BASE_GAP * i,\n            \"ph_raw_score\": 0.0,\n            \"ph_norm_score\": 0.0,\n            \"score\": -PH_BLEND_BASE_GAP * i,\n        } for i in range(n)]\n        return candidates, dbg\n\n    ref_feat = _make_reference_ph_feature(candidates)\n\n    raw_rows = []\n    for i, c in enumerate(candidates):\n        try:\n            ph_raw, aux = ph_score_candidate(c, ref_feat)\n        except Exception as e:\n            ph_raw = -1e18\n            aux = {\n                \"ph_dist\": np.inf,\n                \"bond_pen\": np.inf,\n                \"clash_pen\": np.inf,\n                \"score\": ph_raw,\n                \"error\": str(e),\n            }\n        raw_rows.append((i, c, ph_raw, aux))\n\n    ph_vals = np.array([r[2] for r in raw_rows], dtype=np.float32)\n    finite_mask = np.isfinite(ph_vals) & (ph_vals > -1e17)\n\n    ph_norm = np.zeros(n, dtype=np.float32)\n    if finite_mask.sum() >= 2:\n        mu = float(ph_vals[finite_mask].mean())\n        sd = float(ph_vals[finite_mask].std())\n        if sd > 1e-8:\n            ph_norm[finite_mask] = (ph_vals[finite_mask] - mu) / sd\n    elif finite_mask.sum() == 1:\n        ph_norm[finite_mask] = 0.0\n\n    scored = []\n    for i, c, ph_raw, aux in raw_rows:\n        base_score = -PH_BLEND_BASE_GAP * float(i)\n        final_score = base_score + PH_BLEND_LAMBDA * float(ph_norm[i])\n        row_aux = {\n            **aux,\n            \"base_score\": base_score,\n            \"ph_raw_score\": float(ph_raw) if np.isfinite(ph_raw) else ph_raw,\n            \"ph_norm_score\": float(ph_norm[i]),\n            \"score\": final_score,\n        }\n        scored.append((final_score, i, c, row_aux))\n\n    scored.sort(key=lambda z: z[0], reverse=True)\n    reranked = [z[2] for z in scored]\n\n    dbg = []\n    for rank, (final_score, old_idx, _, aux) in enumerate(scored, start=1):\n        dbg.append({\"rank\": rank, \"old_idx\": old_idx, **aux})\n\n    if verbose:\n        tag = f\"[PH blend rerank][{target_id}]\" if target_id is not None else \"[PH blend rerank]\"\n        print(tag)\n        for d in dbg[:min(5, len(dbg))]:\n            print(\n                f\"  rank={d['rank']} old_idx={d['old_idx']} \"\n                f\"final={d['score']:.4f} \"\n                f\"base={d['base_score']:.4f} \"\n                f\"ph_norm={d['ph_norm_score']:.4f} \"\n                f\"ph_raw={d['ph_raw_score']:.4f}\"\n            )\n\n    # ── パーシステント図を描画（全候補） ─────────────────────────────────\n    if GUDHI_AVAILABLE:\n        tid_label = target_id.replace(\"::\", \"_\") if target_id else \"unknown\"\n        for rank_info in dbg:\n            s_idx   = rank_info[\"old_idx\"]\n            c_orig  = candidates[s_idx] if s_idx < len(candidates) else reranked[0]\n            h1p     = _get_h1_pairs_for_coords(c_orig)\n            plot_persistence_diagram(h1p, tid_label, s_idx)\n\n    return reranked, dbg\n\n\n# ─────────────── Main ────────────────────────────────────────────────────────\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}. \"\n            \"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    ensure_required_files(root_dir)\n    seed_everything(SEED)\n\n    # ── Load test data ──────────────────────────────────────────────────────\n    test_df_full = pd.read_csv(test_csv)\n    test_df      = (test_df_full.head(LOCAL_N_SAMPLES) if not IS_KAGGLE\n                    else test_df_full).reset_index(drop=True)\n    print(f\"Test targets : {len(test_df)}\"\n          + (\" (LOCAL MODE)\" if not IS_KAGGLE else \"\"))\n\n    seq_by_id = dict(zip(test_df[\"target_id\"], test_df[\"sequence\"]))\n\n    # Truncated copy for Protenix (Protenix has token limits)\n    test_df_trunc = test_df.copy()\n    test_df_trunc[\"sequence\"] = test_df_trunc[\"sequence\"].str[:MAX_SEQ_LEN]\n\n    # ── Load training data for TBM ──────────────────────────────────────────\n    print(\"\\nLoading training data for TBM …\")\n    train_seqs   = pd.read_csv(DEFAULT_TRAIN_CSV)\n    val_seqs     = pd.read_csv(DEFAULT_VAL_CSV)\n    train_labels = pd.read_csv(DEFAULT_TRAIN_LBLS)\n    val_labels   = pd.read_csv(DEFAULT_VAL_LBLS)\n\n    combined_seqs   = pd.concat([train_seqs,   val_seqs],    ignore_index=True)\n    combined_labels = pd.concat([train_labels, val_labels],  ignore_index=True)\n    train_coords    = process_labels(combined_labels)\n    segments_map, _ = build_segments_map(test_df)\n\n    print(f\"Template pool: {len(combined_seqs)} sequences, {len(train_coords)} structures\")\n\n    # ─── PHASE 1: TBM ──────────────────────────────────────────────────────\n    template_preds, protenix_queue = tbm_phase(\n        test_df, combined_seqs, train_coords, segments_map\n    )\n\n    # ─── PHASE 2: Protenix (only for targets that need extra predictions) ──\n    protenix_preds: dict = {}   # target_id -> np.ndarray (n_needed, seq_len, 3)\n\n    if protenix_queue and USE_PROTENIX:\n        print(f\"\\n{'='*60}\")\n        print(f\"PHASE 2: Protenix for {len(protenix_queue)} targets\")\n        print(f\"{'='*60}\")\n\n        work_dir = Path(\"/kaggle/working\")\n        work_dir.mkdir(parents=True, exist_ok=True)\n\n        # ── 1. Preparation: create tasks for all sequences/chunks ────────────\n        tasks = []          # list of dict for input_json\n        chunk_info = {}     # target_id -> list of {\"name\": chunk_name, \"range\": (s, e)}\n        \n        for target_id, (n_needed, full_seq) in protenix_queue.items():\n            seq_len = len(full_seq)\n            if seq_len <= MAX_SEQ_LEN:\n                tasks.append({\"target_id\": target_id, \"sequence\": full_seq})\n                chunk_info[target_id] = [{\"name\": target_id, \"range\": (0, seq_len)}]\n                print(f\"  {target_id} ({seq_len} nt): single pass queued\")\n            else:\n                chunks = split_into_chunks(seq_len, MAX_SEQ_LEN, CHUNK_OVERLAP)\n                print(f\"  {target_id} ({seq_len} nt): {len(chunks)} chunks queued \"\n                      f\"{[(s, e) for s, e in chunks]}\")\n                \n                chunk_info[target_id] = []\n                for ci, (cs, ce) in enumerate(chunks):\n                    chunk_name = f\"{target_id}_chunk{ci}\"\n                    sub_seq = full_seq[cs:ce]\n                    tasks.append({\"target_id\": chunk_name, \"sequence\": sub_seq})\n                    chunk_info[target_id].append({\"name\": chunk_name, \"range\": (cs, ce)})\n\n        # Build combined input JSON\n        tasks_df = pd.DataFrame(tasks)\n        input_json_path = str(work_dir / \"protenix_queue_input.json\")\n        build_input_json(tasks_df, input_json_path)\n\n        from protenix.data.inference.infer_dataloader import InferenceDataset\n        from runner.inference import (InferenceRunner,\n                                      update_gpu_compatible_configs,\n                                      update_inference_configs)\n\n        # Initialize model exactly ONCE\n        configs = build_configs(input_json_path, str(work_dir / \"outputs\"), MODEL_NAME)\n        configs = update_gpu_compatible_configs(configs)\n        runner  = InferenceRunner(configs)\n        dataset = InferenceDataset(configs)\n\n        # ── 2. Inference: process dataset and collect predictions ────────────\n        raw_predictions = {}  # sample_name -> coords (np.ndarray or None)\n\n        def _extract_c1_coords(prediction, feat, chunk_seq_len, raw_coords):\n            if \"centre_atom_mask\" in feat:\n                mask = (feat[\"centre_atom_mask\"] == 1).to(raw_coords.device)\n            elif \"atom_to_tokatom_idx\" in feat:\n                m11 = (feat[\"atom_to_tokatom_idx\"] == 11).to(raw_coords.device)\n                m12 = (feat[\"atom_to_tokatom_idx\"] == 12).to(raw_coords.device)\n                c11, c12 = m11.sum(), m12.sum()\n                mask = m11 if abs(c11 - chunk_seq_len) < abs(c12 - chunk_seq_len) else m12\n            else:\n                mask = torch.zeros(raw_coords.shape[1], dtype=torch.bool, device=raw_coords.device)\n            \n            coords = raw_coords[:, mask, :].detach().cpu().numpy()\n            \n            # Collapse check\n            if coords.shape[1] > 1:\n                diffs = np.linalg.norm(coords[0, 1:] - coords[0, :-1], axis=-1)\n                if np.all(diffs < 1e-4):\n                    print(f\"    WARNING: Collapsed coordinates detected\")\n                    return None\n            \n            if coords.shape[1] != chunk_seq_len:\n                if coords.shape[1] == 1 and chunk_seq_len > 1:\n                    return None\n                padded = np.zeros((coords.shape[0], chunk_seq_len, 3), dtype=np.float32)\n                ml = min(coords.shape[1], chunk_seq_len)\n                padded[:, :ml, :] = coords[:, :ml, :]\n                coords = padded\n            return coords\n\n        for i in tqdm(range(len(dataset)), desc=\"Protenix Inference\"):\n            data, atom_array, err = dataset[i]\n            sample_name = data.get(\"sample_name\", f\"sample_{i}\")\n            \n            if err:\n                print(f\"  {sample_name} data error: {err}\")\n                raw_predictions[sample_name] = None\n                del data, atom_array, err\n                gc.collect(); torch.cuda.empty_cache(); gc.collect()\n                continue\n            \n            # Find how many samples are needed for this specific query\n            target_id = sample_name.split(\"_chunk\")[0] if \"_chunk\" in sample_name else sample_name\n            n_needed = protenix_queue.get(target_id, (N_SAMPLE, \"\"))[0]\n            sub_seq_len = data[\"N_token\"].item() # roughly correct\n            \n            try:\n                new_cfg = update_inference_configs(configs, sub_seq_len)\n                new_cfg.sample_diffusion.N_sample = n_needed\n                runner.update_model_configs(new_cfg)\n                \n                pred = runner.predict(data)\n                raw_coords = pred[\"coordinate\"]\n                \n                coords = _extract_c1_coords(pred, data[\"input_feature_dict\"], \n                                            sub_seq_len, raw_coords)\n\n                # ── pLDDT抽出してサンプルを降順ソート ──────────────────────\n                if coords is not None:\n                    plddt_raw = None\n                    for key in (\"plddt\", \"atom_plddt\", \"confidence_score\",\n                                \"predicted_lddt\", \"plddts\"):\n                        if key in pred:\n                            plddt_raw = pred[key]\n                            print(f\"[DEBUG] pLDDT key='{key}' shape={plddt_raw.shape}\")\n                            break\n                    if plddt_raw is None:\n                        print(\"[DEBUG] pLDDT not found:\", list(pred.keys()))\n\n                    plddt_scores = None\n                    if plddt_raw is not None:\n                        try:\n                            p = plddt_raw.detach().cpu().numpy()\n                            n_s = coords.shape[0]\n                            if p.ndim == 2 and p.shape[0] == n_s:\n                                plddt_scores = p.mean(axis=-1)\n                            elif p.ndim == 2 and p.shape[1] == n_s:\n                                plddt_scores = p.T.mean(axis=-1)\n                            elif p.ndim == 1 and len(p) == n_s:\n                                plddt_scores = p\n                            else:\n                                print(f\"[DEBUG] pLDDT shape {p.shape} incompatible\")\n                        except Exception as e:\n                            print(f\"[DEBUG] pLDDT failed: {e}\")\n\n                    if plddt_scores is not None and len(plddt_scores) == coords.shape[0]:\n                        order  = np.argsort(plddt_scores)[::-1]\n                        coords = coords[order]\n                        print(f\"  {sample_name}: sorted by pLDDT \"\n                              f\"{plddt_scores[order].round(1).tolist()}\")\n\n                raw_predictions[sample_name] = coords\n            except Exception as exc:\n                print(f\"  {sample_name} inference failed: {exc}\")\n                import traceback; traceback.print_exc()\n                raw_predictions[sample_name] = None\n            finally:\n                try: del pred, data, atom_array, raw_coords\n                except: pass\n                gc.collect(); torch.cuda.empty_cache(); gc.collect()\n\n        # ── 3. Post-processing: Stitching and final formatting ───────────────\n        for target_id, (n_needed, full_seq) in protenix_queue.items():\n            seq_len = len(full_seq)\n            chunks = chunk_info.get(target_id, [])\n            \n            if not chunks:\n                continue\n\n            if len(chunks) == 1:\n                # Single pass\n                coords = raw_predictions.get(target_id)\n                protenix_preds[target_id] = coords\n                if coords is not None:\n                    print(f\"  {target_id}: {coords.shape[0]} predictions generated\")\n                else:\n                    print(f\"  {target_id}: FAILED\")\n            else:\n                # Stitch chunks together\n                chunk_results_per_sample = {s: [] for s in range(n_needed)}\n                all_ok = True\n                \n                for ci, cinfo in enumerate(chunks):\n                    cname = cinfo[\"name\"]\n                    crange = cinfo[\"range\"]\n                    ccoords = raw_predictions.get(cname)\n                    \n                    if ccoords is None:\n                        all_ok = False\n                        break\n                    \n                    for s_idx in range(n_needed):\n                        if s_idx < ccoords.shape[0]:\n                            chunk_results_per_sample[s_idx].append((ccoords[s_idx], crange))\n                        else:\n                            chunk_results_per_sample[s_idx].append((ccoords[-1], crange))\n                \n                if not all_ok:\n                    print(f\"  {target_id}: chunked inference incomplete, using fallback\")\n                    protenix_preds[target_id] = None\n                    continue\n                \n                stitched_samples = []\n                for s_idx in range(n_needed):\n                    items = chunk_results_per_sample[s_idx]\n                    coords_list = [c for c, _ in items]\n                    ranges_list = [r for _, r in items]\n                    full_coords = stitch_chunk_coords(coords_list, ranges_list, seq_len)\n                    stitched_samples.append(full_coords)\n                \n                result = np.stack(stitched_samples, axis=0)\n                protenix_preds[target_id] = result\n                print(f\"  {target_id}: {result.shape[0]} stitched predictions generated\")\n# ...existing code...\n\n    elif protenix_queue and not USE_PROTENIX:\n        print(f\"\\nPHASE 2 skipped (USE_PROTENIX=False). \"\n              f\"De-novo fallback will cover {len(protenix_queue)} targets.\")\n\n    # ─── PHASE 3: Combine everything ───────────────────────────────────────\n    print(f\"\\n{'='*60}\")\n    print(\"PHASE 3: Combine TBM + Protenix + de-novo fallback\")\n    print(f\"{'='*60}\")\n\n    all_rows = []\n\n    for _, row in test_df.iterrows():\n        tid = row[\"target_id\"]\n        seq = row[\"sequence\"]\n\n        combined: list = list(template_preds.get(tid, []))  # TBM predictions (keep original order)\n\n        # Append Protenix predictions to fill remaining slots.\n        # PH is applied ONLY inside the Protenix candidate block, so TBM ordering is preserved.\n        ptx = protenix_preds.get(tid)\n        if ptx is not None and ptx.ndim == 3:\n            ptx_list = [ptx[j] for j in range(ptx.shape[0])]\n\n            if PH_RERANK_ENABLED and len(ptx_list) > 1:\n                ptx_list, ph_dbg = ph_rerank_candidates(\n                    ptx_list,\n                    target_id=f\"{tid}::Protenix\",\n                    verbose=True,   # Always draw diagram\n                )\n\n            for cand in ptx_list:\n                if len(combined) >= N_SAMPLE:\n                    break\n                combined.append(cand)  # (seq_len, 3)\n\n        # De-novo fallback for any still-empty slots\n        n_denovo = 0\n        while len(combined) < N_SAMPLE:\n            seed_val = row.name * 1000000 + len(combined) * 1000\n            dn       = generate_rna_structure(seq, seed=seed_val)\n            combined.append(adaptive_rna_constraints(dn, tid, segments_map, confidence=0.2))\n            n_denovo += 1\n\n        if n_denovo:\n            print(f\"  {tid}: {n_denovo} slot(s) filled with de-novo fallback\")\n\n        # Stack to (N_SAMPLE, seq_len, 3) and write rows\n        stacked = np.stack(combined[:N_SAMPLE], axis=0)\n        all_rows.extend(coords_to_rows(tid, seq, stacked))\n\n    # ── Save ───────────────────────────────────────────────────────────────\n    sub = pd.DataFrame(all_rows)\n    cols = [\"ID\", \"resname\", \"resid\"] + [\n        f\"{c}_{i}\" for i in range(1, N_SAMPLE + 1) for c in [\"x\", \"y\", \"z\"]\n    ]\n    coord_cols = [c for c in cols if c.startswith((\"x_\", \"y_\", \"z_\"))]\n    sub[coord_cols] = sub[coord_cols].clip(-999.999, 9999.999)\n    sub[cols].to_csv(output_csv, index=False)\n\n    print(f\"\\n✓ Saved submission to {output_csv}  ({len(sub):,} rows)\")\n","metadata":{"_cell_guid":"6ed2b9a0-71f7-4881-9338-57e86955b821","_uuid":"908e861c-11fa-40ac-80d4-08bed075a80a","collapsed":false,"execution":{"iopub.status.busy":"2026-03-15T10:49:50.490597Z","iopub.execute_input":"2026-03-15T10:49:50.490789Z","iopub.status.idle":"2026-03-15T10:49:50.761702Z","shell.execute_reply.started":"2026-03-15T10:49:50.49077Z","shell.execute_reply":"2026-03-15T10:49:50.760916Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.254346,"end_time":"2026-03-09T10:22:24.435419","exception":false,"start_time":"2026-03-09T10:22:24.181073","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Main","metadata":{"_cell_guid":"d31c3343-caae-456b-8115-14c70473c17a","_uuid":"37cba882-9668-4c81-8435-00ce4c903142","collapsed":false,"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.004174,"end_time":"2026-03-09T10:22:24.443976","exception":false,"start_time":"2026-03-09T10:22:24.439802","status":"completed"},"tags":[]}},{"cell_type":"code","source":"\nif __name__ == \"__main__\":\n    main()","metadata":{"_cell_guid":"d9d3b108-731d-422e-8f9b-ca8a6e1a3e77","_uuid":"3d1b164f-da40-4a91-ad1e-55d29a38f63a","collapsed":false,"execution":{"iopub.status.busy":"2026-03-15T10:49:50.762711Z","iopub.execute_input":"2026-03-15T10:49:50.763171Z"},"jupyter":{"outputs_hidden":false},"papermill":{"duration":4679.440125,"end_time":"2026-03-09T11:40:23.888556","exception":false,"start_time":"2026-03-09T10:22:24.448431","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#read submission.csv\nsubmission_path = \"/kaggle/working/submission.csv\"\nsubmission_df = pd.read_csv(submission_path)\nprint(submission_df.head(20))","metadata":{"_cell_guid":"c899176c-ab4e-4196-8fc3-1e9b4dc5d4e6","_uuid":"a237c1a0-aa74-46a3-a316-df8395580177","collapsed":false,"jupyter":{"outputs_hidden":false},"papermill":{"duration":0.071092,"end_time":"2026-03-09T11:40:23.971561","exception":false,"start_time":"2026-03-09T11:40:23.900469","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"'''# ============================================================\n# Post-submission noise augmentation (reusable)\n# ============================================================\n# The main pipeline outputs submission_org.csv.\n# This cell adds noise to generate submission.csv (for final submission).\n# The noise is saved to noise_record.npz and can be reused on subsequent runs.\n# ============================================================\n\nimport numpy as np\nimport pandas as pd\nimport os\n\nNOISE_SEED     = 42          # ← Change this for a different noise pattern\nNOISE_STD      = 0.1         # Standard deviation of added noise (Å)\nSUBMISSION_IN  = \"/kaggle/working/submission.csv\"   # メイン出力（ノイズなし）\nSUBMISSION_OUT = \"/kaggle/working/submission.csv\"        # 最終提出用（ノイズあり）\nNOISE_SAVE     = \"/kaggle/working/noise_record.npz\"\nNOISE_LOAD     = \"/kaggle/input/datasets/takahiro3110/v32-noise/noise_record.npz\" \n\nsub = pd.read_csv(SUBMISSION_IN)\n\ncoord_cols = [c for c in sub.columns if c.startswith((\"x_\", \"y_\", \"z_\"))]\nvalues = sub[coord_cols].values.astype(np.float64)\n\nif NOISE_LOAD and os.path.exists(NOISE_LOAD):\n    loaded = np.load(NOISE_LOAD)\n    noise  = loaded[\"noise\"]\n    assert noise.shape == values.shape, (\n        f\"Noise shape mismatch: saved={noise.shape}, current={values.shape}\"\n    )\n    print(f\"✓ Noise reused: {NOISE_LOAD}  shape={noise.shape}\")\nelse:\n    rng   = np.random.default_rng(NOISE_SEED)\n    noise = rng.normal(0.0, NOISE_STD, values.shape)\n    np.savez_compressed(\n        NOISE_SAVE,\n        noise      = noise,\n        coord_cols = np.array(coord_cols),\n        seed       = np.array([NOISE_SEED]),\n        std        = np.array([NOISE_STD]),\n    )\n    print(f\"✓ Noise generated and saved: {NOISE_SAVE}  shape={noise.shape}  std={NOISE_STD}\")\n\nsub_out = sub.copy()\nsub_out[coord_cols] = values + noise\nsub_out.to_csv(SUBMISSION_OUT, index=False)\n\nprint(f\"✓ Saved: {SUBMISSION_OUT}\")\nprint(f\"  Rows: {len(sub_out)}\")\nprint(f\"  Noise stats: mean={noise.mean():.4f}  std={noise.std():.4f}  max_abs={np.abs(noise).max():.4f}\")\nprint()\nprint(\"How to reuse:\")\nprint(\"  Change NOISE_LOAD = \\\"/kaggle/working/noise_record.npz\\\" and re-run\")'''","metadata":{"papermill":{"duration":0.393116,"end_time":"2026-03-09T11:40:24.378045","exception":false,"start_time":"2026-03-09T11:40:23.984929","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Display the 5 most recently created persistent diagrams\n# ============================================================\nimport glob, os\nfrom IPython.display import display, Image\n\npng_files = sorted(\n    glob.glob(\"/kaggle/working/ph_diagram_*.png\"),\n    key=os.path.getmtime,\n    reverse=True\n)[:5]\n\nif png_files:\n    print(f\"Displaying {len(png_files)} most recent file(s) (newest first)\")\n    for path in png_files:\n        print(f\"  {os.path.basename(path)}\")\n        display(Image(filename=path))\nelse:\n    print(\"No ph_diagram_*.png files found.\")\n    print(\"Please ensure PH reranking is enabled (GUDHI_AVAILABLE=True) and has been run.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}