{"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":6888177,"sourceType":"datasetVersion","datasetId":3889652},{"sourceId":7073078,"sourceType":"datasetVersion","datasetId":4073585}],"dockerImageVersionId":30559,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"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-29T20:03:17.254870Z","iopub.execute_input":"2023-11-29T20:03:17.255426Z","iopub.status.idle":"2023-11-29T20:03:17.260092Z","shell.execute_reply.started":"2023-11-29T20:03:17.255396Z","shell.execute_reply":"2023-11-29T20:03:17.259171Z"},"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-29T20:03:17.276747Z","iopub.execute_input":"2023-11-29T20:03:17.277032Z","iopub.status.idle":"2023-11-29T20:03:32.784592Z","shell.execute_reply.started":"2023-11-29T20:03:17.277006Z","shell.execute_reply":"2023-11-29T20:03:32.783728Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEBUG = False","metadata":{"execution":{"iopub.status.busy":"2023-11-29T20:03:32.786065Z","iopub.execute_input":"2023-11-29T20:03:32.786686Z","iopub.status.idle":"2023-11-29T20:03:32.790576Z","shell.execute_reply.started":"2023-11-29T20:03:32.786647Z","shell.execute_reply":"2023-11-29T20:03:32.789676Z"},"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-29T20:03:32.791507Z","iopub.execute_input":"2023-11-29T20:03:32.791768Z","iopub.status.idle":"2023-11-29T20:03:32.828140Z","shell.execute_reply.started":"2023-11-29T20:03:32.791743Z","shell.execute_reply":"2023-11-29T20:03:32.827291Z"},"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, added layer norm to no decay parameters, initializer range 0.02, 20 epochs, 0.1 wd, with errors scaling, structure\",\n    #     config=dict((name, getattr(config, name)) for name in dir(config) if not name.startswith('__')) \n    )","metadata":{"execution":{"iopub.status.busy":"2023-11-29T20:03:32.830221Z","iopub.execute_input":"2023-11-29T20:03:32.830592Z","iopub.status.idle":"2023-11-29T20:04:11.408668Z","shell.execute_reply.started":"2023-11-29T20:03:32.830563Z","shell.execute_reply":"2023-11-29T20:04:11.407710Z"},"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        # 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\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-29T20:04:11.410221Z","iopub.execute_input":"2023-11-29T20:04:11.410528Z","iopub.status.idle":"2023-11-29T20:04:11.464393Z","shell.execute_reply.started":"2023-11-29T20:04:11.410500Z","shell.execute_reply":"2023-11-29T20:04:11.463269Z"},"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-29T20:04:11.465899Z","iopub.execute_input":"2023-11-29T20:04:11.466250Z","iopub.status.idle":"2023-11-29T20:04:11.491409Z","shell.execute_reply.started":"2023-11-29T20:04:11.466205Z","shell.execute_reply":"2023-11-29T20:04:11.490444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_parquet('/kaggle/input/rna-dataset/train.parquet')\n\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    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').reset_index(drop=True)\n        df_2A3_MaP = df[df['experiment_type'] == '2A3_MaP'].sort_values('sequence_id').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-29T20:04:11.492808Z","iopub.execute_input":"2023-11-29T20:04:11.493124Z","iopub.status.idle":"2023-11-29T20:04:25.567267Z","shell.execute_reply.started":"2023-11-29T20:04:11.493093Z","shell.execute_reply":"2023-11-29T20:04:25.566177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = RNAModel(config)","metadata":{"execution":{"iopub.status.busy":"2023-11-29T20:04:25.568565Z","iopub.execute_input":"2023-11-29T20:04:25.568954Z","iopub.status.idle":"2023-11-29T20:04:31.097642Z","shell.execute_reply.started":"2023-11-29T20:04:25.568917Z","shell.execute_reply":"2023-11-29T20:04:31.096605Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ntrain_ds = RibonanzaDatasetTrain(df, config, is_train=True)\nvalid_ds = RibonanzaDatasetTrain(df, config, is_train=False)","metadata":{"execution":{"iopub.status.busy":"2023-11-29T20:04:31.099004Z","iopub.execute_input":"2023-11-29T20:04:31.099377Z","iopub.status.idle":"2023-11-29T20:05:28.629713Z","shell.execute_reply.started":"2023-11-29T20:04:31.099342Z","shell.execute_reply":"2023-11-29T20:05:28.628565Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds[100]","metadata":{"execution":{"iopub.status.busy":"2023-11-29T20:05:28.635393Z","iopub.execute_input":"2023-11-29T20:05:28.636104Z","iopub.status.idle":"2023-11-29T20:05:28.689965Z","shell.execute_reply.started":"2023-11-29T20:05:28.636063Z","shell.execute_reply":"2023-11-29T20:05:28.688780Z"},"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='constant_with_warmup',\n    warmup_ratio=0.2\n)","metadata":{"execution":{"iopub.status.busy":"2023-11-29T20:05:28.691304Z","iopub.execute_input":"2023-11-29T20:05:28.691646Z","iopub.status.idle":"2023-11-29T20:05:28.700997Z","shell.execute_reply.started":"2023-11-29T20:05:28.691611Z","shell.execute_reply":"2023-11-29T20:05:28.699945Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","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-28T10:38:26.201013Z","iopub.status.idle":"2023-11-28T10:38:26.201944Z","shell.execute_reply.started":"2023-11-28T10:38:26.201662Z","shell.execute_reply":"2023-11-28T10:38:26.201689Z"}}},{"cell_type":"markdown","source":"torch.save(model.state_dict(), 'last_epoch_model.bin')","metadata":{"execution":{"iopub.status.busy":"2023-11-28T10:38:26.203462Z","iopub.status.idle":"2023-11-28T10:38:26.204264Z","shell.execute_reply.started":"2023-11-28T10:38:26.204016Z","shell.execute_reply":"2023-11-28T10:38:26.204040Z"}}},{"cell_type":"code","source":"print(df.columns)","metadata":{"execution":{"iopub.status.busy":"2023-11-29T20:05:28.702400Z","iopub.execute_input":"2023-11-29T20:05:28.702751Z","iopub.status.idle":"2023-11-29T20:05:28.710658Z","shell.execute_reply.started":"2023-11-29T20:05:28.702717Z","shell.execute_reply":"2023-11-29T20:05:28.709610Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df  = pd.read_parquet('/kaggle/input/rna-dataset/test.parquet')","metadata":{"execution":{"iopub.status.busy":"2023-11-29T20:05:28.712140Z","iopub.execute_input":"2023-11-29T20:05:28.712552Z","iopub.status.idle":"2023-11-29T20:05:34.375693Z","shell.execute_reply.started":"2023-11-29T20:05:28.712514Z","shell.execute_reply":"2023-11-29T20:05:34.374598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RibonanzaDatasetTest():\n    def __init__(self, df, config):\n        self.config = config\n        self.df = df.sort_values('sequence_id').reset_index(drop=True)\n        self.sequences = self.df[['sequence', 'structure']].apply(\n            lambda x: [config.vocab_map[c] for c in zip(x['sequence'], x['structure'])], axis=1\n        ).values\n\n    def __len__(self):\n        return len(self.sequences)\n\n    def __getitem__(self, idx):\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\n        return outputs\n","metadata":{"execution":{"iopub.status.busy":"2023-11-29T20:05:34.377123Z","iopub.execute_input":"2023-11-29T20:05:34.377534Z","iopub.status.idle":"2023-11-29T20:05:34.388903Z","shell.execute_reply.started":"2023-11-29T20:05:34.377497Z","shell.execute_reply":"2023-11-29T20:05:34.387846Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_ds = RibonanzaDatasetTest(df, config)","metadata":{"execution":{"iopub.status.busy":"2023-11-29T20:05:34.390208Z","iopub.execute_input":"2023-11-29T20:05:34.390787Z","iopub.status.idle":"2023-11-29T20:06:51.820001Z","shell.execute_reply.started":"2023-11-29T20:05:34.390753Z","shell.execute_reply":"2023-11-29T20:06:51.818817Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def collate_fn_test(batch):\n    \"\"\"\n    Collate function for the test dataset without padding.\n    \"\"\"\n    new_batch = dict()\n    for k in batch[0].keys():\n        new_batch[k] = [i[k] for i in batch]\n\n    return new_batch\n","metadata":{"execution":{"iopub.status.busy":"2023-11-29T20:06:51.821305Z","iopub.execute_input":"2023-11-29T20:06:51.821648Z","iopub.status.idle":"2023-11-29T20:06:51.828590Z","shell.execute_reply.started":"2023-11-29T20:06:51.821605Z","shell.execute_reply":"2023-11-29T20:06:51.827414Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataloader = DataLoader(test_ds, batch_size=512, collate_fn=collate_fn_test)\n\n# Put the model in evaluation mode\nmodel.eval()\n","metadata":{"execution":{"iopub.status.busy":"2023-11-29T20:06:51.830063Z","iopub.execute_input":"2023-11-29T20:06:51.830858Z","iopub.status.idle":"2023-11-29T20:06:51.844895Z","shell.execute_reply.started":"2023-11-29T20:06:51.830830Z","shell.execute_reply":"2023-11-29T20:06:51.843898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(df.columns)","metadata":{"execution":{"iopub.status.busy":"2023-11-29T20:06:51.846007Z","iopub.execute_input":"2023-11-29T20:06:51.846347Z","iopub.status.idle":"2023-11-29T20:06:51.855916Z","shell.execute_reply.started":"2023-11-29T20:06:51.846308Z","shell.execute_reply":"2023-11-29T20:06:51.854771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"checkpoint = torch.load('/kaggle/input/nodekf/last_epoch_model.bin')\nmodel.load_state_dict(checkpoint)\n\n# Step 4: Put the model in evaluation mode\nmodel.eval()","metadata":{"execution":{"iopub.status.busy":"2023-11-29T20:06:51.856970Z","iopub.execute_input":"2023-11-29T20:06:51.857276Z","iopub.status.idle":"2023-11-29T20:06:52.110179Z","shell.execute_reply.started":"2023-11-29T20:06:51.857246Z","shell.execute_reply":"2023-11-29T20:06:52.109061Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.eval()\nall_predictions = []","metadata":{"execution":{"iopub.status.busy":"2023-11-29T20:06:52.111589Z","iopub.execute_input":"2023-11-29T20:06:52.111978Z","iopub.status.idle":"2023-11-29T20:06:52.119188Z","shell.execute_reply.started":"2023-11-29T20:06:52.111941Z","shell.execute_reply":"2023-11-29T20:06:52.118194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def collate_fn_test(batch):\n    \"\"\"\n    Collate function for the test dataset without padding.\n    \"\"\"\n    new_batch = dict()\n    for k in batch[0].keys():\n        new_batch[k] = [i[k] for i in batch]\n\n    return new_batch\n","metadata":{"execution":{"iopub.status.busy":"2023-11-29T20:06:52.121856Z","iopub.execute_input":"2023-11-29T20:06:52.122202Z","iopub.status.idle":"2023-11-29T20:06:52.131550Z","shell.execute_reply.started":"2023-11-29T20:06:52.122176Z","shell.execute_reply":"2023-11-29T20:06:52.130531Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RibonanzaDatasetTest():\n    def __init__(self, df, config):\n        self.config = config\n        df = df.sort_values('sequence_id').reset_index(drop=True)\n        self.sequences = df[['sequence', 'structure']].apply(lambda x: [config.vocab_map[c] for c in zip(x['sequence'], x['structure'])], axis=1).values\n\n    def __len__(self):\n        return len(self.sequences)\n\n    def __getitem__(self, idx):\n        seq = self.sequences[idx]\n\n        if self.config.attn_sink_with_pos_extension:\n            seq = [self.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\n        return outputs\n\n   \n\n","metadata":{"execution":{"iopub.status.busy":"2023-11-29T20:06:52.132914Z","iopub.execute_input":"2023-11-29T20:06:52.134131Z","iopub.status.idle":"2023-11-29T20:06:52.147469Z","shell.execute_reply.started":"2023-11-29T20:06:52.134091Z","shell.execute_reply":"2023-11-29T20:06:52.146047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = RibonanzaDatasetTest(df, config)\ntest_dataloader = DataLoader(test_dataset, batch_size=512, collate_fn=collate_fn)\n","metadata":{"execution":{"iopub.status.busy":"2023-11-29T20:06:52.148697Z","iopub.execute_input":"2023-11-29T20:06:52.149579Z","iopub.status.idle":"2023-11-29T20:08:10.864657Z","shell.execute_reply.started":"2023-11-29T20:06:52.149546Z","shell.execute_reply":"2023-11-29T20:08:10.863648Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\n# Now you can iterate over the test set and perform inference\nmodel.eval()\n# Make predictions on the test dataset\n# Make predictions on the test dataset\nall_predictions = []\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")  # Get the available device\n# Move model to the appropriate device\nmodel.to(device)\n\n# Move config to the appropriate device if necessary\n\nwith torch.no_grad():\n    a=1\n    for sequence in tqdm(test_ds, desc=\"Inference Progress\", unit=\"batch\"):\n        input_ids = sequence['input_ids'].unsqueeze(0).to(device)  # Add a batch dimension and move to the device\n        attention_mask = sequence['attention_mask'].unsqueeze(0).to(device)  # Add a batch dimension and move to the device\n\n        # Forward pass\n        predictions = model(input_ids=input_ids, attention_mask=attention_mask)\n        if(a==1):\n            print(predictions)\n        a=2\n        # Process predictions as needed and append to the list\n        # For example, if the model returns a tuple with the first element being the predicted labels\n        all_predictions.append(predictions[0].cpu().numpy())\n        \n\n\n# Concatenate predictions from all batches\nall_predictions = np.concatenate(all_predictions, axis=0)\n\n# Now 'all_predictions' contains the model predictions for the test set\n","metadata":{"execution":{"iopub.status.busy":"2023-11-29T20:17:26.458859Z","iopub.execute_input":"2023-11-29T20:17:26.459211Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Flatten the 3D array to 2D\nflat_predictions = all_predictions.reshape(-1, all_predictions.shape[-1])\n\n# Create DataFrame\nsubmission = pd.DataFrame({\n    'id': np.arange(0, len(flat_predictions), 1),\n    'reactivity_DMS_MaP': flat_predictions[:, 1],\n    'reactivity_2A3_MaP': flat_predictions[:, 0]\n})\n\nsubmission.to_csv('submission.csv', index=False)\nsubmission\n","metadata":{"execution":{"iopub.status.busy":"2023-11-29T20:14:08.339195Z","iopub.execute_input":"2023-11-29T20:14:08.340131Z","iopub.status.idle":"2023-11-29T20:14:08.364572Z","shell.execute_reply.started":"2023-11-29T20:14:08.340088Z","shell.execute_reply":"2023-11-29T20:14:08.363456Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"execution":{"iopub.status.busy":"2023-11-29T20:12:33.384290Z","iopub.execute_input":"2023-11-29T20:12:33.384732Z","iopub.status.idle":"2023-11-29T20:12:33.545596Z","shell.execute_reply.started":"2023-11-29T20:12:33.384699Z","shell.execute_reply":"2023-11-29T20:12:33.544170Z"},"trusted":true},"execution_count":null,"outputs":[]}]}