{"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":11403143,"sourceType":"competition"},{"sourceId":10855324,"sourceType":"datasetVersion","datasetId":6742586},{"sourceId":10880419,"sourceType":"datasetVersion","datasetId":6760509},{"sourceId":11065669,"sourceType":"datasetVersion","datasetId":6889817}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"from datetime import datetime\nimport pytz\nprint('LOGGING TIME OF START:',  datetime.strftime(datetime.now(pytz.timezone('Asia/Singapore')), \"%Y-%m-%d %H:%M:%S\"))\n\n\ntry:\n    import Bio\nexcept:\n    #for drfold2 --------\n    #!pip install biopython\n    !pip install /kaggle/input/biopython/biopython-1.85-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n\nprint('PIP INSTALL OK !!!!')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-17T21:27:56.154230Z","iopub.execute_input":"2025-03-17T21:27:56.154562Z","iopub.status.idle":"2025-03-17T21:27:56.160138Z","shell.execute_reply.started":"2025-03-17T21:27:56.154532Z","shell.execute_reply":"2025-03-17T21:27:56.159231Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os,sys\n\nimport pandas as pd\npd.set_option('display.max_columns', 20)\npd.set_option('display.expand_frame_repr', False)\n\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom timeit import default_timer as timer\nimport re\n\nimport matplotlib \nimport matplotlib.pyplot as plt\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\ndef time_to_str(t, mode='min'):\n\tif mode=='min':\n\t\tt  = int(t)/60\n\t\thr = t//60\n\t\tmin = t%60\n\t\treturn '%2d hr %02d min'%(hr,min) \n\telif mode=='sec':\n\t\tt   = int(t)\n\t\tmin = t//60\n\t\tsec = t%60\n\t\treturn '%2d min %02d sec'%(min,sec)\n\n\telse:\n\t\traise NotImplementedError\n\ndef gpu_memory_use():\n    if torch.cuda.is_available():\n        device = torch.device(0)\n        free, total = torch.cuda.mem_get_info(device)\n        used= (total - free) / 1024 ** 3\n        return int(round(used))\n    else:\n        return 0\n\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\nprint('torch',torch.__version__)\nprint('torch.cuda',torch.version.cuda)\n\nprint('IMPORT OK!!!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-17T21:27:56.162737Z","iopub.execute_input":"2025-03-17T21:27:56.162943Z","iopub.status.idle":"2025-03-17T21:27:56.177869Z","shell.execute_reply.started":"2025-03-17T21:27:56.162924Z","shell.execute_reply":"2025-03-17T21:27:56.177216Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MODE = 'submit' #'local' # submit\n\nDATA_KAGGLE_DIR = '/kaggle/input/stanford-rna-3d-folding'\nif MODE == 'local':\n    valid_df = pd.read_csv(f'{DATA_KAGGLE_DIR}/validation_sequences.csv')\n    label_df = pd.read_csv(f'{DATA_KAGGLE_DIR}/validation_labels.csv')\n    label_df['target_id'] = label_df['ID'].apply(lambda x: '_'.join(x.split('_')[:-1]))\n\n    valid_df = valid_df.iloc[[0,1,6,7]].reset_index(drop=True) #for speedup debug\n\nif MODE == 'submit':\n\tvalid_df = pd.read_csv(f'{DATA_KAGGLE_DIR}/test_sequences.csv')\n\nprint('len(valid_df)',len(valid_df))\nprint(valid_df.iloc[0])\nprint('')\n\n\n# cfg = dotdict(\n#     num_conf = 5,\n#     max_length=480,\n# )\nNUM_CONF=5\nMAX_LENGTH=480\nDEVICE='cuda' #'cpu'\n\nprint('MODE:', MODE)\nprint('SETTING OK!!!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-17T21:27:56.188425Z","iopub.execute_input":"2025-03-17T21:27:56.188688Z","iopub.status.idle":"2025-03-17T21:27:56.240811Z","shell.execute_reply.started":"2025-03-17T21:27:56.188666Z","shell.execute_reply":"2025-03-17T21:27:56.240025Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sys.path.append('/kaggle/input/hengck23-drfold2-dummy-00/drfold2/cfg_97')\nfrom EvoMSA2XYZ.Model import MSA2XYZ\nfrom RNALM2.Model import RNA2nd\nfrom data import parse_seq, Get_base, BASE_COOR\nfrom data import write_frame_coor_to_pdb, parse_pdb_to_xyz\n\n\n###########################################################3\nKAGGLE_TRUTH_PDB_DIR ='/kaggle/input/hengck23-drfold2-dummy-00/kaggle-casp15-truth'\nUSALIGN = '/kaggle/working/USalign' \nos.system('cp /kaggle/input/usalign/USalign /kaggle/working/')\nos.system('sudo chmod u+x /kaggle/working/USalign')\n\n# evaluate helper\ndef get_truth_df(target_id, 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_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    return np.array(rotation_matrix)\n\n\n\n# data helper\ndef make_data(seq):\n    aa_type = parse_seq(seq)\n    base = Get_base(seq, BASE_COOR)\n    seq_idx = np.arange(len(seq)) + 1\n\n    msa = aa_type[None, :]\n    msa = torch.from_numpy(msa)\n    msa = torch.cat([msa, msa], 0) #???\n    msa = F.one_hot(msa.long(), 6).float()\n\n    base_x  = torch.from_numpy(base).float()\n    seq_idx = torch.from_numpy(seq_idx).long()\n    return msa, base_x, seq_idx\n    \ndef make_dummy_solution():\n    solution=dotdict()\n    for i, row in valid_df.iterrows():\n        target_id = row.target_id\n        sequence = row.sequence\n        solution[target_id]=dotdict(\n            target_id=target_id,\n            sequence=sequence,\n            coord=[],\n        )\n    return solution\n\ndef solution_to_submit_df(solution):\n    submit_df = []\n    for k,s in solution.items():\n        df = coord_to_df(s.sequence, s.coord, s.target_id)\n        submit_df.append(df)\n    \n    submit_df = pd.concat(submit_df)\n    return submit_df\n \n\ndef coord_to_df(sequence, coord, target_id):\n    L = len(sequence)\n    df = pd.DataFrame()\n    df['ID'] = [f'{target_id}_{i + 1}' for i in range(L)]\n    df['resname'] = [s for s in sequence]\n    df['resid'] = [i + 1 for i in range(L)]\n\n    num_coord = len(coord)\n    for j in range(num_coord):\n        df[f'x_{j+1}'] = coord[j][:, 0]\n        df[f'y_{j+1}'] = coord[j][:, 1]\n        df[f'z_{j+1}'] = coord[j][:, 2]\n    return df\n\n################### start here !!! #######################################################3\n\n\nout_dir = '/kaggle/working/model-output'\nos.makedirs(out_dir, exist_ok=True)\nsolution = make_dummy_solution()\n\n\n#load model (these are moified versions, not the same from their github repo)\nrnalm = RNA2nd(dict(\n    s_in_dim=5,\n    z_in_dim=2,\n    s_dim= 512,\n    z_dim= 128,\n    N_elayers=18,\n))\nrnalm_file = '/kaggle/input/hengck23-drfold2-dummy-00/RCLM/epoch_67000'\nprint(rnalm_file)\nprint(\n    rnalm.load_state_dict(torch.load(rnalm_file, map_location='cpu', weights_only=True), strict=False)\n    #Unexpected key(s) in state_dict: \"ss_head.linear.weight\", \"ss_head.linear.bias\".\n)\nrnalm = rnalm.to(DEVICE)\nrnalm = rnalm.eval()\n\n\ntotal_time_taken = 0\nmax_gpu_mem_used = 0\nfor c in range(NUM_CONF):\n\n    msa2xyz = MSA2XYZ(dict(\n        seq_dim=6,\n        msa_dim=7,\n        N_ensemble=1,#3\n        N_cycle=8,\n        m_dim=64,\n        s_dim=64,\n        z_dim=64,\n    ))\n    msa2xyz_file = [\n        f'/kaggle/input/hengck23-drfold2-dummy-00/cfg_97/model_{k}' for k in [0,1,2,8,9]\n    ][c]\n    print(msa2xyz_file)\n    print(\n        msa2xyz.load_state_dict(torch.load(msa2xyz_file, map_location='cpu', weights_only=True), strict=True)\n    )\n    msa2xyz.msaxyzone.premsa.rnalm = rnalm\n    msa2xyz = msa2xyz.to(DEVICE)\n    msa2xyz = msa2xyz.eval()\n \n    for i,row in valid_df.iterrows():\n        start_timer = timer()\n        \n        target_id = row.target_id\n        sequence = row.sequence\n        seq = row.sequence    \n        \n        L = len(sequence)\n        if L>MAX_LENGTH:\n            i0 = np.random.choice(L-MAX_LENGTH+1)\n            i1 = i0 + MAX_LENGTH\n        else:\n            i0 = 0\n            i1 = L\n        \n        seq = sequence[i0:i1]\n        print(c,i,target_id, L, seq[:75]+'...')\n        \n        msa, base_x, seq_idx = make_data(seq)\n        msa, base_x, seq_idx = msa.to(DEVICE), base_x.to(DEVICE), seq_idx.to(DEVICE)\n        secondary = None #secondary structure\n    \n        with torch.no_grad(): \n            out = msa2xyz.pred(msa, seq_idx, secondary, base_x, np.array(list(seq)))\n\n        # key = list(out.keys()) # plddt(L,L), coor(L,3,3), dist_p(L,L,38), dist_c, dist_n,\n        # for k in key:\n        #     print(k, type(out[k]), out[k].shape)\n \n        \n        if L!=len(seq):\n             out['coor'] = np.pad(out['coor'] ,((i0, L - i1), (0, 0), (0, 0)), 'constant', constant_values=0)\n\n\n        print('out:',  out['coor'].shape)\n        time_taken = timer()-start_timer\n        total_time_taken += time_taken\n        print('time_taken:', time_to_str(time_taken, mode='sec')) \n        \n        gpu_mem_used = gpu_memory_use()\n        max_gpu_mem_used = max(max_gpu_mem_used,gpu_mem_used)\n        print('gpu_mem_used:', gpu_mem_used, 'GB')\n\n        torch.cuda.empty_cache() \n        pdb_file = f'{out_dir}/{target_id}-coor.{c}.pdb'\n        write_frame_coor_to_pdb(out['coor'], sequence, pdb_file) \n        xyz, resname, resid = parse_pdb_to_xyz(pdb_file)\n        #assert(resname==row.sequence)\n        #assert(resid==list(np.arange(L)+1))\n\n        solution[target_id].coord.append(xyz)\n        \n        if MODE == 'local':\n            pass  # save for local cv\n        else:\n            os.remove(pdb_file)\n    print('')\n    \n#-----end of conformation generation ----\nprint('MAX_LENGTH', MAX_LENGTH)\nprint('### total_time_taken:', time_to_str(total_time_taken, mode='min'))\nprint('### max_gpu_mem_used:', max_gpu_mem_used, 'GB')\nprint('')\n\nsubmit_df = solution_to_submit_df(solution)\nsubmit_df.to_csv(f'submission.csv', index=False)\nprint(submit_df)\nprint('SUBMIT OK!!!!!!')\nprint('')\n\n\nif 1: \n    print('debug: show first perdict')\n    solution = list(solution.values())\n    s = solution[0]\n    target_id = s.target_id\n    \n    fig = plt.figure(figsize=(10, 10))\n    ax = fig.add_subplot(111, projection='3d')\n \n\n\n    if MODE=='local':\n        truth_df  = get_truth_df(target_id, label_df)\n        truth = truth_df[['x_1', 'y_1', 'z_1']].to_numpy().astype('float32')\n        truth_pdb = f'{KAGGLE_TRUTH_PDB_DIR}/kaggle_truth_{target_id}_C1.pdb' \n        # print(os.path.isfile(truth_pdb))\n \n        x, y, z = truth[:, 0], truth[:, 1], truth[:, 2]\n        ax.scatter(x, y, z, c='black', s=30, alpha=1)\n        ax.plot(x, y, z, color='black', linewidth=1, alpha=1, label=f'truth')\n        \n    else:\n        truth = s.coord[0] #align to first one\n        truth_pdb = f'{out_dir}/{target_id}-coor.{c}.pdb' \n        # print(os.path.isfile(truth_pdb))\n         \n    aligned = []\n    tm_score = []\n    for c in range(5):\n        predict_pdb = f'{out_dir}/{target_id}-coor.{c}.pdb'\n        # print(os.path.isfile(predict_pdb))\n         \n        command = f'{USALIGN} {predict_pdb} {truth_pdb} -atom \" C1\\'\" -m -'\n        output = os.popen(command).read()\n        tm = parse_usalign_for_tm_score(output)\n        transform = parse_usalign_for_transform(output)\n        aligned.append(s.coord[c]@transform[:,1:].T + transform[:,[0]].T)\n        tm_score.append(tm)\n\n    \n    if MODE!='local':\n        tm_score =['?']*5\n\n\n    max_c = np.array(tm_score).argmax() \n    for c in range(5):\n        x, y, z = aligned[c][:, 0], aligned[c][:, 1], aligned[c][:, 2]\n        alpha =1 if c==max_c else 0.2\n        ax.scatter(x, y, z, c='RED', s=30, alpha=alpha)\n        ax.plot(x, y, z, color='RED', linewidth=1, alpha=alpha, label=f'{c}: tm {tm_score[c]}:0.5f')\n        \n    set_aspect_equal(ax)\n    plt.legend()\n    plt.show() \n    plt.close()\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-17T21:27:56.281531Z","iopub.execute_input":"2025-03-17T21:27:56.281765Z","execution_failed":"2025-03-17T21:34:21.970Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if MODE=='local':\n    # local validation\n \n    tm_score=[]\n    for i,row in valid_df.iterrows(): \n        target_id = row.target_id#'R1116' #casp15 R1116: len(157)\n        seq = row.sequence \n        #-----------------------------------------------\n        print(i,target_id, len(seq), seq[:75]+'...')\n    \n        truth_pdb = f'{KAGGLE_TRUTH_PDB_DIR}/kaggle_truth_{target_id}_C1.pdb'\n        # print(os.path.isfile(truth_pdb))\n        \n        tm = []\n        for c in range(NUM_CONF):\n            predict_pdb = f'{out_dir}/{target_id}-coor.{c}.pdb'\n            # print(os.path.isfile(predict_pdb))\n        \n            command = f'{USALIGN} {predict_pdb} {truth_pdb} -atom \" C1\\'\" -m -'\n            output = os.popen(command).read()\n            # print(output)\n            try:\n                tm_c = parse_usalign_for_tm_score(output)\n            except:\n                tm_c = 0\n            tm.append(tm_c)\n        print('### tm:', tm)\n        tm_score.append(max(tm))\n    \n    print('ALL\\n',tm_score)\n    print('MEAN', np.array(tm_score).mean())\n\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-03-17T21:34:21.971Z"}},"outputs":[],"execution_count":null}]}