{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":87793,"databundleVersionId":12276181,"sourceType":"competition"},{"sourceId":5123458,"sourceType":"datasetVersion","datasetId":2975803},{"sourceId":10923077,"sourceType":"datasetVersion","datasetId":6785143},{"sourceId":11118830,"sourceType":"datasetVersion","datasetId":6933267},{"sourceId":11641448,"sourceType":"datasetVersion","datasetId":7304841},{"sourceId":11644010,"sourceType":"datasetVersion","datasetId":7306643},{"sourceId":11695366,"sourceType":"datasetVersion","datasetId":6791615},{"sourceId":11837219,"sourceType":"datasetVersion","datasetId":7436926},{"sourceId":11855248,"sourceType":"datasetVersion","datasetId":7449300},{"sourceId":11855563,"sourceType":"datasetVersion","datasetId":7449485},{"sourceId":11891296,"sourceType":"datasetVersion","datasetId":7446443},{"sourceId":11956872,"sourceType":"datasetVersion","datasetId":7491577},{"sourceId":11969392,"sourceType":"datasetVersion","datasetId":7526656},{"sourceId":13297027,"sourceType":"datasetVersion","datasetId":8427792},{"sourceId":13307215,"sourceType":"datasetVersion","datasetId":8435012},{"sourceId":13328497,"sourceType":"datasetVersion","datasetId":8450079},{"sourceId":13332155,"sourceType":"datasetVersion","datasetId":8452932},{"sourceId":13361789,"sourceType":"datasetVersion","datasetId":8475411}],"dockerImageVersionId":31154,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":2379.646057,"end_time":"2025-05-25T12:08:20.685554","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2025-05-25T11:28:41.039497","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"MODEL_TYPE='protenix'\nVALIDATION=False\nlocal = False","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.execute_input":"2025-05-25T11:28:45.334825Z","iopub.status.busy":"2025-05-25T11:28:45.334181Z","iopub.status.idle":"2025-05-25T11:28:45.339940Z","shell.execute_reply":"2025-05-25T11:28:45.339478Z"},"papermill":{"duration":0.015401,"end_time":"2025-05-25T11:28:45.341006","exception":false,"start_time":"2025-05-25T11:28:45.325605","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":0.006051,"end_time":"2025-05-25T11:28:45.353715","exception":false,"start_time":"2025-05-25T11:28:45.347664","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Proteinx","metadata":{"papermill":{"duration":0.005907,"end_time":"2025-05-25T11:28:45.365718","exception":false,"start_time":"2025-05-25T11:28:45.359811","status":"completed"},"tags":[]}},{"cell_type":"code","source":"if local:\n !pip install -q --no-deps protenix\n !pip install -q biopython\n !pip install -q ml-collections\n !pip install -q biotite==1.0.1\n !pip install -q rdkit","metadata":{"execution":{"iopub.execute_input":"2025-05-25T11:28:45.378910Z","iopub.status.busy":"2025-05-25T11:28:45.378389Z","iopub.status.idle":"2025-05-25T11:28:45.383286Z","shell.execute_reply":"2025-05-25T11:28:45.382795Z"},"papermill":{"duration":0.012562,"end_time":"2025-05-25T11:28:45.384342","exception":false,"start_time":"2025-05-25T11:28:45.371780","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import Bio\nimport random\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\nimport pandas as pd\nimport matplotlib\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nimport glob\nimport pickle\nprint('IMPORT OK !!!!')","metadata":{"execution":{"iopub.execute_input":"2025-05-25T11:28:45.397960Z","iopub.status.busy":"2025-05-25T11:28:45.397499Z","iopub.status.idle":"2025-05-25T11:28:51.656495Z","shell.execute_reply":"2025-05-25T11:28:51.655698Z"},"papermill":{"duration":6.26692,"end_time":"2025-05-25T11:28:51.657777","exception":false,"start_time":"2025-05-25T11:28:45.390857","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!export PROTENIX_DATA_ROOT_DIR=/kaggle/input/protenix-checkpoints\n","metadata":{"execution":{"iopub.execute_input":"2025-05-25T11:28:51.671952Z","iopub.status.busy":"2025-05-25T11:28:51.671659Z","iopub.status.idle":"2025-05-25T11:28:51.795861Z","shell.execute_reply":"2025-05-25T11:28:51.794949Z"},"papermill":{"duration":0.132421,"end_time":"2025-05-25T11:28:51.797234","exception":false,"start_time":"2025-05-25T11:28:51.664813","status":"completed"},"tags":[]},"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":{"execution":{"iopub.execute_input":"2025-05-25T11:28:51.811819Z","iopub.status.busy":"2025-05-25T11:28:51.811065Z","iopub.status.idle":"2025-05-25T11:28:52.200199Z","shell.execute_reply":"2025-05-25T11:28:52.199401Z"},"papermill":{"duration":0.397713,"end_time":"2025-05-25T11:28:52.201632","exception":false,"start_time":"2025-05-25T11:28:51.803919","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if MODEL_TYPE=='protenix':\n    \n    \n    from runner.batch_inference import get_default_runner\n    from runner.inference import update_inference_configs, InferenceRunner\n\n    from protenix.data.infer_data_pipeline import InferenceDataset\n\n    np.random.seed(21)\n    torch.random.manual_seed(21)\n    torch.cuda.manual_seed_all(21)\n\n    class DictDataset(InferenceDataset):\n        def __init__(\n            self,\n            seq_list: list,\n            dump_dir: str,\n            id_list: list = None,\n            use_msa: bool = False,\n        ) -> None:\n\n            self.dump_dir = dump_dir\n            self.use_msa = use_msa\n            if isinstance(id_list,type(None)):\n                self.inputs = [{\"sequences\": \n                                [{\"rnaSequence\": \n                                  {\"sequence\": seq, \n                                   \"count\": 1}}],\n                                \"name\": \"query\"} for seq in seq_list]\n            else:\n                self.inputs = [{\"sequences\": \n                                [{\"rnaSequence\": \n                                  {\"sequence\": seq, \n                                   \"count\": 1}}],\n                                \"name\": i} for i, seq in zip(id_list,seq_list)]","metadata":{"execution":{"iopub.execute_input":"2025-05-25T11:28:52.215799Z","iopub.status.busy":"2025-05-25T11:28:52.215470Z","iopub.status.idle":"2025-05-25T11:28:52.224416Z","shell.execute_reply":"2025-05-25T11:28:52.223717Z"},"papermill":{"duration":0.017361,"end_time":"2025-05-25T11:28:52.225578","exception":false,"start_time":"2025-05-25T11:28:52.208217","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nfrom torch.cuda.amp import autocast\n\nif MODEL_TYPE == 'protenix':\n    # Import configs\n    from configs.configs_base import configs as configs_base\n    from configs.configs_data import data_configs\n    from configs.configs_inference import inference_configs\n    from protenix.config.config import parse_configs\n\n    # Enable DeepSpeed Evo Attention based on environment variable\n    configs_base[\"use_deepspeed_evo_attention\"] = (\n        os.environ.get(\"USE_DEEPSPEED_EVO_ATTTENTION\", False) == \"true\"\n    )\n\n    # Set model-specific configs\n    configs_base[\"model\"][\"N_cycle\"] = 10\n    configs_base[\"sample_diffusion\"][\"N_sample\"] = 10\n    configs_base[\"sample_diffusion\"][\"N_step\"] = 200\n\n\n    # Set checkpoint path\n    #inference_configs['load_checkpoint_path'] = '/kaggle/input/17april-proteinx/pytorch/default/17/1999_ema_0.995_casp16-2000steps.pt'\n    inference_configs['load_checkpoint_path'] = '/kaggle/input/13okt-may/15999_ema_0.995.pt'\n    # Merge all configs\n    configs = {**configs_base, **{\"data\": data_configs}, **inference_configs}\n    configs = parse_configs(configs=configs, fill_required_with_null=True)\n\n    # Optional precision field for the runner\n    configs[\"precision\"] = \"bfloat16\"\n\n    # Check GPU and bfloat16 support\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    if not torch.cuda.is_bf16_supported():\n        raise RuntimeError(\"Your GPU does not support bfloat16.\")\n\n   \n\n    # Inference with autocast for bfloat16\n    with autocast(dtype=torch.bfloat16):\n         runner = InferenceRunner(configs)  # Replace with the actual method in your runner, e.g., 'predict()' or 'inference()'","metadata":{"papermill":{"duration":0.00767,"end_time":"2025-05-25T11:34:24.889474","exception":false,"start_time":"2025-05-25T11:34:24.881804","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def parse_c1_coords_to_df(c1_coords_list, sequence, target_id, num_samples=5):\n    \"\"\"\n    Converts list of C1′ atom coordinates from multiple conformations into a DataFrame.\n    Fills missing conformations with NaNs.\n\n    Parameters:\n        c1_coords_list: list of arrays of shape (L, 3)\n        sequence: RNA sequence\n        target_id: string\n        num_samples: number of conformations to output columns for (e.g. 5)\n\n    Returns:\n        pd.DataFrame with columns: ID, resname, resid, x_1, y_1, z_1, ..., x_k, y_k, z_k\n    \"\"\"\n    data = []\n    L = len(sequence)\n    for i in range(L):\n        row = {\n            \"ID\": f\"{target_id}_{i + 1}\",  # Fixed line\n            \"resname\": sequence[i],\n            \"resid\": i + 1,\n        }\n        for j in range(num_samples):\n            if j < len(c1_coords_list):\n                coords = c1_coords_list[j]\n                if i < len(coords):\n                    row[f\"x_{j+1}\"] = round(coords[i][0], 3)\n                    row[f\"y_{j+1}\"] = round(coords[i][1], 3)\n                    row[f\"z_{j+1}\"] = round(coords[i][2], 3)\n                else:\n                    row[f\"x_{j+1}\"] = row[f\"y_{j+1}\"] = row[f\"z_{j+1}\"] = np.nan\n            else:\n                row[f\"x_{j+1}\"] = row[f\"y_{j+1}\"] = row[f\"z_{j+1}\"] = np.nan\n        data.append(row)\n\n    df = pd.DataFrame(data)\n    return df","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def extract_c1_prime_coordinates_batched(coordinate_batch, atom_to_token_idx):\n    \"\"\"\n    Extracts C1′ coordinates from a batch of conformations using atom_to_token_idx == 12.\n    \n    Parameters:\n    - coordinate_batch: (N_samples, N_atoms, 3)\n    - atom_to_token_idx: (N_atoms,) tensor mapping atoms to tokenized atom type indices\n    \n    Returns:\n    - List of np.array of shape (L, 3), one per conformation\n    \"\"\"\n    mask = (atom_to_token_idx == 12).detach().cpu().numpy()\n    coordinate_batch = coordinate_batch.detach().cpu().numpy()\n    \n    all_c1_coords = [coords[mask] for coords in coordinate_batch]  # each is (L, 3)\n    return all_c1_coords","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nimport torch\n\n# --------------------------\n# Your dependencies and imports here\n# Make sure these are defined/imported:\n# DictDataset, runner.predict, update_inference_configs, extract_c1_prime_coordinates_batched, parse_c1_coords_to_df\n# --------------------------\n\n# really infer on testset\nKAGGLE_REEUN = os.getenv('KAGGLE_IS_COMPETITION_RERUN')\n\n\n\ndef kabsch_rmsd(P, Q):\n    \"\"\"\n    Compute the Kabsch RMSD between two coordinate sets P and Q.\n    \"\"\"\n    P = P - P.mean(0)\n    Q = Q - Q.mean(0)\n    C = np.dot(np.transpose(P), Q)\n    V, S, W = np.linalg.svd(C)\n    d = np.sign(np.linalg.det(np.dot(V, W)))\n    U = np.dot(V, np.dot(np.diag([1, 1, d]), W))\n    P_rot = np.dot(P, U)\n    return np.sqrt(np.mean(np.sum((P_rot - Q) ** 2, axis=1)))\n\nif MODEL_TYPE == 'protenix' and not VALIDATION:\n    test_df = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/test_sequences.csv')\n    if not KAGGLE_REEUN:\n       #test_df = test_df.head(1)\n       test_df = test_df[test_df['target_id'] == \"R1107\"].reset_index(drop=True)\n\n\n    \n    dataset = DictDataset(\n        test_df.sequence,\n        dump_dir='output',\n        id_list=test_df.target_id,\n        use_msa=False\n    )\n\n    output_filename = 'submission_proteinx.csv'\n    if os.path.exists(output_filename):\n        os.remove(output_filename)  # Remove old file if exists\n\n    long_seq_count = 0  # To keep track of how many sequences were filled with zeros\n\n    for i, seq in tqdm(enumerate(test_df.sequence), total=len(dataset)):\n        target_id = test_df.target_id[i]\n        print(f\"\\nProcessing {target_id} | length={len(seq)}\")\n\n        if len(seq) > 820:\n            print(f\"⚠️ Sequence too long ({len(seq)} residues), filling with zeros for {target_id}\")\n            long_seq_count += 1\n            num_residues = len(seq)\n            dummy_data = {\n                \"ID\": [f\"{target_id}_{resid}\" for resid in range(num_residues)],\n                \"resname\": [seq[resid] for resid in range(num_residues)],\n                \"resid\": [resid for resid in range(num_residues)],\n            }\n            for conf_idx in range(1, 6):\n                dummy_data[f'x_{conf_idx}'] = [0.0] * num_residues\n                dummy_data[f'y_{conf_idx}'] = [0.0] * num_residues\n                dummy_data[f'z_{conf_idx}'] = [0.0] * num_residues\n\n            result_gpe = pd.DataFrame(dummy_data)\n\n        else:\n            data, atom_array, data_error_message = dataset[i]\n            assert data_error_message == ''\n            assert target_id == data[\"sample_name\"]\n\n            new_configs = update_inference_configs(configs, data[\"N_token\"].item())\n            runner.update_model_configs(new_configs)\n\n            prediction = runner.predict(data)\n            coords_batch = prediction['coordinate']  # (N_conf, N_atoms, 3)\n            summary_conf = prediction['summary_confidence']  # list of dicts\n\n            c1_coords_list = extract_c1_prime_coordinates_batched(\n                coordinate_batch=coords_batch,\n                atom_to_token_idx=data['input_feature_dict']['atom_to_tokatom_idx']\n            )\n\n            c1_coords_list_np = [c.cpu().numpy() if isinstance(c, torch.Tensor) else c for c in c1_coords_list]\n\n            # Step 1: Select top-1 by pLDDT\n            top1_idx = np.argmax([float(conf['chain_plddt']) for conf in summary_conf])\n            selected_indices = [top1_idx]\n\n            # Step 2: Select 4 most diverse conformations\n            while len(selected_indices) < 5:\n                remaining = [j for j in range(len(c1_coords_list_np)) if j not in selected_indices]\n                max_min_rmsd = []\n                for j in remaining:\n                    min_rmsd_to_selected = min(\n                        kabsch_rmsd(c1_coords_list_np[j], c1_coords_list_np[k])\n                        for k in selected_indices\n                    )\n                    max_min_rmsd.append(min_rmsd_to_selected)\n                next_idx = remaining[np.argmax(max_min_rmsd)]\n                selected_indices.append(next_idx)\n\n            top_5_coords = [c1_coords_list[j] for j in selected_indices]\n            \n            result_gpe = parse_c1_coords_to_df(\n                c1_coords_list=top_5_coords,\n                sequence=seq,\n                target_id=target_id,\n                num_samples=5\n            )\n\n        # Write output\n        result_gpe.to_csv(output_filename, index=False, mode='a', header=(i == 0))\n        torch.cuda.empty_cache()\n\n    print(f\"\\n✅ Completed inference. Sequences with length > 720 filled with zeros: {long_seq_count}\")\n\n    # -------------------------------\n    # Final merging with sample_submission\n    # -------------------------------\n    my_submission = pd.read_csv(output_filename)\n    sample_submission = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/sample_submission.csv')\n\n    non_coord_cols = ['ID', 'resname', 'resid']\n    final_submission = sample_submission[non_coord_cols].merge(my_submission, on=non_coord_cols, how='left')\n    final_submission = final_submission.fillna(0.0)\n    final_submission.to_csv(output_filename, index=False)\n    print(\"✅ Final submission aligned, filled, and saved as 'submission.csv'. Ready for upload!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Boltz","metadata":{"papermill":{"duration":0.007514,"end_time":"2025-05-25T11:34:24.904636","exception":false,"start_time":"2025-05-25T11:34:24.897122","status":"completed"},"tags":[]}},{"cell_type":"code","source":"#!pip install --no-index /kaggle/input/boltz-dependencies/*whl --no-deps\n!pip install --no-index /kaggle/input/fairscale-0413/*whl --no-deps\n!pip install -q /kaggle/input/boltz-dependencies/mashumaro-3.14-py3-none-any.whl  --no-deps\n!pip install -q /kaggle/input/boltz-dependencies/modelcif-1.3-py3-none-any.whl  --no-deps\n!pip install -q /kaggle/input/boltz-dependencies/ihm-2.2-py3-none-any.whl --no-deps","metadata":{"execution":{"iopub.execute_input":"2025-05-25T11:34:24.921208Z","iopub.status.busy":"2025-05-25T11:34:24.920931Z","iopub.status.idle":"2025-05-25T11:34:31.356868Z","shell.execute_reply":"2025-05-25T11:34:31.355769Z"},"papermill":{"duration":6.446705,"end_time":"2025-05-25T11:34:31.358959","exception":false,"start_time":"2025-05-25T11:34:24.912254","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%cd /kaggle/working/\n%mkdir inputs_prediction\n%mkdir outputs_prediction\n%cp -rf /kaggle/input/rna-prediction-boltz/boltz/src/boltz .","metadata":{"execution":{"iopub.execute_input":"2025-05-25T11:34:31.376504Z","iopub.status.busy":"2025-05-25T11:34:31.376224Z","iopub.status.idle":"2025-05-25T11:34:32.429354Z","shell.execute_reply":"2025-05-25T11:34:32.428483Z"},"papermill":{"duration":1.063386,"end_time":"2025-05-25T11:34:32.430867","exception":false,"start_time":"2025-05-25T11:34:31.367481","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile inference.py\nimport os\nimport random\nos.environ['CUBLAS_WORKSPACE_CONFIG'] = ':4096:8'\n\nimport pickle\nimport urllib.request\nfrom dataclasses import asdict, dataclass\nfrom pathlib import Path\nfrom typing import Literal, Optional\n\nimport click\nimport torch\nfrom pytorch_lightning import Trainer#, seed_everything\nfrom pytorch_lightning.strategies import DDPStrategy\nfrom pytorch_lightning.utilities import rank_zero_only\nfrom tqdm import tqdm\n\nfrom boltz.data import const\nfrom boltz.data.module.inference import BoltzInferenceDataModule\nfrom boltz.data.msa.mmseqs2 import run_mmseqs2\nfrom boltz.data.parse.a3m import parse_a3m\nfrom boltz.data.parse.csv import parse_csv\nfrom boltz.data.parse.fasta import parse_fasta\nfrom boltz.data.parse.yaml import parse_yaml\nfrom boltz.data.types import MSA, Manifest, Record\nfrom boltz.data.write.writer import BoltzWriter\nfrom boltz.model.model import Boltz1\n\nimport numpy as np\ndef seed_everything(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.benchmark = False\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.enabled = True\n    torch.use_deterministic_algorithms(True)\n\n\nCCD_URL = \"https://huggingface.co/boltz-community/boltz-1/resolve/main/ccd.pkl\"\nMODEL_URL = (\n    \"https://huggingface.co/boltz-community/boltz-1/resolve/main/boltz1_conf.ckpt\"\n)\n\n\n@dataclass\nclass BoltzProcessedInput:\n    \"\"\"Processed input data.\"\"\"\n\n    manifest: Manifest\n    targets_dir: Path\n    msa_dir: Path\n\n\n@dataclass\nclass BoltzDiffusionParams:\n    \"\"\"Diffusion process parameters.\"\"\"\n\n    gamma_0: float = 0.605\n    gamma_min: float = 1.107\n    noise_scale: float = 0.901\n    rho: float = 8\n    step_scale: float = 1.638\n    sigma_min: float = 0.0004\n    sigma_max: float = 160.0\n    sigma_data: float = 16.0\n    P_mean: float = -1.2\n    P_std: float = 1.5\n    coordinate_augmentation: bool = True\n    alignment_reverse_diff: bool = True\n    synchronize_sigmas: bool = True\n    use_inference_model_cache: bool = True\n\n\n@rank_zero_only\ndef download(cache: Path) -> None:\n    \"\"\"Download all the required data.\n\n    Parameters\n    ----------\n    cache : Path\n        The cache directory.\n\n    \"\"\"\n    # Download CCD\n    ccd = cache / \"ccd.pkl\"\n    if not ccd.exists():\n        click.echo(\n            f\"Downloading the CCD dictionary to {ccd}. You may \"\n            \"change the cache directory with the --cache flag.\"\n        )\n        urllib.request.urlretrieve(CCD_URL, str(ccd))  # noqa: S310\n\n    # Download model\n    model = cache / \"boltz1_conf.ckpt\"\n    if not model.exists():\n        click.echo(\n            f\"Downloading the model weights to {model}. You may \"\n            \"change the cache directory with the --cache flag.\"\n        )\n        urllib.request.urlretrieve(MODEL_URL, str(model))  # noqa: S310\n\n\ndef check_inputs(\n    data: Path,\n    outdir: Path,\n    override: bool = False,\n) -> list[Path]:\n    \"\"\"Check the input data and output directory.\n\n    If the input data is a directory, it will be expanded\n    to all files in this directory. Then, we check if there\n    are any existing predictions and remove them from the\n    list of input data, unless the override flag is set.\n\n    Parameters\n    ----------\n    data : Path\n        The input data.\n    outdir : Path\n        The output directory.\n    override: bool\n        Whether to override existing predictions.\n\n    Returns\n    -------\n    list[Path]\n        The list of input data.\n\n    \"\"\"\n    click.echo(\"Checking input data.\")\n\n    # Check if data is a directory\n    if data.is_dir():\n        data: list[Path] = list(data.glob(\"*\"))\n\n        # Filter out non .fasta or .yaml files, raise\n        # an error on directory and other file types\n        filtered_data = []\n        for d in data:\n            if d.suffix in (\".fa\", \".fas\", \".fasta\", \".yml\", \".yaml\"):\n                filtered_data.append(d)\n            elif d.is_dir():\n                msg = f\"Found directory {d} instead of .fasta or .yaml.\"\n                raise RuntimeError(msg)\n            else:\n                msg = (\n                    f\"Unable to parse filetype {d.suffix}, \"\n                    \"please provide a .fasta or .yaml file.\"\n                )\n                raise RuntimeError(msg)\n\n        data = filtered_data\n    else:\n        data = [data]\n\n    # Check if existing predictions are found\n    existing = (outdir / \"predictions\").rglob(\"*\")\n    existing = {e.name for e in existing if e.is_dir()}\n\n    # Remove them from the input data\n    if existing and not override:\n        data = [d for d in data if d.stem not in existing]\n        num_skipped = len(existing) - len(data)\n        msg = (\n            f\"Found some existing predictions ({num_skipped}), \"\n            f\"skipping and running only the missing ones, \"\n            \"if any. If you wish to override these existing \"\n            \"predictions, please set the --override flag.\"\n        )\n        click.echo(msg)\n    elif existing and override:\n        msg = \"Found existing predictions, will override.\"\n        click.echo(msg)\n\n    return data\n\n\ndef compute_msa(\n    data: dict[str, str],\n    target_id: str,\n    msa_dir: Path,\n    msa_server_url: str,\n    msa_pairing_strategy: str,\n) -> None:\n    \"\"\"Compute the MSA for the input data.\n\n    Parameters\n    ----------\n    data : dict[str, str]\n        The input protein sequences.\n    target_id : str\n        The target id.\n    msa_dir : Path\n        The msa directory.\n    msa_server_url : str\n        The MSA server URL.\n    msa_pairing_strategy : str\n        The MSA pairing strategy.\n\n    \"\"\"\n    if len(data) > 1:\n        paired_msas = run_mmseqs2(\n            list(data.values()),\n            msa_dir / f\"{target_id}_paired_tmp\",\n            use_env=True,\n            use_pairing=True,\n            host_url=msa_server_url,\n            pairing_strategy=msa_pairing_strategy,\n        )\n    else:\n        paired_msas = [\"\"] * len(data)\n\n    unpaired_msa = run_mmseqs2(\n        list(data.values()),\n        msa_dir / f\"{target_id}_unpaired_tmp\",\n        use_env=True,\n        use_pairing=False,\n        host_url=msa_server_url,\n        pairing_strategy=msa_pairing_strategy,\n    )\n\n    for idx, name in enumerate(data):\n        # Get paired sequences\n        paired = paired_msas[idx].strip().splitlines()\n        paired = paired[1::2]  # ignore headers\n        paired = paired[: const.max_paired_seqs]\n\n        # Set key per row and remove empty sequences\n        keys = [idx for idx, s in enumerate(paired) if s != \"-\" * len(s)]\n        paired = [s for s in paired if s != \"-\" * len(s)]\n\n        # Combine paired-unpaired sequences\n        unpaired = unpaired_msa[idx].strip().splitlines()\n        unpaired = unpaired[1::2]\n        unpaired = unpaired[: (const.max_msa_seqs - len(paired))]\n        if paired:\n            unpaired = unpaired[1:]  # ignore query is already present\n\n        # Combine\n        seqs = paired + unpaired\n        keys = keys + [-1] * len(unpaired)\n\n        # Dump MSA\n        csv_str = [\"key,sequence\"] + [f\"{key},{seq}\" for key, seq in zip(keys, seqs)]\n\n        msa_path = msa_dir / f\"{name}.csv\"\n        with msa_path.open(\"w\") as f:\n            f.write(\"\\n\".join(csv_str))\n\n\n@rank_zero_only\ndef process_inputs(  # noqa: C901, PLR0912, PLR0915\n    data: list[Path],\n    out_dir: Path,\n    ccd_path: Path,\n    msa_server_url: str,\n    msa_pairing_strategy: str,\n    max_msa_seqs: int = 4096,\n    use_msa_server: bool = False,\n) -> None:\n    \"\"\"Process the input data and output directory.\n\n    Parameters\n    ----------\n    data : list[Path]\n        The input data.\n    out_dir : Path\n        The output directory.\n    ccd_path : Path\n        The path to the CCD dictionary.\n    max_msa_seqs : int, optional\n        Max number of MSA sequences, by default 4096.\n    use_msa_server : bool, optional\n        Whether to use the MMSeqs2 server for MSA generation, by default False.\n\n    Returns\n    -------\n    BoltzProcessedInput\n        The processed input data.\n\n    \"\"\"\n    click.echo(\"Processing input data.\")\n    existing_records = None\n\n    # Check if manifest exists at output path\n    manifest_path = out_dir / \"processed\" / \"manifest.json\"\n    if manifest_path.exists():\n        click.echo(f\"Found a manifest file at output directory: {out_dir}\")\n\n        manifest: Manifest = Manifest.load(manifest_path)\n        input_ids = [d.stem for d in data]\n        existing_records, processed_ids = zip(\n            *[\n                (record, record.id)\n                for record in manifest.records\n                if record.id in input_ids\n            ]\n        )\n\n        if isinstance(existing_records, tuple):\n            existing_records = list(existing_records)\n\n        # Check how many examples need to be processed\n        missing = len(input_ids) - len(processed_ids)\n        if not missing:\n            click.echo(\"All examples in data are processed. Updating the manifest\")\n            # Dump updated manifest\n            updated_manifest = Manifest(existing_records)\n            updated_manifest.dump(out_dir / \"processed\" / \"manifest.json\")\n            return\n\n        click.echo(f\"{missing} missing ids. Preprocessing these ids\")\n        missing_ids = list(set(input_ids).difference(set(processed_ids)))\n        data = [d for d in data if d.stem in missing_ids]\n        assert len(data) == len(missing_ids)\n\n    # Create output directories\n    msa_dir = out_dir / \"msa\"\n    structure_dir = out_dir / \"processed\" / \"structures\"\n    processed_msa_dir = out_dir / \"processed\" / \"msa\"\n    predictions_dir = out_dir / \"predictions\"\n\n    out_dir.mkdir(parents=True, exist_ok=True)\n    msa_dir.mkdir(parents=True, exist_ok=True)\n    structure_dir.mkdir(parents=True, exist_ok=True)\n    processed_msa_dir.mkdir(parents=True, exist_ok=True)\n    predictions_dir.mkdir(parents=True, exist_ok=True)\n\n    # Load CCD\n    with ccd_path.open(\"rb\") as file:\n        ccd = pickle.load(file)  # noqa: S301\n\n    if existing_records is not None:\n        click.echo(f\"Found {len(existing_records)} records. Adding them to records\")\n\n    # Parse input data\n    records: list[Record] = existing_records if existing_records is not None else []\n    for path in tqdm(data):\n        try:\n            # Parse data\n            if path.suffix in (\".fa\", \".fas\", \".fasta\"):\n                target = parse_fasta(path, ccd)\n            elif path.suffix in (\".yml\", \".yaml\"):\n                target = parse_yaml(path, ccd)\n            elif path.is_dir():\n                msg = f\"Found directory {path} instead of .fasta or .yaml, skipping.\"\n                raise RuntimeError(msg)\n            else:\n                msg = (\n                    f\"Unable to parse filetype {path.suffix}, \"\n                    \"please provide a .fasta or .yaml file.\"\n                )\n                raise RuntimeError(msg)\n\n            # Get target id\n            target_id = target.record.id\n\n            # Get all MSA ids and decide whether to generate MSA\n            to_generate = {}\n            prot_id = const.chain_type_ids[\"PROTEIN\"]\n            for chain in target.record.chains:\n                # Add to generate list, assigning entity id\n                if (chain.mol_type == prot_id) and (chain.msa_id == 0):\n                    entity_id = chain.entity_id\n                    msa_id = f\"{target_id}_{entity_id}\"\n                    to_generate[msa_id] = target.sequences[entity_id]\n                    chain.msa_id = msa_dir / f\"{msa_id}.csv\"\n\n                # We do not support msa generation for non-protein chains\n                elif chain.msa_id == 0:\n                    chain.msa_id = -1\n\n            # Generate MSA\n            if to_generate and not use_msa_server:\n                msg = \"Missing MSA's in input and --use_msa_server flag not set.\"\n                raise RuntimeError(msg)\n\n            if to_generate:\n                msg = f\"Generating MSA for {path} with {len(to_generate)} protein entities.\"\n                click.echo(msg)\n                compute_msa(\n                    data=to_generate,\n                    target_id=target_id,\n                    msa_dir=msa_dir,\n                    msa_server_url=msa_server_url,\n                    msa_pairing_strategy=msa_pairing_strategy,\n                )\n\n            \n            # Parse MSA data\n            msas = sorted({c.msa_id for c in target.record.chains if c.msa_id != -1})\n            #print('msas: ', msas)\n            msa_id_map = {}\n            for msa_idx, msa_id in enumerate(msas):\n                # Check that raw MSA exists\n                msa_path = Path(msa_id)\n                if not msa_path.exists():\n                    msg = f\"MSA file {msa_path} not found.\"\n                    raise FileNotFoundError(msg)\n\n                # Dump processed MSA\n                processed = processed_msa_dir / f\"{target_id}_{msa_idx}.npz\"\n                msa_id_map[msa_id] = f\"{target_id}_{msa_idx}\"\n\n                print('processed: ',processed)\n                print('msa_path: ',msa_path)\n                if not processed.exists():\n                    # Parse A3M\n                    if msa_path.suffix == \".a3m\":\n                        msa: MSA = parse_a3m(\n                            msa_path,\n                            taxonomy=None,\n                            max_seqs=max_msa_seqs,\n                        )\n                    elif msa_path.suffix == \".csv\":\n                        msa: MSA = parse_csv(msa_path, max_seqs=max_msa_seqs)\n                    else:\n                        msg = f\"MSA file {msa_path} not supported, only a3m or csv.\"\n                        raise RuntimeError(msg)\n\n                    msa.dump(processed)\n\n            # Modify records to point to processed MSA\n            for c in target.record.chains:\n                if (c.msa_id != -1) and (c.msa_id in msa_id_map):\n                    c.msa_id = msa_id_map[c.msa_id]\n\n            # Keep record\n            records.append(target.record)\n\n            # Dump structure\n            struct_path = structure_dir / f\"{target.record.id}.npz\"\n            target.structure.dump(struct_path)\n\n        except Exception as e:\n            if len(data) > 1:\n                print(f\"Failed to process {path}. Skipping. Error: {e}.\")\n            else:\n                raise e\n\n    # Dump manifest\n    manifest = Manifest(records)\n    manifest.dump(out_dir / \"processed\" / \"manifest.json\")\n\ndef predict(\n    data: str,\n    out_dir: str,\n    cache: str = \"~/.boltz\",\n    checkpoint: Optional[str] = None,\n    devices: int = 1,\n    accelerator: str = \"gpu\",\n    recycling_steps: int = 3,\n    sampling_steps: int = 200,\n    diffusion_samples: int = 1,\n    step_scale: float = 1.638,\n    write_full_pae: bool = False,\n    write_full_pde: bool = False,\n    output_format: Literal[\"pdb\", \"mmcif\"] = \"mmcif\",\n    num_workers: int = 2,\n    override: bool = False,\n    seed: Optional[int] = None,\n    use_msa_server: bool = False,\n    msa_server_url: str = \"https://api.colabfold.com\",\n    msa_pairing_strategy: str = \"greedy\",\n) -> None:\n    \"\"\"Run predictions with Boltz-1.\"\"\"\n    # If cpu, write a friendly warning\n    if accelerator == \"cpu\":\n        msg = \"Running on CPU, this will be slow. Consider using a GPU.\"\n        click.echo(msg)\n\n    # Set no grad\n    torch.set_grad_enabled(False)\n\n    # Ignore matmul precision warning\n    torch.set_float32_matmul_precision(\"highest\")\n\n    # Set seed if desired\n    if seed is not None:\n       seed_everything(int(seed))\n\n    # Set cache path\n    cache = Path(cache).expanduser()\n    cache.mkdir(parents=True, exist_ok=True)\n\n    # Create output directories\n    data = Path(data).expanduser()\n    out_dir = Path(out_dir).expanduser()\n    out_dir = out_dir / f\"boltz_results_{data.stem}\"\n    out_dir.mkdir(parents=True, exist_ok=True)\n\n    # Download necessary data and model\n    download(cache)\n\n    # Validate inputs\n    data = check_inputs(data, out_dir, override)\n    if not data:\n        click.echo(\"No predictions to run, exiting.\")\n        return\n\n    # Set up trainer\n    strategy = \"auto\"\n    if (isinstance(devices, int) and devices > 1) or (\n        isinstance(devices, list) and len(devices) > 1\n    ):\n        strategy = DDPStrategy()\n        if len(data) < devices:\n            msg = (\n                \"Number of requested devices is greater \"\n                \"than the number of predictions.\"\n            )\n            raise ValueError(msg)\n\n    msg = f\"Running predictions for {len(data)} structure\"\n    msg += \"s\" if len(data) > 1 else \"\"\n    click.echo(msg)\n\n    # Process inputs\n    ccd_path = cache / \"ccd.pkl\"\n    process_inputs(\n        data=data,\n        out_dir=out_dir,\n        ccd_path=ccd_path,\n        use_msa_server=use_msa_server,\n        msa_server_url=msa_server_url,\n        msa_pairing_strategy=msa_pairing_strategy,\n    )\n\n    # Load processed data\n    processed_dir = out_dir / \"processed\"\n    processed = BoltzProcessedInput(\n        manifest=Manifest.load(processed_dir / \"manifest.json\"),\n        targets_dir=processed_dir / \"structures\",\n        msa_dir=processed_dir / \"msa\",\n    )\n\n    # Create data module\n    data_module = BoltzInferenceDataModule(\n        manifest=processed.manifest,\n        target_dir=processed.targets_dir,\n        msa_dir=processed.msa_dir,\n        num_workers=num_workers,\n    )\n\n    # Load model\n    if checkpoint is None:\n        checkpoint = cache / \"boltz1_conf.ckpt\"\n\n    predict_args = {\n        \"recycling_steps\": recycling_steps,\n        \"sampling_steps\": sampling_steps,\n        \"diffusion_samples\": diffusion_samples,\n        \"write_confidence_summary\": True,\n        \"write_full_pae\": write_full_pae,\n        \"write_full_pde\": write_full_pde,\n    }\n    diffusion_params = BoltzDiffusionParams()\n    diffusion_params.step_scale = step_scale\n    model_module: Boltz1 = Boltz1.load_from_checkpoint(\n        checkpoint,\n        strict=True,\n        predict_args=predict_args,\n        map_location=\"cpu\",\n        diffusion_process_args=asdict(diffusion_params),\n        ema=False,\n    )\n    model_module.eval()\n\n    # Create prediction writer\n    pred_writer = BoltzWriter(\n        data_dir=processed.targets_dir,\n        output_dir=out_dir / \"predictions\",\n        output_format=output_format,\n    )\n\n    trainer = Trainer(\n        default_root_dir=out_dir,\n        strategy=strategy,\n        callbacks=[pred_writer],\n        accelerator=accelerator,\n        devices=devices,\n        precision=32,\n    )\n\n    # Compute predictions\n    trainer.predict(\n        model_module,\n        datamodule=data_module,\n        return_predictions=False,\n    )\n\n\n\nif __name__ == \"__main__\":\n    \n    predict(data=\"./inputs_prediction\",\n            out_dir=\"./outputs_prediction\",\n            cache=\"/kaggle/input/rna-prediction-boltz/\",\n            diffusion_samples=5,\n            seed=42,\n            override=True)","metadata":{"execution":{"iopub.execute_input":"2025-05-25T11:34:32.449280Z","iopub.status.busy":"2025-05-25T11:34:32.449027Z","iopub.status.idle":"2025-05-25T11:34:32.460661Z","shell.execute_reply":"2025-05-25T11:34:32.459923Z"},"papermill":{"duration":0.022411,"end_time":"2025-05-25T11:34:32.461758","exception":false,"start_time":"2025-05-25T11:34:32.439347","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nsub_file = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/test_sequences.csv')\n\nsub_file.head()\n\nnames = sub_file['target_id'].tolist()\nsequences = sub_file['sequence'].tolist()\n\n# Inference\nidx = 0 \nfor tmp_id, tmp_sequence in zip(names, sequences):\n    with open(f'/kaggle/working/inputs_prediction/{tmp_id}.yaml', 'w') as f:\n        f.write(\"constraints: []\\n\")\n        f.write(\"sequences:\\n\")\n        f.write(\"- rna:\\n\")\n        f.write(\"    id:\\n\")\n        f.write(\"    - A1\\n\")\n        f.write(f\"    sequence: {tmp_sequence}\")","metadata":{"execution":{"iopub.execute_input":"2025-05-25T11:34:32.479040Z","iopub.status.busy":"2025-05-25T11:34:32.478821Z","iopub.status.idle":"2025-05-25T11:34:32.487860Z","shell.execute_reply":"2025-05-25T11:34:32.487317Z"},"papermill":{"duration":0.018914,"end_time":"2025-05-25T11:34:32.489018","exception":false,"start_time":"2025-05-25T11:34:32.470104","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%ls inputs_prediction","metadata":{"execution":{"iopub.execute_input":"2025-05-25T11:34:32.506534Z","iopub.status.busy":"2025-05-25T11:34:32.506318Z","iopub.status.idle":"2025-05-25T11:34:32.631274Z","shell.execute_reply":"2025-05-25T11:34:32.630576Z"},"papermill":{"duration":0.134616,"end_time":"2025-05-25T11:34:32.632364","exception":false,"start_time":"2025-05-25T11:34:32.497748","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\ntorch.cuda.empty_cache()\nimport gc\ngc.collect()","metadata":{"execution":{"iopub.execute_input":"2025-05-25T11:34:32.650151Z","iopub.status.busy":"2025-05-25T11:34:32.649912Z","iopub.status.idle":"2025-05-25T11:34:32.796353Z","shell.execute_reply":"2025-05-25T11:34:32.795601Z"},"papermill":{"duration":0.156803,"end_time":"2025-05-25T11:34:32.797605","exception":false,"start_time":"2025-05-25T11:34:32.640802","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python inference.py","metadata":{"execution":{"iopub.execute_input":"2025-05-25T11:34:32.816978Z","iopub.status.busy":"2025-05-25T11:34:32.816763Z","iopub.status.idle":"2025-05-25T11:54:17.739472Z","shell.execute_reply":"2025-05-25T11:54:17.738694Z"},"papermill":{"duration":1184.93342,"end_time":"2025-05-25T11:54:17.741095","exception":false,"start_time":"2025-05-25T11:34:32.807675","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from Bio.PDB.MMCIF2Dict import MMCIF2Dict\nimport json\n\ndef get_coords(tmp_id, idx):\n    cif_file = f\"outputs_prediction/boltz_results_inputs_prediction/predictions/{tmp_id}/{tmp_id}_model_{idx}.cif\"\n\n    mmcif_dict = MMCIF2Dict(cif_file)\n    \n    entity_poly_seq = mmcif_dict.get(\"_entity_poly_seq.mon_id\", [])\n    sequence = \"\".join(entity_poly_seq)\n    #print(\"RNA sequence:\", sequence)\n    \n    x_coords = mmcif_dict[\"_atom_site.Cartn_x\"]\n    y_coords = mmcif_dict[\"_atom_site.Cartn_y\"]\n    z_coords = mmcif_dict[\"_atom_site.Cartn_z\"]\n    atom_names = mmcif_dict[\"_atom_site.label_atom_id\"]\n    \n    c1_coords = []\n    for i, atom in enumerate(atom_names):\n        if atom == \"C1'\":\n            c1_coords.append((float(x_coords[i]), float(y_coords[i]), float(z_coords[i])))\n\n    conf_file = f\"outputs_prediction/boltz_results_inputs_prediction/predictions/{tmp_id}/confidence_{tmp_id}_model_{idx}.json\"\n    # parse json and read confidence_score from it's dict\n    with open(conf_file, 'r') as f:\n        conf_data = json.load(f)\n\n    confidence_score = conf_data.get('confidence_score', -1)\n        \n    return c1_coords, confidence_score\n\nall_preds = os.listdir('outputs_prediction/boltz_results_inputs_prediction/predictions')\nsubmission = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/sample_submission.csv')","metadata":{"execution":{"iopub.execute_input":"2025-05-25T11:54:17.761821Z","iopub.status.busy":"2025-05-25T11:54:17.761115Z","iopub.status.idle":"2025-05-25T11:54:17.782066Z","shell.execute_reply":"2025-05-25T11:54:17.781198Z"},"papermill":{"duration":0.032553,"end_time":"2025-05-25T11:54:17.783478","exception":false,"start_time":"2025-05-25T11:54:17.750925","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for tmp_id in all_preds:\n    print('#' * 20, f'Inferences for {tmp_id}')\n    coords_with_conf = []\n\n    for model_idx in range(5):\n        coords, conf_score = get_coords(tmp_id, model_idx)\n        coords_with_conf.append((coords, conf_score))\n\n    # 置信度从高到低排序\n    coords_with_conf.sort(key=lambda x: x[1], reverse=True)\n\n    # 写入前 5 个坐标点（按置信度顺序）\n    for rank_idx in range(5):\n        coords = coords_with_conf[rank_idx][0]\n        print(f'{rank_idx} score: ', coords_with_conf[rank_idx][1])\n        submission.loc[\n            submission['ID'].apply(lambda x: tmp_id in x),\n            [f'x_{rank_idx+1}', f'y_{rank_idx+1}', f'z_{rank_idx+1}']\n        ] = coords\n\n    print()","metadata":{"execution":{"iopub.execute_input":"2025-05-25T11:54:17.803377Z","iopub.status.busy":"2025-05-25T11:54:17.803156Z","iopub.status.idle":"2025-05-25T11:54:24.112814Z","shell.execute_reply":"2025-05-25T11:54:24.111937Z"},"papermill":{"duration":6.320816,"end_time":"2025-05-25T11:54:24.114139","exception":false,"start_time":"2025-05-25T11:54:17.793323","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%rm -rf boltz\n%rm -rf inputs_prediction\n%rm -rf outputs_prediction\n%rm -rf inference.py","metadata":{"execution":{"iopub.execute_input":"2025-05-25T11:54:24.135058Z","iopub.status.busy":"2025-05-25T11:54:24.134825Z","iopub.status.idle":"2025-05-25T11:54:24.644942Z","shell.execute_reply":"2025-05-25T11:54:24.643896Z"},"papermill":{"duration":0.521824,"end_time":"2025-05-25T11:54:24.646363","exception":false,"start_time":"2025-05-25T11:54:24.124539","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission.to_csv(\"submission_boltz.csv\", index=False)","metadata":{"execution":{"iopub.execute_input":"2025-05-25T11:54:24.668783Z","iopub.status.busy":"2025-05-25T11:54:24.668510Z","iopub.status.idle":"2025-05-25T11:54:24.714635Z","shell.execute_reply":"2025-05-25T11:54:24.714032Z"},"papermill":{"duration":0.058471,"end_time":"2025-05-25T11:54:24.715988","exception":false,"start_time":"2025-05-25T11:54:24.657517","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Template based matching","metadata":{"papermill":{"duration":0.009448,"end_time":"2025-05-25T11:54:24.774269","exception":false,"start_time":"2025-05-25T11:54:24.764821","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import time\n\nimport pandas as pd\nimport numpy as np\n\nimport random\nfrom Bio import pairwise2\nfrom Bio.Seq import Seq\n\nfrom tqdm import tqdm\n\nfrom scipy.spatial.transform import Rotation as R\nfrom sklearn.preprocessing import normalize\nfrom scipy.spatial import distance_matrix\nimport warnings\nwarnings.filterwarnings('ignore')\n\nprint(\"\\nLoading data files...\")\ntrain_seqs = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/train_sequences.csv')\nvalid_seqs = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/validation_sequences.csv')\ntest_seqs = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/test_sequences.csv')\ntrain_labels = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/train_labels.csv')\nvalid_labels = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/validation_labels.csv')\n\nprint(f\"Loaded {len(train_seqs)} training sequences, {len(valid_seqs)} validation sequences, and {len(test_seqs)} test sequences\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_seqs_v2 = pd.read_csv('/kaggle/input/rna-cif-to-csv/rna_sequences.csv')\ntrain_labels_v2 = pd.read_csv('/kaggle/input/rna-cif-to-csv/rna_coordinates.csv')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\n\n# Function to extend the original dataset with new records from v2\ndef extend_dataset(original_df, v2_df, key_columns, dataset_name):\n    print(f\"Extending {dataset_name}...\")\n    print(f\"  Original size: {len(original_df)} rows\")\n    print(f\"  v2 size: {len(v2_df)} rows\")\n    \n    # Create a composite key for identification if multiple key columns\n    if isinstance(key_columns, list) and len(key_columns) > 1:\n        original_df['temp_key'] = original_df[key_columns].astype(str).agg('_'.join, axis=1)\n        v2_df['temp_key'] = v2_df[key_columns].astype(str).agg('_'.join, axis=1)\n        key_for_identification = 'temp_key'\n    else:\n        key_for_identification = key_columns[0] if isinstance(key_columns, list) else key_columns\n    \n    # Identify unique records in each dataset\n    original_keys = set(original_df[key_for_identification])\n    v2_keys = set(v2_df[key_for_identification])\n    \n    # Calculate stats\n    keys_only_in_original = original_keys - v2_keys\n    keys_only_in_v2 = v2_keys - original_keys \n    common_keys = original_keys.intersection(v2_keys)\n    \n    print(f\"  Keys only in original: {len(keys_only_in_original)}\")\n    print(f\"  Keys only in v2: {len(keys_only_in_v2)}\")\n    print(f\"  Common keys: {len(common_keys)}\")\n    \n    # Create a mask to filter v2 records that don't exist in original\n    new_records_mask = ~v2_df[key_for_identification].isin(original_keys)\n    new_records = v2_df[new_records_mask].copy()\n    \n    # Drop temporary key if it was created\n    if key_for_identification == 'temp_key':\n        new_records.drop('temp_key', axis=1, inplace=True)\n        original_df.drop('temp_key', axis=1, inplace=True)\n    \n    # Combine original with new records from v2\n    extended_df = pd.concat([original_df, new_records], ignore_index=True)\n    \n    # Report final sizes\n    print(f\"  New records added: {len(new_records)}\")\n    print(f\"  Extended dataset size: {len(extended_df)} rows\")\n    print(f\"  Verification - All original keys in extended dataset: {set(original_df[key_columns[0] if isinstance(key_columns, list) else key_columns]).issubset(set(extended_df[key_columns[0] if isinstance(key_columns, list) else key_columns]))}\")\n    \n    # Check for missing values in key columns\n    for col in extended_df.columns:\n        original_missing = original_df[col].isnull().sum()\n        extended_missing = extended_df[col].isnull().sum()\n        if original_missing > 0 or extended_missing > 0:\n            print(f\"  Column '{col}': Missing values - Original: {original_missing}, Extended: {extended_missing}\")\n    \n    # Clean up\n    if key_for_identification == 'temp_key' and 'temp_key' in v2_df.columns:\n        v2_df.drop('temp_key', axis=1, inplace=True)\n        \n    return extended_df\n\n# 1. Extend train_seqs with train_seqs_v2\nprint(\"\\n\" + \"=\"*50)\nprint(\"EXTENDING SEQUENCE DATASETS\")\nprint(\"=\"*50)\ntrain_seqs_extended = extend_dataset(\n    train_seqs, \n    train_seqs_v2,\n    ['target_id'],  # Using target_id as the unique identifier\n    \"train_seqs\"\n)\n\n# 2. Extend train_labels with train_labels_v2\nprint(\"\\n\" + \"=\"*50)\nprint(\"EXTENDING LABELS DATASETS\")\nprint(\"=\"*50)\n# For labels, we need a composite key of ID and resid\ntrain_labels_extended = extend_dataset(\n    train_labels,\n    train_labels_v2,\n    ['ID', 'resid'],  # Using composite key\n    \"train_labels\"\n)\n\n# Verify relationships between extended datasets\nprint(\"\\n\" + \"=\"*50)\nprint(\"VERIFYING RELATIONSHIPS\")\nprint(\"=\"*50)\n\n# Check if all sequence IDs have corresponding labels\nseq_ids = set(train_seqs_extended['target_id'].unique())\nlabel_ids = set(train_labels_extended['ID'].unique())\n\nseq_ids_with_labels = seq_ids.intersection(label_ids)\nseq_ids_without_labels = seq_ids - label_ids\n\nprint(f\"Total unique sequence IDs: {len(seq_ids)}\")\nprint(f\"Sequence IDs with corresponding labels: {len(seq_ids_with_labels)} ({len(seq_ids_with_labels)/len(seq_ids)*100:.2f}%)\")\nprint(f\"Sequence IDs without corresponding labels: {len(seq_ids_without_labels)} ({len(seq_ids_without_labels)/len(seq_ids)*100:.2f}%)\")\n\nif len(seq_ids_without_labels) > 0:\n    print(\"Sample of sequence IDs without labels (up to 5):\")\n    print(list(seq_ids_without_labels)[:5])\n\n# Print summary of extended datasets\nprint(\"\\n\" + \"=\"*50)\nprint(\"SUMMARY OF EXTENDED DATASETS\")\nprint(\"=\"*50)\nprint(f\"Original train_seqs: {len(train_seqs)} rows\")\nprint(f\"Original train_labels: {len(train_labels)} rows\")\nprint(f\"Extended train_seqs: {len(train_seqs_extended)} rows (+{len(train_seqs_extended)-len(train_seqs)})\")\nprint(f\"Extended train_labels: {len(train_labels_extended)} rows (+{len(train_labels_extended)-len(train_labels)})\")\n\n# Save the extended datasets (uncomment to save)\n# train_seqs_extended.to_csv('train_seqs_combined.csv', index=False)\n# train_labels_extended.to_csv('train_labels_combined.csv', index=False)\n\nprint(\"\\n\" + \"=\"*50)\nprint(\"DONE! Extended datasets created.\")\nprint(\"To save the datasets, uncomment the last two lines.\")\nprint(\"=\"*50)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def process_labels(labels_df):\n    coords_dict = {}\n    \n    # Group by target ID and wrap with tqdm for progress tracking\n    id_groups = labels_df.groupby(lambda x: labels_df['ID'][x].rsplit('_', 1)[0])\n    for id_prefix, group in tqdm(id_groups, desc=\"Processing structures\"):\n        # Extract just the coordinates columns for the first structure (x_1, y_1, z_1)\n        coords = []\n        for _, row in group.sort_values('resid').iterrows():\n            coords.append([row['x_1'], row['y_1'], row['z_1']])\n        \n        coords_dict[id_prefix] = np.array(coords)\n    \n    return coords_dict\n\ntrain_coords_dict = process_labels(train_labels_extended)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from Bio.Seq import Seq\nfrom Bio import pairwise2\nimport numpy as np\nfrom sklearn.cluster import KMeans\nfrom sklearn.metrics.pairwise import cosine_similarity\n\ndef find_similar_sequences(query_seq, train_seqs_df, train_coords_dict, top_n=5):\n    \"\"\"\n    Find similar RNA sequences using enhanced scoring and clustering for diversity.\n    \n    Improvements:\n    - Multi-tier length filtering\n    - Enhanced alignment scoring with multiple algorithms\n    - RNA-specific structural features\n    - Adaptive clustering\n    \"\"\"\n    similar_seqs = []\n    query_seq_obj = Seq(query_seq)\n    query_features = _extract_enhanced_rna_features(query_seq)\n    \n    # Step 1: Enhanced candidate selection with multi-tier filtering\n    for _, row in train_seqs_df.iterrows():\n        target_id = row['target_id']\n        train_seq = row['sequence']\n        \n        # Skip if coordinates not available\n        if target_id not in train_coords_dict:\n            continue\n        \n        # Multi-tier length filtering (more permissive for very short/long sequences)\n        len_ratio = abs(len(train_seq) - len(query_seq)) / max(len(train_seq), len(query_seq))\n        if len(query_seq) < 50 or len(train_seq) < 50:  # Short sequences - more permissive\n            if len_ratio > 0.6:\n                continue\n        elif len(query_seq) > 1000 or len(train_seq) > 1000:  # Long sequences - stricter\n            if len_ratio > 0.2:\n                continue\n        else:  # Medium sequences - original threshold\n            if len_ratio > 0.4:\n                continue\n        \n        # Calculate composite similarity score\n        composite_score = _calculate_composite_similarity(query_seq, train_seq, query_features)\n        \n        if composite_score > 0:  # Only keep sequences with positive similarity\n            similar_seqs.append((target_id, train_seq, composite_score, train_coords_dict[target_id]))\n    \n    # Sort by composite score and take top candidates\n    similar_seqs.sort(key=lambda x: x[2], reverse=True)\n    \n    # Adaptive candidate selection based on score distribution\n    candidate_count = min(50, len(similar_seqs))  # Increased initial pool\n    if len(similar_seqs) > 10:\n        # Filter out sequences with very low scores (bottom 20%)\n        score_threshold = np.percentile([x[2] for x in similar_seqs], 80)\n        filtered_candidates = [x for x in similar_seqs if x[2] >= score_threshold]\n        candidate_count = min(candidate_count, len(filtered_candidates))\n        top_candidates = filtered_candidates[:candidate_count]\n    else:\n        top_candidates = similar_seqs[:candidate_count]\n    \n    # If we have fewer sequences than requested clusters, return all\n    if len(top_candidates) <= top_n:\n        return top_candidates[:top_n]\n    \n    # Step 2: Enhanced feature matrix for better clustering\n    feature_matrix = []\n    for _, seq, _, _ in top_candidates:\n        features = _extract_enhanced_rna_features(seq)\n        feature_matrix.append(features)\n    \n    feature_matrix = np.array(feature_matrix)\n    \n    # Step 3: Adaptive clustering\n    n_clusters = min(top_n, len(top_candidates))\n    \n    # Use different clustering approach based on dataset size\n    if len(top_candidates) >= 15:\n        # K-means for larger datasets\n        kmeans = KMeans(n_clusters=n_clusters, random_state=42, n_init=10)\n        cluster_labels = kmeans.fit_predict(feature_matrix)\n    else:\n        # Simple diversity-based selection for smaller datasets\n        cluster_labels = _diversity_based_clustering(feature_matrix, n_clusters)\n    \n    # Step 4: Select best representative from each cluster\n    final_results = []\n    for cluster_id in range(n_clusters):\n        cluster_sequences = [top_candidates[i] for i in range(len(top_candidates)) \n                           if cluster_labels[i] == cluster_id]\n        \n        if cluster_sequences:\n            # Sort by composite score and take the best one\n            cluster_sequences.sort(key=lambda x: x[2], reverse=True)\n            final_results.append(cluster_sequences[0])\n    \n    # Sort final results by similarity score\n    final_results.sort(key=lambda x: x[2], reverse=True)\n    \n    return final_results[:top_n]\n\ndef _calculate_composite_similarity(query_seq, train_seq, query_features):\n    \"\"\"\n    Calculate composite similarity using multiple alignment methods and features.\n    \"\"\"\n    query_seq_obj = Seq(query_seq)\n    \n    # 1. Global alignment (original method)\n    global_alignments = pairwise2.align.globalms(query_seq_obj, train_seq, 2.9, -1, -10, -0.5, one_alignment_only=True)\n    global_score = 0\n    if global_alignments:\n        alignment = global_alignments[0]\n        global_score = alignment.score / (2 * min(len(query_seq), len(train_seq)))\n    \n    # 2. Local alignment for finding similar regions\n    local_alignments = pairwise2.align.localms(query_seq_obj, train_seq, 2.9, -1, -10, -0.5, one_alignment_only=True)\n    local_score = 0\n    if local_alignments:\n        alignment = local_alignments[0]\n        local_score = alignment.score / (2 * min(len(query_seq), len(train_seq)))\n    \n    # 3. Feature-based similarity\n    train_features = _extract_enhanced_rna_features(train_seq)\n    feature_similarity = cosine_similarity([query_features], [train_features])[0][0]\n    \n    # 4. K-mer similarity for sequence motifs\n    kmer_similarity = _calculate_kmer_similarity(query_seq, train_seq, k=3)\n    \n    # Weighted composite score\n    composite_score = (\n        0.4 * global_score + \n        0.3 * local_score + \n        0.2 * feature_similarity + \n        0.1 * kmer_similarity\n    )\n    \n    return composite_score\n\ndef _calculate_kmer_similarity(seq1, seq2, k=3):\n    \"\"\"Calculate k-mer based similarity between sequences.\"\"\"\n    def get_kmers(seq, k):\n        return set(seq[i:i+k] for i in range(len(seq) - k + 1))\n    \n    kmers1 = get_kmers(seq1.upper(), k)\n    kmers2 = get_kmers(seq2.upper(), k)\n    \n    if not kmers1 or not kmers2:\n        return 0\n    \n    intersection = len(kmers1.intersection(kmers2))\n    union = len(kmers1.union(kmers2))\n    \n    return intersection / union if union > 0 else 0\n\ndef _diversity_based_clustering(feature_matrix, n_clusters):\n    \"\"\"Simple diversity-based clustering for small datasets.\"\"\"\n    n_samples = len(feature_matrix)\n    cluster_labels = np.zeros(n_samples, dtype=int)\n    \n    if n_samples <= n_clusters:\n        return np.arange(n_samples)\n    \n    # Select diverse representatives\n    selected_indices = [0]  # Start with first sequence\n    \n    for cluster_id in range(1, n_clusters):\n        max_min_distance = -1\n        best_idx = -1\n        \n        for i in range(n_samples):\n            if i in selected_indices:\n                continue\n            \n            # Find minimum distance to already selected sequences\n            min_distance = min(\n                np.linalg.norm(feature_matrix[i] - feature_matrix[j]) \n                for j in selected_indices\n            )\n            \n            if min_distance > max_min_distance:\n                max_min_distance = min_distance\n                best_idx = i\n        \n        if best_idx != -1:\n            selected_indices.append(best_idx)\n    \n    # Assign remaining sequences to closest cluster centers\n    for i in range(n_samples):\n        if i not in selected_indices:\n            distances = [\n                np.linalg.norm(feature_matrix[i] - feature_matrix[j]) \n                for j in selected_indices\n            ]\n            cluster_labels[i] = np.argmin(distances)\n        else:\n            cluster_labels[i] = selected_indices.index(i)\n    \n    return cluster_labels\n\ndef _extract_enhanced_rna_features(sequence):\n    \"\"\"\n    Extract comprehensive RNA-specific features for better clustering and similarity.\n    \"\"\"\n    seq = sequence.upper()\n    features = []\n    \n    # 1. Basic nucleotide frequencies\n    nucleotides = ['A', 'U', 'G', 'C']\n    for nuc in nucleotides:\n        freq = seq.count(nuc) / len(seq) if len(seq) > 0 else 0\n        features.append(freq)\n    \n    # 2. Dinucleotide frequencies (reduced set - most important for RNA)\n    important_dinucs = ['AU', 'UA', 'GC', 'CG', 'GU', 'UG', 'AA', 'UU', 'GG', 'CC']\n    for dinuc in important_dinucs:\n        count = 0\n        for i in range(len(seq) - 1):\n            if seq[i:i+2] == dinuc:\n                count += 1\n        freq = count / (len(seq) - 1) if len(seq) > 1 else 0\n        features.append(freq)\n    \n    # 3. RNA secondary structure indicators\n    gc_content = (seq.count('G') + seq.count('C')) / len(seq) if len(seq) > 0 else 0\n    au_content = (seq.count('A') + seq.count('U')) / len(seq) if len(seq) > 0 else 0\n    purine_content = (seq.count('A') + seq.count('G')) / len(seq) if len(seq) > 0 else 0\n    pyrimidine_content = (seq.count('U') + seq.count('C')) / len(seq) if len(seq) > 0 else 0\n    \n    features.extend([gc_content, au_content, purine_content, pyrimidine_content])\n    \n    # 4. Sequence complexity measures\n    length_normalized = min(len(seq) / 1000.0, 1.0)  # Capped normalization\n    \n    # Simple entropy calculation\n    entropy = 0\n    for nuc in nucleotides:\n        freq = seq.count(nuc) / len(seq) if len(seq) > 0 else 0\n        if freq > 0:\n            entropy -= freq * np.log2(freq)\n    entropy_normalized = entropy / 2.0  # Max entropy for 4 nucleotides is 2\n    \n    features.extend([length_normalized, entropy_normalized])\n    \n    # 5. Repetitive pattern detection\n    repeat_content = _calculate_repeat_content(seq)\n    features.append(repeat_content)\n    \n    return features\n\ndef _calculate_repeat_content(sequence):\n    \"\"\"Calculate the proportion of repetitive content in the sequence.\"\"\"\n    if len(sequence) < 6:\n        return 0\n    \n    repeat_count = 0\n    window_size = 3\n    \n    for i in range(len(sequence) - window_size + 1):\n        motif = sequence[i:i + window_size]\n        # Look for the same motif in the rest of the sequence\n        for j in range(i + window_size, len(sequence) - window_size + 1):\n            if sequence[j:j + window_size] == motif:\n                repeat_count += 1\n                break\n    \n    return repeat_count / (len(sequence) - window_size + 1) if len(sequence) > window_size else 0","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def adaptive_rna_constraints(coordinates, sequence, confidence=1.0):\n    # Make a copy of coordinates to refine\n    refined_coords = coordinates.copy()\n    n_residues = len(sequence)\n    \n    # Calculate constraint strength (inverse of confidence)\n    # High confidence templates receive gentler constraints\n    constraint_strength = 0.8 * (1.0 - min(confidence, 0.8))\n    \n    # 1. Sequential distance constraints (consecutive nucleotides)\n    # More flexible distance range (statistical distribution from PDB)\n    seq_min_dist = 5.5  # Minimum sequential distance\n    seq_max_dist = 6.5  # Maximum sequential distance\n    \n    for i in range(n_residues - 1):\n        current_pos = refined_coords[i]\n        next_pos = refined_coords[i+1]\n        \n        # Calculate current distance\n        current_dist = np.linalg.norm(next_pos - current_pos)\n        \n        # Only adjust if significantly outside expected range\n        if current_dist < seq_min_dist or current_dist > seq_max_dist:\n            # Calculate target distance (midpoint of range)\n            target_dist = (seq_min_dist + seq_max_dist) / 2\n            \n            # Get direction vector\n            direction = next_pos - current_pos\n            direction = direction / (np.linalg.norm(direction) + 1e-10)\n            \n            # Apply partial adjustment based on constraint strength\n            adjustment = (target_dist - current_dist) * constraint_strength\n            \n            # Only adjust the next position to preserve the overall fold\n            refined_coords[i+1] = current_pos + direction * (current_dist + adjustment)\n    \n    # 2. Steric clash prevention (more conservative)\n    min_allowed_distance = 3.8  # Minimum distance between non-consecutive C1' atoms\n    \n    # Calculate all pairwise distances\n    dist_matrix = distance_matrix(refined_coords, refined_coords)\n    \n    # Find severe clashes (atoms too close)\n    severe_clashes = np.where((dist_matrix < min_allowed_distance) & (dist_matrix > 0))\n    \n    # Fix severe clashes\n    for idx in range(len(severe_clashes[0])):\n        i, j = severe_clashes[0][idx], severe_clashes[1][idx]\n        \n        # Skip consecutive nucleotides and previously processed pairs\n        if abs(i - j) <= 1 or i >= j:\n            continue\n            \n        # Get current positions and distance\n        pos_i = refined_coords[i]\n        pos_j = refined_coords[j]\n        current_dist = dist_matrix[i, j]\n        \n        # Calculate necessary adjustment but scale by constraint strength\n        direction = pos_j - pos_i\n        direction = direction / (np.linalg.norm(direction) + 1e-10)\n        \n        # Calculate partial adjustment\n        adjustment = (min_allowed_distance - current_dist) * constraint_strength\n        \n        # Move points apart\n        refined_coords[i] = pos_i - direction * (adjustment / 2)\n        refined_coords[j] = pos_j + direction * (adjustment / 2)\n    \n    # 3. Very light base-pair constraining (if confidence is low)\n    if constraint_strength > 0.3:  # Only apply if template confidence is low\n        # Simple Watson-Crick base pairs\n        pairs = {'A': 'U', 'U': 'A', 'G': 'C', 'C': 'G'}\n        \n        # Scan for potential base pairs\n        for i in range(n_residues):\n            base_i = sequence[i]\n            complement = pairs.get(base_i)\n            \n            if not complement:\n                continue\n                \n            # Look for complementary bases within a reasonable range\n            for j in range(i + 3, min(i + 20, n_residues)):\n                if sequence[j] == complement:\n                    # Calculate current distance\n                    current_dist = np.linalg.norm(refined_coords[i] - refined_coords[j])\n                    \n                    # Only consider if distance suggests potential pairing\n                    if 8.0 < current_dist < 14.0:\n                        # Target 10.5Å as generic base-pair C1'-C1' distance\n                        target_dist = 10.5\n                        \n                        # Calculate very gentle adjustment (scaled by constraint_strength)\n                        adjustment = (target_dist - current_dist) * (constraint_strength * 0.3)\n                        \n                        # Get direction vector\n                        direction = refined_coords[j] - refined_coords[i]\n                        direction = direction / (np.linalg.norm(direction) + 1e-10)\n                        \n                        # Apply very gentle adjustment to both positions\n                        refined_coords[i] = refined_coords[i] - direction * (adjustment / 2)\n                        refined_coords[j] = refined_coords[j] + direction * (adjustment / 2)\n                        \n                        # Only consider one potential pair per base (closest match)\n                        break\n    \n    return refined_coords","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def adapt_template_to_query(query_seq, template_seq, template_coords, alignment=None):\n    if alignment is None:\n        from Bio.Seq import Seq\n        from Bio import pairwise2\n        \n        query_seq_obj = Seq(query_seq)\n        template_seq_obj = Seq(template_seq)\n        alignments = pairwise2.align.globalms(query_seq_obj, template_seq_obj, 2.9, -1, -10, -0.5, one_alignment_only=True)\n        \n        if not alignments:\n            return generate_improved_rna_structure(query_seq)\n            \n        alignment = alignments[0]\n    \n    aligned_query = alignment.seqA\n    aligned_template = alignment.seqB\n    \n    query_coords = np.zeros((len(query_seq), 3))\n    query_coords.fill(np.nan)\n    \n    # Map template coordinates to query\n    query_idx = 0\n    template_idx = 0\n    \n    for i in range(len(aligned_query)):\n        query_char = aligned_query[i]\n        template_char = aligned_template[i]\n        \n        if query_char != '-' and template_char != '-':\n            if template_idx < len(template_coords):\n                query_coords[query_idx] = template_coords[template_idx]\n            template_idx += 1\n            query_idx += 1\n        elif query_char != '-' and template_char == '-':\n            query_idx += 1\n        elif query_char == '-' and template_char != '-':\n            template_idx += 1\n    \n    # IMPROVED GAP FILLING - maintains RNA backbone geometry\n    backbone_distance = 5.9  # Typical C1'-C1' distance\n    \n    # Fill gaps by maintaining realistic backbone connectivity\n    for i in range(len(query_coords)):\n        if np.isnan(query_coords[i, 0]):\n            # Find nearest valid neighbors\n            prev_valid = next_valid = None\n            \n            for j in range(i-1, -1, -1):\n                if not np.isnan(query_coords[j, 0]):\n                    prev_valid = j\n                    break\n                    \n            for j in range(i+1, len(query_coords)):\n                if not np.isnan(query_coords[j, 0]):\n                    next_valid = j\n                    break\n            \n            if prev_valid is not None and next_valid is not None:\n                # Interpolate along realistic RNA backbone path\n                gap_size = next_valid - prev_valid\n                total_distance = np.linalg.norm(query_coords[next_valid] - query_coords[prev_valid])\n                expected_distance = gap_size * backbone_distance\n                \n                # If gap is compressed, extend it realistically\n                if total_distance < expected_distance * 0.7:\n                    direction = query_coords[next_valid] - query_coords[prev_valid]\n                    direction = direction / (np.linalg.norm(direction) + 1e-10)\n                    \n                    # Place intermediate points along extended path\n                    for k, idx in enumerate(range(prev_valid + 1, next_valid)):\n                        progress = (k + 1) / gap_size\n                        base_pos = query_coords[prev_valid] + direction * expected_distance * progress\n                        \n                        # Add slight curvature for realism\n                        perpendicular = np.cross(direction, [0, 0, 1])\n                        if np.linalg.norm(perpendicular) < 1e-6:\n                            perpendicular = np.cross(direction, [1, 0, 0])\n                        perpendicular = perpendicular / (np.linalg.norm(perpendicular) + 1e-10)\n                        \n                        curve_amplitude = 2.0 * np.sin(progress * np.pi)\n                        query_coords[idx] = base_pos + perpendicular * curve_amplitude\n                else:\n                    # Linear interpolation for normal gaps\n                    for k, idx in enumerate(range(prev_valid + 1, next_valid)):\n                        weight = (k + 1) / gap_size\n                        query_coords[idx] = (1 - weight) * query_coords[prev_valid] + weight * query_coords[next_valid]\n            \n            elif prev_valid is not None:\n                # Extend from previous position\n                if prev_valid > 0 and not np.isnan(query_coords[prev_valid-1, 0]):\n                    direction = query_coords[prev_valid] - query_coords[prev_valid-1]\n                    direction = direction / (np.linalg.norm(direction) + 1e-10)\n                else:\n                    direction = np.array([1.0, 0.0, 0.0])\n                \n                steps_needed = i - prev_valid\n                for step in range(1, steps_needed + 1):\n                    pos_idx = prev_valid + step\n                    if pos_idx < len(query_coords):\n                        query_coords[pos_idx] = query_coords[prev_valid] + direction * backbone_distance * step\n            \n            elif next_valid is not None:\n                # Work backwards from next position\n                direction = np.array([-1.0, 0.0, 0.0])  # Default backward direction\n                steps_needed = next_valid - i\n                for step in range(steps_needed, 0, -1):\n                    pos_idx = next_valid - step\n                    if pos_idx >= 0:\n                        query_coords[pos_idx] = query_coords[next_valid] - direction * backbone_distance * step\n    \n    # Final cleanup\n    query_coords = np.nan_to_num(query_coords)\n    return query_coords","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def generate_improved_rna_structure(sequence):\n    \"\"\"\n    Generate a more realistic RNA structure fallback based on sequence patterns\n    and basic RNA structure principles.\n    \n    Args:\n        sequence: RNA sequence string\n        \n    Returns:\n        Array of 3D coordinates\n    \"\"\"\n    n_residues = len(sequence)\n    coordinates = np.zeros((n_residues, 3))\n    \n    # Analyze sequence to predict structural elements\n    # Look for complementary regions that could form base pairs\n    potential_stems = identify_potential_stems(sequence)\n    \n    # Default parameters\n    radius_helix = 10.0\n    radius_loop = 15.0\n    rise_per_residue_helix = 2.5\n    rise_per_residue_loop = 1.5\n    angle_per_residue_helix = 0.6\n    angle_per_residue_loop = 0.3\n    \n    # Assign structural classifications\n    structure_types = assign_structure_types(sequence, potential_stems)\n    \n    # Generate coordinates based on predicted structure\n    current_pos = np.array([0.0, 0.0, 0.0])\n    current_direction = np.array([0.0, 0.0, 1.0])\n    current_angle = 0.0\n    \n    for i in range(n_residues):\n        if structure_types[i] == 'stem':\n            # Part of a helical stem\n            current_angle += angle_per_residue_helix\n            coordinates[i] = [\n                radius_helix * np.cos(current_angle), \n                radius_helix * np.sin(current_angle), \n                current_pos[2] + rise_per_residue_helix\n            ]\n            current_pos = coordinates[i]\n        elif structure_types[i] == 'loop':\n            # Part of a loop\n            current_angle += angle_per_residue_loop\n            z_shift = rise_per_residue_loop * np.sin(current_angle * 0.5)\n            coordinates[i] = [\n                radius_loop * np.cos(current_angle), \n                radius_loop * np.sin(current_angle), \n                current_pos[2] + z_shift\n            ]\n            current_pos = coordinates[i]\n        else:\n            # Single-stranded region\n            # Add some randomness to make it look more realistic\n            jitter = np.random.normal(0, 1, 3) * 2.0\n            coordinates[i] = current_pos + jitter\n            current_pos = coordinates[i]\n            \n    return coordinates\n\ndef identify_potential_stems(sequence):\n    \"\"\"\n    Identify potential stem regions by looking for self-complementary segments.\n    \n    Args:\n        sequence: RNA sequence string\n        \n    Returns:\n        List of tuples (start1, end1, start2, end2) representing potentially paired regions\n    \"\"\"\n    complementary_bases = {'A': 'U', 'U': 'A', 'G': 'C', 'C': 'G'}\n    min_stem_length = 3\n    potential_stems = []\n    \n    # Simple stem identification\n    for i in range(len(sequence) - min_stem_length):\n        for j in range(i + min_stem_length + 3, len(sequence) - min_stem_length + 1):\n            # Check if regions could form a stem\n            potential_stem_len = min(min_stem_length, len(sequence) - j)\n            is_stem = True\n            \n            for k in range(potential_stem_len):\n                if sequence[i+k] not in complementary_bases or \\\n                   complementary_bases[sequence[i+k]] != sequence[j+potential_stem_len-k-1]:\n                    is_stem = False\n                    break\n            \n            if is_stem:\n                potential_stems.append((i, i+potential_stem_len-1, j, j+potential_stem_len-1))\n    \n    return potential_stems\n\ndef assign_structure_types(sequence, potential_stems):\n    \"\"\"\n    Assign each nucleotide to a structural element type.\n    \n    Args:\n        sequence: RNA sequence string\n        potential_stems: List of tuples representing stem regions\n        \n    Returns:\n        List of structure types ('stem', 'loop', 'single')\n    \"\"\"\n    structure_types = ['single'] * len(sequence)\n    \n    # Mark stem regions\n    for stem in potential_stems:\n        start1, end1, start2, end2 = stem\n        for i in range(end1 - start1 + 1):\n            structure_types[start1 + i] = 'stem'\n            structure_types[end2 - i] = 'stem'\n    \n    # Mark loop regions (regions between paired regions)\n    for i in range(len(potential_stems) - 1):\n        _, end1, start2, _ = potential_stems[i]\n        next_start1, _, _, _ = potential_stems[i+1]\n        \n        if next_start1 > end1 + 1 and start2 > next_start1:\n            for j in range(end1 + 1, next_start1):\n                structure_types[j] = 'loop'\n    \n    return structure_types","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Function to create a more realistic RNA structure when no good templates are found\ndef generate_rna_structure(sequence, seed=None):\n    if seed is not None:\n        np.random.seed(seed)\n        random.seed(seed)\n    \n    n_residues = len(sequence)\n    coordinates = np.zeros((n_residues, 3))\n    \n    # Initialize the first few residues in a helix\n    for i in range(min(3, n_residues)):\n        angle = i * 0.6\n        coordinates[i] = [10.0 * np.cos(angle), 10.0 * np.sin(angle), i * 2.5]\n    \n    # Add more complex folding patterns\n    current_direction = np.array([0.0, 0.0, 1.0])  # Start moving along z-axis\n    \n    # Define base-pairing tendencies (G-C and A-U pairs)\n    for i in range(3, n_residues):\n        # Check for potential base-pairing in the sequence\n        has_pair = False\n        pair_idx = -1\n        \n        # Simple detection of complementary bases (G-C, A-U)\n        complementary = {'G': 'C', 'C': 'G', 'A': 'U', 'U': 'A'}\n        current_base = sequence[i]\n        \n        # Look for potential base-pairing within a window before the current position\n        window_size = min(i, 15)  # Look back up to 15 bases\n        for j in range(i-window_size, i):\n            if j >= 0 and sequence[j] == complementary.get(current_base, 'X'):\n                # Found a potential pair\n                has_pair = True\n                pair_idx = j\n                break\n        \n        if has_pair and i - pair_idx <= 10 and random.random() < 0.7:\n            # Try to create a base-pair by positioning this nucleotide near its pair\n            pair_pos = coordinates[pair_idx]\n            \n            # Create a position that's roughly opposite to the pair\n            random_offset = np.random.normal(0, 1, 3) * 2.0\n            base_pair_distance = 10.0 + random.uniform(-1.0, 1.0)\n            \n            # Calculate a vector from base-pair toward center of structure\n            center = np.mean(coordinates[:i], axis=0)\n            direction = center - pair_pos\n            direction = direction / (np.linalg.norm(direction) + 1e-10)\n            \n            # Position new nucleotide in the general direction of the \"center\"\n            coordinates[i] = pair_pos + direction * base_pair_distance + random_offset\n            \n            # Update direction for next nucleotide\n            current_direction = np.random.normal(0, 0.3, 3)\n            current_direction = current_direction / (np.linalg.norm(current_direction) + 1e-10)\n            \n        else:\n            # No base-pairing detected, continue with the current fold direction\n            # Randomly rotate current direction to simulate RNA flexibility\n            if random.random() < 0.3:\n                # More significant direction change\n                angle = random.uniform(0.2, 0.6)\n                axis = np.random.normal(0, 1, 3)\n                axis = axis / (np.linalg.norm(axis) + 1e-10)\n                rotation = R.from_rotvec(angle * axis)\n                current_direction = rotation.apply(current_direction)\n            else:\n                # Small random changes in direction\n                current_direction += np.random.normal(0, 0.15, 3)\n                current_direction = current_direction / (np.linalg.norm(current_direction) + 1e-10)\n            \n            # Distance between consecutive nucleotides (3.5-4.5Å is typical)\n            step_size = random.uniform(3.5, 4.5)\n            \n            # Update position\n            coordinates[i] = coordinates[i-1] + step_size * current_direction\n    \n    return coordinates","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def predict_rna_structures(sequence, target_id, train_seqs_df, train_coords_dict, n_predictions=5):\n    predictions = []\n    \n    # Find similar sequences in the training data\n    similar_seqs = find_similar_sequences(sequence, train_seqs_df, train_coords_dict, top_n=n_predictions)\n    \n    # If we found any similar sequences, use them as templates\n    if similar_seqs:\n        for i, (template_id, template_seq, similarity_score, template_coords) in enumerate(similar_seqs):\n            # Adapt template coordinates to the query sequence\n            adapted_coords = adapt_template_to_query(sequence, template_seq, template_coords)\n            \n            if adapted_coords is not None:\n                # Apply adaptive constraints based on template similarity\n                # For high similarity templates, apply very gentle constraints\n                refined_coords = adaptive_rna_constraints(adapted_coords, sequence, confidence=similarity_score)\n                \n                # Add some randomness (less for better templates)\n                random_scale = max(0.05, 0.8 - similarity_score)  # Reduced randomness\n                randomized_coords = refined_coords.copy()\n                randomized_coords += np.random.normal(0, random_scale, randomized_coords.shape)\n                \n                predictions.append(randomized_coords)\n                \n                if len(predictions) >= n_predictions:\n                    break\n    \n    # If we don't have enough predictions from templates, generate de novo structures\n    while len(predictions) < n_predictions:\n        seed_value = hash(target_id) % 10000 + len(predictions) * 1000\n        de_novo_coords = generate_rna_structure(sequence, seed=seed_value)\n        \n        # Apply stronger constraints to de novo structures (lower confidence)\n        refined_de_novo = adaptive_rna_constraints(de_novo_coords, sequence, confidence=0.2)\n        \n        predictions.append(refined_de_novo)\n    \n    return predictions[:n_predictions]","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# List to store all prediction records\nall_predictions = []\n\n# Set up time tracking\nstart_time = time.time()\ntotal_targets = len(test_seqs)\n\n# For each sequence in the test set\nfor idx, row in test_seqs.iterrows():\n    target_id = row['target_id']\n    sequence = row['sequence']\n    \n    # Progress tracking\n    if idx % 5 == 0:\n        elapsed = time.time() - start_time\n        targets_processed = idx + 1\n        if targets_processed > 0:\n            avg_time_per_target = elapsed / targets_processed\n            est_time_remaining = avg_time_per_target * (total_targets - targets_processed)\n            print(f\"Processing target {targets_processed}/{total_targets}: {target_id} ({len(sequence)} nt), \"\n                  f\"elapsed: {elapsed:.1f}s, est. remaining: {est_time_remaining:.1f}s\")\n    \n    # Generate 5 different structure predictions\n    predictions = predict_rna_structures(sequence, target_id, train_seqs_extended, train_coords_dict, n_predictions=5)\n    \n    # For each residue in the sequence\n    for j in range(len(sequence)):\n        pred_row = {\n            'ID': f\"{target_id}_{j+1}\",\n            'resname': sequence[j],\n            'resid': j + 1\n        }\n        \n        # Add coordinates from all 5 predictions\n        for i in range(5):\n            pred_row[f'x_{i+1}'] = predictions[i][j][0]\n            pred_row[f'y_{i+1}'] = predictions[i][j][1]\n            pred_row[f'z_{i+1}'] = predictions[i][j][2]\n        \n        all_predictions.append(pred_row)\n\n# Create DataFrame with predictions\nsubmission_df = pd.DataFrame(all_predictions)\n\n# Ensure the submission file has the correct format\ncolumn_order = ['ID', 'resname', 'resid']\nfor i in range(1, 6):\n    for coord in ['x', 'y', 'z']:\n        column_order.append(f'{coord}_{i}')\nsubmission_df = submission_df[column_order]\n\n# Save the submission file\nsubmission_df.to_csv('submission_TBM.csv', index=False)\nprint(f\"Generated predictions for {len(test_seqs)} RNA sequences\")\nprint(f\"Total runtime: {time.time() - start_time:.1f} seconds\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission_df","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Ensemble","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\n\n# ======================================\n# 1️⃣ Load Submissions\n# ======================================\ntbm_path = \"/kaggle/working/submission_TBM.csv\"\nprotenix_path = \"/kaggle/working/submission_proteinx.csv\"\nnufold_path = \"/kaggle/working/submission_boltz.csv\"\n\n# Read CSVs\ntbm = pd.read_csv(tbm_path)\nprotenix = pd.read_csv(protenix_path)\nnufold = pd.read_csv(nufold_path)\n\nprint(\"✅ Files loaded:\")\nprint(\"TBM:\", tbm.shape)\nprint(\"Protenix:\", protenix.shape)\nprint(\"NuFold:\", nufold.shape)\n\n# ======================================\n# 2️⃣ Validate and Extract Target IDs\n# ======================================\nassert set(tbm[\"ID\"]) == set(protenix[\"ID\"]) == set(nufold[\"ID\"]), \\\n    \"❌ Mismatch in residue IDs among models!\"\n\ndef extract_target_id(resid):\n    return resid.split(\"_\")[0]\n\nfor df in [tbm, protenix, nufold]:\n    df[\"target_id\"] = df[\"ID\"].apply(extract_target_id)\n\nprint(\"✅ target_id extracted successfully.\")\n\n# ======================================\n# 3️⃣ Helper Functions\n# ======================================\ndef extract_structures(df, tid):\n    \"\"\"Extract (5, n_res, 3) array of conformations for a target.\"\"\"\n    sub = df[df[\"target_id\"] == tid].sort_values(\"resid\")\n    coords = [sub[[f\"x_{i}\", f\"y_{i}\", f\"z_{i}\"]].values for i in range(1, 6)]\n    return np.stack(coords, axis=0)\n\ndef kabsch_rmsd(P, Q):\n    \"\"\"Compute RMSD between two 3D conformations after alignment.\"\"\"\n    P, Q = P - P.mean(0), Q - Q.mean(0)\n    U, _, Vt = np.linalg.svd(P.T @ Q)\n    R = U @ Vt\n    P_aligned = P @ R\n    return np.sqrt(np.mean(np.sum((P_aligned - Q)**2, axis=1)))\n\ndef mean_rmsd(setA):\n    \"\"\"Mean pairwise RMSD for diversity scoring.\"\"\"\n    if len(setA) < 2: return 0\n    rmsd_vals = []\n    for i in range(len(setA)):\n        for j in range(i+1, len(setA)):\n            rmsd_vals.append(kabsch_rmsd(setA[i], setA[j]))\n    return np.mean(rmsd_vals)\n\n# ======================================\n# 4️⃣ Weighted Agent Search Tree (3 models)\n# ======================================\ndef agent_tree_search(models, tid, weights=(0.45, 0.45, 0.10), w_div=1.0, w_dist=0.5):\n    tbm, protenix, nufold = models\n    tbm_conf = extract_structures(tbm, tid)\n    prot_conf = extract_structures(protenix, tid)\n    nuf_conf  = extract_structures(nufold, tid)\n\n    # Combine candidates (15 conformations)\n    all_conf = np.concatenate([tbm_conf, prot_conf, nuf_conf], axis=0)\n    model_labels = ([\"tbm\"]*5 + [\"protenix\"]*5 + [\"nufold\"]*5)\n\n    # Strong priors: Protenix conf0 + TBM conf0\n    priors = [prot_conf[1], tbm_conf[0]]\n    selected = [prot_conf[1], tbm_conf[0]]\n\n    # Pool of remaining candidates\n    pool = [(conf, label) for conf, label in zip(all_conf, model_labels)\n            if not any(np.allclose(conf, s) for s in selected)]\n\n    # Search loop until 5 conformations selected\n    while len(selected) < 5:\n        best_score, best_conf = -np.inf, None\n        for cand, label in pool:\n            diversity = mean_rmsd(selected + [cand])\n            dist_to_priors = np.mean([kabsch_rmsd(cand, p) for p in priors])\n\n            # Model-specific reliability weighting\n            if label == \"protenix\":\n                model_weight = weights[1]\n            elif label == \"tbm\":\n                model_weight = weights[0]\n            else:  # nufold\n                model_weight = weights[2]\n\n            # Weighted heuristic: balance diversity and proximity\n            score = model_weight * (w_div * diversity - w_dist * dist_to_priors)\n\n            if score > best_score:\n                best_score, best_conf = score, cand\n\n        selected.append(best_conf)\n        pool = [(conf, label) for conf, label in pool if not np.allclose(conf, best_conf)]\n\n    return selected\n\n# ======================================\n# 5️⃣ Build and Save Submission\n# ======================================\ndef build_submission(tbm, protenix, nufold, output_path=\"/kaggle/working/submission.csv\"):\n    models = [tbm, protenix, nufold]\n    all_rows = []\n\n    for tid in tqdm(tbm[\"target_id\"].unique(), desc=\"Weighted Agent Search (0.45,0.45,0.10)\"):\n        selected = agent_tree_search(models, tid, weights=(0.15, 0.75, 0.10))\n        base = tbm[tbm[\"target_id\"] == tid].sort_values(\"resid\")\n        for j, (resid, resname) in enumerate(zip(base[\"resid\"], base[\"resname\"])):\n            row = {\"ID\": f\"{tid}_{resid}\", \"resname\": resname, \"resid\": resid}\n            for k in range(5):\n                row[f\"x_{k+1}\"], row[f\"y_{k+1}\"], row[f\"z_{k+1}\"] = selected[k][j]\n            all_rows.append(row)\n\n    sub = pd.DataFrame(all_rows)\n    sub.to_csv(output_path, index=False)\n    print(f\"✅ 3-model weighted agent ensemble saved to {output_path}\")\n\n# ======================================\n# 6️⃣ Run Ensemble\n# ======================================\nbuild_submission(tbm, protenix, nufold)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}