{"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":290651,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":249017,"modelId":270540}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Hello\nHello, Kagglers. I created this notebook show how to tune the version of Ribonanzanet from multimolecule team(https://huggingface.co/multimolecule/ribonanzanet). It seams they have buch of interesing weights applyable to this competition. Lets check them out. ","metadata":{}},{"cell_type":"markdown","source":"First, Download the package","metadata":{}},{"cell_type":"code","source":"%pip install git+https://github.com/dls5-omics/multimolecule@develop --quiet","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-09T16:19:27.771656Z","iopub.execute_input":"2025-04-09T16:19:27.771966Z","iopub.status.idle":"2025-04-09T16:19:44.220105Z","shell.execute_reply.started":"2025-04-09T16:19:27.771939Z","shell.execute_reply":"2025-04-09T16:19:44.219072Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## RibonanzaNetWith3dHead","metadata":{}},{"cell_type":"markdown","source":"Multimolecule's architecture a bit different than the default one. It uses BOS and EOS indexes as 1 and 2 and different end hot encoding pattern. It seems like that you can turn off  bos and eos indexes somewhere in tokenizer's config but gemini adviced me just to modify the model's output a bit. Exclude first and last predicted coordinates and use others to calculate the loss. I plan to trust the AI this time :). Encoding pattern for ACGU is 6, 7, 8, 9. But I don't need to think about it. This job will be taken by the tokenizer. ","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom multimolecule import RnaTokenizer, RibonanzaNetModel\n\nclass RibonanzaNetWith3DHead(nn.Module):\n    \"\"\"\n    A wrapper model that uses a pre-trained RibonanzaNetModel,\n    removes BOS/EOS token embeddings, and predicts 3D coordinates.\n    \"\"\"\n    def __init__(self, pretrained_model_name=\"multimolecule/ribonanzanet\"):\n        super().__init__()\n        self.pretrained_model_name = pretrained_model_name\n        self.base_model = RibonanzaNetModel.from_pretrained(self.pretrained_model_name)\n        if hasattr(self.base_model, 'config') and hasattr(self.base_model.config, 'hidden_size'):\n             self.hidden_size = self.base_model.config.hidden_size\n             print(f\"Detected hidden size: {self.hidden_size}\")\n        else:\n            # Fallback or raise error if config/hidden_size is not found\n            print(\"Warning: Could not automatically determine hidden size from model config. Assuming 256.\")\n            self.hidden_size = 256 # User mentioned 256, use as fallback\n        self.coord_head = nn.Linear(self.hidden_size, 3) # Output size 3 for (x, y, z)\n\n    def forward(self, input_ids, attention_mask=None, **kwargs):\n        \"\"\"\n        Forward pass through the base model and the coordinate head.\n        Accepts tokenized input (input_ids, attention_mask, etc.)\n        \"\"\"\n        base_model_args = {'input_ids': input_ids}\n        if attention_mask is not None:\n            base_model_args['attention_mask'] = attention_mask.float()\n        outputs = self.base_model(**base_model_args)\n        if hasattr(outputs, 'last_hidden_state'):\n            all_embeddings = outputs.last_hidden_state\n        else:\n            raise AttributeError(\"Model output object does not have 'last_hidden_state'. Inspect the output structure.\")\n        # Slice to remove BOS (index 0) and EOS (index -1) embeddings\n        # Shape becomes: (batch_size, sequence_length_nucleotides, hidden_size)\n        # Note: This assumes BOS is always first and EOS is always last.\n        nucleotide_embeddings = all_embeddings[:, 1:-1, :]\n\n        # Handle potential empty sequence after slicing if input was just [BOS, EOS]\n        if nucleotide_embeddings.shape[1] == 0:\n             # Return an empty tensor with the correct dimensions or handle as needed\n             # For prediction, maybe return shape (batch_size, 0, 3)\n             # During training, this case might need special loss handling (e.g., ignore sample)\n             print(\"Warning: Sequence length is zero after removing BOS/EOS.\")\n             # Example: return empty tensor\n             return torch.zeros(nucleotide_embeddings.shape[0], 0, 3, device=nucleotide_embeddings.device)\n        # Pass nucleotide embeddings through the coordinate prediction head\n        # Shape: (batch_size, sequence_length_nucleotides, 3)\n        predicted_coords = self.coord_head(nucleotide_embeddings)\n\n        return predicted_coords","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-09T16:20:08.340697Z","iopub.execute_input":"2025-04-09T16:20:08.341023Z","iopub.status.idle":"2025-04-09T16:20:32.107450Z","shell.execute_reply.started":"2025-04-09T16:20:08.340997Z","shell.execute_reply":"2025-04-09T16:20:32.106756Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data work. Load and process the data\nNothing tricky here. The code I took from this notebook\nhttps://www.kaggle.com/code/leeheewon01/ribonanzanet-3d-finetune-v2","metadata":{}},{"cell_type":"code","source":"# Load CSVs\nfrom tqdm import tqdm\nimport warnings\ntrain_sequences = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_sequences.csv\")\ntrain_labels = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_labels.csv\")\n\n# Create a pdb_id field\ntrain_labels[\"pdb_id\"] = train_labels[\"ID\"].apply(\n    lambda x: x.split(\"_\")[0] + \"_\" + x.split(\"_\")[1]\n)\n\n# Collect xyz data for each sequence\n# Add warning ignore to make readable output\nwith warnings.catch_warnings():\n    warnings.simplefilter(\"ignore\", RuntimeWarning)\n    all_xyz = []\n    for pdb_id in tqdm(train_sequences[\"target_id\"], desc=\"Collecting XYZ data\"):\n        df = train_labels[train_labels[\"pdb_id\"] == pdb_id]\n        xyz = df[[\"x_1\", \"y_1\", \"z_1\"]].to_numpy().astype(\"float32\")\n        xyz[xyz < -1e17] = float(\"nan\")\n        all_xyz.append(xyz)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-09T16:20:37.505125Z","iopub.execute_input":"2025-04-09T16:20:37.505831Z","iopub.status.idle":"2025-04-09T16:20:46.379787Z","shell.execute_reply.started":"2025-04-09T16:20:37.505798Z","shell.execute_reply":"2025-04-09T16:20:46.379004Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"valid_indices = []\nmax_len_seen = 0\n\nfor i, xyz in enumerate(all_xyz):\n    # Track the maximum length\n    if len(xyz) > max_len_seen:\n        max_len_seen = len(xyz)\n\n    nan_ratio = np.isnan(xyz).mean()\n    seq_len = len(xyz)\n    # Keep sequence if it meets criteria\n    if (nan_ratio <= 0.5) and (10 < seq_len < 99999999):\n        valid_indices.append(i)\n\nprint(f\"Longest sequence in train: {max_len_seen}\")\n\n# Filter sequences & xyz based on valid_indices\ntrain_sequences = train_sequences.loc[valid_indices].reset_index(drop=True)\nall_xyz = [all_xyz[i] for i in valid_indices]\n\n# Prepare final data dictionary\ndata = {\n    \"sequence\": train_sequences[\"sequence\"].tolist(),\n    \"temporal_cutoff\": train_sequences[\"temporal_cutoff\"].tolist(),\n    \"description\": train_sequences[\"description\"].tolist(),\n    \"all_sequences\": train_sequences[\"all_sequences\"].tolist(),\n    \"xyz\": all_xyz,\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-09T16:21:46.626607Z","iopub.execute_input":"2025-04-09T16:21:46.626985Z","iopub.status.idle":"2025-04-09T16:21:46.643677Z","shell.execute_reply.started":"2025-04-09T16:21:46.626954Z","shell.execute_reply":"2025-04-09T16:21:46.642829Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cutoff_date = \"2020-01-01\"\ntest_cutoff_date = \"2022-05-01\"\ncutoff_date = pd.Timestamp(cutoff_date)\ntest_cutoff_date = pd.Timestamp(test_cutoff_date)\ntrain_indices = [i for i, date_str in enumerate(data[\"temporal_cutoff\"]) if pd.Timestamp(date_str) <= cutoff_date]\ntest_indices = [i for i, date_str in enumerate(data[\"temporal_cutoff\"]) if cutoff_date < pd.Timestamp(date_str) <= test_cutoff_date]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-09T16:21:51.181760Z","iopub.execute_input":"2025-04-09T16:21:51.182059Z","iopub.status.idle":"2025-04-09T16:21:51.189304Z","shell.execute_reply.started":"2025-04-09T16:21:51.182035Z","shell.execute_reply":"2025-04-09T16:21:51.188428Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Here we need to modify dataloader to take into account the tokenizer's job. ","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader\nclass RNA3D_Dataset(Dataset):\n    \"\"\"\n    A PyTorch Dataset for 3D RNA structures.\n    \"\"\"\n    def __init__(self, indices, data_dict,  tokenizer, max_len=128):\n        self.indices = indices\n        self.data = data_dict\n        self.max_len = max_len\n\n    def __len__(self):\n        return len(self.indices)\n    \n    def __getitem__(self, idx):\n        data_idx = self.indices[idx]\n        sequence = tokenizer(self.data[\"sequence\"][data_idx], return_tensors=\"pt\", padding=True)\n        sequence[\"input_ids\"] = sequence[\"input_ids\"].squeeze()\n        sequence[\"attention_mask\"] = sequence[\"attention_mask\"].squeeze()\n        if torch.cuda.is_available():\n            sequence = {k: v.cuda() for k, v in sequence.items()}\n        # Convert xyz to torch tensor\n        xyz = torch.tensor(self.data[\"xyz\"][data_idx], dtype=torch.float32)\n\n        # If sequence is longer than max_len, randomly crop\n        if len(sequence[\"input_ids\"]) > self.max_len:\n            crop_start = np.random.randint(len(sequence[\"input_ids\"]) - self.max_len)\n            crop_end = crop_start + self.max_len\n            sequence[\"input_ids\"] = sequence[\"input_ids\"][crop_start:crop_end]\n            sequence[\"attention_mask\"] = sequence[\"attention_mask\"][crop_start:crop_end]\n            sequence[\"input_ids\"][0] = 1\n            sequence[\"input_ids\"][127] = 2 #127 is 128-1. Just save some computation\n            xyz = xyz[crop_start+1:crop_end-1]\n\n        return {\"sequence\": sequence, \"xyz\": xyz, \"shape\": sequence[\"input_ids\"].shape}\n\ntokenizer = RnaTokenizer.from_pretrained(\"multimolecule/ribonanzanet\")\ntrain_dataset = RNA3D_Dataset(train_indices, data, tokenizer) # max_len was 384, I changed to 128 becuase 384 was throwind out of memory error on Kaggle's GPU\nval_dataset = RNA3D_Dataset(test_indices, data, tokenizer)\n\ntrain_loader = DataLoader(train_dataset, batch_size=1, shuffle=True) # Leave batch_size=1 so far, change in the future\nval_loader = DataLoader(val_dataset, batch_size=1, shuffle=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-09T16:22:01.939028Z","iopub.execute_input":"2025-04-09T16:22:01.939369Z","iopub.status.idle":"2025-04-09T16:22:06.087413Z","shell.execute_reply.started":"2025-04-09T16:22:01.939341Z","shell.execute_reply":"2025-04-09T16:22:06.086526Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Loss function","metadata":{}},{"cell_type":"markdown","source":"Again, I took all of the code from this notebook in this cell. https://www.kaggle.com/code/leeheewon01/ribonanzanet-3d-finetune-v2","metadata":{}},{"cell_type":"code","source":"def calculate_distance_matrix(X, Y, epsilon=1e-4):\n    \"\"\"\n    Calculate pairwise distances between every point in X and every point in Y.\n    Shape: (len(X), len(Y))\n    \"\"\"\n    return (torch.square(X[:,None]-Y[None,:])+epsilon).sum(-1).sqrt()\n\ndef dRMSD(pred_x, pred_y, gt_x, gt_y, epsilon=1e-4, Z=10, d_clamp=None):\n    \"\"\"\n    Distance-based RMSD.\n    pred_x, pred_y: predicted coordinates (usually the same tensor for X and Y).\n    gt_x, gt_y: ground truth coordinates.\n    \"\"\"\n    pred_dm = calculate_distance_matrix(pred_x, pred_y)\n    gt_dm = calculate_distance_matrix(gt_x, gt_y)\n\n    mask = ~torch.isnan(gt_dm)\n    mask[torch.eye(mask.shape[0], device=mask.device).bool()] = False\n\n    diff_sq = (pred_dm[mask] - gt_dm[mask])**2 + epsilon\n    if d_clamp is not None:\n        diff_sq = diff_sq.clamp(max=d_clamp**2)\n\n    return diff_sq.sqrt().mean() / Z\n\ndef local_dRMSD(pred_x, pred_y, gt_x, gt_y, epsilon=1e-4, Z=10, d_clamp=30):\n    \"\"\"\n    Local distance-based RMSD, ignoring distances above a clamp threshold.\n    \"\"\"\n    pred_dm = calculate_distance_matrix(pred_x, pred_y)\n    gt_dm = calculate_distance_matrix(gt_x, gt_y)\n\n    mask = (~torch.isnan(gt_dm)) & (gt_dm < d_clamp)\n    mask[torch.eye(mask.shape[0], device=mask.device).bool()] = False\n\n    diff_sq = (pred_dm[mask] - gt_dm[mask])**2 + epsilon\n    return diff_sq.sqrt().mean() / Z\n\ndef dRMAE(pred_x, pred_y, gt_x, gt_y, epsilon=1e-4, Z=10):\n    \"\"\"\n    Distance-based Mean Absolute Error.\n    \"\"\"\n    pred_dm = calculate_distance_matrix(pred_x, pred_y)\n    gt_dm = calculate_distance_matrix(gt_x, gt_y)\n\n    mask = ~torch.isnan(gt_dm)\n    mask[torch.eye(mask.shape[0], device=mask.device).bool()] = False\n\n    diff = torch.abs(pred_dm[mask] - gt_dm[mask])\n    return diff.mean() / Z\n\ndef align_svd_mae(input_coords, target_coords, Z=10):\n    \"\"\"\n    Align input_coords to target_coords via SVD (Kabsch algorithm) and compute MAE.\n    \"\"\"\n    assert input_coords.shape == target_coords.shape, \"Input and target must have the same shape\"\n\n    # Create mask for valid points\n    mask = ~torch.isnan(target_coords.sum(dim=-1))\n    input_coords = input_coords[mask]\n    target_coords = target_coords[mask]\n    \n    # Compute centroids\n    centroid_input = input_coords.mean(dim=0, keepdim=True)\n    centroid_target = target_coords.mean(dim=0, keepdim=True)\n\n    # Center the points\n    input_centered = input_coords - centroid_input\n    target_centered = target_coords - centroid_target\n\n    # Compute covariance matrix\n    cov_matrix = input_centered.T @ target_centered\n\n    # SVD to find optimal rotation\n    U, S, Vt = torch.svd(cov_matrix)\n    R = Vt @ U.T\n\n    # Ensure a proper rotation (determinant R == 1)\n    if torch.det(R) < 0:\n        Vt_adj = Vt.clone()   # Clone to avoid in-place modification issues\n        Vt_adj[-1, :] = -Vt_adj[-1, :]\n        R = Vt_adj @ U.T\n\n    # Rotate input and compute mean absolute error\n    aligned_input = (input_centered @ R.T) + centroid_target\n    return torch.abs(aligned_input - target_coords).mean() / Z","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-09T16:22:23.291536Z","iopub.execute_input":"2025-04-09T16:22:23.291850Z","iopub.status.idle":"2025-04-09T16:22:23.303047Z","shell.execute_reply.started":"2025-04-09T16:22:23.291824Z","shell.execute_reply":"2025-04-09T16:22:23.302059Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training the model","metadata":{}},{"cell_type":"markdown","source":"All the code from this notebook. https://www.kaggle.com/code/leeheewon01/ribonanzanet-3d-finetune-v2 . Special thanks to the author. ","metadata":{}},{"cell_type":"code","source":"def train_model(model, train_dl, val_dl, epochs=1, cos_epoch=35, lr=3e-4, clip=1):\n    \"\"\"Train the model with a CosineAnnealingLR after `cos_epoch` epochs.\"\"\"\n    optimizer = torch.optim.AdamW(model.parameters(), weight_decay=0.0, lr=lr)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n        optimizer,\n        T_max=(epochs - cos_epoch) * len(train_dl),\n    )\n\n    best_val_loss = float(\"inf\")\n    best_preds = None\n\n    for epoch in range(epochs):\n        model.train()\n        train_pbar = tqdm(train_dl, desc=f\"Training Epoch {epoch+1}/{epochs}\")\n        running_loss = 0.0\n\n        for idx, batch in enumerate(train_pbar):\n            #sequence = batch[\"sequence\"].cuda()\n            sequence = batch[\"sequence\"]\n            gt_xyz = batch[\"xyz\"].cuda()\n            gt_xyz = gt_xyz.squeeze()\n\n            pred_xyz = model(**sequence).squeeze()\n\n            # Combine two distance-based losses\n            loss = dRMAE(pred_xyz, pred_xyz, gt_xyz, gt_xyz) + align_svd_mae(pred_xyz, gt_xyz)\n            loss.backward()\n\n            # Gradient clipping\n            torch.nn.utils.clip_grad_norm_(model.parameters(), clip)\n            optimizer.step()\n            optimizer.zero_grad()\n\n            if (epoch + 1) > cos_epoch:\n                scheduler.step()\n\n            running_loss += loss.item()\n            avg_loss = running_loss / (idx + 1)\n            train_pbar.set_description(f\"Epoch {epoch+1} | Loss: {avg_loss:.4f}\")\n\n        # Validation\n        model.eval()\n        val_loss = 0.0\n        val_preds = []\n        with torch.no_grad():\n            for idx, batch in enumerate(val_dl):\n                sequence = batch[\"sequence\"]\n                gt_xyz = batch[\"xyz\"].cuda()\n                gt_xyz = gt_xyz.squeeze()\n\n                pred_xyz = model(**sequence).squeeze()\n                loss = dRMAE(pred_xyz, pred_xyz, gt_xyz, gt_xyz)\n                val_loss += loss.item()\n\n                val_preds.append((gt_xyz.cpu().numpy(), pred_xyz.cpu().numpy()))\n\n            val_loss /= len(val_dl)\n            print(f\"Validation Loss (Epoch {epoch+1}): {val_loss:.4f}\")\n\n            # Check for improvement\n            if val_loss < best_val_loss:\n                best_val_loss = val_loss\n                best_preds = val_preds\n                torch.save(model.state_dict(),\"RibonanzaNet_multimolecule_fine_tuned_good_val.pt\")\n                print(f\"  -> New best model saved at epoch {epoch+1}\")\n\n    # Save final model\n    torch.save(model.state_dict(), \"RibonanzaNet_multimolecule_fine_tuned.pt\")\n    return best_val_loss, best_preds","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-09T16:24:45.599955Z","iopub.execute_input":"2025-04-09T16:24:45.600336Z","iopub.status.idle":"2025-04-09T16:24:45.609303Z","shell.execute_reply.started":"2025-04-09T16:24:45.600301Z","shell.execute_reply":"2025-04-09T16:24:45.608529Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Finally, let's get training ! ","metadata":{}},{"cell_type":"code","source":"model = RibonanzaNetWith3DHead(pretrained_model_name=\"multimolecule/ribonanzanet\").cuda()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-09T16:22:48.788025Z","iopub.execute_input":"2025-04-09T16:22:48.788395Z","iopub.status.idle":"2025-04-09T16:22:51.455567Z","shell.execute_reply.started":"2025-04-09T16:22:48.788364Z","shell.execute_reply":"2025-04-09T16:22:51.454800Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    best_loss, best_predictions = train_model(\n        model=model,\n        train_dl=train_loader,\n        val_dl=val_loader,\n        epochs=10,         # or config[\"epochs\"]\n        cos_epoch=35,      # or config[\"cos_epoch\"]\n        lr=3e-4,\n        clip=1\n    )\n    print(f\"Best Validation Loss: {best_loss:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-09T16:24:50.184577Z","iopub.execute_input":"2025-04-09T16:24:50.184904Z","iopub.status.idle":"2025-04-09T16:49:08.680062Z","shell.execute_reply.started":"2025-04-09T16:24:50.184873Z","shell.execute_reply":"2025-04-09T16:49:08.678700Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Further work\nThe work is not over yet. I plan to check some samples with US-align and prepare the submission as well as experiment with other model parameters. ","metadata":{}}]}