{"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":6784879,"sourceType":"datasetVersion","datasetId":3903927},{"sourceId":6822004,"sourceType":"datasetVersion","datasetId":3719560},{"sourceId":7094457,"sourceType":"datasetVersion","datasetId":4088634},{"sourceId":7094907,"sourceType":"datasetVersion","datasetId":4088873},{"sourceId":7124454,"sourceType":"datasetVersion","datasetId":4109741},{"sourceId":7132723,"sourceType":"datasetVersion","datasetId":4115421},{"sourceId":7143583,"sourceType":"datasetVersion","datasetId":4123327}],"dockerImageVersionId":30587,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pickle\nimport os, gc\nimport numpy as np\nimport random\nimport math\nimport time\nimport pandas as pd\nimport polars as pl\nfrom sklearn.model_selection import KFold\nfrom tqdm import tqdm\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torchmetrics import Metric\nimport csv\nfrom os import path, remove\nimport matplotlib.pyplot as plt","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:08:56.103744Z","iopub.execute_input":"2023-12-07T05:08:56.104620Z","iopub.status.idle":"2023-12-07T05:08:56.111362Z","shell.execute_reply.started":"2023-12-07T05:08:56.104580Z","shell.execute_reply":"2023-12-07T05:08:56.110352Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"PATH = '/kaggle/input/stanford-ribonanza-rna-folding-converted/'\nOUT = './'\nnum_workers = 4\nSEED = 2023\nnfolds = 4\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:08:56.112967Z","iopub.execute_input":"2023-12-07T05:08:56.113255Z","iopub.status.idle":"2023-12-07T05:08:56.123969Z","shell.execute_reply.started":"2023-12-07T05:08:56.113231Z","shell.execute_reply":"2023-12-07T05:08:56.123194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"GENERATE_BPP_FILE = False\nGENERATE_BPP_TRAIN = False\nIN_TRAIN = False\nIN_TEST = True\nIN_TEST1_GENERALIZATION = False\nIN_TEST2_GENERALIZATION = False","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:08:56.124953Z","iopub.execute_input":"2023-12-07T05:08:56.125201Z","iopub.status.idle":"2023-12-07T05:08:56.135013Z","shell.execute_reply.started":"2023-12-07T05:08:56.125179Z","shell.execute_reply":"2023-12-07T05:08:56.134106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**GETTING BPP INFO**","metadata":{}},{"cell_type":"code","source":"if GENERATE_BPP_FILE:\n    root_dir = '/kaggle/input/stanford-ribonanza-rna-folding/Ribonanza_bpp_files/extra_data'\n    seq_id = []\n    file_paths = []\n\n    i = 0\n    for folder, _, files in tqdm(os.walk(root_dir), total=len(os.listdir(root_dir))):\n        for file in files:\n            seq_id.append(file.split('.', 1)[0])\n            file_paths.append(os.path.join(folder, file))\n            i += 1\n            if i % 100000 == 0:\n                print(i)\n                \n    df_file_bpps = pd.DataFrame({'seq_id': seq_id, 'file_path': file_paths})\n    df_file_bpps.set_index('seq_id', inplace=True)\n    df_file_bpps.to_csv('rna-bpp-files.csv')\nelse:\n    df_file_bpps = pd.read_csv('/kaggle/input/standford-ribonanza-rna-bpp-files/rna-bpp-files.csv')\n    df_file_bpps.drop_duplicates(subset=['seq_id'], inplace=True, ignore_index=True)\n    df_file_bpps.set_index('seq_id', inplace=True)","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:08:56.136993Z","iopub.execute_input":"2023-12-07T05:08:56.137370Z","iopub.status.idle":"2023-12-07T05:09:01.815743Z","shell.execute_reply.started":"2023-12-07T05:08:56.137346Z","shell.execute_reply":"2023-12-07T05:09:01.814962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_bpp_max(sequence_id, Lmax, src_mask=None):\n    seq_len = src_mask.sum(-1).max() if src_mask != None else Lmax\n    bpp = np.zeros(Lmax, dtype=np.double)\n    filename = df_file_bpps.loc[sequence_id, 'file_path']\n    file = open(filename, \"r\")\n    \n    for line in file:\n        if line != '\\n':\n            line = line.split(' ')\n            pos1 = int(line[0]) - 1\n            pos2 = int(line[1]) - 1\n            prob = float(line[2])\n            if prob > bpp[pos1]:\n                bpp[pos1] = prob\n            if prob > bpp[pos2]:\n                bpp[pos2] = prob\n            \n    file.close()\n    return bpp","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:01.816724Z","iopub.execute_input":"2023-12-07T05:09:01.817009Z","iopub.status.idle":"2023-12-07T05:09:01.824770Z","shell.execute_reply.started":"2023-12-07T05:09:01.816984Z","shell.execute_reply":"2023-12-07T05:09:01.823877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if GENERATE_BPP_TRAIN:\n    ! rm rna_train_bpp.csv\n    df_train = pd.read_parquet('/kaggle/input/stanford-ribonanza-rna-folding-converted/train_data.parquet')\n    df_train = df_train[df_train.SN_filter == 1]\n    df_train.drop_duplicates(subset=['sequence_id'], inplace=True, ignore_index=True)\n    df_train['L'] = df_train.sequence.apply(len)\n    Lmax = df_train['L'].max()\n    \n    ids, bpps = [],[]\n    for i in range(df_train.shape[0]):\n        sequence_id = df_train.loc[i, 'sequence_id']\n        ids.append(sequence_id)\n        bpps.append(get_bpp_max(sequence_id, Lmax))\n        if i % 10000 == 0:\n            print(i)\n    df_train_bpp = pd.DataFrame(bpps, index=ids) \n    df_train_bpp.to_csv('rna-train-bpp.csv')\n    ! zip rna-train-bpp.zip rna-train-bpp.csv\n    del ids\n    del bpps\n    gc.collect()\nelse:\n    df_train_bpp = pd.read_csv('/kaggle/input/rna-train-bpp/rna-train-bpp.csv', index_col=0)","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:01.826094Z","iopub.execute_input":"2023-12-07T05:09:01.826437Z","iopub.status.idle":"2023-12-07T05:09:10.233250Z","shell.execute_reply.started":"2023-12-07T05:09:01.826405Z","shell.execute_reply":"2023-12-07T05:09:10.232311Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**DISTANCE MATRIX**","metadata":{}},{"cell_type":"code","source":"def get_distance_matrix(Lmax):\n    ## adjacent matrix based on distance on the sequence\n    ## D[i, j] = 1 / (abs(i - j) + 1) ** pow, pow = 1, 2, 4\n    \n    idx = np.arange(Lmax)\n    Ds = []\n    for i in range(len(idx)):\n        d = np.abs(idx[i] - idx)\n        Ds.append(d)\n    \n    Ds = np.array(Ds) + 1\n    Ds = 1 / Ds\n    Ds = Ds[:,:]\n\n    Dss = []\n    for i in [1, 2, 4]: \n        Dss.append(Ds ** i)\n    Ds = np.stack(Dss, axis = 0)\n    Ds = Ds.reshape(3, Lmax, Lmax)\n    return Ds","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:10.236533Z","iopub.execute_input":"2023-12-07T05:09:10.236826Z","iopub.status.idle":"2023-12-07T05:09:10.243439Z","shell.execute_reply.started":"2023-12-07T05:09:10.236789Z","shell.execute_reply":"2023-12-07T05:09:10.242553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**UTILS**","metadata":{}},{"cell_type":"code","source":"def seed_everything(seed=42):\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    \ndef save_weights(model, optimizer, epoch, folder):\n    if os.path.isdir(folder) == False:\n        os.makedirs(folder, exist_ok=True)\n    torch.save(model.state_dict(), folder + '/epoch{}.ckpt'.format(epoch + 1))\n\ndef get_best_weights_from_fold(fold,top=1):\n    csv_file='log_fold{}.csv'.format(fold)\n\n    history=pd.read_csv(csv_file)\n    scores=np.asarray(history.val_acc)\n    top_epochs=scores.argsort()[-3:][::-1]\n    print(scores[top_epochs])\n    os.system('mkdir best_weights')\n\n    for i in range(top):\n        weights_path='checkpoints_fold{}/epoch{}.ckpt'.format(fold,history.epoch[top_epochs[i]])\n        print(weights_path)\n        os.system('cp {} best_weights/fold{}top{}.ckpt'.format(weights_path,fold,i+1))\n    os.system('rm -r checkpoints_fold{}'.format(fold))","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:10.244605Z","iopub.execute_input":"2023-12-07T05:09:10.244944Z","iopub.status.idle":"2023-12-07T05:09:10.257844Z","shell.execute_reply.started":"2023-12-07T05:09:10.244914Z","shell.execute_reply":"2023-12-07T05:09:10.256977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**LOGGER**","metadata":{}},{"cell_type":"code","source":"class CSVLogger:\n    def __init__(self,columns,file):\n        self.columns = columns\n        self.file = file\n        if not self.check_header():\n            self._write_header()\n\n    def check_header(self):\n        if path.exists(self.file):\n            # with open(self.file, 'r') as csvfile:\n            #     sniffer = csv.Sniffer()\n            #     has_header = sniffer.has_header(csvfile.read(2048))\n            #     header=csvfile.seek(0)\n            header = True\n        else:\n            header = False\n        return header\n\n    def _write_header(self):\n        with open(self.file,\"a\") as f:\n            string=\"\"\n            for attrib in self.columns:\n                string += \"{},\".format(attrib)\n            string = string[:len(string)-1]\n            string += \"\\n\"\n            f.write(string)\n        return self\n\n    def log(self,row):\n        if len(row) != len(self.columns):\n            raise Exception(\"Mismatch between row vector and number of columns in logger\")\n        with open(self.file,\"a\") as f:\n            string = \"\"\n            for attrib in row:\n                string += \"{},\".format(attrib)\n            string = string[:len(string)-1]\n            string += \"\\n\"\n            f.write(string)\n        return self","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:10.259121Z","iopub.execute_input":"2023-12-07T05:09:10.259835Z","iopub.status.idle":"2023-12-07T05:09:10.273941Z","shell.execute_reply.started":"2023-12-07T05:09:10.259788Z","shell.execute_reply":"2023-12-07T05:09:10.273016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**LR SCHEDULER**","metadata":{}},{"cell_type":"code","source":"def update_lr(optimizer, lr):\n    for param_group in optimizer.param_groups:\n        param_group['lr'] = lr\n\nclass lr_AIAYN():\n    '''\n    Learning rate scheduler from the paper:\n    Attention is All You Need\n    '''\n    def __init__(self,optimizer, d_model, warmup_steps=4000, factor=1):\n        self.optimizer = optimizer\n        self.d_model = d_model\n        self.warmup_steps = warmup_steps\n        self.step_num = 0\n        self.factor = factor\n\n    def step(self):\n        self.step_num += 1\n        lr=self.d_model ** -0.5 * np.min([self.step_num ** -0.5,\n                                         self.step_num * self.warmup_steps ** -1.5]) * self.factor\n        update_lr(self.optimizer, lr)\n        return lr","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:10.275211Z","iopub.execute_input":"2023-12-07T05:09:10.275821Z","iopub.status.idle":"2023-12-07T05:09:10.288181Z","shell.execute_reply.started":"2023-12-07T05:09:10.275773Z","shell.execute_reply":"2023-12-07T05:09:10.287383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**DATASETS**","metadata":{}},{"cell_type":"code","source":"class RNA_Dataset(Dataset):\n    def __init__(self, df, mode='train', denoise=True, seed=2023, fold=0, nfolds=4, \n                 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.dm = get_distance_matrix(self.Lmax)\n        df_2A3 = df.loc[df.experiment_type=='2A3_MaP']\n        df_DMS = df.loc[df.experiment_type=='DMS_MaP']\n        \n        if fold != -1:    # all dataset\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        if denoise:\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            df_2A3.loc[df_2A3.SN_filter == 0, [c for c in df_2A3.columns if 'reactivity_0' in c]] = None\n            df_DMS.loc[df_DMS.SN_filter == 0, [c for c in df_DMS.columns if 'reactivity_0' in c]] = None\n        \n        self.seq_id = df_2A3['sequence_id'].values\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.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        sn = torch.FloatTensor([self.sn_2A3[idx], self.sn_DMS[idx]])\n        \n        # bpp = get_bpp_max(self.seq_id[idx], self.Lmax)   \n        bpp = df_train_bpp.loc[self.seq_id[idx]].values\n        \n        return {'seq':torch.from_numpy(seq), 'mask':mask, 'bpp': bpp, 'dm': self.dm}, \\\n               {'react':react, 'sn':sn, 'mask':mask}\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-07T05:09:10.289524Z","iopub.execute_input":"2023-12-07T05:09:10.289872Z","iopub.status.idle":"2023-12-07T05:09:10.316403Z","shell.execute_reply.started":"2023-12-07T05:09:10.289838Z","shell.execute_reply":"2023-12-07T05:09:10.315537Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**LOSS & METRIC**","metadata":{}},{"cell_type":"code","source":"def loss_fn(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-12-07T05:09:10.317518Z","iopub.execute_input":"2023-12-07T05:09:10.317784Z","iopub.status.idle":"2023-12-07T05:09:10.331377Z","shell.execute_reply.started":"2023-12-07T05:09:10.317753Z","shell.execute_reply":"2023-12-07T05:09:10.330594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def accuracy_MAE(pred, target, mask):\n    p = pred[mask[:, :pred.shape[1]]]\n    y = target[mask].clip(0, 1)\n    \n    loss = F.l1_loss(p, y, reduction='none')\n    loss = loss[~torch.isnan(loss)].mean()\n    \n    return loss","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:10.332678Z","iopub.execute_input":"2023-12-07T05:09:10.333147Z","iopub.status.idle":"2023-12-07T05:09:10.341744Z","shell.execute_reply.started":"2023-12-07T05:09:10.333110Z","shell.execute_reply":"2023-12-07T05:09:10.340951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**MODEL**","metadata":{}},{"cell_type":"code","source":"# from torch.nn.parameter import Parameter\n\nclass TransformerEncoderLayer(nn.Module):\n    r\"\"\"TransformerEncoderLayer is made up of self-attn and feedforward network.\n    This standard encoder layer is based on the paper \"Attention Is All You Need\".\n    Ashish Vaswani, Noam Shazeer, Niki Parmar, Jakob Uszkoreit, Llion Jones, Aidan N Gomez,\n    Lukasz Kaiser, and Illia Polosukhin. 2017. Attention is all you need. In Advances in\n    Neural Information Processing Systems, pages 6000-6010. Users may modify or implement\n    in a different way during application.\n\n    Args:\n        d_model: the number of expected features in the input (required).\n        nhead: the number of heads in the multiheadattention models (required).\n        dim_feedforward: the dimension of the feedforward network model (default=2048).\n        dropout: the dropout value (default=0.1).\n        activation: the activation function of intermediate layer, relu or gelu (default=relu).\n\n    Examples::\n        >>> encoder_layer = nn.TransformerEncoderLayer(d_model=512, nhead=8)\n        >>> src = torch.rand(10, 32, 512)\n        >>> out = encoder_layer(src)\n    \"\"\"\n\n    def __init__(self, d_model, nhead, dim_feedforward=2048, dropout=0.1, activation=\"relu\"):\n        super(TransformerEncoderLayer, self).__init__()\n        self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)\n        self.linear1 = nn.Linear(d_model, dim_feedforward)\n        self.dropout = nn.Dropout(dropout)\n        self.linear2 = nn.Linear(dim_feedforward, d_model)\n\n        self.norm1 = nn.LayerNorm(d_model)\n        self.norm2 = nn.LayerNorm(d_model)\n        self.dropout1 = nn.Dropout(dropout)\n        self.dropout2 = nn.Dropout(dropout)\n\n        self.activation = nn.ReLU()\n\n\n    def forward(self, src, src_mask=None, src_key_padding_mask=None):\n        src2, attention_weights = self.self_attn(src, src, src, attn_mask=src_mask,\n                                  key_padding_mask=src_key_padding_mask)\n        src = src + self.dropout1(src2)\n        src = self.norm1(src)\n        src2 = self.linear2(self.dropout(self.activation(self.linear1(src))))\n        src = src + self.dropout2(src2)\n        src = self.norm2(src)\n        return src, attention_weights\n\n\nclass LinearDecoder(nn.Module):\n    def __init__(self, num_classes, ninp, dropout, pool=True,):\n        super(LinearDecoder, self).__init__()\n        if pool:\n            self.classifier=nn.Linear(ninp,num_classes)\n        else:\n            self.classifier=nn.Linear(ninp,num_classes)\n        self.pool=pool\n\n    def forward(self, x):\n        if self.pool:\n            x, _ = torch.max(x,dim=1)\n        x = self.classifier(x)\n        return x\n    \nclass K_mer_aggregate(nn.Module):\n    def __init__(self, kmers, kmers_padding, in_dim, out_dim,dropout=0.1):\n        super(K_mer_aggregate, self).__init__()\n        self.dropout = nn.Dropout(dropout)\n        self.convs = []\n        for (i, k) in enumerate(kmers):\n            self.convs.append(nn.Conv1d(in_dim, out_dim, k, padding=kmers_padding[i]))\n        self.convs = nn.ModuleList(self.convs)\n        self.activation = nn.ReLU(inplace=True)\n        self.norm = nn.LayerNorm(out_dim)\n\n    def forward(self, x):\n        outputs=[]\n        for conv in self.convs:\n            outputs.append(conv(x))\n        outputs=torch.cat(outputs,dim=2)\n        outputs=self.norm(outputs.permute(0,2,1)).permute(0,2,1)\n        return outputs\n\n\nclass NucleicTransformer(nn.Module):\n    def __init__(self, ntoken, nclass, ninp, nhead, nhid, nlayers, kmer_aggregation, kmers, kmers_padding, dropout=0.5, return_aw=False):\n        super(NucleicTransformer, self).__init__()\n        self.model_type = 'Transformer'\n        self.src_mask = None\n        self.kmers = kmers\n        self.kmers_padding = kmers_padding\n        self.kmer_aggregation = kmer_aggregation\n        if self.kmer_aggregation:\n            self.k_mer_aggregate = K_mer_aggregate(kmers, kmers_padding, ninp, ninp)\n        else:\n            print(\"No kmer aggregation is chosen\")\n        self.transformer_encoder = []\n        for i in range(nlayers):\n            self.transformer_encoder.append(TransformerEncoderLayer(ninp, nhead, nhid, dropout))\n        self.transformer_encoder= nn.ModuleList(self.transformer_encoder)\n        self.encoder = nn.Embedding(ntoken, ninp)\n        self.ninp = ninp\n        \n        self.adj_learned = nn.Linear(3, 1)\n        \n        self.bpp_aggregate1 = nn.Linear(ninp + 1, ninp * 2)\n        self.bpp_aggregate2 = nn.Linear(ninp * 2, ninp)\n        \n        self.decoder = LinearDecoder(nclass,ninp,dropout,pool=False)\n        self.return_aw = return_aw\n\n    def _generate_square_subsequent_mask(self, sz):\n        mask = (torch.triu(torch.ones(sz, sz)) == 1).transpose(0, 1)\n        mask = mask.float().masked_fill(mask == 0, float('-inf')).masked_fill(mask == 1, float(0.0))\n        print(f'MASK: {mask.shape}, {mask}')\n        return mask\n\n    def forward(self, src, src_mask=None, bpp=None, dm=None):\n        if src_mask != None:\n            Lmax = src_mask.sum(-1).max()\n            mask = src_mask\n        else:\n            mask = None\n\n        src = src.permute(1, 0)\n        src = self.encoder(src)  \n        \n        if self.kmer_aggregation:\n            src = self.k_mer_aggregate(src.permute(1,2,0)).permute(2,0,1)\n            \n        src = src.double()  \n        attention_weights=[]\n        for layer in self.transformer_encoder:\n            src,attention_weights_layer=layer(src, src_mask=dm, src_key_padding_mask=~mask)\n            attention_weights.append(attention_weights_layer)\n            \n        attention_weights=torch.stack(attention_weights).permute(1,0,2,3)\n        encoder_output = src.permute(1,0,2)        \n        \n        if bpp != None:\n            bpp = bpp[:, :, None]\n            output = torch.cat((encoder_output, bpp), dim=2)\n            output = self.bpp_aggregate1(output)\n            output = self.bpp_aggregate2(output)\n        else:\n            output = encoder_output        \n            \n        output = self.decoder(output)\n\n        if self.return_aw:\n            return output,attention_weights\n        else:\n            return output","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:10.343201Z","iopub.execute_input":"2023-12-07T05:09:10.343539Z","iopub.status.idle":"2023-12-07T05:09:10.374592Z","shell.execute_reply.started":"2023-12-07T05:09:10.343509Z","shell.execute_reply":"2023-12-07T05:09:10.373730Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**TRAIN**","metadata":{}},{"cell_type":"code","source":"params = {'ntoken': 4,          # number of tokens to represent DNA nucleotides (should always be 4)\n          'npred': 2,           # number of predictions from the linear decoder\n          'ninp': 192,          # ninp for transformer encoder\n          'nhead': 3,           # nhead for transformer encoder\n          'nhid': 3 * 192,      # nhid for transformer encoder\n          'nlayers': 10,         # nlayers for transformer encoder\n          'weight_decay': 0,    # weight dacay used in optimizer\n          'kmer_aggregation': True,  # when to use kmers\n          'kmers': [5],         # k-mers to be aggregated\n          'kmers_padding': [2], # k-mers padding for conv layer\n          'dropout': 0.1,       # transformer dropout\n          'nfolds': 5,          # number of cross validation folds\n          'fold': -1,           # witch fold to train\n          'batch_size': 64,\n          'shuffle': False, \n          'num_workers': 4,\n          'epochs': 100,        # number of epochs to train\n          'lr': 1e-4,\n          'lr_scale': 0.1,      # learning rate scale\n          'warmup_steps': 3200, # training schedule warmup steps\n          'save_freq': 1        # saving checkpoints per save_freq epochs\n         }","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:10.378145Z","iopub.execute_input":"2023-12-07T05:09:10.378418Z","iopub.status.idle":"2023-12-07T05:09:10.390151Z","shell.execute_reply.started":"2023-12-07T05:09:10.378395Z","shell.execute_reply":"2023-12-07T05:09:10.389363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def validate(model, device, dataset, batch_size=64):\n    batches = len(dataset)\n    model.train(False)\n    total = 0\n    predictions = []\n    loss = 0\n    criterion = loss_fn\n    with torch.no_grad():\n        for data, labels in tqdm(dataset):\n            X = data['seq'].to(device)\n            dm = data['dm'].to(device)\n            dm = dm.permute(1, 0, 2, 3)\n            dm = dm.reshape(-1, dm.shape[2], dm.shape[3])\n            bpp = data['bpp'].to(device)\n            Y = labels \n            output= model(X, data['mask'].to(device), bpp, dm)\n            del X\n            loss += criterion(output, Y)\n            del output\n\n    torch.cuda.empty_cache()\n    val_loss = (loss/batches).cpu()\n    print(f'Val Loss: {val_loss}')\n    return val_loss ","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:10.391320Z","iopub.execute_input":"2023-12-07T05:09:10.391593Z","iopub.status.idle":"2023-12-07T05:09:10.405133Z","shell.execute_reply.started":"2023-12-07T05:09:10.391569Z","shell.execute_reply":"2023-12-07T05:09:10.404380Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if IN_TRAIN:\n    seed_everything(SEED)\n    os.makedirs(OUT, exist_ok=True)\n    df = pd.read_parquet(os.path.join(PATH,'train_data.parquet'))\n    df = df.sort_values(['sequence_id', 'experiment_type', 'signal_to_noise'], ascending=[True, True, False])\n    df.drop_duplicates(subset=['sequence_id', 'experiment_type'], inplace=True, ignore_index=True)\n\n    df.drop([c for c in df.columns if 'reactivity_error_' in c], axis=1, inplace=True)\n\n    for fold in [params['fold']]: # 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=params['batch_size'],\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        if fold != -1:\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=params['batch_size'], \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\n        gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:10.406271Z","iopub.execute_input":"2023-12-07T05:09:10.407037Z","iopub.status.idle":"2023-12-07T05:09:10.419561Z","shell.execute_reply.started":"2023-12-07T05:09:10.406979Z","shell.execute_reply":"2023-12-07T05:09:10.418630Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if IN_TRAIN:\n    # model_to_continue = None\n    model_to_continue = '/kaggle/input/rna-tmodel-v36/epoch32.ckpt'\n    \n    # checkpointing\n    checkpoints_folder = f\"checkpoints_fold{params['fold']}\"\n    csv_file = f\"log_fold{params['fold']}.csv\"\n    columns = ['epoch', 'train_loss', 'train_acc',\n               'val_loss', 'val_acc', 'val_sens', 'val_spec']\n    logger = CSVLogger(columns, csv_file)\n\n    # build model and logger\n    model = NucleicTransformer(params['ntoken'], params['npred'], params['ninp'], params['nhead'], params['nhid'],\n                               params['nlayers'], params['kmer_aggregation'], kmers=params['kmers'], kmers_padding=params['kmers_padding'],\n                               dropout=params['dropout']).to(device)\n    model = model.double()  \n    \n    if model_to_continue != None:\n        model.load_state_dict(torch.load(model_to_continue))\n    \n    optimizer = torch.optim.Adam(model.parameters(), weight_decay=params['weight_decay'])\n    criterion = loss_fn\n    lr_schedule = lr_AIAYN(optimizer, params['ninp'], params['warmup_steps'], params['lr_scale'])\n\n    pytorch_total_params = sum(p.numel() for p in model.parameters())\n    print(f'Total number of parameters: {pytorch_total_params}')\n\n    print(f\"Starting training for fold {params['fold']}/{params['nfolds']}\")\n    #training loop\n    for epoch in range(params['epochs']):\n        model.train(True)\n        t = time.time()\n        total_loss = 0\n        optimizer.zero_grad()\n        total_steps = len(dl_train)\n        masks = []\n\n        for step, (train_features, train_labels) in enumerate(tqdm(dl_train, leave=True)):\n            lr = lr_schedule.step()\n\n            src = train_features['seq'].to(device)  \n            mask = train_features['mask'].to(device) \n            bpp = train_features['bpp'].to(device)\n            dm = train_features['dm'].to(device)\n            dm = dm.permute(1, 0, 2, 3)\n            dm = dm.reshape(-1, dm.shape[2], dm.shape[3])\n\n            output = model(src, mask, bpp, dm)\n\n            loss = torch.mean(criterion(output, train_labels))\n            loss.backward()\n\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1)\n            optimizer.step()\n            optimizer.zero_grad()\n            total_loss += loss\n\n        print('')\n\n        train_loss = total_loss / (step + 1)\n\n        if params['fold'] != -1:\n            val_loss = validate(model, device, dl_val, batch_size=params['batch_size'] * 2)\n        else:\n            val_loss = 0\n\n        to_log = [epoch + 1, train_loss, 0, val_loss, 0, 0, 0]\n        logger.log(to_log) \n\n        if (epoch + 1) % params['save_freq'] == 0:\n            save_weights(model, optimizer, epoch, checkpoints_folder)\n\n    get_best_weights_from_fold(params['fold'])","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:10.421135Z","iopub.execute_input":"2023-12-07T05:09:10.421576Z","iopub.status.idle":"2023-12-07T05:09:10.440217Z","shell.execute_reply.started":"2023-12-07T05:09:10.421541Z","shell.execute_reply":"2023-12-07T05:09:10.439170Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**INFERENCE**","metadata":{}},{"cell_type":"code","source":"if IN_TEST:\n    MODELS = ['/kaggle/input/rna-tmodel-v36-2/epoch26.ckpt']\n    PATH = '/kaggle/input/stanford-ribonanza-rna-folding-converted/'\n    num_workers = 4\n    device = 'cuda' if torch.cuda.is_available() else 'cpu'","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:10.441570Z","iopub.execute_input":"2023-12-07T05:09:10.442057Z","iopub.status.idle":"2023-12-07T05:09:10.455192Z","shell.execute_reply.started":"2023-12-07T05:09:10.442021Z","shell.execute_reply":"2023-12-07T05:09:10.454217Z"},"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.dm = get_distance_matrix(self.Lmax)\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_id, seq = self.df.loc[idx, ['id_min','id_max', 'sequence_id', '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        bpp = get_bpp_max(seq_id, self.Lmax)   \n        dm = get_distance_matrix(self.Lmax)\n        \n        return {'seq':torch.from_numpy(seq), 'mask':mask, 'bpp':bpp, 'dm': self.dm}, \\\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-12-07T05:09:10.456390Z","iopub.execute_input":"2023-12-07T05:09:10.456721Z","iopub.status.idle":"2023-12-07T05:09:10.473739Z","shell.execute_reply.started":"2023-12-07T05:09:10.456694Z","shell.execute_reply":"2023-12-07T05:09:10.472648Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if IN_TEST:\n    df_test = pd.read_parquet(os.path.join(PATH,'test_sequences.parquet'))\n\n    if IN_TEST1_GENERALIZATION:\n        df_test = df_test[(df_test.id_min >= 269545321) & (df_test.id_max <= 269724007)].reset_index(drop=True)\n\n    ds = RNA_Dataset_Test(df_test)\n    dl = DeviceDataLoader(torch.utils.data.DataLoader(ds, batch_size=params['batch_size'], \n                          shuffle=False, drop_last=False, num_workers=num_workers), device)\n    del df_test\n    gc.collect()\n\n    models = []\n    for m in MODELS:\n        model = NucleicTransformer(params['ntoken'], params['npred'], params['ninp'], params['nhead'], params['nhid'],\n                                   params['nlayers'], params['kmer_aggregation'], kmers=params['kmers'], kmers_padding=params['kmers_padding'],\n                                   dropout=params['dropout']).to(device)\n        model = model.to(device).double()\n        model.load_state_dict(torch.load(m, map_location=torch.device('cpu')))\n        model.eval()\n        models.append(model)","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:10.474979Z","iopub.execute_input":"2023-12-07T05:09:10.475276Z","iopub.status.idle":"2023-12-07T05:09:13.995835Z","shell.execute_reply.started":"2023-12-07T05:09:10.475252Z","shell.execute_reply":"2023-12-07T05:09:13.994862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if IN_TEST:\n    ! rm \"submission.csv\"\n    first_write = True\n    for x, y in tqdm(dl):\n\n        dm = x['dm'].to(device)\n        dm = dm.permute(1, 0, 2, 3)\n        dm = dm.reshape(-1, dm.shape[2], dm.shape[3])\n\n        with torch.no_grad(),torch.cuda.amp.autocast():\n            p = torch.stack([torch.nan_to_num(model(x['seq'].to(device), x['mask'].to(device), x['bpp'].to(device), dm)) for model in models]\n                            ,0).mean(0).clip(0,1)\n\n        ids, preds = [],[]\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\n        ids = torch.concat(ids)\n        preds = torch.concat(preds)\n\n        df = pd.DataFrame({'id':ids.numpy(), 'reactivity_DMS_MaP':preds[:,1].numpy(), \n                           'reactivity_2A3_MaP':preds[:,0].numpy()})\n        df.to_csv('submission.csv', index=False, float_format='%.4f', mode='a', header=first_write) # 6.5GB\n        first_write = False","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:13.997593Z","iopub.execute_input":"2023-12-07T05:09:13.997931Z","iopub.status.idle":"2023-12-07T05:09:24.190447Z","shell.execute_reply.started":"2023-12-07T05:09:13.997902Z","shell.execute_reply":"2023-12-07T05:09:24.189343Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**TEST1 TO CHECK GENERALIZATION**","metadata":{}},{"cell_type":"code","source":"if IN_TEST1_GENERALIZATION:\n    import matplotlib.pyplot as plt\n    #read your sub here\n    df = pd.read_csv(\"/kaggle/working/submission.csv\")\n    #some parameters\n    font_size = 6\n    id1 = 269545321\n    id2 = 269724007\n    reshape1 = 391\n    reshape2 = 457\n    #get predictions\n    pred_DMS = df['reactivity_DMS_MaP'].to_numpy().reshape(reshape1,reshape2)\n    pred_2A3 = df['reactivity_2A3_MaP'].to_numpy().reshape(reshape1,reshape2)\n    #plot mutate and map\n    fig = plt.figure()\n    plt.subplot(121)\n    plt.title(f'reactivity_DMS_MaP', fontsize=font_size)\n    plt.imshow(pred_DMS, vmin=0, vmax=1, cmap='gray_r')\n    plt.subplot(122)\n    plt.title(f'reactivity_2A3_MaP', fontsize=font_size)\n    plt.imshow(pred_2A3, vmin=0, vmax=1, cmap='gray_r')\n    plt.tight_layout()\n    plt.savefig(f\"plot_test1.png\",dpi=500)\n    plt.clf()\n    plt.close()","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:24.192108Z","iopub.execute_input":"2023-12-07T05:09:24.192519Z","iopub.status.idle":"2023-12-07T05:09:26.480342Z","shell.execute_reply.started":"2023-12-07T05:09:24.192474Z","shell.execute_reply":"2023-12-07T05:09:26.479245Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if IN_TEST1_GENERALIZATION:\n    from IPython.display import FileLink\n\n    !zip plot_test1.zip plot_test1.png\n\n    os.chdir(r'/kaggle/working')\n    FileLink(r'plot_test1.zip')","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:26.481740Z","iopub.execute_input":"2023-12-07T05:09:26.482099Z","iopub.status.idle":"2023-12-07T05:09:27.690397Z","shell.execute_reply.started":"2023-12-07T05:09:26.482070Z","shell.execute_reply":"2023-12-07T05:09:27.689148Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**R1138 TEST**","metadata":{}},{"cell_type":"code","source":"if IN_TEST2_GENERALIZATION:\n    root_dir = '/kaggle/input/r1138-bpp/R1138_bpp_files'\n    seq_id = []\n    file_paths = []\n\n    i = 0\n    for folder, _, files in tqdm(os.walk(root_dir), total=len(os.listdir(root_dir))):\n        for file in files:\n            seq_id.append(file.split('.', 1)[0])\n            file_paths.append(os.path.join(folder, file))\n            i += 1\n            if i % 100000 == 0:\n                print(i)","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:27.692399Z","iopub.execute_input":"2023-12-07T05:09:27.692860Z","iopub.status.idle":"2023-12-07T05:09:27.983316Z","shell.execute_reply.started":"2023-12-07T05:09:27.692803Z","shell.execute_reply":"2023-12-07T05:09:27.982357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if IN_TEST2_GENERALIZATION:\n    df_file_bpps = pd.DataFrame({'seq_id': seq_id, 'file_path': file_paths})\n    df_file_bpps.set_index('seq_id', inplace=True)\n    df_file_bpps.to_csv('rna-r1138-bpp-files.csv')\n    \n    df_test = pd.read_csv('/kaggle/input/r1138v1-m2/R1138v1_m2.csv')\n    df_test['L'] = df_test.sequence.apply(len)\n    \n    SEQ_LEN = 720\n    for i in range(df_test.shape[0]):\n        df_test.loc[i, 'id_min'] = i * SEQ_LEN \n        df_test.loc[i, 'id_max'] = ((i + 1) * SEQ_LEN) - 1","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:27.984595Z","iopub.execute_input":"2023-12-07T05:09:27.984907Z","iopub.status.idle":"2023-12-07T05:09:28.460242Z","shell.execute_reply.started":"2023-12-07T05:09:27.984880Z","shell.execute_reply":"2023-12-07T05:09:28.459118Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if IN_TEST2_GENERALIZATION:\n    MODELS = ['/kaggle/input/rna-tmodel-v36-2/epoch26.ckpt']\n    \n    ds = RNA_Dataset_Test(df_test)\n    dl = DeviceDataLoader(torch.utils.data.DataLoader(ds, batch_size=params['batch_size'], \n                          shuffle=False, drop_last=False, num_workers=num_workers), device)\n    # del df_test\n    gc.collect()\n\n    models = []\n    for m in MODELS:\n        model = NucleicTransformer(params['ntoken'], params['npred'], params['ninp'], params['nhead'], params['nhid'],\n                                   params['nlayers'], params['kmer_aggregation'], kmers=params['kmers'], kmers_padding=params['kmers_padding'],\n                                   dropout=params['dropout']).to(device)\n        model = model.to(device).double()\n        model.load_state_dict(torch.load(m, map_location=torch.device('cpu')))\n        model.eval()\n        models.append(model)","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:28.461781Z","iopub.execute_input":"2023-12-07T05:09:28.462095Z","iopub.status.idle":"2023-12-07T05:09:28.816588Z","shell.execute_reply.started":"2023-12-07T05:09:28.462070Z","shell.execute_reply":"2023-12-07T05:09:28.815569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if IN_TEST2_GENERALIZATION:\n    m2_preds = torch.tensor([]).to(device)\n    for x, y in tqdm(dl):\n\n        dm = x['dm'].to(device)\n        dm = dm.permute(1, 0, 2, 3)\n        dm = dm.reshape(-1, dm.shape[2], dm.shape[3])\n\n        with torch.no_grad(),torch.cuda.amp.autocast():\n            p = torch.stack([torch.nan_to_num(model(x['seq'].to(device), x['mask'].to(device), x['bpp'].to(device), dm)) for model in models]\n                            ,0).mean(0).clip(0,1)\n            m2_preds = torch.concat([m2_preds, p])","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:28.818640Z","iopub.execute_input":"2023-12-07T05:09:28.818977Z","iopub.status.idle":"2023-12-07T05:09:40.392889Z","shell.execute_reply.started":"2023-12-07T05:09:28.818950Z","shell.execute_reply":"2023-12-07T05:09:40.391400Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if IN_TEST2_GENERALIZATION:\n    #read your sub here\n    font_size = 6\n    #get predictions\n    pred_DMS = m2_preds[:,:,1]\n    pred_2A3 = m2_preds[:,:,0]\n    #plot mutate and map\n    fig = plt.figure()\n    plt.subplot(121)\n    plt.title(f'reactivity_DMS_MaP', fontsize=font_size)\n    plt.imshow(pred_DMS.cpu(), vmin=0, vmax=1, cmap='gray_r')\n    plt.subplot(122)\n    plt.title(f'reactivity_2A3_MaP', fontsize=font_size)\n    plt.imshow(pred_2A3.cpu(), vmin=0, vmax=1, cmap='gray_r')\n    plt.tight_layout()\n    plt.savefig(f\"plot_test2.png\",dpi=500)\n    plt.clf()\n    plt.close()","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:40.394619Z","iopub.execute_input":"2023-12-07T05:09:40.395496Z","iopub.status.idle":"2023-12-07T05:09:42.699119Z","shell.execute_reply.started":"2023-12-07T05:09:40.395454Z","shell.execute_reply":"2023-12-07T05:09:42.698284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if IN_TEST2_GENERALIZATION:\n    !zip plot_test2.zip plot_test2.png\n\n    os.chdir(r'/kaggle/working')\n    from IPython.display import FileLink\n\n    FileLink(r'plot_test2.zip')","metadata":{"execution":{"iopub.status.busy":"2023-12-07T05:09:42.700446Z","iopub.execute_input":"2023-12-07T05:09:42.700873Z","iopub.status.idle":"2023-12-07T05:09:43.882249Z","shell.execute_reply.started":"2023-12-07T05:09:42.700835Z","shell.execute_reply":"2023-12-07T05:09:43.881185Z"},"trusted":true},"execution_count":null,"outputs":[]}]}