{"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":47317,"databundleVersionId":5799376,"sourceType":"competition"},{"sourceId":14786434,"sourceType":"datasetVersion","datasetId":9452824}],"dockerImageVersionId":31259,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"%%writefile train.py\nimport os\nimport time\nimport argparse\nimport random\nimport numpy as np\n\nimport torch\nfrom torch.utils.data import DataLoader\n\nfrom vesuvius_data_2 import VesuviusDatasetConfig, VesuviusSegPatchDataset, debug_print_batch_stats_once\nfrom unetr import VideoMAEUNETR2D, load_videomae_encoder_from_mae_ckpt, masked_bce_dice_loss, logits_stats\n\ndef set_seed(seed: int = 42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.benchmark = True\n\ndef make_loaders(args):\n    train_cfg = VesuviusDatasetConfig(\n        data_root=args.data_root,\n        split=\"train\",\n        fragment_ids=tuple(args.train_ids),\n        tile_size=args.tile_size,\n        stride=args.stride,\n        num_frames=args.num_frames,\n        depth_mode=args.depth_mode,\n        clip_min=0.0,\n        clip_max=200.0,\n        pos_ratio=0.5,                # 50/50\n        pos_tile_min_frac=args.pos_tile_min_frac,\n        valid_tile_min_frac=args.valid_tile_min_frac,\n        repeat=args.repeat,\n    )\n    val_cfg = VesuviusDatasetConfig(\n        data_root=args.data_root,\n        split=\"train\",\n        fragment_ids=tuple(args.val_ids),\n        tile_size=args.tile_size,\n        stride=args.stride,\n        num_frames=args.num_frames,\n        depth_mode=\"center_contig\",\n        clip_min=0.0,\n        clip_max=200.0,\n        pos_ratio=0.0,                # not used in val\n        pos_tile_min_frac=args.pos_tile_min_frac,\n        valid_tile_min_frac=args.valid_tile_min_frac,\n        repeat=1,\n    )\n\n    train_ds = VesuviusSegPatchDataset(train_cfg, is_train=True)\n    val_ds = VesuviusSegPatchDataset(val_cfg, is_train=False)\n\n    train_loader = DataLoader(\n        train_ds,\n        batch_size=args.batch_size,\n        shuffle=True,\n        num_workers=args.num_workers,\n        pin_memory=True,\n        drop_last=True,\n        persistent_workers=(args.num_workers > 0),\n    )\n    val_loader = DataLoader(\n        val_ds,\n        batch_size=args.batch_size,\n        shuffle=False,\n        num_workers=args.num_workers,\n        pin_memory=True,\n        drop_last=False,\n        persistent_workers=(args.num_workers > 0),\n    )\n    return train_loader, val_loader\n\ndef evaluate(model, loader, device, args):\n    model.eval()\n    losses = []\n    with torch.no_grad():\n        for step, (x, y, _) in enumerate(loader):\n            x = x.to(device, non_blocking=True)\n            y = y.to(device, non_blocking=True)\n\n            logits = model(x)\n            loss, _, _ = masked_bce_dice_loss(\n                logits, y,\n                pos_weight=args.pos_weight,\n                bce_weight=args.bce_weight,\n                dice_weight=args.dice_weight,\n            )\n            if torch.isfinite(loss):\n                losses.append(float(loss.detach().cpu()))\n    return float(np.mean(losses)) if len(losses) else float(\"nan\")\n\ndef train_one_epoch(model, loader, optimizer, device, args, epoch, scaler=None):\n    model.train()\n    t0 = time.time()\n\n    running = []\n    skip_steps = 0\n    grad_bad_steps = 0\n\n    for step, (x, y, _) in enumerate(loader):\n        x = x.to(device, non_blocking=True)\n        y = y.to(device, non_blocking=True)\n\n        debug_print_batch_stats_once(\"train\", x, y)\n\n        optimizer.zero_grad(set_to_none=True)\n\n        # forward (we already force encoder fp32 inside model)\n        logits = model(x)\n\n        # check logits finite first (hard gate)\n        if not torch.isfinite(logits).all():\n            skip_steps += 1\n            continue\n\n        loss, bce, dice = masked_bce_dice_loss(\n            logits, y,\n            pos_weight=args.pos_weight,\n            bce_weight=args.bce_weight,\n            dice_weight=args.dice_weight,\n        )\n\n        if not torch.isfinite(loss):\n            skip_steps += 1\n            continue\n\n        loss.backward()\n\n        # grad clip + finite check\n        total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), args.max_grad_norm)\n        if not torch.isfinite(total_norm):\n            grad_bad_steps += 1\n            optimizer.zero_grad(set_to_none=True)\n            continue\n\n        optimizer.step()\n\n        running.append(float(loss.detach().cpu()))\n\n        if (step + 1) % args.log_every == 0:\n            stats = logits_stats(logits)\n            lr0 = optimizer.param_groups[0][\"lr\"]\n            msg = (\n                f\"Epoch {epoch:02d} | step {step+1:04d}/{len(loader)} \"\n                f\"| loss={np.mean(running[-args.log_every:]):.4f} \"\n                f\"(bce={float(bce):.4f}, dice={float(dice):.4f}) \"\n                f\"| grad_norm={float(total_norm):.3f} | lr={lr0:.2e} \"\n                f\"| logits(mean={stats['mean']:.3f}, std={stats['std']:.3f}, min={stats['min']:.2f}, max={stats['max']:.2f}) \"\n                f\"| skip={skip_steps} grad_bad={grad_bad_steps}\"\n            )\n            print(msg)\n\n    dt = time.time() - t0\n    train_loss = float(np.mean(running)) if len(running) else float(\"nan\")\n    print(f\"[Epoch {epoch:02d}] train_loss={train_loss:.6f}  skip_steps={skip_steps}  grad_bad_steps={grad_bad_steps}  time={dt:.1f}s\")\n    return train_loss, skip_steps, grad_bad_steps\n\ndef freeze_encoder(model: VideoMAEUNETR2D, freeze: bool = True):\n    for p in model.encoder.parameters():\n        p.requires_grad = not freeze\n\ndef main():\n    ap = argparse.ArgumentParser()\n    ap.add_argument(\"--data_root\", type=str, default=\"/kaggle/input/vesuvius-challenge-ink-detection\")\n    ap.add_argument(\"--train_ids\", nargs=\"+\", default=[\"1\", \"2\", \"3\"])\n    ap.add_argument(\"--val_ids\", nargs=\"+\", default=[\"1\"])\n\n    ap.add_argument(\"--tile_size\", type=int, default=64)\n    ap.add_argument(\"--stride\", type=int, default=64)\n    ap.add_argument(\"--num_frames\", type=int, default=24)\n    ap.add_argument(\"--depth_mode\", type=str, default=\"rand_contig\")\n\n    ap.add_argument(\"--batch_size\", type=int, default=16)\n    ap.add_argument(\"--num_workers\", type=int, default=2)\n\n    ap.add_argument(\"--epochs\", type=int, default=8)\n    ap.add_argument(\"--lr\", type=float, default=1e-4)\n    ap.add_argument(\"--encoder_lr\", type=float, default=1e-5)\n    ap.add_argument(\"--weight_decay\", type=float, default=1e-2)\n\n    ap.add_argument(\"--pos_tile_min_frac\", type=float, default=0.01)\n    ap.add_argument(\"--valid_tile_min_frac\", type=float, default=0.5)\n    ap.add_argument(\"--repeat\", type=int, default=1)\n\n    ap.add_argument(\"--pos_weight\", type=float, default=10.0)  # 50/50 下建议 5~20\n    ap.add_argument(\"--bce_weight\", type=float, default=0.5)\n    ap.add_argument(\"--dice_weight\", type=float, default=0.5)\n\n    ap.add_argument(\"--max_grad_norm\", type=float, default=1.0)\n    ap.add_argument(\"--log_every\", type=int, default=500)\n\n    ap.add_argument(\"--seed\", type=int, default=42)\n    ap.add_argument(\"--mae_ckpt\", type=str, default=\"/kaggle/working/mae_outputs/best_mae.pt\")\n    ap.add_argument(\"--out_dir\", type=str, default=\"/kaggle/working/seg_outputs_run\")\n    ap.add_argument(\"--freeze_encoder_epochs\", type=int, default=1)\n\n    args = ap.parse_args()\n\n    set_seed(args.seed)\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    os.makedirs(args.out_dir, exist_ok=True)\n    print(\"Device:\", device)\n\n    train_loader, val_loader = make_loaders(args)\n\n    model = VideoMAEUNETR2D(tile_size=args.tile_size, num_frames=args.num_frames).to(device)\n\n    if args.mae_ckpt and os.path.exists(args.mae_ckpt):\n        print(\"Loading MAE encoder from:\", args.mae_ckpt)\n        load_videomae_encoder_from_mae_ckpt(model.encoder, args.mae_ckpt)\n\n    # param groups: encoder lr smaller\n    enc_params = [p for p in model.encoder.parameters() if p.requires_grad]\n    dec_params = [p for n, p in model.named_parameters() if (not n.startswith(\"encoder.\")) and p.requires_grad]\n\n    # (start frozen optionally)\n    if args.freeze_encoder_epochs > 0:\n        freeze_encoder(model, True)\n        enc_params = []  # frozen => empty\n        print(f\"[Warmup] Freeze encoder for {args.freeze_encoder_epochs} epoch(s)\")\n\n    optimizer = torch.optim.AdamW(\n        [\n            {\"params\": dec_params, \"lr\": args.lr},\n            {\"params\": enc_params, \"lr\": args.encoder_lr},\n        ],\n        weight_decay=args.weight_decay,\n    )\n\n    best_val = float(\"inf\")\n\n    # print val batch stats once\n    for xb, yb, _ in val_loader:\n        debug_print_batch_stats_once(\"val\", xb, yb)\n        break\n\n    for epoch in range(1, args.epochs + 1):\n        # unfreeze after warmup\n        if (epoch == args.freeze_encoder_epochs + 1) and (args.freeze_encoder_epochs > 0):\n            freeze_encoder(model, False)\n            # rebuild optimizer with encoder params\n            enc_params = [p for p in model.encoder.parameters() if p.requires_grad]\n            dec_params = [p for n, p in model.named_parameters() if (not n.startswith(\"encoder.\")) and p.requires_grad]\n            optimizer = torch.optim.AdamW(\n                [\n                    {\"params\": dec_params, \"lr\": args.lr},\n                    {\"params\": enc_params, \"lr\": args.encoder_lr},\n                ],\n                weight_decay=args.weight_decay,\n            )\n            print(\"[Warmup] Encoder unfrozen. Optimizer rebuilt.\")\n\n        train_loss, skip_steps, grad_bad_steps = train_one_epoch(model, train_loader, optimizer, device, args, epoch)\n        val_loss = evaluate(model, val_loader, device, args)\n        print(f\"Epoch {epoch:02d}/{args.epochs} | train_loss={train_loss:.6f} | val_loss={val_loss:.6f} | skip={skip_steps} | grad_bad={grad_bad_steps}\")\n\n        # save best\n        if np.isfinite(val_loss) and val_loss < best_val:\n            best_val = val_loss\n            ckpt_path = os.path.join(args.out_dir, \"best.pt\")\n            torch.save({\"model\": model.state_dict(), \"epoch\": epoch, \"val_loss\": val_loss}, ckpt_path)\n            print(\"  [Saved] best ->\", ckpt_path)\n\nif __name__ == \"__main__\":\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T10:04:10.901492Z","iopub.execute_input":"2026-02-10T10:04:10.902160Z","iopub.status.idle":"2026-02-10T10:04:10.910622Z","shell.execute_reply.started":"2026-02-10T10:04:10.902136Z","shell.execute_reply":"2026-02-10T10:04:10.910067Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile vesuvius_data_1.py\n# vesuvius_data.py\nimport os\nimport glob\nimport random\nfrom dataclasses import dataclass\nfrom typing import Dict, List, Tuple, Optional\n\nimport cv2\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset\n\ntry:\n    import tifffile\nexcept Exception:\n    tifffile = None\n\n\nIGNORE_INDEX = 127  # keep consistent with your losses\n\n\n# ---- per-process cache (DataLoader workers each has its own process) ----\n_MEMMAP_CACHE: Dict[str, np.ndarray] = {}\n\n\ndef _read_slice_memmap(path: str) -> np.ndarray:\n    arr = _MEMMAP_CACHE.get(path, None)\n    if arr is not None:\n        return arr\n    if tifffile is not None:\n        # Kaggle /kaggle/input is read-only, so MUST use mode=\"r\"\n        try:\n            arr = tifffile.memmap(path, mode=\"r\")\n        except TypeError:\n            # some tifffile versions don't expose mode; fallback to imread\n            arr = tifffile.imread(path)\n        except PermissionError:\n            # fallback to normal read if memmap fails\n            arr = tifffile.imread(path)\n    else:\n        arr = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    _MEMMAP_CACHE[path] = arr\n    return arr\n\n\n\ndef _load_png_gray(path: str) -> np.ndarray:\n    img = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    if img is None:\n        raise FileNotFoundError(path)\n    return img\n\n\ndef _normalize_volume(x: np.ndarray, clip_min=0.0, clip_max=200.0) -> np.ndarray:\n    # follow your existing normalization style (clip then scale to [-1, 1])\n    x = np.clip(x, clip_min, clip_max)\n    x = x / 255.0\n    x = (x - 0.5) / 0.5\n    return x.astype(np.float32)\n\n\ndef sample_depth_indices(total_slices: int, num_frames: int, mode: str) -> List[int]:\n    \"\"\"\n    total_slices: usually 65 (00..64)\n    num_frames: e.g. 24\n    mode:\n      - rand_contig: random consecutive window\n      - center_contig: fixed center window\n      - odd_subsample: pick 24 roughly-uniform indices from odd slices\n      - rand_stride2: random start + stride 2 (if feasible), else fall back to rand_contig\n    \"\"\"\n    if num_frames > total_slices:\n        raise ValueError(f\"num_frames={num_frames} > total_slices={total_slices}\")\n\n    if mode == \"rand_contig\":\n        s = random.randint(0, total_slices - num_frames)\n        return list(range(s, s + num_frames))\n\n    if mode == \"center_contig\":\n        s = (total_slices - num_frames) // 2\n        return list(range(s, s + num_frames))\n\n    if mode == \"odd_subsample\":\n        odds = list(range(1, total_slices, 2))  # 1,3,5,...\n        if len(odds) >= num_frames:\n            # uniform pick num_frames from odds\n            idx = np.linspace(0, len(odds) - 1, num_frames).round().astype(int)\n            return [odds[i] for i in idx]\n        # fallback\n        return sample_depth_indices(total_slices, num_frames, \"rand_contig\")\n\n    if mode == \"rand_stride2\":\n        stride = 2\n        needed = 1 + (num_frames - 1) * stride\n        if needed <= total_slices:\n            s = random.randint(0, total_slices - needed)\n            return [s + i * stride for i in range(num_frames)]\n        return sample_depth_indices(total_slices, num_frames, \"rand_contig\")\n\n    raise ValueError(f\"Unknown depth mode: {mode}\")\n\n\ndef build_tile_coords(roi_mask: np.ndarray, tile_size: int, stride: int, min_roi_frac: float = 0.05) -> List[Tuple[int, int]]:\n    \"\"\"\n    roi_mask: 0/255 or 0/1\n    returns list of (x, y) upper-left\n    min_roi_frac: tile must have at least this fraction of ROI pixels\n    \"\"\"\n    h, w = roi_mask.shape\n    roi = (roi_mask > 0).astype(np.uint8)\n\n    coords = []\n    tile_area = tile_size * tile_size\n    thr = int(tile_area * min_roi_frac)\n\n    # integral image for fast sum\n    integ = cv2.integral(roi)\n\n    def rect_sum(x1, y1, x2, y2):\n        # sum over [y1:y2, x1:x2]\n        return integ[y2, x2] - integ[y1, x2] - integ[y2, x1] + integ[y1, x1]\n\n    for y in range(0, h - tile_size + 1, stride):\n        y2 = y + tile_size\n        for x in range(0, w - tile_size + 1, stride):\n            x2 = x + tile_size\n            s = rect_sum(x, y, x2, y2)\n            if s >= thr:\n                coords.append((x, y))\n    return coords\n\n\n@dataclass\nclass VesuviusDatasetConfig:\n    data_root: str = \"/kaggle/input/vesuvius-challenge-ink-detection\"\n    split: str = \"train\"  # train/test\n    fragment_ids: Tuple[str, ...] = (\"1\", \"2\", \"3\")\n\n    tile_size: int = 64\n    stride: int = 64\n\n    # depth\n    num_frames: int = 24\n    depth_mode: str = \"rand_contig\"\n    total_slices: int = 65  # expect 00..64\n\n    # normalization\n    clip_min: float = 0.0\n    clip_max: float = 200.0\n\n    # training sampling\n    repeat: int = 1\n    pos_prob: float = 0.5   # for seg train: probability sampling from ink-positive tiles\n    min_roi_frac: float = 0.05\n\n\nclass VesuviusMAEPatchDataset(Dataset):\n    \"\"\"\n    Returns only video tensor: (T, 1, H, W) for MAE pretraining.\n    \"\"\"\n    def __init__(self, cfg: VesuviusDatasetConfig, is_train: bool = True):\n        self.cfg = cfg\n        self.is_train = is_train\n\n        self.fragments = []\n        self.slice_paths: Dict[str, List[str]] = {}\n        self.roi_masks: Dict[str, np.ndarray] = {}\n        self.coords_all: List[Tuple[str, int, int]] = []\n\n        for fid in cfg.fragment_ids:\n            fid = str(fid).replace(\"Frag\", \"\")\n            base = os.path.join(cfg.data_root, cfg.split, fid)\n            vol_dir = os.path.join(base, \"surface_volume\")\n            mask_path = os.path.join(base, \"mask.png\")\n\n            paths = sorted(glob.glob(os.path.join(vol_dir, \"*.tif\")))\n            if len(paths) == 0:\n                raise FileNotFoundError(vol_dir)\n            self.slice_paths[fid] = paths\n\n            roi = _load_png_gray(mask_path)\n            self.roi_masks[fid] = roi\n\n            coords = build_tile_coords(roi, cfg.tile_size, cfg.stride, cfg.min_roi_frac)\n            for (x, y) in coords:\n                self.coords_all.append((fid, x, y))\n\n        if len(self.coords_all) == 0:\n            raise RuntimeError(\"No tiles found. Check mask.png / tile_size / stride.\")\n\n        self._len = len(self.coords_all) * max(1, cfg.repeat)\n\n    def __len__(self):\n        return self._len\n\n    def __getitem__(self, idx: int):\n        fid, x, y = self.coords_all[idx % len(self.coords_all)]\n        paths = self.slice_paths[fid]\n        total = len(paths)\n\n        depth_mode = self.cfg.depth_mode if self.is_train else \"center_contig\"\n        z_idx = sample_depth_indices(total, self.cfg.num_frames, depth_mode)\n\n        # read patch stack: (H, W, T)\n        tile = np.empty((self.cfg.tile_size, self.cfg.tile_size, len(z_idx)), dtype=np.float32)\n        for i, z in enumerate(z_idx):\n            arr = _read_slice_memmap(paths[z])\n            patch = np.asarray(arr[y:y+self.cfg.tile_size, x:x+self.cfg.tile_size], dtype=np.float32)\n            tile[..., i] = patch\n\n        tile = _normalize_volume(tile, self.cfg.clip_min, self.cfg.clip_max)  # (H,W,T)\n        # (T,1,H,W)\n        video = torch.from_numpy(tile).permute(2, 0, 1).unsqueeze(1)\n        return video\n\n\nclass VesuviusSegPatchDataset(Dataset):\n    \"\"\"\n    Returns (video, mask, (fid, x, y))\n      video: (T,1,H,W)\n      mask:  (1,H,W) with {0,1,IGNORE_INDEX}\n    \"\"\"\n    def __init__(self, cfg: VesuviusDatasetConfig, is_train: bool = True):\n        self.cfg = cfg\n        self.is_train = is_train\n\n        self.slice_paths: Dict[str, List[str]] = {}\n        self.roi_masks: Dict[str, np.ndarray] = {}\n        self.ink_labels: Dict[str, np.ndarray] = {}\n\n        self.coords_all: List[Tuple[str, int, int]] = []\n        self.coords_pos: List[Tuple[str, int, int]] = []\n\n        for fid in cfg.fragment_ids:\n            fid = str(fid).replace(\"Frag\", \"\")\n            base = os.path.join(cfg.data_root, cfg.split, fid)\n            vol_dir = os.path.join(base, \"surface_volume\")\n            mask_path = os.path.join(base, \"mask.png\")\n            ink_path = os.path.join(base, \"inklabels.png\")\n\n            paths = sorted(glob.glob(os.path.join(vol_dir, \"*.tif\")))\n            if len(paths) == 0:\n                raise FileNotFoundError(vol_dir)\n            self.slice_paths[fid] = paths\n\n            roi = _load_png_gray(mask_path)\n            ink = _load_png_gray(ink_path)\n\n            self.roi_masks[fid] = roi\n            self.ink_labels[fid] = ink\n\n            coords = build_tile_coords(roi, cfg.tile_size, cfg.stride, cfg.min_roi_frac)\n            for (x, y) in coords:\n                self.coords_all.append((fid, x, y))\n                ink_patch = ink[y:y+cfg.tile_size, x:x+cfg.tile_size]\n                if (ink_patch > 0).mean() > 0.001:\n                    self.coords_pos.append((fid, x, y))\n\n        if len(self.coords_all) == 0:\n            raise RuntimeError(\"No tiles found. Check mask.png / tile_size / stride.\")\n\n        self._len = len(self.coords_all) * max(1, cfg.repeat)\n\n    def __len__(self):\n        return self._len\n\n    def __getitem__(self, idx: int):\n        # oversample positives in training\n        if self.is_train and len(self.coords_pos) > 0 and random.random() < self.cfg.pos_prob:\n            fid, x, y = random.choice(self.coords_pos)\n        else:\n            fid, x, y = self.coords_all[idx % len(self.coords_all)]\n\n        paths = self.slice_paths[fid]\n        total = len(paths)\n\n        depth_mode = self.cfg.depth_mode if self.is_train else \"center_contig\"\n        z_idx = sample_depth_indices(total, self.cfg.num_frames, depth_mode)\n\n        tile = np.empty((self.cfg.tile_size, self.cfg.tile_size, len(z_idx)), dtype=np.float32)\n        for i, z in enumerate(z_idx):\n            arr = _read_slice_memmap(paths[z])\n            patch = np.asarray(arr[y:y+self.cfg.tile_size, x:x+self.cfg.tile_size], dtype=np.float32)\n            tile[..., i] = patch\n\n        tile = _normalize_volume(tile, self.cfg.clip_min, self.cfg.clip_max)\n        video = torch.from_numpy(tile).permute(2, 0, 1).unsqueeze(1)  # (T,1,H,W)\n\n        ink = self.ink_labels[fid][y:y+self.cfg.tile_size, x:x+self.cfg.tile_size]\n        roi = self.roi_masks[fid][y:y+self.cfg.tile_size, x:x+self.cfg.tile_size]\n        mask = (ink > 0).astype(np.uint8)\n        mask[roi == 0] = IGNORE_INDEX\n        mask_t = torch.from_numpy(mask).unsqueeze(0).float()  # (1,H,W)\n\n        return video, mask_t, (fid, x, y)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile vesuvius_data_2.py\nimport os\nimport glob\nimport random\nfrom dataclasses import dataclass\nfrom typing import Dict, List, Tuple\n\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset\n\nimport tifffile\nimport cv2\n\nIGNORE_INDEX = 127\n\n# print-once control\n_DEBUG_PRINTED = set()\n\ndef debug_print_batch_stats_once(prefix: str, x: torch.Tensor, y: torch.Tensor, ignore_index: int = IGNORE_INDEX):\n    key = str(prefix)\n    if key in _DEBUG_PRINTED:\n        return\n    _DEBUG_PRINTED.add(key)\n\n    with torch.no_grad():\n        # x: (B,T,1,H,W), y: (B,1,H,W)\n        valid = (y != float(ignore_index)).float()\n        valid_frac = float(valid.mean().cpu().item())\n\n        pos = (y > 0.5).float()\n        pos_frac_all = float(pos.mean().cpu().item())\n        pos_frac_valid = float((pos * valid).sum().cpu().item() / (valid.sum().cpu().item() + 1e-6))\n\n        uniq = torch.unique(y).detach().cpu().tolist()\n\n    print(f\"[DebugBatch:{prefix}] x={tuple(x.shape)} y={tuple(y.shape)} uniq={uniq}\")\n    print(f\"[DebugBatch:{prefix}] valid_frac={valid_frac:.6f}  pos_frac_all={pos_frac_all:.6f}  pos_frac_valid={pos_frac_valid:.6f}\")\n\n\n# ---- per-process cache (each DataLoader worker is a process) ----\n_MEMMAP_CACHE: Dict[str, np.ndarray] = {}\n\ndef read_slice_memmap(path: str) -> np.ndarray:\n    \"\"\"\n    Kaggle input is read-only => memmap MUST be mode='r'\n    \"\"\"\n    arr = _MEMMAP_CACHE.get(path, None)\n    if arr is not None:\n        return arr\n    arr = tifffile.memmap(path, mode=\"r\")\n    _MEMMAP_CACHE[path] = arr\n    return arr\n\n\ndef load_png_gray(path: str) -> np.ndarray:\n    img = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    if img is None:\n        raise FileNotFoundError(path)\n    return img\n\n\ndef normalize_tile(tile: np.ndarray, clip_min=0.0, clip_max=200.0) -> np.ndarray:\n    \"\"\"\n    Follow your previous style: clip then map to [-1, 1].\n    tile: (H,W,T) float32\n    \"\"\"\n    tile = np.clip(tile, clip_min, clip_max)\n    tile = tile / 255.0\n    tile = (tile - 0.5) / 0.5\n    return tile.astype(np.float32)\n\n\ndef sample_depth_indices(total_slices: int, num_frames: int, mode: str) -> List[int]:\n    if num_frames > total_slices:\n        raise ValueError(f\"num_frames={num_frames} > total_slices={total_slices}\")\n\n    if mode == \"rand_contig\":\n        s = random.randint(0, total_slices - num_frames)\n        return list(range(s, s + num_frames))\n\n    if mode == \"center_contig\":\n        s = (total_slices - num_frames) // 2\n        return list(range(s, s + num_frames))\n\n    if mode == \"odd_subsample\":\n        odds = list(range(1, total_slices, 2))\n        if len(odds) >= num_frames:\n            idx = np.linspace(0, len(odds) - 1, num_frames).round().astype(int)\n            return [odds[i] for i in idx]\n        return sample_depth_indices(total_slices, num_frames, \"rand_contig\")\n\n    if mode == \"rand_stride2\":\n        stride = 2\n        needed = 1 + (num_frames - 1) * stride\n        if needed <= total_slices:\n            s = random.randint(0, total_slices - needed)\n            return [s + i * stride for i in range(num_frames)]\n        return sample_depth_indices(total_slices, num_frames, \"rand_contig\")\n\n    raise ValueError(f\"Unknown depth mode: {mode}\")\n\n\ndef integral_image_uint8(mask01: np.ndarray) -> np.ndarray:\n    \"\"\"\n    mask01: {0,1} uint8/bool\n    return ii int64 shape (H+1,W+1)\n    \"\"\"\n    m = (mask01 > 0).astype(np.int64)\n    ii = np.pad(m, ((1, 0), (1, 0)), constant_values=0)\n    ii = ii.cumsum(0).cumsum(1)\n    return ii\n\n\ndef rect_sum(ii: np.ndarray, x: int, y: int, sz: int) -> int:\n    \"\"\"\n    ii: (H+1,W+1) integral image\n    sum over [y:y+sz, x:x+sz]\n    \"\"\"\n    x2 = x + sz\n    y2 = y + sz\n    return int(ii[y2, x2] - ii[y, x2] - ii[y2, x] + ii[y, x])\n\n\n@dataclass\nclass VesuviusDatasetConfig:\n    data_root: str = \"/kaggle/input/vesuvius-challenge-ink-detection\"\n    split: str = \"train\"\n    fragment_ids: Tuple[str, ...] = (\"1\", \"2\", \"3\")\n\n    tile_size: int = 64\n    stride: int = 64\n\n    num_frames: int = 24\n    depth_mode: str = \"rand_contig\"\n    total_slices: int = 65\n\n    clip_min: float = 0.0\n    clip_max: float = 200.0\n\n    # --- segmentation sampling ---\n    pos_ratio: float = 0.5          # 50/50 sampling (train only)\n    pos_tile_min_frac: float = 0.01 # hard positive tile\n    valid_tile_min_frac: float = 0.5\n    repeat: int = 1\n\n\nclass VesuviusSegPatchDataset(Dataset):\n    \"\"\"\n    Returns:\n      video: (T,1,H,W) float32\n      mask:  (1,H,W) float32 in {0,1,127}\n      meta:  (fid, x, y)\n    \"\"\"\n    def __init__(self, cfg: VesuviusDatasetConfig, is_train: bool = True):\n        super().__init__()\n        self.cfg = cfg\n        self.is_train = is_train\n\n        self.slice_paths: Dict[str, List[str]] = {}\n        self.ink_labels: Dict[str, np.ndarray] = {}\n        self.roi_masks: Dict[str, np.ndarray] = {}\n\n        self.coords_all: List[Tuple[str, int, int]] = []\n        self.coords_pos: List[Tuple[str, int, int]] = []\n        self.coords_neg: List[Tuple[str, int, int]] = []\n\n        ts = cfg.tile_size\n        tile_area = ts * ts\n        pos_thr = int(tile_area * cfg.pos_tile_min_frac)\n        valid_thr = int(tile_area * cfg.valid_tile_min_frac)\n\n        for raw_fid in cfg.fragment_ids:\n            fid = str(raw_fid).replace(\"Frag\", \"\")\n            base = os.path.join(cfg.data_root, cfg.split, fid)\n            vol_dir = os.path.join(base, \"surface_volume\")\n            mask_path = os.path.join(base, \"mask.png\")\n            ink_path = os.path.join(base, \"inklabels.png\")\n\n            paths = sorted(glob.glob(os.path.join(vol_dir, \"*.tif\")))\n            if len(paths) == 0:\n                raise FileNotFoundError(vol_dir)\n\n            roi = load_png_gray(mask_path)\n            ink = load_png_gray(ink_path)\n\n            self.slice_paths[fid] = paths\n            self.roi_masks[fid] = roi\n            self.ink_labels[fid] = ink\n\n            H, W = roi.shape\n            gx = (W - ts) // cfg.stride + 1\n            gy = (H - ts) // cfg.stride + 1\n            total_grid = gx * gy\n\n            print(f\"[Index] frag={fid} scanning tiles... grid={gx}x{gy} (~{total_grid}) pos_thr={pos_thr}/{tile_area} valid_thr={valid_thr}/{tile_area}\")\n\n            roi01 = (roi > 0).astype(np.uint8)\n            ink01 = ((ink > 0) & (roi > 0)).astype(np.uint8)  # only count ink inside ROI\n\n            ii_roi = integral_image_uint8(roi01)\n            ii_ink = integral_image_uint8(ink01)\n\n            for y in range(0, H - ts + 1, cfg.stride):\n                for x in range(0, W - ts + 1, cfg.stride):\n                    vcnt = rect_sum(ii_roi, x, y, ts)\n                    if vcnt < valid_thr:\n                        continue\n\n                    pcnt = rect_sum(ii_ink, x, y, ts)\n                    self.coords_all.append((fid, x, y))\n                    if pcnt >= pos_thr:\n                        self.coords_pos.append((fid, x, y))\n                    else:\n                        self.coords_neg.append((fid, x, y))\n\n            print(f\"[Index] frag={fid} done. tiles_all={len(self.coords_all)} tiles_pos={len(self.coords_pos)} (pos_tile_min_frac={cfg.pos_tile_min_frac}, pos_ratio={cfg.pos_ratio})\")\n\n        if len(self.coords_all) == 0:\n            raise RuntimeError(\"No tiles found. Check mask.png / tile_size / stride / valid_tile_min_frac.\")\n        if self.is_train and (len(self.coords_pos) == 0 or len(self.coords_neg) == 0):\n            print(\"[WARN] pos or neg tiles empty. You may need to adjust pos_tile_min_frac / valid_tile_min_frac.\")\n\n        self._len = len(self.coords_all) * max(1, int(cfg.repeat))\n\n    def __len__(self):\n        return self._len\n\n    def pick_coord(self, idx: int) -> Tuple[str, int, int]:\n        # 50/50 sampling only for train\n        if self.is_train and len(self.coords_pos) > 0 and len(self.coords_neg) > 0:\n            if random.random() < float(self.cfg.pos_ratio):\n                return random.choice(self.coords_pos)\n            return random.choice(self.coords_neg)\n        return self.coords_all[idx % len(self.coords_all)]\n\n    def __getitem__(self, idx: int):\n        fid, x, y = self.pick_coord(idx)\n        paths = self.slice_paths[fid]\n        total = len(paths)\n\n        depth_mode = self.cfg.depth_mode if self.is_train else \"center_contig\"\n        z_idx = sample_depth_indices(total, self.cfg.num_frames, depth_mode)\n\n        # tile: (H,W,T)\n        tile = np.empty((self.cfg.tile_size, self.cfg.tile_size, len(z_idx)), dtype=np.float32)\n        for i, z in enumerate(z_idx):\n            arr = read_slice_memmap(paths[z])\n            patch = np.asarray(arr[y:y+self.cfg.tile_size, x:x+self.cfg.tile_size], dtype=np.float32)\n            tile[..., i] = patch\n\n        tile = normalize_tile(tile, self.cfg.clip_min, self.cfg.clip_max)\n        video = torch.from_numpy(tile).permute(2, 0, 1).unsqueeze(1)  # (T,1,H,W)\n\n        ink = self.ink_labels[fid][y:y+self.cfg.tile_size, x:x+self.cfg.tile_size]\n        roi = self.roi_masks[fid][y:y+self.cfg.tile_size, x:x+self.cfg.tile_size]\n        mask = (ink > 0).astype(np.uint8)\n        mask[roi == 0] = IGNORE_INDEX\n        mask_t = torch.from_numpy(mask).unsqueeze(0).float()\n\n        return video, mask_t, (fid, x, y)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile unetr.py\nimport os\nimport math\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom transformers import VideoMAEConfig, VideoMAEModel\n\nIGNORE_INDEX = 127\n\ndef _strip_prefix(k: str) -> str:\n    for p in [\"model.\", \"net.\", \"encoder.\", \"videomae.\", \"module.\", \"model_state.\"]:\n        if k.startswith(p):\n            return k[len(p):]\n    return k\n\ndef _pick_state_dict(ckpt: dict):\n    if isinstance(ckpt, dict):\n        if \"state_dict\" in ckpt and isinstance(ckpt[\"state_dict\"], dict):\n            return ckpt[\"state_dict\"]\n        if \"model_state\" in ckpt and isinstance(ckpt[\"model_state\"], dict):\n            return ckpt[\"model_state\"]\n    return ckpt\n\n@torch.no_grad()\ndef load_videomae_encoder_from_mae_ckpt(encoder: VideoMAEModel, ckpt_path: str):\n    ckpt = torch.load(ckpt_path, map_location=\"cpu\", weights_only=False)\n    state = _pick_state_dict(ckpt)\n    if not isinstance(state, dict):\n        raise TypeError(f\"Checkpoint does not contain a state dict: {type(state)}\")\n\n    new_state = {}\n    for k, v in state.items():\n        if not torch.is_tensor(v):\n            continue\n        kk = _strip_prefix(k)\n\n        # drop obvious decoder heads / bridges\n        if any(s in kk for s in [\"decoder\", \"mask_token\", \"decoder_pos_embed\", \"encoder_to_decoder\"]):\n            continue\n\n        # normalize possible nesting\n        for p in [\"videomae.videomae.\", \"model.videomae.\", \"videomae.\"]:\n            if kk.startswith(p):\n                kk = kk[len(p):]\n                break\n\n        new_state[kk] = v\n\n    missing, unexpected = encoder.load_state_dict(new_state, strict=False)\n    print(f\"[load_videomae_encoder_from_mae_ckpt] loaded from {ckpt_path}\")\n    print(f\"  missing={len(missing)} unexpected={len(unexpected)}\")\n    if len(unexpected) > 0:\n        print(\"  unexpected (first 5):\", unexpected[:5])\n    return missing, unexpected\n\ndef _gn_groups(ch: int, max_groups: int = 8) -> int:\n    g = min(max_groups, ch)\n    while g > 1:\n        if ch % g == 0:\n            return g\n        g -= 1\n    return 1\n\nclass ConvGNAct(nn.Module):\n    def __init__(self, in_ch, out_ch, k=3, p=1, max_groups=8):\n        super().__init__()\n        g = _gn_groups(out_ch, max_groups)\n        self.net = nn.Sequential(\n            nn.Conv2d(in_ch, out_ch, k, padding=p, bias=False),\n            nn.GroupNorm(g, out_ch),\n            nn.GELU(),\n        )\n    def forward(self, x):\n        return self.net(x)\n\nclass UpBlock(nn.Module):\n    def __init__(self, in_ch, skip_ch, out_ch, max_groups=8):\n        super().__init__()\n        self.up = nn.ConvTranspose2d(in_ch, out_ch, kernel_size=2, stride=2)\n        self.conv1 = ConvGNAct(out_ch + skip_ch, out_ch, max_groups=max_groups)\n        self.conv2 = ConvGNAct(out_ch, out_ch, max_groups=max_groups)\n\n    def forward(self, x, skip):\n        x = self.up(x)\n        if skip is not None:\n            if skip.shape[-2:] != x.shape[-2:]:\n                skip = F.interpolate(skip, size=x.shape[-2:], mode=\"bilinear\", align_corners=False)\n            x = torch.cat([x, skip], dim=1)\n        x = self.conv2(self.conv1(x))\n        return x\n\nclass VideoMAEUNETR2D(nn.Module):\n    \"\"\"\n    Input:  (B, T, 1, H, W)\n    Output: (B, 1, H, W)\n    \"\"\"\n    def __init__(self, tile_size=64, num_frames=24, intermediate_size=3072, use_layers=(3, 6, 9, 12), ch=256, max_gn_groups=8):\n        super().__init__()\n        tubelet_size = 2 if (num_frames % 2 == 0) else 1\n\n        self.vcfg = VideoMAEConfig(\n            image_size=tile_size,\n            patch_size=16,\n            num_channels=1,\n            num_frames=num_frames,\n            tubelet_size=tubelet_size,\n            hidden_size=768,\n            num_hidden_layers=12,\n            num_attention_heads=12,\n            intermediate_size=intermediate_size,\n        )\n        self.encoder = VideoMAEModel(self.vcfg)\n        self.use_layers = tuple(use_layers)\n\n        self.patch = self.vcfg.patch_size\n        self.Hp = tile_size // self.patch\n        self.Wp = tile_size // self.patch\n        self.Tp = num_frames // tubelet_size\n        self.D = self.vcfg.hidden_size\n\n        self.proj = nn.ModuleDict({str(l): nn.Conv2d(self.D, ch, kernel_size=1) for l in self.use_layers})\n\n        self.up1 = UpBlock(ch, ch, ch, max_groups=max_gn_groups)\n        self.up2 = UpBlock(ch, ch, ch, max_groups=max_gn_groups)\n        self.up3 = UpBlock(ch, ch, ch, max_groups=max_gn_groups)\n        self.up4 = UpBlock(ch, 0,  ch, max_groups=max_gn_groups)\n        self.head = nn.Conv2d(ch, 1, kernel_size=1)\n\n    def tokens_to_2d(self, hs: torch.Tensor) -> torch.Tensor:\n        # hs: (B, L, D)\n        B, L, D = hs.shape\n        n_hw = self.Hp * self.Wp\n\n        # expected tokens (without cls)\n        expected = self.Tp * n_hw\n\n        if L == expected + 1:\n            x = hs[:, 1:, :]\n            Tp = self.Tp\n        elif L == expected:\n            x = hs\n            Tp = self.Tp\n        else:\n            # try infer Tp from L\n            if (L - 1) % n_hw == 0:\n                x = hs[:, 1:, :]\n                Tp = (L - 1) // n_hw\n            elif L % n_hw == 0:\n                x = hs\n                Tp = L // n_hw\n            else:\n                raise RuntimeError(f\"Cannot reshape tokens: hs={hs.shape}, Hp={self.Hp},Wp={self.Wp},Tp={self.Tp}\")\n\n        x = x.reshape(B, Tp, self.Hp, self.Wp, D)\n        x = x.mean(dim=1)                     # (B,Hp,Wp,D)\n        x = x.permute(0, 3, 1, 2).contiguous() # (B,D,Hp,Wp)\n        return x\n\n    def forward(self, video: torch.Tensor) -> torch.Tensor:\n        # --- Force encoder in fp32 for stability (even if outer AMP is enabled) ---\n        with torch.cuda.amp.autocast(enabled=False):\n            out = self.encoder(video.float(), output_hidden_states=True, return_dict=True)\n            hss = out.hidden_states  # len=13\n\n            feats = {}\n            for l in self.use_layers:\n                f2d = self.tokens_to_2d(hss[l])\n                feats[l] = self.proj[str(l)](f2d)\n\n        x = feats[self.use_layers[-1]]\n        x = self.up1(x, feats[self.use_layers[-2]])\n        x = self.up2(x, feats[self.use_layers[-3]])\n        x = self.up3(x, feats[self.use_layers[-4]])\n        x = self.up4(x, None)\n        return self.head(x)\n\ndef masked_bce_dice_loss(\n    logits: torch.Tensor,\n    y: torch.Tensor,\n    ignore_index: int = IGNORE_INDEX,\n    pos_weight: float = 10.0,\n    bce_weight: float = 0.5,\n    dice_weight: float = 0.5,\n    eps: float = 1e-6,\n):\n    \"\"\"\n    logits,y in any dtype -> compute in fp32\n    y: (B,1,H,W) with {0,1,127}\n    \"\"\"\n    logits = logits.float()\n    y = y.float()\n\n    valid = (y != float(ignore_index)).float()\n    y_bin = (y > 0.5).float() * valid\n\n    pw = torch.tensor([pos_weight], device=logits.device, dtype=torch.float32)\n\n    bce = F.binary_cross_entropy_with_logits(logits, y_bin, reduction=\"none\", pos_weight=pw)\n    bce = (bce * valid).sum() / (valid.sum() + eps)\n\n    p = torch.sigmoid(logits) * valid\n    inter = (p * y_bin).sum()\n    den = p.sum() + y_bin.sum()\n    dice = (2.0 * inter + eps) / (den + eps)\n    dice_loss = 1.0 - dice\n\n    loss = bce_weight * bce + dice_weight * dice_loss\n    return loss, bce.detach(), dice_loss.detach()\n\n@torch.no_grad()\ndef logits_stats(logits: torch.Tensor) -> dict:\n    t = logits.float()\n    return {\n        \"mean\": float(t.mean().cpu()),\n        \"std\": float(t.std().cpu()),\n        \"min\": float(t.min().cpu()),\n        \"max\": float(t.max().cpu()),\n    }\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile mae.py\n\nimport os\nimport csv\nimport math\nimport argparse\n\nimport numpy as np\nimport torch\nfrom torch.utils.data import DataLoader\nfrom tqdm.auto import tqdm\n\nfrom transformers import VideoMAEConfig, VideoMAEForPreTraining\n\n# You should already have this file in /kaggle/working\nfrom vesuvius_data_1 import VesuviusDatasetConfig, VesuviusMAEPatchDataset\n\n\ndef set_seed(seed: int = 42):\n    import random\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\n\ndef save_history_csv(path: str, rows):\n    os.makedirs(os.path.dirname(path), exist_ok=True)\n    with open(path, \"w\", newline=\"\") as f:\n        w = csv.writer(f)\n        w.writerow([\"epoch\", \"train_loss\", \"val_loss\", \"lr\"])\n        for r in rows:\n            w.writerow(r)\n\n\ndef build_lr_scheduler(optimizer, total_steps, warmup_steps, min_lr=1e-6):\n    \"\"\"\n    Linear warmup + Cosine decay to min_lr.\n    Returns a function step() that updates lr each optimizer step.\n    \"\"\"\n    base_lrs = [pg[\"lr\"] for pg in optimizer.param_groups]\n\n    def lr_at(step):\n        if step < warmup_steps:\n            scale = (step + 1) / max(1, warmup_steps)\n        else:\n            t = (step - warmup_steps) / max(1, total_steps - warmup_steps)\n            scale = 0.5 * (1.0 + math.cos(math.pi * t))\n        return scale\n\n    def step_fn(step):\n        scale = lr_at(step)\n        for i, pg in enumerate(optimizer.param_groups):\n            pg[\"lr\"] = min_lr + (base_lrs[i] - min_lr) * scale\n\n    return step_fn\n\n\ndef parse_args():\n    ap = argparse.ArgumentParser()\n\n    ap.add_argument(\"--data_root\", type=str, default=\"/kaggle/input/vesuvius-challenge-ink-detection\")\n    ap.add_argument(\"--train_ids\", nargs=\"+\", default=[\"1\", \"2\", \"3\"])\n    ap.add_argument(\"--valid_ids\", nargs=\"+\", default=[\"1\"])\n\n    ap.add_argument(\"--tile_size\", type=int, default=64)\n    ap.add_argument(\"--stride\", type=int, default=64)\n\n    ap.add_argument(\"--num_frames\", type=int, default=24)\n    ap.add_argument(\n        \"--depth_mode\", type=str, default=\"rand_contig\",\n        choices=[\"rand_contig\", \"center_contig\", \"odd_subsample\", \"rand_stride2\"]\n    )\n\n    ap.add_argument(\"--mask_ratio\", type=float, default=0.85)\n\n    ap.add_argument(\"--batch_size\", type=int, default=32)\n    ap.add_argument(\"--epochs\", type=int, default=20)\n\n    ap.add_argument(\"--lr\", type=float, default=1e-4)\n    ap.add_argument(\"--weight_decay\", type=float, default=0.05)\n    ap.add_argument(\"--warmup_epochs\", type=int, default=2)\n    ap.add_argument(\"--min_lr\", type=float, default=1e-6)\n\n    ap.add_argument(\"--repeat\", type=int, default=8)\n    ap.add_argument(\"--num_workers\", type=int, default=2)\n\n    ap.add_argument(\"--out_dir\", type=str, default=\"/kaggle/working/mae_outputs\")\n    ap.add_argument(\"--seed\", type=int, default=42)\n\n    # by default True (you can disable by passing --no_fp16 if you want)\n    ap.add_argument(\"--fp16\", action=\"store_true\", default=True)\n    return ap.parse_args()\n\n\n@torch.no_grad()\ndef make_bool_masked_pos(batch_size: int, num_patches: int, mask_ratio: float, device: torch.device):\n    \"\"\"\n    Create bool_masked_pos of shape (B, num_patches) where about mask_ratio patches are masked.\n    \"\"\"\n    num_mask = int(mask_ratio * num_patches)\n    num_mask = max(1, min(num_mask, num_patches - 1))  # safe range\n\n    bool_masked_pos = torch.zeros((batch_size, num_patches), dtype=torch.bool, device=device)\n    for i in range(batch_size):\n        idx = torch.randperm(num_patches, device=device)[:num_mask]\n        bool_masked_pos[i, idx] = True\n    return bool_masked_pos\n\n\ndef main():\n    args = parse_args()\n    set_seed(args.seed)\n    os.makedirs(args.out_dir, exist_ok=True)\n\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(\"device =\", device)\n\n    # ------------------------\n    # Dataset / Loader\n    # ------------------------\n    train_cfg = VesuviusDatasetConfig(\n        data_root=args.data_root,\n        split=\"train\",\n        fragment_ids=tuple(args.train_ids),\n        tile_size=args.tile_size,\n        stride=args.stride,\n        num_frames=args.num_frames,\n        depth_mode=args.depth_mode,\n        repeat=args.repeat,\n    )\n    valid_cfg = VesuviusDatasetConfig(\n        data_root=args.data_root,\n        split=\"train\",\n        fragment_ids=tuple(args.valid_ids),\n        tile_size=args.tile_size,\n        stride=args.tile_size,\n        num_frames=args.num_frames,\n        depth_mode=\"center_contig\",\n        repeat=1,\n    )\n    train_cfg.mask_ratio = args.mask_ratio\n    valid_cfg.mask_ratio = args.mask_ratio\n\n    train_ds = VesuviusMAEPatchDataset(train_cfg, is_train=True)\n    val_ds = VesuviusMAEPatchDataset(valid_cfg, is_train=False)\n\n    train_loader = DataLoader(\n        train_ds, batch_size=args.batch_size, shuffle=True,\n        num_workers=args.num_workers, pin_memory=True, drop_last=True\n    )\n    val_loader = DataLoader(\n        val_ds, batch_size=args.batch_size, shuffle=False,\n        num_workers=args.num_workers, pin_memory=True, drop_last=False\n    )\n\n    # ------------------------\n    # Model (IMPORTANT: config must match finetune encoder)\n    # ------------------------\n    tubelet_size = 2 if (args.num_frames % 2 == 0) else 1\n\n    vcfg = VideoMAEConfig(\n        image_size=args.tile_size,\n        patch_size=16,\n        num_channels=1,\n        num_frames=args.num_frames,\n        tubelet_size=tubelet_size,\n\n        hidden_size=768,\n        num_hidden_layers=12,\n        num_attention_heads=12,\n        intermediate_size=3072,\n\n        decoder_num_hidden_layers=4,\n        decoder_hidden_size=512,\n        decoder_num_attention_heads=8,\n        decoder_intermediate_size=2048,\n\n        norm_pix_loss=True,\n        mask_ratio=args.mask_ratio,\n    )\n\n    model = VideoMAEForPreTraining(vcfg).to(device)\n    model.train()\n\n    # ---- compute num_patches for bool_masked_pos ----\n    patch = vcfg.patch_size\n    tube = vcfg.tubelet_size\n    Hp = args.tile_size // patch\n    Wp = args.tile_size // patch\n    Tp = args.num_frames // tube\n    num_patches = Hp * Wp * Tp\n    print(f\"[info] patch_size={patch} tubelet_size={tube} Hp={Hp} Wp={Wp} Tp={Tp} num_patches={num_patches}\")\n\n    optimizer = torch.optim.AdamW(\n        model.parameters(), lr=args.lr, weight_decay=args.weight_decay, betas=(0.9, 0.95)\n    )\n\n    total_steps = args.epochs * len(train_loader)\n    warmup_steps = args.warmup_epochs * len(train_loader)\n    lr_step = build_lr_scheduler(\n        optimizer, total_steps=total_steps, warmup_steps=warmup_steps, min_lr=args.min_lr\n    )\n\n    # new AMP API (still works on kaggle)\n    use_amp = bool(args.fp16 and device.type == \"cuda\")\n    scaler = torch.amp.GradScaler(\"cuda\", enabled=use_amp)\n\n    history = []\n    best_val = float(\"inf\")\n    global_step = 0\n\n    history_path = os.path.join(args.out_dir, f\"training_history_videomae_{args.tile_size}_{args.num_frames}.csv\")\n    best_ckpt_path = os.path.join(args.out_dir, \"best_mae.pt\")\n    last_ckpt_path = os.path.join(args.out_dir, \"last_mae.pt\")\n\n    for epoch in range(1, args.epochs + 1):\n        # ------------------------\n        # Train\n        # ------------------------\n        model.train()\n        train_losses = []\n\n        pbar = tqdm(train_loader, desc=f\"[Train] epoch {epoch}/{args.epochs}\", leave=False)\n        for batch in pbar:\n            batch = batch.to(device, non_blocking=True)  # (B,T,1,H,W)\n\n            lr_step(global_step)\n            optimizer.zero_grad(set_to_none=True)\n\n            # create bool_masked_pos for this batch\n            B = batch.size(0)\n            bool_masked_pos = make_bool_masked_pos(B, num_patches, args.mask_ratio, device)\n\n            with torch.amp.autocast(\"cuda\", enabled=use_amp):\n                out = model(pixel_values=batch, bool_masked_pos=bool_masked_pos)\n                loss = out.loss\n\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n\n            train_losses.append(loss.item())\n            global_step += 1\n\n            pbar.set_postfix(loss=float(np.mean(train_losses)), lr=optimizer.param_groups[0][\"lr\"])\n\n        train_loss = float(np.mean(train_losses))\n\n        # ------------------------\n        # Val\n        # ------------------------\n        model.eval()\n        val_losses = []\n        with torch.no_grad():\n            pbar = tqdm(val_loader, desc=f\"[Val] epoch {epoch}/{args.epochs}\", leave=False)\n            for batch in pbar:\n                batch = batch.to(device, non_blocking=True)\n                B = batch.size(0)\n                bool_masked_pos = make_bool_masked_pos(B, num_patches, args.mask_ratio, device)\n\n                with torch.amp.autocast(\"cuda\", enabled=use_amp):\n                    out = model(pixel_values=batch, bool_masked_pos=bool_masked_pos)\n                    loss = out.loss\n\n                val_losses.append(loss.item())\n                pbar.set_postfix(val_loss=float(np.mean(val_losses)))\n\n        val_loss = float(np.mean(val_losses))\n        lr_now = optimizer.param_groups[0][\"lr\"]\n        print(f\"Epoch {epoch:02d} | train_loss={train_loss:.6f} | val_loss={val_loss:.6f} | lr={lr_now:.2e}\")\n\n        history.append([epoch, train_loss, val_loss, lr_now])\n        save_history_csv(history_path, history)\n\n        # save last\n        torch.save(\n            {\n                \"epoch\": epoch,\n                \"model_state\": model.state_dict(),\n                \"config\": vcfg.to_dict(),\n                \"args\": vars(args),\n            },\n            last_ckpt_path\n        )\n\n        # save best\n        if val_loss < best_val:\n            best_val = val_loss\n            torch.save(\n                {\n                    \"epoch\": epoch,\n                    \"model_state\": model.state_dict(),\n                    \"config\": vcfg.to_dict(),\n                    \"args\": vars(args),\n                },\n                best_ckpt_path\n            )\n            print(\"  -> saved BEST to\", best_ckpt_path)\n\n    print(\"Done.\")\n    print(\"Best ckpt:\", best_ckpt_path)\n    print(\"History csv:\", history_path)\n\n\nif __name__ == \"__main__\":\n    main()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python mae.py --data_root /kaggle/input/vesuvius-challenge-ink-detection --train_ids  2 3 --valid_ids 1 --tile_size 64 --stride 96 --num_frames 24 --depth_mode rand_contig --mask_ratio 0.85 --batch_size 32 --epochs 13 --lr 1e-4 --weight_decay 0.05 --warmup_epochs 2 --repeat 1 --out_dir /kaggle/working/mae_outputs\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T09:04:42.374043Z","iopub.execute_input":"2026-02-10T09:04:42.374622Z","iopub.status.idle":"2026-02-10T09:31:04.746390Z","shell.execute_reply.started":"2026-02-10T09:04:42.374595Z","shell.execute_reply":"2026-02-10T09:31:04.745672Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python train.py ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T10:04:19.279674Z","iopub.execute_input":"2026-02-10T10:04:19.279974Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile infer_submit.py\nimport os\nimport glob\nimport argparse\nfrom typing import List, Tuple, Optional, Dict, Any\nimport re\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\n\nimport cv2\nimport tifffile\nimport csv\nos.environ.setdefault(\"TRANSFORMERS_NO_TF\", \"1\")\nos.environ.setdefault(\"TRANSFORMERS_NO_FLAX\", \"1\")\nos.environ.setdefault(\"TRANSFORMERS_NO_JAX\", \"1\")\nos.environ.setdefault(\"TF_CPP_MIN_LOG_LEVEL\", \"3\")\n\ntry:\n    from tqdm.auto import tqdm\nexcept Exception:\n    def tqdm(x, **kwargs):\n        return x\n\n# ---- your model code ----\nfrom unetr import VideoMAEUNETR2D, load_videomae_encoder_from_mae_ckpt\n\n\ndef rle_encode(mask01: np.ndarray) -> str:\n    \"\"\"Fortran order RLE: transpose then flatten.\"\"\"\n    m = mask01.astype(np.uint8)\n    pixels = m.T.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    changes = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs = changes.copy()\n    runs[1::2] -= runs[::2]\n    return \" \".join(str(x) for x in runs)\n\n\ndef _list_tif_slices(frag_root: str) -> List[str]:\n    vol_dir = os.path.join(frag_root, \"surface_volume\")\n    if not os.path.isdir(vol_dir):\n        alt = os.path.join(frag_root, \"surface_volumn\")  # tolerate typo\n        if os.path.isdir(alt):\n            vol_dir = alt\n    paths = sorted(glob.glob(os.path.join(vol_dir, \"*.tif\")))\n    if len(paths) == 0:\n        raise FileNotFoundError(f\"No tif slices found in: {vol_dir}\")\n    return paths\ndef write_submission_csv(path: str, ids: List[str], rles: List[str]) -> None:\n    # 把所有空白（含 \\n \\r \\t）压成一个空格，彻底消灭“拆行”\n    clean = []\n    for s in rles:\n        s = \"\" if s is None else str(s)\n        s = re.sub(r\"\\s+\", \" \", s).strip()\n        clean.append(s)\n\n    with open(path, \"w\", newline=\"\", encoding=\"utf-8\") as f:\n        w = csv.writer(f, delimiter=\",\", quotechar='\"', quoting=csv.QUOTE_ALL)\n        w.writerow([\"Id\", \"Predicted\"])\n        for fid, rle in zip(ids, clean):\n            w.writerow([str(fid), rle])\n\n\ndef _choose_z_indices(total_slices: int, num_frames: int, mode: str, start_z: Optional[int]) -> List[int]:\n    if num_frames > total_slices:\n        raise ValueError(f\"num_frames={num_frames} > total_slices={total_slices}\")\n    if start_z is not None:\n        s = max(0, min(int(start_z), total_slices - num_frames))\n        return list(range(s, s + num_frames))\n    if mode == \"front_contig\":\n        return list(range(0, num_frames))\n    s = (total_slices - num_frames) // 2\n    return list(range(s, s + num_frames))\n\n\ndef load_mask_png(path: str) -> np.ndarray:\n    m = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    if m is None:\n        raise FileNotFoundError(path)\n    return (m > 0).astype(np.uint8)\n\n\ndef build_coords(mask01: np.ndarray, tile_size: int, stride: int) -> List[Tuple[int, int]]:\n    H, W = mask01.shape\n    ys = list(range(0, max(H - tile_size + 1, 1), stride))\n    xs = list(range(0, max(W - tile_size + 1, 1), stride))\n    if len(ys) == 0: ys = [0]\n    if len(xs) == 0: xs = [0]\n    y_last = max(H - tile_size, 0)\n    x_last = max(W - tile_size, 0)\n    if ys[-1] != y_last: ys.append(y_last)\n    if xs[-1] != x_last: xs.append(x_last)\n\n    coords: List[Tuple[int, int]] = []\n    for y in ys:\n        for x in xs:\n            if mask01[y:y+tile_size, x:x+tile_size].sum() > 0:\n                coords.append((x, y))\n    return coords\n\n\ndef normalize_patch_u16_to_model(patch_u16: np.ndarray, clip_min: float, clip_max: float) -> np.ndarray:\n    # match your train normalization: clip -> /255 -> (x-0.5)/0.5\n    x = patch_u16.astype(np.float32)\n    x = np.clip(x, clip_min, clip_max)\n    x = x / 255.0\n    x = (x - 0.5) / 0.5\n    return x.astype(np.float16)\n\n\ndef load_volume_slices_to_ram(slice_paths: List[str], z_idx: List[int]) -> np.ndarray:\n    imgs = [tifffile.imread(slice_paths[z]) for z in z_idx]\n    return np.stack(imgs, axis=0)  # (T,H,W)\n\n\nclass TestPatchDataset(Dataset):\n    def __init__(\n        self,\n        vol_u16: np.ndarray,      # (T,H,W)\n        coords: List[Tuple[int, int]],\n        tile_size: int,\n        clip_min: float,\n        clip_max: float,\n        cache_float16: bool = True,\n    ):\n        self.vol_u16 = vol_u16\n        self.coords = coords\n        self.tile_size = int(tile_size)\n        self.clip_min = float(clip_min)\n        self.clip_max = float(clip_max)\n\n        self.vol_f16 = None\n        if cache_float16:\n            try:\n                x = vol_u16.astype(np.float32)\n                x = np.clip(x, self.clip_min, self.clip_max)\n                x = x / 255.0\n                x = (x - 0.5) / 0.5\n                self.vol_f16 = x.astype(np.float16)\n            except MemoryError:\n                self.vol_f16 = None\n\n    def __len__(self):\n        return len(self.coords)\n\n    def __getitem__(self, idx: int):\n        x0, y0 = self.coords[idx]\n        ts = self.tile_size\n\n        if self.vol_f16 is not None:\n            patch = self.vol_f16[:, y0:y0+ts, x0:x0+ts]  # (T,ts,ts)\n        else:\n            patch_u16 = self.vol_u16[:, y0:y0+ts, x0:x0+ts]\n            patch = normalize_patch_u16_to_model(patch_u16, self.clip_min, self.clip_max)\n\n        x = torch.from_numpy(patch).unsqueeze(1)  # (T,1,ts,ts)\n        coord = torch.tensor([x0, y0], dtype=torch.int32)\n        return x, coord\n\n\ndef _strip_outer_prefix(k: str) -> str:\n    # do NOT strip \"encoder.\" because it is a real submodule in your UNETR\n    for p in (\"model.\", \"net.\", \"module.\", \"pl_module.\", \"lit_model.\", \"vmae.\"):\n        if k.startswith(p):\n            return k[len(p):]\n    return k\n\n\ndef find_ckpt_path(ckpt_dir_or_file: str) -> str:\n    \"\"\"\n    Support:\n      - a file path (xxx.pt / xxx.pth / xxx.ckpt)\n      - a directory containing those files\n    Preference: best > last > newest\n    \"\"\"\n    p = ckpt_dir_or_file.rstrip(\"/\")\n    if os.path.isfile(p) and p.lower().endswith((\".pt\", \".pth\", \".ckpt\")):\n        return p\n\n    if not os.path.isdir(p):\n        raise FileNotFoundError(f\"ckpt_dir not found: {p}\")\n\n    exts = (\"*.pt\", \"*.pth\", \"*.ckpt\")\n    cands = []\n    for ext in exts:\n        cands.extend(glob.glob(os.path.join(p, \"**\", ext), recursive=True))\n        cands.extend(glob.glob(os.path.join(p, ext)))\n\n    # dedup\n    cands = sorted(list(set(cands)))\n    if len(cands) == 0:\n        raise FileNotFoundError(f\"No .pt/.pth/.ckpt found under: {p}\")\n\n    def score(path: str) -> Tuple[int, int]:\n        name = os.path.basename(path).lower()\n        # higher is better\n        s = 0\n        if \"best\" in name: s += 30\n        if \"last\" in name: s += 20\n        if \"final\" in name: s += 10\n        # newer is better\n        mtime = int(os.path.getmtime(path))\n        return (s, mtime)\n\n    cands = sorted(cands, key=score, reverse=True)\n    return cands[0]\n\n\ndef _extract_state_and_hp(ckpt_obj: Any) -> Tuple[Dict[str, torch.Tensor], Dict[str, Any]]:\n    \"\"\"\n    Handle:\n      1) {\"state_dict\": ...}\n      2) pure state_dict\n      3) {\"model\": state_dict} / {\"net\": state_dict} (rare)\n    \"\"\"\n    hp = {}\n    if isinstance(ckpt_obj, dict):\n        if \"hyper_parameters\" in ckpt_obj and isinstance(ckpt_obj[\"hyper_parameters\"], dict):\n            hp = ckpt_obj[\"hyper_parameters\"]\n\n        for key in (\"state_dict\", \"model_state\", \"model\", \"net\"):\n            if key in ckpt_obj and isinstance(ckpt_obj[key], dict):\n                return ckpt_obj[key], hp\n\n        # if it already looks like a state dict\n        if all(isinstance(v, torch.Tensor) for v in ckpt_obj.values()):\n            return ckpt_obj, hp\n\n    raise TypeError(f\"Unrecognized checkpoint format: {type(ckpt_obj)}\")\n\n\ndef load_seg_model(ckpt_path: str, tile_size: int, num_frames: int, mae_ckpt: str):\n    ckpt = torch.load(ckpt_path, map_location=\"cpu\", weights_only=False)\n    state, hp = _extract_state_and_hp(ckpt)\n\n    tile_size = int(hp.get(\"tile_size\", tile_size))\n    num_frames = int(hp.get(\"num_frames\", num_frames))\n\n    net = VideoMAEUNETR2D(tile_size=tile_size, num_frames=num_frames)\n\n    # optional MAE fill: helpful if ckpt lacks encoder weights\n    if mae_ckpt and os.path.exists(mae_ckpt):\n        try:\n            load_videomae_encoder_from_mae_ckpt(net.encoder, mae_ckpt)\n        except Exception as e:\n            print(f\"[WARN] failed to load mae_ckpt='{mae_ckpt}': {e}\")\n\n    new_state = {}\n    for k, v in state.items():\n        if torch.is_tensor(v):\n            new_state[_strip_outer_prefix(k)] = v\n\n    missing, unexpected = net.load_state_dict(new_state, strict=False)\n    print(f\"[CKPT] {ckpt_path}\")\n    print(f\"  loaded with missing={len(missing)} unexpected={len(unexpected)}\")\n    if unexpected:\n        print(\"  unexpected (first 10):\", unexpected[:10])\n    if missing:\n        print(\"  missing (first 10):\", missing[:10])\n\n    return net, tile_size, num_frames\n\n\n@torch.no_grad()\ndef predict_fragment(\n    net: VideoMAEUNETR2D,\n    vol_u16: np.ndarray,     # (T,H,W)\n    mask01: np.ndarray,      # (H,W)\n    tile_size: int,\n    stride: int,\n    clip_min: float,\n    clip_max: float,\n    batch_size: int,\n    num_workers: int,\n    cache_float16: bool,\n    use_amp: bool,\n) -> np.ndarray:\n    device = next(net.parameters()).device\n    coords = build_coords(mask01, tile_size=tile_size, stride=stride)\n    print(f\"[Predict] HxW={mask01.shape} tiles={len(coords)} tile={tile_size} stride={stride}\")\n\n    ds = TestPatchDataset(\n        vol_u16=vol_u16,\n        coords=coords,\n        tile_size=tile_size,\n        clip_min=clip_min,\n        clip_max=clip_max,\n        cache_float16=cache_float16,\n    )\n    dl = DataLoader(\n        ds,\n        batch_size=batch_size,\n        shuffle=False,\n        num_workers=num_workers,\n        pin_memory=(device.type == \"cuda\"),\n        drop_last=False,\n    )\n\n    H, W = mask01.shape\n    pred = np.zeros((H, W), dtype=np.float32)\n    cnt = np.zeros((H, W), dtype=np.float32)\n\n    for xb, cb in tqdm(dl, desc=\"tiles\", leave=False):\n        xb = xb.to(device, non_blocking=True)  # (B,T,1,ts,ts)\n\n        if use_amp and device.type == \"cuda\":\n            with torch.cuda.amp.autocast(dtype=torch.float16):\n                logits = net(xb)\n        else:\n            logits = net(xb)\n\n        probs = torch.sigmoid(logits.float()).squeeze(1).cpu().numpy()\n        cb = cb.cpu().numpy()\n\n        for i in range(probs.shape[0]):\n            x0, y0 = int(cb[i, 0]), int(cb[i, 1])\n            pred[y0:y0+tile_size, x0:x0+tile_size] += probs[i]\n            cnt[y0:y0+tile_size, x0:x0+tile_size] += 1.0\n\n    pred = pred / np.maximum(cnt, 1.0)\n    pred = pred * mask01.astype(np.float32)\n    return pred\n\n\ndef parse_args(argv=None):\n    ap = argparse.ArgumentParser()\n    ap.add_argument(\"--data_root\", type=str, default=\"/kaggle/input/vesuvius-challenge-ink-detection\")\n    ap.add_argument(\"--ckpt_dir\", type=str, default=\"/kaggle/working/seg_outputs_run3\")\n    ap.add_argument(\"--mae_ckpt\", type=str, default=\"/kaggle/working/mae_outputs/best_mae.pt\")\n    ap.add_argument(\"--out_csv\", type=str, default=\"/kaggle/working/submission.csv\")\n\n    ap.add_argument(\"--tile_size\", type=int, default=64)\n    ap.add_argument(\"--stride\", type=int, default=64)\n    ap.add_argument(\"--num_frames\", type=int, default=24)\n    ap.add_argument(\"--depth_mode\", type=str, default=\"center_contig\", choices=[\"center_contig\", \"front_contig\"])\n    ap.add_argument(\"--start_z\", type=int, default=-1)\n\n    ap.add_argument(\"--clip_min\", type=float, default=0.0)\n    ap.add_argument(\"--clip_max\", type=float, default=200.0)\n\n    ap.add_argument(\"--batch_size\", type=int, default=8)\n    ap.add_argument(\"--num_workers\", type=int, default=0)\n    ap.add_argument(\"--threshold\", type=float, default=0.5)\n\n    ap.add_argument(\"--cache_float16\", type=int, default=1)\n    ap.add_argument(\"--no_amp\", action=\"store_true\")\n\n    # notebook execution -> no argv -> use defaults\n    if argv is None:\n        import sys\n        argv = sys.argv[1:]\n        if (\"ipykernel\" in sys.argv[0]) or (\"colab_kernel_launcher\" in sys.argv[0]):\n            argv = []\n    return ap.parse_args(argv)\n\n\ndef main(argv=None):\n    args = parse_args(argv)\n\n    ckpt_path = find_ckpt_path(args.ckpt_dir)\n    start_z = None if args.start_z < 0 else int(args.start_z)\n\n    net, tile_size, num_frames = load_seg_model(\n        ckpt_path=ckpt_path,\n        tile_size=args.tile_size,\n        num_frames=args.num_frames,\n        mae_ckpt=args.mae_ckpt,\n    )\n\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    net = net.to(device).eval()\n\n    sample_sub_path = os.path.join(args.data_root, \"sample_submission.csv\")\n    sub = pd.read_csv(sample_sub_path)\n    frag_ids = sub[\"Id\"].tolist()\n\n    test_root = os.path.join(args.data_root, \"test\")\n    out_rle = []\n\n    for fid in frag_ids:\n        frag_root = os.path.join(test_root, str(fid))\n        mask_path = os.path.join(frag_root, \"mask.png\")\n        mask01 = load_mask_png(mask_path)\n\n        slice_paths = _list_tif_slices(frag_root)\n        z_idx = _choose_z_indices(len(slice_paths), num_frames, args.depth_mode, start_z)\n\n        print(f\"\\n=== Fragment {fid} | slices={len(slice_paths)} use={z_idx[0]}..{z_idx[-1]} (T={num_frames}) ===\")\n        vol_u16 = load_volume_slices_to_ram(slice_paths, z_idx)\n\n        prob = predict_fragment(\n            net=net,\n            vol_u16=vol_u16,\n            mask01=mask01,\n            tile_size=int(tile_size),\n            stride=int(args.stride),\n            clip_min=float(args.clip_min),\n            clip_max=float(args.clip_max),\n            batch_size=int(args.batch_size),\n            num_workers=int(args.num_workers),\n            cache_float16=bool(int(args.cache_float16)),\n            use_amp=(not args.no_amp),\n        )\n\n        pred_bin = (prob > float(args.threshold)).astype(np.uint8)\n        out_rle.append(rle_encode(pred_bin))\n\n    sub[\"Id\"] = sub[\"Id\"].astype(str).str.strip()\n    frag_ids = sub[\"Id\"].tolist()\n\n    write_submission_csv(args.out_csv, frag_ids, out_rle)\n\n# 自检：必须只有 3 行（header + a + b）\n    print(\"\\n[CHECK] first 2 lines:\")\n    with open(args.out_csv, \"r\", encoding=\"utf-8\") as f:\n        for _ in range(2):\n           print(f.readline().rstrip(\"\\n\"))\n\n    print(\"[CHECK] line count =\", sum(1 for _ in open(args.out_csv, \"r\", encoding=\"utf-8\")))\n    print(f\"\\n✅ saved: {args.out_csv}\")\n    print(sub.head())\n\n\nif __name__ == \"__main__\":\n    main()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python infer_submit.py \\\n  --ckpt_dir /kaggle/working/seg_outputs_run \\\n  --mae_ckpt /kaggle/working/mae_outputs/best_mae.pt \\\n  --out_csv /kaggle/working/submission.csv \\\n  --tile_size 64 --stride 64 --num_frames 24 \\\n  --depth_mode center_contig \\\n  --batch_size 8 --num_workers 0 \\\n  --threshold 0.5\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}