{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"It is an example of inference with model trained in [RNA starter kernel](https://www.kaggle.com/code/iafoss/rna-starter)","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport os, gc\nimport numpy as np\nfrom tqdm.notebook import tqdm\nimport math\nfrom sklearn.model_selection import KFold\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-11-06T03:20:57.474703Z","iopub.execute_input":"2023-11-06T03:20:57.475070Z","iopub.status.idle":"2023-11-06T03:21:01.420875Z","shell.execute_reply.started":"2023-11-06T03:20:57.475036Z","shell.execute_reply":"2023-11-06T03:21:01.419871Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#import torch_xla\n#import torch_xla.core.xla_model as xm","metadata":{"execution":{"iopub.status.busy":"2023-11-06T03:21:01.423010Z","iopub.execute_input":"2023-11-06T03:21:01.423552Z","iopub.status.idle":"2023-11-06T03:21:01.427796Z","shell.execute_reply.started":"2023-11-06T03:21:01.423514Z","shell.execute_reply":"2023-11-06T03:21:01.426821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MODELS = ['/kaggle/input/rna-model/example0_0.pth']\nPATH = '/kaggle/input/stanford-ribonanza-rna-folding-converted/'\nbs = 256\nnum_workers = 2\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'","metadata":{"execution":{"iopub.status.busy":"2023-11-06T03:21:01.428879Z","iopub.execute_input":"2023-11-06T03:21:01.429136Z","iopub.status.idle":"2023-11-06T03:21:01.462249Z","shell.execute_reply.started":"2023-11-06T03:21:01.429113Z","shell.execute_reply":"2023-11-06T03:21:01.461358Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#!pip3 install fastparquet pyarrow","metadata":{"execution":{"iopub.status.busy":"2023-11-06T03:21:01.464279Z","iopub.execute_input":"2023-11-06T03:21:01.464623Z","iopub.status.idle":"2023-11-06T03:21:01.479619Z","shell.execute_reply.started":"2023-11-06T03:21:01.464583Z","shell.execute_reply":"2023-11-06T03:21:01.478876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#device = xm.xla_device()\ndevice","metadata":{"execution":{"iopub.status.busy":"2023-11-06T03:21:01.480682Z","iopub.execute_input":"2023-11-06T03:21:01.480950Z","iopub.status.idle":"2023-11-06T03:21:01.492765Z","shell.execute_reply.started":"2023-11-06T03:21:01.480926Z","shell.execute_reply":"2023-11-06T03:21:01.491868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data","metadata":{}},{"cell_type":"code","source":"def create_extended_sequence(seq, struct):\n    mapping = {\n        ('G', '('): 0,\n        ('G', '.'): 1,\n        ('G', ')'): 2,\n        ('A', '('): 3,\n        ('A', '.'): 4,\n        ('A', ')'): 5,\n        ('C', '('): 6,\n        ('C', '.'): 7,\n        ('C', ')'): 8,\n        ('U', '('): 9,\n        ('U', '.'): 10,\n        ('U', ')'): 11\n    }\n    extended_seq = [mapping[(n, s)] for n, s in zip(seq, struct)]\n\n    return extended_seq","metadata":{"execution":{"iopub.status.busy":"2023-11-06T03:21:01.493961Z","iopub.execute_input":"2023-11-06T03:21:01.494286Z","iopub.status.idle":"2023-11-06T03:21:01.501558Z","shell.execute_reply.started":"2023-11-06T03:21:01.494256Z","shell.execute_reply":"2023-11-06T03:21:01.500640Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RNA_Dataset_Test(Dataset):\n    def __init__(self, df, mask_only=False, **kwargs):\n        self.seq_map = {'A':0,'C':1,'G':2,'U':3}\n        df['L'] = df.sequence.apply(len)\n        self.Lmax = df['L'].max()\n        self.df = df\n        self.mask_only = mask_only\n        \n    def __len__(self):\n        return len(self.df)  \n    \n    def __getitem__(self, idx):\n        id_min, id_max, seq ,structure = self.df.loc[idx, ['id_min','id_max','sequence','complex_structure_features']]\n        mask = torch.zeros(self.Lmax, dtype=torch.bool)\n        L = len(seq)\n        mask[:L] = True\n        if self.mask_only: return {'mask':mask},{}\n        ids = np.arange(id_min,id_max+1)\n        \n        seq = create_extended_sequence(seq,structure)\n        seq = np.array(seq)\n        seq = np.pad(seq,(0,self.Lmax-L))\n        ids = np.pad(ids,(0,self.Lmax-L), constant_values=-1)\n        \n        return {'seq':torch.from_numpy(seq), 'mask':mask}, \\\n               {'ids':ids}\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 tuple(dict_to(x, self.device) for x in batch)","metadata":{"execution":{"iopub.status.busy":"2023-11-06T03:26:06.529581Z","iopub.execute_input":"2023-11-06T03:26:06.529945Z","iopub.status.idle":"2023-11-06T03:26:06.542850Z","shell.execute_reply.started":"2023-11-06T03:26:06.529919Z","shell.execute_reply":"2023-11-06T03:26:06.541861Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Swish(nn.Module):\n    def forward(self, x):\n        return x * torch.sigmoid(x)\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\n\nclass RNA_Model(nn.Module):\n    def __init__(self, dim=192, depth=14, head_size=32, **kwargs):\n        super().__init__()\n        self.emb = nn.Embedding(12,dim)\n        \n        #self.transformer = Encoder(depth,d_model=dim, nhead=dim//head_size, dim_feedforward=4*dim,\n        #        dropout=0.15, activation=Swish())\n        \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.15, activation=Swish(), batch_first=True, norm_first=True), depth)\n        \n        self.ffn = nn.Sequential(\n            nn.Linear(dim, dim * 4),  \n            Swish(), \n            nn.Linear(dim * 4, dim), \n        )\n    \n        self.dropout = nn.Dropout(0.15)\n        self.proj_out = nn.Linear(dim,2)\n    \n    def forward(self, x0):\n        mask = x0['mask']\n        Lmax = mask.sum(-1).max()\n        mask = mask[:,:Lmax]\n        x = x0['seq'][:,:Lmax]\n        \n        pos = torch.arange(Lmax, device=x.device).unsqueeze(0)\n        pos = self.pos_enc(pos)\n        x = self.emb(x)\n        x = x + pos\n        \n        x = self.transformer(x,src_key_padding_mask=~mask)\n        x = self.ffn(x)\n        x = self.dropout(x)\n        x = self.proj_out(x)\n        \n        return x","metadata":{"execution":{"iopub.status.busy":"2023-11-06T03:21:01.516855Z","iopub.execute_input":"2023-11-06T03:21:01.517263Z","iopub.status.idle":"2023-11-06T03:21:01.530584Z","shell.execute_reply.started":"2023-11-06T03:21:01.517232Z","shell.execute_reply":"2023-11-06T03:21:01.529847Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model0 = RNA_Model()\n\ndef count_parameters(model):\n    \"\"\"Count the total number of trainable parameters in a PyTorch model.\"\"\"\n    return sum(p.numel() for p in model.parameters() if p.requires_grad)\n\nprint(f\"Total Num of Parameters: {count_parameters(model0)}\")\n\n","metadata":{"execution":{"iopub.status.busy":"2023-11-06T03:21:01.531840Z","iopub.execute_input":"2023-11-06T03:21:01.532172Z","iopub.status.idle":"2023-11-06T03:21:01.626352Z","shell.execute_reply.started":"2023-11-06T03:21:01.532141Z","shell.execute_reply":"2023-11-06T03:21:01.625340Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"df_test = pd.read_parquet('/kaggle/input/rna-dataset/reduced_test_structure.parquet')","metadata":{"execution":{"iopub.status.busy":"2023-11-06T03:25:15.219587Z","iopub.execute_input":"2023-11-06T03:25:15.220228Z","iopub.status.idle":"2023-11-06T03:25:19.350151Z","shell.execute_reply.started":"2023-11-06T03:25:15.220196Z","shell.execute_reply":"2023-11-06T03:25:19.349091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#df_test = pd.read_parquet(os.path.join(PATH,'test_sequences.parquet'))\nds = RNA_Dataset_Test(df_test)\ndl = DeviceDataLoader(torch.utils.data.DataLoader(ds, batch_size=bs, \n               shuffle=False, drop_last=False, num_workers=num_workers), device)\ndel df_test\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-11-06T03:26:19.330068Z","iopub.execute_input":"2023-11-06T03:26:19.330450Z","iopub.status.idle":"2023-11-06T03:26:30.531427Z","shell.execute_reply.started":"2023-11-06T03:26:19.330419Z","shell.execute_reply":"2023-11-06T03:26:30.530199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#state_dict1 = torch.load('/kaggle/input/rna-model/example0_0.pth')","metadata":{"execution":{"iopub.status.busy":"2023-11-06T03:22:46.324701Z","iopub.execute_input":"2023-11-06T03:22:46.325546Z","iopub.status.idle":"2023-11-06T03:22:46.329543Z","shell.execute_reply.started":"2023-11-06T03:22:46.325515Z","shell.execute_reply":"2023-11-06T03:22:46.328530Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#state_dict2 = torch.load('/kaggle/input/rna-model/example0_0.pth')","metadata":{"execution":{"iopub.status.busy":"2023-11-06T03:22:46.519898Z","iopub.execute_input":"2023-11-06T03:22:46.520674Z","iopub.status.idle":"2023-11-06T03:22:46.524763Z","shell.execute_reply.started":"2023-11-06T03:22:46.520639Z","shell.execute_reply":"2023-11-06T03:22:46.523823Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\"\"\"new_state_dict = {}\nfor (key_dp, value_dp), (key_no_dp, value_no_dp) in zip(state_dict2.items(), state_dict1.items()):\n# Assuming that the keys match between the two state dictionaries\n#assert key_dp == key_no_dp, \"Keys do not match between state dictionaries\"\n\n# Calculate the mean of the weights\nmean_weight = (value_dp + value_no_dp) / 2.0\n\n# Assign the mean weight to the new state dictionary\nnew_state_dict[key_no_dp] = mean_weight\"\"\"","metadata":{"execution":{"iopub.status.busy":"2023-11-06T03:22:46.734750Z","iopub.execute_input":"2023-11-06T03:22:46.735465Z","iopub.status.idle":"2023-11-06T03:22:46.741408Z","shell.execute_reply.started":"2023-11-06T03:22:46.735433Z","shell.execute_reply":"2023-11-06T03:22:46.740483Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"models = []\nfor m in MODELS:\n    model = RNA_Model()   \n    model = model.to(device)\n    #model = nn.DataParallel(model)\n    model.load_state_dict(torch.load(m,map_location=torch.device('cpu')))\n    model.eval()\n    models.append(model)\nprint(\"All Models Loaded!\")","metadata":{"execution":{"iopub.status.busy":"2023-11-06T03:22:46.919100Z","iopub.execute_input":"2023-11-06T03:22:46.919430Z","iopub.status.idle":"2023-11-06T03:22:50.060712Z","shell.execute_reply.started":"2023-11-06T03:22:46.919406Z","shell.execute_reply":"2023-11-06T03:22:50.059689Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ids, preds = [], []\nprint('Prediction Precess Started!!!')\nfor idx, (x, y) in enumerate(dl):\n    with torch.no_grad(), torch.cuda.amp.autocast():\n        p = torch.stack([torch.nan_to_num(model(x)) for model in models], 0).mean(0).clip(0, 1)\n\n    for sample_id, mask, pi in zip(y['ids'].cpu(), x['mask'].cpu(), p.cpu()):\n        ids.append(sample_id[mask])\n        preds.append(pi[mask[:pi.shape[0]]])\n\n    if idx % 500 == 0:\n        print(f'{idx} Samples Predicted')","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-11-06T03:26:34.364460Z","iopub.execute_input":"2023-11-06T03:26:34.364818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ids = torch.concat(ids)\npreds = torch.concat(preds)\n\ndf = pd.DataFrame({'id':ids.numpy(), 'reactivity_DMS_MaP':preds[:,1].numpy(), \n                   'reactivity_2A3_MaP':preds[:,0].numpy()})\ndf.to_csv('submission.csv', index=False, float_format='%.4f') # 6.5GB\ndf.head()","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-11-06T03:12:03.118836Z","iopub.status.idle":"2023-11-06T03:12:03.119244Z","shell.execute_reply.started":"2023-11-06T03:12:03.119049Z","shell.execute_reply":"2023-11-06T03:12:03.119075Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}