{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":51294,"databundleVersionId":6923401,"sourceType":"competition"},{"sourceId":6888177,"sourceType":"datasetVersion","datasetId":3889652}],"dockerImageVersionId":30559,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Based on https://github.com/huggingface/transformers/blob/main/src/transformers/models/mistral/modeling_mistral.py\n# data based on but for full train https://www.kaggle.com/code/crimson206/rna-seq-struct-flexibletransformer/notebook\n# I did not implemented sliding window attention since it is not necessary here with max seq len ~ 400\n# current LB is 0.15760","metadata":{"execution":{"iopub.status.busy":"2023-11-27T07:05:10.746320Z","iopub.execute_input":"2023-11-27T07:05:10.746795Z","iopub.status.idle":"2023-11-27T07:05:10.780291Z","shell.execute_reply.started":"2023-11-27T07:05:10.746757Z","shell.execute_reply":"2023-11-27T07:05:10.778983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nfrom pathlib import Path\nimport gc\nimport math\nfrom tqdm.auto import tqdm\nfrom typing import List, Dict, Tuple, Union, Any, Optional\nimport os\nos.environ[\"TOKENIZERS_PARALLELISM\"] = \"false\"\n\nfrom transformers import TrainingArguments, Trainer\nfrom transformers.modeling_outputs import TokenClassifierOutput\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.nn.utils.rnn import pad_sequence\n\nfrom fastai.vision.all import *\nimport fastai","metadata":{"execution":{"iopub.status.busy":"2023-11-27T07:05:10.965287Z","iopub.execute_input":"2023-11-27T07:05:10.965749Z","iopub.status.idle":"2023-11-27T07:05:31.026268Z","shell.execute_reply.started":"2023-11-27T07:05:10.965714Z","shell.execute_reply":"2023-11-27T07:05:31.024658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_parquet('/kaggle/input/rna-dataset/train.parquet')\nprint(len(df))","metadata":{"execution":{"iopub.status.busy":"2023-11-27T07:12:51.008076Z","iopub.execute_input":"2023-11-27T07:12:51.011874Z","iopub.status.idle":"2023-11-27T07:13:01.129982Z","shell.execute_reply.started":"2023-11-27T07:12:51.011771Z","shell.execute_reply":"2023-11-27T07:13:01.128483Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_select = df.groupby(['sequence_id']).agg({'reads': 'max'})\ndf_select = df_select[df_select['reads'] > 100]\ndf_select\n\ndf = df[df['sequence_id'].isin(set(df_select.index.values))].reset_index(drop=True)\nlen(df)","metadata":{"execution":{"iopub.status.busy":"2023-11-27T07:13:01.133395Z","iopub.execute_input":"2023-11-27T07:13:01.133919Z","iopub.status.idle":"2023-11-27T07:13:08.328047Z","shell.execute_reply.started":"2023-11-27T07:13:01.133874Z","shell.execute_reply":"2023-11-27T07:13:08.325667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEBUG = False","metadata":{"execution":{"iopub.status.busy":"2023-11-27T07:13:08.330009Z","iopub.execute_input":"2023-11-27T07:13:08.330477Z","iopub.status.idle":"2023-11-27T07:13:08.340042Z","shell.execute_reply.started":"2023-11-27T07:13:08.330433Z","shell.execute_reply":"2023-11-27T07:13:08.338697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RNAConfig:\n    pad_token_id = 0\n    vocab_map = {\n        'PAD': pad_token_id, \n        ('G', '('): 1,\n        ('G', '.'): 2,\n        ('G', ')'): 3,\n        ('A', '('): 4,\n        ('A', '.'): 5,\n        ('A', ')'): 6,\n        ('C', '('): 7,\n        ('C', '.'): 8,\n        ('C', ')'): 9,\n        ('U', '('): 10,\n        ('U', '.'): 11,\n        ('U', ')'): 12,\n        'SINK_TOKEN': 13}\n    vocab_size = len(vocab_map)\n    hidden_size = 128\n    intermediate_size = 256\n    rms_norm_eps = 1e-6\n    num_attention_heads = 8\n    rope_theta = 10000.0\n    max_position_embeddings = 512\n    num_hidden_layers = 12\n    initializer_range = 0.02\n    attn_sink_with_pos_extension = True\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nPATH = Path(\"/kaggle/input/rna-dataset\")\n    \nconfig = RNAConfig()","metadata":{"execution":{"iopub.status.busy":"2023-11-27T07:13:08.343826Z","iopub.execute_input":"2023-11-27T07:13:08.344924Z","iopub.status.idle":"2023-11-27T07:13:08.359776Z","shell.execute_reply.started":"2023-11-27T07:13:08.344865Z","shell.execute_reply":"2023-11-27T07:13:08.358224Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not DEBUG:\n    from kaggle_secrets import UserSecretsClient\n    import wandb\n    user_secrets = UserSecretsClient()\n    secret_value_0 = user_secrets.get_secret(\"wandb\")\n\n    os.environ['WANDB_WATCH'] = 'all'\n    !wandb login $secret_value_0\n\n    wandb.init(\n        # set the wandb project where this run will be logged\n        project=\"Ribonanza RNA\",\n        name='mistral architecture',\n        notes=\"with filtering based on error < 1, reads filter, added layer norm to no decay parameters, initializer range 0.02, 20 epochs, 0.1 wd, structure, cosine scheduler\",\n    #     config=dict((name, getattr(config, name)) for name in dir(config) if not name.startswith('__')) \n    )","metadata":{"execution":{"iopub.status.busy":"2023-11-27T07:13:08.361732Z","iopub.execute_input":"2023-11-27T07:13:08.362734Z","iopub.status.idle":"2023-11-27T07:13:50.429410Z","shell.execute_reply.started":"2023-11-27T07:13:08.362684Z","shell.execute_reply":"2023-11-27T07:13:50.428361Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RNARotaryEmbedding(nn.Module):\n    def __init__(self, dim: int, max_position_embeddings: int, base: float):\n        super().__init__()\n        \n        self.dim = dim\n        self.max_position_embeddings = max_position_embeddings\n        self.base = base\n        inv_freq = 1.0 / (self.base ** (torch.arange(0, self.dim, 2).float().to(device) / self.dim))\n        self.register_buffer(\"inv_freq\", inv_freq)\n\n        self._set_cos_sin_cache(\n            seq_len=self.max_position_embeddings, device=self.inv_freq.device, dtype=torch.get_default_dtype()\n        )\n\n    def _set_cos_sin_cache(self, seq_len: int, device, dtype):\n        self.max_seq_len_cached = seq_len\n        t = torch.arange(self.max_seq_len_cached, device=device, dtype=self.inv_freq.dtype) # type: ignore\n\n        freqs = torch.einsum(\"i,j->ij\", t, self.inv_freq)\n        emb = torch.cat((freqs, freqs), dim=-1)\n        self.register_buffer(\"cos_cache\", emb.cos()[None, None, :, :].to(dtype), persistent=False)\n        self.register_buffer(\"sin_cache\", emb.sin()[None, None, :, :].to(dtype), persistent=False)\n\n    def forward(self, x, seq_len=None):\n        if seq_len and seq_len > self.max_seq_len_cached:\n            self._set_cos_sin_cache(seq_len=seq_len, device=x.device, dtype=x.dtype)\n\n        return (\n            self.cos_cache[:, :, :seq_len, ...].to(x.dtype), # type: ignore\n            self.sin_cache[:, :, :seq_len, ...].to(x.dtype), # type: ignore\n        )\n    \ndef rotate_half(x):\n    x1 = x[..., : x.shape[-1] // 2]\n    x2 = x[..., x.shape[-1] // 2 :]\n    return torch.cat((-x2, x1), dim=-1)\n    \ndef apply_rotary_pos_emb(q, k, cos, sin, position_idx):\n    cos = cos.squeeze(1).squeeze(0) # [seq_len, dim]\n    sin = sin.squeeze(1).squeeze(0) # [seq_len, dim]\n    cos = cos[position_idx].unsqueeze(1) # [bs, 1, seq_len, dim]\n    sin = sin[position_idx].unsqueeze(1) # [bs, 1, seq_len, dim]\n    q_embed = (q * cos) + (rotate_half(q) * sin)\n    k_embed = (k * cos) + (rotate_half(k) * sin)\n    return q_embed, k_embed\n\nclass RNAAttention(nn.Module):\n    def __init__(self, config: RNAConfig):\n        super().__init__()\n        self.config = config\n        self.hidden_size = config.hidden_size\n        self.num_heads = config.num_attention_heads\n        self.head_dim = self.hidden_size // self.num_heads\n        self.max_position_embeddings = config.max_position_embeddings\n        self.rope_theta = config.rope_theta\n\n        self.q_proj = nn.Linear(self.hidden_size, self.hidden_size, bias=False)\n        self.k_proj = nn.Linear(self.hidden_size, self.hidden_size, bias=False)\n        self.v_proj = nn.Linear(self.hidden_size, self.hidden_size, bias=False)\n        self.out_proj = nn.Linear(self.hidden_size, self.hidden_size, bias=False)\n\n        self.rotary_emb = RNARotaryEmbedding(\n            dim=self.head_dim,\n            max_position_embeddings=self.max_position_embeddings,\n            base=self.rope_theta,\n        )\n\n    def forward(\n        self,\n        hidden_states: torch.Tensor,\n        attention_mask: Optional[torch.Tensor] = None,\n        position_ids: Optional[torch.LongTensor] = None,\n        padding_mask: Optional[torch.Tensor] = None,\n    ) -> torch.Tensor:\n        bsz, q_len, _ = hidden_states.size()\n\n        query_states = self.q_proj(hidden_states)\n        key_states = self.k_proj(hidden_states)\n        value_states = self.v_proj(hidden_states)\n\n        query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)\n        key_states = key_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)\n        value_states = value_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)\n\n        cos, sin = self.rotary_emb(value_states, seq_len=position_ids.max() + 1)\n        query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin, position_ids)\n\n        attn_weights = torch.matmul(query_states, key_states.transpose(2, 3)) / math.sqrt(self.head_dim)\n        if attention_mask is not None:\n            attn_weights = attn_weights + attention_mask\n\n        attn_weights = F.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query_states.dtype)\n        attn_output = torch.matmul(attn_weights, value_states) # [bs, num_heads, q_len, head_dim]\n\n        attn_output = attn_output.transpose(1, 2).contiguous().view(bsz, q_len, self.hidden_size)\n        attn_output = self.out_proj(attn_output)\n\n        return attn_output\n\nclass RNAMLP(nn.Module):\n    def __init__(self, config: RNAConfig):\n        super().__init__()\n        self.config = config\n        self.hidden_size = config.hidden_size\n        self.intermediate_size = config.intermediate_size\n        self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size)\n        self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size)\n        self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size)\n        self.act_fn = nn.GELU()\n\n    def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:\n        gate = self.act_fn(self.gate_proj(hidden_states))\n        up = self.act_fn(self.up_proj(hidden_states))\n        down = self.down_proj(gate * up)\n\n        return down\n\nclass RNALayerNorm(nn.Module):\n    def __init__(self, config: RNAConfig, eps: float):\n        super().__init__()\n        self.weight = nn.Parameter(torch.ones(config.hidden_size))\n        self.variance_epsilon = eps\n\n    def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:\n        input_dtype = hidden_states.dtype\n        hidden_states = hidden_states.to(torch.float32)\n        variance = hidden_states.pow(2).mean(-1, keepdim=True)\n        hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)\n        return self.weight * hidden_states.to(input_dtype)\n\nclass RNADecoderLayer(nn.Module):\n    def __init__(self, config: RNAConfig):\n        super().__init__()\n        self.hidden_size = config.hidden_size\n        self.self_attn = RNAAttention(config)\n        self.mlp = RNAMLP(config)\n        self.input_layernorm = RNALayerNorm(config, eps=config.rms_norm_eps)\n        self.post_attention_layernorm = RNALayerNorm(config, eps=config.rms_norm_eps)\n\n    def forward(\n        self,\n        hidden_states: torch.Tensor,\n        attention_mask: Optional[torch.Tensor] = None,\n        position_ids: Optional[torch.LongTensor] = None,\n        padding_mask: Optional[torch.Tensor] = None,\n    ) -> torch.Tensor:\n        \n        residual = hidden_states\n\n        hidden_states = self.input_layernorm(hidden_states)\n\n        hidden_states = self.self_attn(\n            hidden_states,\n            attention_mask,\n            position_ids,\n            padding_mask\n        )\n        hidden_states = residual + hidden_states\n\n        residual = hidden_states\n        hidden_states = self.post_attention_layernorm(hidden_states)\n        hidden_states = self.mlp(hidden_states)\n        outputs = residual + hidden_states\n\n        return outputs\n\n\nclass RNATransformer(nn.Module):\n    def __init__(self, config: RNAConfig):\n        super().__init__()\n        self.config = config\n        self.padding_idx = config.pad_token_id\n        self.vocab_size = config.vocab_size\n\n        self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)\n        self.layers = nn.ModuleList([RNADecoderLayer(config) for _ in range(config.num_hidden_layers)])\n        self.norm = RNALayerNorm(config, eps=config.rms_norm_eps)\n\n        self.gradient_checkpointing = False\n#         self.emb_layer_norm = RNALayerNorm(config, eps=config.rms_norm_eps)\n        # Initialize weights and apply final processing\n\n    def forward(\n        self,\n        input_ids: Optional[torch.LongTensor] = None,\n        attention_mask: Optional[torch.Tensor] = None,\n        position_ids: Optional[torch.LongTensor] = None,\n        inputs_embeds: Optional[torch.FloatTensor] = None,\n        output_hidden_states: Optional[bool] = None,\n    ): \n        if input_ids is not None:\n            batch_size, seq_length = input_ids.shape\n        elif inputs_embeds is not None:\n            batch_size, seq_length, _ = inputs_embeds.shape\n        else:\n            raise ValueError(\"You have to specify either input_ids or inputs_embeds\")\n        \n        if position_ids is None:\n            device = input_ids.device if input_ids is not None else inputs_embeds.device # type: ignore\n            position_ids = torch.arange(seq_length, dtype=torch.long, device=device) # type: ignore\n            position_ids = position_ids.unsqueeze(0).view(-1, seq_length) # type: ignore\n        else:\n            position_ids = position_ids.view(-1, seq_length).long() # type: ignore\n        if self.config.attn_sink_with_pos_extension and self.training:\n            bs = position_ids.size()[0]\n            shift = torch.randint(high=200, size=(bs, ), dtype=position_ids.dtype, device=position_ids.device).view(bs, -1)\n            position_ids[:, 1:] = position_ids[:, 1:] + shift            \n\n        if inputs_embeds is None:\n            inputs_embeds = self.embed_tokens(input_ids)\n#         inputs_embeds = self.emb_layer_norm(inputs_embeds)\n\n        attention_mask = attention_mask.unsqueeze(1).unsqueeze(2).log() # type: ignore\n\n        hidden_states = inputs_embeds\n\n        all_hidden_states = () if output_hidden_states else None\n\n        for layer in self.layers:\n            if output_hidden_states:\n                all_hidden_states += (hidden_states,) # type: ignore\n\n            hidden_states = layer(\n                hidden_states,\n                attention_mask,\n                position_ids\n            )\n\n        hidden_states = self.norm(hidden_states)\n\n        if output_hidden_states:\n            all_hidden_states += (hidden_states,) # type: ignore\n            return all_hidden_states\n        else:\n            return hidden_states\n\ndef loss_fn(outputs: torch.Tensor, targets: torch.Tensor) -> torch.Tensor:\n    outputs = outputs[~targets.isnan()]\n    targets = targets[~targets.isnan()].clip(0,1)\n    return F.l1_loss(outputs, targets)\n\nclass RNAModel(nn.Module):\n    def __init__(self, config: RNAConfig):\n        super().__init__()\n        self._config = config\n        self.backbone = RNATransformer(config)\n        self.head = nn.Linear(config.hidden_size, 2)\n        self.post_init()\n\n    def forward(\n        self,\n        input_ids: Optional[torch.LongTensor] = None,\n        attention_mask: Optional[torch.Tensor] = None,\n        position_ids: Optional[torch.LongTensor] = None,\n        inputs_embeds: Optional[torch.FloatTensor] = None,\n        output_hidden_states: Optional[bool] = None,\n        labels: Optional[torch.FloatTensor] = None,\n    ) -> TokenClassifierOutput:\n        \n        outputs = self.backbone(\n            input_ids,\n            attention_mask,\n            position_ids,\n            inputs_embeds,\n            output_hidden_states\n        )\n        predictions = self.head(outputs)\n\n        loss = None\n        if labels is not None:\n            loss = loss_fn(predictions, labels)\n\n        return TokenClassifierOutput(\n            loss=loss, # type: ignore\n            logits=predictions,\n        )\n    \n        self.post_init()\n    \n    def post_init(self):\n        self.apply(self._init_weights)\n        \n    def _init_weights(self, module):\n        if isinstance(module, (nn.Linear, nn.Embedding)):\n            # Slightly different from the TF version which uses truncated_normal for initialization\n            # cf https://github.com/pytorch/pytorch/pull/5617\n            module.weight.data.normal_(mean=0.0, std=self._config.initializer_range)\n        if isinstance(module, nn.Linear) and module.bias is not None:\n            module.bias.data.zero_()\n","metadata":{"execution":{"iopub.status.busy":"2023-11-27T07:13:50.434632Z","iopub.execute_input":"2023-11-27T07:13:50.435237Z","iopub.status.idle":"2023-11-27T07:13:50.512012Z","shell.execute_reply.started":"2023-11-27T07:13:50.435203Z","shell.execute_reply":"2023-11-27T07:13:50.510778Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ALL_LAYERNORM_LAYERS = [nn.LayerNorm, RNALayerNorm]\n\ndef get_parameter_names(model, forbidden_layer_types):\n    \"\"\"\n    Returns the names of the model parameters that are not inside a forbidden layer.\n    \"\"\"\n    result = []\n    for name, child in model.named_children():\n        result += [\n            f\"{name}.{n}\"\n            for n in get_parameter_names(child, forbidden_layer_types)\n            if not isinstance(child, tuple(forbidden_layer_types))\n        ]\n    # Add model specific parameters (defined with nn.Parameter) since they are not in any child.\n    result += list(model._parameters.keys())\n    return result\n\nclass CustomTrainer(Trainer):\n    def get_decay_parameter_names(self, model):\n        \"\"\"\n        Get all parameter names that weight decay will be applied to\n\n        Note that some models implement their own layernorm instead of calling nn.LayerNorm, weight decay could still\n        apply to those modules since this function only filter out instance of nn.LayerNorm\n        \"\"\"\n        decay_parameters = get_parameter_names(model, ALL_LAYERNORM_LAYERS)\n        decay_parameters = [name for name in decay_parameters if \"bias\" not in name]\n        return decay_parameters\n    \n    def create_optimizer(self):\n        \"\"\"\n        Setup the optimizer.\n\n        We provide a reasonable default that works well. If you want to use something else, you can pass a tuple in the\n        Trainer's init through `optimizers`, or subclass and override this method in a subclass.\n        \"\"\"\n        opt_model = self.model\n\n        if self.optimizer is None:\n            decay_parameters = self.get_decay_parameter_names(opt_model)\n            optimizer_grouped_parameters = [\n                {\n                    \"params\": [\n                        p for n, p in opt_model.named_parameters() if (n in decay_parameters and p.requires_grad)\n                    ],\n                    \"weight_decay\": self.args.weight_decay,\n                    'lr': self.args.learning_rate,\n                },\n                {\n                    \"params\": [\n                        p for n, p in opt_model.named_parameters() if (n not in decay_parameters and p.requires_grad)\n                    ],\n                    \"weight_decay\": 0.0,\n                    'lr': self.args.learning_rate,\n                },\n            ]\n\n            optimizer_cls, optimizer_kwargs = Trainer.get_optimizer_cls_and_kwargs(self.args)\n\n            self.optimizer = optimizer_cls(optimizer_grouped_parameters, **optimizer_kwargs)\n\n        return self.optimizer","metadata":{"execution":{"iopub.status.busy":"2023-11-27T07:13:50.513972Z","iopub.execute_input":"2023-11-27T07:13:50.514825Z","iopub.status.idle":"2023-11-27T07:13:50.535893Z","shell.execute_reply.started":"2023-11-27T07:13:50.514782Z","shell.execute_reply":"2023-11-27T07:13:50.534563Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_parquet('/kaggle/input/rna-dataset/train.parquet')\n\nlabel_cols = [c for c in df.columns if c.startswith('reactivity_0')]\nerror_cols = [c for c in df.columns if c.startswith('reactivity_error_0')]\n\ndef prepare_targets(df):\n    targets = df[label_cols].values.astype(np.float32)\n    mask = df[error_cols].values.astype(np.float32) > 1.0\n    targets[mask] = np.nan\n    mask2 = df['reads'] < 100\n    targets[mask2] = np.nan\n    return targets\n    \nclass RibonanzaDatasetTrain():\n    def __init__(self, df, config, is_train):\n        self.config = config\n        self.is_train = is_train\n        df = df[df.is_train == is_train].reset_index(drop=True)\n        df_DMS_MaP = df[df['experiment_type'] == 'DMS_MaP'].sort_values(['sequence_id', 'dataset_name']).reset_index(drop=True)\n        df_2A3_MaP = df[df['experiment_type'] == '2A3_MaP'].sort_values(['sequence_id', 'dataset_name']).reset_index(drop=True)\n        self.sequences = df_DMS_MaP[['sequence', 'structure']].apply(lambda x: [config.vocab_map[c] for c in zip(x['sequence'], x['structure'])], axis=1).values\n        self.targets_DMS_MaP = prepare_targets(df_DMS_MaP)\n        self.targets_2A3_MaP = prepare_targets(df_2A3_MaP)\n        \n        if config.attn_sink_with_pos_extension:\n            x = np.full((self.targets_DMS_MaP.shape[0], 1), np.nan)\n            self.targets_DMS_MaP = np.concatenate((x, self.targets_DMS_MaP), axis=1)\n            self.targets_2A3_MaP = np.concatenate((x, self.targets_2A3_MaP), axis=1)\n#             x = np.full((self.sequences.shape[0], 1), config.vocab_map['SINK_TOKEN'], dtype=np.int64)\n#             self.sequences = np.concatenate((x, self.sequences), axis=1)\n        \n    def __len__(self):\n        return len(self.sequences)\n    \n    def __getitem__(self, idx):\n        \n        seq = self.sequences[idx]\n        \n        if config.attn_sink_with_pos_extension:\n            seq = [config.vocab_map['SINK_TOKEN'], *seq]\n\n        outputs = {\n            'input_ids': torch.tensor(seq, dtype=torch.long),\n            'attention_mask': torch.ones(len(seq), dtype=torch.float),\n        }\n        outputs['labels'] = torch.from_numpy(np.stack(\n                    (self.targets_DMS_MaP[idx][:len(seq)], \n                     self.targets_2A3_MaP[idx][:len(seq)]), axis=1\n                ))\n        \n        return outputs\n    \ndef collate_fn(batch):\n    new_batch = dict()\n    for k in batch[0].keys():\n        new_batch[k] = pad_sequence((i[k] for i in batch), batch_first=True, padding_value=0)\n        \n    return new_batch","metadata":{"execution":{"iopub.status.busy":"2023-11-27T07:13:50.538420Z","iopub.execute_input":"2023-11-27T07:13:50.539656Z","iopub.status.idle":"2023-11-27T07:13:59.593921Z","shell.execute_reply.started":"2023-11-27T07:13:50.539590Z","shell.execute_reply":"2023-11-27T07:13:59.592852Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = RNAModel(config)\n\ntrain_ds = RibonanzaDatasetTrain(df, config, is_train=True)\nvalid_ds = RibonanzaDatasetTrain(df, config, is_train=False)","metadata":{"execution":{"iopub.status.busy":"2023-11-27T07:13:59.598948Z","iopub.execute_input":"2023-11-27T07:13:59.600990Z","iopub.status.idle":"2023-11-27T07:15:21.031308Z","shell.execute_reply.started":"2023-11-27T07:13:59.600936Z","shell.execute_reply":"2023-11-27T07:15:21.030204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_args = TrainingArguments(\n    output_dir=\"trainer\",\n    evaluation_strategy=\"epoch\",\n    save_strategy=\"epoch\",\n    logging_strategy=\"steps\",\n    logging_steps=int(1),\n    learning_rate=1e-3,\n    adam_beta2=0.98,\n    num_train_epochs=20,\n    weight_decay=0.1,\n    per_device_train_batch_size=256, \n    per_device_eval_batch_size=512,\n    load_best_model_at_end=True,\n    save_total_limit=1,\n    report_to='wandb' if not DEBUG else 'none',\n    dataloader_num_workers=2,\n    lr_scheduler_type='cosine',\n    warmup_ratio=0.2\n)","metadata":{"execution":{"iopub.status.busy":"2023-11-27T07:15:21.035272Z","iopub.execute_input":"2023-11-27T07:15:21.036844Z","iopub.status.idle":"2023-11-27T07:15:21.054425Z","shell.execute_reply.started":"2023-11-27T07:15:21.036788Z","shell.execute_reply":"2023-11-27T07:15:21.052757Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = CustomTrainer(\n    args=training_args, \n    model=model, \n    data_collator=collate_fn, \n    train_dataset=train_ds, \n    eval_dataset=valid_ds\n)\ntrainer.train()","metadata":{"execution":{"iopub.status.busy":"2023-11-27T07:15:21.056667Z","iopub.execute_input":"2023-11-27T07:15:21.057822Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.state_dict(), 'last_epoch_model.bin')","metadata":{"execution":{"iopub.status.busy":"2023-11-26T21:32:02.835127Z","iopub.status.idle":"2023-11-26T21:32:02.835611Z","shell.execute_reply.started":"2023-11-26T21:32:02.835376Z","shell.execute_reply":"2023-11-26T21:32:02.835402Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}