{"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":"none","dataSources":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":15231210,"isSourceIdPinned":false}],"dockerImageVersionId":31328,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-25T04:30:40.604321Z","iopub.execute_input":"2026-03-25T04:30:40.604661Z","iopub.status.idle":"2026-03-25T04:30:54.067959Z","shell.execute_reply.started":"2026-03-25T04:30:40.604634Z","shell.execute_reply":"2026-03-25T04:30:54.066702Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Stanford RNA 3D Folding Challenge 2\n# TM-oriented, memory-safe offline notebook baseline\n# with fuzzy Kaggle file discovery\n# ============================================================\n\nimport os\nimport glob\nimport random\nfrom dataclasses import dataclass\nfrom collections import Counter\n\nimport numpy as np\nimport pandas as pd\n\n# -----------------------------\n# Config\n# -----------------------------\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\n\nN_PREDS = 5\nMAX_TEMPLATES_FOR_RERANK = 12\nTOP_TEMPLATES_TO_BLEND = 4\nMAX_MSA_SEQS = 150\nRUN_VALIDATION = False  # safest for Kaggle submit runs\n\n# -----------------------------\n# File / directory discovery\n# -----------------------------\ndef list_all_files():\n    paths = []\n    search_roots = [\".\", \"/kaggle/input\", \"/kaggle/working\", \"/mnt/data\"]\n    for root in search_roots:\n        if os.path.exists(root):\n            for p in glob.glob(os.path.join(root, \"**\", \"*\"), recursive=True):\n                if os.path.isfile(p):\n                    paths.append(p)\n    return sorted(set(paths))\n\nALL_FILES = list_all_files()\n\ndef show_available_files(limit=200):\n    print(\"Available files:\")\n    for p in ALL_FILES[:limit]:\n        print(\" \", p)\n\ndef find_dir(dirname):\n    candidates = [\n        dirname,\n        f\"/kaggle/working/{dirname}\",\n        f\"/mnt/data/{dirname}\",\n    ]\n    for p in candidates:\n        if os.path.isdir(p):\n            return p\n\n    matches = glob.glob(f\"/kaggle/input/**/{dirname}\", recursive=True)\n    matches = [p for p in matches if os.path.isdir(p)]\n    if matches:\n        return matches[0]\n\n    matches = glob.glob(f\"**/{dirname}\", recursive=True)\n    matches = [p for p in matches if os.path.isdir(p)]\n    if matches:\n        return matches[0]\n\n    return None\n\ndef find_file_exact(filename):\n    candidates = [\n        filename,\n        f\"/kaggle/working/{filename}\",\n        f\"/mnt/data/{filename}\",\n    ]\n    for p in candidates:\n        if os.path.exists(p):\n            return p\n\n    matches = glob.glob(f\"/kaggle/input/**/{filename}\", recursive=True)\n    if matches:\n        return matches[0]\n\n    matches = glob.glob(f\"**/{filename}\", recursive=True)\n    if matches:\n        return matches[0]\n\n    return None\n\ndef find_file_fuzzy(required_terms, exclude_terms=None, required=True):\n    exclude_terms = exclude_terms or []\n    required_terms = [t.lower() for t in required_terms]\n    exclude_terms = [t.lower() for t in exclude_terms]\n\n    candidates = []\n    for p in ALL_FILES:\n        low = os.path.basename(p).lower()\n        if not low.endswith(\".csv\"):\n            continue\n        if all(t in low for t in required_terms) and not any(t in low for t in exclude_terms):\n            candidates.append(p)\n\n    candidates = sorted(\n        candidates,\n        key=lambda x: (\n            0 if x.startswith(\"/kaggle/input\") else 1,\n            len(x),\n            x,\n        )\n    )\n\n    if candidates:\n        return candidates[0]\n\n    if required:\n        print(\"\\nCould not find CSV with required terms:\", required_terms)\n        print(\"CSV files I can see:\")\n        for p in sorted([f for f in ALL_FILES if f.lower().endswith(\".csv\")])[:200]:\n            print(\" \", p)\n        raise FileNotFoundError(f\"Could not find CSV matching terms: {required_terms}\")\n    return None\n\ndef find_train_sequences():\n    p = find_file_exact(\"train_sequences.csv\")\n    if p:\n        return p\n    return find_file_fuzzy([\"train\", \"sequence\"], exclude_terms=[\"validation\", \"test\", \"sample\"])\n\ndef find_train_labels():\n    p = find_file_exact(\"train_labels.csv\")\n    if p:\n        return p\n    return find_file_fuzzy([\"train\", \"label\"], exclude_terms=[\"validation\", \"sample\"])\n\ndef find_validation_sequences():\n    p = find_file_exact(\"validation_sequences.csv\")\n    if p:\n        return p\n    return find_file_fuzzy([\"validation\", \"sequence\"], exclude_terms=[\"train\", \"test\"], required=False)\n\ndef find_validation_labels():\n    p = find_file_exact(\"validation_labels.csv\")\n    if p:\n        return p\n    return find_file_fuzzy([\"validation\", \"label\"], exclude_terms=[\"train\", \"sample\"], required=False)\n\ndef find_test_sequences():\n    p = find_file_exact(\"test_sequences.csv\")\n    if p:\n        return p\n    return find_file_fuzzy([\"test\", \"sequence\"], exclude_terms=[\"train\", \"validation\", \"sample\"])\n\nTRAIN_SEQ_PATH = find_train_sequences()\nTRAIN_LABELS_PATH = find_train_labels()\nTEST_SEQ_PATH = find_test_sequences()\nVAL_SEQ_PATH = find_validation_sequences()\nVAL_LABELS_PATH = find_validation_labels()\nMSA_DIR = find_dir(\"MSA\")\n\nprint(\"TRAIN_SEQ_PATH:\", TRAIN_SEQ_PATH)\nprint(\"TRAIN_LABELS_PATH:\", TRAIN_LABELS_PATH)\nprint(\"TEST_SEQ_PATH:\", TEST_SEQ_PATH)\nprint(\"VAL_SEQ_PATH:\", VAL_SEQ_PATH)\nprint(\"VAL_LABELS_PATH:\", VAL_LABELS_PATH)\nprint(\"MSA_DIR:\", MSA_DIR)\n\n# -----------------------------\n# Helpers\n# -----------------------------\ndef exists(path):\n    return path is not None and os.path.exists(path)\n\ndef extract_target_id_from_row_id(row_id: str) -> str:\n    return row_id.rsplit(\"_\", 1)[0]\n\ndef get_conformer_indices(df):\n    ks = []\n    k = 1\n    while f\"x_{k}\" in df.columns and f\"y_{k}\" in df.columns and f\"z_{k}\" in df.columns:\n        ks.append(k)\n        k += 1\n    return ks\n\ndef is_valid_xyz_block(xyz):\n    xyz = np.asarray(xyz, dtype=np.float32)\n    if not np.all(np.isfinite(xyz)):\n        return False\n    if np.max(np.abs(xyz)) > 1e10:\n        return False\n    return True\n\ndef safe_center(coords):\n    coords = np.asarray(coords, dtype=np.float32)\n    if len(coords) == 0:\n        return coords\n    return coords - coords.mean(axis=0, keepdims=True)\n\ndef clip_coords(coords):\n    return np.clip(coords, -999.999, 9999.999).astype(np.float32)\n\ndef moving_average_1d(x, w=3):\n    x = np.asarray(x, dtype=np.float32)\n    n = len(x)\n\n    if w <= 1 or n <= 2:\n        return x.copy()\n\n    # force odd window so output length matches input length\n    if w % 2 == 0:\n        w += 1\n\n    pad = w // 2\n    xp = np.pad(x, (pad, pad), mode=\"edge\")\n    kernel = np.ones(w, dtype=np.float32) / float(w)\n    y = np.convolve(xp, kernel, mode=\"valid\")\n\n    # safety guard\n    if len(y) != n:\n        y = y[:n]\n    return y\n\ndef smooth_coords(coords, w=3):\n    coords = np.asarray(coords, dtype=np.float32)\n    if len(coords) == 0:\n        return coords\n\n    out = np.zeros_like(coords)\n    for j in range(coords.shape[1]):\n        out[:, j] = moving_average_1d(coords[:, j], w)\n    return out\n\ndef arc_length_resample(coords, new_len):\n    coords = np.asarray(coords, dtype=np.float32)\n    n = len(coords)\n\n    if new_len <= 0:\n        return np.zeros((0, 3), dtype=np.float32)\n    if n == 0:\n        return np.zeros((new_len, 3), dtype=np.float32)\n    if n == 1:\n        return np.repeat(coords, new_len, axis=0)\n\n    seg = np.linalg.norm(coords[1:] - coords[:-1], axis=1)\n    cum = np.concatenate([[0.0], np.cumsum(seg)]).astype(np.float32)\n    total = float(cum[-1])\n\n    if total < 1e-8:\n        return np.repeat(coords[:1], new_len, axis=0)\n\n    t_new = np.linspace(0.0, total, new_len, dtype=np.float32)\n    out = np.zeros((new_len, 3), dtype=np.float32)\n\n    j = 0\n    for i, t in enumerate(t_new):\n        while j + 1 < len(cum) and cum[j + 1] < t:\n            j += 1\n        if j + 1 >= len(cum):\n            out[i] = coords[-1]\n        else:\n            left, right = cum[j], cum[j + 1]\n            if right - left < 1e-8:\n                out[i] = coords[j]\n            else:\n                a = (t - left) / (right - left)\n                out[i] = (1 - a) * coords[j] + a * coords[j + 1]\n    return out\n\ndef compose_to_dict(seq):\n    cnt = Counter(seq)\n    total = max(1, len(seq))\n    return {b: cnt.get(b, 0) / total for b in \"ACGU\"}\n\ndef comp_l1(a, b):\n    da, db = compose_to_dict(a), compose_to_dict(b)\n    return sum(abs(da[k] - db[k]) for k in \"ACGU\")\n\ndef sequence_identity_prefix(a, b):\n    if len(a) == 0 or len(b) == 0:\n        return 0.0\n    m = min(len(a), len(b))\n    return sum(1 for i in range(m) if a[i] == b[i]) / max(len(a), len(b))\n\ndef longest_common_kmer_score(a, b, k=5):\n    if len(a) < k or len(b) < k:\n        return 0.0\n    s1 = {a[i:i+k] for i in range(len(a) - k + 1)}\n    s2 = {b[i:i+k] for i in range(len(b) - k + 1)}\n    inter = len(s1 & s2)\n    denom = max(1, min(len(s1), len(s2)))\n    return inter / denom\n\ndef simple_edit_similarity(a, b):\n    from difflib import SequenceMatcher\n    return SequenceMatcher(None, a, b).ratio()\n\ndef estimate_step_stats(coord_list):\n    ds = []\n    for coords in coord_list:\n        if len(coords) >= 2:\n            d = np.linalg.norm(coords[1:] - coords[:-1], axis=1)\n            d = d[np.isfinite(d)]\n            d = d[(d > 0.1) & (d < 20)]\n            if len(d):\n                ds.append(d)\n    if not ds:\n        return 5.5, 0.5\n    d = np.concatenate(ds)\n    return float(d.mean()), float(d.std() + 1e-6)\n\n# -----------------------------\n# Alignment / TM-like metric\n# -----------------------------\ndef kabsch_align(P, Q):\n    P = np.asarray(P, dtype=np.float64)\n    Q = np.asarray(Q, dtype=np.float64)\n\n    Pc = P - P.mean(axis=0, keepdims=True)\n    Qc = Q - Q.mean(axis=0, keepdims=True)\n\n    H = Pc.T @ Qc\n    U, _, 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\n    P_aligned = Pc @ R\n    return P_aligned, Qc\n\ndef tm_like_score(pred, truth, cutoff_angs=30.0):\n    pred = np.asarray(pred, dtype=np.float64)\n    truth = np.asarray(truth, dtype=np.float64)\n\n    if len(pred) != len(truth) or len(pred) == 0:\n        return 0.0\n\n    P_aligned, Qc = kabsch_align(pred, truth)\n    d = np.sqrt(np.sum((P_aligned - Qc) ** 2, axis=1))\n\n    mask = d <= cutoff_angs\n    if mask.sum() == 0:\n        return 0.0\n\n    d_mask = d[mask]\n    L = len(d_mask)\n    d0 = max(0.5, 1.24 * max(L - 15, 1) ** (1/3) - 1.8)\n    return float(np.mean(1.0 / (1.0 + (d_mask / d0) ** 2)))\n\ndef best_of_five_tm_like_score(preds, truths):\n    best = -1.0\n    for p in preds:\n        for t in truths:\n            s = tm_like_score(p, t)\n            if s > best:\n                best = s\n    return best\n\n# -----------------------------\n# MSA features\n# -----------------------------\nmsa_feature_cache = {}\n\ndef parse_fasta_limited(path, max_seqs=MAX_MSA_SEQS):\n    seqs = []\n    if not exists(path):\n        return seqs\n\n    name = None\n    buf = []\n    with open(path, \"r\", encoding=\"utf-8\", errors=\"ignore\") as f:\n        for line in f:\n            line = line.strip()\n            if not line:\n                continue\n            if line.startswith(\">\"):\n                if name is not None:\n                    seqs.append((name, \"\".join(buf)))\n                    if len(seqs) >= max_seqs:\n                        break\n                name = line[1:]\n                buf = []\n            else:\n                buf.append(line)\n        if name is not None and len(seqs) < max_seqs:\n            seqs.append((name, \"\".join(buf)))\n    return seqs\n\ndef get_msa_path(target_id):\n    if MSA_DIR is None:\n        return None\n    path = os.path.join(MSA_DIR, f\"{target_id}.MSA.fasta\")\n    return path if exists(path) else None\n\ndef msa_features(target_id):\n    if target_id in msa_feature_cache:\n        return msa_feature_cache[target_id]\n\n    path = get_msa_path(target_id)\n    if path is None:\n        feat = {\n            \"msa_depth\": 0,\n            \"msa_eff_depth\": 0.0,\n            \"msa_gap_frac\": 1.0,\n            \"msa_cons_mean\": 0.0,\n            \"msa_has_data\": 0,\n        }\n        msa_feature_cache[target_id] = feat\n        return feat\n\n    rows = parse_fasta_limited(path, max_seqs=MAX_MSA_SEQS)\n    seqs = [s for _, s in rows if s]\n    if not seqs:\n        feat = {\n            \"msa_depth\": 0,\n            \"msa_eff_depth\": 0.0,\n            \"msa_gap_frac\": 1.0,\n            \"msa_cons_mean\": 0.0,\n            \"msa_has_data\": 0,\n        }\n        msa_feature_cache[target_id] = feat\n        return feat\n\n    L = len(seqs[0])\n    total_gap = 0\n    cons = []\n\n    ungapped_counter = Counter()\n    for s in seqs:\n        ungapped_counter[\"\".join(ch for ch in s if ch != \"-\")] += 1\n\n    for i in range(L):\n        counts = Counter()\n        non_gap = 0\n        for s in seqs:\n            ch = s[i] if i < len(s) else \"-\"\n            if ch == \"-\":\n                total_gap += 1\n            else:\n                counts[ch] += 1\n                non_gap += 1\n        cons.append((max(counts.values()) / non_gap) if non_gap > 0 else 0.0)\n\n    eff = float(sum(1.0 / c for c in ungapped_counter.values()))\n    gap_frac = total_gap / max(1, len(seqs) * L)\n\n    feat = {\n        \"msa_depth\": len(seqs),\n        \"msa_eff_depth\": eff,\n        \"msa_gap_frac\": float(gap_frac),\n        \"msa_cons_mean\": float(np.mean(cons)),\n        \"msa_has_data\": 1,\n    }\n    msa_feature_cache[target_id] = feat\n    return feat\n\n# -----------------------------\n# Template record\n# -----------------------------\n@dataclass\nclass TemplateRecord:\n    target_id: str\n    sequence: str\n    coords: np.ndarray\n    conformer_idx: int\n    length: int\n    msa_feat: dict\n\n# -----------------------------\n# Read data\n# -----------------------------\ntrain_seq = pd.read_csv(TRAIN_SEQ_PATH)\ntrain_labels = pd.read_csv(TRAIN_LABELS_PATH, low_memory=False)\ntest_seq = pd.read_csv(TEST_SEQ_PATH)\n\nhas_val = exists(VAL_SEQ_PATH) and exists(VAL_LABELS_PATH)\nif has_val:\n    val_seq = pd.read_csv(VAL_SEQ_PATH)\n    val_labels = pd.read_csv(VAL_LABELS_PATH, low_memory=False)\nelse:\n    val_seq = None\n    val_labels = None\n\ntrain_seq_map = dict(zip(train_seq[\"target_id\"], train_seq[\"sequence\"]))\n\nprint(\"\\nShapes:\")\nprint(\"train_seq:\", train_seq.shape)\nprint(\"train_labels:\", train_labels.shape)\nprint(\"test_seq:\", test_seq.shape)\nif has_val:\n    print(\"val_seq:\", val_seq.shape)\n    print(\"val_labels:\", val_labels.shape)\n\n# -----------------------------\n# Parse labels into target -> conformers\n# -----------------------------\ndef build_target_conformers(labels_df):\n    labels_df = labels_df.copy()\n    labels_df[\"target_id\"] = labels_df[\"ID\"].map(extract_target_id_from_row_id)\n\n    conformers = get_conformer_indices(labels_df)\n    sort_cols = [c for c in [\"target_id\", \"copy\", \"chain\", \"resid\"] if c in labels_df.columns]\n    labels_df = labels_df.sort_values(sort_cols).reset_index(drop=True)\n\n    target_to_truths = {}\n    target_to_resnames = {}\n\n    for tid, g in labels_df.groupby(\"target_id\", sort=False):\n        g = g.sort_values([c for c in [\"copy\", \"chain\", \"resid\"] if c in g.columns]).reset_index(drop=True)\n        target_to_resnames[tid] = \"\".join(g[\"resname\"].astype(str).tolist())\n\n        truths = []\n        for k in conformers:\n            xyz = g[[f\"x_{k}\", f\"y_{k}\", f\"z_{k}\"]].values.astype(np.float32)\n            valid_rows = np.array([is_valid_xyz_block(v) for v in xyz], dtype=bool)\n            if valid_rows.all():\n                truths.append(safe_center(xyz))\n        target_to_truths[tid] = truths\n\n    return target_to_truths, target_to_resnames\n\ntrain_truths, train_resnames = build_target_conformers(train_labels)\nif has_val:\n    val_truths, val_resnames = build_target_conformers(val_labels)\n\n# -----------------------------\n# Build template bank\n# -----------------------------\ntemplate_bank = []\nfor tid, truths in train_truths.items():\n    seq = train_seq_map.get(tid, train_resnames.get(tid, \"\"))\n    feat = msa_features(tid)\n    for ci, coords in enumerate(truths, start=1):\n        template_bank.append(\n            TemplateRecord(\n                target_id=tid,\n                sequence=seq,\n                coords=np.asarray(coords, dtype=np.float32),\n                conformer_idx=ci,\n                length=len(coords),\n                msa_feat=feat,\n            )\n        )\n\nMEAN_STEP, STEP_STD = estimate_step_stats([t.coords for t in template_bank])\n\nprint(f\"\\nTrain targets: {len(train_truths)}\")\nprint(f\"Template conformers: {len(template_bank)}\")\nprint(f\"Mean step length: {MEAN_STEP:.3f}, std: {STEP_STD:.3f}\")\n\n# -----------------------------\n# Retrieval\n# -----------------------------\ndef template_score(query_seq, query_msa, templ: TemplateRecord):\n    tseq = templ.sequence\n    qlen = len(query_seq)\n    tlen = len(tseq)\n\n    len_ratio = min(qlen, tlen) / max(qlen, tlen)\n    sim = simple_edit_similarity(query_seq, tseq)\n    pos = sequence_identity_prefix(query_seq, tseq)\n    k = min(5, max(3, min(qlen, tlen) // 6 if min(qlen, tlen) > 0 else 3))\n    kmer = longest_common_kmer_score(query_seq, tseq, k=k)\n    comp = 1.0 - 0.5 * comp_l1(query_seq, tseq)\n\n    tm = templ.msa_feat\n    msa_prior = 0.0\n    if tm[\"msa_has_data\"]:\n        msa_prior += 0.02 * np.tanh(tm[\"msa_eff_depth\"] / 30.0)\n        msa_prior += 0.015 * tm[\"msa_cons_mean\"]\n        msa_prior -= 0.01 * tm[\"msa_gap_frac\"]\n\n    query_bonus = 0.0\n    if query_msa[\"msa_has_data\"]:\n        query_bonus += 0.01 * np.tanh(query_msa[\"msa_eff_depth\"] / 30.0)\n        query_bonus += 0.005 * query_msa[\"msa_cons_mean\"]\n\n    return float(\n        0.52 * sim\n        + 0.20 * pos\n        + 0.16 * len_ratio\n        + 0.08 * kmer\n        + 0.04 * comp\n        + msa_prior\n        + 0.5 * query_bonus\n    )\n\ndef get_ranked_templates(query_target_id, query_seq, topn=MAX_TEMPLATES_FOR_RERANK):\n    qmsa = msa_features(query_target_id)\n    scored = [(template_score(query_seq, qmsa, t), t) for t in template_bank]\n    scored.sort(key=lambda x: x[0], reverse=True)\n    return [t for _, t in scored[:topn]], qmsa\n\ndef pick_template_buckets(query_target_id, query_seq, topn=MAX_TEMPLATES_FOR_RERANK):\n    ranked, qmsa = get_ranked_templates(query_target_id, query_seq, topn=topn)\n    if not ranked:\n        return [], qmsa\n\n    best = ranked[0]\n    best_len = max(ranked, key=lambda t: min(len(query_seq), t.length) / max(len(query_seq), t.length))\n    best_msa = max(\n        ranked,\n        key=lambda t: (t.msa_feat.get(\"msa_eff_depth\", 0.0), t.msa_feat.get(\"msa_cons_mean\", 0.0))\n    )\n    alt = ranked[min(4, len(ranked) - 1)]\n\n    uniq = []\n    seen = set()\n    for t in [best, best_len, best_msa, alt] + ranked:\n        key = (t.target_id, t.conformer_idx)\n        if key not in seen:\n            uniq.append(t)\n            seen.add(key)\n\n    return uniq[:6], qmsa\n\n# -----------------------------\n# Geometry\n# -----------------------------\ndef normalize_step(coords, target_mean_step):\n    coords = np.asarray(coords, dtype=np.float32)\n    if len(coords) < 2:\n        return coords\n    d = np.linalg.norm(coords[1:] - coords[:-1], axis=1)\n    cur = float(np.mean(d))\n    if cur < 1e-8:\n        return coords\n    return coords * (target_mean_step / cur)\n\ndef coords_from_template(templ_coords, target_len):\n    out = arc_length_resample(templ_coords, target_len)\n    out = safe_center(out)\n    out = smooth_coords(out, w=3)\n    out = normalize_step(out, MEAN_STEP)\n    out = safe_center(out)\n    return out.astype(np.float32)\n\ndef blend_templates(templ_list, weights, target_len):\n    weights = np.asarray(weights, dtype=np.float32)\n    weights = weights / max(1e-8, weights.sum())\n\n    coords_list = [coords_from_template(t.coords, target_len) for t in templ_list]\n    out = np.zeros((target_len, 3), dtype=np.float32)\n    for w, c in zip(weights, coords_list):\n        out += w * c\n\n    out = smooth_coords(out, w=3)\n    out = normalize_step(out, MEAN_STEP)\n    out = safe_center(out)\n    return out\n\ndef fallback_shape(seq_len):\n    if seq_len <= 0:\n        return np.zeros((0, 3), dtype=np.float32)\n    r = 8.5\n    rise = 2.0\n    turn = 0.72\n    coords = np.zeros((seq_len, 3), dtype=np.float32)\n    for i in range(seq_len):\n        a = i * turn\n        coords[i] = [r * np.cos(a), r * np.sin(a), i * rise]\n    coords = normalize_step(coords, MEAN_STEP)\n    return safe_center(coords)\n\ndef compactify(coords, target_radius_scale=1.15):\n    coords = np.asarray(coords, dtype=np.float32)\n    if len(coords) == 0:\n        return coords\n\n    centered = coords - coords.mean(axis=0, keepdims=True)\n    r = np.linalg.norm(centered, axis=1)\n    med = np.median(r) + 1e-6\n    limit = target_radius_scale * med\n\n    scale = np.ones_like(r)\n    far = r > limit\n    scale[far] = limit / r[far]\n\n    out = centered * scale[:, None]\n    return out.astype(np.float32)\n\ndef perturb_structure(coords, idx, msa_q=None):\n    c = np.asarray(coords, dtype=np.float32).copy()\n    n = len(c)\n    if n == 0:\n        return c\n\n    msa_strength = 0.0\n    if msa_q is not None and msa_q[\"msa_has_data\"]:\n        msa_strength = np.tanh(msa_q[\"msa_eff_depth\"] / 40.0)\n\n    base_noise = STEP_STD * (0.09 - 0.03 * msa_strength)\n    base_noise = max(0.012, base_noise)\n\n    t = np.linspace(0, 2 * np.pi, n, dtype=np.float32)\n\n    if idx == 0:\n        out = c\n    elif idx == 1:\n        out = smooth_coords(c, w=3)\n    elif idx == 2:\n        out = c.copy()\n    elif idx == 3:\n        out = c.copy()\n        out[:, 0] += 0.18 * base_noise * np.sin(t)\n        out[:, 1] += 0.18 * base_noise * np.cos(2 * t)\n        out[:, 2] += 0.10 * base_noise * np.sin(3 * t)\n    else:\n        out = c.copy()\n        noise = np.random.normal(0.0, base_noise, size=(n, 3)).astype(np.float32)\n        noise = np.cumsum(noise, axis=0)\n        noise -= noise.mean(axis=0, keepdims=True)\n        out += 0.16 * noise\n        out = smooth_coords(out, w=3)\n\n    out = normalize_step(out, MEAN_STEP)\n    out = compactify(out, target_radius_scale=1.15)\n    out = safe_center(out)\n    out = clip_coords(out)\n    return out\n\n# -----------------------------\n# Prediction\n# -----------------------------\ndef predict_target(target_id, sequence):\n    buckets, qmsa = pick_template_buckets(target_id, sequence, topn=MAX_TEMPLATES_FOR_RERANK)\n    L = len(sequence)\n\n    if not buckets:\n        base = fallback_shape(L)\n        return [perturb_structure(base, i, qmsa) for i in range(N_PREDS)]\n\n    blend_candidates = buckets[:min(TOP_TEMPLATES_TO_BLEND, len(buckets))]\n    raw_scores = np.array([template_score(sequence, qmsa, t) for t in blend_candidates], dtype=np.float32)\n    x = raw_scores - raw_scores.max()\n    weights = np.exp(6.0 * x)\n    weights /= weights.sum()\n    base_blend = blend_templates(blend_candidates, weights, L)\n\n    preds = []\n    preds.append(perturb_structure(base_blend, 0, qmsa))\n    preds.append(perturb_structure(coords_from_template(buckets[0].coords, L), 1, qmsa))\n\n    t3 = buckets[1] if len(buckets) > 1 else buckets[0]\n    preds.append(perturb_structure(coords_from_template(t3.coords, L), 2, qmsa))\n\n    t4 = buckets[2] if len(buckets) > 2 else buckets[0]\n    preds.append(perturb_structure(coords_from_template(t4.coords, L), 3, qmsa))\n\n    if len(blend_candidates) >= 3:\n        alt_w = np.array([0.55, 0.30, 0.15], dtype=np.float32)\n        alt_c = blend_templates(blend_candidates[:3], alt_w, L)\n    elif len(blend_candidates) >= 2:\n        alt_w = np.array([0.70, 0.30], dtype=np.float32)\n        alt_c = blend_templates(blend_candidates[:2], alt_w, L)\n    else:\n        alt_c = coords_from_template(blend_candidates[0].coords, L)\n    preds.append(perturb_structure(alt_c, 4, qmsa))\n\n    return preds[:N_PREDS]\n\n# -----------------------------\n# Submission\n# -----------------------------\ndef make_submission_df(test_df):\n    rows = []\n    for _, row in test_df.iterrows():\n        tid = row[\"target_id\"]\n        seq = row[\"sequence\"]\n        preds = predict_target(tid, seq)\n\n        for resid, resname in enumerate(seq, start=1):\n            rec = {\n                \"ID\": f\"{tid}_{resid}\",\n                \"resname\": resname,\n                \"resid\": resid,\n            }\n            for k in range(N_PREDS):\n                rec[f\"x_{k+1}\"] = float(preds[k][resid - 1, 0])\n                rec[f\"y_{k+1}\"] = float(preds[k][resid - 1, 1])\n                rec[f\"z_{k+1}\"] = float(preds[k][resid - 1, 2])\n            rows.append(rec)\n\n    cols = [\"ID\", \"resname\", \"resid\"]\n    for k in range(1, N_PREDS + 1):\n        cols += [f\"x_{k}\", f\"y_{k}\", f\"z_{k}\"]\n    return pd.DataFrame(rows)[cols]\n\n# -----------------------------\n# Optional validation\n# -----------------------------\ndef evaluate_on_validation(val_seq_df, val_truths_dict):\n    scores = []\n    details = []\n\n    for _, row in val_seq_df.iterrows():\n        tid = row[\"target_id\"]\n        seq = row[\"sequence\"]\n        truths = val_truths_dict.get(tid, [])\n        if not truths:\n            continue\n\n        preds = predict_target(tid, seq)\n        score = best_of_five_tm_like_score(preds, truths)\n        scores.append(score)\n        details.append({\n            \"target_id\": tid,\n            \"length\": len(seq),\n            \"n_truth_conformers\": len(truths),\n            \"tm_like_proxy\": score,\n        })\n\n    if not details:\n        return float(\"nan\"), pd.DataFrame()\n\n    detail_df = pd.DataFrame(details).sort_values(\"tm_like_proxy\", ascending=False).reset_index(drop=True)\n    mean_score = float(np.mean(scores)) if scores else float(\"nan\")\n    return mean_score, detail_df\n\nif RUN_VALIDATION and has_val:\n    print(\"\\nRunning TM-like validation...\")\n    val_mean, val_detail = evaluate_on_validation(val_seq, val_truths)\n    print(\"Validation TM-like proxy:\", round(val_mean, 6))\n    print(val_detail.head(10))\nelse:\n    print(\"\\nSkipping validation.\")\n\nsubmission = make_submission_df(test_seq)\nsubmission.to_csv(\"submission.csv\", index=False)\n\nprint(\"\\nSaved submission.csv\")\nprint(submission.head())\nprint(\"Submission shape:\", submission.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-25T04:30:54.070044Z","iopub.execute_input":"2026-03-25T04:30:54.07048Z","iopub.status.idle":"2026-03-25T04:38:33.073813Z","shell.execute_reply.started":"2026-03-25T04:30:54.070435Z","shell.execute_reply":"2026-03-25T04:38:33.07253Z"}},"outputs":[],"execution_count":null}]}