{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":87793,"databundleVersionId":11553390,"sourceType":"competition"},{"sourceId":10855324,"sourceType":"datasetVersion","datasetId":6742586},{"sourceId":11173012,"sourceType":"datasetVersion","datasetId":6973060},{"sourceId":224830487,"sourceType":"kernelVersion"},{"sourceId":302884,"sourceType":"modelInstanceVersion","modelInstanceId":258592,"modelId":279808}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"MODEL_TYPE='RF2NA'\nVALIDATION=False","metadata":{"trusted":true,"execution":{"execution_failed":"2025-03-26T16:12:29.905Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Install requirements ","metadata":{}},{"cell_type":"code","source":"if MODEL_TYPE=='RF2NA' and VALIDATION:\n    !pip install torch==2.1\n    !pip install  torchdata==0.7.0 \n    !pip install dgl -f https://data.dgl.ai/wheels/torch-2.1/cu121/repo.html\n    !pip install biopython\n    !pip install e3nn\n\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-03-26T16:12:29.905Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! tar xvfz /kaggle/input/rf2na-weights/RF2NA_apr23.tgz ","metadata":{"trusted":true,"execution":{"execution_failed":"2025-03-26T16:12:29.905Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Helper scripts","metadata":{}},{"cell_type":"code","source":"\nfrom copy import deepcopy\nimport pandas as pd\n\nimport os, sys\nimport re\nimport numpy as np\nimport torch\n\nimport matplotlib\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\n\nimport Bio\nfrom Bio.PDB import Atom, Model, Chain, Residue, Structure, PDBParser\nfrom Bio import SeqIO\n\nprint('IMPORT OK !!!!')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"execution_failed":"2025-03-26T16:12:29.905Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sys.path.append('/kaggle/input/rf2na/pytorch/default/1/network')\nsys.path.append('/kaggle/input/rf2na/pytorch/default/1/SE3Transformer')","metadata":{"trusted":true,"execution":{"execution_failed":"2025-03-26T16:12:29.905Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PYTHON = sys.executable\nprint('PYTHON',PYTHON)\n\n\nUSALIGN = \\\n'/kaggle/working//USalign'\n#'<your us align path>/USalign'\n\nos.system('cp /kaggle/input/usalign/USalign /kaggle/working/')\nos.system('sudo chmod u+x /kaggle/working//USalign')\n\n\nDATA_KAGGLE_DIR = '/kaggle/input/stanford-rna-3d-folding'\n\n\n# helper ----\nclass dotdict(dict):\n\t__setattr__ = dict.__setitem__\n\t__delattr__ = dict.__delitem__\n\n\tdef __getattr__(self, name):\n\t\ttry:\n\t\t\treturn self[name]\n\t\texcept KeyError:\n\t\t\traise AttributeError(name)\n\n# visualisation helper ----\ndef set_aspect_equal(ax):\n\tx_limits = ax.get_xlim()\n\ty_limits = ax.get_ylim()\n\tz_limits = ax.get_zlim()\n\n\t# Compute the mean of each axis\n\tx_middle = np.mean(x_limits)\n\ty_middle = np.mean(y_limits)\n\tz_middle = np.mean(z_limits)\n\n\t# Compute the max range across all axes\n\tmax_range = max(x_limits[1] - x_limits[0],\n\t\t\t\t\ty_limits[1] - y_limits[0],\n\t\t\t\t\tz_limits[1] - z_limits[0]) / 2.0\n\n\t# Set the new limits to ensure equal scaling\n\tax.set_xlim(x_middle - max_range, x_middle + max_range)\n\tax.set_ylim(y_middle - max_range, y_middle + max_range)\n\tax.set_zlim(z_middle - max_range, z_middle + max_range)\n\n\n\n\n# xyz df helper --------------------\ndef get_truth_df(target_id):\n    truth_df = LABEL_DF[LABEL_DF['target_id'] == target_id]\n    truth_df = truth_df.reset_index(drop=True)\n    return truth_df\n\ndef parse_output_to_df(output, seq, target_id):\n    df = []\n    chain_data = []\n    for i, res in enumerate(seq):\n        d=dict(ID = target_id,\n                    resname=res,\n                    resid=i+1)\n        for n in range(len(output)):\n            d={**d, f'x_{n+1}': round(output[n,i,0].item(),3),\n                     f'y_{n+1}': round(output[n,i,1].item(),3),\n                     f'z_{n+1}': round(output[n,i,2].item(),3)}\n        chain_data.append(d)\n\n    if len(chain_data)!=0:\n        chain_df = pd.DataFrame(chain_data)\n        df.append(chain_df)\n        ##print(chain_df)\n    return df\n\ndef parse_pdb_to_df(pdb_file, target_id):\n    parser = PDBParser()\n    structure = parser.get_structure('', pdb_file)\n\n    df = []\n    for model in structure:\n        for chain in model:\n            print(chain)\n            chain_data = []\n            for residue in chain:\n                # print(residue)\n                if residue.get_resname() in ['A', 'U', 'G', 'C']:\n                    # Check if the residue has a C1' atom\n                    if 'C1\\'' in residue:\n                        atom = residue['C1\\'']\n                        xyz = atom.get_coord()\n                        resname = residue.get_resname()\n                        resid = residue.get_id()[1]\n\n                        #todo detect discontinous: resid = prev_resid+1\n                        #ID\tresname\tresid\tx_1\ty_1\tz_1\n                        chain_data.append(dict(\n                            ID = target_id+'_'+str(resid),\n                            resname=resname,\n                            resid=resid,\n                            x_1=xyz[0],\n                            y_1=xyz[1],\n                            z_1=xyz[2],\n                        ))\n                        ##print(f\"Residue {resname} {resid}, Atom: {atom.get_name()}, xyz: {xyz}\")\n\n            if len(chain_data)!=0:\n                chain_df = pd.DataFrame(chain_data)\n                df.append(chain_df)\n                ##print(chain_df)\n    return df\n\n# usalign helper --------------------\ndef write_target_line(\n    atom_name, atom_serial, residue_name, chain_id, residue_num, x_coord, y_coord, z_coord, occupancy=1.0, b_factor=0.0, atom_type='P'\n):\n    \"\"\"\n    Writes a single line of PDB format based on provided atom information.\n\n    Args:\n        atom_name (str): Name of the atom (e.g., \"N\", \"CA\").\n        atom_serial (int): Atom serial number.\n        residue_name (str): Residue name (e.g., \"ALA\").\n        chain_id (str): Chain identifier.\n        residue_num (int): Residue number.\n        x_coord (float): X coordinate.\n        y_coord (float): Y coordinate.\n        z_coord (float): Z coordinate.\n        occupancy (float, optional): Occupancy value (default: 1.0).\n        b_factor (float, optional): B-factor value (default: 0.0).\n\n    Returns:\n        str: A single line of PDB string.\n    \"\"\"\n    return f'ATOM  {atom_serial:>5d}  {atom_name:<5s} {residue_name:<3s} {residue_num:>3d}    {x_coord:>8.3f}{y_coord:>8.3f}{z_coord:>8.3f}{occupancy:>6.2f}{b_factor:>6.2f}           {atom_type}\\n'\n\ndef write_xyz_to_pdb(df, pdb_file, xyz_id = 1):\n    resolved_cnt = 0\n    with open(pdb_file, 'w') as target_file:\n        for _, row in df.iterrows():\n            x_coord = row[f'x_{xyz_id}']\n            y_coord = row[f'y_{xyz_id}']\n            z_coord = row[f'z_{xyz_id}']\n\n            if x_coord > -1e17 and y_coord > -1e17 and z_coord > -1e17:\n                resolved_cnt += 1\n                target_line = write_target_line(\n                    atom_name=\"C1'\",\n                    atom_serial=int(row['resid']),\n                    residue_name=row['resname'],\n                    chain_id='0',\n                    residue_num=int(row['resid']),\n                    x_coord=x_coord,\n                    y_coord=y_coord,\n                    z_coord=z_coord,\n                    atom_type='C',\n                )\n                target_file.write(target_line)\n    return resolved_cnt\n\ndef parse_usalign_for_tm_score(output):\n    # Extract TM-score based on length of reference structure (second)\n    tm_score_match = re.findall(r'TM-score=\\s+([\\d.]+)', output)[1]\n    if not tm_score_match:\n        raise ValueError('No TM score found')\n    return float(tm_score_match)\n\ndef parse_usalign_for_transform(output):\n    # Locate the rotation matrix section\n    matrix_lines = []\n    found_matrix = False\n\n    for line in output.splitlines():\n        if \"The rotation matrix to rotate Structure_1 to Structure_2\" in line:\n            found_matrix = True\n        elif found_matrix and re.match(r'^\\d+\\s+[-\\d.]+\\s+[-\\d.]+\\s+[-\\d.]+\\s+[-\\d.]+$', line):\n            matrix_lines.append(line)\n        elif found_matrix and not line.strip():\n            break  # Stop parsing if an empty line is encountered after the matrix\n\n    # Parse the rotation matrix values\n    rotation_matrix = []\n    for line in matrix_lines:\n        parts = line.split()\n        row_values = list(map(float, parts[1:]))  # Skip the first column (index)\n        rotation_matrix.append(row_values)\n\n    return np.array(rotation_matrix)\n\ndef call_usalign(predict_df, truth_df, verbose=1):\n    truth_pdb = '~truth.pdb'\n    predict_pdb = '~predict.pdb'\n    write_xyz_to_pdb(predict_df, predict_pdb, xyz_id=1)\n    write_xyz_to_pdb(truth_df, truth_pdb, xyz_id=1)\n\n    command = f'{USALIGN} {predict_pdb} {truth_pdb} -atom \" C1\\'\" -m -'\n    output = os.popen(command).read()\n    if verbose==1:\n        print(output)\n    tm_score = parse_usalign_for_tm_score(output)\n    transform = parse_usalign_for_transform(output)\n    return tm_score, transform\n\nprint('HELPER OK!!!')","metadata":{"trusted":true,"execution":{"execution_failed":"2025-03-26T16:12:29.905Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from predict import *    \n\nclass RF2NAPredictor(Predictor):\n\n    def __init__(self, model_weights, device):\n        super().__init__(model_weights, device)\n        self.xyz_converter = self.xyz_converter.to(self.device)\n    \n    def predict_rna_seq(self, seq, n_templ=4):\n        # pass 1, combined MSA\n        Ls, msas, inss, types = [], [], [], []\n        has_paired = False\n                       \n        L=len(seq)\n        msa_i = [seq]\n        ins_i = [np.zeros((L))]\n        alphabet = np.array(list(\"00000000000000000000-000000ACGUN\"), dtype='|S1').view(np.uint8)\n        msa_i = np.array([list(s) for s in msa_i], dtype='|S1').view(np.uint8)\n\n        for i in range(alphabet.shape[0]):\n            msa_i[msa_i == alphabet[i]] = i\n\n        msa_i[msa_i == ord(\"T\")] = 30\n        assert (np.all(msa_i<=31))\n\n        ins_i = np.array(ins_i, dtype=np.uint8)\n\n        Ls.append(L)\n\n        msa_i = torch.tensor(msa_i).long()\n        ins_i = torch.tensor(ins_i).long()\n\n        msas.append(msa_i)\n        inss.append(ins_i)\n        types.append('R')\n\n        msa_orig = {'msa':msas[0],'ins':inss[0]}\n        if (has_paired):\n            if (len(Ls)!=2 or len(msas)!=1):\n                print (\"ERROR: Paired Protein/NA fastas can not be combined with other inputs!\")\n                assert (False)\n        else:\n            for i in range(1,len(Ls)):\n                msa_orig = merge_a3m_hetero(msa_orig, {'msa':msas[i],'ins':inss[i]}, [sum(Ls[:i]),Ls[i]])\n\n        msa_orig, ins_orig = msa_orig['msa'], msa_orig['ins']\n\n        # pass 2, templates\n        L = sum(Ls)\n        xyz_t = INIT_CRDS.reshape(1,1,NTOTAL,3).repeat(n_templ,L,1,1) + torch.rand(n_templ,L,1,3)*5.0 - 2.5\n        is_NA = util.is_nucleic(msa_orig[0])\n        xyz_t[:,is_NA] = INIT_NA_CRDS.reshape(1,1,NTOTAL,3)\n\n        mask_t = torch.full((n_templ, L, NTOTAL), False) \n        t1d = torch.nn.functional.one_hot(torch.full((n_templ, L), 20).long(), num_classes=NAATOKENS-1).float() # all gaps\n        t1d = torch.cat((t1d, torch.zeros((n_templ,L,1)).float()), -1)\n\n        maxtmpl=1\n        \n        # template features\n        xyz_t = xyz_t[:maxtmpl].float().unsqueeze(0).to(self.device)\n        mask_t = mask_t[:maxtmpl].unsqueeze(0).to(self.device)\n        t1d = t1d[:maxtmpl].float().unsqueeze(0).to(self.device)\n\n        same_chain = torch.ones((1,L,L), dtype=torch.bool, device=xyz_t.device)\n\n        mask_t_2d = mask_t[:,:,:,:3].all(dim=-1) # (B, T, L)\n        mask_t_2d = mask_t_2d[:,:,None]*mask_t_2d[:,:,:,None] # (B, T, L, L)\n        mask_t_2d = mask_t_2d.float()*same_chain.float()[:,None] # (ignore inter-chain region)\n        t2d = xyz_to_t2d(xyz_t, mask_t_2d)\n\n        seq_tmp = t1d[...,:-1].argmax(dim=-1).reshape(-1,L)\n        alpha, _, alpha_mask, _ = self.xyz_converter.get_torsions(xyz_t.reshape(-1,L,NTOTAL,3), seq_tmp, mask_in=mask_t.reshape(-1,L,NTOTAL))\n        alpha_mask = torch.logical_and(alpha_mask, ~torch.isnan(alpha[...,0]))\n\n        alpha[torch.isnan(alpha)] = 0.0\n        alpha = alpha.reshape(1,-1,L,NTOTALDOFS,2)\n        alpha_mask = alpha_mask.reshape(1,-1,L,NTOTALDOFS,1)\n        alpha_t = torch.cat((alpha, alpha_mask), dim=-1).reshape(1, -1, L, 3*NTOTALDOFS)\n\n        self.model.eval()\n        pred=self._run_model_single(Ls, msa_orig, ins_orig, t1d, t2d, xyz_t, xyz_t[:,0], alpha_t, same_chain, mask_t_2d)\n        torch.cuda.empty_cache()\n        return pred\n\n    def _run_model_single(self, L_s, msa_orig, ins_orig, t1d, t2d, xyz_t, xyz, alpha_t, same_chain, mask_t_2d):\n        self.xyz_converter = self.xyz_converter.to(self.device)\n        with torch.no_grad():\n            seq, msa_seed_orig, msa_seed, msa_extra, mask_msa = MSAFeaturize(\n                msa_orig, ins_orig, p_mask=0.0, params={'MAXLAT': MAXLAT, 'MAXSEQ': MAXSEQ, 'MAXCYCLE': MAX_CYCLE})\n\n            _, N, L = msa_seed.shape[:3]\n            B = 1   \n            #\n            idx_pdb = torch.arange(L).long().view(1, L)\n            for i in range(len(L_s)-1):\n                idx_pdb[ :, sum(L_s[:(i+1)]): ] += 100\n\n            #\n            seq = seq.unsqueeze(0)\n            msa_seed = msa_seed.unsqueeze(0)\n            msa_extra = msa_extra.unsqueeze(0)\n\n            t1d = t1d.to(self.device)\n            t2d = t2d.to(self.device)\n            idx_pdb = idx_pdb.to(self.device)\n            xyz_t = xyz_t.to(self.device)\n            alpha_t = alpha_t.to(self.device)\n            xyz = xyz.to(self.device)\n            same_chain = same_chain.to(self.device)\n            mask_t_2d = mask_t_2d.to(self.device)\n\n            msa_prev = None\n            pair_prev = None\n            alpha_prev = torch.zeros((1,L,NTOTALDOFS,2), device=self.device)\n            xyz_prev=xyz\n            state_prev = None\n\n            best_lddt = torch.tensor([-1.0], device=self.device)\n            best_xyz = None\n            best_logit = None\n            best_aa = None\n            for i_cycle in range(MAX_CYCLE):\n                msa_seed_i = msa_seed[:,i_cycle].to(self.device)\n                msa_extra_i = msa_extra[:,i_cycle].to(self.device)\n                seq_i = seq[:,i_cycle].to(self.device)\n                with torch.cuda.amp.autocast(True):\n                    logit_s, logit_aa_s, logit_pae, p_bind, init_crds, alpha_prev, _, pred_lddt_binned, msa_prev, pair_prev, state_prev = self.model(\n                        msa_latent=msa_seed_i, \n                        msa_full=msa_extra_i,\n                        seq=seq_i, \n                        seq_unmasked=seq_i, \n                        xyz=xyz_prev, \n                        sctors=alpha_prev,\n                        idx=idx_pdb,\n                        t1d=t1d, \n                        t2d=t2d,\n                        xyz_t=xyz_t[:,:,:,1],\n                        mask_t=mask_t_2d,\n                        alpha_t=alpha_t,\n                        msa_prev=msa_prev,\n                        pair_prev=pair_prev,\n                        state_prev=state_prev,\n                        same_chain=same_chain\n                    )\n\n                    logit_aa_s = logit_aa_s.reshape(B,-1,N,L)[:,:,0].permute(0,2,1)\n\n                xyz_prev = init_crds[-1]\n                alpha_prev = alpha_prev[-1]\n                pred_lddt = lddt_unbin(pred_lddt_binned)\n                pae = pae_unbin(logit_pae)\n\n                _, all_crds = self.xyz_converter.compute_all_atom(seq[:,i_cycle], init_crds[-1], alpha_prev)\n\n                if pred_lddt.mean() < best_lddt.mean():\n                    continue\n\n                best_xyz = all_crds.clone()\n                best_logit = logit_s\n                best_aa = logit_aa_s\n                best_lddt = pred_lddt.clone()\n                best_pae = pae.clone()\n\n            prob_s = list()\n            for logit in logit_s:\n                prob = self.active_fn(logit.float()) # distogram\n                prob = prob.reshape(-1, L, L) #.permute(1,2,0).cpu().numpy()\n                prob_s.append(prob)\n        \n        end = time.time()\n        return(best_xyz[0])\n\n\n\nprint('RF2NA OK')","metadata":{"trusted":true,"execution":{"execution_failed":"2025-03-26T16:12:29.905Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if MODEL_TYPE=='RF2NA':\n    ckpt='/kaggle/working/weights/RF2NA_apr23.pt'\n    ffdb=None\n    prefix='rf2na_outputs'\n\n    if (torch.cuda.is_available()):\n        print (\"Running on GPU\")\n        pred = RF2NAPredictor(ckpt, torch.device(\"cuda:0\"))\n    else:\n        print (\"Running on CPU\")\n        pred = RF2NAPredictor(ckpt, torch.device(\"cpu\"))\n\nprint('Model OK')","metadata":{"trusted":true,"execution":{"execution_failed":"2025-03-26T16:12:29.906Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if VALIDATION:\n    LABEL_DF = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/train_labels.csv')\n    LABEL_DF['target_id'] = LABEL_DF['ID'].apply(lambda x: '_'.join(x.split('_')[:-1]))\n    train_df=pd.read_csv('/kaggle/input/stanford-rna-3d-folding/train_sequences.csv')\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-03-26T16:12:29.906Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"'''\ninputs=[f'R:/kaggle/input/stanford-rna-3d-folding/MSA/{x}.MSA.fasta' for x in train_df.target_id]\npred.predict(inputs=inputs[:2], out_prefix=prefix, ffdb=ffdb)\n'''\n\nif MODEL_TYPE=='RF2NA' and VALIDATION:\n\n    train_df['rf2na_tm_score']=None\n    num_data=len(train_df)\n\n    for i, seq in tqdm(enumerate(train_df.sequence),total=num_data):\n        if len(seq)>300:\n            continue\n        target_id=train_df.target_id[i]\n        truth_df = get_truth_df(target_id)\n        if sum(~np.isnan(truth_df.x_1))<3:\n            continue\n        try:\n            prediction = pred.predict_rna_seq(seq)\n            prediction=prediction[None,:,10,:]\n        except:\n            continue   \n        result = parse_output_to_df(prediction, seq, target_id)[0]\n        try:\n            tm_score, transform = call_usalign(result, truth_df, verbose=0)\n            train_df.loc[i,'rf2na_tm_score']=tm_score\n        except:\n            pass\n    train_df.to_csv('rf2na_tm_scores.csv', index=False)\n    display(train_df.rf2na_tm_score.hist())\n    print(train_df.rf2na_tm_score.mean())","metadata":{"trusted":true,"execution":{"execution_failed":"2025-03-26T16:12:29.906Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if MODEL_TYPE=='RF2NA' and not VALIDATION:\n    test_df=pd.read_csv('/kaggle/input/stanford-rna-3d-folding/test_sequences.csv')\n\n    num_data = len(test_df)\n    for i, seq in tqdm(enumerate(test_df.sequence),total=num_data):\n        try:\n            target_id=test_df.target_id[i]\n            prediction=[]\n            for k in range(5):\n                p=pred.predict_rna_seq(seq)\n                prediction.append(p[:,10,:])\n            prediction=torch.stack(prediction,axis=0)\n            result = parse_output_to_df(prediction, seq, target_id)[0]\n        except:\n            target_id==test_df.target_id[i]\n            print('Failed to predict', target_id)\n            result=pd.DataFrame(columns=['ID', 'resname', 'resid', \n                                         'x_1', 'y_1', 'z_1', \n                                         'x_2', 'y_2', 'z_2',\n                                         'x_3', 'y_3', 'z_3', \n                                         'x_4', 'y_4', 'z_4', \n                                         'x_5', 'y_5', 'z_5'], \n                                         data=[[target_id, x, j+1] + [0.0]*15 for j, x in enumerate(seq)])\n            \n        result['ID']=result.apply(lambda x: x.ID + '_' + str(x.resid), axis=1)\n        result.to_csv('submission.csv', index=False, mode='a', header=(i==0))\n        torch.cuda.empty_cache()\n\n    display(pd.read_csv('submission.csv'))","metadata":{"trusted":true,"execution":{"execution_failed":"2025-03-26T16:12:29.906Z"}},"outputs":[],"execution_count":null}]}