{"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":6233380,"sourceType":"datasetVersion","datasetId":3580819},{"sourceId":7395079,"sourceType":"datasetVersion","datasetId":4299455},{"sourceId":7639698,"sourceType":"datasetVersion","datasetId":4299272},{"sourceId":10197079,"sourceType":"datasetVersion","datasetId":6300687},{"sourceId":10855324,"sourceType":"datasetVersion","datasetId":6742586},{"sourceId":10878276,"sourceType":"datasetVersion","datasetId":6758842},{"sourceId":10878463,"sourceType":"datasetVersion","datasetId":6759157},{"sourceId":10880297,"sourceType":"datasetVersion","datasetId":6760419},{"sourceId":10880353,"sourceType":"datasetVersion","datasetId":6760463},{"sourceId":10880374,"sourceType":"datasetVersion","datasetId":6760482},{"sourceId":10880419,"sourceType":"datasetVersion","datasetId":6760509},{"sourceId":11419352,"sourceType":"datasetVersion","datasetId":891011},{"sourceId":224703571,"sourceType":"kernelVersion"},{"sourceId":224896926,"sourceType":"kernelVersion"},{"sourceId":232583388,"sourceType":"kernelVersion"},{"sourceId":235531625,"sourceType":"kernelVersion"}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### 💡Main Idea of This Notebook💡\n\n\nHigh-score notebooks are coming public today, so some of you might be wondering how to combine them to make better predictions. Here is a baseline for this!\n\n・Call Rhofold class\n\n\n・Precompute embeddings\n\nFor Rhofold, apply hook to recycle.embnet . Recycle.embnet encodes outputs from previous cycles into embeddings reused to iteratively refine RNA 3D structure predictions. \n\n\nFor RibonanzaNet, simply use get_embeddings function(this somehow does not work in rhofold)\n\n\n・Fusion Module \n\nIn this notebook, simple linear layer\n\nRhofold's seq_emb is (1, L, 256) L is sequence length, and RibonanzaNet's embedding before xyz_predictor is also (1, L,256).\n\nFusion Module in this notebook converts these into a fused (1, L, 256) embeddings and pass them to xyz_predictor.\n\n\n・Training\n\nIn this notebook, train fusion layer only\n\nIn this notebook, xyz_predictor loads weights from RibonanzaNet finetune notebook.\n\n\n・Inference&Visualization","metadata":{"_kg_hide-input":false}},{"cell_type":"markdown","source":"### Why do I started this?\n\nI started from 2 questions:\n\n・Is ensembling possible in structure prediction?\n\n・RibonanzaNet is originally for reactivity prediction, but can those embeddings be useful for structure prediction when supplemented by other models?","metadata":{}},{"cell_type":"markdown","source":"### 📚Insights obtained from this notebook📚\n\nFor visualization result, see the bottom of this notebook.\n\nTo compare with Ribonanzabet's prediction and Rhofold's prediction, visualization is stored here:https://www.kaggle.com/code/nanacat0520/stanfordrna-submission-visualizer-ribonanza\n\nhttps://www.kaggle.com/code/nanacat0520/stanfordrna-submission-visualizer-rho-ribonanza\n\n\n\n・Although score is almost the same as RibonanzaNet, ensembled predictions look like they recognize more secondary structures\n\n・Stronger decoder stack will improve predictions\n","metadata":{}},{"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 argparse\nimport os\nimport sys\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T16:58:02.565074Z","iopub.execute_input":"2025-04-22T16:58:02.565396Z","iopub.status.idle":"2025-04-22T16:58:02.569476Z","shell.execute_reply.started":"2025-04-22T16:58:02.565369Z","shell.execute_reply":"2025-04-22T16:58:02.568620Z"},"scrolled":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### data processing& loss function","metadata":{"_kg_hide-input":true,"_kg_hide-output":true}},{"cell_type":"code","source":"!pip install /kaggle/input/openmm/OpenMM-8.2.0-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T16:58:02.570652Z","iopub.execute_input":"2025-04-22T16:58:02.570873Z","iopub.status.idle":"2025-04-22T16:58:06.042141Z","shell.execute_reply.started":"2025-04-22T16:58:02.570854Z","shell.execute_reply":"2025-04-22T16:58:06.041087Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install /kaggle/input/simtk-0-1/simtk-0.1.0-py2.py3-none-any.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T16:58:06.044011Z","iopub.execute_input":"2025-04-22T16:58:06.044338Z","iopub.status.idle":"2025-04-22T16:58:09.382723Z","shell.execute_reply.started":"2025-04-22T16:58:06.044307Z","shell.execute_reply":"2025-04-22T16:58:09.381898Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install /kaggle/input/pytest-runner/pytest_runner-6.0.1-py3-none-any.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T16:58:09.384488Z","iopub.execute_input":"2025-04-22T16:58:09.384869Z","iopub.status.idle":"2025-04-22T16:58:12.706206Z","shell.execute_reply.started":"2025-04-22T16:58:09.384834Z","shell.execute_reply":"2025-04-22T16:58:12.705128Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install /kaggle/input/biopython/biopython-1.85-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T16:58:12.707311Z","iopub.execute_input":"2025-04-22T16:58:12.707654Z","iopub.status.idle":"2025-04-22T16:58:16.082278Z","shell.execute_reply.started":"2025-04-22T16:58:12.707629Z","shell.execute_reply":"2025-04-22T16:58:16.081182Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install /kaggle/input/ml-collections/ml_collections-1.0.0-py3-none-any.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T16:58:16.083336Z","iopub.execute_input":"2025-04-22T16:58:16.083612Z","iopub.status.idle":"2025-04-22T16:58:19.428854Z","shell.execute_reply.started":"2025-04-22T16:58:16.083589Z","shell.execute_reply":"2025-04-22T16:58:19.428044Z"}},"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-22T16:58:19.431478Z","iopub.execute_input":"2025-04-22T16:58:19.431747Z","iopub.status.idle":"2025-04-22T16:58:19.437404Z","shell.execute_reply.started":"2025-04-22T16:58:19.431725Z","shell.execute_reply":"2025-04-22T16:58:19.436632Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_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-22T16:58:19.438680Z","iopub.execute_input":"2025-04-22T16:58:19.438908Z","iopub.status.idle":"2025-04-22T16:58:19.645513Z","shell.execute_reply.started":"2025-04-22T16:58:19.438888Z","shell.execute_reply":"2025-04-22T16:58:19.644627Z"}},"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-22T16:58:19.646592Z","iopub.execute_input":"2025-04-22T16:58:19.646917Z","iopub.status.idle":"2025-04-22T16:58:19.738279Z","shell.execute_reply.started":"2025-04-22T16:58:19.646891Z","shell.execute_reply":"2025-04-22T16:58:19.737178Z"}},"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\": 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-22T16:58:19.739332Z","iopub.execute_input":"2025-04-22T16:58:19.739673Z","iopub.status.idle":"2025-04-22T16:58:19.745214Z","shell.execute_reply.started":"2025-04-22T16:58:19.739639Z","shell.execute_reply":"2025-04-22T16:58:19.744074Z"}},"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\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, pred_y, gt_x, gt_y, epsilon=1e-4, Z=10, d_clamp=None):\n    \"\"\"\n    De-centered Root Mean Absolute Error loss with optional distance clamping.\n    \"\"\"\n    # shape adjustment\n    for tensor in [pred_x, pred_y, gt_x, gt_y]:\n        assert tensor.ndim in [2, 3], f\"Expected 2D or 3D tensor but got shape {tensor.shape}\"\n    \n    # remove batch dim if exists\n    if pred_x.ndim == 3:\n        pred_x = pred_x.squeeze(0)\n    if pred_y.ndim == 3:\n        pred_y = pred_y.squeeze(0)\n    if gt_x.ndim == 3:\n        gt_x = gt_x.squeeze(0)\n    if gt_y.ndim == 3:\n        gt_y = gt_y.squeeze(0)\n\n    # shape check\n    assert pred_x.shape == gt_x.shape, f\"Shape mismatch: pred_x {pred_x.shape}, gt_x {gt_x.shape}\"\n\n    # pairwise distance matrix\n    pred_dm = torch.cdist(pred_x, pred_y)  # [L, L]\n    gt_dm = torch.cdist(gt_x, gt_y)        # [L, L]\n\n    # create mask (exclude diagonal)\n    mask = torch.ones_like(gt_dm, dtype=torch.bool)\n    mask.fill_diagonal_(False)\n\n    if d_clamp is not None:\n        pred_dm = pred_dm.clamp(0, d_clamp)\n        gt_dm = gt_dm.clamp(0, d_clamp)\n\n    rmsd = torch.abs(pred_dm[mask] - gt_dm[mask])\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-22T16:58:19.746348Z","iopub.execute_input":"2025-04-22T16:58:19.746741Z","iopub.status.idle":"2025-04-22T16:58:19.764289Z","shell.execute_reply.started":"2025-04-22T16:58:19.746701Z","shell.execute_reply":"2025-04-22T16:58:19.763466Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Precomputing embeddings from Rhofold&RibonanzaNet","metadata":{}},{"cell_type":"code","source":"\nimport os\nimport sys\nimport shutil\nimport torch\nimport numpy as np\nimport pandas as pd\nimport yaml\nimport random\nfrom tqdm import tqdm\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import GradScaler\n\n#set seed\ntorch.manual_seed(0)\nnp.random.seed(0)\nrandom.seed(0)\n\n# ============================\n# importing RhoFold \n# ============================\nsrc_dir = \"/kaggle/input/rhofold-repo/rhofold\"\ndst_dir = \"/kaggle/working/rhofold\"\n\nif not os.path.exists(dst_dir):\n    shutil.copytree(src_dir, dst_dir)\n\ndef rewrite_imports(filepath, from_level=1):\n    if not os.path.exists(filepath):\n        return\n    with open(filepath, \"r\", encoding=\"utf-8\") as f:\n        lines = f.readlines()\n    new_lines = []\n    for line in lines:\n        if \"from rhofold.\" in line:\n            new_lines.append(line.replace(\"from rhofold.\", \"from \" + \".\" * from_level))\n        elif \"import rhofold.\" in line:\n            new_lines.append(line.replace(\"import rhofold.\", \"import \" + \".\" * from_level))\n        else:\n            new_lines.append(line)\n    with open(filepath, \"w\", encoding=\"utf-8\") as f:\n        f.writelines(new_lines)\n\nrewrite_imports(os.path.join(dst_dir, \"rhofold.py\"), from_level=1)\nrewrite_imports(os.path.join(dst_dir, \"data\", \"data_pipeline.py\"), from_level=2)\nrewrite_imports(os.path.join(dst_dir, \"utils\", \"tensor_utils.py\"), from_level=2)\nsys.path.insert(0, \"/kaggle/working\")\n\n# ============================\n# RhoFold&RibonanzaNet&Fusion Layer\n# ============================\nfrom rhofold.rhofold import RhoFold\nfrom rhofold.config import rhofold_config as rho_config\nfrom rhofold.utils.alphabet import get_features\n\nsys.path.append(\"/kaggle/input/ribonanzanet2d-final\")\nfrom Network import RibonanzaNet\n\nclass Config:\n    def __init__(self, **entries):\n        self.__dict__.update(entries)\n        self.entries=entries\n\n    def print(self):\n        print(self.entries)\n\ndef load_config_from_yaml(file_path):\n    with open(file_path, 'r') as file:\n        config = yaml.safe_load(file)\n    return Config(**config)\n\nclass finetuned_RibonanzaNet(RibonanzaNet):\n    def __init__(self, config, pretrained=False):\n        config.dropout=0.2\n        super(finetuned_RibonanzaNet, self).__init__(config)\n        if pretrained:\n            self.load_state_dict(torch.load(\"/kaggle/input/ribonanzanet-3d-finetune/RibonanzaNet-3D-final.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        self.xyz_predictor=nn.Linear(256,3)\n    def get_finetuned_embeddings(self, src):\n        if src.dim() == 1:\n            src = src.unsqueeze(0)\n        seq_feat, pair_feat = self.get_embeddings(src, torch.ones_like(src).long().to(src.device))\n        return {\n            'sequence_features': seq_feat,\n            'pairwise_features': pair_feat,\n            'xyz_predictions': self.xyz_predictor(seq_feat)\n        }\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        xyz=self.xyz_predictor(sequence_features)\n\n        return xyz\n\nclass FusionLayer(nn.Module):\n    def __init__(self, input_dim=256):\n        super().__init__()\n        self.rhofold_transform = nn.Linear(input_dim, input_dim)\n        self.ribonanza_transform = nn.Linear(input_dim, input_dim)\n        self.fusion_layer = nn.Linear(input_dim * 2, input_dim)\n\n    def forward(self, rhofold_emb, ribonanza_emb):\n        r_emb = self.rhofold_transform(rhofold_emb)\n        r2_emb = self.ribonanza_transform(ribonanza_emb)\n        concat = torch.cat([r_emb, r2_emb], dim=-1)\n        return self.fusion_layer(concat)\n\n\ndef get_final_embeddings(fasta_path, ckpt_path=\"/kaggle/input/rhofold-repo/pretrained/RhoFold_pretrained.pt\"):\n    model_rho = RhoFold(rho_config).to(\"cuda\").eval()\n    model_rho.structure_module.refinenet = None\n    checkpoint = torch.load(ckpt_path, map_location=\"cuda\")\n    model_rho.load_state_dict(checkpoint[\"model\"])\n    embeddings = {}\n    def hook_fn(name):\n        def hook(module, input, output):\n            embeddings[name] = output.detach().clone() if isinstance(output, torch.Tensor) else output[0].detach().clone()\n        return hook\n    hooks = [module.register_forward_hook(hook_fn(name)) for name, module in model_rho.named_modules() if 'recycle_embnet' in name]\n    data_dict = get_features(fasta_path, fasta_path)\n    with torch.no_grad():\n        _ = model_rho(tokens=data_dict['tokens'].to(\"cuda\"), rna_fm_tokens=data_dict['rna_fm_tokens'].to(\"cuda\"), seq=data_dict['seq'])\n    for hook in hooks:\n        hook.remove()\n    return {k: v.cpu() for k, v in embeddings.items()}\n\n\n# ============================\n# DataLoader\n# ============================\ntrain_sequences = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_sequences.csv\")\ntrain_labels = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_labels.csv\")\ntrain_labels[\"pdb_id\"] = train_labels[\"ID\"].apply(lambda x: x.split(\"_\")[0] + \"_\" + x.split(\"_\")[1])\n\nall_xyz = []\nfor pdb_id in tqdm(train_sequences['target_id']):\n    df = train_labels[train_labels[\"pdb_id\"] == pdb_id]\n    xyz = df[['x_1', 'y_1', 'z_1']].to_numpy().astype('float32')\n    xyz[xyz < -1e17] = float('nan')\n    all_xyz.append(xyz)\n\nconfig_data = {'cutoff_date': '2021-12-01', 'test_cutoff_date': '2022-08-01', 'min_len_filter': 50, 'max_len_filter': 400, 'max_len': 400}\n\nfilter_nan = [(np.isnan(xyz).mean() <= 0.5) and (config_data['min_len_filter'] < len(xyz) < config_data['max_len_filter']) for xyz in all_xyz]\nnon_nan_indices = np.where(filter_nan)[0]\ntrain_sequences = train_sequences.loc[non_nan_indices].reset_index(drop=True)\nall_xyz = [all_xyz[i] for i in non_nan_indices]\n\ndata = {\n    \"sequence\": train_sequences['sequence'].tolist(),\n    \"temporal_cutoff\": train_sequences['temporal_cutoff'].tolist(),\n    \"xyz\": all_xyz\n}\n\ncutoff = pd.Timestamp(config_data['cutoff_date'])\ntest_cutoff = pd.Timestamp(config_data['test_cutoff_date'])\ntrain_index = [i for i, d in enumerate(data['temporal_cutoff']) if pd.Timestamp(d) <= cutoff]\n\nclass RNA3D_Dataset(Dataset):\n    def __init__(self, indices, data):\n        self.indices = indices\n        self.data = data\n        self.tokens = {nt: i for i, nt in enumerate('ACGU')}\n\n    def __len__(self):\n        return len(self.indices)\n\n    def __getitem__(self, idx):\n        idx = self.indices[idx]\n        sequence = torch.tensor([self.tokens[nt] for nt in self.data['sequence'][idx]], dtype=torch.long)\n        sequence = sequence.unsqueeze(0)\n        xyz = torch.tensor(self.data['xyz'][idx], dtype=torch.float32)\n        if len(sequence) > config_data['max_len']:\n            start = np.random.randint(len(sequence) - config_data['max_len'])\n            sequence = sequence[start:start + config_data['max_len']]\n            xyz = xyz[start:start + config_data['max_len']]\n        return {'sequence': sequence, 'xyz': xyz, 'target_id': train_sequences.iloc[idx]['target_id']}\n\ntrain_dataset = RNA3D_Dataset(train_index, data)\ntrain_loader = DataLoader(train_dataset, batch_size=1, shuffle=True)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T16:58:19.765304Z","iopub.execute_input":"2025-04-22T16:58:19.765680Z","iopub.status.idle":"2025-04-22T16:58:28.311302Z","shell.execute_reply.started":"2025-04-22T16:58:19.765650Z","shell.execute_reply":"2025-04-22T16:58:28.310582Z"},"_kg_hide-output":true,"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_ribo = finetuned_RibonanzaNet(load_config_from_yaml(\"/kaggle/input/ribonanzanet2d-final/configs/pairwise.yaml\"), pretrained=False).cuda()\nmodel_ribo.load_state_dict(torch.load(\"/kaggle/input/ribonanzanet-3d-finetune/RibonanzaNet-3D.pt\"))\nmodel_ribo.eval()\nfusion_module = FusionLayer().cuda().train()\noptimizer = torch.optim.Adam(fusion_module.parameters(), lr=1e-4)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T16:58:28.312152Z","iopub.execute_input":"2025-04-22T16:58:28.312423Z","iopub.status.idle":"2025-04-22T16:58:29.640856Z","shell.execute_reply.started":"2025-04-22T16:58:28.312391Z","shell.execute_reply":"2025-04-22T16:58:29.640177Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Precompute embeddings\ndef precompute_embeddings():\n    print(\"Calculating embeddings...\")\n    embeddings_cache = {}\n    \n    tokens = {nt: i for i, nt in enumerate('ACGU')}\n    \n    for idx in tqdm(train_index):\n        target_id = train_sequences.iloc[idx]['target_id']\n        sequence_tensor = torch.tensor([tokens[nt] for nt in data['sequence'][idx]], dtype=torch.long).unsqueeze(0)\n        \n        \n        if target_id in embeddings_cache:\n            continue\n            \n        fasta_path = f\"/kaggle/input/stanford-rna-3d-folding/MSA/{target_id}.MSA.fasta\"\n        if not os.path.exists(fasta_path):\n            continue\n            \n        # RhoFold embeddings\n        try:\n            rho_embs = get_final_embeddings(fasta_path)\n            recycle_embnet = next(iter(rho_embs.values())).cpu()  # CPUに保存してGPUメモリを節約\n            \n            # Ribonanza embeddings\n            with torch.no_grad():\n                sequence_cuda = sequence_tensor.cuda()\n                ribo_embs = model_ribo.get_finetuned_embeddings(sequence_cuda)\n                ribo_seq_emb = ribo_embs['sequence_features'].cpu()  # CPUに保存\n                \n            #saving to cache\n            embeddings_cache[target_id] = {\n                'rho_emb': recycle_embnet,\n                'ribo_emb': ribo_seq_emb\n            }\n            \n            del rho_embs, recycle_embnet, ribo_embs, ribo_seq_emb\n            torch.cuda.empty_cache()\n            \n        except Exception as e:\n            print(f\"Error processing {target_id}: {e}\")\n            continue\n    \n    print(f\"embeddings successfully calculated! {len(embeddings_cache)}entries\")\n    return embeddings_cache\n#Calculate embeddings\nembeddings_cache = precompute_embeddings()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T16:58:29.641635Z","iopub.execute_input":"2025-04-22T16:58:29.641856Z","execution_failed":"2025-04-22T17:02:43.782Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pickle\nimport os\nimport torch\n\n\noutput_dir = \"/kaggle/working/embeddings_cache\"\nos.makedirs(output_dir, exist_ok=True)\n\n# 1. Saving embeddings to Pickle file \ndef save_cache_pickle(embeddings_cache, filename=\"embeddings_cache.pkl\"):\n    file_path = os.path.join(output_dir, filename)\n    print(f\"saving as pickle: {file_path}\")\n    \n    with open(file_path, \"wb\") as f:\n        pickle.dump(embeddings_cache, f)\n    \n    print(f\"Embeddings saved! size: {os.path.getsize(file_path) / (1024 * 1024):.2f} MB\")\n    return file_path\n\n\n\n\n\n\ndef load_cache_pickle(filename=\"embeddings_cache.pkl\"):\n    file_path = os.path.join(output_dir, filename)\n    print(f\"Loading cache from Pickle: {file_path}\")\n    \n    with open(file_path, \"rb\") as f:\n        embeddings_cache = pickle.load(f)\n    \n    print(f\":Successfully loaded cache! size: {len(embeddings_cache)}\")\n    return embeddings_cache\n\n\n\n\ntry:\n    pickle_path = save_cache_pickle(embeddings_cache)\n    \n    print(\"Embeddings Successfully Saved!\")\nexcept Exception as e:\n    print(f\"File save error: {e}\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-04-22T17:02:43.783Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"エンベディングキャッシュをロード中...\")\nembeddings_cache = load_cache_pickle()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-04-22T17:02:43.783Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Training","metadata":{}},{"cell_type":"code","source":"# ============================\n# Dataset to load embeddings \n# ============================\nclass RNA3D_Dataset_WithCache(Dataset):\n    def __init__(self, indices, data, embeddings_cache):\n        self.indices = indices\n        self.data = data\n        self.tokens = {nt: i for i, nt in enumerate('ACGU')}\n        self.embeddings_cache = embeddings_cache\n        \n    def __len__(self):\n        return len(self.indices)\n        \n    def __getitem__(self, idx):\n        idx = self.indices[idx]\n        sequence = torch.tensor([self.tokens[nt] for nt in self.data['sequence'][idx]], dtype=torch.long)\n        sequence = sequence.unsqueeze(0)\n        xyz = torch.tensor(self.data['xyz'][idx], dtype=torch.float32)\n        target_id = train_sequences.iloc[idx]['target_id']\n        \n        # get embeddings\n        embeddings = None\n        if target_id in self.embeddings_cache:\n            embeddings = self.embeddings_cache[target_id]\n        \n        if len(sequence) > config_data['max_len']:\n            start = np.random.randint(len(sequence) - config_data['max_len'])\n            sequence = sequence[start:start + config_data['max_len']]\n            xyz = xyz[start:start + config_data['max_len']]\n            \n        return {'sequence': sequence, 'xyz': xyz, 'target_id': target_id, 'embeddings': embeddings}\n\n# create dataset and dataloader\ntrain_dataset = RNA3D_Dataset_WithCache(train_index, data, embeddings_cache)\ntrain_loader = DataLoader(train_dataset, batch_size=1, shuffle=True)\n\n# load models\nmodel_ribo = finetuned_RibonanzaNet(load_config_from_yaml(\"/kaggle/input/ribonanzanet2d-final/configs/pairwise.yaml\"), pretrained=False).cuda()\nmodel_ribo.load_state_dict(torch.load(\"/kaggle/input/ribonanzanet-3d-finetune/RibonanzaNet-3D.pt\"))\nmodel_ribo.eval()\n\n# prepare fusion module\nfusion_module = FusionLayer().cuda().train()\noptimizer = torch.optim.Adam(fusion_module.parameters(), lr=1e-4)\n\n\nscaler = GradScaler()\n\n\nepochs = 50\ncos_epoch = 30\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=(epochs - cos_epoch) * len(train_loader))\n\n# ============================\n# loading models\n# ============================\nmodel_ribo = finetuned_RibonanzaNet(load_config_from_yaml(\"/kaggle/input/ribonanzanet2d-final/configs/pairwise.yaml\"), pretrained=False).cuda()\nmodel_ribo.load_state_dict(torch.load(\"/kaggle/input/ribonanzanet-3d-finetune/RibonanzaNet-3D.pt\"))\nmodel_ribo.eval()\nfusion_module = FusionLayer().cuda().train()\noptimizer = torch.optim.Adam(fusion_module.parameters(), lr=1e-4)\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-04-22T17:02:43.783Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================\n# training loop\n# ============================\nepochs = 50\ncos_epoch = 30\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=(epochs - cos_epoch) * len(train_loader))\nbest_loss = float('inf')\n\nfor epoch in range(epochs):\n    model_ribo.eval()\n    fusion_module.train()\n    tbar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{epochs}\")\n    total_loss = 0\n    oom = 0\n    for idx, batch in enumerate(tbar): \n        try:\n            sequence = batch['sequence'].squeeze(0).cuda()\n            gt_xyz = batch['xyz'].squeeze(0).cuda()\n\n            if torch.isnan(gt_xyz).any():\n                print(f\"⚠️ NaN detected in gt_xyz at batch {idx}, skipping\")\n                continue\n\n            embeddings = batch['embeddings']\n            \n            \n            if embeddings is None:\n                continue\n                \n            \n            rho_emb = embeddings['rho_emb'].cuda().squeeze(0)\n            ribo_seq_emb = embeddings['ribo_emb'].cuda().squeeze(0)\n\n            fused_embedding = fusion_module(rho_emb, ribo_seq_emb)\n            pred_xyz = model_ribo.xyz_predictor(fused_embedding).squeeze(0)\n            if epoch == 0 :\n                print(\"📏 pred_xyz.shape:\", pred_xyz.shape)\n                print(\"📏 gt_xyz.shape:\", gt_xyz.shape)\n                print(\"❓ pred_xyz NaN:\", torch.isnan(pred_xyz).any().item())\n                print(\"❓ gt_xyz NaN:\", torch.isnan(gt_xyz).any().item())\n                print(\"❓ rho_emb NaN:\", torch.isnan(rho_emb).any().item())\n                print(\"❓ ribo_emb NaN:\", torch.isnan(ribo_seq_emb).any().item())\n\n\n            loss = dRMAE(pred_xyz, pred_xyz, gt_xyz, gt_xyz) + align_svd_mae(pred_xyz, gt_xyz)\n\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(fusion_module.parameters(), 1.0)\n            optimizer.step()\n            optimizer.zero_grad()\n            \n            if (epoch + 1) > cos_epoch:\n                scheduler.step()\n\n            total_loss += loss.item()\n            tbar.set_postfix(loss=total_loss / (tbar.n + 1), OOM=oom)\n\n            \n\n        except RuntimeError as e:\n            print(f\"❌ OOM or Error at batch {tbar.n}: {e}\")\n            oom += 1\n            torch.cuda.empty_cache()\n            import gc; gc.collect()\n            continue\n\n    print(f\"✅ Epoch {epoch+1} completed. Avg Loss: {total_loss / len(tbar):.6f}\")\n\n\ntorch.save(fusion_module.state_dict(), \"fusion_module_finetuned.pt\")\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-04-22T17:02:43.783Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nprint(\"✅ 最終 pred_xyz の shape:\", pred_xyz.shape)\nprint(\"🧾 pred_xyz の中身（先頭5個）:\", pred_xyz[:5].detach().cpu().numpy())\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-04-22T17:02:43.783Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Inference & Visualization","metadata":{}},{"cell_type":"code","source":"fusion_module = FusionLayer().to(\"cuda\")\nfusion_ckpt_path = \"/kaggle/working/fusion_module_finetuned.pt\"\nfusion_module.load_state_dict(torch.load(fusion_ckpt_path, map_location=\"cuda\"))\nfusion_module.eval()\n\n# ====================================\n# main function\n# ====================================\nimport gc\nimport torch\n\ndef tokenize_sequence(input_sequence):\n    tokens = {nt: i for i, nt in enumerate('ACGU')}\n    sequence = [tokens[nt] for nt in input_sequence]\n    sequence = torch.tensor(sequence).unsqueeze(0)\n    return sequence\n\ndef process_dataset(sequence_dict, output_dir=\"/kaggle/working/predictions\"):\n    os.makedirs(output_dir, exist_ok=True)\n\n    for i, (name, sequence) in enumerate(sequence_dict.items()):\n        print(f\"\\n🧬 [{i+1}/{len(sequence_dict)}] Processing {name}\")\n\n        try:\n            fasta_path = f\"/kaggle/input/stanford-rna-3d-folding/MSA/{name}.MSA.fasta\"\n            if not os.path.exists(fasta_path):\n                print(f\"❌ {fasta_path} not found. Skipping.\")\n                continue\n            # ---- RhoFold embeddings----\n            rho_embs = get_final_embeddings(fasta_path)\n            recycle_embnet = next(iter(rho_embs.values())).to(\"cuda\")\n\n            # ---- Ribonanza embeddings ----\n            input_tensor = tokenize_sequence(sequence).to(\"cuda\")\n            with torch.no_grad():\n                ribo_embs = model_ribo.get_finetuned_embeddings(input_tensor)\n            ribo_seq_emb = ribo_embs['sequence_features']\n            # ---- fusion&prediction----\n            with torch.no_grad():\n                fused_embedding = fusion_module(recycle_embnet, ribo_seq_emb)\n                xyz = model_ribo.xyz_predictor(fused_embedding).squeeze(0).cpu().numpy()\n\n            # ---- save prediction ----\n            df_xyz = pd.DataFrame(xyz, columns=[\"x_1\", \"y_1\", \"z_1\"])\n            csv_path = os.path.join(output_dir, f\"{name}_predicted.csv\")\n            df_xyz.to_csv(csv_path, index=False)\n            print(f\"✅ Saved: {csv_path}\")\n\n        except RuntimeError as e:\n            print(f\"❌ RuntimeError in {name}: {e}\")\n            torch.cuda.empty_cache()\n            gc.collect()\n            continue\n        \n        del rho_embs, recycle_embnet, input_tensor, ribo_embs, ribo_seq_emb, fused_embedding, xyz\n        torch.cuda.empty_cache()\n        gc.collect()\n\ndf = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/test_sequences.csv\")\nsequence_dict = dict(zip(df[\"target_id\"], df[\"sequence\"]))\nprocess_dataset(sequence_dict)\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-04-22T17:02:43.783Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample_submission = pd.read_csv('/kaggle/input/ribonanzanet-3d-inference/submission.csv')\n# load R1138 prediction from Ribonanzanet inference notebook because it runs oom in rhofold\ninput_csv = \"/kaggle/input/ribonanzanet-3d-inference/submission.csv\"\noutput_csv = \"/kaggle/working/predictions/R1138_predicted.csv\"\n\n\n\n\ndf = pd.read_csv(input_csv)\n\n\ndf_r1138 = df[df[\"ID\"].str.startswith(\"R1138\")].copy()\n\ndf_r1138 = df_r1138[[\"x_1\", \"y_1\", \"z_1\"]]\ndf_r1138.to_csv(output_csv, index=False)\nprint(f\"✅ R1138 のデータを {output_csv} に保存しました。\")\nimport glob\nimport re\n\n\ninput_dir = \"/kaggle/working/predictions\"\n\n\nrows = []\n\n# converting predictions to visualizer format\nfor csv_path in sorted(glob.glob(os.path.join(input_dir, \"*_predicted.csv\"))):\n    name = os.path.basename(csv_path).replace(\"_predicted.csv\", \"\")  # 例: \"R1138\"\n    df = pd.read_csv(csv_path)\n\n    sequence = sequence_dict.get(name, \"\")\n    for i in range(len(df)):\n        row = {\n            \"ID\": f\"{name}_{i+1}\",\n            \"resname\": sequence[i] if i < len(sequence) else \"N\",\n            \"resid\": i + 1,\n            \"x_1\": df.loc[i, \"x_1\"],\n            \"y_1\": df.loc[i, \"y_1\"],\n            \"z_1\": df.loc[i, \"z_1\"],\n        }\n        rows.append(row)\n\n\ndf_all = pd.DataFrame(rows)\n\ndef extract_sort_keys(id_str):\n    match = re.match(r\"R(\\d+)_(\\d+)\", id_str)\n    return (int(match.group(1)), int(match.group(2))) if match else (float('inf'), float('inf'))\n\ndf_all['sort_key'] = df_all['ID'].apply(extract_sort_keys)\ndf_all = df_all.sort_values(by='sort_key').drop(columns='sort_key')\n\n\ndf_all.reset_index(drop=True, inplace=True)\n\n\ndf_all.to_csv(\"/kaggle/working/fused_predictions_all.csv\", index=False)\nprint(\"✅ 並び替え＋インデックスリセット済みで fused_predictions_all.csv に保存しました！\")\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-04-22T17:02:43.783Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\nfused_path = \"/kaggle/working/fused_predictions_all.csv\"\nsample_path = \"/kaggle/input/ribonanzanet-3d-inference/submission.csv\"\noutput_path = \"/kaggle/working/submission_filled.csv\"\n\nnum_copies = 5\n\n\ndf_fused = pd.read_csv(fused_path)\ndf_sample = pd.read_csv(sample_path)\n\nfor i in range(1, num_copies + 1):\n    df_fused[f\"x_{i}\"] = df_fused[\"x_1\"]\n    df_fused[f\"y_{i}\"] = df_fused[\"y_1\"]\n    df_fused[f\"z_{i}\"] = df_fused[\"z_1\"]\n\n\nreplace_cols = [f\"{axis}_{i}\" for i in range(1, num_copies + 1) for axis in ['x', 'y', 'z']]\n\n\ndf_merged = df_sample.drop(columns=replace_cols).merge(\n    df_fused[[\"ID\"] + replace_cols],\n    on=\"ID\",\n    how=\"left\"\n)\n\n\ndf_merged.to_csv(output_path, index=False)\nprint(f\"✅ 置換後の submission を {output_path} に保存しました。\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-04-22T17:02:43.783Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"try:\n    import Bio\nexcept:\n    #for rhofold+ #####################\n    !pip install biopython\n    !pip install ml-collections\n    !pip install python-box\n    !pip install dm-tree\n    !pip install openmm[cuda12]\n\nfrom copy import deepcopy\n\nimport pandas as pd\nfrom Bio.PDB import Atom, Model, Chain, Residue, Structure, PDBParser\nfrom Bio import SeqIO\nimport os, sys\nimport re\nimport numpy as np\nimport pandas as pd\n\nimport matplotlib\nimport matplotlib.pyplot as plt\n!pip install kagglehub\nimport kagglehub\n\nprint('IMPORT OK !!!!')\nPYTHON = sys.executable\nprint('PYTHON',PYTHON)\n\nusalign_path = kagglehub.dataset_download('metric/usalign')\nUSALIGN = \\\n'/kaggle/working//USalign'\n#'<your us align path>/USalign'\n\nos.system('cp /kaggle/input/usalign/USalign /kaggle/working/')\nos.system(' chmod u+x /kaggle/working//USalign')\n\n\nsubmission = df_merged\nLABEL_DF = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/train_labels.csv')\n\nLABEL_DF[\"pdb_id\"] = LABEL_DF[\"ID\"].apply(lambda x: x.split(\"_\")[0]+'_'+x.split(\"_\")[1])\nsubmission['submission_id'] = submission['ID'].str.split('_').str[0]","metadata":{"trusted":true,"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sys.path.append('/kaggle/usr/lib/stanfordrna_submission_visualizer_ribonanza')\nfrom stanfordrna_submission_visualizer_ribonanza import submission_to_visual_2","metadata":{"trusted":true,"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission_to_visual_2(submission)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}