{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":87793,"databundleVersionId":12276181,"sourceType":"competition"},{"sourceId":5123458,"sourceType":"datasetVersion","datasetId":2975803},{"sourceId":10855324,"sourceType":"datasetVersion","datasetId":6742586},{"sourceId":10880419,"sourceType":"datasetVersion","datasetId":6760509},{"sourceId":10923077,"sourceType":"datasetVersion","datasetId":6785143},{"sourceId":11118830,"sourceType":"datasetVersion","datasetId":6933267},{"sourceId":11695366,"sourceType":"datasetVersion","datasetId":6791615},{"sourceId":224830487,"sourceType":"kernelVersion"}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n#for dirname, _, filenames in os.walk('/kaggle/input'):\n#    for filename in filenames:\n#        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-10-28T22:55:49.599585Z","iopub.execute_input":"2025-10-28T22:55:49.599979Z","iopub.status.idle":"2025-10-28T22:55:50.506512Z","shell.execute_reply.started":"2025-10-28T22:55:49.599940Z","shell.execute_reply":"2025-10-28T22:55:50.505702Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Environments\n","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%ls /kaggle/input/boltz-dependencies","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T22:57:20.426543Z","iopub.execute_input":"2025-10-28T22:57:20.426993Z","iopub.status.idle":"2025-10-28T22:57:20.554857Z","shell.execute_reply.started":"2025-10-28T22:57:20.426958Z","shell.execute_reply":"2025-10-28T22:57:20.553715Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install --no-index /kaggle/input/boltz-dependencies/*whl --no-deps","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T22:57:24.859996Z","iopub.execute_input":"2025-10-28T22:57:24.860393Z","iopub.status.idle":"2025-10-28T22:57:33.952307Z","shell.execute_reply.started":"2025-10-28T22:57:24.860355Z","shell.execute_reply":"2025-10-28T22:57:33.951426Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install --no-index /kaggle/input/fairscale-0413/*whl --no-deps","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T22:57:39.450834Z","iopub.execute_input":"2025-10-28T22:57:39.451259Z","iopub.status.idle":"2025-10-28T22:57:41.155318Z","shell.execute_reply.started":"2025-10-28T22:57:39.451220Z","shell.execute_reply":"2025-10-28T22:57:41.154148Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install --no-index /kaggle/input/biopython/*whl --no-deps","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T22:57:47.110274Z","iopub.execute_input":"2025-10-28T22:57:47.110608Z","iopub.status.idle":"2025-10-28T22:57:49.424237Z","shell.execute_reply.started":"2025-10-28T22:57:47.110576Z","shell.execute_reply":"2025-10-28T22:57:49.423187Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip show protenix","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T22:58:01.052168Z","iopub.execute_input":"2025-10-28T22:58:01.052477Z","iopub.status.idle":"2025-10-28T22:58:03.939117Z","shell.execute_reply.started":"2025-10-28T22:58:01.052450Z","shell.execute_reply":"2025-10-28T22:58:03.938267Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Prepare scripts","metadata":{}},{"cell_type":"code","source":"%cd /kaggle/working/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T22:58:26.512897Z","iopub.execute_input":"2025-10-28T22:58:26.513242Z","iopub.status.idle":"2025-10-28T22:58:26.519073Z","shell.execute_reply.started":"2025-10-28T22:58:26.513209Z","shell.execute_reply":"2025-10-28T22:58:26.518004Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%mkdir inputs_prediction\n%mkdir outputs_prediction","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T22:58:29.345699Z","iopub.execute_input":"2025-10-28T22:58:29.346030Z","iopub.status.idle":"2025-10-28T22:58:29.578205Z","shell.execute_reply.started":"2025-10-28T22:58:29.346005Z","shell.execute_reply":"2025-10-28T22:58:29.577248Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%cp -rf /kaggle/input/rna-prediction-boltz/boltz/src/boltz .","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T22:58:34.524824Z","iopub.execute_input":"2025-10-28T22:58:34.525138Z","iopub.status.idle":"2025-10-28T22:58:35.569153Z","shell.execute_reply.started":"2025-10-28T22:58:34.525110Z","shell.execute_reply":"2025-10-28T22:58:35.567945Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%ls boltz","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T22:58:39.670494Z","iopub.execute_input":"2025-10-28T22:58:39.670829Z","iopub.status.idle":"2025-10-28T22:58:39.791249Z","shell.execute_reply.started":"2025-10-28T22:58:39.670798Z","shell.execute_reply":"2025-10-28T22:58:39.790401Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Write file","metadata":{}},{"cell_type":"code","source":"%%writefile inference.py\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\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            # Parse MSA data\n            msas = sorted({c.msa_id for c in target.record.chains if c.msa_id != -1})\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                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\nif __name__ == \"__main__\":\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":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T23:41:51.442010Z","iopub.execute_input":"2025-10-28T23:41:51.442313Z","iopub.status.idle":"2025-10-28T23:41:51.450586Z","shell.execute_reply.started":"2025-10-28T23:41:51.442289Z","shell.execute_reply":"2025-10-28T23:41:51.449850Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Prepare inputs","metadata":{}},{"cell_type":"code","source":"sub_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":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T23:41:59.788906Z","iopub.execute_input":"2025-10-28T23:41:59.789209Z","iopub.status.idle":"2025-10-28T23:41:59.799596Z","shell.execute_reply.started":"2025-10-28T23:41:59.789184Z","shell.execute_reply":"2025-10-28T23:41:59.798775Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%ls inputs_prediction","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T23:42:03.263905Z","iopub.execute_input":"2025-10-28T23:42:03.264208Z","iopub.status.idle":"2025-10-28T23:42:03.471075Z","shell.execute_reply.started":"2025-10-28T23:42:03.264185Z","shell.execute_reply":"2025-10-28T23:42:03.469919Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%ls outputs_prediction","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T23:42:05.726634Z","iopub.execute_input":"2025-10-28T23:42:05.727045Z","iopub.status.idle":"2025-10-28T23:42:05.921885Z","shell.execute_reply.started":"2025-10-28T23:42:05.727010Z","shell.execute_reply":"2025-10-28T23:42:05.920889Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Exec inference","metadata":{}},{"cell_type":"code","source":"import torch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T23:42:09.146032Z","iopub.execute_input":"2025-10-28T23:42:09.146386Z","iopub.status.idle":"2025-10-28T23:42:09.150520Z","shell.execute_reply.started":"2025-10-28T23:42:09.146353Z","shell.execute_reply":"2025-10-28T23:42:09.149569Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.cuda.empty_cache()\nimport gc\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T23:42:11.467188Z","iopub.execute_input":"2025-10-28T23:42:11.467526Z","iopub.status.idle":"2025-10-28T23:42:11.813148Z","shell.execute_reply.started":"2025-10-28T23:42:11.467497Z","shell.execute_reply":"2025-10-28T23:42:11.812386Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import subprocess\nimport logging\n\nlogging.basicConfig(level=logging.INFO)\nlogger = logging.getLogger(__name__)\n\nresult = subprocess.run(['python', 'inference.py'], capture_output=True, text=True)\nlogger.info(f\"Command output: {result.stdout}\")\nlogger.error(f\"Command error: {result.stderr}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T23:42:14.727659Z","iopub.execute_input":"2025-10-28T23:42:14.728021Z","iopub.status.idle":"2025-10-29T00:01:00.662120Z","shell.execute_reply.started":"2025-10-28T23:42:14.727993Z","shell.execute_reply":"2025-10-29T00:01:00.661416Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Read RNA files","metadata":{}},{"cell_type":"code","source":"result","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T00:04:03.135000Z","iopub.execute_input":"2025-10-29T00:04:03.135326Z","iopub.status.idle":"2025-10-29T00:04:03.140769Z","shell.execute_reply.started":"2025-10-29T00:04:03.135304Z","shell.execute_reply":"2025-10-29T00:04:03.139855Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Gather results","metadata":{"execution":{"iopub.status.busy":"2025-03-04T21:51:26.943982Z","iopub.execute_input":"2025-03-04T21:51:26.94429Z","iopub.status.idle":"2025-03-04T21:51:26.949385Z","shell.execute_reply.started":"2025-03-04T21:51:26.944265Z","shell.execute_reply":"2025-03-04T21:51:26.948635Z"}}},{"cell_type":"code","source":"from Bio.PDB.MMCIF2Dict import MMCIF2Dict\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    return c1_coords\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":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T00:04:16.973832Z","iopub.execute_input":"2025-10-29T00:04:16.974165Z","iopub.status.idle":"2025-10-29T00:04:17.000974Z","shell.execute_reply.started":"2025-10-29T00:04:16.974138Z","shell.execute_reply":"2025-10-29T00:04:17.000236Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"idx = 0\nfor tmp_id in all_preds:\n    print('#'*20, f'inferences for {tmp_id}')\n    for idx in range(5):\n        c1_coords = get_coords(tmp_id, idx)\n        submission.loc[submission['ID'].apply(lambda x: tmp_id in x), [f'x_{idx+1}', f'y_{idx+1}', f'z_{idx+1}']] = c1_coords\n    print()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T00:04:20.611969Z","iopub.execute_input":"2025-10-29T00:04:20.612271Z","iopub.status.idle":"2025-10-29T00:04:27.483318Z","shell.execute_reply.started":"2025-10-29T00:04:20.612248Z","shell.execute_reply":"2025-10-29T00:04:27.482529Z"}},"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},{"cell_type":"code","source":"%ls","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T00:04:34.052271Z","iopub.execute_input":"2025-10-29T00:04:34.052558Z","iopub.status.idle":"2025-10-29T00:04:34.254404Z","shell.execute_reply.started":"2025-10-29T00:04:34.052536Z","shell.execute_reply":"2025-10-29T00:04:34.253152Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%rm -rf boltz","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T00:04:40.417536Z","iopub.execute_input":"2025-10-29T00:04:40.417968Z","iopub.status.idle":"2025-10-29T00:04:40.621345Z","shell.execute_reply.started":"2025-10-29T00:04:40.417922Z","shell.execute_reply":"2025-10-29T00:04:40.620107Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%rm -rf inputs_prediction","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T00:04:43.342880Z","iopub.execute_input":"2025-10-29T00:04:43.343330Z","iopub.status.idle":"2025-10-29T00:04:43.539552Z","shell.execute_reply.started":"2025-10-29T00:04:43.343294Z","shell.execute_reply":"2025-10-29T00:04:43.538407Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%rm -rf outputs_prediction","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T00:04:45.625867Z","iopub.execute_input":"2025-10-29T00:04:45.626190Z","iopub.status.idle":"2025-10-29T00:04:45.832439Z","shell.execute_reply.started":"2025-10-29T00:04:45.626165Z","shell.execute_reply":"2025-10-29T00:04:45.831334Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%rm -rf inference.py","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T00:04:47.577317Z","iopub.execute_input":"2025-10-29T00:04:47.577643Z","iopub.status.idle":"2025-10-29T00:04:47.775507Z","shell.execute_reply.started":"2025-10-29T00:04:47.577616Z","shell.execute_reply":"2025-10-29T00:04:47.774373Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission.to_csv(\"submissionnew1.csv\", index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T00:04:51.087714Z","iopub.execute_input":"2025-10-29T00:04:51.088119Z","iopub.status.idle":"2025-10-29T00:04:51.147411Z","shell.execute_reply.started":"2025-10-29T00:04:51.088086Z","shell.execute_reply":"2025-10-29T00:04:51.146466Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%ls","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T00:04:53.365027Z","iopub.execute_input":"2025-10-29T00:04:53.365355Z","iopub.status.idle":"2025-10-29T00:04:53.561700Z","shell.execute_reply.started":"2025-10-29T00:04:53.365329Z","shell.execute_reply":"2025-10-29T00:04:53.560798Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Pronetix","metadata":{}},{"cell_type":"code","source":"MODEL_TYPE='protenix'\nVALIDATION=False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T22:58:49.728374Z","iopub.execute_input":"2025-10-28T22:58:49.728794Z","iopub.status.idle":"2025-10-28T22:58:49.732649Z","shell.execute_reply.started":"2025-10-28T22:58:49.728746Z","shell.execute_reply":"2025-10-28T22:58:49.731831Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!export PROTENIX_DATA_ROOT_DIR=/kaggle/input/protenix-checkpoints","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T22:58:52.316842Z","iopub.execute_input":"2025-10-28T22:58:52.317145Z","iopub.status.idle":"2025-10-28T22:58:52.435352Z","shell.execute_reply.started":"2025-10-28T22:58:52.317121Z","shell.execute_reply":"2025-10-28T22:58:52.434179Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! mkdir /af3-dev \n! ln -s /kaggle/input/protenix-checkpoints /af3-dev/release_data\n! ls /af3-dev/release_data/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T22:59:20.985501Z","iopub.execute_input":"2025-10-28T22:59:20.985831Z","iopub.status.idle":"2025-10-28T22:59:21.344704Z","shell.execute_reply.started":"2025-10-28T22:59:20.985803Z","shell.execute_reply":"2025-10-28T22:59:21.343607Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import Bio\n\nfrom copy import deepcopy\n\nimport pandas as pd\nfrom Bio.PDB import Atom, Model, Chain, Residue, Structure, PDBParser\nfrom Bio import SeqIO\nimport os, sys\nimport re\nimport numpy as np\nimport torch\n\nimport matplotlib\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nimport time\ntime0=time.time()\n\nprint('IMPORT OK !!!!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T22:59:23.784281Z","iopub.execute_input":"2025-10-28T22:59:23.784631Z","iopub.status.idle":"2025-10-28T22:59:27.004685Z","shell.execute_reply.started":"2025-10-28T22:59:23.784599Z","shell.execute_reply":"2025-10-28T22:59:27.003850Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PYTHON = sys.executable\nprint('PYTHON',PYTHON)\n\nRHONET_DIR=\\\n'/kaggle/input/data-for-demo-for-rhofold-plus-with-kaggle-msa/RhoFold-main'\n#'<your downloaded rhofold repo>/RhoFold-main'\n\nUSALIGN = \\\n'/kaggle/working//USalign'\n#'<your us align path>/USalign'\n\nos.system('cp /kaggle/input/usalign/USalign /kaggle/working/')\nos.system('sudo chmod u+x /kaggle/working//USalign')\nsys.path.append(RHONET_DIR)\n\n\nDATA_KAGGLE_DIR = '/kaggle/input/stanford-rna-3d-folding'\n\n\n# helper ----\nclass dotdict(dict):\n\t__setattr__ = dict.__setitem__\n\t__delattr__ = dict.__delitem__\n\n\tdef __getattr__(self, name):\n\t\ttry:\n\t\t\treturn self[name]\n\t\texcept KeyError:\n\t\t\traise AttributeError(name)\n\n# visualisation helper ----\ndef set_aspect_equal(ax):\n\tx_limits = ax.get_xlim()\n\ty_limits = ax.get_ylim()\n\tz_limits = ax.get_zlim()\n\n\t# Compute the mean of each axis\n\tx_middle = np.mean(x_limits)\n\ty_middle = np.mean(y_limits)\n\tz_middle = np.mean(z_limits)\n\n\t# Compute the max range across all axes\n\tmax_range = max(x_limits[1] - x_limits[0],\n\t\t\t\t\ty_limits[1] - y_limits[0],\n\t\t\t\t\tz_limits[1] - z_limits[0]) / 2.0\n\n\t# Set the new limits to ensure equal scaling\n\tax.set_xlim(x_middle - max_range, x_middle + max_range)\n\tax.set_ylim(y_middle - max_range, y_middle + max_range)\n\tax.set_zlim(z_middle - max_range, z_middle + max_range)\n\n\n\n\n# xyz df helper --------------------\ndef get_truth_df(target_id):\n    truth_df = LABEL_DF[LABEL_DF['target_id'] == target_id]\n    truth_df = truth_df.reset_index(drop=True)\n    return truth_df\n\ndef parse_output_to_df(output, seq, target_id):\n    df = []\n    chain_data = []\n    for i, res in enumerate(seq):\n        d=dict(ID = target_id,\n                    resname=res,\n                    resid=i+1)\n        for n in range(len(output)):\n            d={**d, f'x_{n+1}': round(output[n,i,0].item(),3),\n                     f'y_{n+1}': round(output[n,i,1].item(),3),\n                     f'z_{n+1}': round(output[n,i,2].item(),3)}\n        chain_data.append(d)\n\n    if len(chain_data)!=0:\n        chain_df = pd.DataFrame(chain_data)\n        df.append(chain_df)\n        ##print(chain_df)\n    return df\n\ndef parse_pdb_to_df(pdb_file, target_id):\n    parser = PDBParser()\n    structure = parser.get_structure('', pdb_file)\n\n    df = []\n    for model in structure:\n        for chain in model:\n            print(chain)\n            chain_data = []\n            for residue in chain:\n                # print(residue)\n                if residue.get_resname() in ['A', 'U', 'G', 'C']:\n                    # Check if the residue has a C1' atom\n                    if 'C1\\'' in residue:\n                        atom = residue['C1\\'']\n                        xyz = atom.get_coord()\n                        resname = residue.get_resname()\n                        resid = residue.get_id()[1]\n\n                        #todo detect discontinous: resid = prev_resid+1\n                        #ID\tresname\tresid\tx_1\ty_1\tz_1\n                        chain_data.append(dict(\n                            ID = target_id+'_'+str(resid),\n                            resname=resname,\n                            resid=resid,\n                            x_1=xyz[0],\n                            y_1=xyz[1],\n                            z_1=xyz[2],\n                        ))\n                        ##print(f\"Residue {resname} {resid}, Atom: {atom.get_name()}, xyz: {xyz}\")\n\n            if len(chain_data)!=0:\n                chain_df = pd.DataFrame(chain_data)\n                df.append(chain_df)\n                ##print(chain_df)\n    return df\n\n# usalign helper --------------------\ndef write_target_line(\n    atom_name, atom_serial, residue_name, chain_id, residue_num, x_coord, y_coord, z_coord, occupancy=1.0, b_factor=0.0, atom_type='P'\n):\n    \"\"\"\n    Writes a single line of PDB format based on provided atom information.\n\n    Args:\n        atom_name (str): Name of the atom (e.g., \"N\", \"CA\").\n        atom_serial (int): Atom serial number.\n        residue_name (str): Residue name (e.g., \"ALA\").\n        chain_id (str): Chain identifier.\n        residue_num (int): Residue number.\n        x_coord (float): X coordinate.\n        y_coord (float): Y coordinate.\n        z_coord (float): Z coordinate.\n        occupancy (float, optional): Occupancy value (default: 1.0).\n        b_factor (float, optional): B-factor value (default: 0.0).\n\n    Returns:\n        str: A single line of PDB string.\n    \"\"\"\n    return f'ATOM  {atom_serial:>5d}  {atom_name:<5s} {residue_name:<3s} {residue_num:>3d}    {x_coord:>8.3f}{y_coord:>8.3f}{z_coord:>8.3f}{occupancy:>6.2f}{b_factor:>6.2f}           {atom_type}\\n'\n\ndef write_xyz_to_pdb(df, pdb_file, xyz_id = 1):\n    resolved_cnt = 0\n    with open(pdb_file, 'w') as target_file:\n        for _, row in df.iterrows():\n            x_coord = row[f'x_{xyz_id}']\n            y_coord = row[f'y_{xyz_id}']\n            z_coord = row[f'z_{xyz_id}']\n\n            if x_coord > -1e17 and y_coord > -1e17 and z_coord > -1e17:\n                resolved_cnt += 1\n                target_line = write_target_line(\n                    atom_name=\"C1'\",\n                    atom_serial=int(row['resid']),\n                    residue_name=row['resname'],\n                    chain_id='0',\n                    residue_num=int(row['resid']),\n                    x_coord=x_coord,\n                    y_coord=y_coord,\n                    z_coord=z_coord,\n                    atom_type='C',\n                )\n                target_file.write(target_line)\n    return resolved_cnt\n\ndef parse_usalign_for_tm_score(output):\n    # Extract TM-score based on length of reference structure (second)\n    tm_score_match = re.findall(r'TM-score=\\s+([\\d.]+)', output)[1]\n    if not tm_score_match:\n        raise ValueError('No TM score found')\n    return float(tm_score_match)\n\ndef parse_usalign_for_transform(output):\n    # Locate the rotation matrix section\n    matrix_lines = []\n    found_matrix = False\n\n    for line in output.splitlines():\n        if \"The rotation matrix to rotate Structure_1 to Structure_2\" in line:\n            found_matrix = True\n        elif found_matrix and re.match(r'^\\d+\\s+[-\\d.]+\\s+[-\\d.]+\\s+[-\\d.]+\\s+[-\\d.]+$', line):\n            matrix_lines.append(line)\n        elif found_matrix and not line.strip():\n            break  # Stop parsing if an empty line is encountered after the matrix\n\n    # Parse the rotation matrix values\n    rotation_matrix = []\n    for line in matrix_lines:\n        parts = line.split()\n        row_values = list(map(float, parts[1:]))  # Skip the first column (index)\n        rotation_matrix.append(row_values)\n\n    return np.array(rotation_matrix)\n\ndef call_usalign(predict_df, truth_df, verbose=1):\n    truth_pdb = '~truth.pdb'\n    predict_pdb = '~predict.pdb'\n    write_xyz_to_pdb(predict_df, predict_pdb, xyz_id=1)\n    write_xyz_to_pdb(truth_df, truth_pdb, xyz_id=1)\n\n    command = f'{USALIGN} {predict_pdb} {truth_pdb} -atom \" C1\\'\" -m -'\n    output = os.popen(command).read()\n    if verbose==1:\n        print(output)\n    tm_score = parse_usalign_for_tm_score(output)\n    transform = parse_usalign_for_transform(output)\n    return tm_score, transform\n\nprint('HELPER OK!!!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T22:59:33.977656Z","iopub.execute_input":"2025-10-28T22:59:33.978210Z","iopub.status.idle":"2025-10-28T22:59:34.081709Z","shell.execute_reply.started":"2025-10-28T22:59:33.978180Z","shell.execute_reply":"2025-10-28T22:59:34.080931Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"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(0)\n    torch.random.manual_seed(0)\n    torch.cuda.manual_seed_all(0)\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":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T22:59:43.516611Z","iopub.execute_input":"2025-10-28T22:59:43.516920Z","iopub.status.idle":"2025-10-28T22:59:45.475043Z","shell.execute_reply.started":"2025-10-28T22:59:43.516894Z","shell.execute_reply":"2025-10-28T22:59:45.474138Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if 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\"] = 5\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/protenix-checkpoints/model_v0.2.0.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   \n    runner = InferenceRunner(configs)  # Replace with the actual method in your runner, e.g., 'predict()' or 'inference()'\n\nimport random\ndef set_seed(seed):\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed)\n        torch.cuda.manual_seed_all(seed) \n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False\n    random.seed(seed)\n    np.random.seed(seed)\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T23:02:37.043519Z","iopub.execute_input":"2025-10-28T23:02:37.043843Z","iopub.status.idle":"2025-10-28T23:02:54.887400Z","shell.execute_reply.started":"2025-10-28T23:02:37.043818Z","shell.execute_reply":"2025-10-28T23:02:54.886761Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_df=pd.read_csv('/kaggle/input/stanford-rna-3d-folding/test_sequences.csv')\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")  ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T23:04:45.307315Z","iopub.execute_input":"2025-10-28T23:04:45.307663Z","iopub.status.idle":"2025-10-28T23:04:45.327468Z","shell.execute_reply.started":"2025-10-28T23:04:45.307632Z","shell.execute_reply":"2025-10-28T23:04:45.326826Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataset = DictDataset(test_df.sequence, dump_dir='output', id_list=test_df.target_id, use_msa=True)\nnum_data = len(dataset)\nfor i, seq in tqdm(enumerate(test_df.sequence),total=num_data):\n    try:\n        data, atom_array, data_error_message=dataset[i]\n        target_id = data[\"sample_name\"]\n        assert target_id==test_df.target_id[i]\n        assert data_error_message==''\n        new_configs = update_inference_configs(configs, data[\"N_token\"].item())\n        prediction3 = []\n        set_seed(2023)\n        data, atom_array, data_error_message=dataset[i]\n        target_id = data[\"sample_name\"]\n        assert target_id==test_df.target_id[i]\n        assert data_error_message==''\n        new_configs = update_inference_configs(configs, data[\"N_token\"].item())\n        runner=InferenceRunner(configs)\n        runner.update_model_configs(new_configs)\n        prediction = runner.predict(data)\n        prediction=prediction['coordinate'][:,data['input_feature_dict']['atom_to_tokatom_idx']==12]\n        result = parse_output_to_df(prediction, seq, target_id)[0]\n\n    except:\n        target_id==test_df.target_id[i]\n        print('Failed to predict', target_id)\n        result=pd.DataFrame(columns=['ID', 'resname', 'resid', \n                                     'x_1', 'y_1', 'z_1', \n                                     'x_2', 'y_2', 'z_2',\n                                     'x_3', 'y_3', 'z_3', \n                                     'x_4', 'y_4', 'z_4', \n                                     'x_5', 'y_5', 'z_5'], \n                                     data=[[target_id, x, j+1] + [0.0]*15 for j, x in enumerate(seq)])\n\n    result['ID']=result.apply(lambda x: x.ID + '_' + str(x.resid), axis=1)\n    result.to_csv('submissionnew2.csv', index=False, mode='a', header=(i==0))\n    torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T23:04:50.703064Z","iopub.execute_input":"2025-10-28T23:04:50.703363Z","iopub.status.idle":"2025-10-28T23:41:24.210572Z","shell.execute_reply.started":"2025-10-28T23:04:50.703341Z","shell.execute_reply":"2025-10-28T23:41:24.209819Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nsub1 = pd.read_csv('submissionnew1.csv').set_index('ID')\nsub12 = pd.read_csv('submissionnew2.csv').set_index('ID')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T00:05:05.492547Z","iopub.execute_input":"2025-10-29T00:05:05.492881Z","iopub.status.idle":"2025-10-29T00:05:05.524318Z","shell.execute_reply.started":"2025-10-29T00:05:05.492853Z","shell.execute_reply":"2025-10-29T00:05:05.523430Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for key in sub1.index:\n    if key in sub12.index:\n        for c in [1,2,3,4,5]:\n            if abs(sub1.loc[key,f'x_{c}'])+abs(sub1.loc[key,f'y_{c}'])+abs(sub1.loc[key,f'z_{c}'])==0:\n                sub1.loc[key,f'x_{c}'] = sub12.loc[key,f'x_{c}']\n                sub1.loc[key,f'y_{c}'] = sub12.loc[key,f'y_{c}']\n                sub1.loc[key,f'z_{c}'] = sub12.loc[key,f'z_{c}']\n    \n        \n        sub1.loc[key,'x_2'] = sub12.loc[key,'x_1']\n        sub1.loc[key,'y_2'] = sub12.loc[key,'y_1']\n        sub1.loc[key,'z_2'] = sub12.loc[key,'z_1']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T00:05:07.839504Z","iopub.execute_input":"2025-10-29T00:05:07.839857Z","iopub.status.idle":"2025-10-29T00:05:10.479559Z","shell.execute_reply.started":"2025-10-29T00:05:07.839828Z","shell.execute_reply":"2025-10-29T00:05:10.478889Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub1.to_csv('submission.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T00:05:13.958448Z","iopub.execute_input":"2025-10-29T00:05:13.958784Z","iopub.status.idle":"2025-10-29T00:05:14.010865Z","shell.execute_reply.started":"2025-10-29T00:05:13.958718Z","shell.execute_reply":"2025-10-29T00:05:14.010172Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! rm -rf submissionnew1.csv\n! rm -rf submissionnew2.csv","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T00:05:16.959377Z","iopub.execute_input":"2025-10-29T00:05:16.959671Z","iopub.status.idle":"2025-10-29T00:05:17.347175Z","shell.execute_reply.started":"2025-10-29T00:05:16.959647Z","shell.execute_reply":"2025-10-29T00:05:17.346020Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%ls","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T00:05:18.975198Z","iopub.execute_input":"2025-10-29T00:05:18.975518Z","iopub.status.idle":"2025-10-29T00:05:19.171180Z","shell.execute_reply.started":"2025-10-29T00:05:18.975492Z","shell.execute_reply":"2025-10-29T00:05:19.170252Z"}},"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},{"cell_type":"code","source":"","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},{"cell_type":"code","source":"","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},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}