{"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":11228175,"sourceType":"competition"},{"sourceId":7639698,"sourceType":"datasetVersion","datasetId":4299272},{"sourceId":8318191,"sourceType":"datasetVersion","datasetId":4459124}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"This is a copy of https://www.kaggle.com/code/shujun717/ribonanzanet-3d-finetune with adding pl for training loop\n\n* added training with batch size equal to 2","metadata":{}},{"cell_type":"code","source":"import torch\nimport random\nimport pickle\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nfrom tqdm import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-07T09:21:11.849537Z","iopub.execute_input":"2025-03-07T09:21:11.849885Z","iopub.status.idle":"2025-03-07T09:21:15.402963Z","shell.execute_reply.started":"2025-03-07T09:21:11.849853Z","shell.execute_reply":"2025-03-07T09:21:15.402046Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"config = {\n    \"seed\": 0,\n    \"cutoff_date\": \"2020-01-01\",\n    \"test_cutoff_date\": \"2022-05-01\",\n    \"max_len\": 384,\n    \"batch_size\": 1,\n    \"learning_rate\": 1e-4,\n    \"weight_decay\": 0.0,\n    \"mixed_precision\": \"bf16\",\n    \"model_config_path\": \"../working/configs/pairwise.yaml\",  # Adjust path as needed\n    \"epochs\": 10,\n    \"cos_epoch\": 5,\n    \"loss_power_scale\": 1.0,\n    \"max_cycles\": 1,\n    \"grad_clip\": 0.1,\n    \"gradient_accumulation_steps\": 1,\n    \"d_clamp\": 30,\n    \"max_len_filter\": 9999999,\n    \"min_len_filter\": 10, \n    \"structural_violation_epoch\": 50,\n    \"balance_weight\": False,\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-07T09:21:15.404018Z","iopub.execute_input":"2025-03-07T09:21:15.404451Z","iopub.status.idle":"2025-03-07T09:21:15.408847Z","shell.execute_reply.started":"2025-03-07T09:21:15.404428Z","shell.execute_reply":"2025-03-07T09:21:15.408001Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Get data and do some data processing¶\n","metadata":{"execution":{"iopub.status.busy":"2025-02-27T00:35:07.639563Z","iopub.execute_input":"2025-02-27T00:35:07.63984Z","iopub.status.idle":"2025-02-27T00:35:07.643454Z","shell.execute_reply.started":"2025-02-27T00:35:07.639817Z","shell.execute_reply":"2025-02-27T00:35:07.64259Z"}}},{"cell_type":"code","source":"# Load data\n\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\ntrain_labels[\"pdb_id\"] = train_labels[\"ID\"].apply(lambda x: x.split(\"_\")[0]+'_'+x.split(\"_\")[1])\ntrain_labels[\"pdb_id\"] ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-07T09:21:15.410074Z","iopub.execute_input":"2025-03-07T09:21:15.410271Z","iopub.status.idle":"2025-03-07T09:21:15.879962Z","shell.execute_reply.started":"2025-03-07T09:21:15.410253Z","shell.execute_reply":"2025-03-07T09:21:15.879242Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_xyz = []\n\nfor pdb_id in tqdm(train_sequences['target_id']):\n    df = train_labels[train_labels[\"pdb_id\"] == pdb_id]\n    xyz = df[['x_1','y_1','z_1']].to_numpy().astype('float32')\n    if not np.isnan(xyz).any(): xyz[xyz<-1e17] = float('Nan')\n\n    all_xyz.append(xyz)\n\ndf","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-07T09:21:16.227181Z","iopub.execute_input":"2025-03-07T09:21:16.227480Z","iopub.status.idle":"2025-03-07T09:21:24.658765Z","shell.execute_reply.started":"2025-03-07T09:21:16.227456Z","shell.execute_reply":"2025-03-07T09:21:24.657904Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# filter the data\n# Filter and process data\nfilter_nan = []\nmax_len = 0\nfor xyz in all_xyz:\n    if len(xyz) > max_len:\n        max_len = len(xyz)\n\n    #fill -1e18 masked sequences to nans\n\n    #sugar_xyz = np.stack([nt_xyz['sugar_ring'] for nt_xyz in xyz], axis=0)\n    filter_nan.append((np.isnan(xyz).mean() <= 0.5) & \\\n                      (len(xyz)<config['max_len_filter']) & \\\n                      (len(xyz)>config['min_len_filter']))\n\nprint(f\"Longest sequence in train: {max_len}\")\n\nfilter_nan = np.array(filter_nan)\nnon_nan_indices = np.arange(len(filter_nan))[filter_nan]\n\ntrain_sequences = train_sequences.loc[non_nan_indices].reset_index(drop=True)\nall_xyz=[all_xyz[i] for i in non_nan_indices]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-07T09:21:24.659900Z","iopub.execute_input":"2025-03-07T09:21:24.660159Z","iopub.status.idle":"2025-03-07T09:21:24.678366Z","shell.execute_reply.started":"2025-03-07T09:21:24.660137Z","shell.execute_reply":"2025-03-07T09:21:24.677453Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#pack data into a dictionary\n\ndata={\n    \"sequence\":train_sequences['sequence'].to_list(),\n    \"temporal_cutoff\": train_sequences['temporal_cutoff'].to_list(),\n    \"description\": train_sequences['description'].to_list(),\n    \"all_sequences\": train_sequences['all_sequences'].to_list(),\n    \"xyz\": all_xyz\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-07T09:21:40.862642Z","iopub.execute_input":"2025-03-07T09:21:40.862934Z","iopub.status.idle":"2025-03-07T09:21:40.867384Z","shell.execute_reply.started":"2025-03-07T09:21:40.862911Z","shell.execute_reply":"2025-03-07T09:21:40.866524Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Split train data into train/val/test¶\nWe will simply do a temporal split, because that's how testing is done in structural biology in general (in actual blind tests)","metadata":{}},{"cell_type":"code","source":"# Split data into train and test\nall_index = np.arange(len(data['sequence']))\ncutoff_date = pd.Timestamp(config['cutoff_date'])\ntest_cutoff_date = pd.Timestamp(config['test_cutoff_date'])\ntrain_index = [i for i, d in enumerate(data['temporal_cutoff']) if pd.Timestamp(d) <= cutoff_date]\ntest_index = [i for i, d in enumerate(data['temporal_cutoff']) if pd.Timestamp(d) > cutoff_date and pd.Timestamp(d) <= test_cutoff_date]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-07T09:21:44.461455Z","iopub.execute_input":"2025-03-07T09:21:44.461751Z","iopub.status.idle":"2025-03-07T09:21:44.469157Z","shell.execute_reply.started":"2025-03-07T09:21:44.461729Z","shell.execute_reply":"2025-03-07T09:21:44.468259Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"Train size: {len(train_index)}\")\nprint(f\"Test size: {len(test_index)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-07T09:21:45.095330Z","iopub.execute_input":"2025-03-07T09:21:45.095615Z","iopub.status.idle":"2025-03-07T09:21:45.100441Z","shell.execute_reply.started":"2025-03-07T09:21:45.095593Z","shell.execute_reply":"2025-03-07T09:21:45.099527Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Get pytorch dataset¶","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader\n\nclass RNA3D_Dataset(Dataset):\n    def __init__(self,indices,data):\n        self.indices=indices\n        self.data=data\n        self.tokens={nt:i for i,nt in enumerate('ACGU')}\n\n    def __len__(self):\n        return len(self.indices)\n    \n    def __getitem__(self, idx):\n        idx=self.indices[idx]\n        sequence=[self.tokens[nt] for nt in (self.data['sequence'][idx])]\n        sequence=np.array(sequence)\n        sequence=torch.tensor(sequence)\n\n        #get C1' xyz\n        xyz=self.data['xyz'][idx]\n        xyz=torch.tensor(np.array(xyz))\n\n        if len(sequence)>config['max_len']:\n            crop_start = np.random.randint(len(sequence)-config['max_len'])\n            crop_end = crop_start+config['max_len']\n\n            sequence = sequence[crop_start:crop_end]\n            xyz = xyz[crop_start:crop_end]\n\n        return {\n            'sequence': sequence,\n            'xyz': xyz,\n        }\n\ntrain_dataset = RNA3D_Dataset(train_index,data)\nval_dataset = RNA3D_Dataset(test_index,data)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-07T09:21:48.725154Z","iopub.execute_input":"2025-03-07T09:21:48.725452Z","iopub.status.idle":"2025-03-07T09:21:48.732173Z","shell.execute_reply.started":"2025-03-07T09:21:48.725431Z","shell.execute_reply":"2025-03-07T09:21:48.731403Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Get RibonanzaNet¶\nWe will add a linear layer to predict xyz of C1' atoms","metadata":{}},{"cell_type":"code","source":"import sys\n\nsys.path.append(\"/kaggle/input/ribonanzanet2d-final\")\n\nfrom Network import *\nimport yaml\n\nclass Config:\n    def __init__(self, **entries):\n        self.__dict__.update(entries)\n        self.entries=entries\n\n    def print(self):\n        print(self.entries)\n\ndef load_config_from_yaml(file_path):\n    with open(file_path, 'r') as file:\n        config = yaml.safe_load(file)\n    return Config(**config)\n\nclass finetuned_RibonanzaNet(RibonanzaNet):\n    def __init__(self, config, pretrained=False):\n        config.dropout=0.1\n        super(finetuned_RibonanzaNet, self).__init__(config)\n        if pretrained: self.load_state_dict(torch.load(\"/kaggle/input/ribonanzanet-weights/RibonanzaNet.pt\",map_location='cpu'))\n\n        self.xyz_predictor=nn.Linear(256,3)\n\n    def forward(self, src, mask):\n        sequence_features, pairwise_features=self.get_embeddings(src, mask)\n        xyz=self.xyz_predictor(sequence_features)\n\n        return xyz","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-07T09:21:50.145098Z","iopub.execute_input":"2025-03-07T09:21:50.145419Z","iopub.status.idle":"2025-03-07T09:21:52.021282Z","shell.execute_reply.started":"2025-03-07T09:21:50.145392Z","shell.execute_reply":"2025-03-07T09:21:52.020559Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training loop¶\nwe will use dRMSD loss on the predicted xyz. the loss function is invariant to translations, rotations, and reflections. because dRMSD is invariant to reflections, it cannot distinguish chiral structures, so there may be better loss functions","metadata":{}},{"cell_type":"code","source":"def calculate_distance_matrix(X,Y,epsilon=1e-4):\n    return (torch.square(X[:,None]-Y[None,:])+epsilon).sum(-1).sqrt()\n\ndef dRMSD(pred_x,\n          gt_x,\n          epsilon=1e-4,Z=10,d_clamp=None):\n    pred_dm=calculate_distance_matrix(pred_x,pred_x)\n    gt_dm=calculate_distance_matrix(gt_x,gt_x)\n\n    mask=~torch.isnan(gt_dm)\n    mask[torch.eye(mask.shape[0]).bool()]=False\n\n    if d_clamp is not None:\n        rmsd=(torch.square(pred_dm[mask]-gt_dm[mask])+epsilon).clip(0,d_clamp**2)\n    else:\n        rmsd=torch.square(pred_dm[mask]-gt_dm[mask])+epsilon\n\n    return rmsd.sqrt().mean()/Z\n\ndef local_dRMSD(pred_x,\n          pred_y,\n          gt_x,\n          gt_y,\n          epsilon=1e-4,Z=10,d_clamp=30):\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]).bool()]=False\n\n    rmsd=torch.square(pred_dm[mask]-gt_dm[mask])+epsilon\n    # rmsd=(torch.square(pred_dm[mask]-gt_dm[mask])+epsilon).sqrt()/Z\n    #rmsd=torch.abs(pred_dm[mask]-gt_dm[mask])/Z\n    return rmsd.sqrt().mean()/Z\n\ndef dRMAE(pred_x, gt_x, mask, epsilon=1e-4, Z=10, d_clamp=None):\n    pred_dm = torch.cdist(pred_x, pred_x, p=2) * mask\n    gt_dm = torch.cdist(gt_x, gt_x, p=2) * mask\n\n    mask = ~torch.isnan(gt_dm)  \n    diff = torch.abs(pred_dm - gt_dm)[mask].mean()\n\n    return diff / Z \n\ndef align_svd_mae(input, target, mask_, Z=10):\n    \"\"\"\n    Aligns the input (Nx3) to target (Nx3) using SVD-based Procrustes alignment\n    and computes RMSD loss.\n    \n    Args:\n        input (torch.Tensor): Nx3 tensor representing the input points.\n        target (torch.Tensor): Nx3 tensor representing the target points.\n    \n    Returns:\n        aligned_input (torch.Tensor): Nx3 aligned input.\n        rmsd_loss (torch.Tensor): RMSD loss.\n    \"\"\"\n    assert input.shape == target.shape, \"Input and target must have the same shape\"\n\n    input=input[mask_ == 1]\n    target=target[mask_ == 1]\n    \n    #mask \n    mask=~torch.isnan(target.sum(-1))\n\n    input=input[mask]\n    target=target[mask]\n    \n    # Compute centroids\n    centroid_input = input.mean(dim=0, keepdim=True)\n    centroid_target = target.mean(dim=0, keepdim=True)\n\n    # Center the points\n    input_centered = input - centroid_input.detach()\n    target_centered = target - 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\n    # Compute rotation matrix\n    R = Vt @ U.T\n\n    # Ensure a proper rotation (det(R) = 1, no reflection)\n    if torch.det(R) < 0:\n        Vt[-1, :] *= -1\n        R = Vt @ U.T\n\n    # Rotate input\n    aligned_input = (input_centered @ R.T.detach()) + centroid_target.detach()\n\n    # # Compute RMSD loss\n    # rmsd_loss = torch.sqrt(((aligned_input - target) ** 2).mean())\n\n    # rmsd_loss = torch.sqrt(((aligned_input - target) ** 2).mean())\n    \n    # return aligned_input, rmsd_loss\n    return torch.abs(aligned_input-target).mean()/Z","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-07T09:34:50.135652Z","iopub.execute_input":"2025-03-07T09:34:50.136041Z","iopub.status.idle":"2025-03-07T09:34:50.146899Z","shell.execute_reply.started":"2025-03-07T09:34:50.135992Z","shell.execute_reply":"2025-03-07T09:34:50.145998Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pytorch_lightning as pl\nfrom pytorch_lightning.callbacks import ModelCheckpoint\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\nclass RibonanzaModel(pl.LightningModule):\n    def __init__(self, model, cos_epoch=35, epochs=50):\n        super().__init__()\n        self.model = model\n        self.cos_epoch = cos_epoch\n        self.epochs = epochs\n        self.automatic_optimization = False\n        \n    def training_step(self, batch, batch_idx):\n        sequence = batch['sequence']\n        mask = batch['mask']\n\n        pred_xyz = self.model(sequence, mask)\n        gt_xyz = batch['xyz']\n\n        row = mask.unsqueeze(2).expand(-1, -1, mask.size(1))\n        col = mask.unsqueeze(1).expand(-1, mask.size(1), -1)\n        mask_ = (row & col)\n\n        loss1 = dRMAE(pred_xyz, gt_xyz, mask_)\n\n        with torch.autocast(device_type='cuda', dtype=torch.float32):\n            loss2 = align_svd_mae(pred_xyz[0], gt_xyz[0], mask[0])\n\n        loss = loss1 + loss2\n        self.manual_backward(loss)\n\n        torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1)\n\n        opt = self.optimizers()\n        opt.step()\n        opt.zero_grad()\n        \n        if self.current_epoch >= self.cos_epoch:\n            sch = self.lr_schedulers()\n            sch.step()\n\n        self.log('train_loss', loss, prog_bar=True)\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        sequence = batch['sequence']\n        mask = batch['mask']\n\n        pred_xyz = self.model(sequence, mask)\n        gt_xyz = batch['xyz']\n\n        row = mask.unsqueeze(2).expand(-1, -1, mask.size(1))\n        col = mask.unsqueeze(1).expand(-1, mask.size(1), -1)\n        mask_ = (row & col)\n\n        loss = dRMAE(pred_xyz, gt_xyz, mask_)\n        self.log('val_loss', loss, prog_bar=True)\n        return loss\n\n    def configure_optimizers(self):\n        optimizer = torch.optim.Adam(self.parameters(), lr=0.0001, weight_decay=0.0)\n        scheduler = CosineAnnealingLR(\n            optimizer,\n            T_max=(self.epochs - self.cos_epoch) * self.trainer.estimated_stepping_batches\n        )\n\n        return [optimizer], [{'scheduler': scheduler, 'interval': 'step'}]\n\n    def configure_callbacks(self):\n        return [\n            ModelCheckpoint(monitor='val_loss', filename='best', save_top_k=1),\n            ModelCheckpoint(filename='last', save_last=True),\n        ]\n\nmodel = finetuned_RibonanzaNet(\n    load_config_from_yaml(\"/kaggle/input/ribonanzanet2d-final/configs/pairwise.yaml\"),\n    pretrained=True\n)\n\nbatch_size = 2\n\ndef collate_fn(batch):\n    mlen = max([ len(elem['sequence']) for elem in batch ])\n\n    for elem in batch:\n        mask = torch.zeros(mlen).long()\n        mask[:len(elem['sequence'])] = 1\n\n        elem['mask'] = mask\n        elem['xyz'] = torch.nn.functional.pad(elem['xyz'], (0, 0, 0, mlen-len(elem['sequence'])))\n        elem['sequence'] = torch.nn.functional.pad(elem['sequence'], (0, mlen-len(elem['sequence'])))\n\n    batch_ = {\n        'mask': torch.stack([ elem['mask'] for elem in batch ]),\n        'sequence': torch.stack([ elem['sequence'] for elem in batch ]),\n        'xyz': torch.stack([ elem['xyz'] for elem in batch ]),\n    }\n\n    return batch_\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=batch_size,\n    shuffle=True,\n    collate_fn=collate_fn,\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=batch_size,\n    shuffle=False,\n    collate_fn=collate_fn,\n)\n\ntrainer = pl.Trainer(\n    max_epochs=50,\n    accelerator='gpu',\n    devices=1,\n    precision='bf16-mixed',\n)\n\ntorch.manual_seed(0)\nnp.random.seed(0)\nrandom.seed(0)\n\nplmodel = RibonanzaModel(model)\ntrainer.fit(plmodel, train_loader, val_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-07T09:38:31.974034Z","iopub.execute_input":"2025-03-07T09:38:31.974340Z"}},"outputs":[],"execution_count":null}]}