{"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":51294,"databundleVersionId":6923401,"sourceType":"competition"},{"sourceId":7098396,"sourceType":"datasetVersion","datasetId":4091471},{"sourceId":7107895,"sourceType":"datasetVersion","datasetId":4098015},{"sourceId":7115156,"sourceType":"datasetVersion","datasetId":4103253},{"sourceId":7123551,"sourceType":"datasetVersion","datasetId":4109118},{"sourceId":7132471,"sourceType":"datasetVersion","datasetId":4115255},{"sourceId":7142015,"sourceType":"datasetVersion","datasetId":4122256}],"dockerImageVersionId":30587,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport os, gc\nimport numpy as np\nfrom sklearn.model_selection import KFold\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport csv\nfrom tqdm import tqdm\nfrom fastai.vision.all import *","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-12-08T14:07:55.178191Z","iopub.execute_input":"2023-12-08T14:07:55.178820Z","iopub.status.idle":"2023-12-08T14:08:02.783738Z","shell.execute_reply.started":"2023-12-08T14:07:55.178782Z","shell.execute_reply":"2023-12-08T14:08:02.782980Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'\n\nclass SinusoidalPosEmb(nn.Module):\n    def __init__(self, dim=16, M=10000):\n        super().__init__()\n        self.dim = dim\n        self.M = M\n\n    def forward(self, x):\n        device = x.device\n        half_dim = self.dim // 2\n        emb = math.log(self.M) / half_dim\n        emb = torch.exp(torch.arange(half_dim, device=device) * (-emb))\n        emb = x[...,None] * emb[None,...]\n        emb = torch.cat((emb.sin(), emb.cos()), dim=-1)\n        return emb\n\nclass RNA_Model(nn.Module):\n    def __init__(self, dim=192, depth=12, head_size=32, **kwargs):\n        super().__init__()\n        self.emb = nn.Embedding(4,dim)\n#       Convolution to the bpp + Distancesx3 + Diagonal channels of scores for the embedding transformation\n#       A single vonvolution, I want this step simple for a fast convergence\n        self.conv = nn.Conv2d(5,1,1)\n        self.pos_enc = SinusoidalPosEmb(dim)\n        self.transformer = nn.TransformerEncoder(\n                nn.TransformerEncoderLayer(d_model=dim, nhead=dim//head_size, dim_feedforward=4*dim,\n                dropout=0.1, activation=nn.GELU(), batch_first=True, norm_first=True, device=device), depth)\n        self.proj_out = nn.Linear(dim,2)\n    \n    def forward(self, x0):\n        mask = x0['mask']\n        L = mask.sum(-1).max()\n        mask = mask[:,:L]\n        x = x0['seq'][:,:L]\n#       Reinforced scores as described in:\n#       RNAdegformer: accurate prediction of mRNA degradation at nucleotide resolution with deep learning\n#       Shujun He, Baizhen Gao, Rushant Sabnis and Qing Sun\n#       Corresponding author. Qing Sun, Department of Chemical Engineering, Texas A&M University, 100 Spence St., 77843 TX, USA. Tel.: 979-845-3401;\n#       E-mail: sunqing@tamu.edu\n#       BPP\n        bpp = x0['bpp'][:,:L,:L]\n        scr = torch.zeros(list(bpp.shape)+[5],dtype=torch.float32,device=device)\n        scr[:,:,:,0] = bpp\n#       Distance matrix\n#         0   1 1/2 1/3 ...\n#         1   0   1 1/2 ...\n#       1/2   1   0   1 ...\n#       1/3 1/2   1   0 ...\n#       ... ... ... ... ...\n        distances = torch.arange(1,L)\n        distances[1:] = 1/distances[1:]\n        distance_mat = torch.zeros((L,L,3),dtype=torch.float32)\n        for i in range(L):\n            dist = distances[:L-i-1]\n            distance_mat[i,i+1:L,0] = dist\n            distance_mat[i+1:L,i,0] = dist\n\n        distance_mat[:L,:L,1] = distance_mat[:L,:L,0]*distance_mat[:L,:L,0]# Squares\n        distance_mat[:L,:L,2] = distance_mat[:L,:L,1]*distance_mat[:L,:L,0]# Cubes\n\n        scr[:,:,:,1:4] = distance_mat\n#       Diagonal\n        scr[:,:,:,-1] = torch.diag(torch.ones(L,dtype=torch.float32))\n\n        scr = torch.swapaxes(scr,1,-1)\n        bpp = self.conv(scr).squeeze(1)\n#       bpp = nn.Relu()(bpp)\n#       bpp = nn.Softmax()(bpp)\n        \n        pos = torch.arange(L, device=x.device).unsqueeze(0)\n        pos = self.pos_enc(pos)\n        x = self.emb(x)\n        x = torch.matmul(bpp,x)\n        x = x + pos\n        \n        x = self.transformer(x,src_key_padding_mask=~mask)\n        x = self.proj_out(x)\n        \n        return x\n    \nmodel = RNA_Model(dim=192, depth=24, head_size=32) \nmodel.load_state_dict(torch.load('/kaggle/input/fold-1-train/model_31.pth')['model'])\nmodel.to(device)\nmodel.eval()","metadata":{"execution":{"iopub.status.busy":"2023-12-08T14:08:52.405415Z","iopub.execute_input":"2023-12-08T14:08:52.406274Z","iopub.status.idle":"2023-12-08T14:08:57.701889Z","shell.execute_reply.started":"2023-12-08T14:08:52.406237Z","shell.execute_reply":"2023-12-08T14:08:57.700860Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RNA_test(Dataset):\n    def __init__(self, df, **kwargs):\n        self.seq_map = {'A':0,'C':1,'G':2,'U':3}\n        self.seq = df['sequence'].values\n        self.L = df['LEN'].values\n        self.filepaths = df['path'].values\n        self.sequence_id = df['sequence_id'].values\n        self.bpp_root_dir = bpp_root_dir\n        self.id_min = df.id_min.values\n        self.id_max = df.id_max.values\n        self.seq_id = df.sequence_id.values\n        \n    def __len__(self):\n        return len(self.seq)  \n    \n    def __getitem__(self, idx):\n        L = self.L[idx]\n        seq = self.seq[idx]\n        seq = [self.seq_map[s] for s in seq]\n        seq = np.array(seq)\n        mask = torch.ones(L,dtype=torch.bool)\n\n        filepath = self.bpp_root_dir + '/' + self.filepaths[idx] + '/' + self.sequence_id[idx] + '.txt'\n        df_bpp = pd.read_csv(filepath,header=None,sep=' ')\n        df_bpp = pd.concat((df_bpp,pd.DataFrame({0:df_bpp[1],1:df_bpp[0],2:df_bpp[2]})))# Adding transpose indices.\n        df_bpp.loc[:,[0,1]] -= 1# Correcting indices.\n        indices = (df_bpp[[0,1]].values).swapaxes(0,1)\n        values = df_bpp[2].values\n        bpp = torch.sparse_coo_tensor(indices, values, [L, L],dtype=torch.float32).to_dense()\n        \n        return {'seq':torch.from_numpy(seq), 'bpp':bpp, 'mask':mask}\n    \ndef dict_to(x, device='cuda'):\n    return {k:x[k].to(device) for k in x}\n\ndef to_device(x, device='cuda'):\n    return tuple(dict_to(e,device) for e in x)\n\nclass DeviceDataLoader:\n    def __init__(self, dataloader, device='cuda'):\n        self.dataloader = dataloader\n        self.device = device\n    \n    def __len__(self):\n        return len(self.dataloader)\n    \n    def __iter__(self):\n        for batch in self.dataloader:\n            yield dict_to(batch, self.device)","metadata":{"execution":{"iopub.status.busy":"2023-12-03T20:39:47.097999Z","iopub.execute_input":"2023-12-03T20:39:47.098306Z","iopub.status.idle":"2023-12-03T20:39:47.109769Z","shell.execute_reply.started":"2023-12-03T20:39:47.098281Z","shell.execute_reply":"2023-12-03T20:39:47.108390Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_csv('/kaggle/input/stanford-ribonanza-rna-folding/test_sequences.csv')\ntest_df['LEN'] = test_df.sequence.str.len()\ntest_df['PREDS'] = 0\ntest_df['COUNT'] = 0\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-12-03T21:05:33.670161Z","iopub.execute_input":"2023-12-03T21:05:33.670585Z","iopub.status.idle":"2023-12-03T21:05:37.487600Z","shell.execute_reply.started":"2023-12-03T21:05:33.670550Z","shell.execute_reply":"2023-12-03T21:05:37.486929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"bpp_df = pd.read_csv('/kaggle/input/bpp-files-corrected/bpp.csv').set_index('sequence_id',drop=True).drop(columns='Unnamed: 0')\nbpp_df['sequence_id'] = bpp_df.index\nbpp_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-12-03T20:40:42.553189Z","iopub.execute_input":"2023-12-03T20:40:42.554276Z","iopub.status.idle":"2023-12-03T20:40:44.688567Z","shell.execute_reply.started":"2023-12-03T20:40:42.554240Z","shell.execute_reply":"2023-12-03T20:40:44.687518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"bpp_root_dir = '/kaggle/input/stanford-ribonanza-rna-folding/Ribonanza_bpp_files'","metadata":{"execution":{"iopub.status.busy":"2023-12-03T20:41:07.619342Z","iopub.execute_input":"2023-12-03T20:41:07.620762Z","iopub.status.idle":"2023-12-03T20:41:07.625577Z","shell.execute_reply.started":"2023-12-03T20:41:07.620707Z","shell.execute_reply":"2023-12-03T20:41:07.624509Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE = 256\nACCUMULATED_IDX = 0\n\nwith open(\"output.csv\", \"a\",newline='') as csv_file:\n    writer = csv.writer(csv_file, delimiter=',')\n    writer.writerow(['id','reactivity_DMS_MaP','reactivity_2A3_MaP'])\n    with torch.no_grad():\n        for BATCH_LEN in test_df['LEN'].unique():\n            print('Sequence len: ',BATCH_LEN)\n            df = test_df[test_df['LEN']==BATCH_LEN].reset_index(drop=True).reset_index(drop=True)\n            STEPS = len(df)//BATCH_SIZE + 1\n            START = 0\n            END = 0\n            for STEP in tqdm(range(STEPS)):\n                END += BATCH_SIZE\n                BATCH_test_df = df[START:END]\n                PREDS = pd.Series(BATCH_test_df.PREDS.values,index=BATCH_test_df.sequence_id).to_dict()\n                COUNT = pd.Series(BATCH_test_df.COUNT.values,index=BATCH_test_df.sequence_id).to_dict()\n                BATCH_bpp_df = bpp_df.loc[list(BATCH_test_df['sequence_id'])].reset_index(drop=True)\n                BATCH_df = pd.merge(BATCH_test_df,BATCH_bpp_df,on='sequence_id').reset_index(drop=True)\n                ds = RNA_test(BATCH_df)\n                dl = DeviceDataLoader(torch.utils.data.DataLoader(ds,batch_size=len(BATCH_df)),device)\n                for BATCH in dl:\n                    OUTPUT_BATCH = model(BATCH)\n                    for i in range(len(BATCH_df)):\n                        PREDS[BATCH_df['sequence_id'][i]] += OUTPUT_BATCH[i]\n                        COUNT[BATCH_df['sequence_id'][i]] += 1\n                   \n                    for SEQUENCE in BATCH_test_df['sequence_id']:\n                        if COUNT[SEQUENCE] > 1:\n                            SEQUENCE = PREDS[SEQUENCE]/COUNT[SEQUENCE]\n                        else:\n                            SEQUENCE = PREDS[SEQUENCE]\n                            \n                        writer.writerows([[ACCUMULATED_IDX + index, i[1].item(), i[0].item()] for index, i in enumerate(SEQUENCE)])\n                        ACCUMULATED_IDX += BATCH_LEN\n                \n                START = END\n               \n                del BATCH_test_df, PREDS, COUNT, BATCH_bpp_df, BATCH_df, ds, dl\n            del df\n            print('Final prediction: ',ACCUMULATED_IDX)","metadata":{"execution":{"iopub.status.busy":"2023-12-03T21:12:08.519595Z","iopub.execute_input":"2023-12-03T21:12:08.520014Z","iopub.status.idle":"2023-12-03T21:12:10.341129Z","shell.execute_reply.started":"2023-12-03T21:12:08.519983Z","shell.execute_reply":"2023-12-03T21:12:10.339902Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Conv weights for the score channels\nimport numpy as np\nweights = np.array([[0.5241,-0.0928,-0.0836,-0.0736,1.0798],\n                    [0.5324,-0.0924,-0.0832,-0.0732,1.0797],\n                    [0.5317,-0.0891,-0.0799,-0.0699,1.0901],\n                    [0.5283,-0.0943,-0.0851,-0.0751,1.0932]])","metadata":{"execution":{"iopub.status.busy":"2023-12-08T14:31:29.854547Z","iopub.execute_input":"2023-12-08T14:31:29.854961Z","iopub.status.idle":"2023-12-08T14:31:29.862292Z","shell.execute_reply.started":"2023-12-08T14:31:29.854927Z","shell.execute_reply":"2023-12-08T14:31:29.860409Z"},"trusted":true},"execution_count":null,"outputs":[]}]}