{"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":7639698,"sourceType":"datasetVersion","datasetId":4299272},{"sourceId":8318191,"sourceType":"datasetVersion","datasetId":4459124},{"sourceId":11316837,"sourceType":"datasetVersion","datasetId":7078712},{"sourceId":230614988,"sourceType":"kernelVersion"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport torch\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport json\nimport torch\nimport random\nimport pickle\nfrom tqdm import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-01T05:54:21.100122Z","iopub.execute_input":"2025-04-01T05:54:21.100405Z","iopub.status.idle":"2025-04-01T05:54:24.601734Z","shell.execute_reply.started":"2025-04-01T05:54:21.100374Z","shell.execute_reply":"2025-04-01T05:54:24.601016Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#set seed for everything\ntorch.manual_seed(0)\nnp.random.seed(0)\nrandom.seed(0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T05:54:24.603183Z","iopub.execute_input":"2025-04-01T05:54:24.603581Z","iopub.status.idle":"2025-04-01T05:54:24.611775Z","shell.execute_reply.started":"2025-04-01T05:54:24.603555Z","shell.execute_reply":"2025-04-01T05:54:24.610991Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"CONTINUE_TRAINING = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T05:54:24.612982Z","iopub.execute_input":"2025-04-01T05:54:24.613258Z","iopub.status.idle":"2025-04-01T05:54:24.627174Z","shell.execute_reply.started":"2025-04-01T05:54:24.613236Z","shell.execute_reply":"2025-04-01T05:54:24.626557Z"}},"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-04-01T05:54:24.627988Z","iopub.execute_input":"2025-04-01T05:54:24.628194Z","iopub.status.idle":"2025-04-01T05:54:24.642352Z","shell.execute_reply.started":"2025-04-01T05:54:24.628175Z","shell.execute_reply":"2025-04-01T05:54:24.641503Z"}},"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\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\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T05:54:24.643172Z","iopub.execute_input":"2025-04-01T05:54:24.643444Z","iopub.status.idle":"2025-04-01T05:54:25.034793Z","shell.execute_reply.started":"2025-04-01T05:54:24.643411Z","shell.execute_reply":"2025-04-01T05:54:25.034113Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_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-04-01T05:54:25.037264Z","iopub.execute_input":"2025-04-01T05:54:25.037487Z","iopub.status.idle":"2025-04-01T05:54:25.126377Z","shell.execute_reply.started":"2025-04-01T05:54:25.037467Z","shell.execute_reply":"2025-04-01T05:54:25.125726Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"float('Nan')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T05:54:25.127563Z","iopub.execute_input":"2025-04-01T05:54:25.127805Z","iopub.status.idle":"2025-04-01T05:54:25.132551Z","shell.execute_reply.started":"2025-04-01T05:54:25.127785Z","shell.execute_reply":"2025-04-01T05:54:25.131897Z"}},"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    #break\n    xyz=df[['x_1','y_1','z_1']].to_numpy().astype('float32')\n    xyz[xyz<-1e17]=float('Nan');\n    all_xyz.append(xyz)\n\n\ndf","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T05:54:25.133308Z","iopub.execute_input":"2025-04-01T05:54:25.133511Z","iopub.status.idle":"2025-04-01T05:54:33.547367Z","shell.execute_reply.started":"2025-04-01T05:54:25.133492Z","shell.execute_reply":"2025-04-01T05:54:33.546697Z"}},"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-04-01T05:54:33.548279Z","iopub.execute_input":"2025-04-01T05:54:33.548607Z","iopub.status.idle":"2025-04-01T05:54:33.567358Z","shell.execute_reply.started":"2025-04-01T05:54:33.548575Z","shell.execute_reply":"2025-04-01T05:54:33.566656Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_json = '/kaggle/input/standfordrna-data-extraction/RNA_Data.json'\nwith open(data_json, 'r') as file:\n    data_loaded = json.load(file)\nprint(\"Data loaded from file\")\n\nxyz_loaded = np.load('/kaggle/input/standfordrna-data-extraction/RNA_Labels.npy', allow_pickle=True)\nlen(xyz_loaded)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T05:54:33.568241Z","iopub.execute_input":"2025-04-01T05:54:33.568458Z","iopub.status.idle":"2025-04-01T05:54:33.630060Z","shell.execute_reply.started":"2025-04-01T05:54:33.568426Z","shell.execute_reply":"2025-04-01T05:54:33.629236Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#pack data into a dictionary\n\ndata={\n      \"sequence\": train_sequences['sequence'].to_list() + data_loaded['sequences'],\n      \"temporal_cutoff\": train_sequences['temporal_cutoff'].to_list() + data_loaded['temporal_cutoffs'],\n      \"description\": train_sequences['description'].to_list() + data_loaded['descriptions'],\n      #\"all_sequences\": train_sequences['all_sequences'].to_list(),\n      \"xyz\": all_xyz + xyz_loaded.tolist(),\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T05:54:33.630930Z","iopub.execute_input":"2025-04-01T05:54:33.631187Z","iopub.status.idle":"2025-04-01T05:54:33.635274Z","shell.execute_reply.started":"2025-04-01T05:54:33.631159Z","shell.execute_reply":"2025-04-01T05:54:33.634423Z"}},"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-04-01T05:54:33.636088Z","iopub.execute_input":"2025-04-01T05:54:33.636319Z","iopub.status.idle":"2025-04-01T05:54:33.660889Z","shell.execute_reply.started":"2025-04-01T05:54:33.636285Z","shell.execute_reply":"2025-04-01T05:54:33.660261Z"}},"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-04-01T05:54:33.661805Z","iopub.execute_input":"2025-04-01T05:54:33.662096Z","iopub.status.idle":"2025-04-01T05:54:33.668587Z","shell.execute_reply.started":"2025-04-01T05:54:33.662067Z","shell.execute_reply":"2025-04-01T05:54:33.667922Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Get pytorch dataset¶","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader\nfrom ast import literal_eval\n\ndef get_ct(bp,s):\n    ct_matrix=np.zeros((len(s),len(s)))\n    for b in bp:\n        ct_matrix[b[0]-1,b[1]-1]=1\n    return ct_matrix\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\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\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\n        return {'sequence':sequence,\n                'xyz':xyz}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T05:54:33.669295Z","iopub.execute_input":"2025-04-01T05:54:33.669545Z","iopub.status.idle":"2025-04-01T05:54:33.678965Z","shell.execute_reply.started":"2025-04-01T05:54:33.669524Z","shell.execute_reply":"2025-04-01T05:54:33.678225Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset=RNA3D_Dataset(train_index,data)\nval_dataset=RNA3D_Dataset(test_index,data)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T05:54:33.679682Z","iopub.execute_input":"2025-04-01T05:54:33.679952Z","iopub.status.idle":"2025-04-01T05:54:33.695922Z","shell.execute_reply.started":"2025-04-01T05:54:33.679920Z","shell.execute_reply":"2025-04-01T05:54:33.695190Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import plotly.graph_objects as go\nimport numpy as np\n\n\n\n# Example: Generate an Nx3 matrix\nxyz = train_dataset[200]['xyz']  # Replace this with your actual Nx3 data\nN = len(xyz)\n\n\nfor _ in range(2): #plot twice because it doesnt show up on first try for some reason\n    # Extract columns\n    x, y, z = xyz[:, 0], xyz[:, 1], xyz[:, 2]\n    \n    # Create the 3D scatter plot\n    fig = go.Figure(data=[go.Scatter3d(\n        x=x, y=y, z=z,\n        mode='markers',\n        marker=dict(\n            size=5,\n            color=z,  # Coloring based on z-value\n            colorscale='Viridis',  # Choose a colorscale\n            opacity=0.8\n        )\n    )])\n    \n    # Customize layout\n    fig.update_layout(\n        scene=dict(\n            xaxis_title=\"X\",\n            yaxis_title=\"Y\",\n            zaxis_title=\"Z\"\n        ),\n        title=\"3D Scatter Plot\"\n    )\n\nfig.show()\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T05:54:33.696700Z","iopub.execute_input":"2025-04-01T05:54:33.696928Z","iopub.status.idle":"2025-04-01T05:54:34.205464Z","shell.execute_reply.started":"2025-04-01T05:54:33.696907Z","shell.execute_reply":"2025-04-01T05:54:34.204664Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_loader=DataLoader(train_dataset,batch_size=1,shuffle=True)\nval_loader=DataLoader(val_dataset,batch_size=1,shuffle=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T05:54:34.206140Z","iopub.execute_input":"2025-04-01T05:54:34.206394Z","iopub.status.idle":"2025-04-01T05:54:34.210334Z","shell.execute_reply.started":"2025-04-01T05:54:34.206358Z","shell.execute_reply":"2025-04-01T05:54:34.209698Z"}},"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\nimport torch.nn as nn\nfrom Network import RibonanzaNet, MultiHeadAttention\nimport yaml\n\nclass SimpleStructureModule(nn.Module):\n\n    def __init__(self, d_model, nhead, \n                 dim_feedforward, pairwise_dimension, dropout=0.1,\n                 ):\n        super(SimpleStructureModule, self).__init__()\n        #self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)\n        self.self_attn = MultiHeadAttention(d_model, nhead, d_model//nhead, d_model//nhead, dropout=dropout)\n\n\n        self.linear1 = nn.Linear(d_model, dim_feedforward)\n        self.dropout = nn.Dropout(dropout)\n        self.linear2 = nn.Linear(dim_feedforward, d_model)\n\n        self.norm1 = nn.LayerNorm(d_model)\n        self.norm2 = nn.LayerNorm(d_model)\n        self.norm3 = nn.LayerNorm(d_model)\n        #self.norm4 = nn.LayerNorm(d_model)\n        self.dropout1 = nn.Dropout(dropout)\n        self.dropout2 = nn.Dropout(dropout)\n        self.dropout3 = nn.Dropout(dropout)\n        #self.dropout4 = nn.Dropout(dropout)\n\n        self.pairwise2heads=nn.Linear(pairwise_dimension,nhead,bias=False)\n        self.pairwise_norm=nn.LayerNorm(pairwise_dimension)\n\n        self.distance2heads=nn.Linear(1,nhead,bias=False)\n        #self.pairwise_norm=nn.LayerNorm(pairwise_dimension)\n\n        self.activation = nn.GELU()\n\n        \n    def custom(self, module):\n        def custom_forward(*inputs):\n            inputs = module(*inputs)\n            return inputs\n        return custom_forward\n\n    def forward(self, input):\n        src , pairwise_features, pred_t, src_mask = input\n        \n        #src = src*src_mask.float().unsqueeze(-1)\n\n        pairwise_bias=self.pairwise2heads(self.pairwise_norm(pairwise_features)).permute(0,3,1,2)\n\n        \n        distance_matrix=pred_t[None,:,:]-pred_t[:,None,:]\n        distance_matrix=(distance_matrix**2).sum(-1).clip(2,37**2).sqrt()\n        distance_matrix=distance_matrix[None,:,:,None]\n        distance_bias=self.distance2heads(distance_matrix).permute(0,3,1,2)\n\n                    \n        \n        pairwise_bias=pairwise_bias+distance_bias\n\n        #print(src.shape)\n        src2,attention_weights = self.self_attn(src, src, src, mask=pairwise_bias, src_mask=src_mask)\n        \n\n        src = src + self.dropout1(src2)\n        src = self.norm1(src)\n        src2 = self.linear2(self.dropout(self.activation(self.linear1(src))))\n        src = src + self.dropout2(src2)\n        src = self.norm2(src)\n\n\n        return src\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\n\n\nclass finetuned_RibonanzaNet(RibonanzaNet):\n    def __init__(self, config, pretrained=False):\n        config.dropout=0.1\n        config.use_grad_checkpoint=True\n        super(finetuned_RibonanzaNet, self).__init__(config)\n        if pretrained:\n            self.load_state_dict(torch.load(\"/kaggle/input/ribonanzanet-weights/RibonanzaNet.pt\",map_location='cpu'))\n        # self.ct_predictor=nn.Sequential(nn.Linear(64,256),\n        #                                 nn.ReLU(),\n        #                                 nn.Linear(256,64),\n        #                                 nn.ReLU(),\n        #                                 nn.Linear(64,1)) \n        self.dropout=nn.Dropout(0.0)\n\n        self.structure_module=SimpleStructureModule(d_model=256, nhead=8, \n                 dim_feedforward=1024, pairwise_dimension=64)\n        \n        self.xyz_predictor=nn.Linear(256,3)\n\n    def custom(self, module):\n        def custom_forward(*inputs):\n            inputs = module(*inputs)\n            return inputs\n        return custom_forward\n    \n    def forward(self,src):\n        \n        #with torch.no_grad():\n        sequence_features, pairwise_features=self.get_embeddings(src, torch.ones_like(src).long().to(src.device))\n        \n        xyzs=[]\n        xyz=torch.zeros(sequence_features.shape[1],3).cuda().float()\n        #print(xyz.shape)\n        #xyz=self.xyz_predictor(sequence_features)\n\n        for i in range(9):\n            sequence_features=self.structure_module([sequence_features,pairwise_features,xyz,None])\n            xyz=xyz+self.xyz_predictor(sequence_features).squeeze(0)\n            xyzs.append(xyz)\n            \n        \n        return xyzs\n\nmodel=finetuned_RibonanzaNet(load_config_from_yaml(\"/kaggle/input/ribonanzanet2d-final/configs/pairwise.yaml\"),pretrained=True).cuda()\n#model(torch.ones(1,10).long().cuda())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T05:54:34.213415Z","iopub.execute_input":"2025-04-01T05:54:34.213627Z","iopub.status.idle":"2025-04-01T05:54:37.276491Z","shell.execute_reply.started":"2025-04-01T05:54:34.213609Z","shell.execute_reply":"2025-04-01T05:54:37.275519Z"}},"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\n\ndef dRMSD(pred_x,\n          pred_y,\n          gt_x,\n          gt_y,\n          epsilon=1e-4,Z=10,d_clamp=None):\n    pred_dm=calculate_distance_matrix(pred_x,pred_y)\n    gt_dm=calculate_distance_matrix(gt_x,gt_y)\n\n\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\n\n    mask=(~torch.isnan(gt_dm))*(gt_dm<d_clamp)\n    mask[torch.eye(mask.shape[0]).bool()]=False\n\n\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,\n          pred_y,\n          gt_x,\n          gt_y,\n          epsilon=1e-4,Z=10,d_clamp=None):\n    pred_dm=calculate_distance_matrix(pred_x,pred_y)\n    gt_dm=calculate_distance_matrix(gt_x,gt_y)\n\n\n\n    mask=~torch.isnan(gt_dm)\n    mask[torch.eye(mask.shape[0]).bool()]=False\n\n    rmsd=torch.abs(pred_dm[mask]-gt_dm[mask])\n\n    return rmsd.mean()/Z\n\nimport torch\n\ndef align_svd_mae(input, target, 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    #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-04-01T05:54:37.277584Z","iopub.execute_input":"2025-04-01T05:54:37.277980Z","iopub.status.idle":"2025-04-01T05:54:37.288816Z","shell.execute_reply.started":"2025-04-01T05:54:37.277955Z","shell.execute_reply":"2025-04-01T05:54:37.288032Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#pred_xyz=model(sequence)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T05:54:37.290355Z","iopub.execute_input":"2025-04-01T05:54:37.290676Z","iopub.status.idle":"2025-04-01T05:54:37.310716Z","shell.execute_reply.started":"2025-04-01T05:54:37.290654Z","shell.execute_reply":"2025-04-01T05:54:37.309997Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"start_epoch = 60\noptimizer = torch.optim.Adam(model.parameters(), weight_decay=0.001, lr=0.01) #no weight decay following AF\n\nif CONTINUE_TRAINING:\n    checkpoint_path = \"/kaggle/input/srna3df-data-checkpoints/RibonanzaNet-3D_final.pt\"\n    if torch.cuda.is_available():\n        checkpoint = torch.load(checkpoint_path)  # Load on GPU if available\n    else:\n        checkpoint = torch.load(checkpoint_path, map_location=torch.device('cpu'))  # Load on CPU\n    \n    # Restore model and optimizer state\n    model.load_state_dict(checkpoint[\"model_state_dict\"])\n    optimizer.load_state_dict(checkpoint[\"optimizer_state_dict\"])\n    start_epoch = checkpoint[\"epoch\"]\n    loss = checkpoint[\"loss\"]\n    \n    print(f\"Resuming training from Epoch {start_epoch}, Loss: {loss}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T05:54:37.311529Z","iopub.execute_input":"2025-04-01T05:54:37.311740Z","iopub.status.idle":"2025-04-01T05:54:37.329336Z","shell.execute_reply.started":"2025-04-01T05:54:37.311710Z","shell.execute_reply":"2025-04-01T05:54:37.328487Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\nfrom torch.amp import GradScaler\n\nepochs=60 + start_epoch\ncos_epoch=20 + start_epoch\n\n\nbest_loss=np.inf\n\nbatch_size=1\n\n#for cycle in range(2):\n\ncriterion=torch.nn.BCEWithLogitsLoss(reduction='none')\n\nscaler = GradScaler()\n\nschedule=torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=(epochs-cos_epoch)*len(train_loader)//batch_size)\n\nbest_val_loss=99999999999\nfor epoch in range(start_epoch, epochs):\n    model.train()\n    tbar=tqdm(train_loader)\n    total_loss=0\n    oom=0\n    for idx, batch in enumerate(tbar):\n        #try:\n        sequence=batch['sequence'].cuda()\n        gt_xyz=batch['xyz'].cuda().squeeze().to(torch.float32)\n        #with torch.autocast(device_type='cuda', dtype=torch.float16):\n        pred_xyzs=model(sequence)#.squeeze()\n\n        loss=0\n        for pred_xyz in pred_xyzs:\n            loss+=dRMAE(pred_xyz,pred_xyz,gt_xyz,gt_xyz) \n            loss+=align_svd_mae(pred_xyz, gt_xyz)\n             #local_dRMSD(pred_xyz,pred_xyz,gt_xyz,gt_xyz)\n\n        if loss!=loss:\n            stop\n\n        \n        (loss/batch_size).backward()\n\n        if (idx+1)%batch_size==0 or idx+1 == len(tbar):\n\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 1)\n            optimizer.step()\n            optimizer.zero_grad()\n            # scaler.scale(loss/batch_size).backward()\n            # scaler.unscale_(optimizer)\n            # torch.nn.utils.clip_grad_norm_(model.parameters(), 1)\n            # scaler.step(optimizer)\n            # scaler.update()\n\n            \n            if (epoch+1)>cos_epoch:\n                schedule.step()\n        #schedule.step()\n        total_loss+=loss.item()\n        \n        tbar.set_description(f\"Epoch {epoch + 1} Loss: {total_loss/(idx+1)} OOMs: {oom}\")\n\n\n\n        # except Exception:\n        #     #print(Exception)\n        #     oom+=1\n    tbar=tqdm(val_loader)\n    model.eval()\n    val_preds=[]\n    val_loss=0\n    for idx, batch in enumerate(tbar):\n        sequence=batch['sequence'].cuda()\n        gt_xyz=batch['xyz'].cuda().squeeze().to(torch.float32)\n\n        with torch.no_grad():\n            pred_xyz=model(sequence)[-1].squeeze()\n            loss=dRMAE(pred_xyz,pred_xyz,gt_xyz,gt_xyz)\n            \n        val_loss+=loss.item()\n        val_preds.append([gt_xyz.cpu().numpy(),pred_xyz.cpu().numpy()])\n    val_loss=val_loss/len(tbar)\n    print(f\"val loss: {val_loss}\")\n    \n    \n    \n    if val_loss<best_val_loss:\n        best_val_loss=val_loss\n        best_preds=val_preds\n        torch.save({\n            \"epoch\": epoch + 1,\n            \"model_state_dict\": model.state_dict(),\n            \"optimizer_state_dict\": optimizer.state_dict(),\n            \"loss\": loss.item(),\n        }, \"RibonanzaNet-3D_{}.pt\".format(epoch))\n\ntorch.save({\n    \"epoch\": epochs,\n    \"model_state_dict\": model.state_dict(),\n    \"optimizer_state_dict\": optimizer.state_dict(),\n    \"loss\": loss.item(),\n}, 'RibonanzaNet-3D_final.pt')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-01T05:54:37.330212Z","iopub.execute_input":"2025-04-01T05:54:37.330471Z","execution_failed":"2025-04-01T06:05:34.365Z"}},"outputs":[],"execution_count":null}]}