{"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":87793,"databundleVersionId":11553390,"sourceType":"competition"},{"sourceId":11447883,"sourceType":"datasetVersion","datasetId":7172166}],"dockerImageVersionId":30918,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Introduction\nI'd first like to extend my gratitude to Stanford for organizing these inspiring and challenging competitions!\n\nMy name is Faris, and I'm an AI engineer deeply passionate about tackling complex problems through innovative AI solutions. I was particularly drawn to this competition due to its intriguing difficulty and the unique challenges it presents.\n\n### Competition Goal\n\nThe primary objective of this competition is to accurately predict the coordinates of each nucleoid within a given sequence. Each sequence contains five distinct output structures, adding an exciting layer of complexity to the task.\n\n### My Experemnts\n| Sequence vectorization | AI model | Loss functino| Output Standraization |Augmentaion | Status| CV | LB\n| -- | -- | --|--|--|--|--|--|\n| Embedding | Transformers | DMSE | Z-score| False |finished| - | 0.137 |\n| Embedding + PSSM + Coveriance matrix + Frequency matrix| Transformers | dRMAE + Torsion + SVD | None| True | finished| - | - |\n| Embedding + PSSM + Coveriance matrix + Frequency matrix| Transformers | dRMAE + Torsion + SVD | None| False | finished| - | - |\n| Embedding | Transformers | FAPE | Z-score| False| in progress| - | - |\n| Embedding | Transformers | DMSE + FAPE | Z-score| False| in progress| - | - |\n| Embedding + PSSM | Transformers | DMSE | Z-score| False |not started| - | - |\n| Embedding + MSA Covariance Features | Transformers | DMSE| Z-score| False |not started| - | - |\n| MSA Covariance Features | GNN | DMSE| Z-score| False|not started| - | - |","metadata":{}},{"cell_type":"code","source":"%%capture\n! pip install /kaggle/input/bio-whl/biopython-1.85-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:03:10.315321Z","iopub.execute_input":"2025-04-18T18:03:10.315659Z","iopub.status.idle":"2025-04-18T18:03:15.742328Z","shell.execute_reply.started":"2025-04-18T18:03:10.315622Z","shell.execute_reply":"2025-04-18T18:03:15.741128Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nfrom Bio import SeqIO\nimport torch\nfrom torch.utils.data import Dataset\nimport math\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom einops import rearrange\nfrom os import rmdir\nfrom torch.utils.data import DataLoader\nimport pandas as pd","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:03:15.743776Z","iopub.execute_input":"2025-04-18T18:03:15.744077Z","iopub.status.idle":"2025-04-18T18:03:19.263947Z","shell.execute_reply.started":"2025-04-18T18:03:15.744050Z","shell.execute_reply":"2025-04-18T18:03:19.262836Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class PositionEmbeddingLayer(nn.Module):\n    def __init__(self, sequence_length, vocab_size, output_dim):\n        super(PositionEmbeddingLayer, self).__init__()\n        self.embedding_layer = nn.Embedding(vocab_size, output_dim, padding_idx=0)\n        self.position_embedding_layer = nn.Embedding(sequence_length, output_dim)\n\n    def forward(self, inputs):\n        # inputs shape: (batch, seq_length)\n        seq_length = inputs.size(1)\n        positions = torch.arange(seq_length, device=inputs.device).unsqueeze(0)  # (1, seq_length)\n        embedded = self.embedding_layer(inputs)                            # (batch, seq_length, output_dim)\n        pos_emb = self.position_embedding_layer(positions)                 # (1, seq_length, output_dim)\n        return embedded + pos_emb\n\nclass DotProductAttention(nn.Module):\n    def __init__(self):\n        super(DotProductAttention, self).__init__()\n\n    def forward(self, queries, keys, values, mask=None, covariance=None):\n        # queries: (batch, heads, seq_len, head_dim)\n        scale = math.sqrt(queries.size(-1))\n        scores = torch.matmul(queries, keys.transpose(-2, -1)) / scale  # (batch, heads, seq_len, seq_len)\n        if covariance is not None:\n            scores = scores + covariance.unsqueeze(1)  # make sure covariance is broadcastable\n        if mask is not None:\n            scores = scores * mask\n        weights = torch.softmax(scores, dim=-1)\n        return torch.matmul(weights, values)\n\nclass MultiHeadAttention(nn.Module):\n    def __init__(self, h, d_k, d_v, d_model):\n        \"\"\"\n        h: number of heads\n        d_k: output dimension for query and key projection (assumed to be divisible by h)\n        d_v: output dimension for value projection (similarly, heads × (d_v/h))\n        d_model: final output dimension (from W_o)\n        \"\"\"\n        super(MultiHeadAttention, self).__init__()\n        self.heads = h\n        self.d_k = d_k  # note: d_k should be chosen so that head_dim = d_k // h\n        self.W_q = nn.Linear(d_model, d_k)\n        self.W_k = nn.Linear(d_model, d_k)\n        self.W_v = nn.Linear(d_model, d_v)\n        self.W_o = nn.Linear(d_v, d_model)\n        self.attention = DotProductAttention()\n\n    def reshape_tensor(self, x, flag):\n        # x: (batch, seq_len, d) where d will be split over heads\n\n        if flag:\n            batch, seq_len, d = x.size()\n            head_dim = d // self.heads\n            # reshape to (batch, seq_len, heads, head_dim) then permute to (batch, heads, seq_len, head_dim)\n            return x.view(batch, seq_len, self.heads, head_dim).permute(0, 2, 1, 3)\n        else:\n            batch, h,seq_len, d = x.size()\n            # reverse: from (batch, heads, seq_len, head_dim) to (batch, seq_len, d)\n            return x.permute(0, 2, 1, 3).contiguous().view(batch, seq_len, d*h)\n\n    def forward(self, query, key, value, attention_mask=None, covariance=None):\n        # Linear projections: assume input shape (batch, seq_len, d_model)\n        q = self.W_q(query)  # (batch, seq_len, d_k)\n        k = self.W_k(key)    # (batch, seq_len, d_k)\n        v = self.W_v(value)  # (batch, seq_len, d_v)\n\n        # Reshape for multi-head attention.\n        q_reshaped = self.reshape_tensor(q, flag=True)  # (batch, heads, seq_len, d_k/head)\n        k_reshaped = self.reshape_tensor(k, flag=True)\n        v_reshaped = self.reshape_tensor(v, flag=True)\n\n        # Compute attention output.\n        out = self.attention(q_reshaped, k_reshaped, v_reshaped, mask=attention_mask, covariance=covariance)\n        # Reverse reshape: (batch, seq_len, d_v)\n        output = self.reshape_tensor(out, flag=False)\n        return self.W_o(output)\n\nclass AddNormalization(nn.Module):\n    def __init__(self, d_model):\n        super(AddNormalization, self).__init__()\n        self.layer_norm = nn.LayerNorm(d_model)\n\n    def forward(self, x, sublayer_x):\n        return self.layer_norm(x + sublayer_x)\n\nclass FeedForward(nn.Module):\n    def __init__(self, d_ff, d_model):\n        super(FeedForward, self).__init__()\n        self.fc1 = nn.Linear(d_model, d_ff)\n        self.fc2 = nn.Linear(d_ff, d_model)\n        self.activation = nn.SiLU()  # swish activation\n\n    def forward(self, x):\n        return self.fc2(self.activation(self.fc1(x)))\n\nclass EncoderLayer(nn.Module):\n    def __init__(self, h, d_k, d_v, d_model, d_ff, dropout_rate):\n        super(EncoderLayer, self).__init__()\n        self.multihead_attention = MultiHeadAttention(h, d_k, d_v, d_model)\n        self.dropout1 = nn.Dropout(dropout_rate)\n        self.add_norm1 = AddNormalization(d_model)\n        self.feed_forward = FeedForward(d_ff, d_model)\n        self.dropout2 = nn.Dropout(dropout_rate)\n        self.add_norm2 = AddNormalization(d_model)\n\n    def forward(self, x, covariance, padding_mask):\n        # x: (batch, seq_len, d_model)\n        mha_out = self.multihead_attention(x, x, x, attention_mask=padding_mask, covariance=covariance)\n        mha_out = self.dropout1(mha_out)\n        addnorm_out = self.add_norm1(x, mha_out)\n        ff_out = self.feed_forward(addnorm_out)\n        ff_out = self.dropout2(ff_out)\n        return self.add_norm2(addnorm_out, ff_out)\n\nclass Encoder(nn.Module):\n    def __init__(self, vocab_size, sequence_length, h, d_k, d_v, d_model, d_ff, n_layers, dropout_rate):\n        super(Encoder, self).__init__()\n        self.pos_encoding = PositionEmbeddingLayer(sequence_length, vocab_size, d_model)\n        self.dropout = nn.Dropout(dropout_rate)\n        self.dense = nn.Linear(4, d_model)\n        self.activation = nn.SiLU()\n        self.layer_norm = nn.LayerNorm(d_model)\n        self.encoder_layers = nn.ModuleList([\n            EncoderLayer(h, d_k, d_v, d_model, d_ff, dropout_rate) for _ in range(n_layers)\n        ])\n\n    def forward(self, input_sequence, covariance, pssm, freq_matrix):\n        # Generate a padding mask: positions not equal to zero.\n        # (batch, seq_len)\n        mask = (input_sequence != 0).float()\n        # Compute pairwise mask: (batch, seq_len, seq_len)\n        padding_mask = torch.matmul(mask.unsqueeze(2), mask.unsqueeze(1))\n\n        # Positional encoding.\n        pos_enc_out = self.pos_encoding(input_sequence)\n        x = self.dropout(pos_enc_out)\n\n        # Project extra features and add.\n        pssm_proj = self.activation(self.dense(pssm))\n        freq_proj = self.activation(self.dense(freq_matrix))\n        x = x + pssm_proj + freq_proj\n        x = self.layer_norm(x)\n\n        for layer in self.encoder_layers:\n            x = layer(x, covariance, padding_mask)\n        return x\n\nclass RNA3DModel(nn.Module):\n    def __init__(self, vocab_size, sequence_length, h, d_k, d_v, d_model, d_ff, n, dropout_rate):\n        super(RNA3DModel, self).__init__()\n        self.encoder = Encoder(vocab_size, sequence_length, h, d_k, d_v, d_model, d_ff, n, dropout_rate)\n        self.linear = nn.Linear(d_model, 3)\n\n    def forward(self, inputs):\n        X = inputs[0]\n\n        covariance = inputs[1] if len(inputs) > 1 else None\n        freq_matrix = inputs[2] if len(inputs) > 2 else None\n        pssm = inputs[3] if len(inputs) > 3 else None\n\n        output = self.encoder(X, covariance, pssm, freq_matrix)\n        output = self.linear(output)\n        return output","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:03:19.265822Z","iopub.execute_input":"2025-04-18T18:03:19.266180Z","iopub.status.idle":"2025-04-18T18:03:19.286292Z","shell.execute_reply.started":"2025-04-18T18:03:19.266157Z","shell.execute_reply":"2025-04-18T18:03:19.285499Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def parse_msa_fasta(msa_filepath, alphabet=('A', 'C', 'G', 'U'), pseudocount=1e-6):\n    \"\"\"\n    Parse an MSA FASTA file and compute three features:\n      - Frequency matrix: shape (L, |alphabet|) \n      - PSSM: log-odds score matrix of shape (L, |alphabet|)\n      - Covariance matrix: a scalar covariance between columns, shape (L, L)\n\n    Args:\n        msa_filepath (str): Path to the MSA FASTA file.\n        alphabet (tuple): Symbols to consider (default for RNA: A, C, G, U).\n        pseudocount (float): Small value added to avoid log(0).\n\n    Returns:\n        tuple: (frequency_matrix, pssm, covariance)\n            frequency_matrix: np.array of shape (L, len(alphabet))\n            pssm: np.array of shape (L, len(alphabet))\n            covariance: np.array of shape (L, L)\n    \"\"\"\n    # Parse all sequences from the MSA file.\n    msa_records = list(SeqIO.parse(msa_filepath, \"fasta\"))\n    if len(msa_records) == 0:\n        raise ValueError(\"No sequences found in the provided MSA file.\")\n    \n    # All sequences should be the same length.\n    seq_length = len(msa_records[0].seq)\n    M = len(msa_records)\n    K = len(alphabet)\n    \n    # Initialize frequency matrix and one-hot encoded array.\n    freq_matrix = np.zeros((seq_length, K), dtype=np.float32)\n    one_hot = np.zeros((M, seq_length, K), dtype=np.float32)\n    \n    for m, record in enumerate(msa_records):\n        seq = str(record.seq).upper()\n        if len(seq) != seq_length:\n            raise ValueError(\"All sequences in the MSA must have the same length.\")\n        for i, letter in enumerate(seq):\n            if letter in alphabet:\n                idx = alphabet.index(letter)\n                freq_matrix[i, idx] += 1\n                one_hot[m, i, idx] = 1.0\n            # Optionally: handle gaps ('-') or ambiguous symbols here.\n    \n    # Normalize frequencies along each position.\n    freq_matrix /= (M + 1e-8)\n    \n    # Compute PSSM:\n    # Background probabilities, assuming a uniform background.\n    background = np.full((K,), 1.0 / K, dtype=np.float32)\n    # Add pseudocount so we never divide by zero.\n    freq_with_pc = freq_matrix + pseudocount\n    background = background + pseudocount\n    # Compute log-odds score. The result is the PSSM.\n    pssm = np.log(freq_with_pc / background)\n    \n    # Compute Covariance Matrix between positions:\n    # For each position, we already have a frequency distribution computed via one-hot averages.\n    # First, compute the mean one-hot vector per position.\n    mean_one_hot = np.mean(one_hot, axis=0)  # shape: (L, K)\n    # Compute the difference from the mean for each sequence.\n    diff = one_hot - mean_one_hot[None, :, :]  # shape: (M, L, K)\n    # Compute covariance between each pair of positions.\n    # This gives a scalar covariance value for each pair (i, j).\n    covariance = np.einsum('mik,mjk->ij', diff, diff) / (M - 1 + 1e-8)  # shape: (L, L)\n    \n    return freq_matrix, pssm, covariance\n\n\nclass RNADataset(Dataset):\n    def __init__(self, data, tokenizer, freq_matrix=True, pssm=True, covariance=True,\n                 max_length=512, mode='train', augment_coords=True, noise_std=0.1):\n        \"\"\"\n        data: a pandas DataFrame with columns:\n            - \"target_id\", \"sequence\", \"freq_matrix\", \"ppsm\", \"coverience\"\n            - coordinate columns \"x_1\", \"y_1\", \"z_1\"\n        tokenizer: dict mapping sequence tokens to indices.\n        freq_matrix, pssm, covariance: flags to include respective features.\n        max_length: maximum sequence length for padding/truncation.\n        mode: 'train' or 'test'.\n        augment_coords: whether to apply random Gaussian noise to the 3D coordinates.\n        noise_std: standard deviation for coordinate noise.\n        \"\"\"\n        self.data = data.copy()\n        self.tokenizer = tokenizer\n        self.max_length = max_length\n        self.freq_matrix = freq_matrix\n        self.pssm = pssm\n        self.covariance = covariance\n        self.mode = mode\n        self.augment_coords = augment_coords\n        self.noise_std = noise_std\n\n        # Prepare samples aggregated by target_id\n        self.target_ids = self.data[\"target_id\"].unique()\n        self.samples = []\n        for tid in self.target_ids:\n            sample = {}\n            sample_df = self.data[self.data[\"target_id\"] == tid]\n            row = sample_df.iloc[0]\n\n            # Tokenize and pad sequence\n            seq = list(row[\"sequence\"])\n            tokenized = [self.tokenizer.get(tok, 0) for tok in seq]\n            if len(tokenized) < self.max_length:\n                tokenized = tokenized + [0] * (self.max_length - len(tokenized))\n            else:\n                tokenized = tokenized[:self.max_length]\n            sample[\"sequence\"] = np.array(tokenized, dtype=np.int32)\n\n            if self.mode == 'train':\n                # Coordinates: pad with NaN\n                coords = sample_df[[\"x_1\", \"y_1\", \"z_1\"]].values\n                padded_coords = np.full((self.max_length, 3), np.nan, dtype=np.float32)\n                padded_coords[:coords.shape[0], :] = coords\n                sample[\"coords\"] = padded_coords\n\n                # Frequency matrix\n                if self.freq_matrix:\n                    fm = row[\"freq_matrix\"]  # shape (L, K)\n                    padded_fm = np.zeros((self.max_length, fm.shape[1]), dtype=np.float32)\n                    padded_fm[:fm.shape[0], :] = fm\n                    sample[\"freq_matrix\"] = padded_fm\n\n                # PSSM\n                if self.pssm:\n                    ps = row[\"ppsm\"]  # shape (L, K)\n                    padded_ps = np.zeros((self.max_length, ps.shape[1]), dtype=np.float32)\n                    padded_ps[:ps.shape[0], :] = ps\n                    sample[\"pssm\"] = padded_ps\n\n                # Covariance\n                if self.covariance:\n                    cov = row[\"coverience\"]  # shape (L, L)\n                    padded_cov = np.zeros((self.max_length, self.max_length), dtype=np.float32)\n                    padded_cov[:cov.shape[0], :cov.shape[1]] = cov\n                    sample[\"covariance\"] = padded_cov\n            else:\n                if self.freq_matrix:\n                    fm = row[\"freq_matrix\"]  # shape (L, K)\n                    padded_fm = np.zeros((self.max_length, fm.shape[1]), dtype=np.float32)\n                    padded_fm[:fm.shape[0], :] = fm\n                    sample[\"freq_matrix\"] = padded_fm\n\n                # PSSM\n                if self.pssm:\n                    ps = row[\"ppsm\"]  # shape (L, K)\n                    padded_ps = np.zeros((self.max_length, ps.shape[1]), dtype=np.float32)\n                    padded_ps[:ps.shape[0], :] = ps\n                    sample[\"pssm\"] = padded_ps\n\n                # Covariance\n                if self.covariance:\n                    cov = row[\"coverience\"]  # shape (L, L)\n                    padded_cov = np.zeros((self.max_length, self.max_length), dtype=np.float32)\n                    padded_cov[:cov.shape[0], :cov.shape[1]] = cov\n                    sample[\"covariance\"] = padded_cov\n\n            self.samples.append(sample)\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        seq_tensor = torch.tensor(sample[\"sequence\"], dtype=torch.int32)\n\n        if self.mode == 'train':\n            # Assemble inputs\n            inputs = [seq_tensor]\n            if self.covariance:\n                inputs.append(torch.tensor(sample[\"covariance\"], dtype=torch.float))\n            if self.freq_matrix:\n                inputs.append(torch.tensor(sample[\"freq_matrix\"], dtype=torch.float))\n            if self.pssm:\n                inputs.append(torch.tensor(sample[\"pssm\"], dtype=torch.float))\n\n            # Coordinates\n            coords = torch.tensor(sample[\"coords\"], dtype=torch.float)\n\n            # Random augmentation: add Gaussian noise\n            if self.augment_coords:\n                mask = ~torch.isnan(coords[:, 0])\n                noise = torch.randn(mask.sum().item(), 3, device=coords.device) * self.noise_std\n                coords[mask] = coords[mask] + noise\n\n            return inputs, coords\n        else:\n            inputs = [seq_tensor]\n            if self.covariance:\n                inputs.append(torch.tensor(sample[\"covariance\"], dtype=torch.float))\n            if self.freq_matrix:\n                inputs.append(torch.tensor(sample[\"freq_matrix\"], dtype=torch.float))\n            if self.pssm:\n                inputs.append(torch.tensor(sample[\"pssm\"], dtype=torch.float))\n            return inputs","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:03:19.287497Z","iopub.execute_input":"2025-04-18T18:03:19.287762Z","iopub.status.idle":"2025-04-18T18:03:19.308444Z","shell.execute_reply.started":"2025-04-18T18:03:19.287742Z","shell.execute_reply":"2025-04-18T18:03:19.307691Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DMSELoss(nn.Module):\n    def __init__(self):\n        super(DMSELoss, self).__init__()\n\n    def forward(self, y_true, y_pred):\n        d_true = self._pairwise_distance_matrix(y_true)\n        d_pred = self._pairwise_distance_matrix(y_pred)\n        \n        mask=~torch.isnan(d_true)\n        B, N, _ = d_true.shape\n        diag_mask = torch.eye(N, device=d_true.device, dtype=torch.bool).unsqueeze(0).expand(B, N, N)\n        mask[diag_mask] = False\n        \n        mse = (d_true[mask] - d_pred[mask]) ** 2\n        return mse.mean()\n\n    def pairwise_diff_matrix(self, X):\n        return X.unsqueeze(2) - X.unsqueeze(1)  # (batch, seq_len, seq_len, 3)\n\n    def _pairwise_distance_matrix(self, coords):\n        diff = self.pairwise_diff_matrix(coords)\n        dists = torch.sqrt((diff ** 2).sum(dim=-1) + 1e-8)\n        return dists\n\nclass DMAELoss(nn.Module):\n    def __init__(self):\n        super(DMAELoss, self).__init__()\n\n    def forward(self, y_true, y_pred):  \n        d_true = self._pairwise_distance_matrix(y_true)\n        d_pred = self._pairwise_distance_matrix(y_pred)\n        \n        mask=~torch.isnan(d_true)\n        B, N, _ = d_true.shape\n        diag_mask = torch.eye(N, device=d_true.device, dtype=torch.bool).unsqueeze(0).expand(B, N, N)\n        mask[diag_mask] = False\n        mae = torch.abs(d_true[mask] - d_pred[mask])\n        return mae.mean()/10\n\n    def pairwise_diff_matrix(self, X):\n        return X.unsqueeze(2) - X.unsqueeze(1)  # (batch, seq_len, seq_len, 3)\n\n    def _pairwise_distance_matrix(self, coords):\n        diff = self.pairwise_diff_matrix(coords)\n        dists = torch.sqrt((diff ** 2).sum(dim=-1) + 1e-8)\n        return dists\n\nclass TorsionLoss(nn.Module):\n    def __init__(self):\n        super(TorsionLoss, self).__init__()\n\n    def _compute_dihedral_angle(self, p0, p1, p2, p3):\n        b0 = p1 - p0\n        b1 = p2 - p1\n        b2 = p3 - p2\n        b1_norm = torch.norm(b1, dim=-1, keepdim=True)\n        b1_norm_clipped = torch.clamp(b1_norm, min=1e-8)\n        b1 = b1 / b1_norm_clipped\n        v = b0 - torch.sum(b0 * b1, dim=-1, keepdim=True) * b1\n        w = b2 - torch.sum(b2 * b1, dim=-1, keepdim=True) * b1\n        v_norm = torch.norm(v, dim=-1, keepdim=True)\n        w_norm = torch.norm(w, dim=-1, keepdim=True)\n        v = torch.where(v_norm < 1e-8, torch.zeros_like(v), v)\n        w = torch.where(w_norm < 1e-8, torch.zeros_like(w), w)\n        x = torch.sum(v * w, dim=-1)\n        y = torch.sum(torch.cross(b1, v, dim=-1) * w, dim=-1)\n        angle = torch.atan2(y + 1e-8, x + 1e-8)\n        return angle\n\n    def forward(self, y_true, y_pred):\n        mask = (~torch.isnan(y_true[..., 0])).float()  # (batch, seq_len)\n        y_true_clean = torch.where(torch.isnan(y_true), torch.zeros_like(y_true), y_true)\n        y_pred_clean = torch.where(torch.isnan(y_true), torch.zeros_like(y_pred), y_pred)\n        # Create 4-point sliding windows\n        p0_true = y_true_clean[:, :-3]\n        p1_true = y_true_clean[:, 1:-2]\n        p2_true = y_true_clean[:, 2:-1]\n        p3_true = y_true_clean[:, 3:]\n        p0_pred = y_pred_clean[:, :-3]\n        p1_pred = y_pred_clean[:, 1:-2]\n        p2_pred = y_pred_clean[:, 2:-1]\n        p3_pred = y_pred_clean[:, 3:]\n        angle_true = self._compute_dihedral_angle(p0_true, p1_true, p2_true, p3_true)\n        angle_pred = self._compute_dihedral_angle(p0_pred, p1_pred, p2_pred, p3_pred)\n        mask_torsion = mask[:, :-3] * mask[:, 1:-2] * mask[:, 2:-1] * mask[:, 3:]\n        angle_diff = (angle_pred - angle_true + np.pi) % (2 * np.pi) - np.pi\n        loss = 0.5 * (1.0 - torch.cos(angle_diff))\n        loss = (loss * mask_torsion).sum() / (mask_torsion.sum() + 1e-8)\n        return loss\n\n\ndef safe_normalize(x, dim=-1, eps=1e-8):\n    norm = torch.norm(x, dim=dim, keepdim=True)\n    return x / (norm + eps)\n\nclass FAPEloss(nn.Module):\n    def __init__(self, Z=10.0, clamp=10.0):\n        super().__init__()\n        self.Z = Z\n        self.clamp = clamp\n\n    def compute_local_frames(self, coords):\n        B, N, _ = coords.shape\n        R = torch.eye(3, device=coords.device).unsqueeze(0).unsqueeze(0).expand(B, N, 3, 3).clone()\n\n        if N > 2:\n            p_prev = coords[:, :-2, :]  # (B, N-2, 3)\n            p = coords[:, 1:-1, :]      # (B, N-2, 3)\n            p_next = coords[:, 2:, :]   # (B, N-2, 3)\n\n            v1 = p - p_prev             # (B, N-2, 3)\n            v2 = p_next - p             # (B, N-2, 3)\n\n            x_axis = safe_normalize(v1, dim=-1)  # (B, N-2, 3)\n            z_axis = torch.cross(v1, v2, dim=-1)\n            z_axis = safe_normalize(z_axis, dim=-1)\n            y_axis = torch.cross(z_axis, x_axis, dim=-1)  # (B, N-2, 3)\n\n            local_R = torch.stack([x_axis, y_axis, z_axis], dim=-1)  # (B, N-2, 3, 3)\n            R[:, 1:-1, :, :] = local_R\n\n        T = torch.zeros((B, N, 4, 4), device=coords.device)\n        T[:, :, :3, :3] = R\n        T[:, :, :3, 3] = coords   \n        T[:, :, 3, 3] = 1.0\n        return R, T\n\n    def forward(self, y_pred, y_true):\n        R_pred, T_pred = self.compute_local_frames(y_pred)\n        R_true, T_true = self.compute_local_frames(y_true)\n\n\n        delta_pred = y_pred.unsqueeze(2) - y_pred.unsqueeze(1)  # (B, N, N, 3)\n        delta_true = y_true.unsqueeze(2) - y_true.unsqueeze(1)  # (B, N, N, 3)\n        \n        X_pred = torch.einsum('b i c d, b i j d -> b i j c', R_pred, delta_pred)\n        X_true = torch.einsum('b i c d, b i j d -> b i j c', R_true, delta_true)\n\n        trans_error = torch.norm(X_pred - X_true, dim=-1)  # (B, N, N)\n        trans_error = torch.clamp(trans_error, max=self.clamp) / self.Z\n        translation_loss = torch.mean(trans_error)\n\n        R_diff = torch.matmul(R_pred.transpose(-2, -1), R_true)  # (B, N, 3, 3)\n        trace_val = R_diff.diagonal(offset=0, dim1=-2, dim2=-1).sum(-1)  # (B, N)\n        angle_error = torch.acos(torch.clamp((trace_val - 1) / 2, min=-1.0, max=1.0))\n        rotation_loss = torch.mean(angle_error)\n\n        transformation_loss = torch.mean((T_pred - T_true) ** 2)\n\n        total_loss = translation_loss + rotation_loss + transformation_loss\n        return total_loss\n\nclass AlignSVDMSELoss(nn.Module):\n\n    def __init__(self, Z=10.0):\n        super(AlignSVDMSELoss, self).__init__()\n        self.Z = Z\n\n    def forward(self, y_true,y_pred):\n        assert y_pred.shape == y_true.shape\n        mask=~torch.isnan(y_true.sum(-1))\n\n        y_pred=y_pred[mask]\n        y_true=y_true[mask]\n        \n\n        centroid_y_pred = y_pred.mean(dim=0, keepdim=True)\n        centroid_y_true = y_true.mean(dim=0, keepdim=True)\n\n\n        y_pred_centered = y_pred - centroid_y_pred.detach()\n        y_true_centered = y_true - centroid_y_true\n\n        cov_matrix = y_pred_centered.T @ y_true_centered\n\n        U, S, Vt = torch.svd(cov_matrix)\n\n        R = Vt @ U.T\n\n        if torch.det(R) < 0:\n            Vt[-1, :] *= -1\n            R = Vt @ U.T\n\n        aligned_y_pred = (y_pred_centered @ R.T.detach()) + centroid_y_true.detach()\n\n        return torch.abs(aligned_y_pred-y_true).mean()/self.Z\n\nclass dRMAE(nn.Module):\n    def __init__(self):\n        super(dRMAE, self).__init__()\n        self.Z = 10\n\n    def pairwise_diff_matrix(self, X):\n        return X.unsqueeze(2) - X.unsqueeze(1)  # (batch, seq_len, seq_len, 3)\n\n    def _pairwise_distance_matrix(self, coords):\n        diff = self.pairwise_diff_matrix(coords[:,:,:-1])\n        dists = torch.sqrt((diff ** 2).sum(dim=-1) + 1e-8)\n        return dists\n\n    def forward(self, y_true,y_pred, epsilon=1e-4,d_clamp=None): \n\n        y_true_d=self._pairwise_distance_matrix(y_true)\n        y_pred_d=self._pairwise_distance_matrix(y_pred)\n\n        mask=~torch.isnan(y_true_d)\n        diag_eye = torch.stack([torch.eye(mask.shape[1]).bool() for i in range(mask.shape[0])])\n        mask[diag_eye]=False\n\n        rmsd=torch.abs(y_pred_d[mask]-y_true_d[mask])\n\n        return rmsd.mean()/self.Z\nclass CompositeLoss(nn.Module):\n    def __init__(self,\n                 torsion_weight=1,\n                 dmse_weight=1,\n                 dmae_weight=1,\n                 fape_weight=1,\n                 allignsvd_weight=1,\n                 drmae_weight=1):\n        super(CompositeLoss, self).__init__()\n        self.torsion_weight = torsion_weight\n        self.dmse_weight = dmse_weight\n        self.fape_weight = fape_weight\n        self.allignsvd_weight = allignsvd_weight\n        self.drmae_weight = drmae_weight\n        self.dmae_weight = dmae_weight\n\n        self.torsion_loss = TorsionLoss()\n        # self.dmse_loss = DMSELoss()\n        # self.fape_loss = FAPEloss()\n        self.allignsvd_loss = AlignSVDMSELoss()\n        self.drmae_loss = dRMAE()\n        # self.dmae_loss = DMAELoss()\n\n    def forward(self, y_true, y_pred):\n        torsion = self.torsion_loss(y_true, y_pred)\n        # dmse = self.dmse_loss(y_true, y_pred)\n        # fape = self.fape_loss(y_true, y_pred)\n        allignsvd = self.allignsvd_loss(y_true,y_pred)\n        drmae = self.drmae_loss(y_true,y_pred)\n        # dmae = self.dmae_loss(y_true,y_pred)\n\n        return self.allignsvd_weight * allignsvd \\\n        + self.drmae_weight * drmae \\\n        + self.torsion_weight * torsion","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:03:19.309224Z","iopub.execute_input":"2025-04-18T18:03:19.309458Z","iopub.status.idle":"2025-04-18T18:03:19.336881Z","shell.execute_reply.started":"2025-04-18T18:03:19.309439Z","shell.execute_reply":"2025-04-18T18:03:19.336008Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def compute_rmsd(y_true, y_pred):\n    rmsd_values = []\n    B = y_true.shape[0]\n    for i in range(B):\n        yt = y_true[i]  # shape (N, 3)\n        yp = y_pred[i]  # shape (N, 3)\n\n        # Check for non-finite values. You can also use torch.nan_to_num to replace them.\n        if (not torch.isfinite(yt).all()) or (not torch.isfinite(yp).all()):\n            # Skip this sample or set its RMSD to zero.\n            rmsd_values.append(torch.tensor(0.0, device=yt.device))\n            continue\n\n        # Center the coordinates by subtracting the mean (per sample).\n        X = yp - torch.mean(yp, dim=0, keepdim=True)\n        Y = yt - torch.mean(yt, dim=0, keepdim=True)\n\n        # Compute the covariance matrix between the centered predicted and true coordinates.\n        # This gives a 3x3 matrix.\n        C = torch.matmul(X.t(), Y)\n        # Check for non-finite values in C.\n        if not torch.isfinite(C).all():\n            rmsd_values.append(torch.tensor(0.0, device=yt.device))\n            continue\n\n        # Compute the SVD of C. We use torch.linalg.svd which is robust and returns (U, S, Vh).\n        try:\n            U, S, Vt = torch.linalg.svd(C)\n        except RuntimeError as e:\n            # If SVD fails, skip this sample (or return 0 for it)\n            rmsd_values.append(torch.tensor(0.0, device=yt.device))\n            continue\n\n        # Reflection fix: if the determinant of (U * Vt) is negative, fix the sign on the last column of U.\n        det = torch.det(torch.matmul(U, Vt))\n        d = torch.sign(det)\n        # Fix the sign of the last column of U if needed.\n        U[:, -1] = U[:, -1] * d\n        # Compute the optimal rotation matrix.\n        R = torch.matmul(U, Vt)\n\n        # Align the predicted coordinates.\n        X_aligned = torch.matmul(X, R)\n        diff = X_aligned - Y\n        # Compute RMSD over all residues.\n        rmsd = torch.sqrt(torch.mean(torch.sum(diff ** 2, dim=1)))\n        rmsd_values.append(rmsd)\n\n    # Compute average RMSD over the batch.\n    if len(rmsd_values) > 0:\n        return torch.stack(rmsd_values).mean()\n    else:\n        return torch.tensor(0.0, device=y_true.device)\n\ndef compute_tm_score(y_true, y_pred):\n    \"\"\"\n    Compute TM-score between predicted and true 3D coordinates.\n\n    This function supports both single-sample input of shape (N, 3)\n    and batched input of shape (B, N, 3). If batched, the TM-score is averaged\n    over samples.\n\n    Args:\n        y_true (ndarray): Ground-truth coordinates. Can be (N, 3) or (B, N, 3).\n        y_pred (ndarray): Predicted coordinates. Can be (N, 3) or (B, N, 3).\n\n    Returns:\n        float: TM-score (average over batch if batched).\n    \"\"\"\n    # If batched input, process each sample individually.\n    if y_true.ndim == 3:\n        tm_scores = []\n        for i in range(y_true.shape[0]):\n            tm = compute_tm_score(y_true[i], y_pred[i])\n            tm_scores.append(tm)\n        return np.mean(tm_scores)\n\n    # Ensure single sample inputs are at least 2D.\n    if y_true.ndim == 1:\n        y_true = y_true.reshape(1, -1)\n    if y_pred.ndim == 1:\n        y_pred = y_pred.reshape(1, -1)\n\n    # Remove padded residues. We assume padded rows contain NaNs.\n    mask = ~np.isnan(y_true[:, 0])\n    y_true = y_true[mask]\n    y_pred = y_pred[mask]\n\n    L = len(y_true)\n    if L == 0:\n        return 0.0  # Nothing to compare.\n\n    # Center the coordinates.\n    X = y_pred - np.mean(y_pred, axis=0)\n    Y = y_true - np.mean(y_true, axis=0)\n\n    # Compute the covariance matrix.\n    C = np.dot(X.T, Y)\n    if C.ndim < 2:\n        return 0.0\n\n    # Compute SVD in a try/except block to catch convergence issues.\n    try:\n        U, S, Vt = np.linalg.svd(C)\n    except np.linalg.LinAlgError:\n        return 0.0\n\n    # Compute optimal rotation R.\n    R = np.dot(U, Vt)\n\n    # If the rotation matrix is improper (reflection), fix it.\n    if np.linalg.det(R) < 0:\n        Vt[-1, :] *= -1\n        R = np.dot(U, Vt)\n\n    # Align predicted coordinates.\n    X_aligned = np.dot(X, R)\n    dists = np.sqrt(np.sum((X_aligned - Y) ** 2, axis=1))\n\n    # Compute d0 parameter as defined (with a floor value of 0.5).\n    if L > 15:\n        d0 = 1.24 * ((L - 15) ** (1/3)) - 1.8\n    else:\n        d0 = 0.5\n    d0 = max(d0, 0.5)\n\n    tm = np.sum(1 / (1 + (dists / d0) ** 2)) / L\n    return tm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:03:19.337758Z","iopub.execute_input":"2025-04-18T18:03:19.338020Z","iopub.status.idle":"2025-04-18T18:03:19.355703Z","shell.execute_reply.started":"2025-04-18T18:03:19.338000Z","shell.execute_reply":"2025-04-18T18:03:19.354884Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# A, C, G, and U\nTOKENIZER = {\"A\":1, \"C\":2, \"G\":3, \"U\":4,\"-\":5,\"X\":6}\nROOT_DIR = \"/content/drive/MyDrive/Kaggle/Stanford RNA 3D/\"\nCUTOFF_DATE = \"2020-01-01\"\nTEST_CUTOFF_DATE = \"2022-05-01\"\nLEARNING_RATE = 0.0001\nWD = 0.001\nEPOCHS = 50\nWARMUP_RATIO = 0.25\nLEARNING_RATE_MAX = 0.001\nVER = 0.1\nTRAIN = True\nCOVARIENCE = True\nPSSM = True\nFREQ_MATRIX = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:03:19.356399Z","iopub.execute_input":"2025-04-18T18:03:19.356625Z","iopub.status.idle":"2025-04-18T18:03:19.371275Z","shell.execute_reply.started":"2025-04-18T18:03:19.356606Z","shell.execute_reply":"2025-04-18T18:03:19.370503Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_seq = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_sequences.csv\")\nval_seq = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/validation_sequences.csv\")\ntrain_labels = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_labels.csv\")\nval_labels = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/validation_labels.csv\")\ntest_seq = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/test_sequences.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:03:19.373804Z","iopub.execute_input":"2025-04-18T18:03:19.374027Z","iopub.status.idle":"2025-04-18T18:03:19.747452Z","shell.execute_reply.started":"2025-04-18T18:03:19.373997Z","shell.execute_reply":"2025-04-18T18:03:19.746789Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_seq = train_seq.sample(frac=1)\nval_seq = val_seq.sample(frac=1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:03:19.748908Z","iopub.execute_input":"2025-04-18T18:03:19.749154Z","iopub.status.idle":"2025-04-18T18:03:19.758331Z","shell.execute_reply.started":"2025-04-18T18:03:19.749133Z","shell.execute_reply":"2025-04-18T18:03:19.757384Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_seq[[\"freq_matrix\",\"ppsm\",\"coverience\"]] = train_seq[\"target_id\"].apply(lambda x: parse_msa_fasta(\"/kaggle/input/stanford-rna-3d-folding/MSA/\" + x + \".MSA.fasta\")).apply(pd.Series)\nval_seq[[\"freq_matrix\",\"ppsm\",\"coverience\"]] = val_seq[\"target_id\"].apply(lambda x: parse_msa_fasta(\"/kaggle/input/stanford-rna-3d-folding/MSA/\" + x + \".MSA.fasta\")).apply(pd.Series)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:03:19.759121Z","iopub.execute_input":"2025-04-18T18:03:19.759325Z","iopub.status.idle":"2025-04-18T18:08:10.591762Z","shell.execute_reply.started":"2025-04-18T18:03:19.759307Z","shell.execute_reply":"2025-04-18T18:08:10.591063Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_labels[\"target_id\"] = train_labels[\"ID\"].str.split(\"_\",expand=True).apply(lambda x: \"_\".join(x[0:-1]),axis=1)\nval_labels[\"target_id\"] = val_labels[\"ID\"].str.split(\"_\",expand=True).apply(lambda x: \"_\".join(x[0:-1]),axis=1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:10.592492Z","iopub.execute_input":"2025-04-18T18:08:10.592734Z","iopub.status.idle":"2025-04-18T18:08:13.401748Z","shell.execute_reply.started":"2025-04-18T18:08:10.592714Z","shell.execute_reply":"2025-04-18T18:08:13.401066Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"to_remove_train = train_labels.groupby(\"target_id\").apply(lambda x: x[\"x_1\"].isna().sum() > (x[\"resid\"].max()/2)).reset_index().rename(columns={0:\"to_remove\"})\nto_remove_val = val_labels.groupby(\"target_id\").apply(lambda x: x[\"x_1\"].isna().sum() > (x[\"resid\"].max()/2)).reset_index().rename(columns={0:\"to_remove\"})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:13.402463Z","iopub.execute_input":"2025-04-18T18:08:13.402731Z","iopub.status.idle":"2025-04-18T18:08:13.577058Z","shell.execute_reply.started":"2025-04-18T18:08:13.402710Z","shell.execute_reply":"2025-04-18T18:08:13.576191Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_labels = train_labels.merge(to_remove_train[to_remove_train[\"to_remove\"]==False][[\"target_id\"]],on=\"target_id\",how=\"right\")\nval_labels = val_labels.merge(to_remove_val[to_remove_val[\"to_remove\"]==False][[\"target_id\"]],on=\"target_id\",how=\"right\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:13.577804Z","iopub.execute_input":"2025-04-18T18:08:13.578082Z","iopub.status.idle":"2025-04-18T18:08:13.637330Z","shell.execute_reply.started":"2025-04-18T18:08:13.578059Z","shell.execute_reply":"2025-04-18T18:08:13.636663Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data = train_seq.merge(train_labels,on=\"target_id\",how=\"right\")\nval_data = val_seq.merge(val_labels,on=\"target_id\",how=\"right\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:13.638022Z","iopub.execute_input":"2025-04-18T18:08:13.638241Z","iopub.status.idle":"2025-04-18T18:08:13.709133Z","shell.execute_reply.started":"2025-04-18T18:08:13.638223Z","shell.execute_reply":"2025-04-18T18:08:13.708459Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MAX_LENGTH = train_data[\"resid\"].max()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:13.709961Z","iopub.execute_input":"2025-04-18T18:08:13.710244Z","iopub.status.idle":"2025-04-18T18:08:13.714068Z","shell.execute_reply.started":"2025-04-18T18:08:13.710205Z","shell.execute_reply":"2025-04-18T18:08:13.713346Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data.index = train_data[\"target_id\"]\nval_data.index = val_data[\"target_id\"]\ntest_seq.index = test_seq[\"target_id\"]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:13.715122Z","iopub.execute_input":"2025-04-18T18:08:13.715448Z","iopub.status.idle":"2025-04-18T18:08:13.729176Z","shell.execute_reply.started":"2025-04-18T18:08:13.715415Z","shell.execute_reply":"2025-04-18T18:08:13.728557Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MEAN = np.nanmean(train_data[[\"x_1\",\"y_1\",\"z_1\"]].values)\nSTD = np.nanstd(train_data[[\"x_1\",\"y_1\",\"z_1\"]].values)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:13.729875Z","iopub.execute_input":"2025-04-18T18:08:13.730165Z","iopub.status.idle":"2025-04-18T18:08:13.757148Z","shell.execute_reply.started":"2025-04-18T18:08:13.730143Z","shell.execute_reply":"2025-04-18T18:08:13.756494Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# train_data[[\"x_1\",\"y_1\",\"z_1\"]] = (train_data[[\"x_1\",\"y_1\",\"z_1\"]] - MEAN) / STD","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:13.757888Z","iopub.execute_input":"2025-04-18T18:08:13.758138Z","iopub.status.idle":"2025-04-18T18:08:13.761539Z","shell.execute_reply.started":"2025-04-18T18:08:13.758118Z","shell.execute_reply":"2025-04-18T18:08:13.760723Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data = train_data.fillna(np.nan)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:13.762237Z","iopub.execute_input":"2025-04-18T18:08:13.762432Z","iopub.status.idle":"2025-04-18T18:08:13.851136Z","shell.execute_reply.started":"2025-04-18T18:08:13.762414Z","shell.execute_reply":"2025-04-18T18:08:13.850487Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data_final = train_data[pd.to_datetime(train_data[\"temporal_cutoff\"]) <= pd.to_datetime(CUTOFF_DATE)]\nval_data_final = train_data[(pd.to_datetime(train_data[\"temporal_cutoff\"]) > pd.to_datetime(CUTOFF_DATE)) & (pd.to_datetime(train_data[\"temporal_cutoff\"]) <= pd.to_datetime(TEST_CUTOFF_DATE))]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:13.851970Z","iopub.execute_input":"2025-04-18T18:08:13.852174Z","iopub.status.idle":"2025-04-18T18:08:13.902612Z","shell.execute_reply.started":"2025-04-18T18:08:13.852157Z","shell.execute_reply":"2025-04-18T18:08:13.901624Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Train","metadata":{}},{"cell_type":"code","source":"if TRAIN:\n    VOCAB_SIZE = 6\n    H = 8\n    D_K = 64\n    D_V = 64\n    D_MODEL = 128\n    D_FF = 768\n    N_LAYERS = 6\n    DROPOUT_RATE = 0.1\n    LEARNING_RATE = 1e-4\n    WEIGHT_DECAY = 1e-2\n    EPOCHS = 50\n    \n    \n    \n    train_dataset = RNADataset(train_data_final, TOKENIZER, freq_matrix=True, pssm=True, covariance=True,\n                               max_length=MAX_LENGTH, mode='train',augment_coords=False)\n    val_dataset = RNADataset(val_data_final, TOKENIZER, freq_matrix=True, pssm=True, covariance=True,\n                             max_length=MAX_LENGTH, mode='train',augment_coords=False)\n    \n    train_loader = DataLoader(train_dataset, batch_size=1, shuffle=True)\n    val_loader = DataLoader(val_dataset, batch_size=1, shuffle=False)\n    \n    model = RNA3DModel(VOCAB_SIZE, MAX_LENGTH, H, D_K, D_V, D_MODEL, D_FF, N_LAYERS, DROPOUT_RATE)\n    model = model.to(torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\"))\n    \n    optimizer = torch.optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=len(train_loader)*EPOCHS, eta_min=1e-5)\n    \n    loss_fn = AlignSVDMSELoss()\n    loss_fn = loss_fn.to(torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\"))\n    \n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    model.train()\n    for epoch in range(EPOCHS):\n        running_loss = 0.0\n        tm_scores = 0\n        rmsd_scores = 0\n    \n        for batch_idx, (inputs, y_true) in enumerate(train_loader):\n            # Move data to the device.\n            inputs = tuple(inp.to(device) for inp in inputs if inp is not None)\n            y_true = y_true.to(device)\n    \n            optimizer.zero_grad()\n            # Forward pass.\n            y_pred = model(inputs)\n            # Here y_pred is assumed to have shape (B, seq_len, 3)\n            loss = loss_fn(y_true, y_pred)\n            loss.backward()\n            optimizer.step()\n    \n            running_loss += loss.item()\n            rmsd_scores += compute_rmsd(y_true.detach().cpu(), y_pred.detach().cpu())\n            tm_scores += compute_tm_score(y_true.detach().cpu().numpy(), y_pred.detach().cpu().numpy())\n    \n            if batch_idx % 100 == 0:\n                print(f\"Epoch [{epoch+1}/{EPOCHS}], Step [{batch_idx}/{len(train_loader)}], Loss: {loss.item():.4f}\")\n    \n        avg_loss = running_loss / len(train_loader)\n        avg_rmsd = rmsd_scores / len(train_loader)\n        avg_tm = tm_scores / len(train_loader)\n        print(f\"Epoch [{epoch+1}/{EPOCHS}] Average Loss: {avg_loss:.4f}, RMSD: {avg_rmsd:.4f}, TM: {avg_tm:.4f}\")\n    \n        # Evaluate on validation set and compute TM-score.\n        tm_scores = []\n        model.eval()\n        with torch.no_grad():\n            for inputs, y_true in val_loader:\n                inputs = tuple(inp.to(device) for inp in inputs if inp is not None)\n                y_true = y_true.to(device)\n                y_pred = model(inputs)\n                # Convert to numpy and compute TM-score per sample.\n                y_true_np = y_true.cpu().numpy()\n                y_pred_np = y_pred.cpu().numpy()\n                tm = compute_tm_score(y_true_np, y_pred_np)\n                tm_scores.append(tm)\n        avg_tm = np.mean(tm_scores) if tm_scores else 0.0\n        print(f\"Epoch [{epoch+1}/{EPOCHS}] Validation TM-score: {avg_tm:.4f}\")\n        model.train()\n    \n        scheduler.step()\n    \n    # Save the model weights\n    torch.save(model.state_dict(), f\"StanfordRNA3D_model_{VER}.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:13.903705Z","iopub.execute_input":"2025-04-18T18:08:13.904034Z","iopub.status.idle":"2025-04-18T18:08:45.792314Z","shell.execute_reply.started":"2025-04-18T18:08:13.904006Z","shell.execute_reply":"2025-04-18T18:08:45.791392Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"code","source":"del val_loader, train_loader, train_dataset, val_dataset,model,inputs,y_true, y_pred","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:45.793230Z","iopub.execute_input":"2025-04-18T18:08:45.793671Z","iopub.status.idle":"2025-04-18T18:08:45.813792Z","shell.execute_reply.started":"2025-04-18T18:08:45.793638Z","shell.execute_reply":"2025-04-18T18:08:45.812983Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import gc\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:45.817071Z","iopub.execute_input":"2025-04-18T18:08:45.817278Z","iopub.status.idle":"2025-04-18T18:08:45.964820Z","shell.execute_reply.started":"2025-04-18T18:08:45.817260Z","shell.execute_reply":"2025-04-18T18:08:45.964019Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\ntorch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:45.966039Z","iopub.execute_input":"2025-04-18T18:08:45.966324Z","iopub.status.idle":"2025-04-18T18:08:46.046271Z","shell.execute_reply.started":"2025-04-18T18:08:45.966299Z","shell.execute_reply":"2025-04-18T18:08:46.045641Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.cuda.reset_max_memory_allocated()\ntorch.cuda.reset_peak_memory_stats()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:46.047013Z","iopub.execute_input":"2025-04-18T18:08:46.047256Z","iopub.status.idle":"2025-04-18T18:08:46.052712Z","shell.execute_reply.started":"2025-04-18T18:08:46.047228Z","shell.execute_reply":"2025-04-18T18:08:46.051971Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/sample_submission.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:46.053365Z","iopub.execute_input":"2025-04-18T18:08:46.053631Z","iopub.status.idle":"2025-04-18T18:08:46.078975Z","shell.execute_reply.started":"2025-04-18T18:08:46.053601Z","shell.execute_reply":"2025-04-18T18:08:46.078180Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_seq[[\"freq_matrix\",\"ppsm\",\"coverience\"]] = test_seq[\"target_id\"].apply(lambda x: parse_msa_fasta(\"/kaggle/input/stanford-rna-3d-folding/MSA/\" + x + \".MSA.fasta\")).apply(pd.Series)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:46.079779Z","iopub.execute_input":"2025-04-18T18:08:46.080062Z","iopub.status.idle":"2025-04-18T18:08:49.984196Z","shell.execute_reply.started":"2025-04-18T18:08:46.080034Z","shell.execute_reply":"2025-04-18T18:08:49.983545Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_dataset = RNADataset(test_seq, TOKENIZER, freq_matrix=True, pssm=True, covariance=True,\n                             max_length=MAX_LENGTH,augment_coords=False, mode='test')\n    \ntest_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:49.985044Z","iopub.execute_input":"2025-04-18T18:08:49.985283Z","iopub.status.idle":"2025-04-18T18:08:50.009314Z","shell.execute_reply.started":"2025-04-18T18:08:49.985262Z","shell.execute_reply":"2025-04-18T18:08:50.008757Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = RNA3DModel(VOCAB_SIZE, MAX_LENGTH, H, D_K, D_V, D_MODEL, D_FF, N_LAYERS, DROPOUT_RATE)\nmodel = model.to(torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:50.009965Z","iopub.execute_input":"2025-04-18T18:08:50.010157Z","iopub.status.idle":"2025-04-18T18:08:50.040457Z","shell.execute_reply.started":"2025-04-18T18:08:50.010140Z","shell.execute_reply":"2025-04-18T18:08:50.039865Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.load_state_dict(torch.load('/kaggle/working/StanfordRNA3D_model_0.1.pth',weights_only=True))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:50.041200Z","iopub.execute_input":"2025-04-18T18:08:50.041506Z","iopub.status.idle":"2025-04-18T18:08:50.082222Z","shell.execute_reply.started":"2025-04-18T18:08:50.041459Z","shell.execute_reply":"2025-04-18T18:08:50.081581Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import gc\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:50.082996Z","iopub.execute_input":"2025-04-18T18:08:50.083196Z","iopub.status.idle":"2025-04-18T18:08:50.220827Z","shell.execute_reply.started":"2025-04-18T18:08:50.083179Z","shell.execute_reply":"2025-04-18T18:08:50.220030Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\ntorch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:50.221446Z","iopub.execute_input":"2025-04-18T18:08:50.221711Z","iopub.status.idle":"2025-04-18T18:08:50.236170Z","shell.execute_reply.started":"2025-04-18T18:08:50.221691Z","shell.execute_reply":"2025-04-18T18:08:50.235317Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"preds = torch.zeros(len(test_loader),MAX_LENGTH,5,3)\nfor i in range(5):\n    for j,inputs in enumerate(test_loader):\n        inputs = tuple(inp.to(device) for inp in inputs if inp is not None)\n        with torch.inference_mode():\n            pred = model(inputs)\n        del inputs\n        torch.cuda.empty_cache()\n        preds[j,:,i,:] = pred","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:20:54.422236Z","iopub.execute_input":"2025-04-18T18:20:54.422550Z","iopub.status.idle":"2025-04-18T18:21:08.544452Z","shell.execute_reply.started":"2025-04-18T18:20:54.422525Z","shell.execute_reply":"2025-04-18T18:21:08.543557Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"preds.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:21:39.588132Z","iopub.execute_input":"2025-04-18T18:21:39.588446Z","iopub.status.idle":"2025-04-18T18:21:39.593254Z","shell.execute_reply.started":"2025-04-18T18:21:39.588418Z","shell.execute_reply":"2025-04-18T18:21:39.592502Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pred_cords = preds.reshape(-1,MAX_LENGTH,15)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:22:06.273999Z","iopub.execute_input":"2025-04-18T18:22:06.274282Z","iopub.status.idle":"2025-04-18T18:22:06.278158Z","shell.execute_reply.started":"2025-04-18T18:22:06.274261Z","shell.execute_reply":"2025-04-18T18:22:06.277281Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"c = 0\nfor i,(target_id,row) in enumerate(test_seq.iterrows()):\n    sequence = list(row[\"sequence\"])\n    seq_len = len(sequence)\n    for j in range(seq_len):\n        submission.loc[c,:] = [target_id+f\"_{j+1}\",sequence[j],j+1] + pred_cords[i,j,:].tolist()\n        c += 1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:22:09.216339Z","iopub.execute_input":"2025-04-18T18:22:09.216681Z","iopub.status.idle":"2025-04-18T18:22:14.438996Z","shell.execute_reply.started":"2025-04-18T18:22:09.216642Z","shell.execute_reply":"2025-04-18T18:22:14.438088Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission.to_csv(\"submission.csv\",index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T18:08:50.383696Z","iopub.status.idle":"2025-04-18T18:08:50.384015Z","shell.execute_reply":"2025-04-18T18:08:50.383905Z"}},"outputs":[],"execution_count":null}]}