{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":118765,"databundleVersionId":15231210,"sourceType":"competition"}],"dockerImageVersionId":31239,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-10T09:21:36.633796Z","iopub.execute_input":"2026-01-10T09:21:36.634152Z","iopub.status.idle":"2026-01-10T09:21:46.765228Z","shell.execute_reply.started":"2026-01-10T09:21:36.634121Z","shell.execute_reply":"2026-01-10T09:21:46.764151Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, math, random\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T09:21:46.766136Z","iopub.execute_input":"2026-01-10T09:21:46.76639Z","iopub.status.idle":"2026-01-10T09:21:46.771738Z","shell.execute_reply.started":"2026-01-10T09:21:46.76637Z","shell.execute_reply":"2026-01-10T09:21:46.77071Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MODE = \"infer\"   # \"train\" or \"infer\"\nMODEL_PATH = \"rna_foldnet.pt\"\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T09:21:46.772739Z","iopub.execute_input":"2026-01-10T09:21:46.773097Z","iopub.status.idle":"2026-01-10T09:21:46.789168Z","shell.execute_reply.started":"2026-01-10T09:21:46.773066Z","shell.execute_reply":"2026-01-10T09:21:46.788215Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nset_seed()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T09:21:46.790156Z","iopub.execute_input":"2026-01-10T09:21:46.790497Z","iopub.status.idle":"2026-01-10T09:21:46.8087Z","shell.execute_reply.started":"2026-01-10T09:21:46.790469Z","shell.execute_reply":"2026-01-10T09:21:46.807451Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BASE_TO_IDX = {'A':0,'U':1,'G':2,'C':3}\nIDX_TO_BASE = {v:k for k,v in BASE_TO_IDX.items()}\n\ndef encode_sequence(seq):\n    return torch.tensor([BASE_TO_IDX[b] for b in seq], dtype=torch.long)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T09:21:46.810079Z","iopub.execute_input":"2026-01-10T09:21:46.810421Z","iopub.status.idle":"2026-01-10T09:21:46.824465Z","shell.execute_reply.started":"2026-01-10T09:21:46.810391Z","shell.execute_reply":"2026-01-10T09:21:46.823342Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def dotbracket_to_pairs(db):\n    stack, pairs = [], {}\n    for i,c in enumerate(db):\n        if c == '(':\n            stack.append(i)\n        elif c == ')':\n            j = stack.pop()\n            pairs[i] = j\n            pairs[j] = i\n    return pairs\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T09:21:46.825546Z","iopub.execute_input":"2026-01-10T09:21:46.826316Z","iopub.status.idle":"2026-01-10T09:21:46.8422Z","shell.execute_reply.started":"2026-01-10T09:21:46.826282Z","shell.execute_reply":"2026-01-10T09:21:46.841095Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RNAFoldNet(nn.Module):\n    def __init__(self, d=128, layers=6, heads=8):\n        super().__init__()\n        self.embed = nn.Embedding(4, d)\n\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d,\n            nhead=heads,\n            dim_feedforward=256,\n            dropout=0.0,\n            batch_first=True\n        )\n        self.transformer = nn.TransformerEncoder(\n            encoder_layer, num_layers=layers\n        )\n\n        self.coord_head = nn.Linear(d, 3)\n\n    def forward(self, seq):\n        x = self.embed(seq)\n        x = self.transformer(x)\n        return self.coord_head(x)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T09:21:46.843221Z","iopub.execute_input":"2026-01-10T09:21:46.843527Z","iopub.status.idle":"2026-01-10T09:21:46.861352Z","shell.execute_reply.started":"2026-01-10T09:21:46.843501Z","shell.execute_reply":"2026-01-10T09:21:46.86007Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def pairwise_distance_loss(pred, gt):\n    return torch.mean(torch.abs(\n        torch.cdist(pred, pred) - torch.cdist(gt, gt)\n    ))\n\ndef bond_loss(coords, ideal=3.8):\n    return torch.mean((torch.norm(coords[1:] - coords[:-1], dim=-1) - ideal)**2)\n\ndef clash_loss(coords, min_dist=2.0):\n    d = torch.cdist(coords, coords)\n    mask = (d > 0) & (d < min_dist)\n    return torch.sum((min_dist - d[mask])**2)\n\ndef total_loss(pred, gt):\n    return (\n        nn.functional.mse_loss(pred, gt)\n        + 0.5 * pairwise_distance_loss(pred, gt)\n        + 0.1 * bond_loss(pred)\n        + 0.01 * clash_loss(pred)\n    )\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T09:21:46.862575Z","iopub.execute_input":"2026-01-10T09:21:46.862837Z","iopub.status.idle":"2026-01-10T09:21:46.883292Z","shell.execute_reply.started":"2026-01-10T09:21:46.862818Z","shell.execute_reply":"2026-01-10T09:21:46.882249Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RNATrainDataset(Dataset):\n    def __init__(self, df):\n        self.df = df\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        return (\n            encode_sequence(row.sequence),\n            torch.tensor(row.coords, dtype=torch.float32)\n        )\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T09:21:46.884459Z","iopub.execute_input":"2026-01-10T09:21:46.884753Z","iopub.status.idle":"2026-01-10T09:21:46.900899Z","shell.execute_reply.started":"2026-01-10T09:21:46.884731Z","shell.execute_reply":"2026-01-10T09:21:46.899985Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if MODE == \"train\":\n    train_df = pd.read_pickle(\"/kaggle/input/traindata/train_sequences.csv\")\n    dataset = RNATrainDataset(train_df)\n    loader = DataLoader(dataset, batch_size=1, shuffle=True)\n\n    model = RNAFoldNet().to(DEVICE)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)\n\n    for epoch in range(10):\n        model.train()\n        for seq, gt in loader:\n            seq, gt = seq.to(DEVICE), gt.to(DEVICE)\n\n            pred = model(seq)\n            loss = total_loss(pred[0], gt[0])\n\n            optimizer.zero_grad()\n            loss.backward()\n            optimizer.step()\n\n        print(f\"Epoch {epoch+1} | Loss {loss.item():.4f}\")\n\n    torch.save(model.state_dict(), MODEL_PATH)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T09:21:46.90189Z","iopub.execute_input":"2026-01-10T09:21:46.902232Z","iopub.status.idle":"2026-01-10T09:21:46.919296Z","shell.execute_reply.started":"2026-01-10T09:21:46.902204Z","shell.execute_reply":"2026-01-10T09:21:46.918221Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if MODE == \"infer\":\n    test_df = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv\")\n\n    model = RNAFoldNet(layers=4).to(DEVICE)  # fewer layers = faster\n\nimport os\nfrom torch.cuda.amp import autocast\n\nrows = []\n\nmodel.eval()\n\nwith torch.no_grad():\n    for idx, row in test_df.iterrows():\n        seq_id = f\"RNA_{idx+1}\"\n        seq = row[\"sequence\"]\n\n        x = encode_sequence(seq).unsqueeze(0).to(DEVICE)\n\n        with autocast():\n            coords = model(x)[0].cpu()\n\n        preds = []\n        for i in range(5):\n            preds.append(coords + torch.randn_like(coords) * (0.04 + i * 0.01))\n\n        for i, base in enumerate(seq):\n            out = {\n                \"ID\": f\"{seq_id}_{i+1}\",\n                \"resname\": base,\n                \"resid\": i + 1\n            }\n            for p in range(5):\n                out[f\"x_{p+1}\"] = preds[p][i, 0].item()\n                out[f\"y_{p+1}\"] = preds[p][i, 1].item()\n                out[f\"z_{p+1}\"] = preds[p][i, 2].item()\n            rows.append(out)\n\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T09:21:46.920344Z","iopub.execute_input":"2026-01-10T09:21:46.920682Z","iopub.status.idle":"2026-01-10T09:21:50.657112Z","shell.execute_reply.started":"2026-01-10T09:21:46.920642Z","shell.execute_reply":"2026-01-10T09:21:50.656152Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Total rows written:\", len(rows))\nsubmission = pd.DataFrame(rows)\nsubmission.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T09:21:50.657982Z","iopub.execute_input":"2026-01-10T09:21:50.658231Z","iopub.status.idle":"2026-01-10T09:21:50.709893Z","shell.execute_reply.started":"2026-01-10T09:21:50.658211Z","shell.execute_reply":"2026-01-10T09:21:50.709003Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}