{"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":31260,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport random\nimport warnings\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader\nfrom torch_cluster import knn_graph\nimport math\n\nimport matplotlib.pyplot as plt\nfrom mpl_toolkits.mplot3d import Axes3D\n\nwarnings.filterwarnings('ignore')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-31T13:28:43.856061Z","iopub.execute_input":"2026-01-31T13:28:43.85686Z","iopub.status.idle":"2026-01-31T13:28:43.861553Z","shell.execute_reply.started":"2026-01-31T13:28:43.856828Z","shell.execute_reply":"2026-01-31T13:28:43.860651Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def seed_everything(seed=28):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nseed_everything(28)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T12:55:35.077121Z","iopub.execute_input":"2026-01-31T12:55:35.077732Z","iopub.status.idle":"2026-01-31T12:55:35.088087Z","shell.execute_reply.started":"2026-01-31T12:55:35.077704Z","shell.execute_reply":"2026-01-31T12:55:35.087251Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"NUC2IDX = {'A': 0, 'U': 1, 'G': 2, 'C': 3}\nNUCLEOTIDE_MAPPING = {\n    'A': 'A', 'U': 'U', 'G': 'G', 'C': 'C',\n    'I': 'A', '1MA': 'A', 'PSU': 'U', 'M2G': 'G', '5MC': 'C', 'T': 'U',\n}\n\ndef clean_sequence(seq: str) -> str:\n    return \"\".join(NUCLEOTIDE_MAPPING.get(b, 'A') for b in seq)\n\n\ndef encode_sequence(seq: str) -> torch.LongTensor:\n    return torch.tensor([NUC2IDX[b] for b in seq], dtype=torch.long)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T12:55:37.309078Z","iopub.execute_input":"2026-01-31T12:55:37.309704Z","iopub.status.idle":"2026-01-31T12:55:37.315404Z","shell.execute_reply.started":"2026-01-31T12:55:37.309676Z","shell.execute_reply":"2026-01-31T12:55:37.31438Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def preprocess_labels(train_labels: pd.DataFrame):\n    train_labels = train_labels.copy()\n\n    train_labels['target_id'] = train_labels['ID'].str.split('_').str[0]\n\n    grouped = {}\n    for tid, g in train_labels.groupby('target_id'):\n        g = g.sort_values('resid')\n\n        coords = g[['x_1', 'y_1', 'z_1']].values.astype('float32')\n        mask = ~np.isnan(coords).any(axis=1)\n\n        coords[~mask] = 0.0\n\n        grouped[tid] = {\n            \"coords\": coords,\n            \"mask\": mask,\n        }\n\n    return grouped","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T12:55:39.050867Z","iopub.execute_input":"2026-01-31T12:55:39.051477Z","iopub.status.idle":"2026-01-31T12:55:39.05627Z","shell.execute_reply.started":"2026-01-31T12:55:39.051445Z","shell.execute_reply":"2026-01-31T12:55:39.055456Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RNADiffusionDataset(torch.utils.data.Dataset):\n    def __init__(\n        self,\n        seq_df: pd.DataFrame,\n        label_dict: dict,\n        normalize: bool = True,\n        max_len: int | None = None,\n    ):\n        self.records = []\n\n        for _, row in seq_df.iterrows():\n            tid = row['target_id']\n            if tid not in label_dict:\n                continue\n\n            seq = clean_sequence(row['sequence'])\n            coords = label_dict[tid]['coords']\n            mask = label_dict[tid]['mask']\n\n            if len(seq) != coords.shape[0]:\n                continue  # safety\n\n            if max_len and len(seq) > max_len:\n                continue\n\n            self.records.append((seq, coords, mask))\n\n        self.normalize = normalize\n\n    def __len__(self):\n        return len(self.records)\n\n    def __getitem__(self, idx):\n        seq, coords, mask = self.records[idx]\n\n        seq = encode_sequence(seq)\n        coords = torch.tensor(coords, dtype=torch.float32)\n        mask = torch.tensor(mask, dtype=torch.bool)\n\n        if self.normalize:\n            valid = mask\n            mean = coords[valid].mean(0, keepdim=True)\n            std = coords[valid].std(0, keepdim=True) + 1e-6\n\n            coords = (coords - mean) / std\n\n        return {\n            \"seq\": seq,\n            \"coords\": coords,\n            \"mask\": mask,\n            \"length\": len(seq),\n        }\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T12:55:40.897091Z","iopub.execute_input":"2026-01-31T12:55:40.898083Z","iopub.status.idle":"2026-01-31T12:55:40.906133Z","shell.execute_reply.started":"2026-01-31T12:55:40.898046Z","shell.execute_reply":"2026-01-31T12:55:40.905279Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def collate_rna(batch):\n    max_len = max(x['length'] for x in batch)\n\n    seqs   = []\n    coords = []\n    masks  = []\n\n    for b in batch:\n        L = b['length']\n        pad = max_len - L\n\n        seqs.append(\n            F.pad(b['seq'], (0, pad), value=0)\n        )\n\n        coords.append(\n            F.pad(b['coords'], (0, 0, 0, pad))\n        )\n\n        masks.append(\n            F.pad(b['mask'], (0, pad), value=False)\n        )\n\n    return {\n        \"seq\": torch.stack(seqs),\n        \"coords\": torch.stack(coords),\n        \"mask\": torch.stack(masks),\n    }\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T12:55:44.865315Z","iopub.execute_input":"2026-01-31T12:55:44.865872Z","iopub.status.idle":"2026-01-31T12:55:44.871344Z","shell.execute_reply.started":"2026-01-31T12:55:44.865841Z","shell.execute_reply":"2026-01-31T12:55:44.870608Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def cosine_beta_schedule(T, s=0.008):\n    steps = T + 1\n    x = torch.linspace(0, T, steps)\n    alphas_cumprod = torch.cos(((x / T) + s) / (1 + s) * math.pi / 2) ** 2\n    alphas_cumprod = alphas_cumprod / alphas_cumprod[0]\n    betas = 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1])\n    return torch.clamp(betas, 1e-5, 0.999)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T12:55:46.756882Z","iopub.execute_input":"2026-01-31T12:55:46.75756Z","iopub.status.idle":"2026-01-31T12:55:46.762178Z","shell.execute_reply.started":"2026-01-31T12:55:46.757529Z","shell.execute_reply":"2026-01-31T12:55:46.761507Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Diffusion:\n    def __init__(self, T=1000, device='cuda'):\n        self.T = T\n        self.device = device\n\n        self.betas = cosine_beta_schedule(T).to(device)\n        self.alphas = 1.0 - self.betas\n        self.alphas_cumprod = torch.cumprod(self.alphas, dim=0)\n\n        self.sqrt_acp = torch.sqrt(self.alphas_cumprod)\n        self.sqrt_1m_acp = torch.sqrt(1 - self.alphas_cumprod)\n\n    def q_sample(self, x0, t, noise):\n        a = self.sqrt_acp[t][:, None, None]\n        b = self.sqrt_1m_acp[t][:, None, None]\n        return a * x0 + b * noise\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T12:55:48.417187Z","iopub.execute_input":"2026-01-31T12:55:48.417849Z","iopub.status.idle":"2026-01-31T12:55:48.42306Z","shell.execute_reply.started":"2026-01-31T12:55:48.417819Z","shell.execute_reply":"2026-01-31T12:55:48.422305Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_edges(x, mask, k=16):\n    \"\"\"\n    x: [B,L,3]\n    mask: [B,L]\n    \"\"\"\n    B, L, _ = x.shape\n    edges = []\n\n    for b in range(B):\n        xb = x[b][mask[b]]\n        if xb.shape[0] < 2:\n            edges.append(None)\n            continue\n        edge = knn_graph(xb, k=min(k, xb.shape[0]-1), loop=False)\n        edges.append(edge)\n\n    return edges","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T12:55:50.22292Z","iopub.execute_input":"2026-01-31T12:55:50.223573Z","iopub.status.idle":"2026-01-31T12:55:50.228652Z","shell.execute_reply.started":"2026-01-31T12:55:50.223549Z","shell.execute_reply":"2026-01-31T12:55:50.227806Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EGNNLayer(nn.Module):\n    def __init__(self, d):\n        super().__init__()\n        self.bn = nn.LayerNorm(d)\n        self.phi_e = nn.Sequential(\n            nn.Linear(2*d + 1, d),\n            nn.SiLU(),\n            nn.Linear(d, d)\n        )\n        self.phi_x = nn.Sequential(\n            nn.Linear(d, d),\n            nn.SiLU(),\n            nn.Linear(d, 1, bias=False) \n        )\n        self.phi_h = nn.Sequential(\n            nn.Linear(d + d, d),\n            nn.SiLU(),\n            nn.Linear(d, d)\n        )\n\n    def forward(self, h, x, edge_index):\n        i, j = edge_index\n        \n        diff = x[i] - x[j]\n        dist2 = (diff ** 2).sum(dim=-1, keepdim=True)\n        dist_feat = torch.log(dist2 + 1e-7) \n\n        m_ij = self.phi_e(torch.cat([h[i], h[j], dist_feat], dim=-1))\n        \n        dx = diff * torch.tanh(self.phi_x(m_ij)) \n        x = x + torch.zeros_like(x).index_add(0, i, dx)\n\n        m_i = torch.zeros_like(h).index_add(0, i, m_ij)\n        h = h + self.phi_h(torch.cat([h, m_i], dim=-1)) \n        h = self.bn(h) \n\n        return h, x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T13:11:11.661535Z","iopub.execute_input":"2026-01-31T13:11:11.662127Z","iopub.status.idle":"2026-01-31T13:11:11.669571Z","shell.execute_reply.started":"2026-01-31T13:11:11.662096Z","shell.execute_reply":"2026-01-31T13:11:11.668802Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EGNN(nn.Module):\n    def __init__(self, d, n_layers=4):\n        super().__init__()\n        self.layers = nn.ModuleList([EGNNLayer(d) for _ in range(n_layers)])\n\n    def forward(self, h, x, edge_index):\n        for layer in self.layers:\n            h, x = layer(h, x, edge_index)\n        return h, x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T13:11:15.567351Z","iopub.execute_input":"2026-01-31T13:11:15.567656Z","iopub.status.idle":"2026-01-31T13:11:15.572659Z","shell.execute_reply.started":"2026-01-31T13:11:15.567631Z","shell.execute_reply":"2026-01-31T13:11:15.571903Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RNADiffusionModel(nn.Module):\n    def __init__(self, d=128):\n        super().__init__()\n        self.embed = nn.Embedding(4, d)\n        self.time = nn.Embedding(1000, d)\n\n        self.egnn = EGNN(d, n_layers=6)\n        self.out = nn.Linear(d, 3)\n\n    def forward(self, seq, x_t, t, mask):\n        h = self.embed(seq) + self.time(t)[:, None, :]\n\n        B, L, _ = x_t.shape\n        eps = torch.zeros_like(x_t)\n\n        for b in range(B):\n            idx = mask[b]\n            if idx.sum() < 2:\n                continue\n\n            xb = x_t[b, idx]\n            hb = h[b, idx]\n\n            edge = knn_graph(xb, k=min(16, xb.shape[0]-1))\n            hb, _ = self.egnn(hb, xb, edge)\n\n            eps[b, idx] = self.out(hb)\n\n        return eps","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T13:09:40.28335Z","iopub.execute_input":"2026-01-31T13:09:40.284244Z","iopub.status.idle":"2026-01-31T13:09:40.290519Z","shell.execute_reply.started":"2026-01-31T13:09:40.284212Z","shell.execute_reply":"2026-01-31T13:09:40.289787Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def diffusion_loss(model, diffusion, batch):\n    x0 = batch['coords'].cuda()\n    seq = batch['seq'].cuda()\n    mask = batch['mask'].cuda()\n\n    B = x0.shape[0]\n    t = torch.randint(0, diffusion.T, (B,), device=x0.device)\n    noise = torch.randn_like(x0)\n\n    x_t = diffusion.q_sample(x0, t, noise)\n    eps_pred = model(seq, x_t, t, mask)\n\n    loss = ((eps_pred - noise) ** 2)[mask].mean()\n    return loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T13:09:43.195908Z","iopub.execute_input":"2026-01-31T13:09:43.196671Z","iopub.status.idle":"2026-01-31T13:09:43.201834Z","shell.execute_reply.started":"2026-01-31T13:09:43.196639Z","shell.execute_reply":"2026-01-31T13:09:43.200727Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef sample_structure(model, diffusion, seq, mask):\n    L = seq.shape[0]\n    x = torch.randn(L, 3, device=seq.device)\n\n    for t in reversed(range(diffusion.T)):\n        tt = torch.tensor([t], device=seq.device)\n        eps = model(seq[None], x[None], tt, mask[None])[0]\n\n        alpha = diffusion.alphas[t]\n        acp = diffusion.alphas_cumprod[t]\n\n        x = (x - (1 - alpha) / torch.sqrt(1 - acp) * eps) / torch.sqrt(alpha)\n\n    return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T13:09:44.758958Z","iopub.execute_input":"2026-01-31T13:09:44.759572Z","iopub.status.idle":"2026-01-31T13:09:44.764886Z","shell.execute_reply.started":"2026-01-31T13:09:44.759543Z","shell.execute_reply":"2026-01-31T13:09:44.763919Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(model, loader, optimizer, diffusion, device):\n    model.train()\n    total_loss = 0\n    for batch in loader:\n        optimizer.zero_grad()\n        \n        batch = {k: v.to(device) if isinstance(v, torch.Tensor) else v for k, v in batch.items()}\n        \n        loss = diffusion_loss(model, diffusion, batch)\n        \n        # Проверка на nan в лоссе\n        if torch.isnan(loss):\n            print(\"Warning: nan loss detected, skipping batch\")\n            continue\n            \n        loss.backward()\n        \n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        \n        optimizer.step()\n        total_loss += loss.item()\n    \n    return total_loss / len(loader) if len(loader) > 0 else 0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T13:09:54.517641Z","iopub.execute_input":"2026-01-31T13:09:54.518377Z","iopub.status.idle":"2026-01-31T13:09:54.523883Z","shell.execute_reply.started":"2026-01-31T13:09:54.518323Z","shell.execute_reply":"2026-01-31T13:09:54.523186Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def write_submission(test_seqs, predictions, path='submission.csv'):\n    rows = []\n\n    for tid, pred_list in predictions.items():\n        seq = clean_sequence(\n            test_seqs.loc[test_seqs.target_id == tid, 'sequence'].values[0]\n        )\n\n        for i, base in enumerate(seq):\n            row = {\n                'ID': f'{tid}_{i+1}',\n                'resname': base,\n                'resid': i+1\n            }\n            for k, coords in enumerate(pred_list):\n                x, y, z = coords[i].cpu().numpy()\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\n            rows.append(row)\n\n    pd.DataFrame(rows).to_csv(path, index=False)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T12:56:05.349378Z","iopub.execute_input":"2026-01-31T12:56:05.350047Z","iopub.status.idle":"2026-01-31T12:56:05.355945Z","shell.execute_reply.started":"2026-01-31T12:56:05.350015Z","shell.execute_reply":"2026-01-31T12:56:05.355185Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T12:56:07.604799Z","iopub.execute_input":"2026-01-31T12:56:07.605209Z","iopub.status.idle":"2026-01-31T12:56:07.637419Z","shell.execute_reply.started":"2026-01-31T12:56:07.605185Z","shell.execute_reply":"2026-01-31T12:56:07.636541Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_PATH = '/kaggle/input/stanford-rna-3d-folding-2/'\ntrain_seqs = pd.read_csv(DATA_PATH + 'train_sequences.csv')\ntest_seqs = pd.read_csv(DATA_PATH + 'test_sequences.csv')\ntrain_labels = pd.read_csv(DATA_PATH + 'train_labels.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T12:56:09.021506Z","iopub.execute_input":"2026-01-31T12:56:09.021781Z","iopub.status.idle":"2026-01-31T12:56:18.44179Z","shell.execute_reply.started":"2026-01-31T12:56:09.021761Z","shell.execute_reply":"2026-01-31T12:56:18.441107Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"label_dict = preprocess_labels(train_labels)\ndataset = RNADiffusionDataset(train_seqs, label_dict, max_len=512)\nloader = DataLoader(dataset, batch_size=8, shuffle=True, collate_fn=collate_rna)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T12:56:18.442953Z","iopub.execute_input":"2026-01-31T12:56:18.443208Z","iopub.status.idle":"2026-01-31T12:56:36.625696Z","shell.execute_reply.started":"2026-01-31T12:56:18.443176Z","shell.execute_reply":"2026-01-31T12:56:36.625132Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"batch = next(iter(loader))\n\nprint(batch['seq'].shape)     # [B, L]\nprint(batch['coords'].shape)  # [B, L, 3]\nprint(batch['mask'].shape)    # [B, L]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T12:56:36.626654Z","iopub.execute_input":"2026-01-31T12:56:36.626948Z","iopub.status.idle":"2026-01-31T12:56:36.715125Z","shell.execute_reply.started":"2026-01-31T12:56:36.626918Z","shell.execute_reply":"2026-01-31T12:56:36.714371Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = RNADiffusionModel(d=128).to(device)\ndiffusion = Diffusion(T=1000, device=device)\noptimizer = torch.optim.AdamW(model.parameters(), lr=1e-5, weight_decay=1e-2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T13:11:39.241053Z","iopub.execute_input":"2026-01-31T13:11:39.241796Z","iopub.status.idle":"2026-01-31T13:11:39.262671Z","shell.execute_reply.started":"2026-01-31T13:11:39.241724Z","shell.execute_reply":"2026-01-31T13:11:39.26206Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_epochs = 10\nfor epoch in range(num_epochs):\n    avg_loss = train_one_epoch(model, loader, optimizer, diffusion, device)\n    print(f\"Epoch {epoch+1} | Loss: {avg_loss:.6f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T13:11:41.042779Z","iopub.execute_input":"2026-01-31T13:11:41.043483Z","iopub.status.idle":"2026-01-31T13:22:27.872519Z","shell.execute_reply.started":"2026-01-31T13:11:41.043455Z","shell.execute_reply":"2026-01-31T13:22:27.871785Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef inference_test_set(model, diffusion, test_df, label_dict, device):\n    model.eval()\n    predictions = {}\n\n    for _, row in test_df.iterrows():\n        tid = row['target_id']\n        seq_str = clean_sequence(row['sequence'])\n        seq_idx = encode_sequence(seq_str).to(device)\n        \n        L = len(seq_idx)\n        mask = torch.ones(L, dtype=torch.bool, device=device)\n        \n        coords_pred = sample_structure(model, diffusion, seq_idx, mask)\n\n        predictions[tid] = [coords_pred] \n        \n    return predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T13:22:27.874334Z","iopub.execute_input":"2026-01-31T13:22:27.874569Z","iopub.status.idle":"2026-01-31T13:22:27.879945Z","shell.execute_reply.started":"2026-01-31T13:22:27.874548Z","shell.execute_reply":"2026-01-31T13:22:27.879349Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nT_STEPS = 1000 \n\nmodel = RNADiffusionModel(d=128).to(device)\n\ndiffusion = Diffusion(T=T_STEPS, device=device)\n\ntest_seqs = pd.read_csv(DATA_PATH + 'test_sequences.csv')\nif 'target_id' not in test_seqs.columns:\n    test_seqs['target_id'] = test_seqs['ID'].apply(lambda x: x.split('_')[0])\n    test_seqs = test_seqs.drop_duplicates('target_id')\n\nprint(\"Starting inference...\")\ntest_predictions = inference_test_set(model, diffusion, test_seqs, label_dict, device)\n\nprint(\"Writing submission...\")\nwrite_submission(test_seqs, test_predictions, path='submission.csv')\nprint(\"Done! File saved as submission.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T13:22:27.880663Z","iopub.execute_input":"2026-01-31T13:22:27.880872Z","iopub.status.idle":"2026-01-31T13:27:19.970406Z","shell.execute_reply.started":"2026-01-31T13:22:27.880853Z","shell.execute_reply":"2026-01-31T13:27:19.969622Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_rna_3d(coords, sequence=None, title=\"RNA 3D Structure\"):\n\n    coords = np.asarray(coords)\n    n = coords.shape[0]\n\n    fig = plt.figure(figsize=(8, 7))\n    ax = fig.add_subplot(111, projection='3d')\n\n    colors = np.linspace(0, 1, n)\n\n    sc = ax.scatter(\n        coords[:, 0], coords[:, 1], coords[:, 2],\n        c=colors, cmap='viridis', s=25\n    )\n\n    ax.plot(coords[:, 0], coords[:, 1], coords[:, 2], alpha=0.6)\n\n    ax.set_title(title)\n    ax.set_xlabel(\"X\")\n    ax.set_ylabel(\"Y\")\n    ax.set_zlabel(\"Z\")\n\n    max_range = (coords.max(axis=0) - coords.min(axis=0)).max() / 2\n    mid = coords.mean(axis=0)\n    ax.set_xlim(mid[0] - max_range, mid[0] + max_range)\n    ax.set_ylim(mid[1] - max_range, mid[1] + max_range)\n    ax.set_zlim(mid[2] - max_range, mid[2] + max_range)\n\n    cbar = plt.colorbar(sc, ax=ax, shrink=0.6)\n    cbar.set_label(\"Residue index\")\n\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T13:28:37.026475Z","iopub.execute_input":"2026-01-31T13:28:37.02681Z","iopub.status.idle":"2026-01-31T13:28:37.035494Z","shell.execute_reply.started":"2026-01-31T13:28:37.026782Z","shell.execute_reply":"2026-01-31T13:28:37.034628Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"example_tid = list(test_predictions.keys())[0]\ncoords_to_plot = test_predictions[example_tid][0].detach().cpu().numpy()\nplot_rna_3d(coords_to_plot, title=f\"Predicted RNA Structure: {example_tid}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-31T13:28:55.072931Z","iopub.execute_input":"2026-01-31T13:28:55.073562Z","iopub.status.idle":"2026-01-31T13:28:55.30221Z","shell.execute_reply.started":"2026-01-31T13:28:55.073518Z","shell.execute_reply":"2026-01-31T13:28:55.301354Z"}},"outputs":[],"execution_count":null}]}