{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":118765,"databundleVersionId":15231210,"sourceType":"competition"},{"sourceId":10855324,"sourceType":"datasetVersion","datasetId":6742586},{"sourceId":11118830,"sourceType":"datasetVersion","datasetId":6933267},{"sourceId":11424423,"sourceType":"datasetVersion","datasetId":7154921},{"sourceId":11424460,"sourceType":"datasetVersion","datasetId":7154953},{"sourceId":11424495,"sourceType":"datasetVersion","datasetId":7154981},{"sourceId":14630207,"sourceType":"datasetVersion","datasetId":7155284},{"sourceId":224830487,"sourceType":"kernelVersion"}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"MODEL_TYPE='protenix'\nVALIDATION=False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-26T18:53:46.101841Z","iopub.execute_input":"2026-01-26T18:53:46.102187Z","iopub.status.idle":"2026-01-26T18:53:46.105891Z","shell.execute_reply.started":"2026-01-26T18:53:46.102159Z","shell.execute_reply":"2026-01-26T18:53:46.105015Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Install requirements ","metadata":{}},{"cell_type":"code","source":"if MODEL_TYPE=='protenix' and VALIDATION:\n    !pip install --no-deps protenix\n    !pip install biopython\n    !pip install ml-collections\n    !pip install biotite==1.0.1\n    !pip install rdkit\n!export PROTENIX_DATA_ROOT_DIR=/kaggle/input/protenix-checkpoints","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-26T18:53:46.10739Z","iopub.execute_input":"2026-01-26T18:53:46.107593Z","iopub.status.idle":"2026-01-26T18:53:46.237139Z","shell.execute_reply.started":"2026-01-26T18:53:46.107575Z","shell.execute_reply":"2026-01-26T18:53:46.236052Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! mkdir /af3-dev \n! ln -s /kaggle/input/protenix-checkpoints /af3-dev/release_data\n! ls /af3-dev/release_data/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-26T18:53:46.239126Z","iopub.execute_input":"2026-01-26T18:53:46.239396Z","iopub.status.idle":"2026-01-26T18:53:46.591554Z","shell.execute_reply.started":"2026-01-26T18:53:46.239374Z","shell.execute_reply":"2026-01-26T18:53:46.590477Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install --no-index --find-links=\"/kaggle/input/required-files-protinex/wheels/\" biopython==1.83\n!pip install --no-index --find-links=\"/kaggle/input/required-files-protinex/wheels\" biotite==1.0.1\n!pip install --no-index --find-links=\"/kaggle/input/required-files-protinex/wheels\" rdkit==2024.9.6\n!pip install /kaggle/input/required-files-protinex/wheels/ml_collections-1.0.0-py3-none-any.whl\n!pip install /kaggle/input/required-files-protinex/wheels/modelcif-0.7-py3-none-any.whl\n!pip install /kaggle/input/required-files-protinex/fair_esm-2.0.0-py3-none-any.whl\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-26T18:53:46.593291Z","iopub.execute_input":"2026-01-26T18:53:46.593545Z","iopub.status.idle":"2026-01-26T18:57:26.886256Z","shell.execute_reply.started":"2026-01-26T18:53:46.593522Z","shell.execute_reply":"2026-01-26T18:57:26.88518Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nfrom scipy.spatial.transform import Rotation as _Rot\nimport random as _rnd\nfrom Bio import pairwise2 as _pw2\nfrom Bio.Seq import Seq as _Seq\nimport time as _time\nfrom scipy.spatial import distance_matrix as _distmat\nfrom tqdm.auto import tqdm\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-26T18:57:26.887333Z","iopub.execute_input":"2026-01-26T18:57:26.887624Z","iopub.status.idle":"2026-01-26T18:57:27.542701Z","shell.execute_reply.started":"2026-01-26T18:57:26.8876Z","shell.execute_reply":"2026-01-26T18:57:27.541687Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Helper scripts","metadata":{}},{"cell_type":"code","source":"import Bio\n\nfrom copy import deepcopy\n\nimport pandas as pd\nfrom Bio.PDB import Atom, Model, Chain, Residue, Structure, PDBParser\nfrom Bio import SeqIO\nimport os, sys\nimport re\nimport numpy as np\nimport torch\n\nimport matplotlib\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nimport time\ntime0=time.time()\n\nprint('IMPORT OK !!!!')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2026-01-26T18:57:27.543484Z","iopub.execute_input":"2026-01-26T18:57:27.543804Z","iopub.status.idle":"2026-01-26T18:57:30.83354Z","shell.execute_reply.started":"2026-01-26T18:57:27.543785Z","shell.execute_reply":"2026-01-26T18:57:30.832797Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PYTHON = sys.executable\nprint('PYTHON',PYTHON)\n\nRHONET_DIR=\\\n'/kaggle/input/data-for-demo-for-rhofold-plus-with-kaggle-msa/RhoFold-main'\n#'<your downloaded rhofold repo>/RhoFold-main'\n\nUSALIGN = \\\n'/kaggle/working//USalign'\n#'<your us align path>/USalign'\n\nos.system('cp /kaggle/input/usalign/USalign /kaggle/working/')\nos.system('sudo chmod u+x /kaggle/working//USalign')\nsys.path.append(RHONET_DIR)\n\n\nDATA_KAGGLE_DIR = '/kaggle/input/stanford-rna-3d-folding-2'\n\n\n# helper ----\nclass dotdict(dict):\n\t__setattr__ = dict.__setitem__\n\t__delattr__ = dict.__delitem__\n\n\tdef __getattr__(self, name):\n\t\ttry:\n\t\t\treturn self[name]\n\t\texcept KeyError:\n\t\t\traise AttributeError(name)\n\n# visualisation helper ----\ndef set_aspect_equal(ax):\n\tx_limits = ax.get_xlim()\n\ty_limits = ax.get_ylim()\n\tz_limits = ax.get_zlim()\n\n\t# Compute the mean of each axis\n\tx_middle = np.mean(x_limits)\n\ty_middle = np.mean(y_limits)\n\tz_middle = np.mean(z_limits)\n\n\t# Compute the max range across all axes\n\tmax_range = max(x_limits[1] - x_limits[0],\n\t\t\t\t\ty_limits[1] - y_limits[0],\n\t\t\t\t\tz_limits[1] - z_limits[0]) / 2.0\n\n\t# Set the new limits to ensure equal scaling\n\tax.set_xlim(x_middle - max_range, x_middle + max_range)\n\tax.set_ylim(y_middle - max_range, y_middle + max_range)\n\tax.set_zlim(z_middle - max_range, z_middle + max_range)\n\n\n\n\n# xyz df helper --------------------\ndef get_truth_df(target_id):\n    truth_df = LABEL_DF[LABEL_DF['target_id'] == target_id]\n    truth_df = truth_df.reset_index(drop=True)\n    return truth_df\n\ndef parse_output_to_df(output, seq, target_id):\n    df = []\n    chain_data = []\n    for i, res in enumerate(seq):\n        d=dict(ID = target_id,\n                    resname=res,\n                    resid=i+1)\n        for n in range(len(output)):\n            d={**d, f'x_{n+1}': round(output[n,i,0].item(),3),\n                     f'y_{n+1}': round(output[n,i,1].item(),3),\n                     f'z_{n+1}': round(output[n,i,2].item(),3)}\n        chain_data.append(d)\n\n    if len(chain_data)!=0:\n        chain_df = pd.DataFrame(chain_data)\n        df.append(chain_df)\n        ##print(chain_df)\n    return df\n\ndef parse_pdb_to_df(pdb_file, target_id):\n    parser = PDBParser()\n    structure = parser.get_structure('', pdb_file)\n\n    df = []\n    for model in structure:\n        for chain in model:\n            print(chain)\n            chain_data = []\n            for residue in chain:\n                # print(residue)\n                if residue.get_resname() in ['A', 'U', 'G', 'C']:\n                    # Check if the residue has a C1' atom\n                    if 'C1\\'' in residue:\n                        atom = residue['C1\\'']\n                        xyz = atom.get_coord()\n                        resname = residue.get_resname()\n                        resid = residue.get_id()[1]\n\n                        #todo detect discontinous: resid = prev_resid+1\n                        #ID\tresname\tresid\tx_1\ty_1\tz_1\n                        chain_data.append(dict(\n                            ID = target_id+'_'+str(resid),\n                            resname=resname,\n                            resid=resid,\n                            x_1=xyz[0],\n                            y_1=xyz[1],\n                            z_1=xyz[2],\n                        ))\n                        ##print(f\"Residue {resname} {resid}, Atom: {atom.get_name()}, xyz: {xyz}\")\n\n            if len(chain_data)!=0:\n                chain_df = pd.DataFrame(chain_data)\n                df.append(chain_df)\n                ##print(chain_df)\n    return df\n\n# usalign helper --------------------\ndef write_target_line(\n    atom_name, atom_serial, residue_name, chain_id, residue_num, x_coord, y_coord, z_coord, occupancy=1.0, b_factor=0.0, atom_type='P'\n):\n    \"\"\"\n    Writes a single line of PDB format based on provided atom information.\n\n    Args:\n        atom_name (str): Name of the atom (e.g., \"N\", \"CA\").\n        atom_serial (int): Atom serial number.\n        residue_name (str): Residue name (e.g., \"ALA\").\n        chain_id (str): Chain identifier.\n        residue_num (int): Residue number.\n        x_coord (float): X coordinate.\n        y_coord (float): Y coordinate.\n        z_coord (float): Z coordinate.\n        occupancy (float, optional): Occupancy value (default: 1.0).\n        b_factor (float, optional): B-factor value (default: 0.0).\n\n    Returns:\n        str: A single line of PDB string.\n    \"\"\"\n    return f'ATOM  {atom_serial:>5d}  {atom_name:<5s} {residue_name:<3s} {residue_num:>3d}    {x_coord:>8.3f}{y_coord:>8.3f}{z_coord:>8.3f}{occupancy:>6.2f}{b_factor:>6.2f}           {atom_type}\\n'\n\ndef write_xyz_to_pdb(df, pdb_file, xyz_id = 1):\n    resolved_cnt = 0\n    with open(pdb_file, 'w') as target_file:\n        for _, row in df.iterrows():\n            x_coord = row[f'x_{xyz_id}']\n            y_coord = row[f'y_{xyz_id}']\n            z_coord = row[f'z_{xyz_id}']\n\n            if x_coord > -1e17 and y_coord > -1e17 and z_coord > -1e17:\n                resolved_cnt += 1\n                target_line = write_target_line(\n                    atom_name=\"C1'\",\n                    atom_serial=int(row['resid']),\n                    residue_name=row['resname'],\n                    chain_id='0',\n                    residue_num=int(row['resid']),\n                    x_coord=x_coord,\n                    y_coord=y_coord,\n                    z_coord=z_coord,\n                    atom_type='C',\n                )\n                target_file.write(target_line)\n    return resolved_cnt\n\ndef parse_usalign_for_tm_score(output):\n    # Extract TM-score based on length of reference structure (second)\n    tm_score_match = re.findall(r'TM-score=\\s+([\\d.]+)', output)[1]\n    if not tm_score_match:\n        raise ValueError('No TM score found')\n    return float(tm_score_match)\n\ndef parse_usalign_for_transform(output):\n    # Locate the rotation matrix section\n    matrix_lines = []\n    found_matrix = False\n\n    for line in output.splitlines():\n        if \"The rotation matrix to rotate Structure_1 to Structure_2\" in line:\n            found_matrix = True\n        elif found_matrix and re.match(r'^\\d+\\s+[-\\d.]+\\s+[-\\d.]+\\s+[-\\d.]+\\s+[-\\d.]+$', line):\n            matrix_lines.append(line)\n        elif found_matrix and not line.strip():\n            break  # Stop parsing if an empty line is encountered after the matrix\n\n    # Parse the rotation matrix values\n    rotation_matrix = []\n    for line in matrix_lines:\n        parts = line.split()\n        row_values = list(map(float, parts[1:]))  # Skip the first column (index)\n        rotation_matrix.append(row_values)\n\n    return np.array(rotation_matrix)\n\ndef call_usalign(predict_df, truth_df, verbose=1):\n    truth_pdb = '~truth.pdb'\n    predict_pdb = '~predict.pdb'\n    write_xyz_to_pdb(predict_df, predict_pdb, xyz_id=1)\n    write_xyz_to_pdb(truth_df, truth_pdb, xyz_id=1)\n\n    command = f'{USALIGN} {predict_pdb} {truth_pdb} -atom \" C1\\'\" -m -'\n    output = os.popen(command).read()\n    if verbose==1:\n        print(output)\n    tm_score = parse_usalign_for_tm_score(output)\n    transform = parse_usalign_for_transform(output)\n    return tm_score, transform\n\nprint('HELPER OK!!!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-26T18:57:30.834357Z","iopub.execute_input":"2026-01-26T18:57:30.834744Z","iopub.status.idle":"2026-01-26T18:57:30.933927Z","shell.execute_reply.started":"2026-01-26T18:57:30.834722Z","shell.execute_reply":"2026-01-26T18:57:30.933061Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nsys.path.append(\"/kaggle/input/required-files-protinex/Protenix/runner\")  # Replace with actual path","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-26T18:57:30.935759Z","iopub.execute_input":"2026-01-26T18:57:30.936032Z","iopub.status.idle":"2026-01-26T18:57:30.939298Z","shell.execute_reply.started":"2026-01-26T18:57:30.936011Z","shell.execute_reply":"2026-01-26T18:57:30.938572Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\n\n# Add both the root directory and the build/lib directory to sys.path\nsys.path.append(\"/kaggle/input/required-files-protinex/Protenix/\")  # For direct imports\nsys.path.append(\"/kaggle/input/required-files-protinex/Protenix/build/lib/\")  # For compiled/built modules\n\n# Now try importing\ntry:\n    from runner.batch_inference import get_default_runner\n    from runner.inference import update_inference_configs, InferenceRunner\n    from protenix.data.infer_data_pipeline import InferenceDataset\n    print(\"✅ Imports successful!\")\nexcept ImportError as e:\n    print(f\"❌ Import failed: {e}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-26T18:57:30.941114Z","iopub.execute_input":"2026-01-26T18:57:30.941327Z","iopub.status.idle":"2026-01-26T18:57:33.555734Z","shell.execute_reply.started":"2026-01-26T18:57:30.941309Z","shell.execute_reply":"2026-01-26T18:57:33.554967Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from runner.batch_inference import get_default_runner\nfrom runner.inference import update_inference_configs, InferenceRunner\nfrom protenix.data.infer_data_pipeline import InferenceDataset","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-26T18:57:33.556737Z","iopub.execute_input":"2026-01-26T18:57:33.557392Z","iopub.status.idle":"2026-01-26T18:57:33.561281Z","shell.execute_reply.started":"2026-01-26T18:57:33.557355Z","shell.execute_reply":"2026-01-26T18:57:33.560529Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# RNABaselineModel \n# - Includes: save_pretrained(.pkl) + load_pretrained(.pkl)\n# - Includes: load_data() + build_banks() + make_submission_csv()\n# - Includes: tqdm progress + optional verbose logging\n# ============================================================\n\nimport os, pickle, warnings, time as _time, random as _rnd\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\nfrom scipy.spatial.transform import Rotation as _Rot\nfrom scipy.spatial import distance_matrix as _distmat\nfrom Bio import pairwise2 as _pw2\nfrom Bio.Seq import Seq as _Seq\n\nwarnings.filterwarnings(\"ignore\")\n\n\nclass RNABaselineModel:\n    # -------------------------\n    # init + tunables + logging\n    # -------------------------\n    def __init__(\n        self,\n        n_preds: int = 5,\n        pool: int = 60,\n        rel_len_cap: float = 0.30,\n        aln_match: int = 2,\n        aln_mismatch: int = -1,\n        aln_gap_open: float = -8,\n        aln_gap_ext: float = -0.2,\n        noise_gain: float = 0.25,\n        noise_floor: float = 0.02,\n        log_every: int = 20000,\n        verbose: bool = False,\n    ):\n        # hyperparams\n        self.N_PRED = int(n_preds)\n        self.CAND_POOL = int(pool)\n        self.REL_LEN_CAP = float(rel_len_cap)\n\n        self.ALN_MATCH = aln_match\n        self.ALN_MISMATCH = aln_mismatch\n        self.ALN_GAP_OPEN = aln_gap_open\n        self.ALN_GAP_EXT = aln_gap_ext\n\n        self.NOISE_GAIN = float(noise_gain)\n        self.NOISE_FLOOR = float(noise_floor)\n\n        # logging\n        self.LOG_EVERY = int(log_every)\n        self.VERBOSE = bool(verbose)\n\n        # data holders\n        self.train_df = None\n        self.valid_df = None\n        self.test_df = None\n        self.train_lbl = None\n        self.valid_lbl = None\n\n        # banks\n        self.train_bank = None\n        self.valid_bank = None\n\n        if self.VERBOSE:\n            print(\"[INIT] RNABaselineModel ready\")\n\n    # =========================================================\n    # PUBLIC: Load CSVs\n    # =========================================================\n    def load_data(\n        self,\n        train_seq_path: str,\n        valid_seq_path: str,\n        test_seq_path: str,\n        train_lbl_path: str,\n        valid_lbl_path: str,\n    ):\n        if self.VERBOSE:\n            print(\"[LOAD_DATA] start\")\n        self.train_df = pd.read_csv(train_seq_path)\n        self.valid_df = pd.read_csv(valid_seq_path)\n        self.test_df = pd.read_csv(test_seq_path)\n        self.train_lbl = pd.read_csv(train_lbl_path)\n        self.valid_lbl = pd.read_csv(valid_lbl_path)\n        if self.VERBOSE:\n            print(\"[LOAD_DATA] done\",\n                  len(self.train_df), len(self.valid_df), len(self.test_df),\n                  len(self.train_lbl), len(self.valid_lbl))\n        return self\n\n    # =========================================================\n    # PUBLIC: Build coordinate banks\n    # =========================================================\n    def build_banks(self):\n        if self.VERBOSE:\n            print(\"[BANK] start build_banks()\")\n        self.train_bank = self._pack_label_coords(self.train_lbl, tag=\"train\")\n        self.valid_bank = self._pack_label_coords(self.valid_lbl, tag=\"valid\")\n        if self.VERBOSE:\n            print(\"[BANK] done\", len(self.train_bank), len(self.valid_bank))\n        return self\n\n    # =========================================================\n    # PUBLIC: Save pretrained .pkl\n    # =========================================================\n    def save_pretrained(self, pkl_path: str = \"rna_baseline_pretrained.pkl\"):\n        \"\"\"\n        Save prepared baseline state:\n        - hyperparameters\n        - train_bank / valid_bank\n        - train_df / test_df (kept for template search + inference)\n        \"\"\"\n        if self.train_bank is None or self.train_df is None:\n            raise RuntimeError(\"Call load_data() + build_banks() before save_pretrained().\")\n\n        state = {\n            \"config\": {\n                \"N_PRED\": self.N_PRED,\n                \"CAND_POOL\": self.CAND_POOL,\n                \"REL_LEN_CAP\": self.REL_LEN_CAP,\n                \"ALN_MATCH\": self.ALN_MATCH,\n                \"ALN_MISMATCH\": self.ALN_MISMATCH,\n                \"ALN_GAP_OPEN\": self.ALN_GAP_OPEN,\n                \"ALN_GAP_EXT\": self.ALN_GAP_EXT,\n                \"NOISE_GAIN\": self.NOISE_GAIN,\n                \"NOISE_FLOOR\": self.NOISE_FLOOR,\n                \"LOG_EVERY\": self.LOG_EVERY,\n                \"VERBOSE\": self.VERBOSE,\n            },\n            \"train_bank\": self.train_bank,\n            \"valid_bank\": self.valid_bank,\n            \"train_df\": self.train_df,\n            \"test_df\": self.test_df,\n        }\n\n        with open(pkl_path, \"wb\") as f:\n            pickle.dump(state, f, protocol=pickle.HIGHEST_PROTOCOL)\n\n        print(f\"[PRETRAINED] saved -> {pkl_path}\")\n        return pkl_path\n\n    # =========================================================\n    # PUBLIC: Load pretrained .pkl\n    # =========================================================\n    def load_pretrained(self, pkl_path: str):\n        if not os.path.exists(pkl_path):\n            raise FileNotFoundError(pkl_path)\n\n        with open(pkl_path, \"rb\") as f:\n            state = pickle.load(f)\n\n        cfg = state[\"config\"]\n\n        self.N_PRED = cfg[\"N_PRED\"]\n        self.CAND_POOL = cfg[\"CAND_POOL\"]\n        self.REL_LEN_CAP = cfg[\"REL_LEN_CAP\"]\n        self.ALN_MATCH = cfg[\"ALN_MATCH\"]\n        self.ALN_MISMATCH = cfg[\"ALN_MISMATCH\"]\n        self.ALN_GAP_OPEN = cfg[\"ALN_GAP_OPEN\"]\n        self.ALN_GAP_EXT = cfg[\"ALN_GAP_EXT\"]\n        self.NOISE_GAIN = cfg[\"NOISE_GAIN\"]\n        self.NOISE_FLOOR = cfg[\"NOISE_FLOOR\"]\n        self.LOG_EVERY = cfg[\"LOG_EVERY\"]\n        self.VERBOSE = cfg[\"VERBOSE\"]\n\n        self.train_bank = state[\"train_bank\"]\n        self.valid_bank = state[\"valid_bank\"]\n        self.train_df = state[\"train_df\"]\n        self.test_df = state[\"test_df\"]\n\n        print(f\"[PRETRAINED] loaded <- {pkl_path}\")\n        return self\n\n    # =========================================================\n    # PUBLIC: Build submission CSV\n    # =========================================================\n    def make_submission_csv(self, out_path: str = \"submission.csv\") -> pd.DataFrame:\n        if self.test_df is None:\n            raise RuntimeError(\"test_df missing. Load data or load_pretrained first.\")\n        if self.train_df is None or self.train_bank is None:\n            raise RuntimeError(\"train_df/train_bank missing. Load data+build_banks or load_pretrained first.\")\n\n        t0 = _time.time()\n        rows = []\n\n        for idx, row in enumerate(tqdm(self.test_df.itertuples(index=False), total=len(self.test_df), desc=\"Targets\")):\n            tid = getattr(row, \"target_id\")\n            seq = getattr(row, \"sequence\")\n            cutoff = getattr(row, \"temporal_cutoff\", None)\n\n            preds = self._predict_ensemble(seq, tid, self.train_df, self.train_bank, n_preds=self.N_PRED, cutoff=cutoff)\n\n            for j in range(len(seq)):\n                rec = {\"ID\": f\"{tid}_{j+1}\", \"resname\": seq[j], \"resid\": j + 1}\n                for i in range(self.N_PRED):\n                    rec[f\"x_{i+1}\"] = float(preds[i][j][0])\n                    rec[f\"y_{i+1}\"] = float(preds[i][j][1])\n                    rec[f\"z_{i+1}\"] = float(preds[i][j][2])\n                rows.append(rec)\n\n            if self.VERBOSE and (idx % 5 == 0):\n                print(f\"[RUN] {idx+1}/{len(self.test_df)} tid={tid} L={len(seq)} elapsed={_time.time()-t0:.1f}s\")\n\n        sub = pd.DataFrame(rows)\n        cols = [\"ID\", \"resname\", \"resid\"]\n        for i in range(1, self.N_PRED + 1):\n            cols += [f\"x_{i}\", f\"y_{i}\", f\"z_{i}\"]\n        sub = sub[cols]\n        sub.to_csv(out_path, index=False)\n\n        print(f\"[DONE] wrote {out_path} rows={len(sub)} total_time={_time.time()-t0:.1f}s\")\n        return sub\n\n    # =========================================================\n    # INTERNAL: labels -> coords bank\n    # =========================================================\n    def _pack_label_coords(self, lbl_df: pd.DataFrame, tag: str = \"bank\") -> dict:\n        if lbl_df is None:\n            return {}\n        bank = {}\n        groups = lbl_df.groupby(lambda i: lbl_df[\"ID\"].iloc[i].rsplit(\"_\", 1)[0])\n\n        for g_idx, (tid, grp) in enumerate(tqdm(groups, desc=f\"Packing {tag}\", leave=False)):\n            pts = []\n            for _, r in grp.sort_values(\"resid\").iterrows():\n                pts.append([r[\"x_1\"], r[\"y_1\"], r[\"z_1\"]])\n            bank[tid] = np.asarray(pts, dtype=np.float32)\n\n            if self.VERBOSE and (g_idx % self.LOG_EVERY == 0) and g_idx > 0:\n                print(f\"[PACK] {tag}: packed={g_idx} targets\")\n\n        return bank\n\n    # =========================================================\n    # INTERNAL: template search\n    # =========================================================\n    def _search_templates(self, q_seq: str, db_seqs: pd.DataFrame, coord_bank: dict, cutoff: str = None, top_n: int = None) -> list:\n        if top_n is None:\n            top_n = self.CAND_POOL\n\n        if cutoff is not None and \"temporal_cutoff\" in db_seqs.columns:\n            db = db_seqs[db_seqs[\"temporal_cutoff\"] < cutoff]\n        else:\n            db = db_seqs\n\n        q_obj = _Seq(q_seq)\n        hits = []\n\n        for r_idx, row in enumerate(db.itertuples(index=False)):\n            tid = getattr(row, \"target_id\")\n            t_seq = getattr(row, \"sequence\")\n\n            if tid not in coord_bank:\n                continue\n\n            if abs(len(t_seq) - len(q_seq)) / max(len(t_seq), len(q_seq)) > self.REL_LEN_CAP:\n                continue\n\n            aln = _pw2.align.globalms(\n                q_obj, t_seq,\n                self.ALN_MATCH, self.ALN_MISMATCH,\n                self.ALN_GAP_OPEN, self.ALN_GAP_EXT,\n                one_alignment_only=True\n            )\n            if not aln:\n                continue\n\n            best = aln[0]\n            sim = best.score / (2 * min(len(q_seq), len(t_seq)))\n            hits.append((tid, t_seq, float(sim), coord_bank[tid]))\n\n            if self.VERBOSE and (r_idx % self.LOG_EVERY == 0) and r_idx > 0:\n                print(f\"[TEMPL] scanned={r_idx} hits={len(hits)}\")\n\n        hits.sort(key=lambda x: x[2], reverse=True)\n        return hits[:top_n]\n\n    # =========================================================\n    # INTERNAL: basic helix fallback\n    # =========================================================\n    @staticmethod\n    def _basic_helix(seq: str) -> np.ndarray:\n        n = len(seq)\n        xyz = np.zeros((n, 3), dtype=np.float32)\n        r, rise, ang = 10.0, 2.5, 0.6\n        for i in range(n):\n            a = i * ang\n            xyz[i] = [r * np.cos(a), r * np.sin(a), i * rise]\n        return xyz\n\n    # =========================================================\n    # INTERNAL: adaptive constraints\n    # =========================================================\n    @staticmethod\n    def _refine_with_constraints(xyz: np.ndarray, seq: str, conf: float = 1.0) -> np.ndarray:\n        out = xyz.copy()\n        n = len(seq)\n        strength = 0.8 * (1.0 - min(conf, 0.8))\n\n        # sequential distance constraint\n        dmin, dmax = 5.5, 6.5\n        for i in range(n - 1):\n            a = out[i]\n            b = out[i + 1]\n            dist = float(np.linalg.norm(b - a))\n            if dist < dmin or dist > dmax:\n                target = (dmin + dmax) / 2\n                direction = b - a\n                direction = direction / (np.linalg.norm(direction) + 1e-10)\n                adj = (target - dist) * strength\n                out[i + 1] = a + direction * (dist + adj)\n\n        # steric clash prevention\n        clash_min = 3.8\n        dm = _distmat(out, out)\n        clash = np.where((dm < clash_min) & (dm > 0))\n        for k in range(len(clash[0])):\n            i, j = int(clash[0][k]), int(clash[1][k])\n            if abs(i - j) <= 1 or i >= j:\n                continue\n            pi, pj = out[i], out[j]\n            dist = float(dm[i, j])\n            direction = pj - pi\n            direction = direction / (np.linalg.norm(direction) + 1e-10)\n            adj = (clash_min - dist) * strength\n            out[i] = pi - direction * (adj / 2)\n            out[j] = pj + direction * (adj / 2)\n\n        # light base-pair constraint (only when low confidence)\n        if strength > 0.3:\n            pairs = {\"A\": \"U\", \"U\": \"A\", \"G\": \"C\", \"C\": \"G\"}\n            for i in range(n):\n                comp = pairs.get(seq[i])\n                if comp is None:\n                    continue\n                for j in range(i + 3, min(i + 20, n)):\n                    if seq[j] == comp:\n                        cur = float(np.linalg.norm(out[i] - out[j]))\n                        if 8.0 < cur < 14.0:\n                            target = 10.5\n                            adj = (target - cur) * (strength * 0.3)\n                            direction = out[j] - out[i]\n                            direction = direction / (np.linalg.norm(direction) + 1e-10)\n                            out[i] = out[i] - direction * (adj / 2)\n                            out[j] = out[j] + direction * (adj / 2)\n                            break\n\n        return out\n\n    # =========================================================\n    # INTERNAL: adapt template to query using alignment gaps\n    # =========================================================\n    @staticmethod\n    def _adapt_template(q_seq: str, t_seq: str, t_xyz: np.ndarray, aln=None) -> np.ndarray:\n        if aln is None:\n            q_obj = _Seq(q_seq)\n            t_obj = _Seq(t_seq)\n            alns = _pw2.align.globalms(q_obj, t_obj, 2, -1, -10, -0.5, one_alignment_only=True)\n            if not alns:\n                return RNABaselineModel._basic_helix(q_seq)\n            aln = alns[0]\n\n        aq, at = aln.seqA, aln.seqB\n\n        q_xyz = np.zeros((len(q_seq), 3), dtype=np.float32)\n        q_xyz[:] = np.nan\n\n        qi, ti = 0, 0\n        for k in range(len(aq)):\n            qc, tc = aq[k], at[k]\n            if qc != \"-\" and tc != \"-\":\n                if ti < len(t_xyz):\n                    q_xyz[qi] = t_xyz[ti]\n                ti += 1\n                qi += 1\n            elif qc != \"-\" and tc == \"-\":\n                qi += 1\n            elif qc == \"-\" and tc != \"-\":\n                ti += 1\n\n        # pass1: interpolate\n        for i in range(len(q_xyz)):\n            if np.isnan(q_xyz[i, 0]):\n                p = -1\n                for j in range(i - 1, -1, -1):\n                    if not np.isnan(q_xyz[j, 0]):\n                        p = j\n                        break\n                n = -1\n                for j in range(i + 1, len(q_xyz)):\n                    if not np.isnan(q_xyz[j, 0]):\n                        n = j\n                        break\n                if p >= 0 and n >= 0:\n                    w = (i - p) / (n - p)\n                    q_xyz[i] = (1 - w) * q_xyz[p] + w * q_xyz[n]\n\n        # pass2: fill remaining NaNs\n        step = 4.0\n        for i in range(len(q_xyz)):\n            if np.isnan(q_xyz[i, 0]):\n                if i == 0:\n                    fv = -1\n                    for j in range(1, len(q_xyz)):\n                        if not np.isnan(q_xyz[j, 0]):\n                            fv = j\n                            break\n                    if fv >= 0:\n                        for j in range(fv - 1, -1, -1):\n                            d = np.random.normal(0, 1, 3).astype(np.float32)\n                            d = d / (np.linalg.norm(d) + 1e-10) * step\n                            q_xyz[j] = q_xyz[j + 1] - d\n                    else:\n                        return RNABaselineModel._basic_helix(q_seq)\n                else:\n                    pv = -1\n                    for j in range(i - 1, -1, -1):\n                        if not np.isnan(q_xyz[j, 0]):\n                            pv = j\n                            break\n                    if pv >= 0:\n                        if pv > 0 and not np.isnan(q_xyz[pv - 1, 0]):\n                            d = q_xyz[pv] - q_xyz[pv - 1]\n                            d = d / (np.linalg.norm(d) + 1e-10) * step\n                            q_xyz[i] = q_xyz[pv] + d\n                        else:\n                            d = np.random.normal(0, 1, 3).astype(np.float32)\n                            d = d / (np.linalg.norm(d) + 1e-10) * step\n                            q_xyz[i] = q_xyz[pv] + d\n                    else:\n                        q_xyz[i] = np.random.normal(0, 1, 3).astype(np.float32) * i\n\n        if np.isnan(q_xyz).any():\n            q_xyz = np.nan_to_num(q_xyz).astype(np.float32)\n        return q_xyz\n\n    # =========================================================\n    # INTERNAL: de novo fold\n    # =========================================================\n    @staticmethod\n    def _denovo_fold(seq: str, seed: int = None) -> np.ndarray:\n        if seed is not None:\n            np.random.seed(seed)\n            _rnd.seed(seed)\n\n        n = len(seq)\n        xyz = np.zeros((n, 3), dtype=np.float32)\n\n        for i in range(min(3, n)):\n            ang = i * 0.6\n            xyz[i] = [10.0 * np.cos(ang), 10.0 * np.sin(ang), i * 2.5]\n\n        direction = np.array([0.0, 0.0, 1.0], dtype=np.float32)\n        comp = {\"G\": \"C\", \"C\": \"G\", \"A\": \"U\", \"U\": \"A\"}\n\n        for i in range(3, n):\n            base = seq[i]\n            want = comp.get(base, \"X\")\n\n            found = False\n            pair_j = -1\n            win = min(i, 15)\n            for j in range(i - win, i):\n                if j >= 0 and seq[j] == want:\n                    found = True\n                    pair_j = j\n                    break\n\n            if found and (i - pair_j) <= 10 and _rnd.random() < 0.7:\n                pj = xyz[pair_j]\n                rand_off = np.random.normal(0, 1, 3).astype(np.float32) * 2.0\n                bp = 10.0 + _rnd.uniform(-1.0, 1.0)\n\n                center = np.mean(xyz[:i], axis=0)\n                d = center - pj\n                d = d / (np.linalg.norm(d) + 1e-10)\n\n                xyz[i] = pj + d * bp + rand_off\n\n                direction = np.random.normal(0, 0.3, 3).astype(np.float32)\n                direction = direction / (np.linalg.norm(direction) + 1e-10)\n            else:\n                if _rnd.random() < 0.3:\n                    ang = _rnd.uniform(0.2, 0.6)\n                    axis = np.random.normal(0, 1, 3).astype(np.float32)\n                    axis = axis / (np.linalg.norm(axis) + 1e-10)\n                    rot = _Rot.from_rotvec(ang * axis)\n                    direction = rot.apply(direction)\n                else:\n                    direction = direction + np.random.normal(0, 0.15, 3).astype(np.float32)\n                    direction = direction / (np.linalg.norm(direction) + 1e-10)\n\n                step = _rnd.uniform(3.5, 4.5)\n                xyz[i] = xyz[i - 1] + step * direction\n\n        return xyz\n\n    # =========================================================\n    # INTERNAL: main prediction\n    # =========================================================\n    def _predict_ensemble(self, seq: str, tid: str, db_df: pd.DataFrame, coord_bank: dict, n_preds: int = 5, cutoff: str = None) -> list:\n        preds = []\n        candidates = self._search_templates(seq, db_df, coord_bank, cutoff=cutoff, top_n=self.CAND_POOL)\n\n        if candidates:\n            for (_, tpl_seq, sim, tpl_xyz) in candidates:\n                aligned_xyz = self._adapt_template(seq, tpl_seq, tpl_xyz)\n                refined = self._refine_with_constraints(aligned_xyz, seq, conf=sim)\n\n                noise = max(self.NOISE_FLOOR, self.NOISE_GAIN * (1.0 - sim))\n                out = refined.copy()\n                out += np.random.normal(0, noise, out.shape).astype(np.float32)\n\n                preds.append(out)\n                if len(preds) >= n_preds:\n                    break\n\n        while len(preds) < n_preds:\n            seed = (hash(tid) % 10000) + len(preds) * 1\n            xyz = self._denovo_fold(seq, seed=seed)\n            xyz = self._refine_with_constraints(xyz, seq, conf=0.2)\n            preds.append(xyz)\n\n        return preds[:n_preds]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-26T18:57:33.562249Z","iopub.execute_input":"2026-01-26T18:57:33.562505Z","iopub.status.idle":"2026-01-26T18:57:33.615458Z","shell.execute_reply.started":"2026-01-26T18:57:33.562474Z","shell.execute_reply":"2026-01-26T18:57:33.614477Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# USAGE (TRAIN -> SAVE PKL -> LOAD PKL -> SUBMISSION)\n# ============================================================\n\n# 1) Build + \"pretrain\" (banks) + save pkl\nmodel = RNABaselineModel(\n    n_preds=5,\n    pool=50,\n    rel_len_cap=0.30,\n    log_every=20000,\n    verbose=True\n)\n\nmodel.load_data(\n    train_seq_path=\"/kaggle/input/stanford-rna-3d-folding-2/train_sequences.csv\",\n    valid_seq_path=\"/kaggle/input/stanford-rna-3d-folding-2/validation_sequences.csv\",\n    test_seq_path=\"/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv\",\n    train_lbl_path=\"/kaggle/input/stanford-rna-3d-folding-2/train_labels.csv\",\n    valid_lbl_path=\"/kaggle/input/stanford-rna-3d-folding-2/validation_labels.csv\",\n)\n\nmodel.build_banks()\nmodel.save_pretrained(\"rna_baseline_pretrained.pkl\")\n\n# 2) Fresh object -> load pkl -> generate submission\ninfer = RNABaselineModel(verbose=False)\ninfer.load_pretrained(\"rna_baseline_pretrained.pkl\")\n\nsubmission_df = infer.make_submission_csv(out_path=\"submission.csv\")\nsubmission_df.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-26T18:57:33.616429Z","iopub.execute_input":"2026-01-26T18:57:33.61673Z","execution_failed":"2026-01-26T19:19:45.697Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}