{"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":11553390,"sourceType":"competition"},{"sourceId":10855324,"sourceType":"datasetVersion","datasetId":6742586},{"sourceId":11225224,"sourceType":"datasetVersion","datasetId":7010874},{"sourceId":224830487,"sourceType":"kernelVersion"}],"dockerImageVersionId":30918,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"MODEL_TYPE='chai'\nVALIDATION=False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T10:18:41.955020Z","iopub.execute_input":"2025-03-31T10:18:41.955401Z","iopub.status.idle":"2025-03-31T10:18:41.959958Z","shell.execute_reply.started":"2025-03-31T10:18:41.955369Z","shell.execute_reply":"2025-03-31T10:18:41.958733Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Install requirements ","metadata":{}},{"cell_type":"code","source":"if MODEL_TYPE=='chai' and VALIDATION:\n    !pip install chai_lab==0.6.1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T10:18:45.213017Z","iopub.execute_input":"2025-03-31T10:18:45.213414Z","iopub.status.idle":"2025-03-31T10:19:31.823947Z","shell.execute_reply.started":"2025-03-31T10:18:45.213385Z","shell.execute_reply":"2025-03-31T10:19:31.822622Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!mkdir chai-checkpoints\n!ln -s /kaggle/input/chai-checkpoints/conformers_v1.apkl /kaggle/working/chai-checkpoints/conformers_v1.apkl\n!ln -s /kaggle/input/chai-checkpoints/conformers_v1.download_lock /kaggle/working/chai-checkpoints/conformers_v1.download_lock\n!ln -s /kaggle/input/chai-checkpoints /kaggle/working/chai-checkpoints/models_v2\n!ls /kaggle/working/chai-checkpoints","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T10:19:34.899270Z","iopub.execute_input":"2025-03-31T10:19:34.899626Z","iopub.status.idle":"2025-03-31T10:19:35.483156Z","shell.execute_reply.started":"2025-03-31T10:19:34.899592Z","shell.execute_reply":"2025-03-31T10:19:35.481237Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nos.environ[\"CHAI_DOWNLOADS_DIR\"]='/kaggle/working/chai-checkpoints'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T10:19:36.864018Z","iopub.execute_input":"2025-03-31T10:19:36.864422Z","iopub.status.idle":"2025-03-31T10:19:36.869137Z","shell.execute_reply.started":"2025-03-31T10:19:36.864390Z","shell.execute_reply":"2025-03-31T10:19:36.868004Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! ls $CHAI_DOWNLOADS_DIR","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T10:19:43.879645Z","iopub.execute_input":"2025-03-31T10:19:43.880112Z","iopub.status.idle":"2025-03-31T10:19:43.999985Z","shell.execute_reply.started":"2025-03-31T10:19:43.880061Z","shell.execute_reply":"2025-03-31T10:19:43.998551Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Helper scripts","metadata":{}},{"cell_type":"code","source":"import Bio\n\nfrom copy import deepcopy\n\nimport pandas as pd\nfrom Bio.PDB import Atom, Model, Chain, Residue, Structure, PDBParser\nfrom Bio import SeqIO\nimport os, sys\nimport re\nimport numpy as np\nimport torch\n\nimport matplotlib\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nimport time\ntime0=time.time()\n\nprint('IMPORT OK !!!!')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2025-03-31T10:19:45.552678Z","iopub.execute_input":"2025-03-31T10:19:45.553020Z","iopub.status.idle":"2025-03-31T10:19:50.487899Z","shell.execute_reply.started":"2025-03-31T10:19:45.552990Z","shell.execute_reply":"2025-03-31T10:19:50.486796Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PYTHON = sys.executable\nprint('PYTHON',PYTHON)\n\nRHONET_DIR=\\\n'/kaggle/input/data-for-demo-for-rhofold-plus-with-kaggle-msa/RhoFold-main'\n#'<your downloaded rhofold repo>/RhoFold-main'\n\nUSALIGN = \\\n'/kaggle/working//USalign'\n#'<your us align path>/USalign'\n\nos.system('cp /kaggle/input/usalign/USalign /kaggle/working/')\nos.system('sudo chmod u+x /kaggle/working//USalign')\nsys.path.append(RHONET_DIR)\n\n\nDATA_KAGGLE_DIR = '/kaggle/input/stanford-rna-3d-folding'\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-03-31T10:19:50.489386Z","iopub.execute_input":"2025-03-31T10:19:50.489963Z","iopub.status.idle":"2025-03-31T10:19:50.636328Z","shell.execute_reply.started":"2025-03-31T10:19:50.489924Z","shell.execute_reply":"2025-03-31T10:19:50.635229Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if MODEL_TYPE=='chai':\n\n    from chai_lab.chai1 import raise_if_too_many_tokens\n    from chai_lab.data.dataset.inference_dataset import Input, load_chains_from_raw\n    from chai_lab.data.dataset.structure.all_atom_structure_context import AllAtomStructureContext\n    from chai_lab.data.dataset.msas.msa_context import MSAContext\n    from chai_lab.data.dataset.all_atom_feature_context import (\n        MAX_MSA_DEPTH,\n        MAX_NUM_TEMPLATES,\n        AllAtomFeatureContext,\n    )\n    from chai_lab.data.dataset.templates.context import TemplateContext\n    from chai_lab.data.dataset.embeddings.embedding_context import EmbeddingContext\n    from chai_lab.data.dataset.constraints.restraint_context import RestraintContext\n    \n    def get_context(\n        seq, target_id='query',\n        esm_device: torch.device = torch.device(\"cpu\"),\n    ):\n\n        fasta_inputs=[Input(seq, 1, target_id)]    \n\n        assert len(fasta_inputs) > 0, \"No inputs found in fasta file\"\n\n        # Load structure context\n        chains = load_chains_from_raw(fasta_inputs)\n        del fasta_inputs  # Do not reference inputs after creating chains from them\n\n        merged_context = AllAtomStructureContext.merge(\n            [c.structure_context for c in chains]\n        )\n        n_actual_tokens = merged_context.num_tokens\n        raise_if_too_many_tokens(n_actual_tokens)\n\n        # Generated and/or load MSAs\n        msa_context = MSAContext.create_empty(\n            n_tokens=n_actual_tokens, depth=MAX_MSA_DEPTH\n            )\n        msa_profile_context = MSAContext.create_empty(\n            n_tokens=n_actual_tokens, depth=MAX_MSA_DEPTH\n            )\n\n        assert (\n            msa_context.num_tokens == merged_context.num_tokens\n        ), f\"Discrepant tokens in input and MSA: {merged_context.num_tokens} != {msa_context.num_tokens}\"\n\n        template_context = TemplateContext.empty(\n            n_tokens=n_actual_tokens,\n            n_templates=MAX_NUM_TEMPLATES,\n        )\n        embedding_context = EmbeddingContext.empty(n_tokens=n_actual_tokens)\n\n        restraint_context = RestraintContext.empty()\n\n        merged_context.drop_glycan_leaving_atoms_inplace()\n\n        # Build final feature context\n        feature_context = AllAtomFeatureContext(\n            chains=chains,\n            structure_context=merged_context,\n            msa_context=msa_context,\n            profile_msa_context=msa_profile_context,\n            template_context=template_context,\n            embedding_context=embedding_context,\n            restraint_context=restraint_context,\n        )\n        return feature_context\n    print('DATA LOADER OK!!!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T10:19:53.939991Z","iopub.execute_input":"2025-03-31T10:19:53.940402Z","iopub.status.idle":"2025-03-31T10:20:00.022034Z","shell.execute_reply.started":"2025-03-31T10:19:53.940371Z","shell.execute_reply":"2025-03-31T10:20:00.020830Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if MODEL_TYPE=='chai':\n    from einops import einsum, rearrange, repeat\n    from chai_lab.chai1 import *\n    from chai_lab.data.collate.collate import Collate\n    from chai_lab.data.collate.utils import AVAILABLE_MODEL_SIZES\n    from chai_lab.utils.tensor_utils import move_data_to_device, set_seed, und_self\n    from chai_lab.data.features.generators.token_bond import TokenBondRestraint\n    from chai_lab.model.diffusion_schedules import InferenceNoiseSchedule\n\n    def _tensor_to_atom_names(tensor):\n        return [\n            \"\".join([chr(ord_val + 32) for ord_val in ords_atom]).rstrip()\n            for ords_atom in tensor\n        ]\n    \n    ##\n    ## Load exported models\n    ##\n    device=(torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu'))\n\n    feature_embedding = load_exported(\"feature_embedding.pt\", device)\n    token_input_embedder = load_exported(\"token_embedder.pt\", device)\n    trunk = load_exported(\"trunk.pt\", device)\n    diffusion_module = load_exported(\"diffusion_module.pt\", device)\n\n    @torch.no_grad()\n    def run_folding(\n        seq, target_id,\n        recycle_msa_subsample: int = 0,\n        num_trunk_recycles: int = 3,\n        num_diffn_timesteps: int = 200,\n        num_diffn_samples: int = 5,\n        seed: int = 10,\n        device: torch.device = device,\n):\n   \n        set_seed([seed])\n\n        # Clear memory\n        torch.cuda.empty_cache()\n\n        ##\n        ## Validate inputs\n        ##\n\n        feature_context=get_context(seq, target_id)\n\n        n_actual_tokens = feature_context.structure_context.num_tokens\n        raise_if_too_many_tokens(n_actual_tokens)\n        raise_if_too_many_templates(feature_context.template_context.num_templates)\n        raise_if_msa_too_deep(feature_context.msa_context.depth)\n        # NOTE profile MSA used only for statistics; no depth check\n        feature_context.structure_context.report_bonds()\n\n        ##\n        ## Prepare batch\n        ##\n\n        # Collate inputs into batch\n        collator = Collate(\n            feature_factory=feature_factory,\n            num_key_atoms=128,\n            num_query_atoms=32,\n        )\n\n        feature_contexts = [feature_context]\n        batch_size = len(feature_contexts)\n        batch = collator(feature_contexts)\n\n        batch = move_data_to_device(batch, device=device)\n\n        # Get features and inputs from batch\n        features = {name: feature for name, feature in batch[\"features\"].items()}\n        inputs = batch[\"inputs\"]\n        block_indices_h = inputs[\"block_atom_pair_q_idces\"]\n        block_indices_w = inputs[\"block_atom_pair_kv_idces\"]\n        atom_single_mask = inputs[\"atom_exists_mask\"]\n        atom_token_indices = inputs[\"atom_token_index\"].long()\n        token_single_mask = inputs[\"token_exists_mask\"]\n        token_pair_mask = und_self(token_single_mask, \"b i, b j -> b i j\")\n        token_reference_atom_index = inputs[\"token_ref_atom_index\"]\n        atom_within_token_index = inputs[\"atom_within_token_index\"]\n        msa_mask = inputs[\"msa_mask\"]\n        template_input_masks = und_self(\n            inputs[\"template_mask\"], \"b t n1, b t n2 -> b t n1 n2\"\n        )\n        block_atom_pair_mask = inputs[\"block_atom_pair_mask\"]\n\n        _, _, model_size = msa_mask.shape  \n        assert model_size in AVAILABLE_MODEL_SIZES\n        \n        ##\n        ## Run the features through the feature embedder\n        ##\n\n        embedded_features = feature_embedding.forward(\n            crop_size=model_size,\n            move_to_device=device,\n            return_on_cpu=False,\n            **features,\n        )\n        token_single_input_feats = embedded_features[\"TOKEN\"]\n        token_pair_input_feats, token_pair_structure_input_feats = embedded_features[\n            \"TOKEN_PAIR\"\n        ].chunk(2, dim=-1)\n        atom_single_input_feats, atom_single_structure_input_feats = embedded_features[\n            \"ATOM\"\n        ].chunk(2, dim=-1)\n        block_atom_pair_input_feats, block_atom_pair_structure_input_feats = (\n            embedded_features[\"ATOM_PAIR\"].chunk(2, dim=-1)\n        )\n        template_input_feats = embedded_features[\"TEMPLATES\"]\n        msa_input_feats = embedded_features[\"MSA\"]\n\n        \n        ##\n        ## Run the inputs through the token input embedder\n        ##\n\n        token_input_embedder_outputs: tuple[Tensor, ...] = token_input_embedder.forward(\n            return_on_cpu=False,\n            move_to_device=device,\n            token_single_input_feats=token_single_input_feats,\n            token_pair_input_feats=token_pair_input_feats,\n            atom_single_input_feats=atom_single_input_feats,\n            block_atom_pair_feat=block_atom_pair_input_feats,\n            block_atom_pair_mask=block_atom_pair_mask,\n            block_indices_h=block_indices_h,\n            block_indices_w=block_indices_w,\n            atom_single_mask=atom_single_mask,\n            atom_token_indices=atom_token_indices,\n            crop_size=model_size,\n        )\n        token_single_initial_repr, token_single_structure_input, token_pair_initial_repr = (\n            token_input_embedder_outputs\n        )\n\n        ##\n        ## Run the input representations through the trunk\n        ##\n\n        # Recycle the representations by feeding the output back into the trunk as input for\n        # the subsequent recycle\n        token_single_trunk_repr = token_single_initial_repr\n        token_pair_trunk_repr = token_pair_initial_repr\n        for _ in range(num_trunk_recycles):\n            subsampled_msa_input_feats, subsampled_msa_mask = None, None\n            if recycle_msa_subsample > 0:\n                subsampled_msa_input_feats, subsampled_msa_mask = (\n                    subsample_and_reorder_msa_feats_n_mask(\n                        msa_input_feats,\n                        msa_mask,\n                    )\n                )\n            (token_single_trunk_repr, token_pair_trunk_repr) = trunk.forward(\n                move_to_device=device,\n                token_single_trunk_initial_repr=token_single_initial_repr,\n                token_pair_trunk_initial_repr=token_pair_initial_repr,\n                token_single_trunk_repr=token_single_trunk_repr,  # recycled\n                token_pair_trunk_repr=token_pair_trunk_repr,  # recycled\n                msa_input_feats=(\n                    subsampled_msa_input_feats\n                    if subsampled_msa_input_feats is not None\n                    else msa_input_feats\n                ),\n                msa_mask=(\n                    subsampled_msa_mask if subsampled_msa_mask is not None else msa_mask\n                ),\n                template_input_feats=template_input_feats,\n                template_input_masks=template_input_masks,\n                token_single_mask=token_single_mask,\n                token_pair_mask=token_pair_mask,\n                crop_size=model_size,\n            )\n        torch.cuda.empty_cache()\n\n        ##\n        ## Denoise the trunk representation by passing it through the diffusion module\n        ##\n\n        atom_single_mask = atom_single_mask.to(device)\n\n        static_diffusion_inputs = dict(\n            token_single_initial_repr=token_single_structure_input.float(),\n            token_pair_initial_repr=token_pair_structure_input_feats.float(),\n            token_single_trunk_repr=token_single_trunk_repr.float(),\n            token_pair_trunk_repr=token_pair_trunk_repr.float(),\n            atom_single_input_feats=atom_single_structure_input_feats.float(),\n            atom_block_pair_input_feats=block_atom_pair_structure_input_feats.float(),\n            atom_single_mask=atom_single_mask,\n            atom_block_pair_mask=block_atom_pair_mask,\n            token_single_mask=token_single_mask,\n            block_indices_h=block_indices_h,\n            block_indices_w=block_indices_w,\n            atom_token_indices=atom_token_indices,\n        )\n        static_diffusion_inputs = move_data_to_device(\n            static_diffusion_inputs, device=device\n        )\n\n        def _denoise(atom_pos: Tensor, sigma: Tensor, ds: int) -> Tensor:\n            # verified manually that ds dimension can be arbitrary in diff module\n            atom_noised_coords = rearrange(\n                atom_pos, \"(b ds) ... -> b ds ...\", ds=ds\n            ).contiguous()\n            noise_sigma = repeat(sigma, \" -> b ds\", b=batch_size, ds=ds)\n            return diffusion_module.forward(\n                atom_noised_coords=atom_noised_coords.float(),\n                noise_sigma=noise_sigma.float(),\n                crop_size=model_size,\n                **static_diffusion_inputs,\n            )\n\n        inference_noise_schedule = InferenceNoiseSchedule(\n            s_max=DiffusionConfig.S_tmax,\n            s_min=4e-4,\n            p=7.0,\n            sigma_data=DiffusionConfig.sigma_data,\n        )\n        sigmas = inference_noise_schedule.get_schedule(\n            device=device, num_timesteps=num_diffn_timesteps\n        )\n        gammas = torch.where(\n            (sigmas >= DiffusionConfig.S_tmin) & (sigmas <= DiffusionConfig.S_tmax),\n            min(DiffusionConfig.S_churn / num_diffn_timesteps, math.sqrt(2) - 1),\n            0.0,\n        )\n\n        sigmas_and_gammas = list(zip(sigmas[:-1], sigmas[1:], gammas[:-1]))\n\n        # Initial atom positions\n        _, num_atoms = atom_single_mask.shape\n        atom_pos = sigmas[0] * torch.randn(\n            batch_size * num_diffn_samples, num_atoms, 3, device=device\n        )\n\n        for sigma_curr, sigma_next, gamma_curr in sigmas_and_gammas:\n            # Center coords\n            atom_pos = center_random_augmentation(\n                atom_pos,\n                atom_single_mask=repeat(\n                    atom_single_mask,\n                    \"b a -> (b ds) a\",\n                    ds=num_diffn_samples,\n                ),\n            )\n\n            noise = DiffusionConfig.S_noise * torch.randn(\n                atom_pos.shape, device=atom_pos.device\n            )\n            sigma_hat = sigma_curr + gamma_curr * sigma_curr\n            atom_pos_noise = (sigma_hat**2 - sigma_curr**2).clamp_min(1e-6).sqrt()\n            atom_pos_hat = atom_pos + noise * atom_pos_noise\n\n            denoised_pos = _denoise(\n                atom_pos=atom_pos_hat,\n                sigma=sigma_hat,\n                ds=num_diffn_samples,\n            )\n            d_i = (atom_pos_hat - denoised_pos) / sigma_hat\n            atom_pos = atom_pos_hat + (sigma_next - sigma_hat) * d_i\n\n            if sigma_next != 0 and DiffusionConfig.second_order:  # second order update\n                denoised_pos = _denoise(\n                    atom_pos,\n                    sigma=sigma_next,\n                    ds=num_diffn_samples,\n                )\n                d_i_prime = (atom_pos - denoised_pos) / sigma_next\n                atom_pos = atom_pos + (sigma_next - sigma_hat) * ((d_i_prime + d_i) / 2)\n\n        # We won't be running diffusion anymore\n        del static_diffusion_inputs\n        torch.cuda.empty_cache()\n\n        inputs = move_data_to_device(inputs, torch.device(\"cpu\"))\n        atom_pos = atom_pos.cpu()\n\n        mask=np.array(_tensor_to_atom_names(inputs['atom_ref_name_chars'][0]))==\"C1'\"\n\n        return atom_pos[:,mask,:]\n    print('Model OK')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T10:20:00.023446Z","iopub.execute_input":"2025-03-31T10:20:00.024013Z","iopub.status.idle":"2025-03-31T10:20:41.881208Z","shell.execute_reply.started":"2025-03-31T10:20:00.023973Z","shell.execute_reply":"2025-03-31T10:20:41.880207Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if VALIDATION:\n    LABEL_DF = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/train_labels.csv')\n    LABEL_DF['target_id'] = LABEL_DF['ID'].apply(lambda x: '_'.join(x.split('_')[:-1]))\n    train_df=pd.read_csv('/kaggle/input/stanford-rna-3d-folding/train_sequences.csv')\n    train_df['chai_tm_score']=None\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T10:22:07.861020Z","iopub.execute_input":"2025-03-31T10:22:07.861480Z","iopub.status.idle":"2025-03-31T10:22:08.437137Z","shell.execute_reply.started":"2025-03-31T10:22:07.861447Z","shell.execute_reply":"2025-03-31T10:22:08.436149Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nif MODEL_TYPE=='chai' and VALIDATION:\n    \n    num_data = len(train_df)\n    for i, seq in tqdm(enumerate(train_df.sequence),total=num_data):\n        if train_df.loc[i,'chai_tm_score']!=None:\n            continue\n        if len(seq)>300:\n            continue\n        target_id = train_df.loc[i,'target_id']\n        truth_df = get_truth_df(target_id)\n        if sum(~np.isnan(truth_df.x_1))<3:\n            continue\n        try:\n            prediction=run_folding(seq, target_id, num_diffn_samples=1)\n            assert prediction.shape[1]==len(seq)\n        except KeyboardInterrupt:\n            raise KeyboardInterrupt\n        except:\n            continue\n        result = parse_output_to_df(prediction, seq, target_id)[0]\n        try:\n            tm_score, transform = call_usalign(result, truth_df, verbose=0)\n            train_df.loc[i,'chai_tm_score']=tm_score\n        except:\n            pass\n        if (time.time()-time0)>(12*3600-360):\n            break\n    train_df.to_csv('chai_tm_scores.csv', index=False)\n    print(train_df.chai_tm_score.mean())\n    display(train_df.chai_tm_score.hist())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T10:22:10.604386Z","iopub.execute_input":"2025-03-31T10:22:10.604739Z","execution_failed":"2025-03-31T10:24:32.314Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if MODEL_TYPE=='chai' and not VALIDATION:\n    test_df=pd.read_csv('/kaggle/input/stanford-rna-3d-folding/test_sequences.csv')\n\n    num_data = len(test_df)\n    for i, seq in tqdm(enumerate(test_df.sequence),total=num_data):\n        try:\n            assert (time.time()-time0)<(12*3600-360)\n            target_id = test_df.loc[i,'target_id']\n            prediction=run_folding(seq, target_id)\n            assert prediction.shape[1]==len(seq)\n            result = parse_output_to_df(prediction, seq, target_id)[0]\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('submission.csv', index=False, mode='a', header=(i==0))\n        torch.cuda.empty_cache()\n\n    display(pd.read_csv('submission.csv'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T08:37:07.853174Z","iopub.status.idle":"2025-03-31T08:37:07.853476Z","shell.execute_reply":"2025-03-31T08:37:07.853321Z"}},"outputs":[],"execution_count":null}]}