{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":12276181,"sourceType":"competition"},{"sourceId":10855324,"sourceType":"datasetVersion","datasetId":6742586},{"sourceId":11278691,"sourceType":"datasetVersion","datasetId":7051341},{"sourceId":11469248,"sourceType":"datasetVersion","datasetId":7187409},{"sourceId":11839907,"sourceType":"datasetVersion","datasetId":7438872},{"sourceId":13036374,"sourceType":"datasetVersion","datasetId":8254593},{"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":"# --- Imports and Setup ---\nimport pandas as pd\nimport torch\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport random\nimport pickle\nimport os\nimport sys\nfrom tqdm import tqdm\nimport pandas as pd\nimport numpy as np\nimport pickle\nfrom tqdm import tqdm\nimport os\nimport torch.nn as nn","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.status.busy":"2025-09-12T07:58:20.996770Z","iopub.execute_input":"2025-09-12T07:58:20.997050Z","iopub.status.idle":"2025-09-12T07:58:21.001465Z","shell.execute_reply.started":"2025-09-12T07:58:20.997027Z","shell.execute_reply":"2025-09-12T07:58:21.000450Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"A config dictionary is defined with hyperparameters and file paths for training, such as seed, batch size, learning rate, and model config paths.","metadata":{}},{"cell_type":"code","source":"# Configuration\nconfig = {\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":{"execution":{"iopub.status.busy":"2025-09-12T07:58:21.002710Z","iopub.execute_input":"2025-09-12T07:58:21.002935Z","iopub.status.idle":"2025-09-12T07:58:21.019926Z","shell.execute_reply.started":"2025-09-12T07:58:21.002914Z","shell.execute_reply":"2025-09-12T07:58:21.019208Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data Loading and Preprocessing\nThis section loads the RNA sequence and 3D structure data, either from CSVs or a preprocessed pickle file. It also includes filtering and cleaning steps to prepare the dataset for training.","metadata":{}},{"cell_type":"code","source":"SEQ_CSV = \"/kaggle/input/stanford-rna-3d-folding/train_sequences.v2.csv\"  # Update if needed\nLABEL_CSV = \"/kaggle/input/stanford-rna-3d-folding/train_labels.v2.csv\"  # Update if needed\nPICKLE_OUT = \"/kaggle/working/train_data.pkl\"      # Output pickle file","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T07:58:21.021812Z","iopub.execute_input":"2025-09-12T07:58:21.022097Z","iopub.status.idle":"2025-09-12T07:58:21.033192Z","shell.execute_reply.started":"2025-09-12T07:58:21.022069Z","shell.execute_reply":"2025-09-12T07:58:21.032548Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ==== Filtering thresholds (edit as needed) ====\n# MAX_NAN_FRAC = 0.5\n# MAX_LEN = 9999999\n# MIN_LEN = 10\n\n# # ==== Load CSVs ====\n# train_sequences = pd.read_csv(SEQ_CSV)\n# train_labels = pd.read_csv(LABEL_CSV)\n","metadata":{"execution":{"iopub.status.busy":"2025-09-12T07:58:21.034310Z","iopub.execute_input":"2025-09-12T07:58:21.034543Z","iopub.status.idle":"2025-09-12T07:58:21.049779Z","shell.execute_reply.started":"2025-09-12T07:58:21.034524Z","shell.execute_reply":"2025-09-12T07:58:21.049102Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # Add pdb_id column for grouping\n# train_labels[\"pdb_id\"] = train_labels[\"ID\"].apply(lambda x: x.split(\"_\")[0] + \"_\" + x.split(\"_\")[1])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T07:58:21.050433Z","iopub.execute_input":"2025-09-12T07:58:21.050618Z","iopub.status.idle":"2025-09-12T07:58:21.064334Z","shell.execute_reply.started":"2025-09-12T07:58:21.050601Z","shell.execute_reply":"2025-09-12T07:58:21.063676Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\n# # Collect xyz coordinates for each sequence\n# all_xyz = []\n# for 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#     xyz[xyz < -1e17] = float('nan')\n#     all_xyz.append(xyz)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T07:58:21.065032Z","iopub.execute_input":"2025-09-12T07:58:21.065269Z","iopub.status.idle":"2025-09-12T07:58:21.076960Z","shell.execute_reply.started":"2025-09-12T07:58:21.065250Z","shell.execute_reply":"2025-09-12T07:58:21.076323Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# # Filter out sequences with too many NaNs or invalid length\n# filter_nan = []\n# for xyz in all_xyz:\n#     filter_nan.append((np.isnan(xyz).mean() <= MAX_NAN_FRAC) and\n#                       (len(xyz) < MAX_LEN) and\n#                       (len(xyz) > MIN_LEN))\n# filter_nan = np.array(filter_nan)\n# non_nan_indices = np.arange(len(filter_nan))[filter_nan]\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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T07:58:21.077625Z","iopub.execute_input":"2025-09-12T07:58:21.077805Z","iopub.status.idle":"2025-09-12T07:58:21.090500Z","shell.execute_reply.started":"2025-09-12T07:58:21.077789Z","shell.execute_reply":"2025-09-12T07:58:21.089856Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\n# # Pack data into a dictionary\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# # Save to pickle\n# with open(PICKLE_OUT, \"wb\") as f:\n#     pickle.dump(data, f)\n\n# print(f\"Saved {len(data['sequence'])} sequences to {PICKLE_OUT}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T07:58:21.091253Z","iopub.execute_input":"2025-09-12T07:58:21.091495Z","iopub.status.idle":"2025-09-12T07:58:21.104574Z","shell.execute_reply.started":"2025-09-12T07:58:21.091476Z","shell.execute_reply":"2025-09-12T07:58:21.103950Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load preprocessed data from pickle file\nwith open(\"/kaggle/input/train-data-pkl/train_data.pkl\", \"rb\") as f:\n    data = pickle.load(f)","metadata":{"execution":{"iopub.status.busy":"2025-09-12T07:58:21.106092Z","iopub.execute_input":"2025-09-12T07:58:21.106355Z","iopub.status.idle":"2025-09-12T07:58:21.223901Z","shell.execute_reply.started":"2025-09-12T07:58:21.106324Z","shell.execute_reply":"2025-09-12T07:58:21.223212Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Split data into train and test sets based on temporal cutoff dates\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\nprint(f\"Train size: {len(train_index)}\")\nprint(f\"Test size: {len(test_index)}\")","metadata":{"execution":{"iopub.status.busy":"2025-09-12T07:58:21.225092Z","iopub.execute_input":"2025-09-12T07:58:21.225328Z","iopub.status.idle":"2025-09-12T07:58:21.248296Z","shell.execute_reply.started":"2025-09-12T07:58:21.225308Z","shell.execute_reply":"2025-09-12T07:58:21.247593Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# PyTorch Dataset and DataLoader\nThis section defines a custom PyTorch Dataset for RNA 3D data and sets up DataLoaders for training and validation.","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader\nfrom collections import defaultdict\n\n# Custom Dataset for RNA 3D structure data\nclass RNA3D_Dataset(Dataset):\n    def __init__(self, indices, data):\n        self.indices = indices\n        self.data = data\n        # Map nucleotides to integers, default to 4 for unknowns\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    def __len__(self):\n        return len(self.indices)\n\n    def __getitem__(self, idx):\n        idx = self.indices[idx]\n        # Convert sequence to integer tokens\n        sequence = [self.tokens[nt] for nt in (self.data['sequence'][idx])]\n        sequence = torch.tensor(np.array(sequence))\n\n        # Get C1' atom xyz coordinates\n        xyz = torch.tensor(np.array(self.data['xyz'][idx]))\n\n        # Crop if sequence is too long\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            sequence = sequence[crop_start:crop_end]\n            xyz = xyz[crop_start:crop_end]\n\n        # Center at first valid atom\n        for i in range(len(xyz)):\n            if (~torch.isnan(xyz[i])).all():\n                break\n        xyz = xyz - xyz[i]\n\n        return {'sequence': sequence, 'xyz': xyz}","metadata":{"execution":{"iopub.status.busy":"2025-09-12T07:58:21.249552Z","iopub.execute_input":"2025-09-12T07:58:21.249754Z","iopub.status.idle":"2025-09-12T07:58:21.261342Z","shell.execute_reply.started":"2025-09-12T07:58:21.249737Z","shell.execute_reply":"2025-09-12T07:58:21.260538Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create train and validation DataLoaders\ntrain_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)","metadata":{"execution":{"iopub.status.busy":"2025-09-12T07:58:21.262273Z","iopub.execute_input":"2025-09-12T07:58:21.262558Z","iopub.status.idle":"2025-09-12T07:58:21.278812Z","shell.execute_reply.started":"2025-09-12T07:58:21.262537Z","shell.execute_reply":"2025-09-12T07:58:21.278153Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model Architecture: RibonanzaNet with Diffusion\nThis section defines the neural network architecture for RNA 3D structure prediction, based on a finetuned RibonanzaNet and a diffusion process (DDPM).","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-09-12T07:58:21.279521Z","iopub.execute_input":"2025-09-12T07:58:21.279699Z","iopub.status.idle":"2025-09-12T07:58:23.213281Z","shell.execute_reply.started":"2025-09-12T07:58:21.279683Z","shell.execute_reply":"2025-09-12T07:58:23.212585Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn.functional as F\n\nclass SimpleStructureModuleWithGeoBias(nn.Module):\n    def __init__(self, d_model, nhead, dim_feedforward, pairwise_dimension, dropout=0.1):\n        super(SimpleStructureModuleWithGeoBias, self).__init__()\n        self.self_attn = MultiHeadAttention(d_model, nhead, d_model//nhead, d_model//nhead, dropout=dropout)\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        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        self.pairwise2heads = nn.Linear(pairwise_dimension, nhead, bias=False)\n        self.pairwise_norm = nn.LayerNorm(pairwise_dimension)\n        self.activation = nn.GELU()\n        # New: learnable scalar for distance bias\n        self.distance_bias_weight = nn.Parameter(torch.tensor(1.0))\n\n    def forward(self, input):\n        tgt, src, pairwise_features, xyz, src_mask = input\n        # Compute pairwise bias as before\n        pairwise_bias = self.pairwise2heads(self.pairwise_norm(pairwise_features)).permute(0,3,1,2)\n        # Compute distance matrix (B, L, L)\n        if xyz is not None:\n            dists = torch.cdist(xyz, xyz, p=2)  # (B, L, L)\n            # Normalize and expand for nhead\n            dists = dists / (dists.max() + 1e-6)\n            dists = dists.unsqueeze(1).expand(-1, pairwise_bias.shape[1], -1, -1)  # (B, nhead, L, L)\n            # Add distance-based bias\n            pairwise_bias = pairwise_bias - self.distance_bias_weight * dists\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        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        return tgt\n","metadata":{"execution":{"iopub.status.busy":"2025-09-12T07:58:23.214059Z","iopub.execute_input":"2025-09-12T07:58:23.214532Z","iopub.status.idle":"2025-09-12T07:58:23.223174Z","shell.execute_reply.started":"2025-09-12T07:58:23.214500Z","shell.execute_reply":"2025-09-12T07:58:23.222040Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Config Loading Utility ---\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)","metadata":{"execution":{"iopub.status.busy":"2025-09-12T07:58:23.224119Z","iopub.execute_input":"2025-09-12T07:58:23.224475Z","iopub.status.idle":"2025-09-12T07:58:23.259152Z","shell.execute_reply.started":"2025-09-12T07:58:23.224443Z","shell.execute_reply":"2025-09-12T07:58:23.258367Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Utility: Apply random 3D rotation to a batch of point clouds\ndef random_rotation_point_cloud_torch_batch(point_clouds):\n    \"\"\"\n    Apply a random 3D rotation to a batch of point clouds (PyTorch version).\n    Args:\n        point_clouds (torch.Tensor): BxNx3 tensor of XYZ points.\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","metadata":{"execution":{"iopub.status.busy":"2025-09-12T07:58:23.261271Z","iopub.execute_input":"2025-09-12T07:58:23.261469Z","iopub.status.idle":"2025-09-12T07:58:23.265727Z","shell.execute_reply.started":"2025-09-12T07:58:23.261451Z","shell.execute_reply":"2025-09-12T07:58:23.265006Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Instantiate configs and model\ndiffusion_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":{"execution":{"iopub.status.busy":"2025-09-12T07:58:23.266672Z","iopub.execute_input":"2025-09-12T07:58:23.266861Z","iopub.status.idle":"2025-09-12T07:58:25.805699Z","shell.execute_reply.started":"2025-09-12T07:58:23.266844Z","shell.execute_reply":"2025-09-12T07:58:25.804968Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training Loop and Evaluation\nThis section contains the main training loop, optimizer setup, and evaluation metrics for model performance.","metadata":{}},{"cell_type":"code","source":"# Set up optimizer, loss function, and learning rate scheduler\nepochs = config['epochs']\ncos_epoch = config['cos_epoch']\n\nbest_loss = np.inf\noptimizer = torch.optim.Adam(model.parameters(), weight_decay=0.0, lr=0.0002) # no weight decay following AF\nbatch_size = 1\n\ncriterion = torch.nn.CrossEntropyLoss(reduction='none')\n\nschedule = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=(epochs-cos_epoch)*len(train_loader)//batch_size)","metadata":{"execution":{"iopub.status.busy":"2025-09-12T07:58:25.806519Z","iopub.execute_input":"2025-09-12T07:58:25.806825Z","iopub.status.idle":"2025-09-12T07:58:25.819692Z","shell.execute_reply.started":"2025-09-12T07:58:25.806794Z","shell.execute_reply":"2025-09-12T07:58:25.818837Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Evaluation Metrics: Distance, RMSD, lDDT ---\ndef 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, pred_y, gt_x, gt_y, 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    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 and computes RMSD loss.\n    \"\"\"\n    assert input.shape == target.shape, \"Input and target must have the same shape\"\n\n    mask = ~torch.isnan(target.sum(-1))\n    input = input[mask]\n    target = target[mask]\n    \n    centroid_input = input.mean(dim=0, keepdim=True)\n    centroid_target = target.mean(dim=0, keepdim=True)\n\n    input_centered = input - centroid_input.detach()\n    target_centered = target - centroid_target\n\n    cov_matrix = input_centered.T @ target_centered\n\n    U, S, Vt = torch.svd(cov_matrix)\n    R = Vt @ U.T\n    if torch.det(R) < 0:\n        Vt[-1, :] *= -1\n        R = Vt @ U.T\n\n    aligned_input = (input_centered @ R.T.detach()) + centroid_target.detach()\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    num_atoms = ground_truth_atoms.shape[0]\n    fractions = np.zeros(len(thresholds))\n    for i in range(num_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        mask = (gt_distances > 0) & (gt_distances < cutoff)\n        distance_diff = np.abs(gt_distances[mask] - pred_distances[mask])\n        valid_mask = ~np.isnan(distance_diff)\n        distance_diff = distance_diff[valid_mask]\n        for j, threshold in enumerate(thresholds):\n            if len(distance_diff) > 0:\n                fractions[j] += np.mean(distance_diff < threshold)\n    fractions /= num_atoms\n    lddt_score = np.mean(fractions)\n    return lddt_score","metadata":{"execution":{"iopub.status.busy":"2025-09-12T07:58:25.820478Z","iopub.execute_input":"2025-09-12T07:58:25.820746Z","iopub.status.idle":"2025-09-12T07:58:25.834968Z","shell.execute_reply.started":"2025-09-12T07:58:25.820725Z","shell.execute_reply":"2025-09-12T07:58:25.834198Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Main Training and Validation Loop ---\nfor 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        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        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        pred_noise, distogram_pred = model(sequence, noised_xyz, time_steps)\n\n        loss = torch.square(noise - pred_noise) * loss_weight[:, None, None]\n        loss = loss[mask.repeat(48, 1, 1)].mean()\n\n        distogram_loss = criterion(distogram_pred.squeeze()[distogram_mask], distance_matrix[distogram_mask]).mean()\n        total_distogram_loss += distogram_loss.item()\n\n        if loss != loss:\n            stop\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            torch.nn.utils.clip_grad_norm_(model.parameters(), 1)\n            optimizer.step()\n            optimizer.zero_grad()\n            if (epoch + 1) > cos_epoch:\n                schedule.step()\n        total_loss += loss.item()\n        tbar.set_description(f\"Epoch {epoch + 1} Loss: {total_loss/(idx+1)} Distogram Loss: {total_distogram_loss/(idx+1)}\")\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    for idx, batch in enumerate(tbar):\n        sequence = batch['sequence'].cuda()\n        gt_xyz = batch['xyz'].cuda().squeeze()\n        with torch.no_grad():\n            pred_xyz = model.sample(sequence, 1)[0].squeeze(0)\n            loss = dRMAE(pred_xyz, pred_xyz, gt_xyz, gt_xyz)\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        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    print(f\"val loss: {val_loss}\")\n    print(f\"val_rmsd: {val_rmsd}\")\n    print(f\"val_lddt: {val_lddt}\")","metadata":{"execution":{"iopub.status.busy":"2025-09-12T07:58:25.835775Z","iopub.execute_input":"2025-09-12T07:58:25.836053Z","iopub.status.idle":"2025-09-12T07:58:25.850357Z","shell.execute_reply.started":"2025-09-12T07:58:25.836026Z","shell.execute_reply":"2025-09-12T07:58:25.849591Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.save(model.state_dict(), \"/kaggle/working/mi_ribonanzanet_trained.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-12T07:58:25.851281Z","iopub.execute_input":"2025-09-12T07:58:25.851554Z","iopub.status.idle":"2025-09-12T07:58:26.773060Z","shell.execute_reply.started":"2025-09-12T07:58:25.851526Z","shell.execute_reply":"2025-09-12T07:58:26.772411Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}}]}