{"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":"nvidiaTeslaT4","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":"# ============================================================\n# Stanford RNA 3D Folding Part 2\n# Baseline Transformer Model\n#\n# Licensed under the Apache License, Version 2.0\n# http://www.apache.org/licenses/LICENSE-2.0\n# ============================================================","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-20T13:42:12.441097Z","iopub.execute_input":"2026-01-20T13:42:12.441821Z","iopub.status.idle":"2026-01-20T13:42:12.447663Z","shell.execute_reply.started":"2026-01-20T13:42:12.441791Z","shell.execute_reply":"2026-01-20T13:42:12.447121Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport random\nimport numpy as np\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T14:36:33.161442Z","iopub.execute_input":"2026-01-20T14:36:33.161995Z","iopub.status.idle":"2026-01-20T14:36:33.166557Z","shell.execute_reply.started":"2026-01-20T14:36:33.16196Z","shell.execute_reply":"2026-01-20T14:36:33.16579Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"CUDA available:\", torch.cuda.is_available())\nif torch.cuda.is_available():\n    print(\"GPU:\", torch.cuda.get_device_name(0))\n    device = torch.device(\"cuda\")\nelse:\n    print(\"Using CPU - performance will be poor for long seq!\")\n    device = torch.device(\"cpu\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T14:36:35.675339Z","iopub.execute_input":"2026-01-20T14:36:35.676033Z","iopub.status.idle":"2026-01-20T14:36:35.681074Z","shell.execute_reply.started":"2026-01-20T14:36:35.676006Z","shell.execute_reply":"2026-01-20T14:36:35.680518Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_PATH = \"/kaggle/input/stanford-rna-3d-folding-2\"\n\nprint(os.listdir(DATA_PATH))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T14:36:41.166817Z","iopub.execute_input":"2026-01-20T14:36:41.167428Z","iopub.status.idle":"2026-01-20T14:36:41.173333Z","shell.execute_reply.started":"2026-01-20T14:36:41.167399Z","shell.execute_reply":"2026-01-20T14:36:41.172684Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\ntrain_seq = pd.read_csv(f\"{DATA_PATH}/train_sequences.csv\")\ntrain_labels = pd.read_csv(f\"{DATA_PATH}/train_labels.csv\")\n\nval_seq = pd.read_csv(f\"{DATA_PATH}/validation_sequences.csv\")\nval_labels = pd.read_csv(f\"{DATA_PATH}/validation_labels.csv\")\n\ntest_seq = pd.read_csv(f\"{DATA_PATH}/test_sequences.csv\")\n\ntrain_seq.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T14:36:43.648637Z","iopub.execute_input":"2026-01-20T14:36:43.64893Z","iopub.status.idle":"2026-01-20T14:36:51.419263Z","shell.execute_reply.started":"2026-01-20T14:36:43.648905Z","shell.execute_reply":"2026-01-20T14:36:51.418584Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_labels.columns","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T14:36:55.034368Z","iopub.execute_input":"2026-01-20T14:36:55.035139Z","iopub.status.idle":"2026-01-20T14:36:55.039785Z","shell.execute_reply.started":"2026-01-20T14:36:55.0351Z","shell.execute_reply":"2026-01-20T14:36:55.039091Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Keep only required columns\ntrain_labels = train_labels[[\"ID\", \"resid\", \"x_1\", \"y_1\", \"z_1\"]]\n\n# Rename for consistency\ntrain_labels = train_labels.rename(columns={\n    \"ID\": \"target_id\",\n    \"x_1\": \"x\",\n    \"y_1\": \"y\",\n    \"z_1\": \"z\"\n})\n\ntrain_labels.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T14:36:57.692199Z","iopub.execute_input":"2026-01-20T14:36:57.692914Z","iopub.status.idle":"2026-01-20T14:36:58.217987Z","shell.execute_reply.started":"2026-01-20T14:36:57.692885Z","shell.execute_reply":"2026-01-20T14:36:58.217401Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_seq = train_seq[[\"target_id\", \"sequence\"]]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T14:37:01.937458Z","iopub.execute_input":"2026-01-20T14:37:01.938234Z","iopub.status.idle":"2026-01-20T14:37:01.942823Z","shell.execute_reply.started":"2026-01-20T14:37:01.938205Z","shell.execute_reply":"2026-01-20T14:37:01.942105Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from collections import defaultdict\n\n# 🔹 ID clean: 157D_1 → 157D\ntrain_labels[\"target_id\"] = train_labels[\"target_id\"].str.split(\"_\").str[0]\n\ncoords_dict = defaultdict(list)\n\nfor _, row in train_labels.iterrows():\n    coords_dict[row[\"target_id\"]].append(\n        [row[\"x\"], row[\"y\"], row[\"z\"]]\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T14:37:04.808575Z","iopub.execute_input":"2026-01-20T14:37:04.809129Z","iopub.status.idle":"2026-01-20T14:41:56.22261Z","shell.execute_reply.started":"2026-01-20T14:37:04.809083Z","shell.execute_reply":"2026-01-20T14:41:56.221813Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tid = train_seq.iloc[0][\"target_id\"]\n\nprint(\"Target ID:\", tid)\nprint(\"Sequence length:\", len(train_seq.iloc[0][\"sequence\"]))\nprint(\"Coords length:\", len(coords_dict[tid]))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T14:42:19.442611Z","iopub.execute_input":"2026-01-20T14:42:19.443155Z","iopub.status.idle":"2026-01-20T14:42:19.448124Z","shell.execute_reply.started":"2026-01-20T14:42:19.443101Z","shell.execute_reply":"2026-01-20T14:42:19.447421Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"WINDOW_SIZE = 512   # 🔥 KEY PARAMETER\n\ndef collate_fn(batch):\n    seqs, coords = zip(*batch)\n\n    batch_seqs = []\n    batch_coords = []\n    batch_masks = []\n\n    for seq, coord in zip(seqs, coords):\n        L = len(seq)\n\n        # Random window start (training-time augmentation)\n        if L > WINDOW_SIZE:\n            start = random.randint(0, L - WINDOW_SIZE)\n            end = start + WINDOW_SIZE\n        else:\n            start = 0\n            end = L\n\n        seq_win = seq[start:end]\n        coord_win = coord[start:end]\n\n        valid = ~torch.isnan(coord_win).any(dim=1)\n\n        batch_seqs.append(seq_win)\n        batch_coords.append(torch.nan_to_num(coord_win, nan=0.0))\n        batch_masks.append(valid)\n\n    # Padding (now max_len <= WINDOW_SIZE)\n    max_len = max(len(s) for s in batch_seqs)\n    B = len(batch_seqs)\n\n    padded_seqs = torch.zeros(B, max_len, dtype=torch.long)\n    padded_coords = torch.zeros(B, max_len, 3)\n    mask = torch.zeros(B, max_len, dtype=torch.bool)\n\n    for i in range(B):\n        l = len(batch_seqs[i])\n        padded_seqs[i, :l] = batch_seqs[i]\n        padded_coords[i, :l] = batch_coords[i]\n        mask[i, :l] = batch_masks[i]\n\n    return padded_seqs, padded_coords, mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T14:42:34.679395Z","iopub.execute_input":"2026-01-20T14:42:34.679678Z","iopub.status.idle":"2026-01-20T14:42:34.687205Z","shell.execute_reply.started":"2026-01-20T14:42:34.679654Z","shell.execute_reply":"2026-01-20T14:42:34.686603Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RNADataset(Dataset):\n    def __init__(self, seq_df, coords_dict):\n        self.seq_df = seq_df.reset_index(drop=True)\n        self.coords_dict = coords_dict\n\n        self.vocab = {\"A\": 0, \"U\": 1, \"G\": 2, \"C\": 3}\n\n    def encode(self, seq):\n        return torch.tensor([self.vocab[x] for x in seq], dtype=torch.long)\n\n    def __len__(self):\n        return len(self.seq_df)\n\n    def __getitem__(self, idx):\n        row = self.seq_df.iloc[idx]\n        tid = row[\"target_id\"]\n\n        seq = self.encode(row[\"sequence\"])\n        coords = torch.tensor(self.coords_dict[tid], dtype=torch.float32)\n\n        return seq, coords","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T14:42:38.551574Z","iopub.execute_input":"2026-01-20T14:42:38.552283Z","iopub.status.idle":"2026-01-20T14:42:38.557544Z","shell.execute_reply.started":"2026-01-20T14:42:38.552254Z","shell.execute_reply":"2026-01-20T14:42:38.556847Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = RNADataset(train_seq, coords_dict)\n\n# Validation labels same logic apply karo\nval_labels = val_labels[[\"ID\", \"resid\", \"x_1\", \"y_1\", \"z_1\"]]\nval_labels = val_labels.rename(columns={\n    \"ID\": \"target_id\",\n    \"x_1\": \"x\",\n    \"y_1\": \"y\",\n    \"z_1\": \"z\"\n})\n\nval_labels[\"target_id\"] = val_labels[\"target_id\"].str.split(\"_\").str[0]\n\nfrom collections import defaultdict\nval_coords_dict = defaultdict(list)\n\nfor _, row in val_labels.iterrows():\n    val_coords_dict[row[\"target_id\"]].append([row[\"x\"], row[\"y\"], row[\"z\"]])\n\nval_seq = val_seq[[\"target_id\", \"sequence\"]]\nval_dataset = RNADataset(val_seq, val_coords_dict)\n\nlen(train_dataset), len(val_dataset)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T14:42:45.140349Z","iopub.execute_input":"2026-01-20T14:42:45.14095Z","iopub.status.idle":"2026-01-20T14:42:45.766673Z","shell.execute_reply.started":"2026-01-20T14:42:45.140922Z","shell.execute_reply":"2026-01-20T14:42:45.765937Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_loader = DataLoader(\n    train_dataset,\n    batch_size=4,        # SAFE\n    shuffle=True,\n    collate_fn=collate_fn,\n    num_workers=2,\n    pin_memory=True\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=1,  # changed to 1 for stability on long val seq\n    shuffle=False,\n    collate_fn=collate_fn,\n    num_workers=0,  # 0 for stability\n    pin_memory=True if torch.cuda.is_available() else False\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T14:42:52.008453Z","iopub.execute_input":"2026-01-20T14:42:52.009191Z","iopub.status.idle":"2026-01-20T14:42:52.840641Z","shell.execute_reply.started":"2026-01-20T14:42:52.009147Z","shell.execute_reply":"2026-01-20T14:42:52.839851Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"batch = next(iter(train_loader))\n\nseqs, coords, mask = batch\n\nprint(\"Seqs shape:\", seqs.shape)\nprint(\"Coords shape:\", coords.shape)\nprint(\"Mask shape:\", mask.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T14:42:56.965374Z","iopub.execute_input":"2026-01-20T14:42:56.966083Z","iopub.status.idle":"2026-01-20T14:42:57.20343Z","shell.execute_reply.started":"2026-01-20T14:42:56.966056Z","shell.execute_reply":"2026-01-20T14:42:57.202634Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import math\n\nclass PositionalEncoding(nn.Module):\n    def __init__(self, d_model, dropout=0.1, max_len=5000):\n        super().__init__()\n        self.dropout = nn.Dropout(p=dropout)\n        \n        pe = torch.zeros(1, max_len, d_model)\n        position = torch.arange(0, max_len).unsqueeze(1).float()\n        div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))\n        pe[0, :, 0::2] = torch.sin(position * div_term)\n        pe[0, :, 1::2] = torch.cos(position * div_term[:d_model//2])  # agar d_model odd ho to adjust\n        self.register_buffer('pe', pe)\n\n    def forward(self, x):\n        x = x + self.pe[:, :x.size(1), :]\n        return self.dropout(x)\n\nclass RNATransformer(nn.Module):\n    def __init__(\n        self,\n        vocab_size=4,\n        d_model=128,  # improved: badha diya better embedding ke liye\n        nhead=8,\n        num_layers=6,  # improved: zyada layers for better learning\n        dim_feedforward=512,  # improved: larger FFN\n        dropout=0.1\n    ):\n        super().__init__()\n        self.embedding = nn.Embedding(vocab_size, d_model)\n        self.pos_encoder = PositionalEncoding(d_model, dropout)\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=d_model,\n            nhead=nhead,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout,\n            batch_first=True\n        )\n        self.encoder = nn.TransformerEncoder(\n            encoder_layer,\n            num_layers=num_layers\n        )\n        self.coord_head = nn.Linear(d_model, 3)\n\n    def forward(self, seqs, mask):\n        \"\"\"\n        seqs: (B, L)\n        mask: (B, L) True = valid\n        \"\"\"\n        x = self.embedding(seqs)  # (B, L, d_model)\n        x = self.pos_encoder(x)  # positional add kiya\n        x = self.encoder(\n            x,\n            src_key_padding_mask=~mask\n        )\n        coords = self.coord_head(x)  # (B, L, 3)\n        return coords","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T14:43:05.256574Z","iopub.execute_input":"2026-01-20T14:43:05.257299Z","iopub.status.idle":"2026-01-20T14:43:05.266306Z","shell.execute_reply.started":"2026-01-20T14:43:05.257263Z","shell.execute_reply":"2026-01-20T14:43:05.265643Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def masked_mse_loss(pred, target, mask):\n    mask = mask.unsqueeze(-1).float()  # (B, L, 1)\n    diff = (pred - target) ** 2\n    loss = diff * mask\n    total_loss = loss.sum()\n    valid_count = mask.sum()\n    if valid_count < 1e-6:  # almost no valid tokens\n        return torch.tensor(0.0, device=pred.device, requires_grad=True)\n    return total_loss / valid_count","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T14:43:09.365676Z","iopub.execute_input":"2026-01-20T14:43:09.36629Z","iopub.status.idle":"2026-01-20T14:43:09.370508Z","shell.execute_reply.started":"2026-01-20T14:43:09.366262Z","shell.execute_reply":"2026-01-20T14:43:09.369753Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = RNATransformer().to(device)\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=1e-4,  # improved: thoda badha for faster training\n    weight_decay=1e-5  # improved: kam kiya overfitting avoid karne\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T14:43:16.526804Z","iopub.execute_input":"2026-01-20T14:43:16.527349Z","iopub.status.idle":"2026-01-20T14:43:16.551847Z","shell.execute_reply.started":"2026-01-20T14:43:16.527318Z","shell.execute_reply":"2026-01-20T14:43:16.551321Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"seqs, coords, mask = next(iter(train_loader))\nprint(seqs.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T14:43:20.408722Z","iopub.execute_input":"2026-01-20T14:43:20.409416Z","iopub.status.idle":"2026-01-20T14:43:20.635958Z","shell.execute_reply.started":"2026-01-20T14:43:20.409388Z","shell.execute_reply":"2026-01-20T14:43:20.635191Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.train()\n\nseqs, coords, mask = next(iter(train_loader))\n\nseqs = seqs.to(device)\ncoords = coords.to(device)\nmask = mask.to(device)\n\npred = model(seqs, mask)\nloss = masked_mse_loss(pred, coords, mask)\n\nloss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T14:43:24.663044Z","iopub.execute_input":"2026-01-20T14:43:24.663373Z","iopub.status.idle":"2026-01-20T14:43:24.907805Z","shell.execute_reply.started":"2026-01-20T14:43:24.66334Z","shell.execute_reply":"2026-01-20T14:43:24.907175Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EPOCHS = 10\nGRAD_CLIP = 1.0  # prevent gradient explosion\n\nfor epoch in range(EPOCHS):\n    torch.cuda.empty_cache()\n    model.train()\n    train_loss = 0.0\n    num_batches = 0\n    \n    for seqs, coords, mask in train_loader:\n        seqs = seqs.to(device)\n        coords = coords.to(device)\n        mask = mask.to(device)\n        \n        optimizer.zero_grad(set_to_none=True)\n        \n        pred = model(seqs, mask)\n        loss = masked_mse_loss(pred, coords, mask)\n        \n        if torch.isnan(loss) or torch.isinf(loss):\n            print(\"NaN/Inf loss in train! Skipping batch...\")\n            continue\n        \n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n        optimizer.step()\n        \n        train_loss += loss.item()\n        num_batches += 1\n    \n    if num_batches > 0:\n        train_loss /= num_batches\n    else:\n        train_loss = float('nan')\n    \n    # Validation with sliding window for long sequences\n    model.eval()\n    val_loss = 0.0\n    val_batches = 0\n    \n    with torch.no_grad():\n        for seqs, coords, mask in val_loader:\n            seqs = seqs.to(device)\n            coords = coords.to(device)\n            mask = mask.to(device)\n            \n            B, L = seqs.shape\n            \n            if L <= WINDOW_SIZE:\n                pred = model(seqs, mask)\n            else:\n                pred_full = torch.zeros(B, L, 3, device=device)\n                count = torch.zeros(B, L, device=device)\n                step = WINDOW_SIZE - 128  # good overlap\n                \n                for start in range(0, L - WINDOW_SIZE + 1, step):\n                    end = min(start + WINDOW_SIZE, L)\n                    seq_win = seqs[:, start:end]\n                    mask_win = mask[:, start:end]\n                    pred_win = model(seq_win, mask_win)\n                    \n                    pred_full[:, start:end] += pred_win\n                    count[:, start:end] += 1\n                \n                # Handle positions with zero count (rare edge case)\n                zero_mask = count == 0\n                if zero_mask.any():\n                    # For simplicity, we can leave them as 0 or fill nearest later\n                    # Here just clamp to avoid divide by zero\n                    pass\n                \n                # Safe division\n                pred = pred_full / count.clamp(min=1).unsqueeze(-1)\n            \n            # Extra safety: clip extreme values in predictions\n            pred = torch.nan_to_num(pred, nan=0.0, posinf=1e5, neginf=-1e5)\n            \n            loss = masked_mse_loss(pred, coords, mask)\n            \n            if torch.isnan(loss) or torch.isinf(loss):\n                print(\"NaN/Inf in val after window! Skipping batch...\")\n                continue\n            \n            val_loss += loss.item()\n            val_batches += 1\n    \n    if val_batches > 0:\n        val_loss /= val_batches\n    else:\n        val_loss = float('nan')\n    \n    # Safe printing of val_loss\n    val_str = f\"{val_loss:.2f}\" if not torch.isnan(torch.tensor(val_loss)) else \"NaN\"\n    print(f\"Epoch {epoch+1} | Train Loss: {train_loss:.2f} | Val Loss: {val_str}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T14:50:39.120234Z","iopub.execute_input":"2026-01-20T14:50:39.120865Z","iopub.status.idle":"2026-01-20T14:57:04.468241Z","shell.execute_reply.started":"2026-01-20T14:50:39.120834Z","shell.execute_reply":"2026-01-20T14:57:04.467498Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.train()  # improved: train mode for dropout diversity (MC dropout)\nRNA_VOCAB = {\"A\": 0, \"U\": 1, \"G\": 2, \"C\": 3}\nsubmission_rows = []\ntotal = len(test_seq)\nWINDOW_SIZE = 512  # same as training\nwith torch.no_grad():\n    for idx, row in test_seq.iterrows():\n        seq = row[\"sequence\"]\n        L = len(seq)\n        seq_encoded = torch.tensor(\n            [RNA_VOCAB[x] for x in seq],\n            dtype=torch.long\n        ).to(device)\n        preds = []\n        for _ in range(5):\n            if L <= WINDOW_SIZE:\n                # short seq: direct\n                mask = torch.ones(L, dtype=torch.bool, device=device).unsqueeze(0)\n                seq_in = seq_encoded.unsqueeze(0)\n                coords = model(seq_in, mask)\n                coords = coords.squeeze(0).cpu().numpy()\n            else:\n                # improved: long seq ke liye sliding window (overlap) for better prediction\n                coords_full = np.zeros((L, 3))\n                count = np.zeros(L)\n                step = WINDOW_SIZE // 2\n                for start in range(0, L - WINDOW_SIZE + 1, step):\n                    end = start + WINDOW_SIZE\n                    seq_win = seq_encoded[start:end].unsqueeze(0)\n                    mask = torch.ones(WINDOW_SIZE, dtype=torch.bool, device=device).unsqueeze(0)\n                    coords_win = model(seq_win, mask).squeeze(0).cpu().numpy()\n                    coords_full[start:end] += coords_win\n                    count[start:end] += 1\n                # average\n                coords_full /= count[:, np.newaxis]\n                coords = coords_full\n            preds.append(coords)\n        for i, base in enumerate(seq):\n            out = {\n                \"ID\": row[\"target_id\"],\n                \"resname\": base,\n                \"resid\": i + 1\n            }\n            for k in range(5):\n                out[f\"x_{k+1}\"] = float(preds[k][i][0])\n                out[f\"y_{k+1}\"] = float(preds[k][i][1])\n                out[f\"z_{k+1}\"] = float(preds[k][i][2])\n            submission_rows.append(out)\n        print(f\"Done {idx+1}/{total}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-20T14:59:24.524046Z","iopub.execute_input":"2026-01-20T14:59:24.524929Z","iopub.status.idle":"2026-01-20T14:59:25.624631Z","shell.execute_reply.started":"2026-01-20T14:59:24.52489Z","shell.execute_reply":"2026-01-20T14:59:25.624068Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}