{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.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":118765,"databundleVersionId":15231210,"sourceType":"competition"},{"sourceId":7639698,"sourceType":"datasetVersion","datasetId":4299272},{"sourceId":8318191,"sourceType":"datasetVersion","datasetId":4459124}],"dockerImageVersionId":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"<h1 id=\"setup-and-configuration\"\n    style=\"\n  background: linear-gradient(135deg, #0f2027, #203a43, #2c5364);\n  color: white;\n  padding: 12px 24px;\n  text-align: center;\n  border-radius: 12px;\n  font-size: 22px;\n  font-weight: bold;\n  font-family: Arial, sans-serif;\n  margin: 20px 0;\n\">\n  1. Imports and Configurations\n</h1>","metadata":{}},{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport os\nimport sys\nimport random\nimport pickle\nimport yaml\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport matplotlib.pyplot as plt\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm import tqdm\n\n\nclass Config:\n    sample_sub = \"/kaggle/input/stanford-rna-3d-folding-2/sample_submission.csv\"\n    test_seq = \"/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv\"\n    train_labels = \"/kaggle/input/stanford-rna-3d-folding-2/train_labels.csv\"\n    train_sequences = \"/kaggle/input/stanford-rna-3d-folding-2/train_sequences.csv\"\n    validation_labels = \"/kaggle/input/stanford-rna-3d-folding-2/validation_labels.csv\"\n    validation_sequences = \"/kaggle/input/stanford-rna-3d-folding-2/validation_sequences.csv\"\n    model_config_path = \"/kaggle/input/ribonanzanet2d-final/configs/pairwise.yaml\"\n    pretrained_weights_path = \"/kaggle/input/ribonanzanet-weights/RibonanzaNet.pt\"\n    \n    max_len = 384\n    batch_size = 1\n    max_len_filter = 9999999\n    min_len_filter = 10\n    seed = 42\n\nconfig = Config()\n\n# Set seed for reproducibility\ndef set_seed(seed: int):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nset_seed(config.seed)\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-09T13:16:03.50727Z","iopub.execute_input":"2026-01-09T13:16:03.507559Z","iopub.status.idle":"2026-01-09T13:16:03.516938Z","shell.execute_reply.started":"2026-01-09T13:16:03.507533Z","shell.execute_reply":"2026-01-09T13:16:03.516383Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<h1 id=\"data-load\"\n    style=\"\n  background: linear-gradient(135deg, #0f2027, #203a43, #2c5364);\n  color: white;\n  padding: 12px 24px;\n  text-align: center;\n  border-radius: 12px;\n  font-size: 22px;\n  font-weight: bold;\n  font-family: Arial, sans-serif;\n  margin: 20px 0;\n\">\n  2. Data Loading\n</h1>","metadata":{}},{"cell_type":"code","source":"train_sequences = pd.read_csv(config.train_sequences)\ntrain_labels = pd.read_csv(config.train_labels)\nval_sequences = pd.read_csv(config.validation_sequences)\nval_labels = pd.read_csv(config.validation_labels)\n\ntrain_labels[\"pdb_id\"] = train_labels.ID.str.rsplit('_', n=1, expand=True).iloc[:,0]\nval_labels[\"pdb_id\"] = val_labels.ID.str.rsplit('_', n=1, expand=True).iloc[:,0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-09T13:16:03.531068Z","iopub.execute_input":"2026-01-09T13:16:03.531535Z","iopub.status.idle":"2026-01-09T13:16:19.682395Z","shell.execute_reply.started":"2026-01-09T13:16:03.531515Z","shell.execute_reply":"2026-01-09T13:16:19.681863Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<h1 id=\"data-prep\"\n    style=\"\n  background: linear-gradient(135deg, #0f2027, #203a43, #2c5364);\n  color: white;\n  padding: 12px 24px;\n  text-align: center;\n  border-radius: 12px;\n  font-size: 22px;\n  font-weight: bold;\n  font-family: Arial, sans-serif;\n  margin: 20px 0;\n\">\n  3. Prepairing Data\n</h1>","metadata":{}},{"cell_type":"code","source":"def making_data_dict(sequences_df: pd.DataFrame, labels_df: pd.DataFrame):\n    grouped = labels_df.groupby(\"pdb_id\")\n    data = {}\n    sequences, pdb_ids, all_xyz = [], [], []\n\n    for pdb_id in tqdm(sequences_df[\"target_id\"], desc=\"prepairing data\"):\n        if pdb_id not in grouped.groups:\n            continue\n\n        xyz = grouped.get_group(pdb_id)[[\"x_1\", \"y_1\", \"z_1\"]].to_numpy(dtype=\"float32\")\n\n        sequence = list((sequences_df[sequences_df[\"target_id\"]== pdb_id])[\"sequence\"])[0]\n\n        sequences.append(sequence)\n        pdb_ids.append(pdb_id)\n        all_xyz.append(xyz)\n        \n    data[\"pdb_ids\"] = pdb_ids\n    data[\"sequences\"] = sequences\n    data[\"all_xyz\"] = all_xyz\n    \n    return data\n\n\n#  Create Data Dictionaries\n\ntrain_data_dict = making_data_dict(train_sequences, train_labels)\nval_data_dict = making_data_dict(val_sequences, val_labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-09T13:16:19.683595Z","iopub.execute_input":"2026-01-09T13:16:19.683846Z","iopub.status.idle":"2026-01-09T13:16:27.971258Z","shell.execute_reply.started":"2026-01-09T13:16:19.683825Z","shell.execute_reply":"2026-01-09T13:16:27.970684Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def filter_data_by_nan_and_length(data_dict, max_len_filter=9999999, min_len_filter=10):\n    \"\"\"Filter sequences based on NaN ratio and length constraints.\"\"\"\n    valid_indices = []\n    max_len_seen = 0\n    \n    for i, xyz in enumerate(data_dict[\"all_xyz\"]):\n        # Track the maximum length\n        if len(xyz) > max_len_seen:\n            max_len_seen = len(xyz)\n        \n        nan_ratio = np.isnan(xyz).mean()\n        seq_len = len(xyz)\n        \n        # Keep sequence if it meets criteria\n        if (nan_ratio <= 0.1) and (min_len_filter < seq_len < max_len_filter):\n            valid_indices.append(i)\n    \n    \n    # Filter all fields based on valid_indices\n    filtered_data = {\n        \"pdb_ids\": [data_dict[\"pdb_ids\"][i] for i in valid_indices],\n        \"sequences\": [data_dict[\"sequences\"][i] for i in valid_indices],\n        \"all_xyz\": [data_dict[\"all_xyz\"][i] for i in valid_indices]\n    }\n    \n    return filtered_data\n\n\n# Filter \ntrain_data_dict = filter_data_by_nan_and_length(train_data_dict, \n                                                max_len_filter=config.max_len_filter,min_len_filter=config.min_len_filter)\nval_data_dict = filter_data_by_nan_and_length(val_data_dict,\n                                              max_len_filter=config.max_len_filter,min_len_filter=config.min_len_filter)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-09T13:16:27.972453Z","iopub.execute_input":"2026-01-09T13:16:27.972684Z","iopub.status.idle":"2026-01-09T13:16:28.044323Z","shell.execute_reply.started":"2026-01-09T13:16:27.972665Z","shell.execute_reply":"2026-01-09T13:16:28.043829Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<h1 id=\"loader-set\"\n    style=\"\n  background: linear-gradient(135deg, #0f2027, #203a43, #2c5364);\n  color: white;\n  padding: 12px 24px;\n  text-align: center;\n  border-radius: 12px;\n  font-size: 22px;\n  font-weight: bold;\n  font-family: Arial, sans-serif;\n  margin: 20px 0;\n\">\n  4. Datasets and DataLoaders\n</h1>","metadata":{}},{"cell_type":"code","source":"class RNA3D_Dataset(Dataset):\n    \"\"\"\n    A PyTorch Dataset for 3D RNA structures.\n    \"\"\"\n    def __init__(self, data_dict, max_len=384):\n        self.data = data_dict\n        self.max_len = max_len\n        self.nt_to_idx = {nt: i for i, nt in enumerate(\"ACGU\")}\n\n    def __len__(self):\n        return len(self.data[\"sequences\"])\n    \n    def __getitem__(self, idx):\n        sequence = [self.nt_to_idx[nt] for nt in self.data[\"sequences\"][idx]]\n        sequence = torch.tensor(sequence, dtype=torch.long)\n        xyz = torch.tensor(self.data[\"all_xyz\"][idx], dtype=torch.float32)\n        \n        # If sequence is longer than max_len, randomly crop\n        if len(sequence) > self.max_len:\n            crop_start = np.random.randint(len(sequence) - self.max_len)\n            crop_end = crop_start + self.max_len\n            sequence = sequence[crop_start:crop_end]\n            xyz = xyz[crop_start:crop_end]\n\n        return {\"sequence\": sequence, \"xyz\": xyz}\n\n\n\n# Create Dataset and DataLoaders\n\ntrain_dataset = RNA3D_Dataset(train_data_dict, max_len=config.max_len)\nval_dataset = RNA3D_Dataset(val_data_dict, max_len=config.max_len)\n\ntrain_loader = DataLoader(train_dataset, batch_size=config.batch_size, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=config.batch_size, shuffle=False)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-09T13:16:28.045782Z","iopub.execute_input":"2026-01-09T13:16:28.045989Z","iopub.status.idle":"2026-01-09T13:16:28.053076Z","shell.execute_reply.started":"2026-01-09T13:16:28.045971Z","shell.execute_reply":"2026-01-09T13:16:28.052389Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<h1 id=\"model-diff\"\n    style=\"\n  background: linear-gradient(135deg, #0f2027, #203a43, #2c5364);\n  color: white;\n  padding: 12px 24px;\n  text-align: center;\n  border-radius: 12px;\n  font-size: 22px;\n  font-weight: bold;\n  font-family: Arial, sans-serif;\n  margin: 20px 0;\n\">\n  5.  Model Definition\n</h1>","metadata":{}},{"cell_type":"code","source":"#  Model Configuration Helper Classes\nsys.path.append(\"/kaggle/input/ribonanzanet2d-final\")\nfrom Network import RibonanzaNet\n\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\n\ndef load_config_from_yaml(file_path):\n    with open(file_path, 'r') as file:\n        cfg = yaml.safe_load(file)\n    return Config(**cfg)\n\n\n\n#  Model Definition\nclass FinetunedRibonanzaNet(RibonanzaNet):\n    def __init__(self, config_obj, pretrained=False, dropout=0.1):\n        config_obj.dropout = dropout\n        super(FinetunedRibonanzaNet, self).__init__(config_obj)\n\n        if pretrained:\n            self.load_state_dict(\n                torch.load(config.pretrained_weights_path, map_location=\"cpu\")\n            )\n\n        self.dropout = nn.Dropout(p=0.0)\n        self.xyz_predictor = nn.Linear(256, 3)\n\n    def forward(self, src):\n        sequence_features, _ = self.get_embeddings(\n            src, torch.ones_like(src).long().to(src.device)\n        )\n        xyz_pred = self.xyz_predictor(sequence_features)\n        return xyz_pred\n\n\n\n# Initialize Model\nmodel_cfg = load_config_from_yaml(config.model_config_path)\nmodel = FinetunedRibonanzaNet(model_cfg, pretrained=True).cuda()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-09T13:16:28.05381Z","iopub.execute_input":"2026-01-09T13:16:28.054123Z","iopub.status.idle":"2026-01-09T13:16:28.444177Z","shell.execute_reply.started":"2026-01-09T13:16:28.054105Z","shell.execute_reply":"2026-01-09T13:16:28.443614Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<h1 id=\"loss-fn\"\n    style=\"\n  background: linear-gradient(135deg, #0f2027, #203a43, #2c5364);\n  color: white;\n  padding: 12px 24px;\n  text-align: center;\n  border-radius: 12px;\n  font-size: 22px;\n  font-weight: bold;\n  font-family: Arial, sans-serif;\n  margin: 20px 0;\n\">\n  6. Loss Function\n</h1>","metadata":{}},{"cell_type":"code","source":"def calculate_distance_matrix(X, Y, epsilon=1e-4):\n    return ((X[:, None] - Y[None, :])**2 + epsilon).sum(dim=-1).sqrt()\n\n\ndef dRMSD(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], device=mask.device).bool()] = False\n\n    diff_sq = (pred_dm[mask] - gt_dm[mask])**2 + epsilon\n    if d_clamp is not None:\n        diff_sq = diff_sq.clamp(max=d_clamp**2)\n\n    return diff_sq.sqrt().mean() / Z\n\n\ndef local_dRMSD(pred_x, pred_y, gt_x, gt_y, epsilon=1e-4, Z=10, d_clamp=30):\n    pred_dm = calculate_distance_matrix(pred_x, pred_y)\n    gt_dm = calculate_distance_matrix(gt_x, gt_y)\n\n    mask = (~torch.isnan(gt_dm)) & (gt_dm < d_clamp)\n    mask[torch.eye(mask.shape[0], device=mask.device).bool()] = False\n\n    diff_sq = (pred_dm[mask] - gt_dm[mask])**2 + epsilon\n    return diff_sq.sqrt().mean() / Z\n\n\ndef dRMAE(pred_x, pred_y, gt_x, gt_y, epsilon=1e-4, Z=10):\n    pred_dm = calculate_distance_matrix(pred_x, pred_y)\n    gt_dm = calculate_distance_matrix(gt_x, gt_y)\n\n    mask = ~torch.isnan(gt_dm)\n    mask[torch.eye(mask.shape[0], device=mask.device).bool()] = False\n\n    diff = torch.abs(pred_dm[mask] - gt_dm[mask])\n    return diff.mean() / Z\n\n\ndef align_svd_mae(input_coords, target_coords, Z=10):\n    mask = ~torch.isnan(target_coords.sum(dim=-1))\n    input_coords = input_coords[mask]\n    target_coords = target_coords[mask]\n\n    centroid_input = input_coords.mean(dim=0, keepdim=True)\n    centroid_target = target_coords.mean(dim=0, keepdim=True)\n\n    input_centered = input_coords - centroid_input\n    target_centered = target_coords - centroid_target\n\n    cov_matrix = input_centered.T @ target_centered\n    U, S, Vt = torch.svd(cov_matrix)\n    R = Vt @ U.T\n\n    if torch.det(R) < 0:\n        Vt_adj = Vt.clone()\n        Vt_adj[-1, :] = -Vt_adj[-1, :]\n        R = Vt_adj @ U.T\n\n    aligned_input = (input_centered @ R.T) + centroid_target\n    return torch.abs(aligned_input - target_coords).mean() / Z\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-09T13:16:28.444954Z","iopub.execute_input":"2026-01-09T13:16:28.445215Z","iopub.status.idle":"2026-01-09T13:16:28.454904Z","shell.execute_reply.started":"2026-01-09T13:16:28.445191Z","shell.execute_reply":"2026-01-09T13:16:28.454246Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<h1 id=\"train-fn\"\n    style=\"\n  background: linear-gradient(135deg, #0f2027, #203a43, #2c5364);\n  color: white;\n  padding: 12px 24px;\n  text-align: center;\n  border-radius: 12px;\n  font-size: 22px;\n  font-weight: bold;\n  font-family: Arial, sans-serif;\n  margin: 20px 0;\n\">\n  7. Training Function\n</h1>","metadata":{}},{"cell_type":"code","source":"def train_model(model, train_dl, val_dl, epochs=50, cos_epoch=35, lr=3e-4, clip=1):\n    optimizer = torch.optim.AdamW(model.parameters(), weight_decay=0.0, lr=lr)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n        optimizer,\n        T_max=(epochs - cos_epoch) * len(train_dl),\n    )\n\n    best_val_loss = float(\"inf\")\n    best_preds = None\n\n    for epoch in range(epochs):\n        model.train()\n        train_pbar = tqdm(train_dl, desc=f\"Training Epoch {epoch+1}/{epochs}\")\n        running_loss = 0.0\n\n        for idx, batch in enumerate(train_pbar):\n            sequence = batch[\"sequence\"].cuda()\n            gt_xyz = batch[\"xyz\"].squeeze().cuda()\n\n            pred_xyz = model(sequence).squeeze()\n\n            loss = dRMAE(pred_xyz, pred_xyz, gt_xyz, gt_xyz) + align_svd_mae(pred_xyz, gt_xyz)\n            loss.backward()\n\n            torch.nn.utils.clip_grad_norm_(model.parameters(), clip)\n            optimizer.step()\n            optimizer.zero_grad()\n\n            if (epoch + 1) > cos_epoch:\n                scheduler.step()\n\n            running_loss += loss.item()\n            avg_loss = running_loss / (idx + 1)\n            train_pbar.set_description(f\"Epoch {epoch+1} | Loss: {avg_loss:.4f}\")\n\n        model.eval()\n        val_loss = 0.0\n        val_preds = []\n\n        with torch.no_grad():\n            for batch in val_dl:\n                sequence = batch[\"sequence\"].cuda()\n                gt_xyz = batch[\"xyz\"].squeeze().cuda()\n\n                pred_xyz = model(sequence).squeeze()\n                loss = dRMAE(pred_xyz, pred_xyz, gt_xyz, gt_xyz)\n                val_loss += loss.item()\n\n                val_preds.append((gt_xyz.cpu().numpy(), pred_xyz.cpu().numpy()))\n\n        val_loss /= len(val_dl)\n        print(f\"Validation Loss (Epoch {epoch+1}): {val_loss:.4f}\")\n\n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            best_preds = val_preds\n            torch.save(model.state_dict(), config.save_weights_name)\n            print(f\"  -> New best model saved at epoch {epoch+1}\")\n\n    torch.save(model.state_dict(), config.save_weights_final)\n    return best_val_loss, best_preds","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-09T13:16:28.455741Z","iopub.execute_input":"2026-01-09T13:16:28.455988Z","iopub.status.idle":"2026-01-09T13:16:28.478386Z","shell.execute_reply.started":"2026-01-09T13:16:28.45597Z","shell.execute_reply":"2026-01-09T13:16:28.477661Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<h1 id=\"setup-and-configuration\"\n    style=\"\n  background: linear-gradient(135deg, #0f2027, #203a43, #2c5364);\n  color: white;\n  padding: 12px 24px;\n  text-align: center;\n  border-radius: 12px;\n  font-size: 22px;\n  font-weight: bold;\n  font-family: Arial, sans-serif;\n  margin: 20px 0;\n\">\n  8. Run Training\n</h1>","metadata":{}},{"cell_type":"code","source":"# if __name__ == \"__main__\":\n#     best_loss, best_predictions = train_model(\n#         model=model,\n#         train_dl=train_loader,\n#         val_dl=val_loader,\n#         epochs=20,\n#         cos_epoch=15,\n#         lr=3e-4,\n#         clip=1\n#     )\n#     print(f\"Best Validation Loss: {best_loss:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-09T13:16:28.479195Z","iopub.execute_input":"2026-01-09T13:16:28.479404Z","iopub.status.idle":"2026-01-09T13:17:51.63203Z","shell.execute_reply.started":"2026-01-09T13:16:28.479386Z","shell.execute_reply":"2026-01-09T13:17:51.631098Z"},"_kg_hide-input":false},"outputs":[],"execution_count":null}]}