{"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,"sourceType":"competition"},{"sourceId":7395079,"sourceType":"datasetVersion","datasetId":4299455},{"sourceId":7639698,"sourceType":"datasetVersion","datasetId":4299272},{"sourceId":224703571,"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 random\nimport pickle\nimport yaml\nimport sys\nfrom tqdm import tqdm  # Added for progress tracking\nfrom torch.utils.data import Dataset, DataLoader","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T17:15:49.668936Z","iopub.execute_input":"2025-05-16T17:15:49.669480Z","iopub.status.idle":"2025-05-16T17:15:49.673841Z","shell.execute_reply.started":"2025-05-16T17:15:49.669454Z","shell.execute_reply":"2025-05-16T17:15:49.673001Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Set random seed for reproducibility","metadata":{}},{"cell_type":"code","source":"def set_seed(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\nset_seed(0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T17:15:49.675154Z","iopub.execute_input":"2025-05-16T17:15:49.675390Z","iopub.status.idle":"2025-05-16T17:15:49.693421Z","shell.execute_reply.started":"2025-05-16T17:15:49.675368Z","shell.execute_reply":"2025-05-16T17:15:49.692762Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Configuration","metadata":{}},{"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\",\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-05-16T17:15:49.694150Z","iopub.execute_input":"2025-05-16T17:15:49.694845Z","iopub.status.idle":"2025-05-16T17:15:49.711865Z","shell.execute_reply.started":"2025-05-16T17:15:49.694823Z","shell.execute_reply":"2025-05-16T17:15:49.711204Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Load test data","metadata":{}},{"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-05-16T17:15:49.713231Z","iopub.execute_input":"2025-05-16T17:15:49.713439Z","iopub.status.idle":"2025-05-16T17:15:49.732111Z","shell.execute_reply.started":"2025-05-16T17:15:49.713425Z","shell.execute_reply":"2025-05-16T17:15:49.731566Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset class","metadata":{}},{"cell_type":"code","source":"class RNADataset(Dataset):\n    def __init__(self, data):\n        self.tokens = {'A': 0, 'C': 1, 'G': 2, 'U': 3}\n        self.sequences = [\n            torch.tensor([self.tokens.get(nt, -1) for nt in seq], dtype=torch.long)\n            for seq in data['sequence']\n        ]\n\n    def __len__(self):\n        return len(self.sequences)\n\n    def __getitem__(self, idx):\n        return {'sequence': self.sequences[idx]}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T17:15:49.732891Z","iopub.execute_input":"2025-05-16T17:15:49.733105Z","iopub.status.idle":"2025-05-16T17:15:49.744752Z","shell.execute_reply.started":"2025-05-16T17:15:49.733082Z","shell.execute_reply":"2025-05-16T17:15:49.744107Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Initialize dataset","metadata":{}},{"cell_type":"code","source":"test_dataset = RNADataset(test_data)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T17:15:49.745353Z","iopub.execute_input":"2025-05-16T17:15:49.745511Z","iopub.status.idle":"2025-05-16T17:15:49.760370Z","shell.execute_reply.started":"2025-05-16T17:15:49.745499Z","shell.execute_reply":"2025-05-16T17:15:49.759891Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Verify dataset","metadata":{}},{"cell_type":"code","source":"if len(test_dataset) == 0:\n    raise ValueError(\"Test dataset is empty. Check input data.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T17:15:49.760970Z","iopub.execute_input":"2025-05-16T17:15:49.761176Z","iopub.status.idle":"2025-05-16T17:15:49.774926Z","shell.execute_reply.started":"2025-05-16T17:15:49.761154Z","shell.execute_reply":"2025-05-16T17:15:49.774225Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Load model configuration","metadata":{}},{"cell_type":"code","source":"sys.path.append(\"/kaggle/input/ribonanzanet2d-final\")\nfrom Network import *\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)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T17:15:49.775635Z","iopub.execute_input":"2025-05-16T17:15:49.775856Z","iopub.status.idle":"2025-05-16T17:15:52.377082Z","shell.execute_reply.started":"2025-05-16T17:15:49.775833Z","shell.execute_reply":"2025-05-16T17:15:52.376497Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Define fine-tuned model","metadata":{}},{"cell_type":"code","source":"class 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(\n                \"/kaggle/input/ribonanzanet-weights/RibonanzaNet.pt\",\n                map_location='cpu',\n                weights_only=True  # Added for security\n            ))\n        self.dropout = nn.Dropout(0.0)\n        self.xyz_predictor = nn.Linear(256, 3)\n\n    def forward(self, src):\n        sequence_features, pairwise_features = self.get_embeddings(\n            src, torch.ones_like(src).long().to(src.device)\n        )\n        xyz = self.xyz_predictor(sequence_features)\n        return xyz","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T17:15:52.377798Z","iopub.execute_input":"2025-05-16T17:15:52.378574Z","iopub.status.idle":"2025-05-16T17:15:52.383158Z","shell.execute_reply.started":"2025-05-16T17:15:52.378549Z","shell.execute_reply":"2025-05-16T17:15:52.382446Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Initialize device and model","metadata":{}},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = finetuned_RibonanzaNet(\n    load_config_from_yaml(\"/kaggle/input/ribonanzanet2d-final/configs/pairwise.yaml\"),\n    pretrained=False\n).to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T17:15:52.385842Z","iopub.execute_input":"2025-05-16T17:15:52.386564Z","iopub.status.idle":"2025-05-16T17:15:52.731445Z","shell.execute_reply.started":"2025-05-16T17:15:52.386537Z","shell.execute_reply":"2025-05-16T17:15:52.730917Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Load model weights","metadata":{}},{"cell_type":"code","source":"model.load_state_dict(torch.load(\n    \"/kaggle/input/ribonanzanet-3d-finetune/RibonanzaNet-3D.pt\",\n    weights_only=True  # Added for security\n))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T17:15:52.732073Z","iopub.execute_input":"2025-05-16T17:15:52.732271Z","iopub.status.idle":"2025-05-16T17:15:53.824064Z","shell.execute_reply.started":"2025-05-16T17:15:52.732256Z","shell.execute_reply":"2025-05-16T17:15:53.823470Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Set model to evaluation mode","metadata":{}},{"cell_type":"code","source":"model.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T17:15:53.824701Z","iopub.execute_input":"2025-05-16T17:15:53.824921Z","iopub.status.idle":"2025-05-16T17:15:53.832877Z","shell.execute_reply.started":"2025-05-16T17:15:53.824906Z","shell.execute_reply":"2025-05-16T17:15:53.832225Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Generate predictions","metadata":{}},{"cell_type":"code","source":"preds = []\nnum_ensembles = 5  # Number of ensemble predictions\n\nfor i in tqdm(range(len(test_dataset)), desc=\"Predicting\"):\n    src = test_dataset[i]['sequence'].long().unsqueeze(0).to(device)\n    tmp = []\n    \n    # Generate ensemble predictions\n    with torch.no_grad():\n        for _ in range(num_ensembles):\n            xyz = model(src).squeeze().cpu().numpy()\n            tmp.append(xyz)\n    \n    tmp = np.stack(tmp, axis=0)\n    preds.append(tmp)\n\n# Verify predictions\nif not preds:\n    raise ValueError(\"No predictions generated. Check model or dataset.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T17:15:53.833599Z","iopub.execute_input":"2025-05-16T17:15:53.833846Z","iopub.status.idle":"2025-05-16T17:16:06.602380Z","shell.execute_reply.started":"2025-05-16T17:15:53.833830Z","shell.execute_reply":"2025-05-16T17:16:06.601785Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Create submission DataFrame","metadata":{}},{"cell_type":"code","source":"data = []\nfor i in range(len(test_data)):\n    seq_id = test_data.loc[i, 'target_id']\n    sequence = test_data.loc[i, 'sequence']\n    \n    for j in range(len(sequence)):\n        row = [\n            f\"{seq_id}_{j+1}\",  # ID\n            sequence[j],        # resname\n            j + 1              # resid\n        ]\n        # Add x, y, z for each ensemble\n        for k in range(num_ensembles):\n            row.extend(preds[i][k][j])  # x, y, z\n        data.append(row)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T17:16:06.603081Z","iopub.execute_input":"2025-05-16T17:16:06.603333Z","iopub.status.idle":"2025-05-16T17:16:06.626596Z","shell.execute_reply.started":"2025-05-16T17:16:06.603318Z","shell.execute_reply":"2025-05-16T17:16:06.625787Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Define column names","metadata":{}},{"cell_type":"code","source":"columns = ['ID', 'resname', 'resid']\nfor i in range(1, num_ensembles + 1):\n    columns += [f\"x_{i}\", f\"y_{i}\", f\"z_{i}\"]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T17:16:06.627563Z","iopub.execute_input":"2025-05-16T17:16:06.627847Z","iopub.status.idle":"2025-05-16T17:16:06.647789Z","shell.execute_reply.started":"2025-05-16T17:16:06.627823Z","shell.execute_reply":"2025-05-16T17:16:06.647078Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Create and save submission","metadata":{}},{"cell_type":"code","source":"submission = pd.DataFrame(data, columns=columns)\nsubmission.to_csv('submission.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T17:16:06.648592Z","iopub.execute_input":"2025-05-16T17:16:06.648848Z","iopub.status.idle":"2025-05-16T17:16:06.717226Z","shell.execute_reply.started":"2025-05-16T17:16:06.648827Z","shell.execute_reply":"2025-05-16T17:16:06.716516Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Display submission","metadata":{}},{"cell_type":"code","source":"print(submission.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T17:16:06.717938Z","iopub.execute_input":"2025-05-16T17:16:06.718136Z","iopub.status.idle":"2025-05-16T17:16:06.730113Z","shell.execute_reply.started":"2025-05-16T17:16:06.718112Z","shell.execute_reply":"2025-05-16T17:16:06.729520Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Optional: Visualize one prediction (unchanged)","metadata":{}},{"cell_type":"code","source":"import plotly.graph_objects as go\n\nxyz = preds[7][0]  # Example: First ensemble of 8th prediction\nx, y, z = xyz[:, 0], xyz[:, 1], xyz[:, 2]\n\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,\n        colorscale='Viridis',\n        opacity=0.8\n    )\n)])\n\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\nfig.show(renderer='iframe')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-16T17:16:06.730713Z","iopub.execute_input":"2025-05-16T17:16:06.731054Z","iopub.status.idle":"2025-05-16T17:16:07.255658Z","shell.execute_reply.started":"2025-05-16T17:16:06.731034Z","shell.execute_reply":"2025-05-16T17:16:07.255076Z"}},"outputs":[],"execution_count":null}]}