{"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":[{"sourceId":87793,"databundleVersionId":12024591,"sourceType":"competition"},{"sourceId":10855324,"sourceType":"datasetVersion","datasetId":6742586},{"sourceId":11118830,"sourceType":"datasetVersion","datasetId":6933267},{"sourceId":11243863,"sourceType":"datasetVersion","datasetId":7025274},{"sourceId":11243930,"sourceType":"datasetVersion","datasetId":7025324},{"sourceId":11243958,"sourceType":"datasetVersion","datasetId":7025344},{"sourceId":11244174,"sourceType":"datasetVersion","datasetId":7025496},{"sourceId":11661237,"sourceType":"datasetVersion","datasetId":7317710},{"sourceId":11683946,"sourceType":"datasetVersion","datasetId":7025302},{"sourceId":224830487,"sourceType":"kernelVersion"}],"dockerImageVersionId":30918,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\n# os.environ['CUDA_VISIBLE_DEVICES'] = '0'\n\nMODEL_TYPE='protenix'\nVALIDATION=False\nTRRNA_THRESHOLD=0\nSKIP_PROTENIX=True\nNUM_CPU=os.cpu_count()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Install requirements ","metadata":{}},{"cell_type":"code","source":"!pip install --no-deps '/kaggle/input/dependencies-tr-pr/protenix-0.4.6-py3-none-any.whl'\n!pip install --no-deps '/kaggle/input/dependencies-tr-pr/rdkit-2024.9.6-cp310-cp310-manylinux_2_28_x86_64.whl'\n!pip install --no-deps '/kaggle/input/dependencies-tr-pr/ml_collections-1.1.0-py3-none-any.whl'\n!pip install --no-deps '/kaggle/input/dependencies-tr-pr/biopython-1.85-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl'\n!pip install --no-deps '/kaggle/input/dependencies-tr-pr/pyrosetta-2025.13-cp310-cp310-linux_x86_64.whl'\n!pip install --no-deps '/kaggle/input/dependencies-tr-pr/blosc-1.11.2-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl'\n!pip install --no-deps '/kaggle/input/dependencies-tr-pr/ml_collections-1.1.0-py3-none-any.whl'\n!pip install --no-deps '/kaggle/input/dependencies-tr-pr/biotraj-1.2.2-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl'\n!pip install --no-deps '/kaggle/input/dependencies-tr-pr/biotite-1.0.1-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl'\n!pip install --no-deps '/kaggle/input/dependencies-tr-pr/biopandas-0.5.1-py3-none-any.whl'\n!pip install --no-deps '/kaggle/input/dependencies-tr-pr/looseversion-1.1.2-py3-none-any.whl'\n!export PROTENIX_DATA_ROOT_DIR=/kaggle/input/protenix-checkpoints","metadata":{"trusted":true},"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},"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":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","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# './usalign/USalign'\n\nos.system('cp /kaggle/input/usalign/USalign /kaggle/working/')\nos.system('sudo chmod u+x /kaggle/working//USalign')\n# os.system('chmod u+x ./usalign/USalign')\nsys.path.append(RHONET_DIR)\n\n\nDATA_KAGGLE_DIR = '/kaggle/input/stanford-rna-3d-folding'\n# DATA_KAGGLE_DIR = './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\nLABEL_DF= None\n\n# xyz df helper --------------------\ndef get_truth_df(target_id):\n    global LABEL_DF\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    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                if residue.get_resname() in ['A', 'U', 'G', 'C']:\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                        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\n            if len(chain_data)!=0:\n                chain_df = pd.DataFrame(chain_data)\n                df.append(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},"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},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if MODEL_TYPE=='protenix':\n\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    configs_base[\"use_deepspeed_evo_attention\"] = (\n    os.environ.get(\"USE_DEEPSPEED_EVO_ATTENTION\", False) == \"true\")\n    configs_base[\"model\"][\"N_cycle\"] = 10 #10\n    configs_base[\"sample_diffusion\"][\"N_sample\"] = (1 if VALIDATION else 5)\n    configs_base[\"sample_diffusion\"][\"N_step\"] = 200\n    inference_configs['load_checkpoint_path']='/kaggle/input/protenix-checkpoints/model_v0.2.0.pt'\n    # inference_configs['load_checkpoint_path']='./protenix-checkpoints/model_v0.2.0.pt'\n    configs = {**configs_base, **{\"data\": data_configs}, **inference_configs}\n\n    configs = parse_configs(\n            configs=configs,\n            fill_required_with_null=True,\n        )\n    \n    runner=InferenceRunner(configs)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if VALIDATION:\n    global LABEL_DF\n    LABEL_DF = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/train_labels.csv')\n    # LABEL_DF = pd.read_csv('./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=pd.read_csv('./stanford-rna-3d-folding/train_sequences.csv')\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nif MODEL_TYPE=='protenix' and VALIDATION:\n    import warnings\n    warnings.filterwarnings(\"ignore\")\n    \n    train_df['protenix_tm_score']=None\n    dataset = DictDataset(train_df.sequence, dump_dir='output', id_list=train_df.target_id, use_msa=False)\n    num_data = len(dataset)\n    for i, seq in tqdm(enumerate(train_df.sequence),total=num_data):\n        if train_df.loc[i,'protenix_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        data, atom_array, data_error_message=dataset[i]\n        if data_error_message!='':\n            continue\n        new_configs = update_inference_configs(configs, data[\"N_token\"].item())\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[:1], seq, target_id)[0]\n        try:\n            tm_score, transform = call_usalign(result, truth_df, verbose=0)\n            train_df.loc[i,'protenix_tm_score']=tm_score\n        except:\n            pass\n        if (time.time()-time0)>(12*3600-360):\n            break\n    train_df.to_csv('tm_scores.csv', index=False)\n    print(train_df.protenix_tm_score.mean())\n    display(train_df.protenix_tm_score.hist())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trRNA_input_ids = []\nresults = []\npr_results = {}","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if MODEL_TYPE=='protenix' and not VALIDATION:\n    test_df=pd.read_csv('/kaggle/input/stanford-rna-3d-folding/test_sequences.csv')\n    import warnings\n    warnings.filterwarnings(\"ignore\")  \n    \n    dataset = DictDataset(test_df.sequence, dump_dir='output', id_list=test_df.target_id, use_msa=False)\n    num_data = len(dataset)\n    for i, seq in tqdm(enumerate(test_df.sequence),total=num_data):\n        try:\n            if SKIP_PROTENIX:\n                raise ValueError('Skip Protenix')\n                \n            data, atom_array, data_error_message=dataset[i]\n            \n            target_id = data[\"sample_name\"]\n            assert target_id==test_df.target_id[i]\n            assert data_error_message==''\n            # print(target_id)\n                \n                \n            new_configs = update_inference_configs(configs, data[\"N_token\"].item())\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\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        \n        if(len(seq)<=TRRNA_THRESHOLD):\n            trRNA_input_ids.append(target_id)\n            print('Send to trRNA', target_id)\n            pr_results[target_id] = result\n            # continue\n        else:\n            results.append(result)\n            result.to_csv('submission.csv', index=False, mode='a', header=(i==0))\n        torch.cuda.empty_cache()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if len(results):\n    display(pd.read_csv('submission.csv'))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# trRNA","metadata":{}},{"cell_type":"code","source":"import pdb\nimport os\nimport tempfile\nimport shutil\nimport subprocess\n\nimport tempfile\n\nimport glob\n\nimport string\n\nimport json\nimport os\nimport sys\nimport numpy as np\nimport torch\nimport torch.nn as nn\n\nfrom collections import defaultdict\nfrom argparse import ArgumentParser\nfrom pathlib import Path","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import warnings\nwarnings.simplefilter(\"ignore\", category=FutureWarning)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.cuda.is_available(),torch.cuda.device_count()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport sys\nimport shutil\nfrom pathlib import Path\n\n# 1. 定义需要复制的关键目录\ninput_base = '/kaggle/input'\nwork_base = '/kaggle/working'\n\n# 需要复制的关键组件目录（根据你的实际输入结构调整）\nrequired_dirs = [\n    'folding',\n    'network',\n    'spot-rna',\n    'training',\n    'model-1'\n]\n\n# 2. 清空工作目录并重建结构（防止残留文件干扰）\n# shutil.rmtree(work_base, ignore_errors=True)  # 慎用！确保你知道自己在做什么\n# os.makedirs(work_base, exist_ok=True)\n\n# 3. 执行复制操作（保留目录结构）\nfor dir_name in required_dirs:\n    src = os.path.join(input_base, dir_name)\n    dst = os.path.join(work_base, dir_name)\n    \n    if os.path.exists(src):\n        shutil.copytree(src, dst)\n        print(f\"Copied: {src} => {dst}\")\n    else:\n        raise FileNotFoundError(f\"关键目录缺失: {src}\")\n    \nsys.path.append(\"/kaggle/working/\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.cuda.is_available(),torch.cuda.device_count()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport string\n\n\ndef parse_a3m(filename, limit=20000, rm_query_gap=True):\n    seqs = []\n    table = str.maketrans(dict.fromkeys(string.ascii_lowercase))\n\n    # read file line by line\n    n = 0\n    for line in open(filename, \"r\"):\n        if line[0] != '>' and len(line.strip()) > 0:\n            seqs.append(\n                line.rstrip().replace('W', 'A').replace('R', 'A').replace('Y', 'C').replace('E', 'A').replace('I',\n                                                                                                              'A').replace(\n                    'P', 'G').replace('T', 'U').translate(table))\n            n += 1\n            if n == limit:\n                break\n\n    # convert letters into numbers\n    alphabet = np.array(list(\"AUCG-\"), dtype='|S1').view(np.uint8)\n    msa = np.array([list(s) for s in seqs], dtype='|S1').view(np.uint8)\n    for i in range(alphabet.shape[0]):\n        msa[msa == alphabet[i]] = i\n\n        # treat all unknown characters as gaps\n    msa[msa > 4] = 4\n    if rm_query_gap:\n        return msa[:, msa[0] < 4]\n    return msa\n\n\ndef ss2mat(ss_seq):\n    ss_mat = np.zeros((len(ss_seq), len(ss_seq)))\n    stack = []\n    stack1 = []\n    stack2 = []\n    stack3 = []\n    stack_alpha = {alpha: [] for alpha in string.ascii_lowercase}\n    for i, s in enumerate(ss_seq):\n        if s == '(':\n            stack.append(i)\n        elif s == ')':\n            ss_mat[i, stack.pop()] = 1\n        elif s == '[':\n            stack1.append(i)\n        elif s == ']':\n            ss_mat[i, stack1.pop()] = 1\n        elif s == '{':\n            stack2.append(i)\n        elif s == '}':\n            ss_mat[i, stack2.pop()] = 1\n        elif s == '<':\n            stack3.append(i)\n        elif s == '>':\n            ss_mat[i, stack3.pop()] = 1\n        elif s.isalpha() and s.isupper():\n            stack_alpha[s.lower()].append(i)\n        elif s.isalpha() and s.islower():\n            ss_mat[i, stack_alpha[s].pop()] = 1\n        elif s in ['.', ',', '_', ':', '-']:\n            continue\n        else:\n            raise ValueError(f'unk not: {s}!')\n    allstacks = stack + stack1 + stack2 + stack3\n    for _, stack in stack_alpha.items():\n        allstacks += stack\n    if len(allstacks) > 0:\n        raise ValueError('Provided dot-bracket notation is not completely matched!')\n\n    ss_mat += ss_mat.T\n    return ss_mat\n\n\ndef parse_ct(ct_file, length=None):\n    seq_ct = ''\n    if length is None:\n        length = int(open(ct_file).readlines()[0].split()[0])\n    mat = np.zeros((length, length))\n    for line in open(ct_file):\n        items = line.split()\n        if len(items) >= 6 and items[0].isnumeric() and items[2].isnumeric() and items[3].isnumeric() and items[\n            4].isnumeric():\n            seq_ct += items[1]\n            if int(items[4]) > 0:\n                mat[int(items[4]) - 1, int(items[5]) - 1] = 1\n                mat[int(items[5]) - 1, int(items[4]) - 1] = 1\n    return mat","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pkg_dir = '/kaggle/working'\n# sys.path.insert(0, pkg_dir)\n# sys.path.insert(1, f'{pkg_dir}/network')\n# from utils import *\nfrom network.RNAformer import DistPredictor\nfrom network.config import n_bins, obj\ndevice = torch.device('cuda:0') if torch.cuda.is_available() else torch.device('cpu')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import time\nfrom datetime import datetime\n\ndef predict(model, msa, ss_, window=150, shift=50):\n    start_time = time.time()\n    \n    if ss_.shape[0] != msa.shape[-1]:\n        raise ValueError(f'ss length {ss_.shape[0]}, msa length {msa.shape[1]}!')\n    with torch.no_grad():\n        feat = torch.from_numpy(msa).to(device)\n        ss_ = torch.from_numpy(ss_).to(device)\n        L = msa.shape[-1]\n        res_id = torch.arange(L, device=device).view(1, L)\n        if L > 300:  # predict by crops for long RNA\n            pred_dict = {\n                'contact': torch.zeros((L, L), device=device),\n                'distance': {k: torch.zeros((L, L, n_bins['2D']['distance']), device=device) for k in\n                             obj['2D']['distance']},\n            }\n\n            count_1d = torch.zeros((L)).to(device)\n            count_2d = torch.zeros((L, L)).to(device)\n            \n            grids = np.arange(0, L - window + shift, shift)\n            ngrids = grids.shape[0]\n            print(\"ngrid:     \", ngrids)\n            print(\"grids:     \", grids)\n            print(\"windows:   \", window)\n\n            idx_pdb = torch.arange(L).long().view(1, L)\n            for i in range(ngrids):\n                for j in range(i, ngrids):\n                    start_1 = grids[i]\n                    end_1 = min(grids[i] + window, L)\n                    start_2 = grids[j]\n                    end_2 = min(grids[j] + window, L)\n                    sel = np.zeros((L)).astype(np.bool_)\n                    sel[start_1:end_1] = True\n                    sel[start_2:end_2] = True\n\n                    input_msa = feat[:, sel]\n                    input_ss = ss_[sel][:, sel]\n                    mask = torch.sum(input_msa == 4, dim=-1) < .7 * sel.sum()  # remove too gappy sequences\n\n                    input_msa = input_msa[mask]\n                    input_idx = idx_pdb[:, sel]\n                    input_res_id = res_id[:, sel]\n\n                    print(\"running crop: %d-%d/%d-%d\" % (start_1, end_1, start_2, end_2), input_msa.shape)\n                    pred_gemos = model(input_msa, input_ss, res_id=input_res_id.to(device),\n                                       msa_cutoff=args.nrows)['geoms']\n                    weight = 1\n                    sub_idx = input_idx[0].cpu()\n                    sub_idx_2d = np.ix_(sub_idx, sub_idx)\n                    count_2d[sub_idx_2d] += weight\n                    count_1d[sub_idx] += weight\n\n                    for k in obj['2D']:\n                        if k == 'contact':\n                            pred_dict['contact'][sub_idx_2d] += weight * pred_gemos['contact']\n                        else:\n                            for a in obj['2D'][k]:\n                                pred_dict[k][a][sub_idx_2d] += weight * pred_gemos[k][a]\n                    \n                    current_time = datetime.now().strftime(\"%Y-%m-%d %H:%M:%S\")\n                    print(f\"[{current_time}] Completed crop {i}-{j}, took {time.time() - start_time:.2f} seconds\")\n                    \n            for k in obj['2D']:\n                if k == 'contact':\n                    pred_dict['contact'] /= count_2d\n                else:\n                    for a in obj['2D'][k]:\n                        if pred_dict[k][a].size().__len__() == 3:\n                            pred_dict[k][a] /= count_2d[:, :, None]\n                        else:\n                            pred_dict[k][a] /= count_2d\n        else:\n            pred_dict = model(feat, ss_, res_id=res_id.to(device), msa_cutoff=args.nrows)['geoms']\n            current_time = datetime.now().strftime(\"%Y-%m-%d %H:%M:%S\")\n            print(f\"[{current_time}] Completed prediction, took {time.time() - start_time:.2f} seconds\")\n\n    for l in pred_dict:\n        if isinstance(pred_dict[l], dict):\n            for k in pred_dict[l]:\n                pred_dict[l][k] = pred_dict[l][k].cpu().detach().numpy()\n        else:\n            pred_dict[l] = pred_dict[l].cpu().detach().numpy()\n\n    return pred_dict\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main(args):\n    os.environ[\"CUDA_VISIBLE_DEVICES\"] = str(args.gpu)\n    torch.set_num_threads(args.cpu)\n    device = torch.device('cuda:0') if torch.cuda.is_available() else torch.device('cpu')\n    py = sys.executable\n\n    out_dir = os.path.dirname(os.path.abspath(args.npz))\n    os.makedirs(out_dir, exist_ok=True)\n\n    cwd = os.getcwd()\n    msa = parse_a3m(args.msa, limit=20000)\n    # pdb.set_trace()\n\n    # Check if ss_file is not provided\n    if args.ss_file is None:\n        # Predict SS by SPOT-RNA\n        print('predict SS by SPOT-RNA')\n\n        # Create a temporary directory\n        tmp_dir = tempfile.TemporaryDirectory(prefix=out_dir + '/')\n        spot_out_dir = tmp_dir.name\n\n        # Check if the seq.fasta file doesn't exist, and create it by extracting the first 2 lines from the MSA file\n        if not os.path.isfile(f'{spot_out_dir}/seq.fasta'):\n            with open(f'{spot_out_dir}/seq.fasta', 'w') as fasta_file:\n                with open(args.msa, 'r') as msa_file:\n                    # Write the first 2 lines of MSA file to seq.fasta\n                    for _ in range(2):\n                        line = msa_file.readline()\n                        fasta_file.write(line)\n\n        # Modify Python environment path for SPOT-RNA\n        # spot_py = py.replace('trRNA', 'spot_venv')\n        spot_py = py\n\n        # Change current working directory to SPOT_RNA directory\n        os.chdir(f'{pkg_dir}/spot-rna')\n\n        # If 'utils' directory exists, rename it to 'utils_spot'\n        if os.path.isdir(f'utils'):\n            shutil.move(f'utils', f'utils_spot')\n\n        # Modify SPOT-RNA.py to use the new utils_spot directory\n        with open('SPOT-RNA.py', 'r') as spot_script_file:\n            spot_script = spot_script_file.read()\n\n\n        # Run SPOT-RNA using subprocess with nohup equivalent in Python\n        # print(spot_py)\n        # print(spot_out_dir)\n        with open(f'{out_dir}/spot.log', 'w') as log_file:\n            # pdb.set_trace()\n            subprocess.run(\n                [spot_py, 'SPOT-RNA.py', '--inputs', f'{spot_out_dir}/seq.fasta', '--outputs', spot_out_dir, '--gpu', str(args.gpu)],\n                stdout=log_file, stderr=subprocess.STDOUT\n            )\n\n        # Return to the original working directory\n        os.chdir(cwd)\n\n        prob_files = glob.glob(f'{spot_out_dir}/*.prob')\n        if len(prob_files) == 0: raise ValueError(\n            f'Fails to predict SS! Please refer to {out_dir}/spot.log to see what happened.')\n        ss = np.loadtxt(prob_files[0])\n        if (np.tril(ss) == 0).all() or (np.triu(ss) == 0).all():\n            ss += ss.T\n    else:\n        if args.ss_fmt == 'dot_bracket':\n            ss = ss2mat(open(args.ss_file).read().rstrip().splitlines()[-1].strip())\n        elif args.ss_fmt == 'ct':\n            ss = parse_ct(args.ss_file, length=len(msa[0]))\n        elif args.ss_fmt == 'spot_prob':\n            ss = np.loadtxt(args.ss_file)\n            ss += ss.T\n        if len(ss) != len(msa[0]):\n            raise ValueError(f'The SS shape {ss.shape} mismatches the MSA shape {msa.shape}!')\n\n    print('predict geometries')\n    config = json.load(open(f'{args.model_pth}/config/model_1.json', 'r'))\n\n    model = DistPredictor(dim_2d=config['channels'], layers_2d=config['n_blocks'])\n\n    model_ckpt = torch.load(f'{args.model_pth}/models/model_1.pth.tar', map_location=device)\n    model.load_state_dict(model_ckpt)\n    model.eval()\n    model.to(device)\n\n    pred = predict(model, msa, ss)\n\n    print('done!')\n    print('saving......')\n    np.savez_compressed(args.npz, **pred)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from folding.utils_cst import npz2cst","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\n\nfrom pyrosetta import *\nimport concurrent.futures\nimport os\nimport math\nimport random","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def fold_from_cst(args):\n    global rna_lowres_sf\n    init(\n        '-mute all -hb_cen_soft  -relax:dualspace true -relax:default_repeats 3 -default_max_cycles 200 -detect_disulf -detect_disulf_tolerance 3.0')\n\n    os.environ[\"OPENBLAS_NUM_THREADS\"] = \"1\"\n    rna_lowres_sf = rosetta.core.scoring.ScoreFunctionFactory.create_score_function(\"rna/denovo/rna_lores_with_rnp_aug.wts\")\n\n    global op_score\n    seq = read_fasta(args.FASTA).replace('T', 'U')\n    args.seq = seq.lower()\n\n    op_score = create_score_function('ref2015')\n    op_score.set_weight(rosetta.core.scoring.atom_pair_constraint, 9.0)\n    op_score.set_weight(rosetta.core.scoring.dihedral_constraint, 4.0)\n    op_score.set_weight(rosetta.core.scoring.angle_constraint, 4.0)\n\n    op_score.set_weight(rosetta.core.scoring.fa_rep, 9.0)\n    op_score.set_weight(rosetta.core.scoring.rna_sugar_close, 9.0)\n    op_score.set_weight(rosetta.core.scoring.fa_intra_rep, 9)\n    op_score.set_weight(rosetta.core.scoring.rna_base_pair, 9)\n    op_score.set_weight(rosetta.core.scoring.rna_base_stack, 9)\n\n    cutoff_dist = args.dcut\n    cutoff_angle = 0.75\n    cutoff_dihedral = 0.85\n    cutoff_cont = 0.6\n\n    cstpath = args.tmpdir + f'/cstfile_dist.txt'\n    cstpath_cont = args.tmpdir + f'/cstfile_cont.txt'\n\n    allcst = read_cst(cstpath)\n    allcst_cont = read_cst(cstpath_cont)\n\n    # cst_op=fetch_cst_atomset(allcst,atomset)\n    cst_op = allcst + allcst_cont\n    pcut = {\n        'AtomPair': cutoff_dist,\n        'Dihedral': cutoff_dihedral,\n    }\n\n    \"\"\"  \\mode=1\\\\\"\"\"\n    # using all cst\n    sep1 = 1\n    sep2 = 10000\n\n    minstd = 0.01\n    std_cut = minstd  # +0.01*pcut['AtomPair']\n    cst_all = fetch_cst(cst_op, sep1, sep2, pcut, std_cut)\n\n    cstpath_all = args.tmpdir + f'/cstfile.txt'\n    F = open(cstpath_all, \"w\")\n    for a in cst_all:\n        F.write(a)\n        F.write(\"\\n\")\n    F.close()\n\n    executor = concurrent.futures.ProcessPoolExecutor(args.CPU)\n    futures = [executor.submit(fold_single, cstpath_all, args) for _ in range(args.nmodels)]\n    results = concurrent.futures.wait(futures)\n    poses = list(results[0])\n\n    min_energy = np.inf\n    score_dict = {}\n    for pose in poses:\n        pose = pose.result()\n        energy = op_score(pose)\n        score_dict[pose] = energy\n        if energy < min_energy:\n            best_pose = pose\n            min_energy = energy\n\n    top5 = sorted(score_dict, key=lambda d: score_dict[d])[:args.nmodels]\n    for i, pose in enumerate(top5):\n        pose.dump_pdb(f'{args.tmpdir}/model_{i + 1}.pdb')\n    ermsd = eRMSD(args.NPZ, args.tmpdir, nmodels=args.nmodels)\n    name = args.OUT\n    best_pose.dump_pdb(name)\n    print('\\ndone')\n    print(f'eRMSD = {ermsd:.2f}')\n    lines = open(name).readlines()\n    lines[0] = lines[0].strip() + f'  eRMSD = {ermsd:.2f}\\n'\n    with open(name, 'w') as f:\n        f.write(''.join(lines))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# if len(trRNA_input_ids)==0:\n#     exit()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def fold_single(cst_file, args):\n    mmap = MoveMap()\n    mmap.set_bb(True)  ##Whether the frame dihedral angle changes\n    mmap.set_chi(True)  ##Whether the dihedral angle of the side chain changes\n    mmap.set_jump(True)  ##Relative movement between polypeptide chains\n\n    min_mover = rosetta.protocols.minimization_packing.MinMover(mmap, op_score, 'lbfgs_armijo_nonmonotone', 0.0001,\n                                                                True)\n    min_mover.max_iter(10000)\n\n    repeat_mover = RepeatMover(min_mover, 5)\n\n    mmap_ = MoveMap()\n    mmap_.set_bb(True)  ##Whether the frame dihedral angle changes\n    mmap_.set_chi(True)  ##Whether the dihedral angle of the side chain changes\n    mmap_.set_jump(True)  ##Relative movement between polypeptide chains\n\n    assembler = rosetta.core.import_pose.RNA_HelixAssembler()\n    initpose = assembler.build_init_pose(args.seq, '')  # helix pose\n    pose = basic_folding(initpose)\n    pose.remove_constraints()\n\n    run_min(cst_file, pose, repeat_mover, min_mover)\n    run_refine(pose)\n\n    return pose\n\n\ndef randTrial(your_pose):\n    randNum = random.randint(2, your_pose.total_residue())\n\n    curralpha = your_pose.alpha(randNum)\n    currbeta = your_pose.beta(randNum)\n    currgamma = your_pose.gamma(randNum)\n    currdelta = your_pose.delta(randNum)\n    currepsilon = your_pose.epsilon(randNum)\n    currzeta = your_pose.zeta(randNum)\n    currchi = your_pose.chi(randNum)\n\n    newalpha = random.gauss(curralpha, 25)\n    newbeta = random.gauss(currbeta, 25)\n    newgamma = random.gauss(currgamma, 25)\n    newdelta = random.gauss(currdelta, 25)\n    newepsilon = random.gauss(currepsilon, 25)\n    newzeta = random.gauss(currzeta, 25)\n    newchi = random.gauss(currchi, 25)\n\n    your_pose.set_alpha(randNum, newalpha)\n    your_pose.set_beta(randNum, newbeta)\n    your_pose.set_gamma(randNum, newgamma)\n    your_pose.set_delta(randNum, newdelta)\n    your_pose.set_epsilon(randNum, newepsilon)\n    your_pose.set_zeta(randNum, newzeta)\n    your_pose.set_chi(randNum, newchi)\n\n    return your_pose\n\n\ndef decision(before_pose, after_pose):\n    E = score(after_pose) - score(before_pose)\n    if E < 0:\n        return after_pose\n    elif random.uniform(0, 1) >= math.exp(-E / 1):\n        return before_pose\n    else:\n        return after_pose\n\n\ndef basic_folding(your_pose):\n    lowest_pose = Pose()  # Create an empty pose for tracking the lowest energy pose.\n    for i in range(120):\n        if i == 0:\n            lowest_pose.assign(your_pose)\n\n        before_pose = Pose()\n        before_pose.assign(your_pose)  # keep track of pose before random move\n\n        after_pose = Pose()\n        after_pose.assign(randTrial(your_pose))  # do rand move and store the pose\n\n        your_pose.assign(decision(before_pose, after_pose))  # keep the new pose or old pose\n\n        if score(your_pose) < score(lowest_pose):  # updating lowest pose\n            lowest_pose.assign(your_pose)\n\n    return lowest_pose\n\n\ndef score(your_pose):\n    sf = rna_lowres_sf(your_pose)\n    return sf\n\n\ndef read_fasta(file):\n    fasta = \"\";\n    with open(file, \"r\") as f:\n        for line in f:\n            if (line[0] == \">\"):\n                continue\n            else:\n                line = line.rstrip()\n                fasta = fasta + line;\n    return fasta\n\n\ndef read_cst(file):\n    array = []\n    with open(file, \"r\") as f:\n        for line in f:\n            line = line.rstrip()\n            array.append(line)\n    return array\n\n\ndef add_cst(pose, cstfile):\n    constraints = rosetta.protocols.constraint_movers.ConstraintSetMover()\n    constraints.constraint_file(cstfile)\n    constraints.add_constraints(True)\n    constraints.apply(pose)\n\n\ndef remove_clash(scorefxn, mover, pose):\n    clash_score = float(scorefxn(pose))\n    if (clash_score > 10):\n        for nm in range(0, 10):\n            mover.apply(pose)\n            clash_score = float(scorefxn(pose))\n            if (clash_score < 10): break\n\n\ndef fetch_cst(cst, sep1, sep2, cut, std_cut=None):\n    array = []\n    for line in cst:\n        # print(line)\n        line = line.rstrip()\n        b = line.split()\n        cst_name = b[0]\n        pcut = cut[cst_name]\n\n        if std_cut is not None:\n            if not line.endswith('#cont'):\n                p_std = float(line.split('#')[1].split()[0])\n            else:\n                p_std = 10000\n        if line.endswith('#cont'):\n            prob = 1\n        else:\n            prob = float(b[-1])\n\n        i = int(b[2])\n        j = int(b[4])\n\n        sep = abs(j - i)\n        if (sep < sep1 or sep >= sep2): continue\n        if std_cut is not None and p_std < std_cut: continue\n        if (prob >= pcut):\n            array.append(line)\n    return array\n\n\n# def run_min(array, pose, mover, tmpname):\n# add_cst(pose, array, tmpname)\n# mover.apply(pose)\n\ndef run_min(cstfile, pose, mover1, mover2=None):\n    add_cst(pose, cstfile)\n    mover1.apply(pose)\n    if mover2 is not None:\n        mover2.apply(pose)\n\n\ndef run_refine(pose):\n    idealize = rosetta.protocols.idealize.IdealizeMover()\n    poslist = rosetta.utility.vector1_unsigned_long()\n\n    scorefxn = create_score_function('empty')\n    scorefxn.set_weight(rosetta.core.scoring.cart_bonded, 1.0)\n    scorefxn.score(pose)\n\n    emap = pose.energies()\n    # print(\"idealize...\")\n    for res in range(1, len(pose.residues) + 1):\n        cart = emap.residue_total_energy(res)\n        if cart > 50:\n            poslist.append(res)\n            # print(\"idealize %d %8.3f\" % (res, cart))\n\n    if len(poslist) > 0:\n        idealize.set_pos_list(poslist)\n    try:\n        idealize.apply(pose)\n\n        # cart-minimize\n        scorefxn_min = create_score_function('ref2015_cart')\n\n        mmap = MoveMap()\n        mmap.set_bb(True)\n        mmap.set_chi(True)\n        mmap.set_jump(True)\n        mmap.set_chi(False)\n\n        min_mover = rosetta.protocols.minimization_packing.MinMover(mmap, scorefxn_min, 'lbfgs_armijo_nonmonotone',\n                                                                    0.00001, True)\n        # min_mover = rosetta.protocols.minimization_packing.MinMover(mmap, op_score, 'lbfgs_armijo_nonmonotone', 0.00001, True)\n        min_mover.max_iter(200)\n        min_mover.cartesian(True)\n        # print(\"minimize...\")\n        min_mover.apply(pose)\n\n    except:\n        print('!!! idealization failed !!!')\n\n\ndef get_atom_positions_pdb(pdb_file):\n    # load PDB\n    xyz = []\n    for line in open(pdb_file):\n        if line.startswith('ATOM'):\n            atom = line[12:16].strip()\n            if atom == 'P':\n                x = float(line[30:38])\n                y = float(line[38:46])\n                z = float(line[46:54])\n                xyz.append([x, y, z])\n    return np.array(xyz)\n\n\ndef pmax(pred, sep=12, dist_cutoff=40, top=15):\n    first_bin = 1\n    first_value = 3.5\n    bin_size = 1\n    vmax = 40\n    last_bin = int((vmax - 3) / bin_size)\n    judge_bin = int((dist_cutoff - 3) / bin_size)\n\n    if len(pred.shape) == 4:\n        assert pred.shape[0] == 1\n        dat = pred[0]\n    else:\n        dat = pred\n    n_bins = last_bin - first_bin + 1\n\n    L = int(dat.shape[0])\n\n    idx = np.array([[i, j, *dat[i, j, first_bin:last_bin]] for j in range(L) for i in range(j + sep, L)])\n\n    topn = top * L\n    top_pairs = np.array(sorted(idx, key=lambda x: dat[int(x[0]), int(x[1]), first_bin:judge_bin].sum(-1))[-topn:])\n    bins = np.argmax(top_pairs[:, 2:], axis=1)\n    probs = top_pairs[:, 2:][range(len(top_pairs)), bins]\n\n    means = []\n    for i in range(n_bins):\n        idx = np.where(bins == i)[0]\n        if len(idx) != 0:\n            imean = np.mean(probs[idx])\n            means.append(imean)\n    top_dist = np.mean(means)\n    return top_dist\n\n\ndef STD(pred, pcut=0.45, bin_p=0.005, return_prop=False):\n    first_bin = 1\n    first_value = 3.5\n    bin_size = 1\n    vmax = 40\n    last_bin = int((vmax - 3) / bin_size)\n\n    if len(pred.shape) == 4:\n        assert pred.shape[0] == 1\n        dat = pred[0]\n    else:\n        dat = pred\n    n_bins = last_bin - first_bin + 1\n    bin_values = first_value + bin_size * np.arange(n_bins)\n    p_valid = np.where(dat > bin_p, dat, 0)[..., first_bin:last_bin + 1]\n\n    std_mat = np.where((p_valid.sum(-1) > pcut), p_valid.std(axis=-1), np.nan)\n    std = np.nanmean(std_mat)\n    if return_prop:\n        return std, (p_valid.sum(-1) > pcut).sum() / (p_valid.sum(-1) > -1).sum()\n    return std\n\n\ndef cRMSD(pred_points, true_points):\n    if pred_points.shape[0] == 0:\n        pred_points = pred_points[0]\n    if true_points.shape[0] == 0:\n        true_points = true_points[0]\n\n    nan_idx = np.isnan(true_points[:, 0])\n    true_points = true_points[~nan_idx]\n    pred_points = pred_points[~nan_idx]\n    true_points -= true_points.mean(0)[None]\n    pred_points -= pred_points.mean(0)[None]\n    L = true_points.shape[0]\n\n    h = pred_points.T @ true_points\n    u, s, vt = np.linalg.svd(h)\n    v = vt.T\n    u = u\n\n    d = np.linalg.det(v @ u.T)\n    e = np.array([[1, 0, 0], [0, 1, 0], [0, 0, d]])\n\n    r = v @ e @ u.T\n    rotated = np.einsum('ij,lj->li', r, pred_points)\n    return (((rotated - true_points) ** 2).sum() / L) ** .5\n\n\ndef eRMSD(npz, data_dir,nmodels=5):\n    distP = np.load(npz, allow_pickle=True)['distance'].item()['P']\n    mp40 = pmax(distP, dist_cutoff=40)\n    std, prop = STD(distP, pcut=0, bin_p=0.005, return_prop=True)\n\n    pair_rmsds = []\n    for i in range(1, nmodels+1):\n        for j in range(i, nmodels+1):\n            p1 = get_atom_positions_pdb(f'{data_dir}/model_{i}.pdb')\n            p2 = get_atom_positions_pdb(f'{data_dir}/model_{j}.pdb')\n            pair_rmsds.append(cRMSD(p1, p2))\n    pair_rmsd = np.mean(pair_rmsds)\n\n    ermsd = 0.64279301 * pair_rmsd - 189.42928276 * std - 1.06122834 * prop - 4.01096737 * mp40 + 15.198044839491034\n    return max(0.1, ermsd)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def fold(args):\n    os.makedirs(os.path.dirname(os.path.abspath(args.OUT)), exist_ok=True)\n\n    tmpdir = tempfile.TemporaryDirectory(prefix=args.TMPDIR + '/')\n    args.tmpdir = tmpdir.name\n    print('temp folder:     ', tmpdir.name)\n\n    # parse npz into rosetta-format restraint files\n    npz2cst(args)\n\n    # 修改输出路径，为每个模型创建单独的文件\n    original_out = args.OUT\n    base, ext = os.path.splitext(original_out)\n    \n    # perform energy minimization for each model\n    args.OUT = f\"{base}_model_1{ext}\"\n    fold_from_cst(args)\n    for i in range(2,args.nmodels+1):\n        os.system(f\"cp {tmpdir.name}/model_{i}.pdb {base}_model_{i}{ext}\");\n    \n    # 恢复原始输出路径\n    args.OUT = original_out","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nfrom biopandas.pdb import PandasPdb\n\n\n# 读取所有PDB文件并存储每个模型的C1'坐标\ndef get_c1_coords(base=\"/kaggle/working/output_model_model\", ext=\".pdb\"):\n    c1_coords_all_models = {}  # 存储所有模型的C1'坐标\n    # 遍历5个模型\n    for i in range(1, 6):\n        model_path = f\"{base}_{i}{ext}\"\n        ppdb = PandasPdb().read_pdb(model_path)\n        atom_df = ppdb.df['ATOM']\n        \n        # 按残基分组获取C1'坐标\n        for resid, group in atom_df.groupby('residue_number'):\n            resname = group['residue_name'].iloc[0]\n            c1_coords = group[group['atom_name'] == \"C1'\"][['x_coord', 'y_coord', 'z_coord']].values\n            \n            if resid not in c1_coords_all_models:\n                c1_coords_all_models[resid] = {'resname': resname, 'coords': []}\n                \n            if len(c1_coords) > 0:\n                c1_coords_all_models[resid]['coords'].extend(c1_coords[0])\n            else:\n                c1_coords_all_models[resid]['coords'].extend([0.0, 0.0, 0.0])\n\n    # 输出格式化结果\n    rna_id = os.path.basename(base).split('_')[1]\n    # for resid in sorted(c1_coords_all_models.keys()):\n    #     data = c1_coords_all_models[resid]\n    #     coords = data['coords']\n    #     print(f\"{rna_id}_{resid},{data['resname']},{resid},\" + \",\".join(f\"{x:.3f}\" for x in coords))\n        \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=[[f\"{rna_id}_{resid}\", data['resname'], resid] + \n                                          [data['coords'][j] for j in range(len(data['coords']))] for resid, data in c1_coords_all_models.items()])\n    return result","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ori_dir='/kaggle/input/stanford-rna-3d-folding/MSA'\ntarget_dir='/kaggle/working/simplified-msa/MSA1'\nos.makedirs(target_dir, exist_ok=True)\n# 将ori_dir中所有fasta文件复制到target_dir中, 只保留前两行, 文件名不变\nfor filename in os.listdir(ori_dir):\n    if filename.endswith('.fasta'):\n        with open(os.path.join(ori_dir, filename), 'r') as f:\n            lines = f.readlines()\n        with open(os.path.join(target_dir, filename), 'w') as f:\n            f.writelines(lines[:2])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 示例：手动配置参数\nclass Args:\n    msa=\"/kaggle/input/simplified-msa/MSA1/R1107.MSA.fasta\"\n    npz = \"/kaggle/working/R1107.npz\"\n    ss_file = None\n    ss_fmt = \"dot_bracket\"\n    gpu = \"0\"\n    cpu = NUM_CPU\n    model_pth=\"/kaggle/working/model-1\"\n    nrows=1000","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Fold_Args:\n    NPZ = \"/kaggle/working/R1107.npz\"  # 替换为实际的NPZ文件路径\n    FASTA = \"/kaggle/input/simplified-msa/MSA1/R1107.MSA.fasta\"\n    OUT = \"/kaggle/working/output_model.pdb\"              # 输出PDB文件路径\n    \n    # ==== 可选参数（保留默认值或按需修改）====\n    TMPDIR = \"/kaggle/working\"    # 临时目录\n    nmodels = 5             # 生成模型数量\n    dcut = 0.45            # 距离约束截断值\n    CPU = NUM_CPU                 # 使用的CPU核心数","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"args = Args()\nfold_args = Fold_Args()\n\nfor i,input_id in enumerate(trRNA_input_ids):\n    try:\n        args.npz = f'/kaggle/working/{input_id}.npz'\n        args.msa = f'/kaggle/working/simplified-msa/MSA1/{input_id}.MSA.fasta'\n        fold_args.NPZ = f'/kaggle/working/{input_id}.npz'\n        fold_args.FASTA = f'/kaggle/working/simplified-msa/MSA1/{input_id}.MSA.fasta'\n        fold_args.OUT = f'/kaggle/working/output_{input_id}.pdb'\n        print('args',args)\n        print('fold_args',fold_args)\n        main(args)\n        fold(fold_args)\n        result=get_c1_coords(base=f'/kaggle/working/output_{input_id}_model', ext='.pdb')\n    except:\n        print('Failed to predict', input_id)\n        result=pr_results[input_id]\n        \n        \n    result.to_csv('submission.csv', index=False, mode='a', header=(i==0 and len(results)==0))\n    results.append(result)\n    torch.cuda.empty_cache()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_results = pd.concat(results, ignore_index=True)\nall_results.to_csv('submission.csv', index=False, header=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"display(pd.read_csv('submission.csv'))","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}