{"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":51294,"databundleVersionId":6923401,"sourceType":"competition"},{"sourceId":6822004,"sourceType":"datasetVersion","datasetId":3719560},{"sourceId":6982150,"sourceType":"datasetVersion","datasetId":3869262},{"sourceId":7056821,"sourceType":"datasetVersion","datasetId":3841600}],"dockerImageVersionId":30559,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Description\nWelcome to Stanford Ribonanza RNA Folding challenge. The task of this competition is predicting the chemical reactivity at each position of an RNA molecule. These data are extremely sensitive to the structure that each RNA forms, and an algorithm that could perfectly predict these chemical reactivities would need to have an implicit ‘understanding’ of RNA structure. Such an oracle could be then utilized to predictively model structures of novel RNA molecules. A better understanding of how to manipulate RNA could help usher in an age of programmable medicine, including first cures for pancreatic cancer and Alzheimer’s disease as well as much-needed antibiotics and new biotechnology approaches for climate change. \n\nThis notebook provides a simple baseline that may be used as a starting point for further experiments. Improvment of the baseline may include: \n1. Use of proper loss function to incorporate SN_filter = 0 samples as well as reactivity errors into training\n2. Model improvement and use of additional data, e.g. Ribonanza_bpp_files\n\nFinally, working on this competition keep in mind that train/public LB have different sequence length distribution from private LB, i.e. 115-206 vs. 207-457. Therefore, to avoid a strong shakeup one may need to look into performance vs. sequence length end ensure generalizability.","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport os, gc\nimport numpy as np\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,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-12-07T04:39:35.948686Z","iopub.execute_input":"2023-12-07T04:39:35.948995Z","iopub.status.idle":"2023-12-07T04:39:40.691145Z","shell.execute_reply.started":"2023-12-07T04:39:35.948966Z","shell.execute_reply":"2023-12-07T04:39:40.690381Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Fix fastai bug to enable fp16 training with dictionaries\n\nimport torch\nfrom fastai.vision.all import *\ndef flatten(o):\n    \"Concatenate all collections and items as a generator\"\n    for item in o:\n        if isinstance(o, dict): yield o[item]; continue\n        elif isinstance(item, str): yield item; continue\n        try: yield from flatten(item)\n        except TypeError: yield item\n\nfrom torch.cuda.amp import GradScaler, autocast\n@delegates(GradScaler)\nclass MixedPrecision(Callback):\n    \"Mixed precision training using Pytorch's `autocast` and `GradScaler`\"\n    order = 10\n    def __init__(self, **kwargs): self.kwargs = kwargs\n    def before_fit(self): \n        self.autocast,self.learn.scaler,self.scales = autocast(),GradScaler(**self.kwargs),L()\n    def before_batch(self): self.autocast.__enter__()\n    def after_pred(self):\n        if next(flatten(self.pred)).dtype==torch.float16: self.learn.pred = to_float(self.pred)\n    def after_loss(self): self.autocast.__exit__(None, None, None)\n    def before_backward(self): self.learn.loss_grad = self.scaler.scale(self.loss_grad)\n    def before_step(self):\n        \"Use `self` as a fake optimizer. `self.skipped` will be set to True `after_step` if gradients overflow. \"\n        self.skipped=True\n        self.scaler.step(self)\n        if self.skipped: raise CancelStepException()\n        self.scales.append(self.scaler.get_scale())\n    def after_step(self): self.learn.scaler.update()\n\n    @property \n    def param_groups(self): \n        \"Pretend to be an optimizer for `GradScaler`\"\n        return self.opt.param_groups\n    def step(self, *args, **kwargs): \n        \"Fake optimizer step to detect whether this batch was skipped from `GradScaler`\"\n        self.skipped=False\n    def after_fit(self): self.autocast,self.learn.scaler,self.scales = None,None,None\n        \nimport fastai\nfastai.callback.fp16.MixedPrecision = MixedPrecision\n","metadata":{"_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-12-07T04:39:40.6927Z","iopub.execute_input":"2023-12-07T04:39:40.693103Z","iopub.status.idle":"2023-12-07T04:39:42.644995Z","shell.execute_reply.started":"2023-12-07T04:39:40.693077Z","shell.execute_reply":"2023-12-07T04:39:42.643791Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-12-07T04:39:42.646392Z","iopub.execute_input":"2023-12-07T04:39:42.646757Z","iopub.status.idle":"2023-12-07T04:39:42.653443Z","shell.execute_reply.started":"2023-12-07T04:39:42.646723Z","shell.execute_reply":"2023-12-07T04:39:42.652263Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fname = 'example0'\nPATH = '/kaggle/input/stanford-ribonanza-rna-folding-converted/'\nOUT = './'\nbs = 256\nnum_workers = 2\nSEED = 2023\nnfolds = 4\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'","metadata":{"execution":{"iopub.status.busy":"2023-12-07T04:39:42.655733Z","iopub.execute_input":"2023-12-07T04:39:42.656509Z","iopub.status.idle":"2023-12-07T04:39:42.694191Z","shell.execute_reply.started":"2023-12-07T04:39:42.656472Z","shell.execute_reply":"2023-12-07T04:39:42.69337Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data\n\nThe primary training data is provided in train_data.csv, which contains 821840 RNA sequences and the corresponding reactivity measurements with 2A3_MaP and DMS_MaP methods. The reactivity is reported in columns reactivity_0001 - reactivity_0206 and is set to NaN for the first 26 and the last 21 nucleotides as well as padding for sequences shorter than 206. For faster loading and effective RAM use, I converted the data into a float32 parquet file. \n\nEvaluation in this competition is performed only on samples with SN_filter = 1 for both measurement methods. In this example, I perform training only on samples wiht SN_filter = 1, which gives a noticeable CV boost but uses only 1/4 of the data (i.e. training on noisy SN_filter = 0 data degrades the performance). A proper consideration of all data as well as reactivity errors may boost the performance.\n\nIn this example, I use a simple CV Kfold split. However, given a mismatch in the RNA length between train/public LB vs. private LB data, **it may be important to verify the effect of the sequence length** to avoid a significant shakeup at the private LB.\n\nOne of the tricks, well known in NLP community, which I use here, is length matching batch sampling: composing batches of samples of approximately the same length to minimize the overhead caused by padding tokens.","metadata":{}},{"cell_type":"code","source":"class RNA_Dataset(Dataset):\n    def __init__(self, df, mode='train', seed=2023, fold=0, nfolds=4, \n                 mask_only=False, **kwargs):\n#         self.seq_map = {'A':0,'C':1,'G':2,'U':3,'(':4,')':5,'.':6}\n        self.seq_map = {\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        self.Lmax = 206\n        df['L'] = df.sequence.apply(len)\n        df_2A3 = df.loc[df.experiment_type=='2A3_MaP']\n        df_DMS = df.loc[df.experiment_type=='DMS_MaP']\n        \n        split = list(KFold(n_splits=nfolds, random_state=seed, \n                shuffle=True).split(df_2A3))[fold][0 if mode=='train' else 1]\n        df_2A3 = df_2A3.iloc[split].reset_index(drop=True)\n        df_DMS = df_DMS.iloc[split].reset_index(drop=True)\n        \n        m = (df_2A3['SN_filter'].values > 0) & (df_DMS['SN_filter'].values > 0)\n        #m = (df_2A3['signal_to_noise'].values > 0.5) & (df_DMS['signal_to_noise'].values > 0.5)\n        df_2A3 = df_2A3.loc[m].reset_index(drop=True)\n        df_DMS = df_DMS.loc[m].reset_index(drop=True)\n        \n        self.seq = df_2A3['sequence'].values\n        self.struc = df_2A3['structure'].values\n        self.L = df_2A3['L'].values\n        \n        self.react_2A3 = df_2A3[[c for c in df_2A3.columns if \\\n                                 'reactivity_0' in c]].values\n        self.react_DMS = df_DMS[[c for c in df_DMS.columns if \\\n                                 'reactivity_0' in c]].values\n        \n        self.react_err_2A3 = df_2A3[[c for c in df_2A3.columns if \\\n                                 'reactivity_error_0' in c]].values\n        self.react_err_DMS = df_DMS[[c for c in df_DMS.columns if \\\n                                'reactivity_error_0' in c]].values\n        \n        self.sn_2A3 = df_2A3['signal_to_noise'].values\n        self.sn_DMS = df_DMS['signal_to_noise'].values\n        self.mask_only = mask_only\n        \n    def __len__(self):\n        return len(self.seq)  \n    \n    def __getitem__(self, idx):\n        seq = self.seq[idx]\n        struc = self.struc[idx]\n        #struc = self.struc[idx]\n        if self.mask_only:\n            mask = torch.zeros(self.Lmax, dtype=torch.bool)\n            mask[:len(seq)] = True\n            return {'mask':mask},{'mask':mask}\n#         seq = [self.seq_map[s] for s in seq]\n        seq = [self.seq_map[s] for s in zip(seq,struc)]\n        seq = np.array(seq)\n        mask = torch.zeros(self.Lmax, dtype=torch.bool)\n        mask[:len(seq)] = True\n        seq = np.pad(seq,(0,self.Lmax-len(seq)))\n        \n        react = torch.from_numpy(np.stack([self.react_2A3[idx],\n                                           self.react_DMS[idx]],-1))\n        react_err = torch.from_numpy(np.stack([self.react_err_2A3[idx],\n                                               self.react_err_DMS[idx]],-1))\n        sn = torch.FloatTensor([self.sn_2A3[idx],self.sn_DMS[idx]])\n        \n        return {'seq':torch.from_numpy(seq),\n                #'struc':torch.from_numpy(struc),\n                'mask':mask}, \\\n               {'react':react, 'react_err':react_err,\n                'sn':sn, 'mask':mask}\n\n    \nclass LenMatchBatchSampler(torch.utils.data.BatchSampler):\n    def __iter__(self):\n        buckets = [[]] * 100\n        yielded = 0\n\n        for idx in self.sampler:\n            s = self.sampler.data_source[idx]\n            if isinstance(s,tuple): L = s[0][\"mask\"].sum()\n            else: L = s[\"mask\"].sum()\n            L = max(1,L // 16) \n            if len(buckets[L]) == 0:  buckets[L] = []\n            buckets[L].append(idx)\n            \n            if len(buckets[L]) == self.batch_size:\n                batch = list(buckets[L])\n                yield batch\n                yielded += 1\n                buckets[L] = []\n                \n        batch = []\n        leftover = [idx for bucket in buckets for idx in bucket]\n\n        for idx in leftover:\n            batch.append(idx)\n            if len(batch) == self.batch_size:\n                yielded += 1\n                yield batch\n                batch = []\n\n        if len(batch) > 0 and not self.drop_last:\n            yielded += 1\n            yield batch\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-12-07T04:39:42.695749Z","iopub.execute_input":"2023-12-07T04:39:42.696376Z","iopub.status.idle":"2023-12-07T04:39:42.724362Z","shell.execute_reply.started":"2023-12-07T04:39:42.696344Z","shell.execute_reply":"2023-12-07T04:39:42.723557Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model\nTransformers are ideal for the considered task because they naturally capture long dependences in RNN, which define the secondary structure of the molecule and the corresponding chemical reactivity. For illustration purposes, below I provide a simple S-size transformer model.","metadata":{}},{"cell_type":"code","source":"class 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(12,dim)\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), depth)\n        self.proj_out = nn.Linear(dim,2)\n        self.proj_out1 = nn.Linear(dim,dim//2)\n        self.proj_out2 = nn.Linear(dim//2,dim//4)\n        self.proj_out3 = nn.Linear(dim//4,2)\n        self.relu = nn.Tanh()\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        #x_struc = x0['struc'][:,: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.proj_out(x)\n        \n        return x","metadata":{"execution":{"iopub.status.busy":"2023-12-07T04:39:42.72548Z","iopub.execute_input":"2023-12-07T04:39:42.725754Z","iopub.status.idle":"2023-12-07T04:39:42.7399Z","shell.execute_reply.started":"2023-12-07T04:39:42.725731Z","shell.execute_reply":"2023-12-07T04:39:42.739052Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Loss & Metric\nThe metric accumulates all predictions and then performs average to be consistent with the competition metric. However, the difference with a simple batch-based average is negligible.","metadata":{}},{"cell_type":"code","source":"def loss(pred,target):\n    #p = pred[target['mask'][:,:pred.shape[1]]]\n    #p = p.clip(0,1)\n    # y = target['react'][target['mask']].clip(0,1)\n    #y = target['react'][target['mask']].clip(0,1)\n    y = target['react'][:,0:pred.shape[1],:].clip(0,1)\n    p = pred\n    loss = F.l1_loss(p, y, reduction='none')\n    sn = torch.stack(tuple([target['sn'].float() for i in range(pred.shape[1])]), 1)\n    #loss = loss * sn\n    #loss = loss * target['w']\n    loss = loss[~torch.isnan(loss)]\n    loss = loss.mean()\n    \n    return loss\n\nclass MAE(Metric):\n    def __init__(self): \n        self.reset()\n        \n    def reset(self): \n        self.x,self.y = [],[]\n        \n    def accumulate(self, learn):\n        x = learn.pred[learn.y['mask'][:,:learn.pred.shape[1]]]\n        y = learn.y['react'][learn.y['mask']].clip(0,1)\n        self.x.append(x)\n        self.y.append(y)\n\n    @property\n    def value(self):\n        x,y = torch.cat(self.x,0),torch.cat(self.y,0)\n        loss = F.l1_loss(x, y, reduction='none')\n        loss = loss[~torch.isnan(loss)].mean()\n        return loss","metadata":{"execution":{"iopub.status.busy":"2023-12-07T04:39:42.740903Z","iopub.execute_input":"2023-12-07T04:39:42.741156Z","iopub.status.idle":"2023-12-07T04:39:42.75418Z","shell.execute_reply.started":"2023-12-07T04:39:42.741134Z","shell.execute_reply":"2023-12-07T04:39:42.753374Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"##### seed_everything(SEED)\nos.makedirs(OUT, exist_ok=True)\n# df = pd.read_parquet(os.path.join(PATH,'train_data.parquet'))\ndf1 = pd.read_parquet('/kaggle/input/sr-data/train_data.parquet')\ndf2 = pd.read_parquet('/kaggle/input/sr-data/train_whole_sequences.parquet')\ndf = df1.merge(df2,on=['sequence_id','sequence'],how='inner')\ndel df1\ndel df2\nfor fold in range(nfolds): # running multiple folds at kaggle may cause OOM\n    ds_train = RNA_Dataset(df, mode='train', fold=fold, nfolds=nfolds)\n    ds_train_len = RNA_Dataset(df, mode='train', fold=fold, \n                nfolds=nfolds, mask_only=True)\n    sampler_train = torch.utils.data.RandomSampler(ds_train_len)\n    len_sampler_train = LenMatchBatchSampler(sampler_train, batch_size=bs,\n                drop_last=True)\n    dl_train = DeviceDataLoader(torch.utils.data.DataLoader(ds_train, \n                batch_sampler=len_sampler_train, num_workers=num_workers,\n                persistent_workers=True), device)\n\n    ds_val = RNA_Dataset(df, mode='eval', fold=fold, nfolds=nfolds)\n    ds_val_len = RNA_Dataset(df, mode='eval', fold=fold, nfolds=nfolds, \n               mask_only=True)\n    sampler_val = torch.utils.data.SequentialSampler(ds_val_len)\n    len_sampler_val = LenMatchBatchSampler(sampler_val, batch_size=bs, \n               drop_last=False)\n    dl_val= DeviceDataLoader(torch.utils.data.DataLoader(ds_val, \n               batch_sampler=len_sampler_val, num_workers=num_workers), device)\n    gc.collect()\n    \n    data = DataLoaders(dl_train,dl_val)\n    model = RNA_Model(dim=192,depth=12)   \n    model = model.to(device)\n\n    learn = Learner(data, model, loss_func=loss,cbs=[GradientClip(3.0)],\n                metrics=[MAE()]).to_fp16() \n    #fp16 doesn't help at P100 but gives x1.6-1.8 speedup at modern hardware\n\n    learn.fit_one_cycle(100, lr_max=1e-3, wd=0.18,pct_start=0.01,div=100,\n                        div_final=10000000000000)\n    torch.save(learn.model.state_dict(),os.path.join(OUT,f'{fname}_{fold}.pth'))\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-12-07T04:39:42.75519Z","iopub.execute_input":"2023-12-07T04:39:42.75547Z"},"trusted":true},"outputs":[],"execution_count":null}]}