{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":[{"sourceType":"competition","sourceId":87793,"databundleVersionId":11553390},{"sourceType":"datasetVersion","sourceId":11228135,"datasetId":7006909,"databundleVersionId":11637935},{"sourceType":"datasetVersion","sourceId":11229874,"datasetId":7012586,"databundleVersionId":11639895}],"dockerImageVersionId":30918,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport torch\nimport logging\nimport gc\nimport warnings\nwarnings.filterwarnings('ignore')\ntry:\n    import Bio.PDB\nexcept ImportError:\n    !pip install --no-index --find-links=/kaggle/input/packages biopython ml-collections python-box dm-tree openmm[cuda12]\n    import Bio.PDB\nimport sys\nimport os\nsys.path.append('/kaggle/input/rhofold-custom')\nfrom rhofold.rhofold import RhoFold\nfrom rhofold.config import rhofold_config\nfrom rhofold.utils import get_device, save_ss2ct, timing\n#from rhofold.relax.relax import AmberRelaxation\nfrom rhofold.utils.alphabet import get_features\n\nprint('Finished loading packages')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T18:53:41.322421Z","iopub.execute_input":"2025-03-31T18:53:41.322790Z","iopub.status.idle":"2025-03-31T18:53:41.329666Z","shell.execute_reply.started":"2025-03-31T18:53:41.322767Z","shell.execute_reply":"2025-03-31T18:53:41.328721Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# prepare fasta files\ntest_seqs = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/test_sequences.csv')\nall_RNA_IDs = test_seqs['target_id'].to_list()\nprint(all_RNA_IDs)\n\nmain_output_dir = '/kaggle/working'\nfasta_dir = f'{main_output_dir}/fasta-files'\nos.makedirs(fasta_dir, exist_ok=True)\n\nfor i, row in test_seqs.iterrows():\n    each_RNA_ID = row['target_id']\n    fasta = f'{fasta_dir}/{each_RNA_ID}.fasta'\n    with open(fasta, 'w') as f:\n        f.write(f'>{each_RNA_ID}\\n')\n        f.write(row['sequence'] + '\\n')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T18:53:41.331059Z","iopub.execute_input":"2025-03-31T18:53:41.331382Z","iopub.status.idle":"2025-03-31T18:53:41.352424Z","shell.execute_reply.started":"2025-03-31T18:53:41.331352Z","shell.execute_reply":"2025-03-31T18:53:41.351548Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# do predictions\npretrained = '/kaggle/input/rhofold-custom/pretrained/RhoFold_pretrained.pt'\n\n@torch.no_grad()\ndef predict(fasta, output_dir, a3m=None, ckpt=pretrained, device='cuda:0', \n            single_seq_mode=False, verbose=False):\n    \"\"\"\n    Do inference with RhoFold.\n    \"\"\"\n    os.makedirs(output_dir, exist_ok=True)\n    \n    logger = logging.getLogger('RhoFold Inference')\n    logger.setLevel(level=logging.DEBUG)\n    \n    formatter = logging.Formatter('%(asctime)s - %(levelname)s: %(message)s')\n    file_handler = logging.FileHandler(f'{output_dir}/log.txt', mode='w')\n    file_handler.setLevel(level=logging.DEBUG)\n    file_handler.setFormatter(formatter)\n    \n    stream_handler = logging.StreamHandler(sys.stdout)\n    stream_handler.setLevel(logging.DEBUG)\n    stream_handler.setFormatter(formatter)\n    \n    logger.addHandler(file_handler)\n    logger.addHandler(stream_handler)\n    \n    model = RhoFold(rhofold_config)\n    \n    model.load_state_dict(torch.load(ckpt, map_location=torch.device('cpu'))['model'])\n    model.eval()\n    \n    if single_seq_mode:\n        if verbose:\n            logger.info('Use single sequence mode')\n        a3m = fasta\n    elif a3m is None:\n        # MSA search is too slow so we have to provide MSA file.\n        if verbose:\n            logger.info('Cannot do MSA search currently, please provide MSA file')\n            logger.info('No MSA file provided, use single sequence mode')\n        a3m = fasta\n    else:\n        if verbose:\n            logger.info(f'Use {a3m} as MSA file')\n    \n    with timing('RhoFold Inference', logger=logger):\n        device = get_device(device)\n        if verbose:\n            logger.info(f'    Use device {device}')\n        model = model.to(device)\n        \n        data_dict = get_features(fasta, a3m)\n        \n        outputs = model(tokens=data_dict['tokens'].to(device),\n                       rna_fm_tokens=data_dict['rna_fm_tokens'].to(device),\n                       seq=data_dict['seq'])\n        output = outputs[-1]\n        \n        ss_prob_map = torch.sigmoid(output['ss'][0, 0]).data.cpu().numpy()\n        ss_file = f'{output_dir}/ss.ct'\n        save_ss2ct(ss_prob_map, data_dict['seq'], ss_file, threshold=0.5)\n        \n        npz_file = f'{output_dir}/results.npz'\n        np.savez_compressed(npz_file,\n                            dist_n = torch.softmax(output['n'].squeeze(0), dim=0).data.cpu().numpy(),\n                            dist_p = torch.softmax(output['p'].squeeze(0), dim=0).data.cpu().numpy(),\n                            dist_c = torch.softmax(output['c4_'].squeeze(0), dim=0).data.cpu().numpy(),\n                            ss_prob_map = ss_prob_map,\n                            plddt = output['plddt'][0].data.cpu().numpy(),\n                            )\n        unrelaxed_model = f'{output_dir}/unrelaxed_model.pdb'\n        \n        node_cords_pred = output['cord_tns_pred'][-1].squeeze(0)\n        model.structure_module.converter.export_pdb_file(data_dict['seq'],\n                                                         node_cords_pred.data.cpu().numpy(),\n                                                         path=unrelaxed_model, chain_id=None,\n                                                         confidence=output['plddt'][0].data.cpu().numpy(),\n                                                         logger=logger)\n    \n    # no relaxation\n    return None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T18:53:41.353926Z","iopub.execute_input":"2025-03-31T18:53:41.354130Z","iopub.status.idle":"2025-03-31T18:53:41.363915Z","shell.execute_reply.started":"2025-03-31T18:53:41.354112Z","shell.execute_reply":"2025-03-31T18:53:41.363046Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# skip some chains as they are too long and will cause memory issue\nskip = ['R1126', 'R1136', 'R1138']\n\nfor each_RNA_ID in all_RNA_IDs:\n    if each_RNA_ID in skip:\n        print(f'Skip {each_RNA_ID}')\n        continue\n    fasta = f'{fasta_dir}/{each_RNA_ID}.fasta'\n    output_dir = f'{main_output_dir}/RhoFold-single-seq-predictions/{each_RNA_ID}'\n    print(f'Predict {each_RNA_ID}')\n    predict(fasta, output_dir, single_seq_mode=True)\n    \n    # clear memory\n    gc.collect()\n    torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T18:53:41.364742Z","iopub.execute_input":"2025-03-31T18:53:41.364926Z","iopub.status.idle":"2025-03-31T18:54:42.152883Z","shell.execute_reply.started":"2025-03-31T18:53:41.364910Z","shell.execute_reply":"2025-03-31T18:54:42.152196Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# define function to parse pdb file\n# from github openabc repo\ndef parse_pdb(pdb_file):\n    \"\"\"\n    Load pdb file as pandas dataframe.\n    Note there should be only a single frame in the pdb file.\n    \n    Parameters\n    ----------\n    pdb_file : str\n        Path for the pdb file. \n    \n    Returns\n    -------\n    pdb_atoms : pd.DataFrame\n        A pandas dataframe includes atom information. \n    \n    \"\"\"\n    def pdb_line(line):\n        return dict(recname=str(line[0:6]).strip(),\n                    serial=int(line[6:11]),\n                    name=str(line[12:16]).strip(),\n                    altLoc=str(line[16:17]),\n                    resname=str(line[17:20]).strip(),\n                    chainID=str(line[21:22]),\n                    resSeq=int(line[22:26]),\n                    iCode=str(line[26:27]),\n                    x=float(line[30:38]),\n                    y=float(line[38:46]),\n                    z=float(line[46:54]),\n                    occupancy=0.0 if line[54:60].strip() == '' else float(line[54:60]),\n                    tempFactor=0.0 if line[60:66].strip() == '' else float(line[60:66]),\n                    element=str(line[76:78].strip()),\n                    charge=str(line[78:80].strip()))\n    with open(pdb_file, 'r') as pdb:\n        lines = []\n        for line in pdb:\n            if (len(line) > 6) and (line[:6] in ['ATOM  ', 'HETATM']):\n                lines += [pdb_line(line)]\n    pdb_atoms = pd.DataFrame(lines)\n    pdb_atoms = pdb_atoms[['recname', 'serial', 'name', 'altLoc',\n                           'resname', 'chainID', 'resSeq', 'iCode',\n                           'x', 'y', 'z', 'occupancy', 'tempFactor',\n                           'element', 'charge']]\n    return pdb_atoms","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T18:54:42.153712Z","iopub.execute_input":"2025-03-31T18:54:42.153925Z","iopub.status.idle":"2025-03-31T18:54:42.161146Z","shell.execute_reply.started":"2025-03-31T18:54:42.153907Z","shell.execute_reply":"2025-03-31T18:54:42.160153Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# prepare submission\ndf_sample_submission = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/sample_submission.csv')\ndf_submission_list = []\nfor each_RNA_ID in all_RNA_IDs:\n    if each_RNA_ID in skip:\n        continue\n    print(f'Collect coordinates for {each_RNA_ID}')\n    output_dir = f'{main_output_dir}/RhoFold-single-seq-predictions/{each_RNA_ID}'\n    pdb = f'{output_dir}/unrelaxed_model.pdb'\n    assert os.path.exists(pdb)\n    atoms = parse_pdb(pdb)\n    flag = atoms['name'] == \"C1'\"\n    C1p_atoms = atoms[flag].copy()\n    n_C1p_atoms = len(C1p_atoms.index)\n    C1p_atoms['serial'] = np.arange(1, len(C1p_atoms.index) + 1)\n    coords = C1p_atoms[['x', 'y', 'z']].to_numpy() # in unit angstrom\n    each_df_submission = pd.DataFrame()\n    each_df_submission['ID'] = [f'{each_RNA_ID}_{i + 1}' for i in range(n_C1p_atoms)]\n    each_df_submission['resname'] = C1p_atoms['resname'].to_list()\n    each_df_submission['resid'] = 1 + np.arange(n_C1p_atoms)\n    \n    # just submit one prediction, others are zeros\n    each_df_submission[['x_1', 'y_1', 'z_1']] = coords\n    for i in range(2, 6):\n        each_df_submission[[f'x_{i}', f'y_{i}', f'z_{i}']] = np.zeros((n_C1p_atoms, 3))\n    df_submission_list.append(each_df_submission)\n\n# for those skipped, just use zeros\nval_labels = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/validation_labels.csv')\nval_labels['RNA_ID'] = val_labels['ID'].apply(lambda x: x.split('_')[0])\nunique_val_RNA_IDs = val_labels['RNA_ID'].unique()\n\nfor each_RNA_ID in skip:\n    print(f'Set zero coordinates for {each_RNA_ID}')\n    flag = val_labels['RNA_ID'] == each_RNA_ID\n    cols = ['ID', 'resname', 'resid']\n    each_df_submission = val_labels.loc[flag, cols]\n    coord_cols = []\n    for i in range(1, 6): \n        coord_cols += [f'x_{i}', f'y_{i}', f'z_{i}']\n    n_C1p_atoms = len(each_df_submission.index)\n    each_df_submission.loc[:, coord_cols] = np.zeros((n_C1p_atoms, len(coord_cols)))\n    df_submission_list.append(each_df_submission)\n\ndf_submission = pd.concat(df_submission_list, axis=0, ignore_index=True)\n#df_submission['RNA_ID'] = df_submission['ID'].apply(lambda x: x.split('_')[0])\n#df_submission = df_submission.sort_values(by=['RNA_ID', 'resid'])\n#print(df_submission.head())\n#df_submission = df_submission.drop(columns=['RNA_ID'])\n\n# rearrange rows based on sample submission\nIDs = df_sample_submission['ID'].to_list()\ndf_submission.index = df_submission['ID']\ndf_submission = df_submission.loc[IDs, df_sample_submission.columns]\ndf_submission.index = np.arange(len(df_submission.index))\ndf_submission.to_csv(f'{main_output_dir}/submission.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-31T18:55:27.311907Z","iopub.execute_input":"2025-03-31T18:55:27.312289Z","iopub.status.idle":"2025-03-31T18:55:27.606694Z","shell.execute_reply.started":"2025-03-31T18:55:27.312248Z","shell.execute_reply":"2025-03-31T18:55:27.605648Z"}},"outputs":[],"execution_count":null}]}