{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":117682,"databundleVersionId":15062069,"isSourceIdPinned":false},{"sourceType":"modelInstanceVersion","sourceId":732880,"databundleVersionId":15477237,"modelInstanceId":516822,"modelId":510647,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":290917305,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":298459500,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":299035542,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":299088560,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":299156472,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":299209845,"isSourceIdPinned":false}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"from IPython.display import clear_output\nimport os\n\n# protobuf stability (common Kaggle container mismatch)\nos.environ.setdefault(\"PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION\", \"python\")\nos.environ[\"KERAS_BACKEND\"] = \"jax\"\nos.environ.setdefault(\"OMP_NUM_THREADS\", \"4\")\nos.environ.setdefault(\"MKL_NUM_THREADS\", \"4\")\nos.environ.setdefault(\"OPENBLAS_NUM_THREADS\", \"4\")\nos.environ.setdefault(\"NUMEXPR_NUM_THREADS\", \"4\")\n\nvar=\"/kaggle/input/vsdetection-packages-offline-installer-only/whls\"\n!pip install \\\n  \"$var\"/keras_nightly-*.whl \\\n  \"$var\"/tifffile-*.whl \\\n  \"$var\"/imagecodecs-*.whl \\\n  \"$var\"/medicai-*.whl \\\n  --no-index \\\n  --find-links \"$var\"\nclear_output()\n\nimport time, zipfile\nimport numpy as np\nimport pandas as pd\nimport tifffile\nimport scipy.ndimage as ndi\nfrom skimage.morphology import remove_small_objects\nimport keras\nfrom medicai.transforms import Compose, NormalizeIntensity\nfrom medicai.models import TransUNet\nfrom medicai.utils.inference import SlidingWindowInference\n\nprint(\"Keras backend:\", keras.config.backend(), \"Keras version:\", keras.version())\n\n# ----------------------------\n# CONFIG (minimal, safe)\n# ----------------------------\nCFG = dict(\n    # model\n    kaggle_model_path=\"/kaggle/input/notebooks/tonyai007/train-vesuvius-seresnext50-comboloss-5\",\n    weights_relpath=\"final_model.weights.h5\",\n\n    # SWI overlaps\n    overlap_public=0.42,  # EXACT public 0.55\n    overlap_base=0.48,    # EXACT private 0.55 (important!)\n    overlap_hi=0.6,      # OV06\n    OV06_MAIN_ONLY=True,\n\n    # TTA\n    USE_TTA=True,\n\n    # binary logit definition (match your private 0.55)\n    INK_MODE=\"fg12\",\n\n    # thresholds (match your private 0.55)\n    T_low=0.5,\n    T_high=0.9,\n\n    # topology (match both notebooks' common setting)\n    z_radius=0,\n    xy_radius=2,\n    dust_min_size=200,\n\n    # warmup\n    DO_WARMUP=True,\n)\n\nroot_dir = \"/kaggle/input/vesuvius-challenge-surface-detection\"\ntest_dir = f\"{root_dir}/test_images\"\noutput_dir = \"/kaggle/working/submission_masks\"\nzip_path = \"/kaggle/working/submission.zip\"\nos.makedirs(output_dir, exist_ok=True)\n\ntest_df = pd.read_csv(f\"{root_dir}/test.csv\")\nids = test_df[\"id\"].tolist()\nprint(\"Num test volumes:\", len(ids))\n\nROI = (160, 160, 160)\n\n# ----------------------------\n# Transform (same style)\n# ----------------------------\n_val_pipeline = Compose([\n    NormalizeIntensity(keys=[\"image\"], nonzero=True, channel_wise=False),\n])\ndef val_transformation(image):\n    return _val_pipeline({\"image\": image})[\"image\"]\n\ndef load_volume(path):\n    vol = tifffile.imread(path).astype(np.float32)\n    return vol[None, ..., None]  # (1, D, H, W, 1)\n\n# ----------------------------\n# Numerics\n# ----------------------------\ndef sigmoid_stable(x):\n    x = np.asarray(x, dtype=np.float32)\n    out = np.empty_like(x, dtype=np.float32)\n    pos = x >= 0\n    out[pos] = 1.0 / (1.0 + np.exp(-x[pos]))\n    ex = np.exp(x[~pos])\n    out[~pos] = ex / (1.0 + ex)\n    return out\n\ndef logsumexp2(a, b):\n    a = np.asarray(a, dtype=np.float32)\n    b = np.asarray(b, dtype=np.float32)\n    m = np.maximum(a, b)\n    return m + np.log(np.exp(a - m) + np.exp(b - m) + 1e-12)\n\ndef binary_logit_from_multiclass_logits(logits_5d, mode=\"fg12\"):\n    x = np.asarray(logits_5d, dtype=np.float32)[0]  # (D,H,W,3)\n    L0, L1, L2 = x[...,0], x[...,1], x[...,2]\n    if mode == \"fg12\":\n        return (logsumexp2(L1, L2) - L0).astype(np.float32, copy=False)\n    elif mode == \"class1\":\n        return (L1 - logsumexp2(L0, L2)).astype(np.float32, copy=False)\n    else:\n        raise ValueError(f\"Unknown INK_MODE={mode}\")\n\n# ----------------------------\n# Topology helpers\n# ----------------------------\ndef build_anisotropic_struct(z_radius: int, xy_radius: int):\n    z, r = int(z_radius), int(xy_radius)\n    if z == 0 and r == 0:\n        return None\n    if z == 0 and r > 0:\n        size = 2*r + 1\n        struct = np.zeros((1, size, size), dtype=bool)\n        cy = cx = r\n        for dy in range(-r, r+1):\n            for dx in range(-r, r+1):\n                if dy*dy + dx*dx <= r*r:\n                    struct[0, cy+dy, cx+dx] = True\n        return struct\n    if z > 0 and r == 0:\n        struct = np.zeros((2*z+1, 1, 1), dtype=bool)\n        struct[:, 0, 0] = True\n        return struct\n    depth = 2*z + 1\n    size  = 2*r + 1\n    struct = np.zeros((depth, size, size), dtype=bool)\n    cz = z; cy = cx = r\n    for dz in range(-z, z+1):\n        for dy in range(-r, r+1):\n            for dx in range(-r, r+1):\n                if dy*dy + dx*dx <= r*r:\n                    struct[cz+dz, cy+dy, cx+dx] = True\n    return struct\n\ndef seeded_hysteresis_with_topology(\n    prob, pub_fg_bool,\n    T_low=0.50, T_high=0.90,\n    z_radius=0, xy_radius=2, dust_min_size=200\n):\n    prob = np.asarray(prob, dtype=np.float32)\n    strong = prob >= float(T_high)\n\n    # SAFE expansion: weak includes private weak OR public foreground\n    weak = (prob >= float(T_low)) | pub_fg_bool\n\n    if not strong.any():\n        return np.zeros_like(prob, dtype=np.uint8)\n\n    struct_hyst = ndi.generate_binary_structure(3, 3)\n    mask = ndi.binary_propagation(strong, mask=weak, structure=struct_hyst)\n\n    if not mask.any():\n        return np.zeros_like(prob, dtype=np.uint8)\n\n    struct_close = build_anisotropic_struct(z_radius, xy_radius)\n    if struct_close is not None:\n        mask = ndi.binary_closing(mask, structure=struct_close)\n\n    if int(dust_min_size) > 0:\n        mask = remove_small_objects(mask.astype(bool), min_size=int(dust_min_size))\n\n    return mask.astype(np.uint8)\n\n# ----------------------------\n# Model + SWI (logits)\n# ----------------------------\nweights_path = f\"{CFG['kaggle_model_path']}/{CFG['weights_relpath']}\"\n\nmodel = TransUNet(\n    input_shape=(160, 160, 160, 1),\n    encoder_name=\"seresnext50\",\n    classifier_activation=None,  # TRUE logits\n    num_classes=3,\n)\nmodel.load_weights(weights_path)\nprint(\"Model params (M):\", model.count_params() / 1e6)\n\ndef build_swi(overlap):\n    return SlidingWindowInference(\n        model,\n        num_classes=3,\n        roi_size=ROI,\n        sw_batch_size=1,\n        mode=\"gaussian\",\n        overlap=float(overlap),\n    )\n\nswi_public = build_swi(CFG[\"overlap_public\"])  # public 0.55\nswi_base   = build_swi(CFG[\"overlap_base\"])    # private base 0.55\nswi_hi     = build_swi(CFG[\"overlap_hi\"])      # OV06\n\n# ----------------------------\n# TTA (same set)\n# ----------------------------\ndef iter_tta(volume):\n    yield volume, (lambda y: y)\n    for axis in [1, 2, 3]:\n        v = np.flip(volume, axis=axis)\n        inv = (lambda y, axis=axis: np.flip(y, axis=axis))\n        yield v, inv\n    for k in [1, 2, 3]:\n        v = np.rot90(volume, k=k, axes=(2, 3))\n        inv = (lambda y, k=k: np.rot90(y, k=-k, axes=(2, 3)))\n        yield v, inv\n\n# ----------------------------\n# Predict BOTH streams in one loop (single final path)\n# - Public: mean multiclass logits -> argmax labels\n# - Private: OV06 main-only + mean binary logits -> prob\n# ----------------------------\ndef predict_pub_labels_and_private_prob(volume):\n    mode = CFG[\"INK_MODE\"]\n\n    if not CFG[\"USE_TTA\"]:\n        l_pub = np.asarray(swi_public(volume))\n        pub_labels = l_pub.argmax(-1).astype(np.uint8).squeeze()\n\n        l_prv = np.asarray(swi_hi(volume))\n        s = binary_logit_from_multiclass_logits(l_prv, mode=mode)\n        prob = sigmoid_stable(s)\n        return pub_labels, prob\n\n    logits_sum = None\n    s_sum = None\n    n = 0\n\n    for t, (v, inv) in enumerate(iter_tta(volume)):\n        # public stream\n        l_pub = np.asarray(swi_public(v))\n        l_pub = inv(l_pub)\n        logits_sum = l_pub.astype(np.float32) if logits_sum is None else (logits_sum + l_pub.astype(np.float32))\n\n        # private stream (OV06 main-only)\n        if CFG[\"OV06_MAIN_ONLY\"]:\n            swi_use = swi_hi if (t == 0) else swi_base\n        else:\n            swi_use = swi_hi\n\n        l_prv = np.asarray(swi_use(v))\n        l_prv = inv(l_prv)\n        s = binary_logit_from_multiclass_logits(l_prv, mode=mode)\n        s_sum = s.astype(np.float32) if s_sum is None else (s_sum + s.astype(np.float32))\n\n        n += 1\n\n    mean_logits = logits_sum / float(n)\n    pub_labels = mean_logits.argmax(-1).astype(np.uint8).squeeze()\n\n    s_mean = (s_sum / float(n)).astype(np.float32, copy=False)\n    prob = sigmoid_stable(s_mean)\n    return pub_labels, prob\n\n# ----------------------------\n# Warmup (compile once)\n# ----------------------------\ndef warmup(volume):\n    _ = np.asarray(swi_public(volume))\n    _ = np.asarray(swi_base(volume))\n    _ = np.asarray(swi_hi(volume))\n\n# ----------------------------\n# Run + zip\n# ----------------------------\nprint(\"CFG:\",\n      f\"overlap_public={CFG['overlap_public']}, overlap_base={CFG['overlap_base']}, overlap_hi={CFG['overlap_hi']},\",\n      f\"INK_MODE={CFG['INK_MODE']}, T_low={CFG['T_low']}, T_high={CFG['T_high']},\",\n      f\"OV06_MAIN_ONLY={CFG['OV06_MAIN_ONLY']}\")\n\nt_global0 = time.perf_counter()\n\nwith zipfile.ZipFile(zip_path, \"w\", compression=zipfile.ZIP_DEFLATED) as zf:\n    for i, image_id in enumerate(ids):\n        t0 = time.perf_counter()\n\n        tif_path = f\"{test_dir}/{image_id}.tif\"\n        volume = load_volume(tif_path)\n        volume = val_transformation(volume)\n\n        if i == 0 and CFG[\"DO_WARMUP\"]:\n            print(\"Warming up JAX (compile once)...\")\n            warmup(volume)\n\n        pub_labels, prob = predict_pub_labels_and_private_prob(volume)\n\n        # public fg anchor (no topology here; used only as weak region expansion)\n        pub_fg = (pub_labels != 0)\n\n        # single final path\n        output = seeded_hysteresis_with_topology(\n            prob,\n            pub_fg_bool=pub_fg,\n            T_low=CFG[\"T_low\"],\n            T_high=CFG[\"T_high\"],\n            z_radius=CFG[\"z_radius\"],\n            xy_radius=CFG[\"xy_radius\"],\n            dust_min_size=CFG[\"dust_min_size\"],\n        )\n\n        out_path = f\"{output_dir}/{image_id}.tif\"\n        tifffile.imwrite(out_path, output.astype(np.uint8))\n        zf.write(out_path, arcname=f\"{image_id}.tif\")\n        os.remove(out_path)\n\n        dt = time.perf_counter() - t0\n        elapsed = time.perf_counter() - t_global0\n        print(f\"[{i+1}/{len(ids)}] id={image_id} | {dt/60:.2f} min | elapsed {elapsed/3600:.2f} h | positives={int(output.sum())}\")\n\nprint(\"Submission ZIP:\", zip_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-21T17:06:55.936464Z","iopub.execute_input":"2026-02-21T17:06:55.936753Z","iopub.status.idle":"2026-02-21T17:12:02.785297Z","shell.execute_reply.started":"2026-02-21T17:06:55.936736Z","shell.execute_reply":"2026-02-21T17:12:02.784524Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import matplotlib.pyplot as plt\n# def plot_sample(x, y, sample_idx=0, max_slices=16):\n#     img = np.squeeze(x[sample_idx])\n#     mask = np.squeeze(y[sample_idx])\n#     D = img.shape[0]\n#     step = max(1, D // max_slices)\n#     slices = range(0, D, step)\n#     n_slices = len(slices)\n#     fig, axes = plt.subplots(2, n_slices, figsize=(3*n_slices, 6))\n#     for i, s in enumerate(slices):\n#         axes[0, i].imshow(img[s], cmap='gray')\n#         axes[0, i].set_title(f\"Slice {s}\")\n#         axes[0, i].axis('off')\n#         axes[1, i].imshow(mask[s], cmap='gray')\n#         axes[1, i].set_title(f\"Mask {s}\")\n#         axes[1, i].axis('off')\n#     plt.suptitle(f\"Sample {sample_idx}\")\n#     plt.tight_layout()\n#     plt.show()\n# plot_sample(volume.numpy(), output[None], sample_idx=0, max_slices=5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-21T17:12:02.786288Z","iopub.execute_input":"2026-02-21T17:12:02.78659Z","iopub.status.idle":"2026-02-21T17:12:02.790836Z","shell.execute_reply.started":"2026-02-21T17:12:02.786561Z","shell.execute_reply":"2026-02-21T17:12:02.790096Z"}},"outputs":[],"execution_count":null}]}