{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.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":118765,"databundleVersionId":15231210,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":7639698,"sourceType":"datasetVersion","datasetId":4299272},{"sourceId":8318191,"sourceType":"datasetVersion","datasetId":4459124},{"sourceId":14452625,"sourceType":"datasetVersion","datasetId":9231246}],"dockerImageVersionId":31234,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"<h1 id=\"setup-and-configuration\"\n    style=\"\n  background: linear-gradient(135deg, #0f2027, #203a43, #2c5364);\n  color: white;\n  padding: 12px 24px;\n  text-align: center;\n  border-radius: 12px;\n  font-size: 22px;\n  font-weight: bold;\n  font-family: Arial, sans-serif;\n  margin: 20px 0;\n\">\n  1. Imports and Configurations\n</h1>","metadata":{}},{"cell_type":"code","source":"import sys\nimport random\nimport pickle\nimport yaml\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport matplotlib.pyplot as plt\n\nfrom torch.utils.data import Dataset, DataLoader\nimport plotly.graph_objects as go\n\nclass Config:\n    sample_sub = \"/kaggle/input/stanford-rna-3d-folding-2/sample_submission.csv\"\n    test_seq = \"/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv\"\n    train_labels = \"/kaggle/input/stanford-rna-3d-folding-2/train_labels.csv\"\n    train_sequences = \"/kaggle/input/stanford-rna-3d-folding-2/train_sequences.csv\"\n    validation_labels = \"/kaggle/input/stanford-rna-3d-folding-2/validation_labels.csv\"\n    validation_sequences = \"/kaggle/input/stanford-rna-3d-folding-2/validation_sequences.csv\"\n    model_config_path = \"/kaggle/input/ribonanzanet2d-final/configs/pairwise.yaml\"\n    pretrained_weights_path = \"/kaggle/input/ribonanzanet-weights/RibonanzaNet.pt\"\n    finetuned_weights_path = \"/kaggle/input/ribonanzanet-finetuned/RibonanzaNet-3D.pt\"\n    \n    max_len = 384\n    batch_size = 1\n    seed = 42\n    num_tta = 5  \n\nconfig = Config()\n\n# Set seed for reproducibility\ndef set_seed(seed: int):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nset_seed(config.seed)\n\n#  Load Test Data\ntest_data = pd.read_csv(config.test_seq)\nprint(f\"Test sequences: {len(test_data)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T08:13:45.820097Z","iopub.execute_input":"2026-01-10T08:13:45.820697Z","iopub.status.idle":"2026-01-10T08:13:50.732324Z","shell.execute_reply.started":"2026-01-10T08:13:45.820663Z","shell.execute_reply":"2026-01-10T08:13:50.731504Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<h1 id=\"loader-set\"\n    style=\"\n  background: linear-gradient(135deg, #0f2027, #203a43, #2c5364);\n  color: white;\n  padding: 12px 24px;\n  text-align: center;\n  border-radius: 12px;\n  font-size: 22px;\n  font-weight: bold;\n  font-family: Arial, sans-serif;\n  margin: 20px 0;\n\">\n  2. Test Dataset\n</h1>","metadata":{}},{"cell_type":"code","source":"class RNADataset(Dataset):\n    def __init__(self, data, max_len=384):\n        self.data = data\n        self.max_len = max_len\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 = torch.tensor(np.array(sequence), dtype=torch.long)\n        \n        original_length = len(sequence)\n        \n        return {\n            \"sequence\": sequence,\n            \"original_length\": original_length,\n            \"target_id\": self.data.loc[idx, \"target_id\"]\n        }\n\n\ntest_dataset = RNADataset(test_data, max_len=config.max_len)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T08:13:50.733874Z","iopub.execute_input":"2026-01-10T08:13:50.734493Z","iopub.status.idle":"2026-01-10T08:13:50.741642Z","shell.execute_reply.started":"2026-01-10T08:13:50.734455Z","shell.execute_reply":"2026-01-10T08:13:50.740771Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<h1 id=\"model-diff\"\n    style=\"\n  background: linear-gradient(135deg, #0f2027, #203a43, #2c5364);\n  color: white;\n  padding: 12px 24px;\n  text-align: center;\n  border-radius: 12px;\n  font-size: 22px;\n  font-weight: bold;\n  font-family: Arial, sans-serif;\n  margin: 20px 0;\n\">\n  3.  Model Definition\n</h1>","metadata":{}},{"cell_type":"code","source":"sys.path.append(\"/kaggle/input/ribonanzanet2d-final\")\nfrom Network import RibonanzaNet\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\n\ndef load_config_from_yaml(file_path):\n    with open(file_path, \"r\") as file:\n        cfg = yaml.safe_load(file)\n    return Config(**cfg)\n\n\nclass FinetunedRibonanzaNet(RibonanzaNet):\n    def __init__(self, config_obj, pretrained=False, dropout=0.1):\n        config_obj.dropout = dropout\n        super(FinetunedRibonanzaNet, self).__init__(config_obj)\n\n        if pretrained:\n            self.load_state_dict(\n                torch.load(config.pretrained_weights_path, map_location=\"cpu\")\n            )\n\n        self.dropout = nn.Dropout(p=0.0)\n        self.xyz_predictor = nn.Linear(256, 3)\n\n    def forward(self, src):\n        sequence_features, _ = self.get_embeddings(\n            src, torch.ones_like(src).long().to(src.device)\n        )\n        xyz_pred = self.xyz_predictor(sequence_features)\n        return xyz_pred\n\n\n#  Load Model\nmodel_cfg = load_config_from_yaml(config.model_config_path)\nmodel = FinetunedRibonanzaNet(model_cfg, pretrained=False, dropout=0.2).cuda()\n\n# Load finetuned weights\nmodel.load_state_dict(\n    torch.load(config.finetuned_weights_path, map_location=\"cuda\")\n)\nprint(\"Model loaded successfully!\")\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T08:13:50.742503Z","iopub.execute_input":"2026-01-10T08:13:50.742826Z","iopub.status.idle":"2026-01-10T08:13:55.202862Z","shell.execute_reply.started":"2026-01-10T08:13:50.742795Z","shell.execute_reply":"2026-01-10T08:13:55.202218Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<h1 id=\"inference\"\n    style=\"\n  background: linear-gradient(135deg, #0f2027, #203a43, #2c5364);\n  color: white;\n  padding: 12px 24px;\n  text-align: center;\n  border-radius: 12px;\n  font-size: 22px;\n  font-weight: bold;\n  font-family: Arial, sans-serif;\n  margin: 20px 0;\n\">\n  4.  Inference\n</h1>","metadata":{}},{"cell_type":"code","source":"def predict_sequence(model, sequence, max_len=384):\n    \"\"\"\n    Predict XYZ coordinates for a sequence.\n    Handles sequences longer than max_len by sliding window approach.\n    \"\"\"\n    device = next(model.parameters()).device\n    seq_len = len(sequence)\n    \n    if seq_len <= max_len:\n        # Sequence fits in one pass\n        src = sequence.unsqueeze(0).to(device)\n        \n        model.eval()\n        with torch.no_grad():\n            xyz = model(src).squeeze(0)\n        \n        return xyz.cpu().numpy()\n    \n    else:\n        # Sequence is longer than max_len - use sliding window\n        step_size = max_len // 2  # 50% overlap\n        predictions_sum = np.zeros((seq_len, 3))\n        counts = np.zeros(seq_len)\n        \n        for start in range(0, seq_len - max_len + 1, step_size):\n            end = start + max_len\n            window = sequence[start:end].unsqueeze(0).to(device)\n            \n            model.eval()\n            with torch.no_grad():\n                xyz = model(window).squeeze(0)\n            \n            window_pred = xyz.cpu().numpy()\n            predictions_sum[start:end] += window_pred\n            counts[start:end] += 1\n        \n        # Handle the last segment if it doesn't align\n        if (seq_len - max_len) % step_size != 0:\n            start = seq_len - max_len\n            window = sequence[start:].unsqueeze(0).to(device)\n            \n            model.eval()\n            with torch.no_grad():\n                xyz = model(window).squeeze(0)\n            \n            window_pred = xyz.cpu().numpy()\n            predictions_sum[start:] += window_pred\n            counts[start:] += 1\n        \n        # Average overlapping predictions\n        final_pred = predictions_sum / counts[:, np.newaxis]\n        return final_pred\n\n\n\n# Run Inference\n\nprint(\"Running inference...\")\nall_predictions = []\n\nfor i in range(len(test_dataset)):\n    sample = test_dataset[i]\n    sequence = sample[\"sequence\"]\n    \n    # Predict\n    pred = predict_sequence(model, sequence, max_len=config.max_len)\n    all_predictions.append(pred)\n    \n    if (i + 1) % 10 == 0:\n        print(f\"Processed {i + 1}/{len(test_dataset)} sequences\")\n\nprint(\"Inference complete!\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T08:13:55.204238Z","iopub.execute_input":"2026-01-10T08:13:55.204575Z","iopub.status.idle":"2026-01-10T08:14:09.621673Z","shell.execute_reply.started":"2026-01-10T08:13:55.204551Z","shell.execute_reply":"2026-01-10T08:14:09.621049Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<h1 id=\"submission\"\n    style=\"\n  background: linear-gradient(135deg, #0f2027, #203a43, #2c5364);\n  color: white;\n  padding: 12px 24px;\n  text-align: center;\n  border-radius: 12px;\n  font-size: 22px;\n  font-weight: bold;\n  font-family: Arial, sans-serif;\n  margin: 20px 0;\n\">\n  5.  Submission\n</h1>","metadata":{}},{"cell_type":"code","source":"data = []\n\nfor i in range(len(test_data)):\n    target_id = test_data.loc[i, \"target_id\"]\n    sequence = test_data.loc[i, \"sequence\"]\n    seq_length = len(sequence)\n    \n    # Get predictions for this sequence\n    preds = all_predictions[i]  \n    \n    for j in range(seq_length):\n        row = [\n            f\"{target_id}_{j+1}\",\n            sequence[j],\n            j + 1,\n        ]\n        \n        # Add the same prediction 5 times \n        for k in range(5):\n            row.extend([preds[j, 0], preds[j, 1], preds[j, 2]])\n        \n        data.append(row)\n\n# Create submission DataFrame\ncolumns = [\"ID\", \"resname\", \"resid\"]\nfor i in range(1, 6):\n    columns += [f\"x_{i}\", f\"y_{i}\", f\"z_{i}\"]\n\nsubmission = pd.DataFrame(data, columns=columns)\n\n# Verify submission format\nprint(f\"\\nSubmission shape: {submission.shape}\")\nprint(submission.head())\n\n# Save submission\nsubmission.to_csv(\"submission.csv\", index=False)\nprint(\"\\nSubmission saved to submission.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T08:14:09.622473Z","iopub.execute_input":"2026-01-10T08:14:09.622676Z","iopub.status.idle":"2026-01-10T08:14:09.934573Z","shell.execute_reply.started":"2026-01-10T08:14:09.622656Z","shell.execute_reply":"2026-01-10T08:14:09.933841Z"}},"outputs":[],"execution_count":null}]}