{"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":113558,"databundleVersionId":14878066,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":14403292,"sourceType":"datasetVersion","datasetId":9199005},{"sourceId":14403297,"sourceType":"datasetVersion","datasetId":9199008},{"sourceId":14412957,"sourceType":"datasetVersion","datasetId":9205342},{"sourceId":14428607,"sourceType":"datasetVersion","datasetId":9215849},{"sourceId":14435182,"sourceType":"datasetVersion","datasetId":9220408},{"sourceId":14439084,"sourceType":"datasetVersion","datasetId":9222897},{"sourceId":14439599,"sourceType":"datasetVersion","datasetId":9223185},{"sourceId":14463932,"sourceType":"datasetVersion","datasetId":9238473},{"sourceId":14475319,"sourceType":"datasetVersion","datasetId":9245663},{"sourceId":14475899,"sourceType":"datasetVersion","datasetId":9246050},{"sourceId":14486127,"sourceType":"datasetVersion","datasetId":9252493},{"sourceId":14499765,"sourceType":"datasetVersion","datasetId":9261328},{"sourceId":14508617,"sourceType":"datasetVersion","datasetId":9266650}],"dockerImageVersionId":31236,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"%%writefile test.py\n# -----------------------------\n# UPDATED DDP inference script:\n#   - saves prob_orig as 8-bit PNGs (0..255) to --out_prob_dir\n#   - NO prediction gathering (no big rank0 RAM spike)\n#   - keep tiny dist.reduce for error stats\n# -----------------------------\n\"\"\"\nDDP inference script that SAVES per-image probability maps (prob_orig) as 8-bit PNGs.\n\nOutputs (per image):\n  - {out_prob_dir}/{case_id}.png   # grayscale uint8, where prob = pixel/255.0\n\nKey behavior:\n- SAME preprocessing + TTA as before (4 rotations averaged)\n- SAME mapping back to original image size: resize to square -> crop to orig\n- Supports any file type OpenCV can read; unreadable files produce NO PNG\n  (postprocess cell should treat missing PNG as \"authentic\")\n- DDP inference with torchrun (file sharding)\n- Quant options:\n    --quant none   : normal fp16/bf16 autocast\n    --quant bnb8   : bitsandbytes int8 weights\n    --quant ao8    : TorchAO int8 weight-only\n    --quant ao4    : TorchAO int4 weight-only\n- Attention loading: try flash_attention_2, fallback to sdpa\n\nExample:\n  torchrun --standalone --nproc_per_node=2 submit_hf_save_probs_ddp.py \\\n    --models_dir /kaggle/input/my-model/exp1 \\\n    --test_img_dir /kaggle/input/comp/test \\\n    --out_prob_dir /kaggle/working/prob_maps\n\"\"\"\nfrom __future__ import annotations\n\nimport argparse\nimport contextlib\nimport json\nimport os\nfrom pathlib import Path\nfrom typing import Dict, List, Optional, Tuple\n\nimport cv2\nimport numpy as np\nimport torch\nimport torch.distributed as dist\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset\n\nfrom transformers import AutoImageProcessor, AutoModel\nfrom transformers import BitsAndBytesConfig\n\n# TorchAO (optional)\ntry:\n    from transformers import TorchAoConfig\n    from torchao.quantization import Int8WeightOnlyConfig, Int4WeightOnlyConfig\n\n    _HAS_TORCHAO = True\nexcept Exception:\n    _HAS_TORCHAO = False\n\n\n# -----------------------------\n# DDP utils\n# -----------------------------\ndef ddp_init() -> Tuple[bool, int, int, int]:\n    \"\"\"\n    Returns (is_ddp, rank, world_size, local_rank).\n    Detects torchrun via env vars.\n    \"\"\"\n    if \"RANK\" in os.environ and \"WORLD_SIZE\" in os.environ:\n        rank = int(os.environ[\"RANK\"])\n        world_size = int(os.environ[\"WORLD_SIZE\"])\n        local_rank = int(os.environ.get(\"LOCAL_RANK\", \"0\"))\n\n        if torch.cuda.is_available():\n            torch.cuda.set_device(local_rank)\n            dist.init_process_group(backend=\"nccl\", init_method=\"env://\")\n        else:\n            dist.init_process_group(backend=\"gloo\", init_method=\"env://\")\n        return True, rank, world_size, local_rank\n\n    return False, 0, 1, 0\n\n\ndef ddp_is_main(rank: int) -> bool:\n    return rank == 0\n\n\ndef ddp_barrier(is_ddp: bool) -> None:\n    if is_ddp:\n        dist.barrier()\n\n\ndef ddp_broadcast_object(is_ddp: bool, obj, rank: int):\n    \"\"\"Broadcast a Python object from rank 0 to all ranks.\"\"\"\n    if not is_ddp:\n        return obj\n    obj_list = [obj] if rank == 0 else [None]\n    dist.broadcast_object_list(obj_list, src=0)\n    return obj_list[0]\n\n\n# -----------------------------\n# Image helpers\n# -----------------------------\ndef load_image_rgb(path: Path) -> np.ndarray:\n    \"\"\"\n    Reads image via OpenCV (supports many types). Returns RGB uint8 HWC.\n    Handles grayscale and RGBA. If dtype != uint8 (e.g. 16-bit), scales to uint8.\n    \"\"\"\n    img = cv2.imread(str(path), cv2.IMREAD_UNCHANGED)\n    if img is None:\n        raise FileNotFoundError(f\"Failed to read image: {path}\")\n\n    if img.ndim == 2:\n        img = np.stack([img, img, img], axis=-1)\n    elif img.ndim == 3:\n        if img.shape[2] == 1:\n            img = np.repeat(img, 3, axis=2)\n        elif img.shape[2] >= 3:\n            img = img[:, :, :3]   # works for 3,4,>4 channels\n    else:\n        raise ValueError(f\"Unsupported image shape: {img.shape}\")\n\n    if img.dtype != np.uint8:\n        img_f = img.astype(np.float32)\n        maxv = float(np.max(img_f)) if img_f.size else 1.0\n        if maxv <= 0:\n            img = np.zeros_like(img_f, dtype=np.uint8)\n        else:\n            img = np.clip(img_f * (255.0 / maxv), 0, 255).astype(np.uint8)\n\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    return img\n\n\ndef pad_to_square(img: np.ndarray, pad_value: int = 0):\n    h, w = img.shape[:2]\n    size = max(h, w)\n    pad_top = (size - h) // 2\n    pad_bottom = size - h - pad_top\n    pad_left = (size - w) // 2\n    pad_right = size - w - pad_left\n    img_pad = cv2.copyMakeBorder(\n        img,\n        pad_top,\n        pad_bottom,\n        pad_left,\n        pad_right,\n        borderType=cv2.BORDER_CONSTANT,\n        value=(pad_value, pad_value, pad_value),\n    )\n    meta = {\n        \"orig_h\": int(h),\n        \"orig_w\": int(w),\n        \"pad_top\": int(pad_top),\n        \"pad_left\": int(pad_left),\n        \"square_size\": int(size),\n    }\n    return img_pad, meta\n\n\ndef list_files(folder: Path, recursive: bool) -> List[Path]:\n    if recursive:\n        files = [p for p in folder.rglob(\"*\") if p.is_file()]\n    else:\n        files = [p for p in folder.iterdir() if p.is_file()]\n    files.sort(key=lambda p: p.name.lower())\n    return files\n\n\n# -----------------------------\n# Model wrapper\n# -----------------------------\ndef get_prefix_tokens_from_config(cfg) -> int:\n    return 1 + int(getattr(cfg, \"num_register_tokens\", 0) or 0)\n\n\nclass HFViTSemSeg(nn.Module):\n    def __init__(self, backbone: nn.Module, num_prefix_tokens: int, embed_dim: int):\n        super().__init__()\n        self.backbone = backbone\n        self.num_prefix_tokens = int(num_prefix_tokens)\n        self.embed_dim = int(embed_dim)\n        self.head = nn.Conv2d(self.embed_dim, 1, kernel_size=1)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        try:\n            out = self.backbone(pixel_values=x, return_dict=True, interpolate_pos_encoding=True)\n        except TypeError:\n            out = self.backbone(pixel_values=x, return_dict=True)\n\n        t = out.last_hidden_state  # (B,T,C)\n        patch_tokens = t[:, self.num_prefix_tokens :, :]\n        b, n, c = patch_tokens.shape\n        g = int(np.sqrt(n))\n        if g * g != n:\n            raise RuntimeError(f\"Expected square patch grid, got n={n}\")\n        feat = patch_tokens.transpose(1, 2).reshape(b, c, g, g)\n        return self.head(feat)  # (B,1,g,g)\n\n\n# -----------------------------\n# TTA (4 rotations: 0,90,180,270)\n# -----------------------------\ndef tta_specs() -> List[int]:\n    return [0, 1, 2, 3]\n\n\n# -----------------------------\n# Backbone loading (flash_attention_2 -> sdpa fallback)\n# -----------------------------\ndef load_backbone(models_dir: Path, quant: str, local_rank: int, amp_dtype: torch.dtype, rank0: bool):\n    backbone_dir = models_dir / \"backbone\"\n\n    if torch.cuda.is_available():\n        torch.backends.cuda.enable_flash_sdp(True)\n        torch.backends.cuda.enable_mem_efficient_sdp(True)\n        torch.backends.cuda.enable_math_sdp(True)\n\n    attn_order = (\"flash_attention_2\", \"sdpa\")\n    last_err: Optional[Exception] = None\n\n    for attn in attn_order:\n        try:\n            if quant == \"none\":\n                m = AutoModel.from_pretrained(\n                    backbone_dir,\n                    attn_implementation=attn,\n                )\n                if torch.cuda.is_available():\n                    m.to(device=torch.device(\"cuda\", local_rank))\n                if rank0:\n                    print(f\"[INFO] Backbone attn_implementation={attn} (quant=none)\")\n                return m\n\n            if quant == \"bnb8\":\n                if not torch.cuda.is_available():\n                    raise RuntimeError(\"bnb8 quantization requires CUDA.\")\n                bnb = BitsAndBytesConfig(load_in_8bit=True)\n                try:\n                    m = AutoModel.from_pretrained(\n                        backbone_dir,\n                        quantization_config=bnb,\n                        device_map={\"\": local_rank},\n                        dtype=amp_dtype,\n                        attn_implementation=attn,\n                    )\n                except TypeError:\n                    m = AutoModel.from_pretrained(\n                        backbone_dir,\n                        quantization_config=bnb,\n                        device_map={\"\": local_rank},\n                        torch_dtype=amp_dtype,\n                        attn_implementation=attn,\n                    )\n                if rank0:\n                    print(f\"[INFO] Backbone attn_implementation={attn} (quant=bnb8)\")\n                return m\n\n            if quant in (\"ao8\", \"ao4\"):\n                if not torch.cuda.is_available():\n                    raise RuntimeError(\"TorchAO weight-only quantization requires CUDA here.\")\n                if not _HAS_TORCHAO:\n                    raise RuntimeError(\"TorchAO requested but not available. Install torchao + compatible transformers.\")\n\n                qtype = Int8WeightOnlyConfig() if quant == \"ao8\" else Int4WeightOnlyConfig(group_size=128)\n                qcfg = TorchAoConfig(quant_type=qtype)\n\n                try:\n                    m = AutoModel.from_pretrained(\n                        backbone_dir,\n                        quantization_config=qcfg,\n                        device_map={\"\": local_rank},\n                        dtype=amp_dtype,\n                        attn_implementation=attn,\n                    )\n                except TypeError:\n                    m = AutoModel.from_pretrained(\n                        backbone_dir,\n                        quantization_config=qcfg,\n                        device_map={\"\": local_rank},\n                        torch_dtype=amp_dtype,\n                        attn_implementation=attn,\n                    )\n                if rank0:\n                    print(f\"[INFO] Backbone attn_implementation={attn} (quant={quant})\")\n                return m\n\n            raise ValueError(f\"Unknown --quant {quant}\")\n\n        except Exception as e:\n            last_err = e\n            if rank0:\n                print(\n                    f\"[WARN] Failed loading backbone with attn_implementation='{attn}' \"\n                    f\"(quant={quant}): {type(e).__name__}: {e}\"\n                )\n\n    raise RuntimeError(f\"Failed to load backbone with flash_attention_2 or sdpa. Last error: {last_err}\")\n\n\n# -----------------------------\n# Dataset / DataLoader\n# -----------------------------\nclass TestImageDataset(Dataset):\n    \"\"\"\n    Loads + pads-to-square + resizes on CPU.\n    Returns a uint8 CHW tensor (img_size x img_size) + meta needed to unpad back to original size.\n    Any read/preprocess failure is caught and surfaced as ok=False so we never drop a file.\n    \"\"\"\n\n    def __init__(self, paths: List[Path], img_size: int):\n        self.paths = paths\n        self.img_size = int(img_size)\n\n    def __len__(self) -> int:\n        return len(self.paths)\n\n    def __getitem__(self, idx: int):\n        path = self.paths[idx]\n        case_id = path.stem  # NOTE: intentionally no extension (consistent with your current code)\n        try:\n            img = load_image_rgb(path)\n            img_pad, meta = pad_to_square(img)\n            img_rs = cv2.resize(img_pad, (self.img_size, self.img_size), interpolation=cv2.INTER_LINEAR)\n            x0 = torch.from_numpy(img_rs).permute(2, 0, 1).contiguous()  # uint8 CHW (CPU)\n            return {\n                \"ok\": True,\n                \"case_id\": case_id,\n                \"file_name\": path.name,\n                \"x0_u8_chw\": x0,\n                \"meta\": meta,\n            }\n        except Exception as e:\n            return {\n                \"ok\": False,\n                \"case_id\": case_id,\n                \"file_name\": path.name,\n                \"err\": f\"{type(e).__name__}: {e}\",\n            }\n\n\ndef collate_one(batch):\n    # batch_size is fixed to 1 by design\n    return batch[0]\n\n\n# -----------------------------\n# Prediction: return prob_orig as uint8 PNG-ready map\n# -----------------------------\n@torch.inference_mode()\ndef predict_prob_orig_u16_tta_fast(\n    model: nn.Module,\n    x0_u8_chw: torch.Tensor,\n    meta: Dict[str, int],\n    img_size: int,\n    mean_t: torch.Tensor,\n    std_t: torch.Tensor,\n    device: torch.device,\n    amp_dtype: torch.dtype,\n    use_tta: bool = True,\n    file_name: str = \"\",\n) -> Tuple[Optional[np.ndarray], int]:\n    \"\"\"\n    Returns:\n      - prob_u16: uint16 HxW in [0,65535], where prob = prob_u16/65535.0, at ORIGINAL image size\n      - err_flag: 1 if failed (returns None), else 0\n    \"\"\"\n    try:\n        input_dtype = amp_dtype if device.type == \"cuda\" else torch.float32\n\n        x0 = x0_u8_chw.to(device=device, non_blocking=True).to(dtype=input_dtype)\n        x0 = x0.div_(255.0).unsqueeze(0)\n        x0 = (x0 - mean_t) / std_t\n\n        if use_tta:\n            ks = [0, 1, 2, 3]\n            x = torch.cat([torch.rot90(x0, k=k, dims=(2, 3)) for k in ks], dim=0)\n        else:\n            ks = [0]\n            x = x0\n\n        autocast_ctx = torch.autocast(\"cuda\", dtype=amp_dtype, enabled=True) if device.type == \"cuda\" else contextlib.nullcontext()\n\n        with autocast_ctx:\n            logits = model(x)                           # (B,1,g,g)\n            logits = F.interpolate(logits, size=(img_size, img_size), mode=\"bilinear\", align_corners=False)\n            probs = torch.sigmoid(logits[:, 0])         # (B,H,W)\n\n        if use_tta:\n            probs = torch.stack(\n                [torch.rot90(probs[i], k=(4 - ks[i]) % 4, dims=(0, 1)) for i in range(len(ks))],\n                dim=0,\n            )\n            prob_avg = probs.mean(dim=0)\n        else:\n            prob_avg = probs[0]\n\n        prob_avg_np = prob_avg.float().cpu().numpy()    # float32 (img_size,img_size)\n\n        orig_h, orig_w = int(meta[\"orig_h\"]), int(meta[\"orig_w\"])\n        sq = int(meta[\"square_size\"])\n        pt, pl = int(meta[\"pad_top\"]), int(meta[\"pad_left\"])\n\n        prob_sq = cv2.resize(prob_avg_np, (sq, sq), interpolation=cv2.INTER_LINEAR)\n        prob_orig = prob_sq[pt : pt + orig_h, pl : pl + orig_w]  # float32\n\n        # ---- 16-bit quantization ----\n        prob_u16 = np.clip(np.rint(prob_orig * 65535.0), 0, 65535).astype(np.uint16)\n        return prob_u16, 0\n\n    except Exception as e:\n        if file_name:\n            print(f\"[WARN] Predict failed on {file_name} -> skipping prob PNG ({type(e).__name__}: {e})\")\n        else:\n            print(f\"[WARN] Predict failed -> skipping prob PNG ({type(e).__name__}: {e})\")\n        return None, 1\n\n\ndef main():\n    p = argparse.ArgumentParser()\n    p.add_argument(\"--models_dir\", type=str, required=True)\n    p.add_argument(\"--test_img_dir\", type=str, required=True)\n    p.add_argument(\"--out_prob_dir\", type=str, required=True, help=\"Directory to write uint8 prob PNGs ({case_id}.png)\")\n\n    p.add_argument(\"--quant\", type=str, default=\"none\", choices=[\"none\", \"bnb8\", \"ao8\", \"ao4\"])\n    p.add_argument(\"--amp\", type=str, default=\"fp16\", choices=[\"fp16\", \"bf16\"])\n    p.add_argument(\"--no_tta\", action=\"store_true\")\n    p.add_argument(\"--recursive\", action=\"store_true\")\n\n    # DataLoader knobs (batch_size is fixed to 1 by design)\n    p.add_argument(\"--num_workers\", type=int, default=2)\n    p.add_argument(\"--prefetch_factor\", type=int, default=2)\n\n    # PNG compression (0 fastest/largest .. 9 smallest/slowest)\n    p.add_argument(\"--png_compression\", type=int, default=3)\n\n    args = p.parse_args()\n\n    is_ddp, rank, world_size, local_rank = ddp_init()\n    rank0 = ddp_is_main(rank)\n\n    models_dir = Path(args.models_dir)\n    test_dir = Path(args.test_img_dir)\n    out_prob_dir = Path(args.out_prob_dir)\n\n    if not test_dir.exists():\n        raise FileNotFoundError(f\"test_img_dir not found: {test_dir}\")\n\n    # Device\n    if torch.cuda.is_available():\n        device = torch.device(\"cuda\", local_rank if is_ddp else 0)\n    else:\n        device = torch.device(\"cpu\")\n        if rank0:\n            print(\"[WARN] CUDA not available. Running on CPU (slow).\")\n\n    amp_dtype = torch.float16 if args.amp == \"fp16\" else torch.bfloat16\n\n    # Load config for img_size\n    with open(models_dir / \"config.json\", \"r\") as f:\n        cfg = json.load(f)\n    img_size = int(cfg[\"run\"][\"img_size\"])\n\n    # Load processor mean/std\n    processor = AutoImageProcessor.from_pretrained(models_dir / \"processor\")\n    mean = tuple(float(x) for x in processor.image_mean)\n    std = tuple(float(x) for x in processor.image_std)\n\n    # Load head checkpoint\n    ckpt_path = models_dir / \"model.pt\"\n    if not ckpt_path.exists():\n        raise FileNotFoundError(f\"Missing checkpoint: {ckpt_path}\")\n    ckpt = torch.load(ckpt_path, map_location=\"cpu\")\n\n    # Load backbone (possibly quantized) + build model\n    backbone = load_backbone(\n        models_dir=models_dir,\n        quant=args.quant,\n        local_rank=local_rank if is_ddp else 0,\n        amp_dtype=amp_dtype,\n        rank0=rank0,\n    )\n    backbone.eval()\n\n    num_prefix_tokens = int(ckpt.get(\"num_prefix_tokens\", get_prefix_tokens_from_config(backbone.config)))\n    embed_dim = int(ckpt[\"embed_dim\"])\n    model = HFViTSemSeg(backbone=backbone, num_prefix_tokens=num_prefix_tokens, embed_dim=embed_dim)\n    model.head.load_state_dict(ckpt[\"head_state\"], strict=True)\n\n    # Match your val behavior: head fp16 on CUDA\n    if device.type == \"cuda\":\n        model.head.to(device=device, dtype=torch.float16)\n    else:\n        model.head.to(device=device, dtype=torch.float32)\n    model.eval()\n\n    # Precompute mean/std tensors ONCE per rank\n    input_dtype = amp_dtype if device.type == \"cuda\" else torch.float32\n    mean_t = torch.tensor(mean, device=device, dtype=input_dtype).view(1, 3, 1, 1)\n    std_t = torch.tensor(std, device=device, dtype=input_dtype).view(1, 3, 1, 1)\n\n    # Rank0 lists files and broadcasts to ensure identical ordering across ranks\n    if rank0:\n        files = list_files(test_dir, recursive=args.recursive)\n        if len(files) == 0:\n            raise RuntimeError(f\"No files found in {test_dir}\")\n\n        stems = [p.stem for p in files]\n\n        files_str = [str(p) for p in files]\n    else:\n        files_str = None\n\n    files_str = ddp_broadcast_object(is_ddp, files_str, rank)\n    files = [Path(s) for s in files_str]\n\n    ddp_barrier(is_ddp)\n\n    # Make output dir (safe to call on all ranks)\n    out_prob_dir.mkdir(parents=True, exist_ok=True)\n    ddp_barrier(is_ddp)\n\n    # Shard work\n    local_files = files[rank::world_size] if is_ddp else files\n    if rank0:\n        print(f\"[INFO] Total files: {len(files)} | world_size={world_size} | TTA={'off' if args.no_tta else 'on'}\")\n        print(f\"[INFO] Writing prob PNGs to: {out_prob_dir.resolve()}\")\n        print(f\"[INFO] DataLoader: batch_size=1 | num_workers={args.num_workers} | pin_memory={device.type=='cuda'}\")\n\n    # DataLoader (batch_size fixed to 1)\n    dataset = TestImageDataset(local_files, img_size=img_size)\n    dl_kwargs = dict(\n        batch_size=1,\n        shuffle=False,\n        num_workers=int(args.num_workers),\n        pin_memory=(device.type == \"cuda\"),\n        collate_fn=collate_one,\n        drop_last=False,\n    )\n    if int(args.num_workers) > 0:\n        dl_kwargs[\"persistent_workers\"] = True\n        dl_kwargs[\"prefetch_factor\"] = int(args.prefetch_factor)\n\n    loader = DataLoader(dataset, **dl_kwargs)\n\n    local_errs = 0\n    local_written = 0\n\n    png_params = [cv2.IMWRITE_PNG_COMPRESSION, int(np.clip(args.png_compression, 0, 9))]\n\n    for sample in loader:\n        case_id = sample[\"case_id\"]\n        out_path = out_prob_dir / f\"{case_id}.png\"\n\n        if not sample[\"ok\"]:\n            print(\n                f\"[WARN] Failed on {sample['file_name']} -> no prob PNG written \"\n                f\"({sample.get('err', 'unknown error')})\"\n            )\n            local_errs += 1\n            continue\n\n        prob_u16, err = predict_prob_orig_u16_tta_fast(\n            model=model,\n            x0_u8_chw=sample[\"x0_u8_chw\"],\n            meta=sample[\"meta\"],\n            img_size=img_size,\n            mean_t=mean_t,\n            std_t=std_t,\n            device=device,\n            amp_dtype=amp_dtype,\n            use_tta=(not args.no_tta),\n            file_name=sample[\"file_name\"],\n        )\n        local_errs += err\n        if prob_u16 is None:\n            continue\n\n        ok = cv2.imwrite(str(out_path), prob_u16, png_params)\n        if not ok:\n            print(f\"[WARN] cv2.imwrite failed for {out_path} (skipping)\")\n            local_errs += 1\n            continue\n\n        local_written += 1\n\n    # Reduce stats to rank0 (tiny tensors; no gather_object)\n    if is_ddp:\n        err_dev = device if device.type == \"cuda\" else torch.device(\"cpu\")\n        err_t = torch.tensor([local_errs], device=err_dev, dtype=torch.int64)\n        wr_t = torch.tensor([local_written], device=err_dev, dtype=torch.int64)\n        dist.reduce(err_t, dst=0, op=dist.ReduceOp.SUM)\n        dist.reduce(wr_t, dst=0, op=dist.ReduceOp.SUM)\n    else:\n        err_t = torch.tensor([local_errs], dtype=torch.int64)\n        wr_t = torch.tensor([local_written], dtype=torch.int64)\n\n    if rank0:\n        print(\"\\n[INFO] Done saving probability maps.\")\n        print(f\"[INFO] Total prob PNGs written: {int(wr_t.item())} / {len(files)}\")\n        print(f\"[INFO] Total read/predict/write errors: {int(err_t.item())}\")\n        print(\"[INFO] Next: run the postprocessing cell to create submission.csv.\")\n\n    ddp_barrier(is_ddp)\n    if is_ddp:\n        dist.destroy_process_group()\n\n\nif __name__ == \"__main__\":\n    main()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-10T01:21:33.943943Z","iopub.execute_input":"2026-01-10T01:21:33.944332Z","iopub.status.idle":"2026-01-10T01:21:33.961282Z","shell.execute_reply.started":"2026-01-10T01:21:33.944281Z","shell.execute_reply":"2026-01-10T01:21:33.960474Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!CUDA_VISIBLE_DEVICES=0,1 torchrun --standalone --nproc_per_node=2 test.py \\\n    --models_dir /kaggle/input/dinov3102415epnocut \\\n    --test_img_dir /kaggle/input/recodai-luc-scientific-image-forgery-detection/test_images \\\n    --out_prob_dir /kaggle/working/prob_maps \\\n    --quant bnb8","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -----------------------------\n# PARALLEL INDEPENDENT POSTPROCESS CELL\n#   - Reads saved prob PNGs in parallel\n#   - Applies your postprocessing + RLE\n#   - Writes submission.csv in stable order\n#\n# Notes:\n#   - Uses ProcessPoolExecutor (true parallelism; avoids GIL issues)\n#   - Keeps memory bounded: each worker loads/processes ONE image at a time\n#   - Main process stores only {case_id: annotation} (strings)\n#   - Missing/unreadable prob PNG => \"authentic\"\n# -----------------------------\nimport csv\nimport json\nimport os\nfrom concurrent.futures import ProcessPoolExecutor, as_completed\nfrom pathlib import Path\nfrom typing import Dict, List, Tuple, Optional\n\nimport cv2\nimport numpy as np\n\n\n# -------- RLE (official-style formatting) --------\ndef _rle_encode_numpy(x: np.ndarray, fg_val: int = 1) -> List[int]:\n    \"\"\"\n    Kaggle-style RLE on a 2D mask, using column-major order (\"Fortran\"),\n    equivalent to x.T.flatten().\n    Returns [start1, len1, start2, len2, ...] with 1-indexed starts.\n    \"\"\"\n    if x.ndim != 2:\n        raise ValueError(f\"RLE expects 2D mask, got shape={x.shape}\")\n    pixels = (x == fg_val).flatten(order=\"F\")\n    dots = np.where(pixels)[0]\n    run_lengths: List[int] = []\n    prev = -2\n    for b in dots:\n        if b > prev + 1:\n            run_lengths.extend((int(b) + 1, 0))\n        run_lengths[-1] += 1\n        prev = int(b)\n    return run_lengths\n\n\ndef rle_encode(masks: List[np.ndarray], fg_val: int = 1) -> str:\n    \"\"\"Join multiple instance RLEs with ';'. Each instance is JSON-dumped list[int].\"\"\"\n    return \";\".join([json.dumps(_rle_encode_numpy(m.astype(np.uint8), fg_val=fg_val)) for m in masks])\n\n\n# -------- Postprocessing helpers (same logic as your script) --------\ndef drop_small_connected_components(mask_u8: np.ndarray, min_pixels: int = 256, connectivity: int = 8) -> np.ndarray:\n    \"\"\"\n    Remove small islands (< min_pixels) from a binary mask.\n    mask_u8: uint8 2D array with values {0,1} (or {0,255} also works).\n    Returns a uint8 2D array with values {0,1}.\n    \"\"\"\n    if mask_u8.ndim != 2:\n        raise ValueError(f\"mask must be 2D, got shape={mask_u8.shape}\")\n\n    mask01 = (mask_u8 > 0).astype(np.uint8)\n    num_labels, labels, stats, _ = cv2.connectedComponentsWithStats(mask01, connectivity=connectivity)\n    if num_labels <= 1:\n        return mask01\n\n    areas = stats[:, cv2.CC_STAT_AREA]\n    keep = np.zeros(num_labels, dtype=np.uint8)\n    keep[1:] = (areas[1:] >= int(min_pixels)).astype(np.uint8)\n\n    out = keep[labels]\n    return out.astype(np.uint8)\n\n\ndef keep_top_k_connected_components(mask_u8: np.ndarray, k: int = 1, connectivity: int = 8) -> np.ndarray:\n    \"\"\"\n    Keep only the k largest connected components from a binary mask.\n    Returns uint8 {0,1}.\n    \"\"\"\n    if mask_u8.ndim != 2:\n        raise ValueError(f\"mask must be 2D, got shape={mask_u8.shape}\")\n\n    if k <= 0:\n        return np.zeros_like(mask_u8, dtype=np.uint8)\n\n    mask01 = (mask_u8 > 0).astype(np.uint8)\n    num_labels, labels, stats, _ = cv2.connectedComponentsWithStats(mask01, connectivity=connectivity)\n    if num_labels <= 1:\n        return mask01\n\n    areas = stats[1:, cv2.CC_STAT_AREA]\n    if areas.size <= k:\n        return mask01\n\n    topk_idx = np.argpartition(areas, -k)[-k:]\n    keep_labels = (topk_idx + 1).astype(np.int32)\n\n    keep = np.zeros(num_labels, dtype=np.uint8)\n    keep[keep_labels] = 1\n    out = keep[labels]\n    return out.astype(np.uint8)\n\n\ndef list_files(folder: Path, recursive: bool) -> List[Path]:\n    if recursive:\n        files = [p for p in folder.rglob(\"*\") if p.is_file()]\n    else:\n        files = [p for p in folder.iterdir() if p.is_file()]\n    files.sort(key=lambda p: p.name.lower())\n    return files\n\n\n# -------- Worker function (must be top-level for multiprocessing) --------\ndef _worker_process_one(\n    case_id: str,\n    prob_png_path: str,\n    mask_thr: float,\n    min_area: int,\n    topk_islands: int,\n    small_island_min_pixels: int,\n    connectivity: int,\n) -> Tuple[str, str, int, int]:\n    \"\"\"\n    Returns:\n      (case_id, annotation, missing_flag, error_flag)\n    missing_flag=1 if file missing, error_flag=1 if unreadable/exception.\n    \"\"\"\n    try:\n        p = Path(prob_png_path)\n        if not p.exists():\n            return case_id, \"authentic\", 1, 0\n\n        prob_u16 = cv2.imread(str(p), cv2.IMREAD_UNCHANGED)\n        if prob_u16 is None:\n            return case_id, \"authentic\", 0, 1\n\n        # Handle both 16-bit and accidental 8-bit (robustness)\n        if prob_u16.dtype == np.uint16:\n            prob = prob_u16.astype(np.float32) / 65535.0\n        elif prob_u16.dtype == np.uint8:\n            prob = prob_u16.astype(np.float32) / 255.0\n        else:\n            # unexpected type -> be safe\n            return case_id, \"authentic\", 0, 1\n\n        # Authenticity gate (same as your script): prob>=0.6 area < min_area -> authentic\n        bin_mask_06 = (prob >= 0.75).astype(np.uint8)\n        area = int(bin_mask_06.sum())\n        if area < int(min_area):\n            return case_id, \"authentic\", 0, 0\n\n        # Main threshold\n        bin_mask = (prob >= float(mask_thr)).astype(np.uint8)\n\n        # Drop small islands after passing gate\n        #bin_mask = drop_small_connected_components(\n        #    bin_mask, min_pixels=int(small_island_min_pixels), connectivity=int(connectivity)\n        #)\n\n        # Keep only top-K islands (optional)\n        #if topk_islands and int(topk_islands) > 0:\n        #    bin_mask = keep_top_k_connected_components(\n        #        bin_mask, k=int(topk_islands), connectivity=int(connectivity)\n        #    )\n\n        # If empty -> authentic\n        if int(bin_mask.sum()) == 0:\n            return case_id, \"authentic\", 0, 0\n\n        ann = rle_encode([bin_mask], fg_val=1)\n        return case_id, ann, 0, 0\n\n    except Exception:\n        # Keep it robust: any failure -> authentic\n        return case_id, \"authentic\", 0, 1\n\n\n# -------- Main parallel postprocess --------\ndef make_submission_from_prob_pngs_parallel(\n    test_img_dir: str,\n    prob_dir: str,\n    out_csv: str = \"submission.csv\",\n    recursive: bool = False,\n    mask_thr: float = 0.5,\n    min_area: int = 32,\n    topk_islands: int = 0,\n    small_island_min_pixels: int = 64,\n    connectivity: int = 8,\n    max_workers: Optional[int] = None,\n    chunksize: int = 16,\n):\n    \"\"\"\n    Parallel postprocess:\n      - Enumerates test images to define stable output order\n      - Spawns processes to read prob PNGs and compute annotation\n      - Writes CSV in the same stable order as test listing\n\n    Memory notes:\n      - Workers process one image at a time; they return only strings.\n      - Main process keeps a dict of annotations (strings), size ~O(N).\n    \"\"\"\n    test_dir = Path(test_img_dir)\n    prob_dir = Path(prob_dir)\n    out_csv = Path(out_csv)\n\n    files = list_files(test_dir, recursive=recursive)\n    if not files:\n        raise RuntimeError(f\"No files found in {test_dir}\")\n\n    stems = [p.stem for p in files]\n\n    # Choose a conservative default worker count (avoid RAM blowups)\n    if max_workers is None:\n        cpu = os.cpu_count() or 4\n        max_workers = max(1, min(8, cpu // 2))  # typical safe default on Kaggle\n\n    # Prepare tasks\n    tasks = []\n    for p in files:\n        case_id = p.stem\n        prob_path = str(prob_dir / f\"{case_id}.png\")\n        tasks.append((case_id, prob_path))\n\n    annotations: Dict[str, str] = {}\n    n_missing = 0\n    n_errors = 0\n\n    # Run parallel work\n    with ProcessPoolExecutor(max_workers=max_workers) as ex:\n        # executor.map is memory-friendly and fast; chunksize improves throughput\n        it = ex.map(\n            _worker_process_one,\n            (t[0] for t in tasks),\n            (t[1] for t in tasks),\n            (mask_thr for _ in tasks),\n            (min_area for _ in tasks),\n            (topk_islands for _ in tasks),\n            (small_island_min_pixels for _ in tasks),\n            (connectivity for _ in tasks),\n            chunksize=chunksize,\n        )\n\n        for case_id, ann, miss, err in it:\n            annotations[case_id] = ann\n            n_missing += int(miss)\n            n_errors += int(err)\n\n    # Write CSV in stable order\n    out_csv.parent.mkdir(parents=True, exist_ok=True)\n    n_auth = 0\n    n_mask = 0\n\n    with out_csv.open(\"w\", newline=\"\") as f:\n        w = csv.writer(f)\n        w.writerow([\"case_id\", \"annotation\"])\n        for p in files:\n            case_id = p.stem\n            ann = annotations.get(case_id, \"authentic\")\n            w.writerow([case_id, ann])\n            if ann == \"authentic\":\n                n_auth += 1\n            else:\n                n_mask += 1\n\n    print(f\"Saved submission CSV: {out_csv.resolve()}\")\n    print(f\"Total rows: {len(files)}\")\n    print(f\"authentic: {n_auth} | mask: {n_mask}\")\n    print(f\"missing prob PNGs: {n_missing} | unreadable/exception PNGs: {n_errors}\")\n    print(f\"Parallelism: max_workers={max_workers} | chunksize={chunksize}\")\n\n\n# ---- Example usage (edit paths) ----\nmake_submission_from_prob_pngs_parallel(\n    test_img_dir=\"/kaggle/input/recodai-luc-scientific-image-forgery-detection/test_images\",\n    prob_dir=\"/kaggle/working/prob_maps\",\n    out_csv=\"/kaggle/working/submission.csv\",\n    recursive=False,\n    mask_thr=0.3,\n    min_area=32,\n    topk_islands=0,\n    max_workers=4,   # tune (4-8 is usually safe)\n    chunksize=16,    # tune (16-64)\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}