{"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":[{"sourceType":"competition","sourceId":117682,"databundleVersionId":15062069},{"sourceType":"datasetVersion","sourceId":14972112,"datasetId":9260081,"databundleVersionId":15844375},{"sourceType":"datasetVersion","sourceId":14921998,"datasetId":9253131,"databundleVersionId":15788703},{"sourceType":"datasetVersion","sourceId":14971909,"datasetId":9260044,"databundleVersionId":15844151},{"sourceType":"datasetVersion","sourceId":14976062,"datasetId":9536390,"databundleVersionId":15848775}],"dockerImageVersionId":31234,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# coding=utf-8\nimport sys\nimport zipfile\nimport glob\nimport torch\nimport os\nimport json\nimport numpy as np\nimport tqdm\nimport time\nimport subprocess\nfrom concurrent.futures import ThreadPoolExecutor, as_completed\n\n# sys.path.append(\"/kaggle/input/sdpackages/nnunet/\")\nsys.path.append(\"/kaggle/input/sdpackages/nnUNet_villa/nnUNet_villa\")\nsys.path.append(\"/kaggle/input/sdpackages/packages\")\nsys.path.append(\"/kaggle/input/sdinfer\")\nsys.path.append(\"/kaggle/input/sdinfer/SurfaceDetectionInfer\")\nUSE_LOCAL_ENV = False\nNOTEBOOK_DEBUG = False\nPOST_PROCESS = True\nPOST_PROCESS_MODE = \"adaptive_hysteresis_diffusion\"  # \"none\", \"hysteresis\", \"surfaceness\", \"adaptive_hysteresis\", \"adaptive_hysteresis_v2\", or \"adaptive_hysteresis_diffusion\"\nENABLE_ONNX = False\nPERFORM_EVERYTHING_ON_DEVICE = True\nFUSION_NUM_WORKERS = max(1, min(8, (os.cpu_count() or 1)))\nFUSION_METHOD = \"mean\"  # \"mean\" or \"weighted_logit\"\nFUSION_LOGIT_B = 0.0\nFUSION_LOGIT_EPS = 1e-4\nFUSION_WEIGHT_KEY = \"fusion_weight\"  # each model can set fusion_weight, default=1.0\nMODEL_PROB_CACHE_SUBDIR = \"prob_cache\"  # saved under each model output_path\n# Per-model optional keys in nnUNet_task612_models:\n# - fusion_weight: e.g. 1, 2, 3 (normalized automatically)\n# - infer_crop3: True/False (True means crop/pad border=3 during inference)\n# - infer_crop_border: int (overrides infer_crop3)\n\n# Post-process defaults (aligned with metric-side tuning)\nPP_NONE_THRESHOLD = 0.50\nPP_T_LOW = 0.40\nPP_T_HIGH = 0.90\nPP_Z_RADIUS = 1\nPP_XY_RADIUS = 0\nPP_DUST_MIN_SIZE = 100\nPP_SF_THRESHOLD = 0.75\nPP_SF_GAUSS_SIGMA = 2\nPP_SF_SIGMA = 6\nPP_SF_GAMMA = 1.5\nPP_SF_BETA1 = 0.5\nPP_SF_BETA2 = 0.5\n\n# Adaptive hysteresis gate:\n# start with hysteresis(z=1), fallback to hysteresis(z=0) only for risky morphology.\nPP_ADAPT_LCC_MAX = 0.088\nPP_ADAPT_NCC_NONE_MAX = 40\nPP_ADAPT_NCC_HYS_MIN = 17\nPP_ADAPT_FALLBACK_Z_RADIUS = 0\nPP_ADAPT_FALLBACK_XY_RADIUS = 0\n\n# Adaptive hysteresis v2: confidence filter and candidate guardrails.\nPP_ADAPT_V2_CF_APPLY_NCC_MIN = 8\nPP_ADAPT_V2_CF_MIN_SIZE = 100\nPP_ADAPT_V2_CF_MIN_MEAN_PROB = 0.60\nPP_ADAPT_V2_CF_MIN_Q90_PROB = 0.80\nPP_ADAPT_V2_CF_KEEP_TOPK = 0\nPP_ADAPT_V2_VOX_RATIO_MIN = 0.60\nPP_ADAPT_V2_VOX_RATIO_MAX = 1.40\n\n# Diffusion params for adaptive_hysteresis_diffusion.\nPP_DIFF_FILTER = \"gradient\"  # \"gradient\" or \"curvature\"\nPP_DIFF_ITERATIONS = 1\nPP_DIFF_CONDUCTANCE = 1.5\nPP_DIFF_TIME_STEP = 0.03\n\nif USE_LOCAL_ENV:\n    GPU_IDS = [2, 3]\n    # 对 case 进行切片，(0, 0) 表示全部\n    CASE_RANGE = (0, 0)\n    sys.path.insert(0, \"/home/lichuanpeng/Kaggle/MICCAI22/SurfaceDetection/nnUNet_villa\")\n\n    PATH_ORI_DATA = '/data/lichuanpeng/Kaggle/SurfaceDetection/vesuvius-challenge-surface-detection'\n    PATH_IMAGE_DATA = PATH_ORI_DATA + '/train_images'\n    working_path = '/data/lichuanpeng/Kaggle/SurfaceDetection/working2'\n    SUBMISSION_SAVE_PATH = working_path + \"/submission.zip\"\n    from model_configs_local import nnUNet_task612_models\nelse:\n    # 用于双卡分片推理的物理卡编号（每个模型会在这些卡上并行）\n    GPU_IDS = [0, 1]\n    # 对 case 进行切片，(0, 0) 表示全部\n    CASE_RANGE = (0, 0)\n    !pip install /kaggle/input/sdpackages/nnUNet_install/*.whl --no-deps\n    PATH_ORI_DATA = '/kaggle/input/vesuvius-challenge-surface-detection'\n    if NOTEBOOK_DEBUG:\n        # 复制10个测试\n        PATH_IMAGE_DATA = \"/kaggle/working/train_images\"\n        tif_files = glob.glob(PATH_ORI_DATA + '/train_images/*.tif')[:2]\n        os.makedirs(PATH_IMAGE_DATA, exist_ok=True)\n        for sub_tif in tif_files:\n            os.system(f\"cp {sub_tif} {PATH_IMAGE_DATA}/\")\n\n    else:\n        PATH_IMAGE_DATA = PATH_ORI_DATA + '/test_images'\n\n    working_path = '/kaggle/temp'\n    SUBMISSION_SAVE_PATH = 'submission.zip'\n    base_model_path = \"/kaggle/input/sdmodels\"\n\n    nnUNet_task612_models = [\n        # {\n        #     \"model_path\": base_model_path + '/nnUNetTrainerMedialSurfaceRecall__nnUNetPlansVillaBatch8__3d_fullres',\n        #     \"folds\": [\"all\"],\n        #     \"use_mirroring\": False,\n        #     \"mirroring_axes\": [0, 1,2 ],\n        #     \"step_size\": 0.5,\n        #     \"fusion_weight\": 1,\n        #     \"infer_crop3\":False,\n        #     \"checkpoint_name\": 'checkpoint_epoch997_fp16.pth',\n        #     \"output_path\": working_path + '/nnUNetTrainerMedialSurfaceRecall__nnUNetPlansVillaBatch8__3d_fullres'\n        # },\n        {\n            \"model_path\": base_model_path + '/nnUNetTrainerMedialSurfaceRecallEpoch500_ft__nnUNetPlansVillaPatch224Batch3__3d_fullres',\n            \"folds\": [\"all\"],\n            \"use_mirroring\": True,\n            \"mirroring_axes\": [0, 1, 2],\n            \"step_size\": 0.9,\n            \"fusion_weight\": 1,\n            \"infer_crop3\":False,\n            \"checkpoint_name\": 'checkpoint_epoch497_fp16.pth',\n            \"output_path\": working_path + '/nnUNetTrainerMedialSurfaceRecallEpoch500_ft__nnUNetPlansVillaPatch224Batch3__3d_fullres'\n        },\n        # {\n        #     \"model_path\": base_model_path + '/STUNetTrainer_large_MedialSurfaceRecall_500epochs__nnUNetPlansVilla__3d_fullres',\n        #     \"folds\": [\"all\"],\n        #     \"use_mirroring\": True,\n        #     \"mirroring_axes\": [0, 1, 2],\n        #     \"step_size\": 0.4,\n        #     \"checkpoint_name\": 'checkpoint_epoch498_fp16.pth',\n        #     \"output_path\": working_path + '/STUNetTrainer_large_MedialSurfaceRecall_500epochs__nnUNetPlansVilla__3d_fullres'\n        # },\n        {\n            \"model_path\": base_model_path + '/nnUNetTrainerMedialSurfaceRecallEpoch500_ft__nnUNetPlansVillaPatch256Batch6__3d_fullres',\n            \"folds\": [\"all\"],\n            \"use_mirroring\": True,\n            \"mirroring_axes\": [0, 1, 2],\n            \"step_size\": 0.9,\n            \"fusion_weight\": 1,\n            \"infer_crop3\":False,\n            \"checkpoint_name\": 'checkpoint_epoch500_fp16.pth',\n            \"output_path\": working_path + '/nnUNetTrainerMedialSurfaceRecallEpoch500_ft__nnUNetPlansVillaPatch256Batch6__3d_fullres'\n        },\n        # {\n        #     \"model_path\": base_model_path + '/nnUNetTrainerMedialSurfaceRecall_MedNeXt_M_kernel3_400epochs__nnUNetPlansVillaBatch6__3d_fullres',\n        #     \"folds\": [\"all\"],\n        #     \"use_mirroring\": True,\n        #     \"mirroring_axes\": [0, 1, 2],\n        #     \"step_size\": 0.5,\n        #     \"checkpoint_name\": 'checkpoint_epoch400_fp16.pth',\n        #     \"output_path\": working_path + '/nnUNetTrainerMedialSurfaceRecall_MedNeXt_M_kernel3_400epochs__nnUNetPlansVillaBBatch6__3d_fullres'\n        # },\n        # {\n        #     \"model_path\": base_model_path + '/nnUNetTrainerMedialSurfaceRecall_MedNeXt_L_kernel3_400epochs__nnUNetPlansVillaBBatch6__3d_fullres',\n        #     \"folds\": [\"all\"],\n        #     \"use_mirroring\": True,\n        #     \"mirroring_axes\": [0, 1, 2],\n        #     \"step_size\": 0.5,\n        #     \"checkpoint_name\": 'checkpoint_best.pth',\n        #     \"output_path\": working_path + '/nnUNetTrainerMedialSurfaceRecall_MedNeXt_L_kernel3_400epochs__nnUNetPlansVillaBBatch6__3d_fullres'\n        # },\n        # {\n        #     \"model_path\": base_model_path + '/DataSet613_nnUNetTrainerMedialSurfaceRecallEpoch200_ft__nnUNetPlansVillaPatch224Batch8__3d_fullres',\n        #     \"folds\": [\"all\"],\n        #     \"use_mirroring\": True,\n        #     \"mirroring_axes\": [0, 1, 2],\n        #     \"step_size\": 0.5,\n        #     \"checkpoint_name\": 'checkpoint_epoch111_fp16.pth',\n        #     \"output_path\": working_path + '/DataSet613_nnUNetTrainerMedialSurfaceRecallEpoch200_ft__nnUNetPlansVillaPatch224Batch8__3d_fullres'\n        # },\n        # {\n        #     \"model_path\": base_model_path + '/nnUNetTrainerMedialSurfaceRecallFocalTversky_300epochs_ft__nnUNetPlansVillaBatch8__3d_fullres',\n        #     \"folds\": [\"all\"],\n        #     \"use_mirroring\": True,\n        #     \"mirroring_axes\": [0, 1, 2],\n        #     \"step_size\": 0.5,\n        #     \"checkpoint_name\": 'checkpoint_epoch50_fp16.pth',\n        #     \"output_path\": working_path + '/nnUNetTrainerMedialSurfaceRecallFocalTversky_300epochs_ft__nnUNetPlansVillaBatch8__3d_fullres'\n        # },\n\n        # {\n        #     \"model_path\":  '/kaggle/input/datasets/lingyundev/sfdataset/Dataset615_nnUNetTrainerMedialSurfaceRecallEpoch200_ft__nnUNetPlansVillaPatch224Batch8__3d_fullres',\n        #     \"folds\": [\"all\"],\n        #     \"use_mirroring\": True,\n        #     \"mirroring_axes\": [0, 1, 2],\n        #     \"step_size\": 0.9,\n        #     \"fusion_weight\": 1,\n        #     \"infer_crop3\":True,\n        #     \"checkpoint_name\": 'checkpoint_epoch200_fp16.pth',\n        #     \"output_path\": working_path + '/Dataset615_nnUNetTrainerMedialSurfaceRecallEpoch200_ft__nnUNetPlansVillaPatch224Batch8__3d_fullres'\n        # },\n        # {\n        #     \"model_path\":  '/kaggle/input/datasets/lingyundev/sfdataset/Dataset615_nnUNetTrainerMedialSurfaceRecallEpoch200_ft__nnUNetPlansVillaPatch256Batch6__3d_fullres',\n        #     \"folds\": [\"all\"],\n        #     \"use_mirroring\": True,\n        #     \"mirroring_axes\": [0, 1, 2],\n        #     \"step_size\": 0.9,\n        #     \"fusion_weight\": 1,\n        #     \"infer_crop3\":True,\n        #     \"checkpoint_name\": 'checkpoint_epoch200_fp16.pth',\n        #     \"output_path\": working_path + '/Dataset615_nnUNetTrainerMedialSurfaceRecallEpoch200_ft__nnUNetPlansVillaPatch256Batch6__3d_fullres'\n        # },\n        {\n            \"model_path\":  '/kaggle/input/datasets/lingyundev/sfdataset/Dataset615_nnUNetTrainerMedialSurfaceRecallEpoch500_ft__nnUNetPlansVillaPatch288Batch4__3d_fullres',\n            \"folds\": [\"all\"],\n            \"use_mirroring\": False,\n            \"mirroring_axes\": [0, 1, 2],\n            \"step_size\": 0.5,\n            \"fusion_weight\": 1,\n            \"infer_crop3\":True,\n            \"checkpoint_name\": 'checkpoint_epoch438_fp16.pth',\n            \"output_path\": working_path + '/Dataset615_nnUNetTrainerMedialSurfaceRecallEpoch500_ft__nnUNetPlansVillaPatch288Batch4__3d_fullres'\n        }\n    ]\n\nfrom Tools.data_precess import convert_tif_to_nii_gz, convert_nii_gz_to_tif\nfrom common_tools.dicom.nifti import read_nii_with_param, save_nii_with_param\n\ntry:\n    from postprocess_utils import postprocess_fg_prob, postprocess_fg_prob_adaptive_hysteresis, postprocess_fg_prob_adaptive_hysteresis_v2, postprocess_fg_prob_adaptive_hysteresis_diffusion, postprocess_fg_prob_surfaceness\nexcept ImportError:\n    from SurfaceDetectionInfer.postprocess_utils import postprocess_fg_prob, postprocess_fg_prob_adaptive_hysteresis, postprocess_fg_prob_adaptive_hysteresis_v2, postprocess_fg_prob_adaptive_hysteresis_diffusion, postprocess_fg_prob_surfaceness\n\n\ndef _prob_to_mask(avg_prob):\n    if POST_PROCESS and POST_PROCESS_MODE != \"none\":\n        if POST_PROCESS_MODE == \"surfaceness\":\n            return postprocess_fg_prob_surfaceness(\n                avg_prob,\n                threshold=PP_SF_THRESHOLD,\n                gauss_sigma=PP_SF_GAUSS_SIGMA,\n                sigma=PP_SF_SIGMA,\n                gamma=PP_SF_GAMMA,\n                beta1=PP_SF_BETA1,\n                beta2=PP_SF_BETA2,\n                dust_min_size=PP_DUST_MIN_SIZE,\n            )\n        if POST_PROCESS_MODE == \"hysteresis\":\n            return postprocess_fg_prob(\n                avg_prob,\n                T_low=PP_T_LOW,\n                T_high=PP_T_HIGH,\n                z_radius=PP_Z_RADIUS,\n                xy_radius=PP_XY_RADIUS,\n                dust_min_size=PP_DUST_MIN_SIZE,\n            )\n        if POST_PROCESS_MODE == \"adaptive_hysteresis\":\n            return postprocess_fg_prob_adaptive_hysteresis(\n                avg_prob,\n                T_low=PP_T_LOW,\n                T_high=PP_T_HIGH,\n                z_radius=PP_Z_RADIUS,\n                xy_radius=PP_XY_RADIUS,\n                dust_min_size=PP_DUST_MIN_SIZE,\n                none_threshold=PP_NONE_THRESHOLD,\n                adapt_lcc_max=PP_ADAPT_LCC_MAX,\n                adapt_ncc_none_max=PP_ADAPT_NCC_NONE_MAX,\n                adapt_ncc_hys_min=PP_ADAPT_NCC_HYS_MIN,\n                adapt_fallback_z_radius=PP_ADAPT_FALLBACK_Z_RADIUS,\n                adapt_fallback_xy_radius=PP_ADAPT_FALLBACK_XY_RADIUS,\n            )\n        if POST_PROCESS_MODE == \"adaptive_hysteresis_v2\":\n            return postprocess_fg_prob_adaptive_hysteresis_v2(\n                avg_prob,\n                T_low=PP_T_LOW,\n                T_high=PP_T_HIGH,\n                z_radius=PP_Z_RADIUS,\n                xy_radius=PP_XY_RADIUS,\n                dust_min_size=PP_DUST_MIN_SIZE,\n                none_threshold=PP_NONE_THRESHOLD,\n                adapt_lcc_max=PP_ADAPT_LCC_MAX,\n                adapt_ncc_none_max=PP_ADAPT_NCC_NONE_MAX,\n                adapt_ncc_hys_min=PP_ADAPT_NCC_HYS_MIN,\n                adapt_fallback_z_radius=PP_ADAPT_FALLBACK_Z_RADIUS,\n                adapt_fallback_xy_radius=PP_ADAPT_FALLBACK_XY_RADIUS,\n                cf_apply_ncc_min=PP_ADAPT_V2_CF_APPLY_NCC_MIN,\n                cf_min_size=PP_ADAPT_V2_CF_MIN_SIZE,\n                cf_min_mean_prob=PP_ADAPT_V2_CF_MIN_MEAN_PROB,\n                cf_min_q90_prob=PP_ADAPT_V2_CF_MIN_Q90_PROB,\n                cf_keep_topk=PP_ADAPT_V2_CF_KEEP_TOPK,\n                vox_ratio_min=PP_ADAPT_V2_VOX_RATIO_MIN,\n                vox_ratio_max=PP_ADAPT_V2_VOX_RATIO_MAX,\n            )\n        if POST_PROCESS_MODE == \"adaptive_hysteresis_diffusion\":\n            return postprocess_fg_prob_adaptive_hysteresis_diffusion(\n                avg_prob,\n                T_low=PP_T_LOW,\n                T_high=PP_T_HIGH,\n                z_radius=PP_Z_RADIUS,\n                xy_radius=PP_XY_RADIUS,\n                dust_min_size=PP_DUST_MIN_SIZE,\n                none_threshold=PP_NONE_THRESHOLD,\n                adapt_lcc_max=PP_ADAPT_LCC_MAX,\n                adapt_ncc_none_max=PP_ADAPT_NCC_NONE_MAX,\n                adapt_ncc_hys_min=PP_ADAPT_NCC_HYS_MIN,\n                adapt_fallback_z_radius=PP_ADAPT_FALLBACK_Z_RADIUS,\n                adapt_fallback_xy_radius=PP_ADAPT_FALLBACK_XY_RADIUS,\n                diff_filter=PP_DIFF_FILTER,\n                diff_iterations=PP_DIFF_ITERATIONS,\n                diff_conductance=PP_DIFF_CONDUCTANCE,\n                diff_time_step=PP_DIFF_TIME_STEP,\n            )\n        raise NotImplementedError\n    return (avg_prob >= PP_NONE_THRESHOLD).astype(np.uint8)\n\n\ndef _sum_prob_path(cache_dir, case_name):\n    return os.path.join(cache_dir, f\"{case_name}_sum.npy\")\n\n\ndef _model_prob_dir(model_cfg):\n    output_path = str(model_cfg.get(\"output_path\", \"\")).strip()\n    if not output_path:\n        raise ValueError(\"weighted_logit requires each model to define non-empty output_path\")\n    return os.path.join(output_path, MODEL_PROB_CACHE_SUBDIR)\n\n\ndef _model_prob_path(model_prob_dir, case_name):\n    return os.path.join(model_prob_dir, f\"{case_name}.npy\")\n\n\ndef _safe_logit(prob, eps):\n    p = np.clip(prob.astype(np.float32), eps, 1.0 - eps)\n    return np.log(p) - np.log1p(-p)\n\n\ndef _resolve_fusion_plan(models):\n    method = str(FUSION_METHOD).lower().strip()\n    if method in (\"\", \"mean\", \"avg\", \"average\"):\n        return {\"mode\": \"mean\"}\n    if method != \"weighted_logit\":\n        raise ValueError(f\"Unsupported FUSION_METHOD={FUSION_METHOD}\")\n\n    raw_weights = []\n    model_prob_dirs = []\n    for idx, model in enumerate(models):\n        raw = float(model.get(FUSION_WEIGHT_KEY, 1.0))\n        if raw < 0:\n            raise ValueError(f\"model[{idx}] {FUSION_WEIGHT_KEY} must be >= 0, got {raw}\")\n        raw_weights.append(raw)\n        model_prob_dirs.append(_model_prob_dir(model))\n    total = float(sum(raw_weights))\n    if total <= 0:\n        raise ValueError(f\"sum({FUSION_WEIGHT_KEY}) must be > 0, got {raw_weights}\")\n    weights = [w / total for w in raw_weights]\n    return {\n        \"mode\": \"weighted_logit\",\n        \"weights\": weights,\n        \"raw_weights\": raw_weights,\n        \"b\": float(FUSION_LOGIT_B),\n        \"eps\": float(FUSION_LOGIT_EPS),\n        \"model_prob_dirs\": model_prob_dirs,\n    }\n\n\ndef _load_fused_prob(case_name, prob_cache_dir, num_models, fusion_plan):\n    if fusion_plan[\"mode\"] == \"mean\":\n        sum_path = _sum_prob_path(prob_cache_dir, case_name)\n        if not os.path.exists(sum_path):\n            raise FileNotFoundError(f\"missing prob cache: {sum_path}\")\n        sum_prob = np.load(sum_path).astype(np.float32)\n        return sum_prob / float(num_models)\n\n    fused_logit = None\n    for weight, model_prob_dir in zip(fusion_plan[\"weights\"], fusion_plan[\"model_prob_dirs\"]):\n        if weight <= 0:\n            continue\n        model_prob_path = _model_prob_path(model_prob_dir, case_name)\n        if not os.path.exists(model_prob_path):\n            raise FileNotFoundError(f\"missing model prob cache: {model_prob_path}\")\n        p = np.load(model_prob_path).astype(np.float32)\n        cur_logit = _safe_logit(p, fusion_plan[\"eps\"])\n        if fused_logit is None:\n            fused_logit = weight * cur_logit\n        else:\n            fused_logit += weight * cur_logit\n        del p, cur_logit\n\n    if fused_logit is None:\n        raise ValueError(\"weighted_logit has no active model weights (>0)\")\n    fused_logit += fusion_plan[\"b\"]\n    return 1.0 / (1.0 + np.exp(-np.clip(fused_logit, -20.0, 20.0)))\n\n\ndef _fuse_one_case(sub_image, prob_cache_dir, final_output_path, num_models, fusion_plan):\n    case_t0 = time.time()\n    case_name = os.path.basename(sub_image).split('_00')[0]\n    t_read0 = time.time()\n    _image_array, origin, spacing, direction = read_nii_with_param(sub_image)\n    read_cost = time.time() - t_read0\n    t_load0 = time.time()\n    avg_prob = _load_fused_prob(case_name, prob_cache_dir, num_models, fusion_plan)\n    load_cost = time.time() - t_load0\n    t_pp0 = time.time()\n    mask_array = _prob_to_mask(avg_prob)\n    pp_cost = time.time() - t_pp0\n    t_save0 = time.time()\n    save_nii_with_param(\n        mask_array.astype(np.uint8),\n        origin,\n        spacing,\n        direction,\n        os.path.join(final_output_path, f\"{case_name}.nii.gz\")\n    )\n    save_cost = time.time() - t_save0\n    case_cost = time.time() - case_t0\n    del _image_array, avg_prob, mask_array\n    return case_name, read_cost, load_cost, pp_cost, save_cost, case_cost\n\n\ndef _worker_script_path():\n    local_script = \"/home/lichuanpeng/Kaggle/MICCAI22/SurfaceDetectionInfer/infer_model_worker.py\"\n    kaggle_script = \"/kaggle/input/sdinfer/infer_model_worker.py\"\n    if os.path.exists(local_script):\n        return local_script\n    if os.path.exists(kaggle_script):\n        return kaggle_script\n    raise FileNotFoundError(\n        \"infer_model_worker.py not found. Tried:\\n\"\n        f\"1) {local_script}\\n\"\n        f\"2) {kaggle_script}\"\n    )\n\n\ndef _launch_dual_gpu_workers(\n        model_idx,\n        sub_model,\n        data_nii_path,\n        prob_cache_dir,\n        enable_onnx=False,\n        perform_everything_on_device=False,\n        write_sum_cache=True,\n        model_prob_dir=\"\",\n):\n    available = torch.cuda.device_count()\n    if available <= 0:\n        raise RuntimeError(\"No CUDA device found\")\n    gpu_ids = [g for g in GPU_IDS if g < available]\n    if len(gpu_ids) == 0:\n        gpu_ids = [0]\n    num_shards = len(gpu_ids)\n\n    worker_script = _worker_script_path()\n    print(f\"[INFO] worker_script={worker_script}\")\n\n    model_cfg_str = json.dumps(sub_model, ensure_ascii=False, separators=(\",\", \":\"))\n    procs = []\n    for shard_id, gpu_id in enumerate(gpu_ids):\n        cmd = [\n            sys.executable,\n            worker_script,\n            \"--model_config_str\", model_cfg_str,\n            \"--model_tag\", f\"{model_idx}:{sub_model['checkpoint_name']}\",\n            \"--data_nii_path\", data_nii_path,\n            \"--prob_cache_dir\", prob_cache_dir,\n            \"--case_range\", str(CASE_RANGE[0]), str(CASE_RANGE[1]),\n            \"--shard_id\", str(shard_id),\n            \"--num_shards\", str(num_shards),\n        ]\n        if write_sum_cache:\n            cmd.append(\"--write_sum_cache\")\n        if model_prob_dir:\n            cmd.extend([\"--model_prob_dir\", model_prob_dir])\n        if enable_onnx:\n            cmd.append(\"--enable_onnx\")\n        if perform_everything_on_device:\n            cmd.append(\"--perform_everything_on_device\")\n        if write_sum_cache and model_idx == 0:\n            cmd.append(\"--is_first_model\")\n        env = os.environ.copy()\n        env[\"CUDA_VISIBLE_DEVICES\"] = str(gpu_id)\n        print(\"[LAUNCH]\", \" \".join(cmd), f\"(CUDA_VISIBLE_DEVICES={gpu_id})\")\n        procs.append(subprocess.Popen(cmd, env=env))\n\n    for p in procs:\n        p.wait()\n        if p.returncode != 0:\n            raise RuntimeError(f\"worker failed, returncode={p.returncode}\")\n\n\ntest_img_paths = glob.glob(PATH_IMAGE_DATA + '/*.tif', recursive=True)\n\nif len(test_img_paths) == 1 and NOTEBOOK_DEBUG is False:\n    print('ignore')\n    with zipfile.ZipFile(SUBMISSION_SAVE_PATH, mode=\"w\") as zf:\n        pass\nelse:\n    print(\"Start\")\n    data_nii_path = working_path + '/images_nii'\n    convert_tif_to_nii_gz(PATH_IMAGE_DATA, data_nii_path, num_workers=8)\n    case_paths = sorted(glob.glob(data_nii_path + '/*.nii.gz', recursive=True))\n\n    final_output_path = os.path.join(working_path, \"ensemble_output\")\n    prob_cache_dir = os.path.join(working_path, \"ensemble_prob_cache\")\n    os.makedirs(final_output_path, exist_ok=True)\n    os.makedirs(prob_cache_dir, exist_ok=True)\n    fusion_plan = _resolve_fusion_plan(nnUNet_task612_models)\n    print(f\"[INFO] Fusion plan: {fusion_plan}\")\n    for old_sum in glob.glob(os.path.join(prob_cache_dir, \"*_sum.npy\")):\n        os.remove(old_sum)\n    if fusion_plan[\"mode\"] == \"weighted_logit\":\n        for model_prob_dir in fusion_plan[\"model_prob_dirs\"]:\n            os.makedirs(model_prob_dir, exist_ok=True)\n            for old_prob in glob.glob(os.path.join(model_prob_dir, \"*.npy\")):\n                os.remove(old_prob)\n\n    print(f\"Start model-serial inference, models={len(nnUNet_task612_models)}, cases={len(case_paths)}\")\n    time_s = time.time()\n    for model_idx, sub_model in enumerate(nnUNet_task612_models):\n        print(f\"Model[{model_idx}] start: {sub_model['checkpoint_name']}\")\n        model_time_s = time.time()\n        _launch_dual_gpu_workers(\n            model_idx=model_idx,\n            sub_model=sub_model,\n            data_nii_path=data_nii_path,\n            prob_cache_dir=prob_cache_dir,\n            enable_onnx=ENABLE_ONNX,\n            perform_everything_on_device=PERFORM_EVERYTHING_ON_DEVICE,\n            write_sum_cache=(fusion_plan[\"mode\"] == \"mean\"),\n            model_prob_dir=(\n                fusion_plan[\"model_prob_dirs\"][model_idx]\n                if fusion_plan[\"mode\"] == \"weighted_logit\"\n                else \"\"\n            ),\n        )\n        print(f\"Model[{model_idx}] done, cost: {time.time() - model_time_s:.2f}s\")\n\n    print(f\"Start final fusion, workers={FUSION_NUM_WORKERS}\")\n    num_models = len(nnUNet_task612_models)\n    if FUSION_NUM_WORKERS <= 1:\n        for idx, sub_image in enumerate(case_paths):\n            case_name, read_cost, load_cost, pp_cost, save_cost, case_cost = _fuse_one_case(\n                sub_image=sub_image,\n                prob_cache_dir=prob_cache_dir,\n                final_output_path=final_output_path,\n                num_models=num_models,\n                fusion_plan=fusion_plan,\n            )\n            print(\n                f\"Fusion case[{idx + 1}/{len(case_paths)}] {case_name}: \"\n                f\"read={read_cost:.2f}s load={load_cost:.2f}s pp={pp_cost:.2f}s \"\n                f\"save={save_cost:.2f}s total={case_cost:.2f}s\"\n            )\n            if (idx + 1) % 5 == 0 or (idx + 1) == len(case_paths):\n                print(f\"Fusion progress: {idx + 1}/{len(case_paths)}\")\n    else:\n        done = 0\n        with ThreadPoolExecutor(max_workers=FUSION_NUM_WORKERS) as ex:\n            futures = [\n                ex.submit(\n                    _fuse_one_case,\n                    sub_image,\n                    prob_cache_dir,\n                    final_output_path,\n                    num_models,\n                    fusion_plan,\n                )\n                for sub_image in case_paths\n            ]\n            for fut in as_completed(futures):\n                case_name, read_cost, load_cost, pp_cost, save_cost, case_cost = fut.result()\n                done += 1\n                print(\n                    f\"Fusion case[{done}/{len(case_paths)}] {case_name}: \"\n                    f\"read={read_cost:.2f}s load={load_cost:.2f}s pp={pp_cost:.2f}s \"\n                    f\"save={save_cost:.2f}s total={case_cost:.2f}s\"\n                )\n                if done % 5 == 0 or done == len(case_paths):\n                    print(f\"Fusion progress: {done}/{len(case_paths)}\")\n\n    print(\"ensemble infer cost:\", time.time() - time_s)\n    convert_nii_gz_to_tif(final_output_path, working_path, num_workers=8)\n    if not USE_LOCAL_ENV:\n        with zipfile.ZipFile(SUBMISSION_SAVE_PATH, 'w', zipfile.ZIP_DEFLATED) as zipf:\n            for filename in tqdm.tqdm(glob.glob(working_path + \"/*.tif\"), desc=\"Zipping files\"):\n                zipf.write(filename)\n                os.remove(filename)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-27T10:19:14.116225Z","iopub.execute_input":"2026-02-27T10:19:14.116592Z","iopub.status.idle":"2026-02-27T10:19:17.107390Z","shell.execute_reply.started":"2026-02-27T10:19:14.116559Z","shell.execute_reply":"2026-02-27T10:19:17.106663Z"}},"outputs":[],"execution_count":null}]}