{"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":"none","dataSources":[{"sourceId":87793,"databundleVersionId":12024591,"sourceType":"competition"},{"sourceId":224896926,"sourceType":"kernelVersion"}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Step 1: Imports","metadata":{"execution":{"iopub.status.busy":"2025-04-26T13:12:43.997588Z","iopub.execute_input":"2025-04-26T13:12:43.997890Z","iopub.status.idle":"2025-04-26T13:12:44.002367Z","shell.execute_reply.started":"2025-04-26T13:12:43.997869Z","shell.execute_reply":"2025-04-26T13:12:44.001391Z"}}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nfrom tqdm import tqdm\nimport plotly.graph_objects as go\nfrom plotly.subplots import make_subplots","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-26T13:15:38.431958Z","iopub.execute_input":"2025-04-26T13:15:38.432388Z","iopub.status.idle":"2025-04-26T13:15:40.418466Z","shell.execute_reply.started":"2025-04-26T13:15:38.432254Z","shell.execute_reply":"2025-04-26T13:15:40.417194Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 2: Definne configs","metadata":{}},{"cell_type":"code","source":"root_path                = '/kaggle/input/stanford-rna-3d-folding'\nexperiment_sub_file_path = \"/kaggle/input/ribonanzanet-3d-inference/submission.csv\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-26T13:15:40.419351Z","iopub.execute_input":"2025-04-26T13:15:40.419766Z","iopub.status.idle":"2025-04-26T13:15:40.428429Z","shell.execute_reply.started":"2025-04-26T13:15:40.419741Z","shell.execute_reply":"2025-04-26T13:15:40.427364Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 3 : Load preds and targets","metadata":{}},{"cell_type":"markdown","source":"### Step 3.1 : Load valid data for visualization","metadata":{}},{"cell_type":"code","source":"valid_sequences        = pd.read_csv(f\"{root_path}/validation_sequences.csv\")\nvalid_labels           = pd.read_csv(f\"{root_path}/validation_labels.csv\")\nvalid_labels[\"pdb_id\"] = valid_labels[\"ID\"].apply(lambda x: x.split(\"_\")[0])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-26T13:15:40.434097Z","iopub.execute_input":"2025-04-26T13:15:40.434485Z","iopub.status.idle":"2025-04-26T13:15:40.692460Z","shell.execute_reply.started":"2025-04-26T13:15:40.434461Z","shell.execute_reply":"2025-04-26T13:15:40.691452Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Step 2.2 : Load model predictions","metadata":{}},{"cell_type":"code","source":"preds_df           = pd.read_csv(experiment_sub_file_path)\npreds_df[\"pdb_id\"] = preds_df[\"ID\"].apply(lambda x: x.split(\"_\")[0])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-26T13:15:40.693281Z","iopub.execute_input":"2025-04-26T13:15:40.693518Z","iopub.status.idle":"2025-04-26T13:15:40.731747Z","shell.execute_reply.started":"2025-04-26T13:15:40.693498Z","shell.execute_reply":"2025-04-26T13:15:40.730972Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 3: Format data for visualization","metadata":{}},{"cell_type":"code","source":"all_xyz=[]\nall_preds = []\n\nfor pdb_id in tqdm(valid_sequences['target_id']):\n    df = valid_labels[valid_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\n    temp_arr = []\n    sub_preds_df = preds_df[preds_df['pdb_id']==pdb_id]\n    \n    xyz_preds1=sub_preds_df[['x_1','y_1','z_1']].to_numpy().astype('float32')\n    xyz_preds1[xyz_preds1<-1e17]=float('Nan');\n    temp_arr.append(xyz_preds1)\n\n    xyz_preds2=sub_preds_df[['x_2','y_2','z_2']].to_numpy().astype('float32')\n    xyz_preds2[xyz_preds2<-1e17]=float('Nan');\n    temp_arr.append(xyz_preds2)\n\n    xyz_preds3=sub_preds_df[['x_3','y_3','z_3']].to_numpy().astype('float32')\n    xyz_preds3[xyz_preds3<-1e17]=float('Nan');\n    temp_arr.append(xyz_preds3)\n\n    xyz_preds4=sub_preds_df[['x_4','y_4','z_4']].to_numpy().astype('float32')\n    xyz_preds4[xyz_preds4<-1e17]=float('Nan');\n    temp_arr.append(xyz_preds4)\n\n    xyz_preds5=sub_preds_df[['x_5','y_5','z_5']].to_numpy().astype('float32')\n    xyz_preds5[xyz_preds5<-1e17]=float('Nan');\n    temp_arr.append(xyz_preds5)\n    \n    all_preds.append(temp_arr)\n\nlen(all_xyz), len(all_preds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-26T13:15:40.732614Z","iopub.execute_input":"2025-04-26T13:15:40.732973Z","iopub.status.idle":"2025-04-26T13:15:40.804287Z","shell.execute_reply.started":"2025-04-26T13:15:40.732905Z","shell.execute_reply":"2025-04-26T13:15:40.803447Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#pack data into a dictionary\n\ndata={\n      \"sequence\":valid_sequences['sequence'].to_list(),\n      \"xyz\": all_xyz,\n      \"preds\" : all_preds\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-26T13:15:40.805482Z","iopub.execute_input":"2025-04-26T13:15:40.805823Z","iopub.status.idle":"2025-04-26T13:15:40.810513Z","shell.execute_reply.started":"2025-04-26T13:15:40.805789Z","shell.execute_reply":"2025-04-26T13:15:40.809740Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 4: Plot the graph","metadata":{}},{"cell_type":"code","source":"# Access the target array\ntarget = data['xyz'][2]\npredictions  = data['preds'][2]\n\n# Number of predictions\nnum_preds = len(predictions)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-26T13:15:40.811482Z","iopub.execute_input":"2025-04-26T13:15:40.811739Z","iopub.status.idle":"2025-04-26T13:15:40.826441Z","shell.execute_reply.started":"2025-04-26T13:15:40.811718Z","shell.execute_reply":"2025-04-26T13:15:40.825573Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Step 4.1: targets","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport plotly.graph_objects as go\nfrom plotly.subplots import make_subplots\n\n# Code to plot only the target structure\nif isinstance(target, np.ndarray) and target.ndim == 2 and target.shape[1] == 3:\n    x_t, y_t, z_t = target[:, 0], target[:, 1], target[:, 2]\n    N_t = len(target)\n    sequence_indices_t = np.arange(N_t)\n\n    # Create a figure with a single subplot for the target\n    fig = make_subplots(\n        rows=1,\n        cols=1, # Only one column for the target plot\n        specs=[[{'type': 'scatter3d'}]], # Specify 3D plot type\n        subplot_titles=['Target'] # Title for the single subplot\n    )\n\n    # Add the 3D scatter plot trace for the target structure\n    fig.add_trace(go.Scatter3d(\n        x=x_t, y=y_t, z=z_t,\n        mode='lines+markers', # Show backbone line and markers\n        line=dict(color='red', width=2), # Style the backbone line (e.g., red)\n        marker=dict(\n            size=3, # Size of the markers\n            color=sequence_indices_t, # Color markers sequentially\n            colorscale='Viridis',   # Colorscale for markers\n            opacity=0.8,\n            colorbar=dict(title='Residue Index') # Add the colorbar\n        ),\n        text=[f'Residue {j}' for j in sequence_indices_t], # Hover text\n        hoverinfo='text+x+y+z' # Information to display on hover\n    ), row=1, col=1) # Assign trace to the first (and only) subplot cell\n\n    # Update the layout for the single plot\n    fig.update_layout(\n        height=400,  # Adjust height as needed\n        width=600, # Adjust width for a single plot\n        title_text='Target RNA Structure', # Title for the figure\n        showlegend=False, # Hide legend as colorbar is sufficient\n        margin=dict(l=10, r=10, b=10, t=50) # Adjust margins\n    )\n\n    # Display the figure\n    fig.show(renderer='iframe') # Use 'iframe' or your preferred renderer\n\nelse:\n    # Print a warning if the target data is not in the expected format\n    print(f\"Warning: Target data is not a valid Nx3 numpy array. Cannot plot.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-26T13:15:40.827410Z","iopub.execute_input":"2025-04-26T13:15:40.827854Z","iopub.status.idle":"2025-04-26T13:15:41.743414Z","shell.execute_reply.started":"2025-04-26T13:15:40.827821Z","shell.execute_reply":"2025-04-26T13:15:41.742513Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Step 4.2: Predictions","metadata":{}},{"cell_type":"code","source":"if num_preds != 5:\n    print(f\"Error: Expected 5 predictions in preds[2], but found {num_preds}.\")\nelse:\n    fig = make_subplots(\n        rows=1,\n        cols=num_preds, # Add one column for the target\n        specs=[[{'type': 'scatter3d'}] * (num_preds)], # Specify 3D plot type for each subplot\n        subplot_titles=[f'Prediction {i}' for i in range(num_preds)] # Titles including 'Target'\n    )\n\n    # Iterate through each prediction (xyz coordinate array) and add it to the figure\n    for i, xyz in enumerate(predictions):\n        # Ensure xyz is a numpy array\n        if not isinstance(xyz, np.ndarray) or xyz.ndim != 2 or xyz.shape[1] != 3:\n             print(f\"Warning: Prediction {i+1} is not a valid Nx3 numpy array. Skipping.\")\n             continue\n\n        x, y, z = xyz[:, 0], xyz[:, 1], xyz[:, 2]\n        N = len(xyz)\n        sequence_indices = np.arange(N)\n\n        # Add the 3D scatter plot trace for the current prediction\n        fig.add_trace(go.Scatter3d(\n            x=x, y=y, z=z,\n            mode='lines+markers', # Show backbone line and markers for each residue\n            line=dict(color='grey', width=2), # Style the backbone line\n            marker=dict(\n                size=3, # Size of the markers\n                color=sequence_indices, # Color markers sequentially\n                colorscale='Viridis',   # Colorscale for markers\n                opacity=0.8,\n                # Remove conditional colorbar here, will add one to the target plot\n                # colorbar=dict(title='Residue Index') if i == num_preds - 1 else None\n            ),\n            text=[f'Residue {j}' for j in sequence_indices], # Hover text\n            hoverinfo='text+x+y+z' # Information to display on hover\n        ), row=1, col=i + 1) # Assign trace to the correct subplot cell (cols 1 to 5)\n\n    fig.update_layout(\n        height=400,  # Adjust height as needed\n        width=1800, # Increase width to accommodate 6 plots (approx 300px per plot)\n        title_text='Comparison of 5 RNA Structure Predictions and Target Structure', # Updated title\n        showlegend=False, # Hide legend as colors are self-explanatory with colorbar\n        margin=dict(l=10, r=10, b=10, t=50) # Adjust margins\n    )\n\n    fig.show(renderer='iframe') # Use 'iframe' or your preferred renderer","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-26T13:15:41.745567Z","iopub.execute_input":"2025-04-26T13:15:41.745849Z","iopub.status.idle":"2025-04-26T13:15:41.849078Z","shell.execute_reply.started":"2025-04-26T13:15:41.745825Z","shell.execute_reply":"2025-04-26T13:15:41.848184Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}