{"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":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n\nos.listdir('/kaggle/input')\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-11-12T17:11:11.210194Z","iopub.execute_input":"2023-11-12T17:11:11.210586Z","iopub.status.idle":"2023-11-12T17:11:11.218391Z","shell.execute_reply.started":"2023-11-12T17:11:11.21056Z","shell.execute_reply":"2023-11-12T17:11:11.217388Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport os, gc\nimport numpy as np\nfrom sklearn.model_selection import KFold\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader","metadata":{"execution":{"iopub.status.busy":"2023-11-12T16:55:47.969548Z","iopub.execute_input":"2023-11-12T16:55:47.969881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/stanford-ribonanza-rna-folding/train_data_QUICK_START.csv')\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2023-11-12T16:58:44.715156Z","iopub.execute_input":"2023-11-12T16:58:44.715508Z","iopub.status.idle":"2023-11-12T16:58:52.932201Z","shell.execute_reply.started":"2023-11-12T16:58:44.715482Z","shell.execute_reply":"2023-11-12T16:58:52.931006Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.describe().style.background_gradient(cmap='summer')","metadata":{"execution":{"iopub.status.busy":"2023-11-12T16:59:54.141678Z","iopub.execute_input":"2023-11-12T16:59:54.142036Z","iopub.status.idle":"2023-11-12T16:59:58.406856Z","shell.execute_reply.started":"2023-11-12T16:59:54.142008Z","shell.execute_reply":"2023-11-12T16:59:58.405862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"base_counts = df['sequence'].apply(lambda x: pd.Series(list(x)).value_counts()).sum()\nbase_counts.plot(kind='bar')\nplt.xlabel('Base', fontsize = 12, fontweight = 'bold', color = 'darkblue')\nplt.ylabel('Count', fontsize = 12, fontweight = 'bold', color = 'darkblue')\nplt.title('Base Composition', fontsize = 14, fontweight = 'bold', color = 'darkgreen')\n\n# Save the plot\nplt.savefig('Base Composition.png')\n\n# Show the plot\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-11-12T17:01:13.274344Z","iopub.execute_input":"2023-11-12T17:01:13.274704Z","iopub.status.idle":"2023-11-12T17:03:46.898128Z","shell.execute_reply.started":"2023-11-12T17:01:13.274671Z","shell.execute_reply":"2023-11-12T17:03:46.89718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sequence_lengths = df['sequence'].apply(len)\nsns.histplot(sequence_lengths, kde=True)\nplt.xlabel('Sequence Length', fontsize = 12, fontweight = 'bold', color = 'darkblue')\nplt.ylabel('Frequency', fontsize = 12, fontweight = 'bold', color = 'darkblue')\nplt.title('Sequence Length Distribution', fontsize = 14, fontweight = 'bold', color = 'darkgreen')\n\n# Save the plot\nplt.savefig('Sequence Length Distribution.png')\n\n# Show the plot\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-11-12T17:03:46.899871Z","iopub.execute_input":"2023-11-12T17:03:46.900271Z","iopub.status.idle":"2023-11-12T17:03:48.764363Z","shell.execute_reply.started":"2023-11-12T17:03:46.900249Z","shell.execute_reply":"2023-11-12T17:03:48.763501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"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","metadata":{"execution":{"iopub.status.busy":"2023-11-12T16:56:41.803848Z","iopub.execute_input":"2023-11-12T16:56:41.804194Z","iopub.status.idle":"2023-11-12T16:56:44.812665Z","shell.execute_reply.started":"2023-11-12T16:56:41.804164Z","shell.execute_reply":"2023-11-12T16:56:44.811534Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"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-11-12T16:56:51.479781Z","iopub.execute_input":"2023-11-12T16:56:51.480138Z","iopub.status.idle":"2023-11-12T16:56:51.4859Z","shell.execute_reply.started":"2023-11-12T16:56:51.480106Z","shell.execute_reply":"2023-11-12T16:56:51.484488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"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}\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        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.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        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        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        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 = 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), 'mask':mask}, \\\n               {'react':react, 'react_err':react_err,\n                'sn':sn, 'mask':mask}\n    ","metadata":{"execution":{"iopub.status.busy":"2023-11-12T17:27:04.559008Z","iopub.execute_input":"2023-11-12T17:27:04.559383Z","iopub.status.idle":"2023-11-12T17:27:04.576085Z","shell.execute_reply.started":"2023-11-12T17:27:04.559351Z","shell.execute_reply":"2023-11-12T17:27:04.57481Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/stanford-ribonanza-rna-folding/train_data.csv')\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2023-11-12T17:30:08.887904Z","iopub.execute_input":"2023-11-12T17:30:08.888233Z","iopub.status.idle":"2023-11-12T17:31:44.832312Z","shell.execute_reply.started":"2023-11-12T17:30:08.888208Z","shell.execute_reply":"2023-11-12T17:31:44.83104Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rna_ds = RNA_Dataset(df, mode='train',)","metadata":{"execution":{"iopub.status.busy":"2023-11-12T17:35:05.277077Z","iopub.execute_input":"2023-11-12T17:35:05.277531Z","iopub.status.idle":"2023-11-12T17:35:10.182445Z","shell.execute_reply.started":"2023-11-12T17:35:05.277499Z","shell.execute_reply":"2023-11-12T17:35:10.181189Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rna_ds.seq, rna_ds.seq.shape","metadata":{"execution":{"iopub.status.busy":"2023-11-12T18:48:55.660711Z","iopub.execute_input":"2023-11-12T18:48:55.661082Z","iopub.status.idle":"2023-11-12T18:48:55.671593Z","shell.execute_reply.started":"2023-11-12T18:48:55.661055Z","shell.execute_reply":"2023-11-12T18:48:55.669413Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"item = rna_ds[0]","metadata":{"execution":{"iopub.status.busy":"2023-11-12T18:50:46.542383Z","iopub.execute_input":"2023-11-12T18:50:46.542754Z","iopub.status.idle":"2023-11-12T18:50:46.552129Z","shell.execute_reply.started":"2023-11-12T18:50:46.542729Z","shell.execute_reply":"2023-11-12T18:50:46.550523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"item[0]","metadata":{"execution":{"iopub.status.busy":"2023-11-12T17:37:07.453816Z","iopub.execute_input":"2023-11-12T17:37:07.454174Z","iopub.status.idle":"2023-11-12T17:37:07.463513Z","shell.execute_reply.started":"2023-11-12T17:37:07.454147Z","shell.execute_reply":"2023-11-12T17:37:07.462391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"item[1].keys()","metadata":{"execution":{"iopub.status.busy":"2023-11-12T17:37:56.075345Z","iopub.execute_input":"2023-11-12T17:37:56.075733Z","iopub.status.idle":"2023-11-12T17:37:56.089993Z","shell.execute_reply.started":"2023-11-12T17:37:56.075706Z","shell.execute_reply":"2023-11-12T17:37:56.087478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"item[1]['react'].shape","metadata":{"execution":{"iopub.status.busy":"2023-11-12T18:51:46.502204Z","iopub.execute_input":"2023-11-12T18:51:46.502612Z","iopub.status.idle":"2023-11-12T18:51:46.512764Z","shell.execute_reply.started":"2023-11-12T18:51:46.502585Z","shell.execute_reply":"2023-11-12T18:51:46.510914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"item[1]['react'][:, 0], item[1]['react'][:, 1]","metadata":{"execution":{"iopub.status.busy":"2023-11-12T17:41:24.753115Z","iopub.execute_input":"2023-11-12T17:41:24.753561Z","iopub.status.idle":"2023-11-12T17:41:24.767173Z","shell.execute_reply.started":"2023-11-12T17:41:24.753536Z","shell.execute_reply":"2023-11-12T17:41:24.765742Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class 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-11-12T17:39:07.433363Z","iopub.execute_input":"2023-11-12T17:39:07.433777Z","iopub.status.idle":"2023-11-12T17:39:07.447643Z","shell.execute_reply.started":"2023-11-12T17:39:07.433746Z","shell.execute_reply":"2023-11-12T17:39:07.446225Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"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(4,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    \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.proj_out(x)\n        \n        return x","metadata":{"execution":{"iopub.status.busy":"2023-11-12T17:43:21.253025Z","iopub.execute_input":"2023-11-12T17:43:21.253493Z","iopub.status.idle":"2023-11-12T17:43:21.268355Z","shell.execute_reply.started":"2023-11-12T17:43:21.253466Z","shell.execute_reply":"2023-11-12T17:43:21.266785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = RNA_Model()","metadata":{"execution":{"iopub.status.busy":"2023-11-12T18:57:18.896316Z","iopub.execute_input":"2023-11-12T18:57:18.897476Z","iopub.status.idle":"2023-11-12T18:57:18.929161Z","shell.execute_reply.started":"2023-11-12T18:57:18.897441Z","shell.execute_reply":"2023-11-12T18:57:18.928159Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model","metadata":{"execution":{"iopub.status.busy":"2023-11-12T18:57:21.508911Z","iopub.execute_input":"2023-11-12T18:57:21.510079Z","iopub.status.idle":"2023-11-12T18:57:21.519243Z","shell.execute_reply.started":"2023-11-12T18:57:21.510033Z","shell.execute_reply":"2023-11-12T18:57:21.517907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def loss(pred, target):\n    p = pred[target['mask'][:,:pred.shape[1]]]\n    y = target['react'][target['mask']].clip(0,1)\n    loss = F.l1_loss(p, y, reduction='none')\n    loss = loss[~torch.isnan(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-11-12T17:45:38.998767Z","iopub.execute_input":"2023-11-12T17:45:38.999225Z","iopub.status.idle":"2023-11-12T17:45:39.009995Z","shell.execute_reply.started":"2023-11-12T17:45:38.99919Z","shell.execute_reply":"2023-11-12T17:45:39.008513Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2023-11-12T17:46:36.103731Z","iopub.execute_input":"2023-11-12T17:46:36.10412Z","iopub.status.idle":"2023-11-12T17:46:36.111192Z","shell.execute_reply.started":"2023-11-12T17:46:36.104092Z","shell.execute_reply":"2023-11-12T17:46:36.109395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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-11-12T17:47:39.435895Z","iopub.execute_input":"2023-11-12T17:47:39.436245Z","iopub.status.idle":"2023-11-12T17:47:39.441594Z","shell.execute_reply.started":"2023-11-12T17:47:39.43622Z","shell.execute_reply":"2023-11-12T17:47:39.440575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seed_everything(SEED)\nos.makedirs(OUT, exist_ok=True)\ndf = pd.read_parquet(os.path.join(PATH,'train_data.parquet'))","metadata":{"execution":{"iopub.status.busy":"2023-11-12T17:46:48.568229Z","iopub.execute_input":"2023-11-12T17:46:48.568626Z","iopub.status.idle":"2023-11-12T17:47:00.20699Z","shell.execute_reply.started":"2023-11-12T17:46:48.568599Z","shell.execute_reply":"2023-11-12T17:47:00.206265Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for fold in [0]: # 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()   \n    model = model.to(device)\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(32, lr_max=5e-4, wd=0.05, pct_start=0.02)\n    torch.save(learn.model.state_dict(),os.path.join(OUT,f'{fname}_{fold}.pth'))\n    gc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"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 = self.df.loc[idx, ['id_min','id_max','sequence']]\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 = [self.seq_map[s] for s in seq]\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}","metadata":{"execution":{"iopub.status.busy":"2023-11-12T17:52:31.979747Z","iopub.execute_input":"2023-11-12T17:52:31.980162Z","iopub.status.idle":"2023-11-12T17:52:31.990636Z","shell.execute_reply.started":"2023-11-12T17:52:31.980133Z","shell.execute_reply":"2023-11-12T17:52:31.989139Z"},"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-12T17:53:21.314064Z","iopub.execute_input":"2023-11-12T17:53:21.316472Z","iopub.status.idle":"2023-11-12T17:53:25.533867Z","shell.execute_reply.started":"2023-11-12T17:53:21.31638Z","shell.execute_reply":"2023-11-12T17:53:25.532922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ids,preds = [],[]\nfor x,y in tqdm(dl):\n    with torch.no_grad(),torch.cuda.amp.autocast():\n        p = torch.stack([torch.nan_to_num(model(x)) for model in models]\n                        ,0).mean(0).clip(0,1)\n        \n    for idx, mask, pi in zip(y['ids'].cpu(), x['mask'].cpu(), p.cpu()):\n        ids.append(idx[mask])\n        preds.append(pi[mask[:pi.shape[0]]])\n\nids = 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":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}