{"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":11403143,"sourceType":"competition"},{"sourceId":7395079,"sourceType":"datasetVersion","datasetId":4299455},{"sourceId":7639698,"sourceType":"datasetVersion","datasetId":4299272},{"sourceId":224703571,"sourceType":"kernelVersion"},{"sourceId":228399841,"sourceType":"kernelVersion"}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport torch\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport torch\nimport random\nimport pickle","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-19T07:00:55.967825Z","iopub.execute_input":"2025-03-19T07:00:55.968096Z","iopub.status.idle":"2025-03-19T07:01:00.037009Z","shell.execute_reply.started":"2025-03-19T07:00:55.968066Z","shell.execute_reply":"2025-03-19T07:01:00.036181Z"}},"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    \"structural_violation_epoch\": 50,\n    \"balance_weight\": False,\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-19T07:01:00.038002Z","iopub.execute_input":"2025-03-19T07:01:00.038477Z","iopub.status.idle":"2025-03-19T07:01:00.043207Z","shell.execute_reply.started":"2025-03-19T07:01:00.038449Z","shell.execute_reply":"2025-03-19T07:01:00.042354Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_data=pd.read_csv(\"/kaggle/input/stanford-ribonanza-2-rna-folding-in-3-d/test_sequences.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-19T07:01:00.044042Z","iopub.execute_input":"2025-03-19T07:01:00.044261Z","iopub.status.idle":"2025-03-19T07:01:00.077103Z","shell.execute_reply.started":"2025-03-19T07:01:00.044243Z","shell.execute_reply":"2025-03-19T07:01:00.076366Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader\n\nclass RNADataset(Dataset):\n    def __init__(self,data):\n        self.data=data\n        self.tokens={nt:i for i,nt in enumerate('ACGU')}\n\n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self, idx):\n        sequence=[self.tokens[nt] for nt in (self.data.loc[idx,'sequence'])]\n        sequence=np.array(sequence)\n        sequence=torch.tensor(sequence)\n\n\n\n\n        return {'sequence':sequence}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-19T07:01:00.077860Z","iopub.execute_input":"2025-03-19T07:01:00.078076Z","iopub.status.idle":"2025-03-19T07:01:00.086750Z","shell.execute_reply.started":"2025-03-19T07:01:00.078058Z","shell.execute_reply":"2025-03-19T07:01:00.086069Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_dataset=RNADataset(test_data)\ntest_dataset[0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-19T07:01:00.087522Z","iopub.execute_input":"2025-03-19T07:01:00.087834Z","iopub.status.idle":"2025-03-19T07:01:00.145023Z","shell.execute_reply.started":"2025-03-19T07:01:00.087804Z","shell.execute_reply":"2025-03-19T07:01:00.144180Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\n\nsys.path.append(\"/kaggle/input/ribonanzanet2d-final\")\n\nimport torch.nn as nn\nfrom Network import RibonanzaNet, MultiHeadAttention\nimport yaml\n\nclass SimpleStructureModule(nn.Module):\n\n    def __init__(self, d_model, nhead, \n                 dim_feedforward, pairwise_dimension, dropout=0.1,\n                 ):\n        super(SimpleStructureModule, self).__init__()\n        #self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)\n        self.self_attn = MultiHeadAttention(d_model, nhead, d_model//nhead, d_model//nhead, dropout=dropout)\n\n\n        self.linear1 = nn.Linear(d_model, dim_feedforward)\n        self.dropout = nn.Dropout(dropout)\n        self.linear2 = nn.Linear(dim_feedforward, d_model)\n\n        self.norm1 = nn.LayerNorm(d_model)\n        self.norm2 = nn.LayerNorm(d_model)\n        self.norm3 = nn.LayerNorm(d_model)\n        #self.norm4 = nn.LayerNorm(d_model)\n        self.dropout1 = nn.Dropout(dropout)\n        self.dropout2 = nn.Dropout(dropout)\n        self.dropout3 = nn.Dropout(dropout)\n        #self.dropout4 = nn.Dropout(dropout)\n\n        self.pairwise2heads=nn.Linear(pairwise_dimension,nhead,bias=False)\n        self.pairwise_norm=nn.LayerNorm(pairwise_dimension)\n\n        self.distance2heads=nn.Linear(1,nhead,bias=False)\n        #self.pairwise_norm=nn.LayerNorm(pairwise_dimension)\n\n        self.activation = nn.GELU()\n\n        \n    def custom(self, module):\n        def custom_forward(*inputs):\n            inputs = module(*inputs)\n            return inputs\n        return custom_forward\n\n    def forward(self, input):\n        src , pairwise_features, pred_t, src_mask = input\n        \n        #src = src*src_mask.float().unsqueeze(-1)\n\n        pairwise_bias=self.pairwise2heads(self.pairwise_norm(pairwise_features)).permute(0,3,1,2)\n\n        \n        distance_matrix=pred_t[None,:,:]-pred_t[:,None,:]\n        distance_matrix=(distance_matrix**2).sum(-1).clip(2,37**2).sqrt()\n        distance_matrix=distance_matrix[None,:,:,None]\n        distance_bias=self.distance2heads(distance_matrix).permute(0,3,1,2)\n\n                    \n        \n        pairwise_bias=pairwise_bias+distance_bias\n\n        #print(src.shape)\n        src2,attention_weights = self.self_attn(src, src, src, mask=pairwise_bias, src_mask=src_mask)\n        \n\n        src = src + self.dropout1(src2)\n        src = self.norm1(src)\n        src2 = self.linear2(self.dropout(self.activation(self.linear1(src))))\n        src = src + self.dropout2(src2)\n        src = self.norm2(src)\n\n\n        return src\n\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\ndef load_config_from_yaml(file_path):\n    with open(file_path, 'r') as file:\n        config = yaml.safe_load(file)\n    return Config(**config)\n\n\n\nclass finetuned_RibonanzaNet(RibonanzaNet):\n    def __init__(self, config, pretrained=False):\n        config.dropout=0.1\n        config.use_grad_checkpoint=True\n        super(finetuned_RibonanzaNet, self).__init__(config)\n        if pretrained:\n            self.load_state_dict(torch.load(\"/kaggle/input/ribonanzanet-weights/RibonanzaNet.pt\",map_location='cpu'))\n        # self.ct_predictor=nn.Sequential(nn.Linear(64,256),\n        #                                 nn.ReLU(),\n        #                                 nn.Linear(256,64),\n        #                                 nn.ReLU(),\n        #                                 nn.Linear(64,1)) \n        self.dropout=nn.Dropout(0.0)\n\n        self.structure_module=SimpleStructureModule(d_model=256, nhead=8, \n                 dim_feedforward=1024, pairwise_dimension=64)\n        \n        self.xyz_predictor=nn.Linear(256,3)\n\n    def custom(self, module):\n        def custom_forward(*inputs):\n            inputs = module(*inputs)\n            return inputs\n        return custom_forward\n    \n    def forward(self,src):\n        \n        #with torch.no_grad():\n        sequence_features, pairwise_features=self.get_embeddings(src, torch.ones_like(src).long().to(src.device))\n        \n        xyzs=[]\n        xyz=torch.zeros(sequence_features.shape[1],3).cuda().float()\n        #print(xyz.shape)\n        #xyz=self.xyz_predictor(sequence_features)\n\n        for i in range(18):\n            sequence_features=self.structure_module([sequence_features,pairwise_features,xyz,None])\n            xyz=xyz+self.xyz_predictor(sequence_features).squeeze(0)\n            xyzs.append(xyz)\n            \n        \n        return xyzs\n\nmodel=finetuned_RibonanzaNet(load_config_from_yaml(\"/kaggle/input/ribonanzanet2d-final/configs/pairwise.yaml\"),pretrained=False).cuda()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-19T07:03:27.072174Z","iopub.execute_input":"2025-03-19T07:03:27.072566Z","iopub.status.idle":"2025-03-19T07:03:27.646418Z","shell.execute_reply.started":"2025-03-19T07:03:27.072535Z","shell.execute_reply":"2025-03-19T07:03:27.645497Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model=finetuned_RibonanzaNet(load_config_from_yaml(\"/kaggle/input/ribonanzanet2d-final/configs/pairwise.yaml\"),pretrained=False).cuda()\n\nmodel.load_state_dict(torch.load(\"/kaggle/input/ribonanzanet-3d-finetune-add-structure-module/RibonanzaNet-3D.pt\"))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-19T07:03:29.461182Z","iopub.execute_input":"2025-03-19T07:03:29.461621Z","iopub.status.idle":"2025-03-19T07:03:30.563822Z","shell.execute_reply.started":"2025-03-19T07:03:29.461587Z","shell.execute_reply":"2025-03-19T07:03:30.562939Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_dataset[0]['sequence'].shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-19T07:03:32.037585Z","iopub.execute_input":"2025-03-19T07:03:32.037945Z","iopub.status.idle":"2025-03-19T07:03:32.043543Z","shell.execute_reply.started":"2025-03-19T07:03:32.037916Z","shell.execute_reply":"2025-03-19T07:03:32.042514Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\npreds=[]\nfor i in range(len(test_dataset)):\n    src=test_dataset[i]['sequence'].long()\n    src=src.unsqueeze(0).cuda()\n\n    model.train()\n\n    tmp=[]\n    for i in range(4):\n        with torch.no_grad():\n            xyz=model(src)[-1].squeeze()\n        tmp.append(xyz.cpu().numpy())\n\n    model.eval()\n    with torch.no_grad():\n        xyz=model(src)[-1].squeeze()\n    tmp.append(xyz.cpu().numpy())\n\n    tmp=np.stack(tmp,0)\n    #exit()\n    preds.append(tmp)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-19T07:03:51.275216Z","iopub.execute_input":"2025-03-19T07:03:51.275499Z","iopub.status.idle":"2025-03-19T07:04:04.051150Z","shell.execute_reply.started":"2025-03-19T07:03:51.275476Z","shell.execute_reply":"2025-03-19T07:04:04.050406Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import plotly.graph_objects as go\nimport numpy as np\n\n# Example: Generate an Nx3 matrix\n\nxyz = preds[2][0]  # Replace this with your actual Nx3 data\nN = len(xyz)\n\n# Extract columns\nx, y, z = xyz[:, 0], xyz[:, 1], xyz[:, 2]\n\n# Create the 3D scatter plot\nfig = go.Figure(data=[go.Scatter3d(\n    x=x, y=y, z=z,\n    mode='markers',\n    marker=dict(\n        size=5,\n        color=z,  # Coloring based on z-value\n        colorscale='Viridis',  # Choose a colorscale\n        opacity=0.8\n    )\n)])\n\n# Customize layout\nfig.update_layout(\n    scene=dict(\n        xaxis_title=\"X\",\n        yaxis_title=\"Y\",\n        zaxis_title=\"Z\"\n    ),\n    title=\"3D Scatter Plot\"\n)\n\n# Show figure\nfig.show(renderer='iframe')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-19T07:05:26.641664Z","iopub.execute_input":"2025-03-19T07:05:26.641997Z","iopub.status.idle":"2025-03-19T07:05:26.715136Z","shell.execute_reply.started":"2025-03-19T07:05:26.641971Z","shell.execute_reply":"2025-03-19T07:05:26.714252Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ID=[]\nresname=[]\nresid=[]\nx=[]\ny=[]\nz=[]\n\ndata=[]\n\nfor i in range(len(test_data)):\n    #print(test_data.loc[i])\n\n    \n    for j in range(len(test_data.loc[i,'sequence'])):\n        # ID.append(test_data.loc[i,'sequence_id']+f\"_{j+1}\")\n        # resname.append(test_data.loc[i,'sequence'][j])\n        # resid.append(j+1) # 1 indexed\n        row=[test_data.loc[i,'target_id']+f\"_{j+1}\",\n             test_data.loc[i,'sequence'][j],\n             j+1]\n\n        for k in range(5):\n            for kk in range(3):\n                row.append(preds[i][k][j][kk])\n        data.append(row)\n\ncolumns=['ID','resname','resid']\nfor i in range(1,6):\n    columns+=[f\"x_{i}\"]\n    columns+=[f\"y_{i}\"]\n    columns+=[f\"z_{i}\"]\n\n\nsubmission=pd.DataFrame(data,columns=columns)\n\n\nsubmission\nsubmission.to_csv('submission.csv',index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-19T07:01:02.221617Z","iopub.status.idle":"2025-03-19T07:01:02.221976Z","shell.execute_reply":"2025-03-19T07:01:02.221825Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-19T07:01:02.222622Z","iopub.status.idle":"2025-03-19T07:01:02.222913Z","shell.execute_reply":"2025-03-19T07:01:02.222777Z"}},"outputs":[],"execution_count":null}]}