{"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":"none","dataSources":[{"sourceId":87793,"databundleVersionId":12276181,"sourceType":"competition"},{"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":11247625,"sourceType":"datasetVersion","datasetId":7025847},{"sourceId":11632948,"sourceType":"datasetVersion","datasetId":7037971},{"sourceId":11673760,"sourceType":"datasetVersion","datasetId":7326443},{"sourceId":11683946,"sourceType":"datasetVersion","datasetId":7025302},{"sourceId":11881964,"sourceType":"datasetVersion","datasetId":7467647}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"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":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","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. 清空工作目录并重建结构（防止残留文件干扰）\nshutil.rmtree(work_base, ignore_errors=True)  # 慎用！确保你知道自己在做什么\nos.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":"# 示例：手动配置参数\nclass Args:\n    msa = \"/kaggle/input/testdata/test/R1189.MSA.fasta\"\n    #msa=\"/kaggle/input/simplified-msa/MSA1/R1107.MSA.fasta\"\n    npz = \"/kaggle/working/R1189.npz\"\n    ss_file = None\n    ss_fmt = \"dot_bracket\"\n    gpu = \"0\"\n    cpu = 2\n    model_pth=\"/kaggle/working/model-1\"\n    nrows=1000\n\nargs = Args()\nprint(args.npz)  # 输出: /kaggle/working/XXXX.npz\nprint(\"args1ok\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"main(args)","metadata":{"trusted":true,"scrolled":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ln -s '/kaggle/input/pyrosetta/pyrosetta-2025.13release.80dd00bc09-cp310-cp310-linux_x86_64.whl' '/kaggle/working/pyrosetta-2025.13-cp310-cp310-linux_x86_64.whl'","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install '/kaggle/working/pyrosetta-2025.13-cp310-cp310-linux_x86_64.whl'","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install blosc","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install biopandas","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import tempfile\n\nimport glob\n\nimport sys\nimport os, time\nfrom pathlib import Path\nfrom folding.arguments import get_args\nfrom folding.utils_cst import npz2cst\nfrom folding.utils_ros import fold_from_cst\n\ndef 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    for i in range(1, 2):\n        args.OUT = f\"{base}_model_{i}{ext}\"\n        fold_from_cst(args)\n        print(f\"Output for nmodels={i}:\")\n        print(f\"Model {i} saved to {args.OUT}\")\n    \n    # 恢复原始输出路径\n    args.OUT = original_out","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Args:\n    # ==== 必须手动配置的路径 ====\n    NPZ = \"/kaggle/working/R1189.npz\"  # 替换为实际的NPZ文件路径\n    FASTA = \"/kaggle/input/testdata/test/R1189.MSA.fasta\"         # 替换为FASTA文件路径\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 = 1           # 生成模型数量\n    dcut = 0.45            # 距离约束截断值\n    CPU = 5                 # 使用的CPU核心数\n\nargs = Args()\nprint(args.NPZ)  # 输出: /kaggle/working/XXXX.npz\nprint(\"args2ok\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fold(args)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nfrom biopandas.pdb import PandasPdb\n\n# 定义基本路径和扩展名\nbase = \"/kaggle/working/output_model_model\"  # PDB文件的基本路径\next = \".pdb\"  # 文件扩展名\n\n# 读取所有PDB文件并存储每个模型的C1'坐标\nc1_coords_all_models = {}  # 存储所有模型的C1'坐标\n\n# 遍历nmodels + 1个模型\nfor i in range(1, 2):\n    model_path = f\"{base}_{i}{ext}\"\n    ppdb = PandasPdb().read_pdb(model_path)\n    atom_df = ppdb.df['ATOM']\n\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# 输出格式化结果\nprint(f\"Output for nmodels={i}:\")\nrna_id = os.path.basename(base)\nfor 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))\nprint(\"\\n\" + \"-\"*40)  # 分隔不同模型的输出","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}