{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","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,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":12251485,"sourceType":"datasetVersion","datasetId":7719426}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import sys\nsys.path.append(\"/kaggle/input/rnet3d-ddpm-test80/test80_improved_inference_upload\")\n\n\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-06-23T02:11:48.550794Z","iopub.execute_input":"2025-06-23T02:11:48.551129Z","iopub.status.idle":"2025-06-23T02:11:48.554872Z","shell.execute_reply.started":"2025-06-23T02:11:48.551076Z","shell.execute_reply":"2025-06-23T02:11:48.553950Z"}},"outputs":[],"execution_count":null},{"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\nfrom utils import *\nimport os\n#from Diffusion import Diffusion\nimport argparse\nfrom tqdm import tqdm","metadata":{"execution":{"iopub.status.busy":"2025-06-23T02:11:48.556321Z","iopub.execute_input":"2025-06-23T02:11:48.556568Z","iopub.status.idle":"2025-06-23T02:11:48.570008Z","shell.execute_reply.started":"2025-06-23T02:11:48.556549Z","shell.execute_reply":"2025-06-23T02:11:48.569452Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Get args","metadata":{}},{"cell_type":"code","source":"# If running in a notebook, replace argparse with manual assignment\nclass Args:\n    target_csv = '/kaggle/input/stanford-rna-3d-folding/test_sequences.csv'\n    config = '/kaggle/input/rnet3d-ddpm-test80/test80_improved_inference_upload/recycle.yaml'\n    weights = '/kaggle/input/rnet3d-ddpm-test80/test80_improved_inference_upload/weights/recycle.yaml_RibonanzaNet_3D.pt'\n\nargs = Args()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T02:11:48.570564Z","iopub.execute_input":"2025-06-23T02:11:48.570737Z","iopub.status.idle":"2025-06-23T02:11:48.585833Z","shell.execute_reply.started":"2025-06-23T02:11:48.570723Z","shell.execute_reply":"2025-06-23T02:11:48.585307Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Get dataset","metadata":{}},{"cell_type":"code","source":"test_data=pd.read_csv(args.target_csv)#.loc[2:].reset_index(drop=True)\n\nfrom 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}\n\ntest_dataset=RNADataset(test_data)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T02:11:48.587424Z","iopub.execute_input":"2025-06-23T02:11:48.587602Z","iopub.status.idle":"2025-06-23T02:11:48.616636Z","shell.execute_reply.started":"2025-06-23T02:11:48.587588Z","shell.execute_reply":"2025-06-23T02:11:48.616005Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Get model","metadata":{}},{"cell_type":"code","source":"import torch.nn as nn\nfrom Diffusion import finetuned_RibonanzaNet\n\n\n\nargs.config\n\nconfig=load_config_from_yaml(args.config)\n\nmodel=finetuned_RibonanzaNet(load_config_from_yaml(\"/kaggle/input/rnet3d-ddpm-test80/test80_improved_inference_upload/pairwise.yaml\"),config,pretrained=False).cuda()\n#model.decode(torch.ones(1,10).long().cuda(),torch.ones(1,10).long().cuda())\n\n\nimport torch\nstate_dict=torch.load(args.weights,map_location='cpu')\n#state_dict=torch.load(\"RibonanzaNet-3D-v2.pt\",map_location='cpu')\n\n#get rid of module. from ddp state dict\nnew_state_dict={}\n\nfor key in state_dict:\n    new_state_dict[key[7:]]=state_dict[key]\n\nmodel.load_state_dict(new_state_dict)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T02:11:48.617279Z","iopub.execute_input":"2025-06-23T02:11:48.617463Z","iopub.status.idle":"2025-06-23T02:12:03.352908Z","shell.execute_reply.started":"2025-06-23T02:11:48.617450Z","shell.execute_reply":"2025-06-23T02:12:03.352150Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# inference loop","metadata":{}},{"cell_type":"code","source":"model.eval()\npreds=[]\nfor i in tqdm(range(len(test_dataset))):\n    src=test_dataset[i]['sequence'].long()\n    src=src.unsqueeze(0).cuda()\n    target_id=test_data.loc[i,'target_id']\n\n    predicted_dm=[]\n    #for _ in range(5):\n    with torch.no_grad():\n        #xyz,distogram=model.sample_euler(src,5,200,N_cycle=config.max_cycles)\n        with torch.cuda.amp.autocast():\n            #xyz,distogram=model.sample_euler(src,5,200,N_cycle=10)\n            #xyz,distogram=model.sample_euler(src,5,200,N_cycle=config.max_cycles)\n            #xyz,distogram=model.sample_euler(src,5,200,N_cycle=10)\n            xyz,distogram=model.sample_euler(src,5,200,N_cycle=1)\n    preds.append(xyz.cpu().numpy())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T02:12:03.353796Z","iopub.execute_input":"2025-06-23T02:12:03.354488Z","iopub.status.idle":"2025-06-23T02:14:58.214903Z","shell.execute_reply.started":"2025-06-23T02:12:03.354464Z","shell.execute_reply":"2025-06-23T02:14:58.214144Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Make submission.csv","metadata":{}},{"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\ncsv_filename=args.config.split('/')[-1].replace('.csv','')\nsubmission.to_csv(f'{csv_filename}_predictions.csv',index=False)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T02:14:58.215770Z","iopub.execute_input":"2025-06-23T02:14:58.216305Z","iopub.status.idle":"2025-06-23T02:14:58.327557Z","shell.execute_reply.started":"2025-06-23T02:14:58.216277Z","shell.execute_reply":"2025-06-23T02:14:58.327019Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import plotly.graph_objects as go\nimport numpy as np\n\n# Step 1: Load predicted and native coordinates\nindex = 6\nID = test_data.loc[index, 'target_id']\nxyz = preds[index][0]  # Nx3 predicted coordinates\n\n# Step 5: Plot\nfig = go.Figure()\n\n# Predicted structure (xyz)\nfig.add_trace(go.Scatter3d(\n    x=xyz[:, 0],\n    y=xyz[:, 1],\n    z=xyz[:, 2],\n    mode='markers',\n    marker=dict(size=5, color='blue', opacity=0.7),\n    name='Predicted'\n))\n\n\n# Layout\nfig.update_layout(\n    scene=dict(\n        xaxis_title=\"X\",\n        yaxis_title=\"Y\",\n        zaxis_title=\"Z\"\n    ),\n    title=f\"3D Alignment: {ID}\",\n    legend=dict(x=0.01, y=0.99)\n)\n\n# Show\nfig.show(renderer='iframe')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-23T02:17:37.901308Z","iopub.execute_input":"2025-06-23T02:17:37.901798Z","iopub.status.idle":"2025-06-23T02:17:37.947663Z","shell.execute_reply.started":"2025-06-23T02:17:37.901776Z","shell.execute_reply":"2025-06-23T02:17:37.947113Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}