{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":117682,"databundleVersionId":15062069},{"sourceType":"datasetVersion","sourceId":14924567,"datasetId":9549680,"databundleVersionId":15791545},{"sourceType":"datasetVersion","sourceId":14933969,"datasetId":9556372,"databundleVersionId":15801881},{"sourceType":"datasetVersion","sourceId":14944058,"datasetId":9563841,"databundleVersionId":15813420},{"sourceType":"datasetVersion","sourceId":14940461,"datasetId":9561376,"databundleVersionId":15809216},{"sourceType":"datasetVersion","sourceId":14968646,"datasetId":9568607,"databundleVersionId":15840591},{"sourceType":"datasetVersion","sourceId":14967911,"datasetId":9580650,"databundleVersionId":15839791},{"sourceType":"datasetVersion","sourceId":14974763,"datasetId":9580712,"databundleVersionId":15847302},{"sourceType":"datasetVersion","sourceId":14969973,"datasetId":9582183,"databundleVersionId":15842038},{"sourceType":"datasetVersion","sourceId":14974876,"datasetId":9585454,"databundleVersionId":15847429},{"sourceType":"datasetVersion","sourceId":14746812,"datasetId":9421234,"databundleVersionId":15596411},{"sourceType":"datasetVersion","sourceId":14789097,"datasetId":9454727,"databundleVersionId":15643035},{"sourceType":"kernelVersion","sourceId":288631592}],"dockerImageVersionId":31236,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"cell-0","cell_type":"markdown","source":"# v30: 4-Group Ensemble Blend (Blur + Bright + LowContrast + Base)\n\n4-group quality-adaptive blending with per-model CTNorm:\n\n### Groups\n- **MODELS_BASE** (All): ensemble of base models, weighted average\n- **MODELS_BRIGHT** (D148): bright specialist (mean>=110), 50ep fine-tune\n- **MODELS_LC**: low-contrast specialist (contrast<130), CTNorm recalculated\n- **MODELS_BLUR** (D136): blur specialist (sharpness<400)\n\n### Per-model normalization (`norm_mode`)\n- `\"CTNorm\"` — standard CTNormalization (raw volume -> nnUNet)\n- `\"ct_persample\"` — remap intensity to global stats before CTNorm (v26 approach)\n- `\"zscore\"` — model uses ZScoreNorm internally, no preprocessing needed\n\n### Blend logic (4-way independent)\n```\nw_blur   = clip((400 - sharpness) / 150 + 0.5, 0, 1)\nw_bright = clip((mean - BRIGHT_THRESH) / (2*margin) + 0.5, 0, 1)\nw_lc     = clip((LC_THRESH - contrast) / (2*margin) + 0.5, 0, 1)\ntotal = w_blur + w_bright + w_lc\nif total > 1: normalize proportionally, w_base = 0\nelse: w_base = 1 - total\n```","metadata":{}},{"id":"cell-1","cell_type":"markdown","source":"## 1. Configuration","metadata":{}},{"id":"3b8a46af-ece8-4715-9ed1-ab5e7756938e","cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/vesuvius2-envs3')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T22:49:05.188593Z","iopub.execute_input":"2026-02-26T22:49:05.188935Z","iopub.status.idle":"2026-02-26T22:49:05.195603Z","shell.execute_reply.started":"2026-02-26T22:49:05.188889Z","shell.execute_reply":"2026-02-26T22:49:05.194979Z"}},"outputs":[],"execution_count":null},{"id":"cell-2","cell_type":"code","source":"from pathlib import Path\nimport torch\n\n# ===================== BASE ENSEMBLE =====================\nMODELS_BASE = [\n    {\n        \"name\": \"v24_expC15_curriculum_z128\",\n        \"model_dir\": Path(\"/kaggle/input/datasets/sugupoko/v24-expc15-curriculum-z128\"),\n        \"checkpoint\": \"checkpoint_soup.pth\",\n        \"dataset_id\": 121,\n        \"dataset_name\": \"Phase2\",\n        \"trainer\": \"nnUNetTrainerSkelRecall_500ep_LR005\",\n        \"plans\": \"nnUNetPlans\",\n        \"configuration\": \"3d_fullres\",\n        \"fold\": \"all\",\n        \"target_z\": 128,\n        \"fg_channel\": 1,\n        \"weight\": 1.0,\n        \"norm_mode\": \"CTNorm\",\n    },\n    {\n        \"name\": \"v21_expC11_z160\",\n        \"model_dir\": Path(\"/kaggle/input/datasets/sugupoko/v21-expc11-nnunet-curriculum-z2\"),\n        \"checkpoint\": \"checkpoint_best.pth\",\n        \"dataset_id\": 117,\n        \"dataset_name\": \"VesuviusNewThinHalf\",\n        \"trainer\": \"nnUNetTrainerMedialSurfaceRecall_1200epochs\",\n        \"plans\": \"nnUNetPlans\",\n        \"configuration\": \"3d_fullres\",\n        \"fold\": \"all\",\n        \"target_z\": 160,\n        \"fg_channel\": 1,\n        \"weight\": 1.0,\n        \"norm_mode\": \"CTNorm\",\n    },\n]\n\n# ===================== BRIGHT ENSEMBLE (D148) =====================\nMODELS_BRIGHT = [\n    {\n        \"name\": \"bright_d148_100ep\",\n        \"model_dir\": Path(\"/kaggle/input/datasets/sugupoko/dataset148-brightspecialist-100ep/nnUNetTrainerSkelRecall_GapWeight_3AxisSkel_WeakAug_100ep__nnUNetPlans__3d_fullres/fold_all\"),  # TODO: update after upload\n        \"checkpoint\": \"checkpoint_best.pth\",\n        \"dataset_id\": 148,\n        \"dataset_name\": \"BrightSpecialist\",\n        \"trainer\": \"nnUNetTrainerSkelRecall_GapWeight_3AxisSkel_WeakAug_100ep\",\n        \"plans\": \"nnUNetPlans\",\n        \"configuration\": \"3d_fullres\",\n        \"fold\": \"all\",\n        \"target_z\": 128,\n        \"fg_channel\": 1,\n        \"weight\": 1.0,\n        \"norm_mode\": \"CTNorm\",\n    },\n    # {\n    #     \"name\": \"bright_d150_100ep\",\n    #     \"model_dir\": Path(\"\"),\n    #     \"checkpoint\": \"checkpoint_best.pth\",\n    #     \"dataset_id\": 150,\n    #     \"dataset_name\": \"BlurOnly\",\n    #     \"trainer\": \"nnUNetTrainerSkelRecall_GapWeight_3AxisSkel_WeakAug_100ep\",\n    #     \"plans\": \"nnUNetPlans\",\n    #     \"configuration\": \"3d_fullres\",\n    #     \"fold\": \"all\",\n    #     \"target_z\": 160,\n    #     \"fg_channel\": 1,\n    #     \"weight\": 1.0,\n    #     \"norm_mode\": \"CTNorm\",\n    # },\n]\n\n# ===================== LOW-CONTRAST ENSEMBLE =====================\nMODELS_LC = [\n    {\n        \"name\": \"new_lowcontrast_z160_ft_soft\",\n        \"model_dir\": Path(\"/kaggle/input/datasets/sugupoko/dataset151-lowcontrastonly-z160/nnUNetTrainerSkelRecall_GapWeight_3AxisSkel_WeakAug_100ep__nnUNetPlans__3d_fullres/fold_all\"),\n        \"checkpoint\": \"checkpoint_best.pth\",\n        \"dataset_id\": 151,\n        \"dataset_name\": \"LowContrast\",\n        \"trainer\": \"nnUNetTrainerSkelRecall_GapWeight_3AxisSkel_WeakAug_100ep\",\n        \"plans\": \"nnUNetPlans\",\n        \"configuration\": \"3d_fullres\",\n        \"fold\": \"all\",\n        \"target_z\": 160,\n        \"fg_channel\": 1,\n        \"weight\": 1.0,\n        \"norm_mode\": \"CTNorm\",\n    },\n]\n\n# ===================== BLUR ENSEMBLE =====================\nMODELS_BLUR = [\n    {\n        \"name\": \"blur_d136_30ep\",\n        \"model_dir\": Path(\"/kaggle/input/datasets/sugupoko/dataset136-bluronly-nnunet-results-blur-30ep/nnUNetTrainerSkelRecall_GapWeight_3AxisSkel_WeakAug_30ep__nnUNetPlans__3d_fullres/fold_all\"),\n        \"checkpoint\": \"checkpoint_best.pth\",\n        \"dataset_id\": 136,\n        \"dataset_name\": \"BlurOnly\",\n        \"trainer\": \"nnUNetTrainerSkelRecall_GapWeight_3AxisSkel_WeakAug_30ep\",\n        \"plans\": \"nnUNetPlans\",\n        \"configuration\": \"3d_fullres\",\n        \"fold\": \"all\",\n        \"target_z\": 128,\n        \"fg_channel\": 1,\n        \"weight\": 1.0,\n        \"norm_mode\": \"CTNorm\",\n    },\n    {\n        \"name\": \"blur_d142_50ep\",\n        \"model_dir\": Path(\"/kaggle/input/datasets/sugupoko/dataset142-bluronly-z160/nnUNetTrainerSkelRecall_GapWeight_3AxisSkel_WeakAug__nnUNetPlans__3d_fullres/fold_all\"),\n        \"checkpoint\": \"checkpoint_best.pth\",\n        \"dataset_id\": 136,\n        \"dataset_name\": \"BlurOnly\",\n        \"trainer\": \"nnUNetTrainerSkelRecall_GapWeight_3AxisSkel_WeakAug\",\n        \"plans\": \"nnUNetPlans\",\n        \"configuration\": \"3d_fullres\",\n        \"fold\": \"all\",\n        \"target_z\": 160,\n        \"fg_channel\": 1,\n        \"weight\": 1.0,\n        \"norm_mode\": \"CTNorm\",\n    },\n]\n\n# Flatten for setup\nALL_MODELS = MODELS_BASE + MODELS_BRIGHT + MODELS_LC + MODELS_BLUR\n\n# ===================== PER-SAMPLE NORMALIZATION =====================\nPERSAMPLE_FG_MODE = \"percentile\"\nPERSAMPLE_FG_THRESH = 10\nPERSAMPLE_PERCENTILE_LO = 1.0\nPERSAMPLE_PERCENTILE_HI = 99.0\nGLOBAL_MEAN = 87.56\nGLOBAL_STD = 47.75\nPERSAMPLE_MIN_STD = 1.0\n\n# ===================== BLEND: BRIGHT (D148) =====================\nBLEND_BRIGHT_THRESH = 110\nBLEND_BRIGHT_MARGIN = 15      # ramp 95-125\n\n# ===================== BLEND: LOW-CONTRAST =====================\nBLEND_CONTRAST_THRESH = 130\nBLEND_CONTRAST_MARGIN = 20    # ramp 110-150\n\n# ===================== BLEND: BLUR =====================\nBLEND_SHARP_THRESH = 400\nBLEND_SHARP_RAMP = 150        # s=250->w=1.0, s=400->w=0.5, s=475->w=0.0\n\n# ===================== Z-SHIFT TTA =====================\nZ_SHIFT_N = 1  # OFF (tested: no improvement for z128)\n\n# ===================== SHARED SETTINGS =====================\nDATA_DIR = Path(\"/kaggle/input/vesuvius-challenge-surface-detection\")\nUSE_TTA = True\nTILE_STEP = 0.5\nTHRESHOLD = 0.5\n\nPP_MEDIAN_ITER = 6\nPP_MEDIAN_SIZE = 3\nPP_CLOSING_ITER = 0\nPP_FILL_HOLES = False\nPP_MIN_COMPONENT = 100\n\nSHARP_N_SLICES = 10\n\nNUM_GPUS = torch.cuda.device_count() if torch.cuda.is_available() else 1\nOUTPUT_DIR = Path(\"/kaggle/working\")\n\nprint(f\"4-group ensemble blend:\")\nprint(f\"  Base   ({len(MODELS_BASE)} models):   {[m['name'] for m in MODELS_BASE]}\")\nprint(f\"  Bright ({len(MODELS_BRIGHT)} models): {[m['name'] for m in MODELS_BRIGHT]}\")\nprint(f\"  LC     ({len(MODELS_LC)} models):     {[m['name'] for m in MODELS_LC]}\")\nprint(f\"  Blur   ({len(MODELS_BLUR)} models):   {[m['name'] for m in MODELS_BLUR]}\")\nprint(f\"  Total: {len(ALL_MODELS)} models\")\nprint(f\"  Bright blend: mean {BLEND_BRIGHT_THRESH}+/-{BLEND_BRIGHT_MARGIN}\")\nprint(f\"  LC blend:     contrast {BLEND_CONTRAST_THRESH}+/-{BLEND_CONTRAST_MARGIN}\")\nprint(f\"  Blur blend:   sharp T={BLEND_SHARP_THRESH}, R={BLEND_SHARP_RAMP}\")\nprint(f\"  Z-Shift TTA: {Z_SHIFT_N} ({'OFF' if Z_SHIFT_N <= 1 else f'{Z_SHIFT_N}x'})\")\nprint(f\"  GPUs: {NUM_GPUS}\")\nfor i in range(NUM_GPUS):\n    print(f\"    GPU {i}: {torch.cuda.get_device_name(i)}\")\nprint(f\"  Norm modes: {set(m['norm_mode'] for m in ALL_MODELS)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T22:49:05.197239Z","iopub.execute_input":"2026-02-26T22:49:05.197458Z","iopub.status.idle":"2026-02-26T22:49:09.148081Z","shell.execute_reply.started":"2026-02-26T22:49:05.197437Z","shell.execute_reply":"2026-02-26T22:49:09.147195Z"}},"outputs":[],"execution_count":null},{"id":"cell-3","cell_type":"markdown","source":"## 2. Setup","metadata":{}},{"id":"cell-4","cell_type":"code","source":"# !pip install -q nnunetv2 tifffile","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T22:49:09.149175Z","iopub.execute_input":"2026-02-26T22:49:09.149613Z","iopub.status.idle":"2026-02-26T22:49:09.153146Z","shell.execute_reply.started":"2026-02-26T22:49:09.149587Z","shell.execute_reply":"2026-02-26T22:49:09.152397Z"}},"outputs":[],"execution_count":null},{"id":"cell-5","cell_type":"code","source":"import os\nimport shutil\nimport zipfile\nimport types\nfrom io import BytesIO\nfrom concurrent.futures import ThreadPoolExecutor, as_completed\n\nimport numpy as np\nimport tifffile\nfrom PIL import Image\nfrom scipy.ndimage import (\n    zoom, binary_closing, binary_fill_holes, median_filter,\n    label, sum as ndimage_sum, generate_binary_structure,\n    laplace,\n)\nfrom tqdm.auto import tqdm\n\nprint(f\"PyTorch {torch.__version__}, CUDA {torch.cuda.is_available()}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T22:49:09.154342Z","iopub.execute_input":"2026-02-26T22:49:09.154599Z","iopub.status.idle":"2026-02-26T22:49:09.786244Z","shell.execute_reply.started":"2026-02-26T22:49:09.154579Z","shell.execute_reply":"2026-02-26T22:49:09.785480Z"}},"outputs":[],"execution_count":null},{"id":"cell-6","cell_type":"code","source":"os.environ[\"nnUNet_raw\"] = str(OUTPUT_DIR / \"nnunet_raw\")\nos.environ[\"nnUNet_preprocessed\"] = str(OUTPUT_DIR / \"nnunet_preprocessed\")\nos.environ[\"nnUNet_results\"] = str(OUTPUT_DIR / \"nnunet_model\")\n\n# -- Register custom trainers --\nfrom nnunetv2.training.nnUNetTrainer.nnUNetTrainer import nnUNetTrainer\n\nunique_trainers = {m[\"trainer\"] for m in ALL_MODELS}\ntrainer_stubs = {}\nfor tname in unique_trainers:\n    trainer_stubs[tname] = type(tname, (nnUNetTrainer,), {})\n\nimport nnunetv2.utilities.find_class_by_name as find_module\n\n_orig_code = find_module.recursive_find_python_class.__code__\n_orig_globals = find_module.recursive_find_python_class.__globals__.copy()\n_true_original = types.FunctionType(_orig_code, _orig_globals, 'recursive_find_python_class')\n\ndef _patched_find_class(folder, class_name, current_module):\n    if class_name in trainer_stubs:\n        return trainer_stubs[class_name]\n    return _true_original(folder, class_name, current_module)\n\nfind_module.recursive_find_python_class = _patched_find_class\nprint(f\"Registered {len(trainer_stubs)} trainer stub(s): {list(trainer_stubs.keys())}\")\n\n# -- Build nnUNet folder structures for all models --\nresults_dir = OUTPUT_DIR / \"nnunet_model\"\nmodel_folders = {}\n\nfor m in ALL_MODELS:\n    model_folder = (\n        results_dir\n        / f\"Dataset{m['dataset_id']}_{m['dataset_name']}\"\n        / f\"{m['trainer']}__{m['plans']}__{m['configuration']}\"\n    )\n    fold_folder = model_folder / f\"fold_{m['fold']}\"\n    fold_folder.mkdir(parents=True, exist_ok=True)\n\n    src_dir = m[\"model_dir\"]\n    shutil.copy(src_dir / m[\"checkpoint\"], fold_folder / m[\"checkpoint\"])\n\n    for fname in [\"plans.json\", \"dataset.json\"]:\n        for search_dir in [src_dir, src_dir.parent, src_dir.parent.parent]:\n            if (search_dir / fname).exists():\n                shutil.copy(search_dir / fname, model_folder / fname)\n                break\n\n    model_folders[m[\"name\"]] = model_folder\n    print(f\"  {m['name']}: {model_folder.relative_to(OUTPUT_DIR)} [norm={m['norm_mode']}]\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T22:49:09.788223Z","iopub.execute_input":"2026-02-26T22:49:09.788535Z","iopub.status.idle":"2026-02-26T22:50:21.587229Z","shell.execute_reply.started":"2026-02-26T22:49:09.788511Z","shell.execute_reply":"2026-02-26T22:50:21.586310Z"}},"outputs":[],"execution_count":null},{"id":"cell-7","cell_type":"markdown","source":"## 3. Initialize Predictors","metadata":{}},{"id":"cell-8","cell_type":"code","source":"from nnunetv2.inference.predict_from_raw_data import nnUNetPredictor\n\n# Load all models on EACH GPU for data-parallel inference\npredictors_per_gpu = {}  # gpu_id -> {model_name -> {\"predictor\", \"config\"}}\n\nfor gpu_id in range(NUM_GPUS):\n    predictors_per_gpu[gpu_id] = {}\n    for m in ALL_MODELS:\n        print(f\"Loading {m['name']} on GPU {gpu_id}...\", end=\" \", flush=True)\n\n        pred = nnUNetPredictor(\n            tile_step_size=TILE_STEP,\n            use_gaussian=True,\n            use_mirroring=USE_TTA,\n            perform_everything_on_device=True,\n            device=torch.device(\"cuda\", gpu_id),\n            verbose=False,\n            verbose_preprocessing=False,\n            allow_tqdm=False,\n        )\n        pred.initialize_from_trained_model_folder(\n            str(model_folders[m[\"name\"]]),\n            use_folds=(m[\"fold\"],),\n            checkpoint_name=m[\"checkpoint\"],\n        )\n\n        predictors_per_gpu[gpu_id][m[\"name\"]] = {\"predictor\": pred, \"config\": m}\n        print(\"done\")\n\nprint(f\"\\n{len(ALL_MODELS)} model(s) x {NUM_GPUS} GPU(s) = {len(ALL_MODELS) * NUM_GPUS} predictor(s) ready (TTA={USE_TTA})\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T22:50:21.588438Z","iopub.execute_input":"2026-02-26T22:50:21.589173Z","iopub.status.idle":"2026-02-26T22:51:07.177210Z","shell.execute_reply.started":"2026-02-26T22:50:21.589145Z","shell.execute_reply":"2026-02-26T22:51:07.176397Z"}},"outputs":[],"execution_count":null},{"id":"cell-9","cell_type":"markdown","source":"## 4. Helper Functions","metadata":{}},{"id":"cell-10","cell_type":"code","source":"def resize_z(volume: np.ndarray, target_z: int, order: int = 0) -> np.ndarray:\n    if volume.shape[0] == target_z:\n        return volume\n    return zoom(volume, (target_z / volume.shape[0], 1.0, 1.0), order=order)\n\n\ndef z_compress_with_shift(volume: np.ndarray, target_z: int, shift: int = 0) -> np.ndarray:\n    \"\"\"Z-compress by selecting every Nth slice with offset.\"\"\"\n    original_z = volume.shape[0]\n    ratio = original_z / target_z\n    if ratio <= 1:\n        return volume\n    indices = np.arange(target_z) * ratio + shift * (ratio / Z_SHIFT_N)\n    indices = np.clip(indices, 0, original_z - 1).astype(int)\n    return volume[indices]\n\n\ndef compute_sharpness(volume: np.ndarray) -> float:\n    Z = volume.shape[0]\n    indices = np.linspace(0, Z - 1, SHARP_N_SLICES, dtype=int)\n    variances = []\n    for zi in indices:\n        lap = laplace(volume[zi].astype(np.float32))\n        variances.append(float(np.var(lap)))\n    return float(np.mean(variances))\n\n\ndef compute_quality(volume: np.ndarray) -> dict:\n    Z = volume.shape[0]\n    slices_idx = [Z // 4, Z // 2, 3 * Z // 4]\n    contrasts, means = [], []\n    for zi in slices_idx:\n        slc = volume[zi].astype(np.float32)\n        p5, p95 = np.percentile(slc, [5, 95])\n        contrasts.append(p95 - p5)\n        means.append(float(np.mean(slc)))\n    return {\n        \"contrast\": np.mean(contrasts),\n        \"mean\": np.mean(means),\n        \"sharpness\": compute_sharpness(volume),\n    }\n\n\ndef normalize_persample(volume: np.ndarray) -> tuple:\n    vol_f = volume.astype(np.float32)\n    if PERSAMPLE_FG_MODE == \"percentile\":\n        lo = np.percentile(vol_f, PERSAMPLE_PERCENTILE_LO)\n        hi = np.percentile(vol_f, PERSAMPLE_PERCENTILE_HI)\n        fg_mask = (vol_f >= lo) & (vol_f <= hi)\n    elif PERSAMPLE_FG_MODE == \"threshold\":\n        fg_mask = vol_f > PERSAMPLE_FG_THRESH\n    else:\n        fg_mask = np.ones_like(vol_f, dtype=bool)\n\n    fg_vals = vol_f[fg_mask]\n    if len(fg_vals) < 100:\n        fg_vals = vol_f.ravel()\n\n    s_mean = float(np.mean(fg_vals))\n    s_std = max(float(np.std(fg_vals)), PERSAMPLE_MIN_STD)\n\n    adjusted = (vol_f - s_mean) / s_std * GLOBAL_STD + GLOBAL_MEAN\n    adjusted = np.clip(adjusted, 0, 255).astype(volume.dtype)\n    return adjusted, s_mean, s_std\n\n\ndef prepare_volume(volume_raw: np.ndarray, norm_mode: str) -> np.ndarray:\n    if norm_mode == \"ct_persample\":\n        vol, _, _ = normalize_persample(volume_raw)\n        return vol\n    else:  # \"CTNorm\", \"ct\", \"zscore\"\n        return volume_raw\n\n\ndef compute_blend_weights(quality: dict) -> tuple:\n    \"\"\"Compute 4-group blend weights (w_blur, w_bright, w_lc, w_base). Sum = 1.0.\"\"\"\n    # Blur weight: high when sharpness is low\n    w_blur = float(np.clip(\n        (BLEND_SHARP_THRESH - quality[\"sharpness\"]) / BLEND_SHARP_RAMP + 0.5,\n        0, 1,\n    ))\n\n    # Bright weight: high when mean intensity is high\n    w_bright = float(np.clip(\n        0.5 + (quality[\"mean\"] - BLEND_BRIGHT_THRESH) / (2 * BLEND_BRIGHT_MARGIN),\n        0, 1,\n    ))\n\n    # Low-contrast weight: high when contrast is low\n    w_lc = float(np.clip(\n        0.5 + (BLEND_CONTRAST_THRESH - quality[\"contrast\"]) / (2 * BLEND_CONTRAST_MARGIN),\n        0, 1,\n    ))\n\n    total_specialist = w_blur + w_bright + w_lc\n    if total_specialist > 1.0:\n        w_blur /= total_specialist\n        w_bright /= total_specialist\n        w_lc /= total_specialist\n        w_base = 0.0\n    else:\n        w_base = 1.0 - total_specialist\n\n    return w_blur, w_bright, w_lc, w_base\n\n\ndef predict_probs(volume: np.ndarray, entry: dict) -> np.ndarray:\n    cfg = entry[\"config\"]\n    predictor = entry[\"predictor\"]\n    original_z = volume.shape[0]\n    target_z = cfg[\"target_z\"]\n\n    n_shifts = max(1, Z_SHIFT_N)\n    prob_sum = np.zeros(volume.shape, dtype=np.float32)\n\n    for shift in range(n_shifts):\n        if n_shifts <= 1:\n            vol_z = resize_z(volume, target_z, order=0)\n        else:\n            vol_z = z_compress_with_shift(volume, target_z, shift=shift)\n\n        inp = vol_z[np.newaxis].astype(np.float32)\n        _seg, probs = predictor.predict_single_npy_array(\n            inp, {\"spacing\": (1.0, 1.0, 1.0)}, None, None, True\n        )\n        fg_prob = probs[cfg[\"fg_channel\"]].astype(np.float32)\n        prob_sum += resize_z(fg_prob, original_z, order=1)\n\n    return prob_sum / n_shifts\n\n\ndef predict_group_ensemble(volume_raw: np.ndarray, model_list: list, gpu_id: int) -> np.ndarray:\n    preds = predictors_per_gpu[gpu_id]\n    total_weight = sum(m[\"weight\"] for m in model_list)\n    probs = np.zeros(volume_raw.shape, dtype=np.float32)\n\n    for m in model_list:\n        vol = prepare_volume(volume_raw, m[\"norm_mode\"])\n        entry = preds[m[\"name\"]]\n        p = predict_probs(vol, entry)\n        probs += (m[\"weight\"] / total_weight) * p\n\n    return probs\n\n\ndef postprocess(mask: np.ndarray) -> np.ndarray:\n    mask = mask.copy()\n    for _ in range(PP_MEDIAN_ITER):\n        mask = (median_filter(mask, size=PP_MEDIAN_SIZE) > 0).astype(np.uint8)\n    if PP_CLOSING_ITER > 0:\n        struct = generate_binary_structure(3, 1)\n        for _ in range(PP_CLOSING_ITER):\n            mask = binary_closing(mask, structure=struct).astype(np.uint8)\n    if PP_FILL_HOLES:\n        mask = binary_fill_holes(mask).astype(np.uint8)\n    if PP_MIN_COMPONENT > 0:\n        struct_26 = np.ones((3, 3, 3), dtype=bool)\n        labeled, n_cc = label(mask, structure=struct_26)\n        if n_cc > 0:\n            sizes = ndimage_sum(mask, labeled, range(1, n_cc + 1))\n            for lbl in np.where(np.array(sizes) < PP_MIN_COMPONENT)[0] + 1:\n                mask[labeled == lbl] = 0\n    return mask.astype(np.uint8)\n\n\ndef volume_to_zipbytes(volume: np.ndarray) -> bytes:\n    pages = [Image.fromarray(s.astype(np.uint8)) for s in volume]\n    buf = BytesIO()\n    pages[0].save(buf, format=\"TIFF\", save_all=True, append_images=pages[1:])\n    return buf.getvalue()\n\n\ndef process_one(path: Path, gpu_id: int) -> tuple:\n    volume_raw = tifffile.imread(str(path))\n    quality = compute_quality(volume_raw)\n    w_blur, w_bright, w_lc, w_base = compute_blend_weights(quality)\n\n    EPS = 0.01\n\n    # Fast path: single dominant group\n    if w_base > 1 - EPS:\n        probs = predict_group_ensemble(volume_raw, MODELS_BASE, gpu_id)\n    elif w_blur > 1 - EPS:\n        probs = predict_group_ensemble(volume_raw, MODELS_BLUR, gpu_id)\n    elif w_bright > 1 - EPS:\n        probs = predict_group_ensemble(volume_raw, MODELS_BRIGHT, gpu_id)\n    elif w_lc > 1 - EPS:\n        probs = predict_group_ensemble(volume_raw, MODELS_LC, gpu_id)\n    else:\n        probs = np.zeros(volume_raw.shape, dtype=np.float32)\n        if w_base > EPS:\n            probs += w_base * predict_group_ensemble(volume_raw, MODELS_BASE, gpu_id)\n        if w_bright > EPS:\n            probs += w_bright * predict_group_ensemble(volume_raw, MODELS_BRIGHT, gpu_id)\n        if w_lc > EPS:\n            probs += w_lc * predict_group_ensemble(volume_raw, MODELS_LC, gpu_id)\n        if w_blur > EPS:\n            probs += w_blur * predict_group_ensemble(volume_raw, MODELS_BLUR, gpu_id)\n\n    mask = (probs >= THRESHOLD).astype(np.uint8)\n    mask = postprocess(mask)\n    fg_pct = mask.sum() / mask.size * 100\n    return path.stem, volume_to_zipbytes(mask), fg_pct, (w_blur, w_bright, w_lc, w_base), quality","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T22:51:07.178449Z","iopub.execute_input":"2026-02-26T22:51:07.178828Z","iopub.status.idle":"2026-02-26T22:51:07.205196Z","shell.execute_reply.started":"2026-02-26T22:51:07.178768Z","shell.execute_reply":"2026-02-26T22:51:07.204588Z"}},"outputs":[],"execution_count":null},{"id":"cell-11","cell_type":"markdown","source":"## 5. Run Inference","metadata":{}},{"id":"cell-12","cell_type":"code","source":"test_dir = DATA_DIR / \"test_images\"\ntest_files = sorted(test_dir.glob(\"*.tif\"))\nprint(f\"Test volumes: {len(test_files)}, GPUs: {NUM_GPUS}\")\nprint(f\"Models: base={len(MODELS_BASE)}, bright={len(MODELS_BRIGHT)}, lc={len(MODELS_LC)}, blur={len(MODELS_BLUR)}\")\nprint(f\"Blend Bright: mean {BLEND_BRIGHT_THRESH}+/-{BLEND_BRIGHT_MARGIN}\")\nprint(f\"Blend LC:     contrast {BLEND_CONTRAST_THRESH}+/-{BLEND_CONTRAST_MARGIN}\")\nprint(f\"Blend Blur:   sharp T={BLEND_SHARP_THRESH}, R={BLEND_SHARP_RAMP}\")\n\nsubmission_path = OUTPUT_DIR / \"submission.zip\"\n\nfrom concurrent.futures import ThreadPoolExecutor, as_completed\n\ngpu_assignments = [(path, i % NUM_GPUS) for i, path in enumerate(test_files)]\n\nweight_log = []\nresults = {}\n\nwith ThreadPoolExecutor(max_workers=NUM_GPUS) as pool:\n    futures = {\n        pool.submit(process_one, path, gpu_id): path\n        for path, gpu_id in gpu_assignments\n    }\n    for fut in tqdm(as_completed(futures), total=len(futures), desc=\"Inference\"):\n        path = futures[fut]\n        name, tiff_bytes, fg_pct, weights, quality = fut.result()\n        results[name] = (tiff_bytes, fg_pct, weights, quality)\n\n        w_blur, w_bright, w_lc, w_base = weights\n        weight_log.append(weights)\n\n        if w_base > 0.99:\n            tag = \"BASE\"\n        elif w_blur > 0.99:\n            tag = \"BLUR\"\n        elif w_bright > 0.99:\n            tag = \"BRIGHT\"\n        elif w_lc > 0.99:\n            tag = \"LC\"\n        else:\n            parts = []\n            if w_blur > 0.01:\n                parts.append(f\"blur={w_blur:.2f}\")\n            if w_bright > 0.01:\n                parts.append(f\"brt={w_bright:.2f}\")\n            if w_lc > 0.01:\n                parts.append(f\"lc={w_lc:.2f}\")\n            if w_base > 0.01:\n                parts.append(f\"base={w_base:.2f}\")\n            tag = \"BLEND \" + \" \".join(parts)\n\n        gpu_id = dict(gpu_assignments)[path]\n        print(f\"  {name}: [{tag}] GPU{gpu_id} c={quality['contrast']:.0f} m={quality['mean']:.0f} s={quality['sharpness']:.0f} fg={fg_pct:.1f}%\")\n\nwith zipfile.ZipFile(submission_path, \"w\", zipfile.ZIP_DEFLATED, compresslevel=9) as zf:\n    for path in test_files:\n        name = path.stem\n        tiff_bytes, _, _, _ = results[name]\n        zf.writestr(f\"{name}.tif\", tiff_bytes)\n\nsize_mb = submission_path.stat().st_size / (1024 * 1024)\nn_base = sum(1 for w in weight_log if w[3] > 0.99)\nn_blur = sum(1 for w in weight_log if w[0] > 0.99)\nn_bright = sum(1 for w in weight_log if w[1] > 0.99)\nn_lc = sum(1 for w in weight_log if w[2] > 0.99)\nn_blend = len(weight_log) - n_base - n_blur - n_bright - n_lc\n\nprint(f\"\\nSubmission: {submission_path} ({size_mb:.1f} MB)\")\nprint(f\"Routing: base={n_base}, bright={n_bright}, lc={n_lc}, blur={n_blur}, blended={n_blend}\")\nif weight_log:\n    w_arr = np.array(weight_log)\n    print(f\"Mean weights: blur={w_arr[:,0].mean():.3f}, bright={w_arr[:,1].mean():.3f}, lc={w_arr[:,2].mean():.3f}, base={w_arr[:,3].mean():.3f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T22:51:07.206189Z","iopub.execute_input":"2026-02-26T22:51:07.206985Z","iopub.status.idle":"2026-02-26T22:52:52.233242Z","shell.execute_reply.started":"2026-02-26T22:51:07.206949Z","shell.execute_reply":"2026-02-26T22:52:52.232350Z"}},"outputs":[],"execution_count":null},{"id":"cell-13","cell_type":"markdown","source":"## 6. Visualization (optional)","metadata":{}},{"id":"cell-14","cell_type":"code","source":"# import matplotlib.pyplot as plt\n\n\n# def show_prediction(image_path: Path, mask: np.ndarray, weights: tuple):\n#     vol = tifffile.imread(str(image_path))\n#     d, h, w = vol.shape\n#     cuts = {\n#         \"XY (z-mid)\": (vol[d // 2], mask[d // 2]),\n#         \"XZ (y-mid)\": (vol[:, h // 2], mask[:, h // 2]),\n#         \"YZ (x-mid)\": (vol[:, :, w // 2], mask[:, :, w // 2]),\n#     }\n\n#     fig, axes = plt.subplots(3, 2, figsize=(10, 12))\n#     for i, (title, (img, msk)) in enumerate(cuts.items()):\n#         axes[i, 0].imshow(img, cmap=\"gray\")\n#         axes[i, 0].set_title(f\"{title} - Image\")\n#         axes[i, 0].axis(\"off\")\n\n#         axes[i, 1].imshow(img, cmap=\"gray\")\n#         if msk.any():\n#             overlay = np.zeros((*msk.shape, 4))\n#             overlay[msk > 0] = [1, 0, 0, 0.4]\n#             axes[i, 1].imshow(overlay)\n#         axes[i, 1].set_title(f\"{title} - Prediction\")\n#         axes[i, 1].axis(\"off\")\n\n#     w_blur, w_bright, w_lc, w_base = weights\n#     fig.suptitle(f\"{image_path.stem} (blur={w_blur:.2f} brt={w_bright:.2f} lc={w_lc:.2f} base={w_base:.2f})\", fontsize=14)\n#     plt.tight_layout()\n#     plt.show()\n\n\n# if test_files:\n#     vol_raw = tifffile.imread(str(test_files[0]))\n#     quality = compute_quality(vol_raw)\n#     weights = compute_blend_weights(quality)\n#     w_blur, w_bright, w_lc, w_base = weights\n\n#     probs = np.zeros(vol_raw.shape, dtype=np.float32)\n#     if w_base > 0.01:\n#         probs += w_base * predict_group_ensemble(vol_raw, MODELS_BASE, 0)\n#     if w_bright > 0.01:\n#         probs += w_bright * predict_group_ensemble(vol_raw, MODELS_BRIGHT, 0)\n#     if w_lc > 0.01:\n#         probs += w_lc * predict_group_ensemble(vol_raw, MODELS_LC, 0)\n#     if w_blur > 0.01:\n#         probs += w_blur * predict_group_ensemble(vol_raw, MODELS_BLUR, 0)\n\n#     mask = postprocess((probs >= THRESHOLD).astype(np.uint8))\n#     show_prediction(test_files[0], mask, weights)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T22:52:52.234416Z","iopub.execute_input":"2026-02-26T22:52:52.234986Z","iopub.status.idle":"2026-02-26T22:52:52.239407Z","shell.execute_reply.started":"2026-02-26T22:52:52.234958Z","shell.execute_reply":"2026-02-26T22:52:52.238536Z"}},"outputs":[],"execution_count":null}]}