{"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,"sourceType":"competition"}],"dockerImageVersionId":31259,"isInternetEnabled":true,"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-02-08T10:03:51.709278Z","iopub.execute_input":"2026-02-08T10:03:51.709533Z","iopub.status.idle":"2026-02-08T10:04:12.269943Z","shell.execute_reply.started":"2026-02-08T10:03:51.709499Z","shell.execute_reply":"2026-02-08T10:04:12.269102Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport math\nimport numpy as np\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(\"Using Device: \", DEVICE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T10:04:12.271381Z","iopub.execute_input":"2026-02-08T10:04:12.271739Z","iopub.status.idle":"2026-02-08T10:04:17.85603Z","shell.execute_reply.started":"2026-02-08T10:04:12.271715Z","shell.execute_reply":"2026-02-08T10:04:17.855268Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_DIR = \"/kaggle/input/stanford-rna-3d-folding-2\"\nMSA_DIR = F\"{DATA_DIR}/MSA\"\n\nNUC_MAP = {\n    \"A\": 0,\n    \"C\": 1,\n    \"G\": 2,\n    \"U\": 3,\n    \"-\": 4,\n    \"N\": 4\n}\n\ntrain_seq = pd.read_csv(f\"{DATA_DIR}/train_sequences.csv\")\nval_seq = pd.read_csv(f\"{DATA_DIR}/validation_sequences.csv\")\ntest_seq = pd.read_csv(f\"{DATA_DIR}/test_sequences.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T10:04:17.857042Z","iopub.execute_input":"2026-02-08T10:04:17.857468Z","iopub.status.idle":"2026-02-08T10:04:18.425369Z","shell.execute_reply.started":"2026-02-08T10:04:17.857429Z","shell.execute_reply":"2026-02-08T10:04:18.424517Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_seq[\"target_sequence\"] = train_seq[\"sequence\"]\nval_seq[\"target_sequence\"] = val_seq[\"sequence\"]\ntest_seq[\"target_sequence\"] = test_seq[\"sequence\"]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T10:04:18.42648Z","iopub.execute_input":"2026-02-08T10:04:18.426777Z","iopub.status.idle":"2026-02-08T10:04:18.434915Z","shell.execute_reply.started":"2026-02-08T10:04:18.42675Z","shell.execute_reply":"2026-02-08T10:04:18.434243Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Load labels\n\ntrain_labels = pd.read_csv(f\"{DATA_DIR}/train_labels.csv\")\nval_labels = pd.read_csv(f\"{DATA_DIR}/validation_labels.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T10:04:18.435734Z","iopub.execute_input":"2026-02-08T10:04:18.435925Z","iopub.status.idle":"2026-02-08T10:04:27.247236Z","shell.execute_reply.started":"2026-02-08T10:04:18.435905Z","shell.execute_reply":"2026-02-08T10:04:27.246654Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_coordinates(df):\n    coords = {}\n    for tid, g in df.groupby(df[\"ID\"].str.split(\"_\").str[0]):\n        g = g.sort_values(\"resid\")\n        coords[tid] = g[[\"x_1\", \"y_1\", \"z_1\"]].values.astype(np.float32)\n    return coords","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T10:04:27.249222Z","iopub.execute_input":"2026-02-08T10:04:27.249473Z","iopub.status.idle":"2026-02-08T10:04:27.253639Z","shell.execute_reply.started":"2026-02-08T10:04:27.249429Z","shell.execute_reply":"2026-02-08T10:04:27.253036Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_coords = load_coordinates(train_labels)\nval_coords = load_coordinates(val_labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T10:04:27.254369Z","iopub.execute_input":"2026-02-08T10:04:27.254693Z","iopub.status.idle":"2026-02-08T10:04:41.565279Z","shell.execute_reply.started":"2026-02-08T10:04:27.254661Z","shell.execute_reply":"2026-02-08T10:04:41.5645Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Load MSA\n\ndef load_msa_subsample(target_id, max_msa=32):\n    path = os.path.join(MSA_DIR, f\"{target_id}.MSA.fasta\")\n    seqs, current = [], []\n\n    with open(path) as f:\n        for line in f:\n            line = line.strip()\n            if not line:\n                continue\n            if line.startswith(\">\"):\n                if current:\n                    seqs.append(current)\n                    current = []\n            else:\n                current.extend(line)\n        if current:\n            seqs.append(current)\n\n    if len(seqs) > max_msa:\n        idx = np.random.choice(len(seqs), max_msa, replace=False)\n        seqs = [seqs[i] for i in idx]\n\n    return torch.tensor([[NUC_MAP.get(c, 4) for c in s] for s in seqs])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T10:05:57.26425Z","iopub.execute_input":"2026-02-08T10:05:57.264777Z","iopub.status.idle":"2026-02-08T10:05:57.270339Z","shell.execute_reply.started":"2026-02-08T10:05:57.264748Z","shell.execute_reply":"2026-02-08T10:05:57.269747Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#sanity check\n\ni = 0\ntid = train_seq.loc[i, \"target_id\"]\n\nassert len(train_seq.loc[i, \"target_sequence\"]) == train_coords[tid].shape[0]\nassert load_msa_subsample(tid).shape[1] == len(train_seq.loc[i, \"target_sequence\"])\n\nprint(\"sanity check passed\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T10:06:16.737869Z","iopub.execute_input":"2026-02-08T10:06:16.73821Z","iopub.status.idle":"2026-02-08T10:06:16.872533Z","shell.execute_reply.started":"2026-02-08T10:06:16.738182Z","shell.execute_reply":"2026-02-08T10:06:16.871817Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RNADataset(Dataset):\n    def __init__(self, df, coords=None):\n        self.df = df.reset_index(drop=True)\n        self.coords = coords\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.loc[idx]\n        tid = row.target_id\n        seq = torch.tensor([NUC_MAP[c] for c in row.sequence])\n\n        msa = load_msa_subsample(tid)\n        mask = (seq !=4).float()\n\n        out = {\n            \"seq\": seq,\n            \"msa\": msa,\n            \"mask\": mask,\n            \"id\": tid\n        }\n\n        if self.coords is not None:\n            out[\"coords\"] = torch.tensor(self.coords[tid])\n\n        return out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T10:06:20.554526Z","iopub.execute_input":"2026-02-08T10:06:20.554812Z","iopub.status.idle":"2026-02-08T10:06:20.560466Z","shell.execute_reply.started":"2026-02-08T10:06:20.554787Z","shell.execute_reply":"2026-02-08T10:06:20.559758Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Collate function (No global padding)\n\ndef collate_fn(batch):\n    b = batch[0]\n\n    L = len(b[\"seq\"])\n    M = b[\"msa\"].shape[0]\n\n    return {\n        \"seq\": b[\"seq\"].unsqueeze(0),\n        \"msa\": b[\"msa\"].unsqueeze(0),\n        \"mask\": b[\"mask\"].unsqueeze(0),\n        \"coords\": b.get(\"coords\", None),\n        \"id\": b[\"id\"]\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T10:06:21.562368Z","iopub.execute_input":"2026-02-08T10:06:21.562943Z","iopub.status.idle":"2026-02-08T10:06:21.567599Z","shell.execute_reply.started":"2026-02-08T10:06:21.562906Z","shell.execute_reply":"2026-02-08T10:06:21.566749Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#DataLoaders\n\ntrain_loader = DataLoader(\n    RNADataset(train_seq, train_coords),\n    batch_size = 1,\n    shuffle = True,\n    collate_fn = collate_fn\n)\n\nval_loader = DataLoader(\n    RNADataset(val_seq, val_coords),\n    batch_size = 1,\n    collate_fn = collate_fn\n)\n\ntest_loader = DataLoader(\n    RNADataset(test_seq),\n    batch_size = 1,\n    collate_fn = collate_fn\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T10:06:24.200612Z","iopub.execute_input":"2026-02-08T10:06:24.201274Z","iopub.status.idle":"2026-02-08T10:06:24.208758Z","shell.execute_reply.started":"2026-02-08T10:06:24.201245Z","shell.execute_reply":"2026-02-08T10:06:24.207961Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#MSA based Transformer model\n\nclass SafeMSAModel(nn.Module):\n    def __init__(self, d=128):\n        super().__init__()\n        self.embed = nn.Embedding(5, d)\n\n        self.seq_enc = nn.TransformerEncoder(\n            nn.TransformerEncoderLayer(d, 8, batch_first=True),\n            num_layers = 2\n        )\n        self.msa_enc = nn.TransformerEncoder(\n            nn.TransformerEncoderLayer(d, 8, batch_first=True),\n            num_layers = 1\n        )\n\n        self.out = nn.Linear(d, 3)\n\n    def forward(self, seq, msa):\n        '''\n        seq : (B, L)\n        msa : (B, M, L)\n        '''\n\n        #Sequence path\n        s = self.embed(seq)      #(B, L, D)\n        s = self.seq_enc(s)    #(B, L, D)\n\n        #MSA path\n        m = self.embed(msa)      #(B, M, L, D)\n        m = m.mean(dim=1)      #(B, L, D)    reduce msa rows\n        m = self.msa_enc(m)\n\n        #Fuse \n        x = s + m\n        \n        return self.out(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T10:15:58.061029Z","iopub.execute_input":"2026-02-08T10:15:58.061777Z","iopub.status.idle":"2026-02-08T10:15:58.067584Z","shell.execute_reply.started":"2026-02-08T10:15:58.061749Z","shell.execute_reply":"2026-02-08T10:15:58.066813Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Training Loop","metadata":{}},{"cell_type":"code","source":"def train_epoch(model, loader, opt):\n    model.train()\n    for b in train_loader:\n        seq = b[\"seq\"].to(DEVICE)\n        msa = b[\"msa\"].to(DEVICE)\n        tgt = b[\"coords\"].to(DEVICE)\n\n        opt.zero_grad()\n        pred = model(seq, msa)\n        loss = ((pred - tgt) ** 2).mean()\n        loss.backward()\n        opt.step()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T10:16:03.988698Z","iopub.execute_input":"2026-02-08T10:16:03.989388Z","iopub.status.idle":"2026-02-08T10:16:03.993673Z","shell.execute_reply.started":"2026-02-08T10:16:03.98936Z","shell.execute_reply":"2026-02-08T10:16:03.992909Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Generate 5 structure","metadata":{}},{"cell_type":"code","source":"model = SafeMSAModel().to(DEVICE)\nopt = torch.optim.AdamW(model.parameters(), lr=1e-4)\n\nfor epoch in range(2):\n    train_epoch(model, train_loader, opt)\n    print(f\"Epoch {epoch+1} done\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T10:16:06.084573Z","iopub.execute_input":"2026-02-08T10:16:06.085373Z","execution_failed":"2026-02-08T10:48:37.7Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Create submission\n\nrows = []\n\nmodel.eval()\nfor b in test_loader:\n    seq = b[\"seq\"].to(DEVICE)\n    msa = b[\"msa\"].to(DEVICE)\n    tid = b[\"id\"]\n\n    with torch.no_grad():\n        preds = [model(seq, msa)[0].cpu().numpy() for _ in range(5)]\n\n    L = seq.shape[1]\n    for i in range(L):\n        row = {\"ID\": f\"{tid}_{i+1}\", \"resname\": \"A\", \"resid\": i+1}\n        for k in range(5):\n            x, y, z = preds[k][i]\n            row[f\"x_{k+1}\"] = float(x)\n            row[f\"y_{k+1}\"] = float(y)\n            row[f\"z_{k+1}\"] = float(z)\n        rows.append(row)\n\npd.DataFrame(rows).to_csv(\"submission.csv\", index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T10:04:41.586278Z","iopub.status.idle":"2026-02-08T10:04:41.586549Z","shell.execute_reply.started":"2026-02-08T10:04:41.586402Z","shell.execute_reply":"2026-02-08T10:04:41.586415Z"}},"outputs":[],"execution_count":null}]}