{"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":31286,"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},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nbase = \"/kaggle/input/competitions/stanford-rna-3d-folding-2\"\nprint(os.listdir(base))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-18T12:35:30.210118Z","iopub.execute_input":"2026-03-18T12:35:30.210375Z","iopub.status.idle":"2026-03-18T12:35:30.2217Z","shell.execute_reply.started":"2026-03-18T12:35:30.210347Z","shell.execute_reply":"2026-03-18T12:35:30.220561Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Imports & Global Config","metadata":{}},{"cell_type":"code","source":"from __future__ import annotations\n\nimport re\nimport warnings\nfrom dataclasses import dataclass\nfrom pathlib import Path\nfrom typing import Dict, List, Tuple, Optional\n\nimport numpy as np\nimport pandas as pd\n\nfrom scipy.spatial.transform import Rotation\nfrom sklearn.feature_extraction.text import TfidfVectorizer\nfrom sklearn.metrics.pairwise import cosine_similarity\n\n\nDATA_ROOT = Path(\"/kaggle/input/competitions/stanford-rna-3d-folding-2\")\n\nTRAIN_SEQ_PATH = DATA_ROOT / \"train_sequences.csv\"\nTRAIN_LABEL_PATH = DATA_ROOT / \"train_labels.csv\"\nTEST_SEQ_PATH = DATA_ROOT / \"test_sequences.csv\"\nSAMPLE_SUB_PATH = DATA_ROOT / \"sample_submission.csv\"\n\nRANDOM_SEED = 42\nRNG = np.random.default_rng(RANDOM_SEED)\n\nKMER_SIZE = 3\nTOP_K_RETRIEVAL = 20\nTOP_N_FINAL = 5\nNOISE_STD = 0.01\nINTERPOLATE_MISSING = True\n\nEXPECTED_SUBMISSION_COLUMNS = [\n    \"ID\", \"resname\", \"resid\",\n    \"x_1\", \"y_1\", \"z_1\",\n    \"x_2\", \"y_2\", \"z_2\",\n    \"x_3\", \"y_3\", \"z_3\",\n    \"x_4\", \"y_4\", \"z_4\",\n    \"x_5\", \"y_5\", \"z_5\",\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-18T12:35:36.239197Z","iopub.execute_input":"2026-03-18T12:35:36.240331Z","iopub.status.idle":"2026-03-18T12:35:37.682939Z","shell.execute_reply.started":"2026-03-18T12:35:36.240288Z","shell.execute_reply":"2026-03-18T12:35:37.681868Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Data classes","metadata":{}},{"cell_type":"code","source":"@dataclass\nclass RetrievalCandidate:\n    train_target_id: str\n    train_sequence: str\n    cosine_score: float\n\n\n@dataclass\nclass RankedTemplate:\n    train_target_id: str\n    train_label_id: Optional[str]\n    train_sequence: str\n    cosine_score: float\n    alignment_score: float\n    valid_ratio: float\n    status: str","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-18T12:35:40.246829Z","iopub.execute_input":"2026-03-18T12:35:40.248195Z","iopub.status.idle":"2026-03-18T12:35:40.256689Z","shell.execute_reply.started":"2026-03-18T12:35:40.248158Z","shell.execute_reply":"2026-03-18T12:35:40.255322Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Generic Utilities","metadata":{}},{"cell_type":"code","source":"def infer_target_id_column(df: pd.DataFrame) -> str:\n    candidates = [\"target_id\", \"target\", \"id\", \"ID\"]\n    for col in candidates:\n        if col in df.columns:\n            return col\n    raise KeyError(f\"Could not infer target id column. Available columns: {list(df.columns)}\")\n\n\ndef infer_sequence_column(df: pd.DataFrame) -> str:\n    candidates = [\"sequence\", \"seq\", \"Sequence\", \"SEQ\"]\n    for col in candidates:\n        if col in df.columns:\n            return col\n    raise KeyError(f\"Could not infer sequence column. Available columns: {list(df.columns)}\")\n\n\ndef infer_coordinate_columns(df: pd.DataFrame) -> Tuple[str, str, str]:\n    candidates = [\n        (\"x_1\", \"y_1\", \"z_1\"),\n        (\"x\", \"y\", \"z\"),\n        (\"X\", \"Y\", \"Z\"),\n    ]\n    for cols in candidates:\n        if all(c in df.columns for c in cols):\n            return cols\n    raise KeyError(f\"Could not infer coordinate columns. Available columns: {list(df.columns)}\")\n\n\ndef infer_resid_column(df: pd.DataFrame) -> str:\n    candidates = [\"resid\", \"residue_index\", \"position\", \"idx\"]\n    for col in candidates:\n        if col in df.columns:\n            return col\n    raise KeyError(f\"Could not infer resid column. Available columns: {list(df.columns)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-18T12:35:42.334122Z","iopub.execute_input":"2026-03-18T12:35:42.334448Z","iopub.status.idle":"2026-03-18T12:35:42.343945Z","shell.execute_reply.started":"2026-03-18T12:35:42.334419Z","shell.execute_reply":"2026-03-18T12:35:42.342647Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Step 1: Robust Data Parser","metadata":{}},{"cell_type":"code","source":"def parse_sequences_csv(csv_path: Path) -> pd.DataFrame:\n    \"\"\"\n    Load and standardize sequence CSV into columns:\n        ['target_id', 'sequence']\n    \"\"\"\n    if not csv_path.exists():\n        raise FileNotFoundError(f\"Sequence file not found: {csv_path}\")\n\n    df = pd.read_csv(csv_path)\n    target_col = infer_target_id_column(df)\n    seq_col = infer_sequence_column(df)\n\n    out = df[[target_col, seq_col]].copy()\n    out.columns = [\"target_id\", \"sequence\"]\n    out[\"target_id\"] = out[\"target_id\"].astype(str)\n    out[\"sequence\"] = out[\"sequence\"].astype(str)\n\n    if out[\"target_id\"].duplicated().any():\n        dup_ids = out.loc[out[\"target_id\"].duplicated(), \"target_id\"].tolist()[:10]\n        raise ValueError(f\"Duplicate target_id values found in {csv_path}: {dup_ids}\")\n\n    return out\n\n\ndef parse_labels_to_dict(csv_path: Path) -> Dict[str, np.ndarray]:\n    \"\"\"\n    Parse residue-level labels into:\n        Dict[target_id, np.ndarray of shape (L, 3)]\n\n    Rule:\n    - If target_id column exists, use it directly.\n    - Otherwise extract target_id from ID by removing the final underscore suffix.\n      Example: '157D_1' -> '157D'\n    \"\"\"\n    if not csv_path.exists():\n        raise FileNotFoundError(f\"Label file not found: {csv_path}\")\n\n    df = pd.read_csv(csv_path, low_memory=False)\n\n    resid_col = infer_resid_column(df)\n    x_col, y_col, z_col = infer_coordinate_columns(df)\n\n    df = df.copy()\n    df[resid_col] = pd.to_numeric(df[resid_col], errors=\"coerce\")\n\n    if df[resid_col].isna().any():\n        bad_rows = df[df[resid_col].isna()].head()\n        raise ValueError(f\"Non-numeric resid values found in {csv_path}. Example rows:\\n{bad_rows}\")\n\n    if \"target_id\" in df.columns:\n        df[\"target_id_parsed\"] = df[\"target_id\"].astype(str)\n    elif \"ID\" in df.columns:\n        df[\"ID\"] = df[\"ID\"].astype(str)\n        df[\"target_id_parsed\"] = df[\"ID\"].str.replace(r\"_[^_]+$\", \"\", regex=True)\n    else:\n        raise KeyError(f\"Neither 'target_id' nor 'ID' found in {csv_path}\")\n\n    coords_dict: Dict[str, np.ndarray] = {}\n\n    for target_id, g in df.groupby(\"target_id_parsed\", sort=False):\n        g = g.sort_values(resid_col).reset_index(drop=True)\n\n        coords = g[[x_col, y_col, z_col]].to_numpy(dtype=np.float64)\n\n        if coords.ndim != 2 or coords.shape[1] != 3:\n            raise ValueError(\n                f\"Coordinates for target_id={target_id} do not have shape (L, 3). Got {coords.shape}\"\n            )\n\n        coords_dict[str(target_id)] = coords\n\n    return coords_dict\n\n\ndef build_sequence_to_label_id_map(\n    seq_df: pd.DataFrame,\n    label_dict: Dict[str, np.ndarray],\n) -> Dict[str, str]:\n    \"\"\"\n    Map sequence-level target_id to one label-level ID.\n\n    Strategy:\n    1) exact match if exists\n    2) target_id_1 if exists\n    3) lexicographically smallest target_id_* match\n    \"\"\"\n    if not {\"target_id\", \"sequence\"}.issubset(seq_df.columns):\n        raise ValueError(\"seq_df must contain ['target_id', 'sequence'].\")\n\n    label_keys = list(label_dict.keys())\n    mapping: Dict[str, str] = {}\n\n    for target_id in seq_df[\"target_id\"].astype(str).tolist():\n        if target_id in label_dict:\n            mapping[target_id] = target_id\n            continue\n\n        if f\"{target_id}_1\" in label_dict:\n            mapping[target_id] = f\"{target_id}_1\"\n            continue\n\n        prefix_matches = sorted([k for k in label_keys if k.startswith(f\"{target_id}_\")])\n        if prefix_matches:\n            mapping[target_id] = prefix_matches[0]\n\n    return mapping","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-18T12:35:44.361421Z","iopub.execute_input":"2026-03-18T12:35:44.362059Z","iopub.status.idle":"2026-03-18T12:35:44.382023Z","shell.execute_reply.started":"2026-03-18T12:35:44.36202Z","shell.execute_reply":"2026-03-18T12:35:44.380702Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"TF-IDF k-mer Retriever","metadata":{}},{"cell_type":"code","source":"def sequence_to_kmers(seq: str, k: int = 3) -> str:\n    \"\"\"\n    Convert sequence into whitespace-separated k-mer tokens.\n    Example: AUGCU -> 'AUG UGC GCU'\n    \"\"\"\n    if not isinstance(seq, str):\n        raise TypeError(\"Input sequence must be a string.\")\n    if k <= 0:\n        raise ValueError(\"k must be positive.\")\n\n    if len(seq) < k:\n        return seq\n\n    return \" \".join(seq[i:i+k] for i in range(len(seq) - k + 1))\n\n\nclass TfidfKmerRetriever:\n    def __init__(self, kmer_size: int = 3) -> None:\n        self.kmer_size = kmer_size\n        self.vectorizer: Optional[TfidfVectorizer] = None\n        self.train_matrix = None\n        self.train_target_ids: List[str] = []\n        self.train_sequences: List[str] = []\n\n    def fit(self, train_seq_df: pd.DataFrame) -> \"TfidfKmerRetriever\":\n        if not {\"target_id\", \"sequence\"}.issubset(train_seq_df.columns):\n            raise ValueError(\"train_seq_df must contain ['target_id', 'sequence'].\")\n\n        self.train_target_ids = train_seq_df[\"target_id\"].astype(str).tolist()\n        self.train_sequences = train_seq_df[\"sequence\"].astype(str).tolist()\n\n        corpus = [sequence_to_kmers(seq, self.kmer_size) for seq in self.train_sequences]\n        self.vectorizer = TfidfVectorizer(analyzer=\"word\")\n        self.train_matrix = self.vectorizer.fit_transform(corpus)\n        return self\n\n    def query_top_k(self, query_sequence: str, k: int = 20) -> List[RetrievalCandidate]:\n        if self.vectorizer is None or self.train_matrix is None:\n            raise RuntimeError(\"Retriever must be fit before querying.\")\n\n        if k <= 0:\n            raise ValueError(\"k must be positive.\")\n\n        try:\n            query_text = sequence_to_kmers(query_sequence, self.kmer_size)\n            query_vec = self.vectorizer.transform([query_text])\n            sims = cosine_similarity(query_vec, self.train_matrix)[0]\n\n            top_idx = np.argsort(-sims)[:k]\n            candidates = [\n                RetrievalCandidate(\n                    train_target_id=self.train_target_ids[i],\n                    train_sequence=self.train_sequences[i],\n                    cosine_score=float(sims[i]),\n                )\n                for i in top_idx\n            ]\n            return candidates\n        except Exception as e:\n            warnings.warn(f\"TF-IDF retrieval failed: {e}\")\n            return []","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-18T12:35:47.313506Z","iopub.execute_input":"2026-03-18T12:35:47.314264Z","iopub.status.idle":"2026-03-18T12:35:47.32679Z","shell.execute_reply.started":"2026-03-18T12:35:47.314195Z","shell.execute_reply":"2026-03-18T12:35:47.32586Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Needleman-Wunsch Alignment","metadata":{}},{"cell_type":"code","source":"def needleman_wunsch(\n    seq1: str,\n    seq2: str,\n    match_score: int = 2,\n    mismatch_score: int = -1,\n    gap_score: int = -2,\n) -> Tuple[str, str, int]:\n    \"\"\"\n    Standard global alignment using Needleman-Wunsch.\n    Returns:\n        aligned_seq1, aligned_seq2, alignment_score\n    \"\"\"\n    if not isinstance(seq1, str) or not isinstance(seq2, str):\n        raise TypeError(\"Both seq1 and seq2 must be strings.\")\n\n    n = len(seq1)\n    m = len(seq2)\n\n    score = np.zeros((n + 1, m + 1), dtype=np.int32)\n    traceback = np.zeros((n + 1, m + 1), dtype=np.int8)\n    # 0 = diag, 1 = up, 2 = left\n\n    for i in range(1, n + 1):\n        score[i, 0] = score[i - 1, 0] + gap_score\n        traceback[i, 0] = 1\n\n    for j in range(1, m + 1):\n        score[0, j] = score[0, j - 1] + gap_score\n        traceback[0, j] = 2\n\n    for i in range(1, n + 1):\n        c1 = seq1[i - 1]\n        for j in range(1, m + 1):\n            c2 = seq2[j - 1]\n            diag = score[i - 1, j - 1] + (match_score if c1 == c2 else mismatch_score)\n            up = score[i - 1, j] + gap_score\n            left = score[i, j - 1] + gap_score\n\n            best = max(diag, up, left)\n            score[i, j] = best\n\n            if best == diag:\n                traceback[i, j] = 0\n            elif best == up:\n                traceback[i, j] = 1\n            else:\n                traceback[i, j] = 2\n\n    aligned1: List[str] = []\n    aligned2: List[str] = []\n    i, j = n, m\n\n    while i > 0 or j > 0:\n        if i > 0 and j > 0 and traceback[i, j] == 0:\n            aligned1.append(seq1[i - 1])\n            aligned2.append(seq2[j - 1])\n            i -= 1\n            j -= 1\n        elif i > 0 and (j == 0 or traceback[i, j] == 1):\n            aligned1.append(seq1[i - 1])\n            aligned2.append(\"-\")\n            i -= 1\n        else:\n            aligned1.append(\"-\")\n            aligned2.append(seq2[j - 1])\n            j -= 1\n\n    aligned1.reverse()\n    aligned2.reverse()\n\n    return \"\".join(aligned1), \"\".join(aligned2), int(score[n, m])\n\n\ndef aligned_valid_ratio(aligned_query: str, aligned_template: str) -> float:\n    \"\"\"\n    Ratio of query residues that can receive real template coordinates.\n    \"\"\"\n    if len(aligned_query) != len(aligned_template):\n        raise ValueError(\"Aligned strings must have the same length.\")\n\n    query_res_count = sum(ch != \"-\" for ch in aligned_query)\n    if query_res_count == 0:\n        return 0.0\n\n    valid = sum((q != \"-\" and t != \"-\") for q, t in zip(aligned_query, aligned_template))\n    return float(valid / query_res_count)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-18T12:35:51.673817Z","iopub.execute_input":"2026-03-18T12:35:51.674309Z","iopub.status.idle":"2026-03-18T12:35:51.690794Z","shell.execute_reply.started":"2026-03-18T12:35:51.674267Z","shell.execute_reply":"2026-03-18T12:35:51.689789Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Coordinate Mapping, Interpolation, Noise Fill","metadata":{}},{"cell_type":"code","source":"def map_template_coords_to_query(\n    aligned_query: str,\n    aligned_template: str,\n    template_coords: np.ndarray,\n    fill_value: float = np.nan,\n) -> np.ndarray:\n    \"\"\"\n    Map template coordinates to query positions.\n\n    Rules:\n    - query residue + template residue -> copy template coord\n    - query residue + template gap     -> fill with NaN\n    - query gap     + template residue -> discard template coord\n\n    Output shape must be (L_query, 3)\n    \"\"\"\n    if len(aligned_query) != len(aligned_template):\n        raise ValueError(\"Aligned query and template strings must have same length.\")\n\n    template_coords = np.asarray(template_coords, dtype=np.float64)\n    if template_coords.ndim != 2 or template_coords.shape[1] != 3:\n        raise ValueError(f\"template_coords must have shape (L, 3), got {template_coords.shape}\")\n\n    mapped_coords: List[np.ndarray] = []\n    template_ptr = 0\n\n    for q_char, t_char in zip(aligned_query, aligned_template):\n        q_has_res = q_char != \"-\"\n        t_has_res = t_char != \"-\"\n\n        if q_has_res and t_has_res:\n            if template_ptr >= len(template_coords):\n                raise IndexError(\n                    f\"Template pointer out of range during coordinate mapping: ptr={template_ptr}, len={len(template_coords)}\"\n                )\n            mapped_coords.append(template_coords[template_ptr].copy())\n            template_ptr += 1\n\n        elif q_has_res and not t_has_res:\n            mapped_coords.append(np.array([fill_value, fill_value, fill_value], dtype=np.float64))\n\n        elif not q_has_res and t_has_res:\n            template_ptr += 1\n\n        else:\n            raise ValueError(\"Invalid alignment state: both characters are gaps.\")\n\n    if template_ptr != len(template_coords):\n        raise ValueError(\n            f\"Template pointer mismatch after mapping: ptr={template_ptr}, template_len={len(template_coords)}\"\n        )\n\n    out = np.vstack(mapped_coords) if mapped_coords else np.empty((0, 3), dtype=np.float64)\n    expected_len = sum(ch != \"-\" for ch in aligned_query)\n\n    if out.shape[0] != expected_len:\n        raise ValueError(f\"Mapped coords length mismatch: got {out.shape[0]}, expected {expected_len}\")\n\n    return out\n\n\ndef valid_row_ratio(coords: np.ndarray) -> float:\n    coords = np.asarray(coords, dtype=np.float64)\n    if coords.ndim != 2 or coords.shape[1] != 3:\n        raise ValueError(f\"coords must have shape (L, 3), got {coords.shape}\")\n\n    if coords.shape[0] == 0:\n        return 0.0\n\n    valid_mask = np.all(np.isfinite(coords), axis=1)\n    return float(valid_mask.mean())\n\n\ndef interpolate_nan_coords(coords: np.ndarray) -> np.ndarray:\n    \"\"\"\n    Linear interpolation per axis.\n    Leading/trailing NaNs use nearest valid values.\n    If all rows are NaN, return unchanged.\n    \"\"\"\n    coords = np.asarray(coords, dtype=np.float64)\n    if coords.ndim != 2 or coords.shape[1] != 3:\n        raise ValueError(f\"coords must have shape (L, 3), got {coords.shape}\")\n\n    out = coords.copy()\n    n = out.shape[0]\n    idx = np.arange(n)\n\n    if n == 0:\n        return out\n\n    all_nan_rows = np.all(np.isnan(out), axis=1)\n    if np.all(all_nan_rows):\n        return out\n\n    for dim in range(3):\n        y = out[:, dim]\n        valid = np.isfinite(y)\n        if valid.sum() == 0:\n            continue\n        out[:, dim] = np.interp(idx, idx[valid], y[valid])\n\n    return out\n\n\ndef add_gaussian_noise(\n    coords: np.ndarray,\n    std: float = 0.01,\n    rng: Optional[np.random.Generator] = None,\n) -> np.ndarray:\n    coords = np.asarray(coords, dtype=np.float64)\n    if coords.ndim != 2 or coords.shape[1] != 3:\n        raise ValueError(f\"coords must have shape (L, 3), got {coords.shape}\")\n\n    if rng is None:\n        rng = np.random.default_rng()\n\n    noise = rng.normal(loc=0.0, scale=std, size=coords.shape)\n    return coords + noise","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-18T12:35:55.051667Z","iopub.execute_input":"2026-03-18T12:35:55.052015Z","iopub.status.idle":"2026-03-18T12:35:55.070822Z","shell.execute_reply.started":"2026-03-18T12:35:55.051987Z","shell.execute_reply":"2026-03-18T12:35:55.069242Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Top-K Re-ranking","metadata":{}},{"cell_type":"code","source":"def rerank_templates_with_alignment(\n    query_sequence: str,\n    coarse_candidates: List[RetrievalCandidate],\n    train_label_id_map: Dict[str, str],\n    train_coords_dict: Dict[str, np.ndarray],\n    match_score: int = 2,\n    mismatch_score: int = -1,\n    gap_score: int = -2,\n) -> List[RankedTemplate]:\n    \"\"\"\n    Re-rank Top-K retrieved templates using:\n    1) valid_ratio descending\n    2) alignment_score descending\n    \"\"\"\n    ranked: List[RankedTemplate] = []\n\n    for cand in coarse_candidates:\n        train_label_id = train_label_id_map.get(cand.train_target_id)\n        if train_label_id is None:\n            ranked.append(\n                RankedTemplate(\n                    train_target_id=cand.train_target_id,\n                    train_label_id=None,\n                    train_sequence=cand.train_sequence,\n                    cosine_score=cand.cosine_score,\n                    alignment_score=float(\"-inf\"),\n                    valid_ratio=0.0,\n                    status=\"missing_label_mapping\",\n                )\n            )\n            continue\n\n        template_coords = train_coords_dict.get(train_label_id)\n        if template_coords is None:\n            ranked.append(\n                RankedTemplate(\n                    train_target_id=cand.train_target_id,\n                    train_label_id=train_label_id,\n                    train_sequence=cand.train_sequence,\n                    cosine_score=cand.cosine_score,\n                    alignment_score=float(\"-inf\"),\n                    valid_ratio=0.0,\n                    status=\"missing_template_coords\",\n                )\n            )\n            continue\n\n        try:\n            aligned_q, aligned_t, aln_score = needleman_wunsch(\n                query_sequence,\n                cand.train_sequence,\n                match_score=match_score,\n                mismatch_score=mismatch_score,\n                gap_score=gap_score,\n            )\n            vr = aligned_valid_ratio(aligned_q, aligned_t)\n\n            ranked.append(\n                RankedTemplate(\n                    train_target_id=cand.train_target_id,\n                    train_label_id=train_label_id,\n                    train_sequence=cand.train_sequence,\n                    cosine_score=cand.cosine_score,\n                    alignment_score=float(aln_score),\n                    valid_ratio=float(vr),\n                    status=\"ok\",\n                )\n            )\n        except Exception as e:\n            warnings.warn(f\"Alignment re-ranking failed for template {cand.train_target_id}: {e}\")\n            ranked.append(\n                RankedTemplate(\n                    train_target_id=cand.train_target_id,\n                    train_label_id=train_label_id,\n                    train_sequence=cand.train_sequence,\n                    cosine_score=cand.cosine_score,\n                    alignment_score=float(\"-inf\"),\n                    valid_ratio=0.0,\n                    status=f\"alignment_failed:{e}\",\n                )\n            )\n\n    ranked_sorted = sorted(\n        ranked,\n        key=lambda x: (x.valid_ratio, x.alignment_score),\n        reverse=True,\n    )\n    return ranked_sorted","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-18T12:35:58.37513Z","iopub.execute_input":"2026-03-18T12:35:58.375495Z","iopub.status.idle":"2026-03-18T12:35:58.391564Z","shell.execute_reply.started":"2026-03-18T12:35:58.375465Z","shell.execute_reply":"2026-03-18T12:35:58.389357Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Build One Structure and Top-5 Structures","metadata":{}},{"cell_type":"code","source":"def build_single_structure_from_template(\n    query_sequence: str,\n    template_sequence: str,\n    template_coords: np.ndarray,\n    match_score: int = 2,\n    mismatch_score: int = -1,\n    gap_score: int = -2,\n    interpolate_missing: bool = True,\n) -> Tuple[Optional[np.ndarray], float, str]:\n\n    try:\n        if template_coords.shape[0] != len(template_sequence):\n            return None, 0.0, f\"template_length_mismatch:{template_coords.shape[0]}_vs_{len(template_sequence)}\"\n\n        aligned_q, aligned_t, _ = needleman_wunsch(\n            query_sequence,\n            template_sequence,\n            match_score=match_score,\n            mismatch_score=mismatch_score,\n            gap_score=gap_score,\n        )\n\n        pred_coords = map_template_coords_to_query(\n            aligned_query=aligned_q,\n            aligned_template=aligned_t,\n            template_coords=template_coords,\n            fill_value=np.nan,\n        )\n\n        vr = valid_row_ratio(pred_coords)\n\n        if interpolate_missing:\n            pred_coords = interpolate_nan_coords(pred_coords)\n\n        if pred_coords.shape[0] != len(query_sequence):\n            return None, vr, f\"pred_length_mismatch:{pred_coords.shape[0]}_vs_{len(query_sequence)}\"\n\n        return pred_coords, vr, \"ok\"\n\n    except Exception as e:\n        warnings.warn(f\"Structure build failed: {e}\")\n        return None, 0.0, f\"build_failed:{e}\"\n\ndef build_top5_conformations(\n    query_sequence: str,\n    ranked_templates: List[RankedTemplate],\n    train_coords_dict: Dict[str, np.ndarray],\n    top_n: int = 5,\n    noise_std: float = 0.01,\n    rng: Optional[np.random.Generator] = None,\n    interpolate_missing: bool = True,\n) -> Tuple[Optional[np.ndarray], List[str], str]:\n    \"\"\"\n    Build exactly top_n conformations for one query.\n    Final output shape: (top_n, L_query, 3)\n    \"\"\"\n    if rng is None:\n        rng = np.random.default_rng()\n\n    successful_structures: List[np.ndarray] = []\n    structure_sources: List[str] = []\n\n    for rt in ranked_templates:\n        if len(successful_structures) >= top_n:\n            break\n\n        if rt.status != \"ok\":\n            continue\n        if rt.train_label_id is None:\n            continue\n\n        template_coords = train_coords_dict.get(rt.train_label_id)\n        if template_coords is None:\n            continue\n\n        pred_coords, _, status = build_single_structure_from_template(\n            query_sequence=query_sequence,\n            template_sequence=rt.train_sequence,\n            template_coords=template_coords,\n            interpolate_missing=interpolate_missing,\n        )\n\n        if pred_coords is not None and status == \"ok\":\n            successful_structures.append(pred_coords)\n            structure_sources.append(rt.train_label_id)\n\n    if len(successful_structures) == 0:\n        return None, [], \"no_valid_structure_built\"\n\n    base_structure = successful_structures[0]\n\n    while len(successful_structures) < top_n:\n        noisy_copy = add_gaussian_noise(base_structure, std=noise_std, rng=rng)\n        successful_structures.append(noisy_copy)\n        structure_sources.append(f\"{structure_sources[0]}_noise_fill\")\n\n    try:\n        pred_5 = np.stack(successful_structures[:top_n], axis=0)\n    except Exception as e:\n        warnings.warn(f\"Failed to stack top-{top_n} conformations: {e}\")\n        return None, structure_sources, f\"stack_failed:{e}\"\n\n    if pred_5.ndim != 3 or pred_5.shape[0] != top_n or pred_5.shape[2] != 3:\n        return None, structure_sources, f\"unexpected_output_shape:{pred_5.shape}\"\n\n    if pred_5.shape[1] != len(query_sequence):\n        return None, structure_sources, f\"query_length_mismatch:{pred_5.shape[1]}_vs_{len(query_sequence)}\"\n\n    return pred_5, structure_sources, \"ok\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-18T12:36:01.216737Z","iopub.execute_input":"2026-03-18T12:36:01.217084Z","iopub.status.idle":"2026-03-18T12:36:01.234323Z","shell.execute_reply.started":"2026-03-18T12:36:01.217055Z","shell.execute_reply":"2026-03-18T12:36:01.232889Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Step 2: Test Inference Loop","metadata":{}},{"cell_type":"code","source":"def predict_top5_for_query(\n    query_target_id: str,\n    query_sequence: str,\n    retriever: TfidfKmerRetriever,\n    train_label_id_map: Dict[str, str],\n    train_coords_dict: Dict[str, np.ndarray],\n    coarse_top_k: int = 20,\n    final_top_n: int = 5,\n    noise_std: float = 0.01,\n    rng: Optional[np.random.Generator] = None,\n) -> Tuple[Optional[np.ndarray], str]:\n    \"\"\"\n    Full pipeline for one test query:\n    coarse retrieval -> alignment reranking -> top5 build\n    \"\"\"\n    coarse_candidates = retriever.query_top_k(query_sequence=query_sequence, k=coarse_top_k)\n    if len(coarse_candidates) == 0:\n        return None, \"no_coarse_candidates\"\n\n    ranked_templates = rerank_templates_with_alignment(\n        query_sequence=query_sequence,\n        coarse_candidates=coarse_candidates,\n        train_label_id_map=train_label_id_map,\n        train_coords_dict=train_coords_dict,\n    )\n\n    pred_5, _, status = build_top5_conformations(\n        query_sequence=query_sequence,\n        ranked_templates=ranked_templates[:final_top_n],\n        train_coords_dict=train_coords_dict,\n        top_n=final_top_n,\n        noise_std=noise_std,\n        rng=rng,\n        interpolate_missing=INTERPOLATE_MISSING,\n    )\n\n    if pred_5 is None:\n        return None, status\n\n    return pred_5, \"ok\"\n\n\ndef run_test_inference(\n    test_seq_df: pd.DataFrame,\n    train_seq_df: pd.DataFrame,\n    train_label_id_map: Dict[str, str],\n    train_coords_dict: Dict[str, np.ndarray],\n    coarse_top_k: int = 20,\n    final_top_n: int = 5,\n    kmer_size: int = 3,\n    noise_std: float = 0.01,\n    rng: Optional[np.random.Generator] = None,\n) -> Dict[str, np.ndarray]:\n    \"\"\"\n    Generate predictions for test set:\n        Dict[target_id, np.ndarray shape (5, L, 3)]\n    \"\"\"\n    if rng is None:\n        rng = np.random.default_rng()\n\n    retriever = TfidfKmerRetriever(kmer_size=kmer_size).fit(train_seq_df)\n    test_predictions: Dict[str, np.ndarray] = {}\n\n    for row in test_seq_df.itertuples(index=False):\n        target_id = str(row.target_id)\n        sequence = str(row.sequence)\n\n        pred_5, status = predict_top5_for_query(\n            query_target_id=target_id,\n            query_sequence=sequence,\n            retriever=retriever,\n            train_label_id_map=train_label_id_map,\n            train_coords_dict=train_coords_dict,\n            coarse_top_k=coarse_top_k,\n            final_top_n=final_top_n,\n            noise_std=noise_std,\n            rng=rng,\n        )\n\n        if pred_5 is None:\n            warnings.warn(f\"[TEST] Prediction failed for target_id={target_id}. Status={status}\")\n            continue\n\n        test_predictions[target_id] = pred_5\n\n    return test_predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-18T12:36:04.686658Z","iopub.execute_input":"2026-03-18T12:36:04.687036Z","iopub.status.idle":"2026-03-18T12:36:04.69893Z","shell.execute_reply.started":"2026-03-18T12:36:04.687003Z","shell.execute_reply":"2026-03-18T12:36:04.69788Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Step 3: Safe Submission Assembly Helpers","metadata":{}},{"cell_type":"code","source":"def extract_target_id_from_submission_id(submission_id: str) -> str:\n    \"\"\"\n    Extract target_id prefix from sample_submission ID.\n\n    Common case:\n    - \"157D_1\" -> \"157D\"\n    - \"abc_12\" -> \"abc\"\n\n    We use the part before the final underscore + number if present.\n    Otherwise fallback to the full string.\n    \"\"\"\n    submission_id = str(submission_id)\n\n    m = re.match(r\"^(.*)_([0-9]+)$\", submission_id)\n    if m is not None:\n        return m.group(1)\n\n    if \"_\" in submission_id:\n        return submission_id.rsplit(\"_\", 1)[0]\n\n    return submission_id\n\n\ndef get_residue_index_from_row(row: pd.Series) -> Optional[int]:\n    \"\"\"\n    Convert resid to zero-based index safely.\n    If invalid, return None.\n    \"\"\"\n    try:\n        resid_val = int(row[\"resid\"])\n        idx = resid_val - 1\n        if idx < 0:\n            return None\n        return idx\n    except Exception:\n        return None\n\n\ndef extract_coords_for_submission_row(\n    pred_5: np.ndarray,\n    residue_index: int,\n) -> List[float]:\n    \"\"\"\n    Extract [x_1, y_1, z_1, ..., x_5, y_5, z_5] from pred_5 for one residue index.\n\n    pred_5 shape must be (5, L, 3).\n    If residue_index is out of range, return zeros.\n    \"\"\"\n    zero_fill = [0.0] * 15\n\n    try:\n        pred_5 = np.asarray(pred_5, dtype=np.float64)\n        if pred_5.ndim != 3 or pred_5.shape[0] != 5 or pred_5.shape[2] != 3:\n            warnings.warn(f\"Prediction tensor has unexpected shape: {pred_5.shape}\")\n            return zero_fill\n\n        if residue_index < 0 or residue_index >= pred_5.shape[1]:\n            return zero_fill\n\n        values: List[float] = []\n        for conf_idx in range(5):\n            xyz = pred_5[conf_idx, residue_index, :]\n            if not np.all(np.isfinite(xyz)):\n                xyz = np.array([0.0, 0.0, 0.0], dtype=np.float64)\n            values.extend([float(xyz[0]), float(xyz[1]), float(xyz[2])])\n\n        if len(values) != 15:\n            return zero_fill\n\n        return values\n\n    except Exception as e:\n        warnings.warn(f\"Failed extracting coords for submission row: {e}\")\n        return zero_fill","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-18T12:36:07.831054Z","iopub.execute_input":"2026-03-18T12:36:07.831414Z","iopub.status.idle":"2026-03-18T12:36:07.844312Z","shell.execute_reply.started":"2026-03-18T12:36:07.831384Z","shell.execute_reply":"2026-03-18T12:36:07.843205Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Step 3: Absolute-Safe Submission Assembly","metadata":{}},{"cell_type":"code","source":"def build_submission_from_skeleton(\n    sample_submission_path: Path,\n    test_predictions: Dict[str, np.ndarray],\n) -> pd.DataFrame:\n    \"\"\"\n    Build submission by reading sample_submission.csv as the skeleton.\n\n    IMPORTANT:\n    - Do NOT create submission from scratch.\n    - Never delete any row.\n    - If prediction missing or invalid, fill with 0.0.\n    \"\"\"\n    if not sample_submission_path.exists():\n        raise FileNotFoundError(f\"sample_submission.csv not found: {sample_submission_path}\")\n\n    sub_df = pd.read_csv(sample_submission_path)\n\n    missing_cols = [c for c in EXPECTED_SUBMISSION_COLUMNS if c not in sub_df.columns]\n    if missing_cols:\n        raise ValueError(f\"sample_submission is missing required columns: {missing_cols}\")\n\n    # initialize all prediction columns to 0.0 first for absolute safety\n    coord_cols = EXPECTED_SUBMISSION_COLUMNS[3:]\n    for col in coord_cols:\n        sub_df[col] = 0.0\n\n    for idx, row in sub_df.iterrows():\n        try:\n            submission_id = str(row[\"ID\"])\n            target_id = extract_target_id_from_submission_id(submission_id)\n            residue_index = get_residue_index_from_row(row)\n\n            if residue_index is None:\n                continue\n\n            pred_5 = test_predictions.get(target_id)\n            if pred_5 is None:\n                continue\n\n            values = extract_coords_for_submission_row(pred_5, residue_index)\n\n            sub_df.loc[idx, coord_cols] = values\n\n        except Exception as e:\n            warnings.warn(f\"Failed filling submission row {idx}: {e}\")\n            continue\n\n    return sub_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-18T12:36:12.549729Z","iopub.execute_input":"2026-03-18T12:36:12.55038Z","iopub.status.idle":"2026-03-18T12:36:12.55978Z","shell.execute_reply.started":"2026-03-18T12:36:12.550342Z","shell.execute_reply":"2026-03-18T12:36:12.558274Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Step 4: Save & Validate","metadata":{}},{"cell_type":"code","source":"def validate_and_save_submission(\n    submission_df: pd.DataFrame,\n    sample_submission_path: Path,\n    output_path: Path,\n) -> None:\n    \"\"\"\n    Validate final submission against sample_submission.csv and save submission.csv\n    \"\"\"\n    sample_df = pd.read_csv(sample_submission_path)\n\n    if submission_df.shape != sample_df.shape:\n        raise ValueError(\n            f\"Submission shape mismatch. Got {submission_df.shape}, expected {sample_df.shape}\"\n        )\n\n    if list(submission_df.columns) != list(sample_df.columns):\n        raise ValueError(\n            \"Submission columns do not exactly match sample_submission columns.\\n\"\n            f\"Got: {list(submission_df.columns)}\\n\"\n            f\"Expected: {list(sample_df.columns)}\"\n        )\n\n    if submission_df.isna().any().any():\n        warnings.warn(\"NaN detected in submission. Replacing all NaN with 0.0\")\n        submission_df = submission_df.fillna(0.0)\n\n    # Force numeric prediction columns to finite float\n    coord_cols = EXPECTED_SUBMISSION_COLUMNS[3:]\n    for col in coord_cols:\n        submission_df[col] = pd.to_numeric(submission_df[col], errors=\"coerce\").fillna(0.0)\n        invalid_mask = ~np.isfinite(submission_df[col].to_numpy(dtype=np.float64))\n        if invalid_mask.any():\n            warnings.warn(f\"Non-finite values detected in column {col}. Replacing with 0.0\")\n            arr = submission_df[col].to_numpy(dtype=np.float64)\n            arr[invalid_mask] = 0.0\n            submission_df[col] = arr\n\n    output_path.parent.mkdir(parents=True, exist_ok=True)\n    submission_df.to_csv(output_path, index=False)\n    print(f\"Saved submission to: {output_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-18T12:36:14.59301Z","iopub.execute_input":"2026-03-18T12:36:14.593347Z","iopub.status.idle":"2026-03-18T12:36:14.603936Z","shell.execute_reply.started":"2026-03-18T12:36:14.593318Z","shell.execute_reply":"2026-03-18T12:36:14.602419Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Load Train/Test Data","metadata":{}},{"cell_type":"code","source":"train_seq_df = parse_sequences_csv(TRAIN_SEQ_PATH)\ntest_seq_df = parse_sequences_csv(TEST_SEQ_PATH)\n\nprint(\"train_seq_df shape:\", train_seq_df.shape)\nprint(\"test_seq_df shape:\", test_seq_df.shape)\n\ndisplay(train_seq_df.head())\ndisplay(test_seq_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-18T12:36:18.255539Z","iopub.execute_input":"2026-03-18T12:36:18.255938Z","iopub.status.idle":"2026-03-18T12:36:19.212682Z","shell.execute_reply.started":"2026-03-18T12:36:18.255909Z","shell.execute_reply":"2026-03-18T12:36:19.211337Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Load Train Coordinates and Build Mapping","metadata":{}},{"cell_type":"code","source":"train_coords_dict = parse_labels_to_dict(TRAIN_LABEL_PATH)\ntrain_label_id_map = build_sequence_to_label_id_map(train_seq_df, train_coords_dict)\n\nprint(\"Number of train label IDs:\", len(train_coords_dict))\nprint(\"Mapped train sequence IDs:\", len(train_label_id_map), \"/\", len(train_seq_df))\n\nexample_items = list(train_label_id_map.items())[:10]\nprint(\"\\nExample train sequence -> label mapping:\")\nfor item in example_items:\n    print(item)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-18T12:36:21.67678Z","iopub.execute_input":"2026-03-18T12:36:21.677091Z","iopub.status.idle":"2026-03-18T12:36:50.168231Z","shell.execute_reply.started":"2026-03-18T12:36:21.677064Z","shell.execute_reply":"2026-03-18T12:36:50.167315Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Run Test Inference","metadata":{}},{"cell_type":"code","source":"test_predictions = run_test_inference(\n    test_seq_df=test_seq_df,\n    train_seq_df=train_seq_df,\n    train_label_id_map=train_label_id_map,\n    train_coords_dict=train_coords_dict,\n    coarse_top_k=TOP_K_RETRIEVAL,\n    final_top_n=TOP_N_FINAL,\n    kmer_size=KMER_SIZE,\n    noise_std=NOISE_STD,\n    rng=RNG,\n)\n\nprint(\"Number of successful test predictions:\", len(test_predictions))\n\n# quick sanity check\nif len(test_predictions) > 0:\n    first_key = next(iter(test_predictions))\n    print(\"Example target_id:\", first_key)\n    print(\"Pred shape:\", test_predictions[first_key].shape)  # expected (5, L, 3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-18T12:38:37.977383Z","iopub.execute_input":"2026-03-18T12:38:37.977743Z","iopub.status.idle":"2026-03-18T13:03:20.250744Z","shell.execute_reply.started":"2026-03-18T12:38:37.977714Z","shell.execute_reply":"2026-03-18T13:03:20.249145Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Build Submission from sample_submission Skeleton","metadata":{}},{"cell_type":"code","source":"submission_df = build_submission_from_skeleton(\n    sample_submission_path=SAMPLE_SUB_PATH,\n    test_predictions=test_predictions,\n)\n\nprint(\"submission_df shape:\", submission_df.shape)\ndisplay(submission_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-18T13:03:26.714601Z","iopub.execute_input":"2026-03-18T13:03:26.715702Z","iopub.status.idle":"2026-03-18T13:03:58.717539Z","shell.execute_reply.started":"2026-03-18T13:03:26.715591Z","shell.execute_reply":"2026-03-18T13:03:58.71639Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Final Validate & Save submission.csv","metadata":{}},{"cell_type":"code","source":"OUTPUT_PATH = Path(\"/kaggle/working/submission.csv\")\n\nvalidate_and_save_submission(\n    submission_df=submission_df,\n    sample_submission_path=SAMPLE_SUB_PATH,\n    output_path=OUTPUT_PATH,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-18T13:04:28.36851Z","iopub.execute_input":"2026-03-18T13:04:28.369339Z","iopub.status.idle":"2026-03-18T13:04:28.560172Z","shell.execute_reply.started":"2026-03-18T13:04:28.3693Z","shell.execute_reply":"2026-03-18T13:04:28.558718Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Final Safety Checks","metadata":{}},{"cell_type":"code","source":"final_sub_df = pd.read_csv(OUTPUT_PATH)\n\nprint(\"Final submission shape:\", final_sub_df.shape)\nprint(\"Final submission columns match expected:\", list(final_sub_df.columns) == EXPECTED_SUBMISSION_COLUMNS)\nprint(\"Any NaN in final submission:\", final_sub_df.isna().any().any())\n\ncoord_cols = EXPECTED_SUBMISSION_COLUMNS[3:]\nall_finite = np.isfinite(final_sub_df[coord_cols].to_numpy(dtype=np.float64)).all()\nprint(\"All prediction values finite:\", all_finite)\n\ndisplay(final_sub_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-18T13:04:31.18608Z","iopub.execute_input":"2026-03-18T13:04:31.18658Z","iopub.status.idle":"2026-03-18T13:04:31.235943Z","shell.execute_reply.started":"2026-03-18T13:04:31.186545Z","shell.execute_reply":"2026-03-18T13:04:31.234569Z"}},"outputs":[],"execution_count":null}]}