{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","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":[{"id":"ca684ac8-9819-4eb4-8cfd-3860ecb532e8","cell_type":"markdown","source":"# Stanford RNA 3D Folding — v7.2 (mask + debug + length-safe)\n\nThis version fixes two real issues discovered during debugging:\n\n1) **Missing coordinates (NaNs) in train_labels**\n   - We build a coordinate mask `cmask[L,5]` and compute loss only where labels exist.\n   - We do **NOT** replace NaNs with zeros (zeros are placeholders only).\n\n2) **Very long sequences (max_len_data can be ~125k)**\n   - Position embeddings are capped (`MAX_LEN_CAP`, default 4096).\n   - We **filter training/validation targets** to only those with `len(seq) <= MAX_LEN`.\n     This avoids `nn.Embedding` index-out-of-bounds on GPU.\n\n**Test-time**:\n- We keep all test targets.\n- For test sequences longer than `MAX_LEN`, we do simple **chunked prediction** (no overlap) so you can still submit.\n\nOutputs: `submission.csv` matching `sample_submission.csv`.\n\nTip: If you ever hit `device-side assert triggered`, restart the Kaggle session (GPU state becomes corrupted).\n","metadata":{}},{"id":"708283f1-4125-40cd-a75c-2e409ddf6f72","cell_type":"code","source":"import os, random, time\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\n\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\n\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(\"torch:\", torch.__version__)\nprint(\"cuda available:\", torch.cuda.is_available())\nif torch.cuda.is_available():\n    print(\"gpu:\", torch.cuda.get_device_name(0))\nprint(\"device selected:\", device)\n\npd.set_option(\"display.max_columns\", 200)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T17:36:20.622456Z","iopub.execute_input":"2026-02-10T17:36:20.62285Z","iopub.status.idle":"2026-02-10T17:36:25.479764Z","shell.execute_reply.started":"2026-02-10T17:36:20.622815Z","shell.execute_reply":"2026-02-10T17:36:25.478967Z"}},"outputs":[],"execution_count":null},{"id":"ccfa2f31-d800-4ea7-997e-47945345f4f8","cell_type":"markdown","source":"## 1) Locate dataset + load CSVs\n","metadata":{}},{"id":"7faf5755-7517-4577-aebe-72f2f50cc6e1","cell_type":"code","source":"ROOT = \"/kaggle/input\"\n\ndef find_dir_with_file(root, filename):\n    for d in sorted(os.listdir(root)):\n        p = os.path.join(root, d)\n        if not os.path.isdir(p):\n            continue\n        if os.path.exists(os.path.join(p, filename)):\n            return p\n        for r, _, files in os.walk(p):\n            if filename in files:\n                return r\n    return None\n\nDATA_DIR = find_dir_with_file(ROOT, \"sample_submission.csv\")\nprint(\"DATA_DIR:\", DATA_DIR)\nassert DATA_DIR is not None, \"Could not find sample_submission.csv under /kaggle/input\"\n\npaths = {\n    \"train_sequences\": os.path.join(DATA_DIR, \"train_sequences.csv\"),\n    \"train_labels\": os.path.join(DATA_DIR, \"train_labels.csv\"),\n    \"validation_sequences\": os.path.join(DATA_DIR, \"validation_sequences.csv\"),\n    \"validation_labels\": os.path.join(DATA_DIR, \"validation_labels.csv\"),\n    \"test_sequences\": os.path.join(DATA_DIR, \"test_sequences.csv\"),\n    \"sample_submission\": os.path.join(DATA_DIR, \"sample_submission.csv\"),\n}\n\nfor k,v in paths.items():\n    print(k, \"exists:\", os.path.exists(v), \"->\", v)\nmissing = [k for k,v in paths.items() if not os.path.exists(v)]\nassert not missing, f\"Missing files: {missing}\"\n\ntrain_seq = pd.read_csv(paths[\"train_sequences\"])\ntrain_lab = pd.read_csv(paths[\"train_labels\"], low_memory=False)\nval_seq   = pd.read_csv(paths[\"validation_sequences\"])\nval_lab   = pd.read_csv(paths[\"validation_labels\"], low_memory=False)\ntest_seq  = pd.read_csv(paths[\"test_sequences\"])\nsub       = pd.read_csv(paths[\"sample_submission\"])\n\nprint(\"\\nShapes:\")\nprint(\"train_sequences:\", train_seq.shape)\nprint(\"train_labels   :\", train_lab.shape)\nprint(\"val_sequences  :\", val_seq.shape)\nprint(\"val_labels     :\", val_lab.shape)\nprint(\"test_sequences :\", test_seq.shape)\nprint(\"sample_sub     :\", sub.shape)\n\nprint(\"\\ntrain_sequences columns:\", train_seq.columns.tolist())\nprint(\"train_labels columns    :\", train_lab.columns.tolist())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T17:37:02.734558Z","iopub.execute_input":"2026-02-10T17:37:02.735433Z","iopub.status.idle":"2026-02-10T17:37:14.328813Z","shell.execute_reply.started":"2026-02-10T17:37:02.735391Z","shell.execute_reply":"2026-02-10T17:37:14.32812Z"}},"outputs":[],"execution_count":null},{"id":"5fea4787-8eec-4b86-a919-16ebca6b0345","cell_type":"markdown","source":"## 2) Column detection + NaN diagnostics + max length + model MAX_LEN cap\n","metadata":{}},{"id":"e6467a7b-1574-4443-a52c-bb82b86ee21d","cell_type":"code","source":"def guess_col(df, candidates):\n    for c in candidates:\n        if c in df.columns:\n            return c\n    return None\n\nTID_COL = guess_col(train_seq, [\"target_id\",\"rna_id\",\"sequence_id\",\"id\",\"target\"])\nSEQ_COL = guess_col(train_seq, [\"sequence\",\"seq\",\"rna_sequence\"])\nassert TID_COL is not None, f\"Can't infer target id column. Columns: {train_seq.columns.tolist()}\"\nassert SEQ_COL is not None, f\"Can't infer sequence column. Columns: {train_seq.columns.tolist()}\"\nprint(\"Using TID_COL:\", TID_COL)\nprint(\"Using SEQ_COL:\", SEQ_COL)\n\ndef extract_tid(x: str) -> str:\n    s = str(x)\n    return s.rsplit(\"_\",1)[0] if \"_\" in s else s\n\ntrain_lab[\"target_id\"] = train_lab[\"ID\"].map(extract_tid)\nval_lab[\"target_id\"]   = val_lab[\"ID\"].map(extract_tid)\n\ncols = [\"x_1\",\"y_1\",\"z_1\"]\nprint(\"NaNs in train labels:\", train_lab[cols].isna().sum().to_dict())\nprint(\"NaNs in val labels  :\", val_lab[cols].isna().sum().to_dict())\n\ntrain_map = dict(zip(train_seq[TID_COL].astype(str), train_seq[SEQ_COL].astype(str)))\nval_map   = dict(zip(val_seq[TID_COL].astype(str),   val_seq[SEQ_COL].astype(str)))\ntest_map  = dict(zip(test_seq[TID_COL].astype(str),  test_seq[SEQ_COL].astype(str)))\n\ndef max_len(df):\n    return int(df[SEQ_COL].astype(str).str.len().max())\n\nmax_len_data = max(max_len(train_seq), max_len(val_seq), max_len(test_seq))\nprint(\"max_len_data:\", max_len_data)\n\nMAX_LEN_CAP = 4096\nMAX_LEN = min(int(max_len_data + 8), MAX_LEN_CAP)\nprint(\"MAX_LEN_CAP       :\", MAX_LEN_CAP)\nprint(\"MAX_LEN used model:\", MAX_LEN)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T17:37:21.94086Z","iopub.execute_input":"2026-02-10T17:37:21.941711Z","iopub.status.idle":"2026-02-10T17:37:24.783223Z","shell.execute_reply.started":"2026-02-10T17:37:21.941679Z","shell.execute_reply":"2026-02-10T17:37:24.782367Z"}},"outputs":[],"execution_count":null},{"id":"4c29ffb0-152d-488c-88c9-21478dc42dcb","cell_type":"markdown","source":"## 3) Build dense coords + coordinate mask (ignore NaNs during training)\n","metadata":{}},{"id":"23ea31cb-3d14-49c3-aedb-d066328f45da","cell_type":"code","source":"for c in [\"resid\",\"x_1\",\"y_1\",\"z_1\",\"copy\"]:\n    assert c in train_lab.columns, f\"train_labels missing {c}\"\nfor c in [\"resid\",\"x_1\",\"y_1\",\"z_1\",\"copy\"]:\n    assert c in val_lab.columns, f\"validation_labels missing {c}\"\n\ndef build_coords_and_mask(labels_df: pd.DataFrame, seq_map: dict):\n    coords_out = {}\n    mask_out = {}\n    skipped = 0\n    total_points = 0\n    labeled_points = 0\n\n    for tid, g in labels_df.groupby(\"target_id\", sort=False):\n        tid = str(tid)\n        if tid not in seq_map:\n            skipped += 1\n            continue\n        L = len(seq_map[tid])\n        coords = np.zeros((L,5,3), dtype=np.float32)\n        cmask  = np.zeros((L,5), dtype=np.float32)\n\n        r = g[\"resid\"].astype(int).to_numpy() - 1\n        c = g[\"copy\"].astype(int).to_numpy() - 1\n        ok = (r >= 0) & (r < L) & (c >= 0) & (c < 5)\n\n        xyz = g.loc[ok, [\"x_1\",\"y_1\",\"z_1\"]].to_numpy(np.float32)\n        good = np.isfinite(xyz).all(axis=1)\n\n        r2 = r[ok][good]\n        c2 = c[ok][good]\n        xyz2 = xyz[good]\n\n        coords[r2, c2, :] = xyz2\n        cmask[r2, c2] = 1.0\n\n        coords_out[tid] = coords\n        mask_out[tid] = cmask\n\n        total_points += L * 5\n        labeled_points += int(cmask.sum())\n\n    frac = labeled_points / max(total_points, 1)\n    print(f\"build_coords_and_mask -> targets: {len(coords_out)} skipped: {skipped} labeled_fraction: {frac:.4f}\")\n    return coords_out, mask_out\n\ntrain_coords, train_cmask = build_coords_and_mask(train_lab, train_map)\nval_coords,   val_cmask   = build_coords_and_mask(val_lab,   val_map)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T17:37:30.34089Z","iopub.execute_input":"2026-02-10T17:37:30.341576Z","iopub.status.idle":"2026-02-10T17:37:36.985214Z","shell.execute_reply.started":"2026-02-10T17:37:30.341549Z","shell.execute_reply":"2026-02-10T17:37:36.984551Z"}},"outputs":[],"execution_count":null},{"id":"2674ccf5-badc-40bf-b801-ef495832c7da","cell_type":"markdown","source":"## 4) Dataset + collate (train/val length filtered to avoid pos-embedding OOB)\n","metadata":{}},{"id":"9c4da906-3af6-40f6-b29c-b6e95995a139","cell_type":"code","source":"NUC2I = {\"A\":0,\"C\":1,\"G\":2,\"U\":3}\nPAD_IDX = 4\nVOCAB = 5\n\ndef encode_seq(seq: str):\n    s = str(seq).upper()\n    return [NUC2I.get(ch, 0) for ch in s]\n\nclass RNADataset(Dataset):\n    def __init__(self, seq_map, coords_map=None, cmask_map=None, max_len=None):\n        ids = []\n        dropped = 0\n        for tid, seq in seq_map.items():\n            if coords_map is not None and tid not in coords_map:\n                continue\n            if max_len is not None and len(seq) > max_len:\n                dropped += 1\n                continue\n            ids.append(tid)\n        self.ids = ids\n        self.seq_map = seq_map\n        self.coords_map = coords_map\n        self.cmask_map = cmask_map\n        self.max_len = max_len\n        self.dropped = dropped\n\n    def __len__(self): \n        return len(self.ids)\n\n    def __getitem__(self, idx):\n        tid = self.ids[idx]\n        x = torch.tensor(encode_seq(self.seq_map[tid]), dtype=torch.long)\n        if self.coords_map is None:\n            return tid, x\n        y = torch.tensor(self.coords_map[tid], dtype=torch.float32)\n        cm = torch.tensor(self.cmask_map[tid], dtype=torch.float32)\n        return tid, x, y, cm\n\ndef collate_train(batch):\n    tids, xs, ys, cms = zip(*batch)\n    lens = torch.tensor([len(x) for x in xs], dtype=torch.long)\n    Lmax = int(lens.max())\n    B = len(xs)\n\n    xpad = torch.full((B, Lmax), PAD_IDX, dtype=torch.long)\n    ypad = torch.zeros((B, Lmax, 5, 3), dtype=torch.float32)\n    cmpad = torch.zeros((B, Lmax, 5), dtype=torch.float32)\n    pad_mask = torch.ones((B, Lmax), dtype=torch.bool)\n\n    for i,(x,y,cm) in enumerate(zip(xs,ys,cms)):\n        L = len(x)\n        xpad[i,:L] = x\n        ypad[i,:L] = y\n        cmpad[i,:L] = cm\n        pad_mask[i,:L] = False\n\n    assert int(xpad.max()) < VOCAB, f\"Found token id {int(xpad.max())} but vocab={VOCAB}\"\n    assert int(xpad.min()) >= 0, f\"Found negative token id {int(xpad.min())}\"\n    return list(tids), xpad, ypad, cmpad, pad_mask, lens\n\ndef collate_test(batch):\n    tids, xs = zip(*batch)\n    lens = torch.tensor([len(x) for x in xs], dtype=torch.long)\n    Lmax = int(lens.max())\n    B = len(xs)\n\n    xpad = torch.full((B, Lmax), PAD_IDX, dtype=torch.long)\n    pad_mask = torch.ones((B, Lmax), dtype=torch.bool)\n    for i,x in enumerate(xs):\n        L = len(x)\n        xpad[i,:L] = x\n        pad_mask[i,:L] = False\n\n    assert int(xpad.max()) < VOCAB, f\"Found token id {int(xpad.max())} but vocab={VOCAB}\"\n    assert int(xpad.min()) >= 0, f\"Found negative token id {int(xpad.min())}\"\n    return list(tids), xpad, pad_mask, lens\n\ntrain_ds = RNADataset(train_map, train_coords, train_cmask, max_len=MAX_LEN)\nval_ds   = RNADataset(val_map,   val_coords,   val_cmask,   max_len=MAX_LEN)\ntest_ds  = RNADataset(test_map,  None,         None,        max_len=None)\n\nprint(\"train_ds:\", len(train_ds), \"dropped_long:\", train_ds.dropped)\nprint(\"val_ds  :\", len(val_ds),   \"dropped_long:\", val_ds.dropped)\nprint(\"test_ds :\", len(test_ds))\n\ntrain_lens = [len(train_map[tid]) for tid in train_ds.ids]\nval_lens   = [len(val_map[tid]) for tid in val_ds.ids]\nprint(\"train lengths min/median/max:\", min(train_lens), int(np.median(train_lens)), max(train_lens))\nprint(\"val   lengths min/median/max:\", min(val_lens),   int(np.median(val_lens)),   max(val_lens))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T17:37:46.981413Z","iopub.execute_input":"2026-02-10T17:37:46.981711Z","iopub.status.idle":"2026-02-10T17:37:47.001548Z","shell.execute_reply.started":"2026-02-10T17:37:46.981687Z","shell.execute_reply":"2026-02-10T17:37:47.00075Z"}},"outputs":[],"execution_count":null},{"id":"54ae9ee2-4155-4702-becb-61f28794b12f","cell_type":"markdown","source":"## 5) Transformer model (explicit length guard)\n","metadata":{}},{"id":"6ed1682f-67dd-4970-9f8e-0914c3d9cda0","cell_type":"code","source":"class MiniTransformerCoords(nn.Module):\n    def __init__(self, vocab=5, d_model=128, nhead=8, num_layers=4, dim_ff=256, dropout=0.1, max_len=4096):\n        super().__init__()\n        self.tok = nn.Embedding(vocab, d_model, padding_idx=PAD_IDX)\n        self.pos = nn.Embedding(max_len, d_model)\n        self.max_len = max_len\n\n        enc_layer = nn.TransformerEncoderLayer(\n            d_model=d_model,\n            nhead=nhead,\n            dim_feedforward=dim_ff,\n            dropout=dropout,\n            batch_first=True,\n            activation=\"gelu\",\n        )\n        self.enc = nn.TransformerEncoder(enc_layer, num_layers=num_layers)\n        self.head = nn.Sequential(\n            nn.Linear(d_model, d_model),\n            nn.GELU(),\n            nn.Linear(d_model, 15)\n        )\n\n    def forward(self, x, src_key_padding_mask):\n        B, L = x.shape\n        if L > self.max_len:\n            raise ValueError(f\"Batch sequence length L={L} exceeds model MAX_LEN={self.max_len}. \"\n                             \"Filter lengths or implement chunking.\")\n        pos = torch.arange(L, device=x.device).unsqueeze(0).repeat(B, 1)\n        h = self.tok(x) + self.pos(pos)\n        h = self.enc(h, src_key_padding_mask=src_key_padding_mask)\n        return self.head(h).view(B, L, 5, 3)\n\nmodel = MiniTransformerCoords(max_len=MAX_LEN).to(device)\nn_params = sum(p.numel() for p in model.parameters())\nprint(\"model params:\", n_params, f\"({n_params/1e6:.2f} M)\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T17:37:54.760102Z","iopub.execute_input":"2026-02-10T17:37:54.760492Z","iopub.status.idle":"2026-02-10T17:37:55.058487Z","shell.execute_reply.started":"2026-02-10T17:37:54.760465Z","shell.execute_reply":"2026-02-10T17:37:55.057871Z"}},"outputs":[],"execution_count":null},{"id":"4cd1104a-cebf-4266-a09f-84866a8dfe24","cell_type":"markdown","source":"## 6) Training (masked MSE ignores missing coordinates)\n","metadata":{}},{"id":"10397241-b991-4cd3-a66e-344d20b66c8a","cell_type":"code","source":"from torch.optim import AdamW\n\nBATCH_SIZE = 2\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True, num_workers=0, collate_fn=collate_train)\nval_loader   = DataLoader(val_ds,   batch_size=BATCH_SIZE, shuffle=False, num_workers=0, collate_fn=collate_train)\n\nopt = AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4)\n\ndef masked_mse(pred, true, coord_mask, pad_mask):\n    valid_pos = (~pad_mask).float().unsqueeze(-1)\n    cm = coord_mask * valid_pos\n    cm = cm.unsqueeze(-1)\n    diff2 = (pred - true) ** 2\n    diff2 = diff2 * cm\n    denom = cm.sum() * 3 + 1e-8\n    return diff2.sum() / denom\n\ndef batch_labeled_fraction(coord_mask, pad_mask):\n    valid = (~pad_mask).sum().item() * 5\n    labeled = coord_mask.sum().item()\n    return labeled / max(valid, 1)\n\n@torch.no_grad()\ndef eval_epoch():\n    model.eval()\n    losses = []\n    for step, (tids, x, y, cm, pad_mask, lens) in enumerate(val_loader):\n        x = x.to(device); y = y.to(device); cm = cm.to(device); pad_mask = pad_mask.to(device)\n        pred = model(x, pad_mask)\n        loss = masked_mse(pred, y, cm, pad_mask).item()\n        losses.append(loss)\n        if step == 0:\n            print(\"[VAL] batch0 x:\", tuple(x.shape), \"maxlen:\", int(lens.max()), \"labeled_frac:\", f\"{batch_labeled_fraction(cm,pad_mask):.4f}\")\n    return float(np.mean(losses)) if losses else float(\"nan\")\n\ndef train_epoch():\n    model.train()\n    losses = []\n    t0 = time.time()\n    for step, (tids, x, y, cm, pad_mask, lens) in enumerate(train_loader):\n        x = x.to(device); y = y.to(device); cm = cm.to(device); pad_mask = pad_mask.to(device)\n        pred = model(x, pad_mask)\n        loss = masked_mse(pred, y, cm, pad_mask)\n        opt.zero_grad(set_to_none=True)\n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        opt.step()\n        losses.append(loss.item())\n\n        if step == 0:\n            print(\"[TRAIN] batch0 x:\", tuple(x.shape), \"maxlen:\", int(lens.max()), \"labeled_frac:\", f\"{batch_labeled_fraction(cm,pad_mask):.4f}\")\n        if step % 50 == 0:\n            msg = f\"  step {step:4d} loss {loss.item():.6f}\"\n            if torch.cuda.is_available():\n                mem = torch.cuda.memory_allocated() / 1e9\n                msg += f\" | gpu_mem_alloc {mem:.2f} GB\"\n            print(msg)\n\n    print(f\"train_epoch seconds: {time.time()-t0:.1f}\")\n    return float(np.mean(losses)) if losses else float(\"nan\")\n\nEPOCHS = 2\nfor ep in range(1, EPOCHS+1):\n    print(\"\\n=== EPOCH\", ep, \"===\")\n    tr = train_epoch()\n    va = eval_epoch()\n    print(f\"Epoch {ep}/{EPOCHS}  train_mse={tr:.6f}  val_mse={va:.6f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T17:38:00.919088Z","iopub.execute_input":"2026-02-10T17:38:00.919413Z","iopub.status.idle":"2026-02-10T17:42:43.677852Z","shell.execute_reply.started":"2026-02-10T17:38:00.919386Z","shell.execute_reply":"2026-02-10T17:42:43.677088Z"}},"outputs":[],"execution_count":null},{"id":"c8b66140-dc83-4eab-a027-08501262b399","cell_type":"markdown","source":"## 7) Predict test + build submission.csv (chunk long test sequences)\n","metadata":{}},{"id":"558907f4-6bac-4238-afb0-c2c0f3b21ea6","cell_type":"code","source":"test_loader = DataLoader(test_ds, batch_size=1, shuffle=False, num_workers=0, collate_fn=collate_test)\n\n@torch.no_grad()\ndef predict_one(tid, xpad, pad_mask, lens):\n    L = int(lens[0])\n    if L <= MAX_LEN:\n        x = xpad.to(device)\n        pm = pad_mask.to(device)\n        pred = model(x, pm).detach().cpu().numpy()[0, :L]\n        return pred.astype(np.float32)\n\n    full = np.zeros((L,5,3), dtype=np.float32)\n    counts = np.zeros((L,1,1), dtype=np.float32)\n\n    start = 0\n    while start < L:\n        end = min(start + MAX_LEN, L)\n        xchunk = xpad[:, start:end].to(device)\n        pmchunk = pad_mask[:, start:end].to(device)\n        pred = model(xchunk, pmchunk).detach().cpu().numpy()[0, :end-start]\n        full[start:end] += pred\n        counts[start:end] += 1.0\n        start = end\n\n    full = full / np.maximum(counts, 1.0)\n    return full.astype(np.float32)\n\n@torch.no_grad()\ndef predict_all():\n    model.eval()\n    preds = {}\n    for step, (tids, xpad, pad_mask, lens) in enumerate(test_loader):\n        tid = str(tids[0])\n        pred = predict_one(tid, xpad, pad_mask, lens)\n        preds[tid] = pred\n        if step == 0:\n            print(\"[TEST] first target:\", tid, \"L:\", int(lens[0]), \"pred shape:\", pred.shape)\n        if step % 200 == 0:\n            print(\"  predicted\", step, \"targets...\")\n    return preds\n\ntest_preds = predict_all()\nprint(\"predicted targets:\", len(test_preds))\n\ndef extract_target_id_from_row_id(id_str: str) -> str:\n    s = str(id_str)\n    return s.rsplit(\"_\", 1)[0] if \"_\" in s else s\n\nout = sub.copy()\nout[\"target_id\"] = out[\"ID\"].map(extract_target_id_from_row_id)\n\nfor tid, g in out.groupby(\"target_id\", sort=False):\n    tid = str(tid)\n    coords = test_preds[tid]\n    g2 = g.sort_values(\"resid\")\n    idx = g2.index.to_numpy()\n    resid0 = g2[\"resid\"].to_numpy().astype(int) - 1\n    resid0 = np.clip(resid0, 0, coords.shape[0]-1)\n    for k in range(5):\n        out.loc[idx, f\"x_{k+1}\"] = coords[resid0, k, 0]\n        out.loc[idx, f\"y_{k+1}\"] = coords[resid0, k, 1]\n        out.loc[idx, f\"z_{k+1}\"] = coords[resid0, k, 2]\n\nout = out.drop(columns=[\"target_id\"])\nout.to_csv(\"submission.csv\", index=False)\nprint(\"Saved submission.csv:\", out.shape)\ndisplay(out.head())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T17:42:53.516585Z","iopub.execute_input":"2026-02-10T17:42:53.517662Z","iopub.status.idle":"2026-02-10T17:42:54.14824Z","shell.execute_reply.started":"2026-02-10T17:42:53.517615Z","shell.execute_reply":"2026-02-10T17:42:54.147652Z"}},"outputs":[],"execution_count":null}]}