{"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":10855324,"sourceType":"datasetVersion","datasetId":6742586},{"sourceId":11278691,"sourceType":"datasetVersion","datasetId":7051341},{"sourceId":11839907,"sourceType":"datasetVersion","datasetId":7438872},{"sourceId":11469248,"sourceType":"datasetVersion","datasetId":7187409},{"sourceId":311741,"sourceType":"modelInstanceVersion","modelInstanceId":264400,"modelId":285488}],"dockerImageVersionId":30919,"isInternetEnabled":false,"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 torch\nimport random\nimport pickle\nimport os\nimport sys\nfrom tqdm import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-16T18:49:35.470843Z","iopub.execute_input":"2025-05-16T18:49:35.471118Z","iopub.status.idle":"2025-05-16T18:49:39.058550Z","shell.execute_reply.started":"2025-05-16T18:49:35.471096Z","shell.execute_reply":"2025-05-16T18:49:39.057535Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"config = {\n    \"seed\": 0,\n    \"cutoff_date\": \"2020-01-01\",\n    \"test_cutoff_date\": \"2022-05-01\",\n    \"max_len\": 256,\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\": 1,\n    \"cos_epoch\": 0,\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    \"n_times\": 1000,\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T18:49:39.059862Z","iopub.execute_input":"2025-05-16T18:49:39.060431Z","iopub.status.idle":"2025-05-16T18:49:39.065246Z","shell.execute_reply.started":"2025-05-16T18:49:39.060390Z","shell.execute_reply":"2025-05-16T18:49:39.064377Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# get data","metadata":{}},{"cell_type":"code","source":"# train_sequences=pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_sequences.v2.csv\")\n# train_labels=pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_labels.v2.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T18:49:47.692920Z","iopub.execute_input":"2025-05-16T18:49:47.693232Z","iopub.status.idle":"2025-05-16T18:49:47.697023Z","shell.execute_reply.started":"2025-05-16T18:49:47.693207Z","shell.execute_reply":"2025-05-16T18:49:47.695890Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# train_labels[\"pdb_id\"] = train_labels[\"ID\"].apply(lambda x: x.split(\"_\")[0]+'_'+x.split(\"_\")[1])\n# train_labels[\"pdb_id\"] \n\n\n# # In[6]:\n\n\n# float('Nan')\n\n\n# # In[7]:\n\n\n# all_xyz=[]\n\n# for 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\n# df\n\n\n# # In[8]:\n\n\n# # filter the data\n# # Filter and process data\n# filter_nan = []\n# max_len = 0\n# for 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\n# print(f\"Longest sequence in train: {max_len}\")\n\n# filter_nan = np.array(filter_nan)\n# non_nan_indices = np.arange(len(filter_nan))[filter_nan]\n\n# train_sequences = train_sequences.loc[non_nan_indices].reset_index(drop=True)\n# all_xyz=[all_xyz[i] for i in non_nan_indices]\n\n\n# # In[9]:\n\n\n# #pack data into a dictionary\n\n# data={\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# }\n\n#laod preprocessed pickle file\nwith open(\"/kaggle/input/stanford-3d-train-v2-pickle/train_data.pkl\", \"rb\") as f:\n    data = pickle.load(f)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T18:49:48.721746Z","iopub.execute_input":"2025-05-16T18:49:48.722061Z","iopub.status.idle":"2025-05-16T18:49:50.081957Z","shell.execute_reply.started":"2025-05-16T18:49:48.722034Z","shell.execute_reply":"2025-05-16T18:49:50.081213Z"}},"outputs":[],"execution_count":null},{"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]\n\n\n# In[11]:\n\n\nprint(f\"Train size: {len(train_index)}\")\nprint(f\"Test size: {len(test_index)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T18:49:52.547976Z","iopub.execute_input":"2025-05-16T18:49:52.548288Z","iopub.status.idle":"2025-05-16T18:49:52.571209Z","shell.execute_reply.started":"2025-05-16T18:49:52.548265Z","shell.execute_reply":"2025-05-16T18:49:52.570370Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader\nfrom ast import literal_eval\n\n\nfrom collections import defaultdict\n\nclass RNA3D_Dataset(Dataset):\n    def __init__(self,indices,data):\n        self.indices=indices\n        self.data=data\n        #set default to 4\n        self.tokens=defaultdict(lambda: 4)\n        self.tokens['A']=0\n        self.tokens['C']=1\n        self.tokens['G']=2\n        self.tokens['U']=3\n\n        #{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        #center at first atom if first atom does not exit go until it does\n        for i in range(len(xyz)):\n            if (~torch.isnan(xyz[i])).all():\n                break\n        xyz=xyz-xyz[i]\n\n        # for i in range(len(xyz)):\n\n        #     if torch.isnan(xyz[i]).any():\n        #         if i==0:\n        #             xyz[i]=xyz[i+1]\n        #         else:\n        #             xyz[i]=xyz[i-1]\n\n        return {'sequence':sequence,\n                'xyz':xyz}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T18:49:55.616537Z","iopub.execute_input":"2025-05-16T18:49:55.616835Z","iopub.status.idle":"2025-05-16T18:49:55.624216Z","shell.execute_reply.started":"2025-05-16T18:49:55.616815Z","shell.execute_reply":"2025-05-16T18:49:55.623272Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset=RNA3D_Dataset(train_index,data)\nval_dataset=RNA3D_Dataset(test_index,data)\n\ntrain_loader=DataLoader(train_dataset,batch_size=1,shuffle=True)\nval_loader=DataLoader(val_dataset,batch_size=1,shuffle=False)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T18:50:00.303726Z","iopub.execute_input":"2025-05-16T18:50:00.304009Z","iopub.status.idle":"2025-05-16T18:50:00.308675Z","shell.execute_reply.started":"2025-05-16T18:50:00.303988Z","shell.execute_reply":"2025-05-16T18:50:00.307764Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Network","metadata":{}},{"cell_type":"code","source":"sys.path.append(\"/kaggle/input/ribonanzanet2/pytorch/alpha/1\")\n\nimport torch.nn as nn\nfrom Network import *\n\nclass SinusoidalPosEmb(nn.Module):\n    def __init__(self, dim):\n        super().__init__()\n        self.dim = dim\n\n    def forward(self, x):\n        device = x.device\n        half_dim = self.dim // 2\n        emb = math.log(10000) / (half_dim - 1)\n        emb = torch.exp(torch.arange(half_dim, device=device) * -emb)\n        emb = x[:, None] * emb[None, :]\n        emb = torch.cat((emb.sin(), emb.cos()), dim=-1)\n        return emb\n\nclass finetuned_RibonanzaNet(RibonanzaNet):\n    def __init__(self, rnet_config, config, pretrained=False):\n        rnet_config.dropout=0.1\n        rnet_config.use_grad_checkpoint=True\n        super(finetuned_RibonanzaNet, self).__init__(rnet_config)\n        if pretrained:\n            self.load_state_dict(torch.load(config.pretrained_weight_path,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        decoder_dim=config.decoder_dim\n        self.structure_module=[SimpleStructureModule(d_model=decoder_dim, nhead=config.decoder_nhead, \n                 dim_feedforward=decoder_dim*4, pairwise_dimension=rnet_config.pairwise_dimension, dropout=0.0) for i in range(config.decoder_num_layers)]\n        self.structure_module=nn.ModuleList(self.structure_module)\n\n        self.xyz_embedder=nn.Linear(3,decoder_dim)\n        self.xyz_norm=nn.LayerNorm(decoder_dim)\n        self.xyz_predictor=nn.Linear(decoder_dim,3)\n        \n        self.adaptor=nn.Sequential(nn.Linear(rnet_config.ninp,decoder_dim),nn.LayerNorm(decoder_dim))\n\n        self.distogram_predictor=nn.Sequential(nn.LayerNorm(rnet_config.pairwise_dimension),\n                                                nn.Linear(rnet_config.pairwise_dimension,40))\n\n        self.time_embedder=SinusoidalPosEmb(decoder_dim)\n\n        self.time_mlp=nn.Sequential(nn.Linear(decoder_dim,decoder_dim),\n                                    nn.ReLU(),  \n                                    nn.Linear(decoder_dim,decoder_dim))\n        self.time_norm=nn.LayerNorm(decoder_dim)\n\n        self.distance2pairwise=nn.Linear(1,rnet_config.pairwise_dimension,bias=False)\n\n        self.pair_mlp=nn.Sequential(nn.Linear(rnet_config.pairwise_dimension,rnet_config.pairwise_dimension),\n                                    nn.ReLU(),\n                                    nn.Linear(rnet_config.pairwise_dimension,rnet_config.pairwise_dimension))\n\n\n        #hyperparameters for diffusion\n        self.n_times = config.n_times\n\n        #self.model = model\n        \n        # define linear variance schedule(betas)\n        beta_1, beta_T = config.beta_min, config.beta_max\n        betas = torch.linspace(start=beta_1, end=beta_T, steps=config.n_times)#.to(device) # follows DDPM paper\n        self.sqrt_betas = torch.sqrt(betas)\n                                     \n        # define alpha for forward diffusion kernel\n        self.alphas = 1 - betas\n        self.sqrt_alphas = torch.sqrt(self.alphas)\n        alpha_bars = torch.cumprod(self.alphas, dim=0)\n        self.sqrt_one_minus_alpha_bars = torch.sqrt(1-alpha_bars)\n        self.sqrt_alpha_bars = torch.sqrt(alpha_bars)\n\n        self.data_std=config.data_std\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 embed_pair_distance(self,inputs):\n        pairwise_features,xyz=inputs\n        distance_matrix=xyz[:,None,:,:]-xyz[:,:,None,:]\n        distance_matrix=(distance_matrix**2).sum(-1).clip(2,37**2).sqrt()\n        distance_matrix=distance_matrix[:,:,:,None]\n        pairwise_features=pairwise_features+self.distance2pairwise(distance_matrix)\n\n        return pairwise_features\n\n    def forward(self,src,xyz,t):\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        distogram=self.distogram_predictor(pairwise_features)\n\n        sequence_features=self.adaptor(sequence_features)\n\n        decoder_batch_size=xyz.shape[0]\n        sequence_features=sequence_features.repeat(decoder_batch_size,1,1)\n        \n\n        pairwise_features=pairwise_features.expand(decoder_batch_size,-1,-1,-1)\n\n        pairwise_features= checkpoint.checkpoint(self.custom(self.embed_pair_distance), [pairwise_features,xyz],use_reentrant=False)\n\n        time_embed=self.time_embedder(t).unsqueeze(1)\n        tgt=self.xyz_norm(sequence_features+self.xyz_embedder(xyz)+time_embed)\n\n        tgt=self.time_norm(tgt+self.time_mlp(tgt))\n\n        for layer in self.structure_module:\n            #tgt=layer([tgt, sequence_features,pairwise_features,xyz,None])\n            tgt=checkpoint.checkpoint(self.custom(layer),\n            [tgt, sequence_features,pairwise_features,xyz,None],\n            use_reentrant=False)\n            # xyz=xyz+self.xyz_predictor(sequence_features).squeeze(0)\n            # xyzs.append(xyz)\n            #print(sequence_features.shape)\n        \n        xyz=self.xyz_predictor(tgt).squeeze(0)\n        #.squeeze(0)\n\n        return xyz, distogram\n    \n\n    def denoise(self,sequence_features,pairwise_features,xyz,t):\n        decoder_batch_size=xyz.shape[0]\n        sequence_features=sequence_features.expand(decoder_batch_size,-1,-1)\n        pairwise_features=pairwise_features.expand(decoder_batch_size,-1,-1,-1)\n\n        pairwise_features=self.embed_pair_distance([pairwise_features,xyz])\n\n        sequence_features=self.adaptor(sequence_features)\n        time_embed=self.time_embedder(t).unsqueeze(1)\n        tgt=self.xyz_norm(sequence_features+self.xyz_embedder(xyz)+time_embed)\n        tgt=self.time_norm(tgt+self.time_mlp(tgt))\n        #xyz_batch_size=xyz.shape[0]\n        \n\n\n        for layer in self.structure_module:\n            tgt=layer([tgt, sequence_features,pairwise_features,xyz,None])\n            # xyz=xyz+self.xyz_predictor(sequence_features).squeeze(0)\n            # xyzs.append(xyz)\n            #print(sequence_features.shape)\n        xyz=self.xyz_predictor(tgt).squeeze(0)\n        # print(xyz.shape)\n        # exit()\n        return xyz\n\n\n    def extract(self, a, t, x_shape):\n        \"\"\"\n            from lucidrains' implementation\n                https://github.com/lucidrains/denoising-diffusion-pytorch/blob/beb2f2d8dd9b4f2bd5be4719f37082fe061ee450/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py#L376\n        \"\"\"\n        b, *_ = t.shape\n        out = a.gather(-1, t)\n        return out.reshape(b, *((1,) * (len(x_shape) - 1)))\n    \n    def scale_to_minus_one_to_one(self, x):\n        # according to the DDPMs paper, normalization seems to be crucial to train reverse process network\n        return x * 2 - 1\n    \n    def reverse_scale_to_zero_to_one(self, x):\n        return (x + 1) * 0.5\n    \n    def make_noisy(self, x_zeros, t): \n        # assume we get raw data, so center and scale by 35\n        x_zeros = x_zeros - torch.nanmean(x_zeros,1,keepdim=True)\n        x_zeros = x_zeros/self.data_std\n        #rotate randomly\n        x_zeros = random_rotation_point_cloud_torch_batch(x_zeros)\n\n\n        # perturb x_0 into x_t (i.e., take x_0 samples into forward diffusion kernels)\n        epsilon = torch.randn_like(x_zeros).to(x_zeros.device)\n        \n        sqrt_alpha_bar = self.extract(self.sqrt_alpha_bars.to(x_zeros.device), t, x_zeros.shape)\n        sqrt_one_minus_alpha_bar = self.extract(self.sqrt_one_minus_alpha_bars.to(x_zeros.device), t, x_zeros.shape)\n        \n        # Let's make noisy sample!: i.e., Forward process with fixed variance schedule\n        #      i.e., sqrt(alpha_bar_t) * x_zero + sqrt(1-alpha_bar_t) * epsilon\n        noisy_sample = x_zeros * sqrt_alpha_bar + epsilon * sqrt_one_minus_alpha_bar\n    \n        return noisy_sample.detach(), epsilon\n    \n    \n    # def forward(self, x_zeros):\n    #     x_zeros = self.scale_to_minus_one_to_one(x_zeros)\n        \n    #     B, _, _, _ = x_zeros.shape\n        \n    #     # (1) randomly choose diffusion time-step\n    #     t = torch.randint(low=0, high=self.n_times, size=(B,)).long().to(x_zeros.device)\n        \n    #     # (2) forward diffusion process: perturb x_zeros with fixed variance schedule\n    #     perturbed_images, epsilon = self.make_noisy(x_zeros, t)\n        \n    #     # (3) predict epsilon(noise) given perturbed data at diffusion-timestep t.\n    #     pred_epsilon = self.model(perturbed_images, t)\n        \n    #     return perturbed_images, epsilon, pred_epsilon\n    \n    \n    def denoise_at_t(self, x_t, sequence_features, pairwise_features, timestep, t):\n        B, _, _ = x_t.shape\n        if t > 1:\n            z = torch.randn_like(x_t).to(sequence_features.device)\n        else:\n            z = torch.zeros_like(x_t).to(sequence_features.device)\n        \n        # at inference, we use predicted noise(epsilon) to restore perturbed data sample.\n        epsilon_pred = self.denoise(sequence_features, pairwise_features, x_t, timestep)\n        \n        alpha = self.extract(self.alphas.to(x_t.device), timestep, x_t.shape)\n        sqrt_alpha = self.extract(self.sqrt_alphas.to(x_t.device), timestep, x_t.shape)\n        sqrt_one_minus_alpha_bar = self.extract(self.sqrt_one_minus_alpha_bars.to(x_t.device), timestep, x_t.shape)\n        sqrt_beta = self.extract(self.sqrt_betas.to(x_t.device), timestep, x_t.shape)\n        \n        # denoise at time t, utilizing predicted noise\n        x_t_minus_1 = 1 / sqrt_alpha * (x_t - (1-alpha)/sqrt_one_minus_alpha_bar*epsilon_pred) + sqrt_beta*z\n        \n        return x_t_minus_1#.clamp(-1., 1)\n                \n    def sample(self, src, N):\n        # start from random noise vector, NxLx3\n        x_t = torch.randn((N, src.shape[1], 3)).to(src.device)\n        \n        # autoregressively denoise from x_T to x_0\n        #     i.e., generate image from noise, x_T\n\n        #first get conditioning\n        sequence_features, pairwise_features=self.get_embeddings(src, torch.ones_like(src).long().to(src.device))\n        # sequence_features=sequence_features.expand(N,-1,-1)\n        # pairwise_features=pairwise_features.expand(N,-1,-1,-1)\n        distogram=self.distogram_predictor(pairwise_features).squeeze()\n        distogram=distogram.squeeze()[:,:,2:40]*torch.arange(2,40).float().cuda() \n        distogram=distogram.sum(-1)  \n\n        for t in range(self.n_times-1, -1, -1):\n            timestep = torch.tensor([t]).repeat_interleave(N, dim=0).long().to(src.device)\n            x_t = self.denoise_at_t(x_t, sequence_features, pairwise_features, timestep, t)\n        \n        # denormalize x_0 into 0 ~ 1 ranged values.\n        #x_0 = self.reverse_scale_to_zero_to_one(x_t)\n        x_0 = x_t * self.data_std\n        return x_0, distogram\n\n\n\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        #self.cross_attn = MultiHeadAttention(d_model, nhead, d_model//nhead, d_model//nhead, dropout=dropout)\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.dropout1 = nn.Dropout(dropout)\n        self.dropout2 = 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        tgt , 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\n\n        #print(pairwise_bias.shape,distance_bias.shape)\n\n        #pairwise_bias=pairwise_bias+distance_bias\n\n\n        res=tgt\n        tgt,attention_weights = self.self_attn(tgt, tgt, tgt, mask=pairwise_bias, src_mask=src_mask)\n        tgt = res + self.dropout1(tgt)\n        tgt = self.norm1(tgt)\n\n        # print(tgt.shape,src.shape)\n        # exit()\n\n        res=tgt\n        tgt = self.linear2(self.dropout(self.activation(self.linear1(tgt))))\n        tgt = res + self.dropout2(tgt)\n        tgt = self.norm2(tgt)\n\n\n        return tgt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T18:51:15.879689Z","iopub.execute_input":"2025-05-16T18:51:15.880011Z","iopub.status.idle":"2025-05-16T18:51:15.907427Z","shell.execute_reply.started":"2025-05-16T18:51:15.879986Z","shell.execute_reply":"2025-05-16T18:51:15.906610Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import 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\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T18:51:22.778851Z","iopub.execute_input":"2025-05-16T18:51:22.779292Z","iopub.status.idle":"2025-05-16T18:51:22.805859Z","shell.execute_reply.started":"2025-05-16T18:51:22.779256Z","shell.execute_reply":"2025-05-16T18:51:22.805032Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def random_rotation_point_cloud_torch_batch(point_clouds):\n    \"\"\"\n    Apply a random 3D rotation to a batch of point clouds (PyTorch version).\n    \n    Args:\n        point_clouds (torch.Tensor): BxNx3 tensor of XYZ points.\n\n    Returns:\n        torch.Tensor: Rotated BxNx3 point clouds.\n    \"\"\"\n    B, N, _ = point_clouds.shape\n    device = point_clouds.device\n\n    # Generate a batch of random orthonormal rotation matrices\n    A = torch.randn(B, 3, 3, device=device)\n    Q, R = torch.linalg.qr(A)\n\n    # Ensure det(Q) = +1 for proper rotation\n    det = torch.det(Q)\n    Q[det < 0, :, 0] *= -1\n\n    # Apply batched matrix multiplication\n    rotated = torch.matmul(point_clouds, Q.transpose(1, 2))  # (B, N, 3) x (B, 3, 3)^T -> (B, N, 3)\n\n    return rotated\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T18:51:32.667522Z","iopub.execute_input":"2025-05-16T18:51:32.667801Z","iopub.status.idle":"2025-05-16T18:51:32.672518Z","shell.execute_reply.started":"2025-05-16T18:51:32.667780Z","shell.execute_reply":"2025-05-16T18:51:32.671584Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"diffusion_config=load_config_from_yaml(\"/kaggle/input/ribonanzanet2-ddpm-v2/diffusion_config.yaml\")\nrnet_config=load_config_from_yaml(\"/kaggle/input/ribonanzanet2/pytorch/alpha/1/pairwise.yaml\")\n\nmodel=finetuned_RibonanzaNet(rnet_config,diffusion_config).cuda()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T18:52:51.805349Z","iopub.execute_input":"2025-05-16T18:52:51.805667Z","iopub.status.idle":"2025-05-16T18:52:54.405129Z","shell.execute_reply.started":"2025-05-16T18:52:51.805640Z","shell.execute_reply":"2025-05-16T18:52:54.404484Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training loop","metadata":{}},{"cell_type":"code","source":"from tqdm import tqdm\n\nepochs=config['epochs']\ncos_epoch=config['cos_epoch']\n\n\nbest_loss=np.inf\noptimizer = torch.optim.Adam(model.parameters(), weight_decay=0.0, lr=0.0002) #no weight decay following AF\n\nbatch_size=1\n\n#for cycle in range(2):\n\ncriterion=torch.nn.CrossEntropyLoss(reduction='none')\n\n#scaler = GradScaler()\n\n\nschedule=torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=(epochs-cos_epoch)*len(train_loader)//batch_size)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T18:52:56.284487Z","iopub.execute_input":"2025-05-16T18:52:56.284803Z","iopub.status.idle":"2025-05-16T18:52:56.297811Z","shell.execute_reply.started":"2025-05-16T18:52:56.284778Z","shell.execute_reply":"2025-05-16T18:52:56.296905Z"}},"outputs":[],"execution_count":null},{"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 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\ndef align_svd_rmsd(input, target):\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.square(aligned_input-target).mean().sqrt()\n\ndef compute_lddt(ground_truth_atoms, predicted_atoms, cutoff=30.0, thresholds=[1.0, 2.0, 4.0, 8.0]):\n    \"\"\"\n    Computes the lDDT score between ground truth and predicted atoms.\n    \n    Parameters:\n        ground_truth_atoms (np.array): Nx3 array of ground truth atom coordinates.\n        predicted_atoms (np.array): Nx3 array of predicted atom coordinates.\n        cutoff (float): Distance cutoff in Ångstroms to consider neighbors. Default is 30 Å.\n        thresholds (list): List of thresholds in Ångstroms for the lDDT computation. Default is [0.5, 1.0, 2.0, 4.0].\n    \n    Returns:\n        float: The lDDT score.\n    \"\"\"\n    # Number of atoms\n    num_atoms = ground_truth_atoms.shape[0]\n    \n    # Initialize array to store lDDT fractions for each threshold\n    fractions = np.zeros(len(thresholds))\n    \n    for i in range(num_atoms):\n        # Get the distances from atom i to all other atoms for both ground truth and predicted atoms\n        gt_distances = np.linalg.norm(ground_truth_atoms[i] - ground_truth_atoms, axis=1)\n        pred_distances = np.linalg.norm(predicted_atoms[i] - predicted_atoms, axis=1)\n        \n        # print(gt_distances)\n        # print(pred_distances)\n        # exit()\n        # Apply the cutoff to consider only distances within the cutoff range\n        mask = (gt_distances > 0) & (gt_distances < cutoff)\n        \n        # Calculate the absolute difference between ground truth and predicted distances\n        distance_diff = np.abs(gt_distances[mask] - pred_distances[mask])\n\n        # Filter out any NaN values from the distance difference calculation\n        valid_mask = ~np.isnan(distance_diff)\n        distance_diff = distance_diff[valid_mask]\n\n        # Compute the fractions for each threshold\n        for j, threshold in enumerate(thresholds):\n            if len(distance_diff)>0:\n                fractions[j] += np.mean(distance_diff < threshold)\n    # print(fractions)\n    # print(num_atoms)\n\n    # Average the fractions over the number of atoms\n    fractions /= num_atoms\n    \n    # The final lDDT score is the average of these fractions\n    lddt_score = np.mean(fractions)\n    \n    return lddt_score","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T18:52:58.034159Z","iopub.execute_input":"2025-05-16T18:52:58.034512Z","iopub.status.idle":"2025-05-16T18:52:58.045208Z","shell.execute_reply.started":"2025-05-16T18:52:58.034484Z","shell.execute_reply":"2025-05-16T18:52:58.044419Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for epoch in range(epochs):\n    model.train()\n    tbar=tqdm(train_loader)\n    total_loss=0\n    total_distogram_loss=0\n    oom=0\n    for idx, batch in enumerate(tbar):\n        #try:\n        sequence=batch['sequence'].cuda()\n        gt_xyz=batch['xyz'].squeeze().cuda()\n        mask=~torch.isnan(gt_xyz)\n        gt_xyz[torch.isnan(gt_xyz)]=0\n\n        distance_matrix=calculate_distance_matrix(gt_xyz,gt_xyz)\n        distogram_mask=distance_matrix==distance_matrix\n        distance_matrix=distance_matrix.clip(2,39).long()\n\n\n\n        gt_xyz=gt_xyz.unsqueeze(0).repeat(48,1,1)\n        time_steps=torch.randint(0,config['n_times'],size=(gt_xyz.shape[0],)).to(gt_xyz.device)\n        loss_weight=(1-time_steps/config['n_times'])\n        noised_xyz, noise=model.make_noisy(gt_xyz, time_steps)\n\n        #exit()\n        pred_noise,distogram_pred=model(sequence,noised_xyz,time_steps)#.squeeze()\n        #pred_xyz=aug_xyz[:,1:-1]+pred_displacements[:,1:-1]\n        #exit()\n\n        \n        loss= torch.square(noise-pred_noise)*loss_weight[:,None,None]\n        loss=loss[mask.repeat(48,1,1)].mean()\n        #exit()\n        \n        distogram_loss=criterion(distogram_pred.squeeze()[distogram_mask],distance_matrix[distogram_mask]).mean()\n        total_distogram_loss+=distogram_loss.item()\n\n\n        if loss!=loss:\n            stop\n\n        \n        #(loss/batch_size*len(gt_xyz)).backward()\n\n        #accelerator.backward()\n\n        ((loss+0.2*distogram_loss)/batch_size*len(gt_xyz)).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            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)} Distogram Loss: {total_distogram_loss/(idx+1)}\")\n        #break\n    # visualize_point_cloud_batch(pred_xyz)\n    # visualize_point_cloud_batch(aug_xyz)\n\n\n    \n    total_loss=total_loss/len(tbar)\n    \n    tbar=tqdm(val_loader)\n    model.eval()\n    val_preds=[]\n    val_loss=0\n    val_rmsd=0\n    val_lddt=0\n    #unwrapped_diffusion=accelerator.unwrap_model(diffusion)\n    #unwrapped_model=accelerator.unwrap_model(model)\n    for idx, batch in enumerate(tbar):\n        sequence=batch['sequence'].cuda()\n        gt_xyz=batch['xyz'].cuda().squeeze()\n    \n        with torch.no_grad():\n            # if accelerator.dis\n            #pred_xyz=model.module.decode(sequence,torch.ones_like(sequence).long().cuda()).squeeze()\n            pred_xyz=model.sample(sequence,1)[0].squeeze(0)\n\n            #pred_xyz=model(sequence)[-1].squeeze()\n            loss=dRMAE(pred_xyz,pred_xyz,gt_xyz,gt_xyz)\n    \n        val_rmsd+=align_svd_rmsd(pred_xyz,gt_xyz)\n        val_lddt+=compute_lddt(pred_xyz.cpu().numpy(),gt_xyz.cpu().numpy())\n        val_loss+=loss\n    \n        val_preds.append([gt_xyz.cpu().numpy(),pred_xyz.cpu().numpy()])\n    val_loss=val_loss/len(tbar)\n    val_rmsd=val_rmsd/len(tbar)\n    val_lddt=val_lddt/len(tbar)\n    \n    print(f\"val loss: {val_loss}\")\n    print(f\"val_rmsd: {val_rmsd}\")\n    print(f\"val_lddt: {val_lddt}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T18:53:29.413243Z","iopub.execute_input":"2025-05-16T18:53:29.413649Z"}},"outputs":[],"execution_count":null}]}