refactor(nnu): transition to native volume processing for Pipeline B
Update the nnU-Net pipeline to operate on native (unprocessed) volumes instead of preprocessed ones. This allows nnU-Net to utilize its own `plan_and_preprocess` logic for resampling, cropping, and normalization, ensuring Pipeline B remains distinct from Pipelines A and C. - Update `scripts/05_build_splits.py` to include native `img` and `label` paths in the manifest rows. - Modify `scripts_nnu/01_nnu_prepare_dataset.py` to symlink native volumes and implement a label fix-up mechanism for non-conforming grids. - Update `scripts_nnu/04_nnu_pseudo_label.py` and `05_nnu_eval_test.py` to use native image paths and perform selection/evaluation in physical mm on the native grid. - Refactor `scripts_nnu/nnu_common.py` to handle native-grid label alignment and volume-based selection gates. - Update `README.md` to document the preprocessing differences between Pipelines A/C and Pipeline B.
This commit is contained in:
parent
77adc2b3af
commit
695bbb6e11
7 changed files with 194 additions and 68 deletions
49
README.md
49
README.md
|
|
@ -6,8 +6,10 @@ large unlabeled pool from the same subjects/series to improve a held-out,
|
||||||
patient-level test Dice.
|
patient-level test Dice.
|
||||||
|
|
||||||
All volumes are preprocessed to a uniform 1 mm isotropic, head-cropped,
|
All volumes are preprocessed to a uniform 1 mm isotropic, head-cropped,
|
||||||
percentile-normalized form so that tumor counts are directly comparable in mm³
|
percentile-normalized form (Pipelines A and C) so that tumor counts are directly
|
||||||
across sources and timepoints.
|
comparable in mm³ across sources and timepoints. Pipeline B instead feeds the
|
||||||
|
native volumes straight into nnU-Net, which performs its own preprocessing
|
||||||
|
(see below).
|
||||||
|
|
||||||
## Environment
|
## Environment
|
||||||
|
|
||||||
|
|
@ -123,22 +125,43 @@ pool, and evaluation protocol so results are directly comparable to Pipeline A.
|
||||||
> This installed nnU-Net is a modern fork: preprocessed data is `.b2nd`
|
> This installed nnU-Net is a modern fork: preprocessed data is `.b2nd`
|
||||||
> (blosc2), checkpoints store `network_weights` (not `model`), and
|
> (blosc2), checkpoints store `network_weights` (not `model`), and
|
||||||
> `--save_probabilities` writes a per-case `<case>.npz` (channel-first
|
> `--save_probabilities` writes a per-case `<case>.npz` (channel-first
|
||||||
> `(C, z, y, x)`). The code below targets that build, not the upstream nnU-Net
|
> `(C, z, y, x)`, resampled back to the original **native** input grid —
|
||||||
> docs.
|
> so are the exported argmax segs). The code below targets that build, not the
|
||||||
|
> upstream nnU-Net docs.
|
||||||
|
|
||||||
Dataset: `Dataset210_NTUH_T1C_PL`, single channel `t1c`, 2 classes
|
Dataset: `Dataset210_NTUH_T1C_PL`, single channel `t1c`, 2 classes
|
||||||
(background=0, tumor=1), `3d_fullres` only, `nnUNetPlans`, fold 0.
|
(background=0, tumor=1), `3d_fullres` only, `nnUNetPlans`, fold 0.
|
||||||
|
|
||||||
| # | Script | Purpose |
|
| # | Script | Purpose |
|
||||||
|---|---|---|
|
|---|---|---|
|
||||||
| 01 | `01_nnu_prepare_dataset.py` | Build/rebuild the raw dataset (symlinked `imagesTr`/`labelsTr` + `dataset.json`) from a rows jsonl |
|
| 01 | `01_nnu_prepare_dataset.py` | Build/rebuild the raw dataset (symlinked **native** `imagesTr`, `labelsTr` 0/1 + `dataset.json`) from a rows jsonl |
|
||||||
| 02 | `02_nnu_plan_preprocess.py` | Plan + preprocess (`--clean`), then write subject-level `splits_final.json` |
|
| 02 | `02_nnu_plan_preprocess.py` | Plan + preprocess (`--clean`), then write subject-level `splits_final.json` |
|
||||||
| 03 | `03_nnu_train.py` | Train one round via `NTUHLPLTrainer` (DDP `-num_gpus`), optional warm start |
|
| 03 | `03_nnu_train.py` | Train one round via `NTUHLPLTrainer` (DDP `-num_gpus`), optional warm start |
|
||||||
| 04 | `04_nnu_pseudo_label.py` | Multi-GPU `nnUNetv2_predict` on the pool + selection gates |
|
| 04 | `04_nnu_pseudo_label.py` | Multi-GPU `nnUNetv2_predict` on native pool volumes + selection gates (physical mm) |
|
||||||
| 05 | `05_nnu_eval_test.py` | Held-out test evaluation (probability Dice@0.5 + hard-seg Dice) |
|
| 05 | `05_nnu_eval_test.py` | Held-out test evaluation on native volumes (probability Dice@0.5 + hard-seg Dice) |
|
||||||
| 06 | `06_nnu_run_iterative.py` | Orchestrates rounds 0…K, table + plot |
|
| 06 | `06_nnu_run_iterative.py` | Orchestrates rounds 0…K, table + plot |
|
||||||
| — | `trainers/ntuh_pl_trainer.py` | `NTUHLPLTrainer`: env-driven epochs/LR + full-weight warm start |
|
| — | `trainers/ntuh_pl_trainer.py` | `NTUHLPLTrainer`: env-driven epochs/LR + full-weight warm start |
|
||||||
| — | `nnu_common.py` | Shared paths, env, dataset/split helpers, selection + consistency |
|
| — | `nnu_common.py` | Shared paths, env, dataset/split helpers, native-grid selection + consistency |
|
||||||
|
|
||||||
|
### Preprocessing (native inputs)
|
||||||
|
|
||||||
|
Pipeline B does **no** resampling, cropping, or intensity normalization before
|
||||||
|
nnU-Net. The raw dataset (`01`) symlinks the native T1c NIfTIs from the source
|
||||||
|
manifests (`img` field) into `imagesTr`; `labelsTr` entries are symlinks when
|
||||||
|
the label is already 0/1 on the image grid, otherwise nearest-warped +
|
||||||
|
binarized copies (label fix-up only). All geometric/intensity preprocessing is
|
||||||
|
then nnU-Net's own `plan_and_preprocess`: nonzero-bbox crop, resample to the
|
||||||
|
plan's median-based target spacing, per-channel normalization, `.b2nd` storage.
|
||||||
|
The split/pool row jsonls therefore carry both the processed
|
||||||
|
(`pimg`/`plabel`, Pipelines A/C) and native (`img`/`label`, Pipeline B) paths —
|
||||||
|
rebuild them with `scripts/05_build_splits.py`.
|
||||||
|
|
||||||
|
Downstream of the model (selection gates, consistency filter, evaluation) work
|
||||||
|
in each case's **native grid** with physical units: `nnUNetv2_predict` itself
|
||||||
|
already returns the probability map (and the argmax seg) resampled back to the
|
||||||
|
native input grid, so no extra remapping is needed; tumor volumes are voxel
|
||||||
|
counts × native voxel volume (mm³); centroid offsets are scaled to mm by the
|
||||||
|
native spacing.
|
||||||
|
|
||||||
### Round semantics
|
### Round semantics
|
||||||
|
|
||||||
|
|
@ -205,10 +228,12 @@ python scripts_nnu/05_nnu_eval_test.py --rows data/manifests/split_test.jsonl \
|
||||||
|
|
||||||
### Verified
|
### Verified
|
||||||
|
|
||||||
End-to-end smoke-tested on a scratch 4-case dataset: dataset build → planning →
|
End-to-end smoke-tested on a scratch 4-case native dataset (mixed anisotropic
|
||||||
`splits_final.json` → training 1 epoch → warm start → 2-way sharded prediction →
|
spacing, multi-class + off-grid labels): dataset build (native symlinks +
|
||||||
selection gates → test eval. The consistency filter's keep / reject /
|
label fixup) → planning → `splits_final.json` → training 1 epoch → growing
|
||||||
single-timepoint paths are unit-tested with synthetic volumes.
|
re-plan/re-preprocess → sharded native prediction → selection gates in mm →
|
||||||
|
test eval. The consistency filter's keep / reject / single-timepoint paths are
|
||||||
|
unit-tested with synthetic volumes.
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -5,6 +5,12 @@ Outputs:
|
||||||
data/manifests/split_train.jsonl split_val.jsonl split_test.jsonl (labeled rows, w=1.0)
|
data/manifests/split_train.jsonl split_val.jsonl split_test.jsonl (labeled rows, w=1.0)
|
||||||
data/manifests/unlabeled_pool.jsonl (m6 + lee, labeled subjects removed)
|
data/manifests/unlabeled_pool.jsonl (m6 + lee, labeled subjects removed)
|
||||||
data/vols.json (labeled tumor volume stats, mm3)
|
data/vols.json (labeled tumor volume stats, mm3)
|
||||||
|
|
||||||
|
Each row carries both processed fields (pimg/plabel, data/proc — 1mm cropped/
|
||||||
|
normalized, used by Pipelines A and C) and native fields (img/label, the raw
|
||||||
|
source NIfTIs — used by Pipeline B, which feeds nnU-Net its own
|
||||||
|
plan_and_preprocess). Row membership still requires the processed volume to
|
||||||
|
exist, so the patient-level splits stay identical across pipelines.
|
||||||
"""
|
"""
|
||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
|
|
@ -21,7 +27,8 @@ def proc_row(r, prefix):
|
||||||
p = os.path.join(d("data/proc"), key + ".nii.gz")
|
p = os.path.join(d("data/proc"), key + ".nii.gz")
|
||||||
if not os.path.exists(p):
|
if not os.path.exists(p):
|
||||||
return None
|
return None
|
||||||
row = {"key": key, "subject": f"{prefix}_{r['subject']}", "pimg": p}
|
row = {"key": key, "subject": f"{prefix}_{r['subject']}", "pimg": p,
|
||||||
|
"img": r["img"], "label": r.get("label"), "date": r.get("date")}
|
||||||
lab = os.path.join(d("data/proc"), key + "_label.nii.gz")
|
lab = os.path.join(d("data/proc"), key + "_label.nii.gz")
|
||||||
if r.get("label") and os.path.exists(lab):
|
if r.get("label") and os.path.exists(lab):
|
||||||
row["plabel"] = lab
|
row["plabel"] = lab
|
||||||
|
|
@ -88,7 +95,7 @@ def main():
|
||||||
proc = os.path.join(d("data/proc"), r["key"] + ".nii.gz")
|
proc = os.path.join(d("data/proc"), r["key"] + ".nii.gz")
|
||||||
if os.path.exists(proc):
|
if os.path.exists(proc):
|
||||||
p = {"key": r["key"], "subject": f"lee_{r['sid']}", "pimg": proc,
|
p = {"key": r["key"], "subject": f"lee_{r['sid']}", "pimg": proc,
|
||||||
"date": r["date"], "source": "lee"}
|
"img": nii, "label": None, "date": r["date"], "source": "lee"}
|
||||||
lee_rows[r["key"]] = p
|
lee_rows[r["key"]] = p
|
||||||
pool.extend(lee_rows.values())
|
pool.extend(lee_rows.values())
|
||||||
pool.sort(key=lambda r: (r["subject"], r.get("date", "")))
|
pool.sort(key=lambda r: (r["subject"], r.get("date", "")))
|
||||||
|
|
|
||||||
|
|
@ -1,7 +1,11 @@
|
||||||
"""Build/rebuild the nnU-Net raw dataset (Dataset210_NTUH_T1C_PL) from rows jsonl.
|
"""Build/rebuild the nnU-Net raw dataset (Dataset210_NTUH_T1C_PL) from rows jsonl.
|
||||||
|
|
||||||
rows: {key, pimg, label} where label is an absolute path to a 0/1 nifti mask.
|
rows: {key, img, label} where img is the NATIVE (unprocessed) T1c nifti and
|
||||||
Rebuilds imagesTr/labelsTr as symlinks and writes dataset.json.
|
label a nifti mask on (or warpable to) that grid. No resampling/cropping/
|
||||||
|
normalization is applied here — nnUNetv2_plan_and_preprocess (02) does the
|
||||||
|
preprocessing; imagesTr entries are symlinks to the native volumes, labelsTr
|
||||||
|
entries are symlinks when already 0/1 on the image grid, else nearest-warped
|
||||||
|
+ binarized copies (label fix-up only).
|
||||||
|
|
||||||
Usage: python scripts_nnu/01_nnu_prepare_dataset.py --rows <rows.jsonl>
|
Usage: python scripts_nnu/01_nnu_prepare_dataset.py --rows <rows.jsonl>
|
||||||
"""
|
"""
|
||||||
|
|
@ -16,7 +20,7 @@ from nnu_common import make_raw_dataset, raw_ds
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
ap = argparse.ArgumentParser()
|
ap = argparse.ArgumentParser()
|
||||||
ap.add_argument("--rows", required=True, help="jsonl with {key, pimg, label}")
|
ap.add_argument("--rows", required=True, help="jsonl with {key, img, label} (native paths)")
|
||||||
args = ap.parse_args()
|
args = ap.parse_args()
|
||||||
rows = load_jsonl(args.rows)
|
rows = load_jsonl(args.rows)
|
||||||
n = make_raw_dataset(rows)
|
n = make_raw_dataset(rows)
|
||||||
|
|
|
||||||
|
|
@ -1,15 +1,20 @@
|
||||||
"""Pseudo-label the unlabeled pool with an nnU-Net round model (multi-GPU predict).
|
"""Pseudo-label the unlabeled pool with an nnU-Net round model (multi-GPU predict).
|
||||||
|
|
||||||
1. Symlinks remaining pool volumes (not already in the dataset) into a predict folder.
|
1. Symlinks remaining pool volumes (not already in the dataset) into a predict
|
||||||
|
folder. Volumes are NATIVE (unresampled/uncropped/unnormalized); nnU-Net
|
||||||
|
applies the plan's preprocessing itself at predict time.
|
||||||
2. Runs one nnUNetv2_predict per GPU (-num_parts/-part_id), --save_probabilities.
|
2. Runs one nnUNetv2_predict per GPU (-num_parts/-part_id), --save_probabilities.
|
||||||
3. Selection gates (same as scripts/06_pseudo_label.py):
|
nnU-Net applies the plan's preprocessing at predict time and resamples the
|
||||||
|
probabilities (and argmax seg) back to each case's native grid, so the
|
||||||
|
selection gates run directly on that grid, in physical mm:
|
||||||
pos: p_tumor >= tau_pos, median cleanup, largest-CC fraction >= min_cc_frac,
|
pos: p_tumor >= tau_pos, median cleanup, largest-CC fraction >= min_cc_frac,
|
||||||
volume within [vol_p2, vol_p98] of labeled tumor volumes (data/vols.json)
|
volume within [vol_p2, vol_p98] of labeled tumor volumes (data/vols.json)
|
||||||
neg: >= neg_frac of interior (vol > 0.02) voxels have p_bg >= tau_neg
|
neg: >= neg_frac of head-interior voxels have p_bg >= tau_neg
|
||||||
4. Per-subject longitudinal consistency filter over accepted positive timepoints
|
4. Per-subject longitudinal consistency filter over accepted positive timepoints
|
||||||
(head-relative tumor centroid distance + volume ratio; absolute patient-space
|
(head-relative tumor centroid distance in mm + volume ratio; absolute
|
||||||
grids are not comparable across acquisitions).
|
patient-space grids are not comparable across acquisitions).
|
||||||
Writes <out>/rows.jsonl, <out>/<key>_label.nii.gz, <out>/accepted.jsonl, <out>/summary.json.
|
Writes <out>/rows.jsonl, <out>/<key>_label.nii.gz (native grid), <out>/accepted.jsonl,
|
||||||
|
<out>/summary.json.
|
||||||
|
|
||||||
Usage:
|
Usage:
|
||||||
python scripts_nnu/04_nnu_pseudo_label.py --pool data/manifests/unlabeled_pool.jsonl \
|
python scripts_nnu/04_nnu_pseudo_label.py --pool data/manifests/unlabeled_pool.jsonl \
|
||||||
|
|
@ -34,7 +39,7 @@ def build_input_folder(rows, in_dir):
|
||||||
shutil.rmtree(in_dir, ignore_errors=True)
|
shutil.rmtree(in_dir, ignore_errors=True)
|
||||||
os.makedirs(in_dir, exist_ok=True)
|
os.makedirs(in_dir, exist_ok=True)
|
||||||
for r in rows:
|
for r in rows:
|
||||||
os.symlink(os.path.abspath(r["pimg"]), os.path.join(in_dir, r["key"] + "_0000.nii.gz"))
|
os.symlink(os.path.abspath(r["img"]), os.path.join(in_dir, r["key"] + "_0000.nii.gz"))
|
||||||
|
|
||||||
|
|
||||||
def find_outputs(out_dir, ext):
|
def find_outputs(out_dir, ext):
|
||||||
|
|
@ -74,6 +79,8 @@ def main():
|
||||||
outname = os.path.basename(os.path.normpath(args.out))
|
outname = os.path.basename(os.path.normpath(args.out))
|
||||||
out = d(args.out)
|
out = d(args.out)
|
||||||
pool_rows = load_jsonl(args.pool)
|
pool_rows = load_jsonl(args.pool)
|
||||||
|
if any("img" not in r for r in pool_rows):
|
||||||
|
raise SystemExit(f"{args.pool} rows lack native 'img'; rerun scripts/05_build_splits.py")
|
||||||
already = set()
|
already = set()
|
||||||
if args.already and os.path.exists(args.already):
|
if args.already and os.path.exists(args.already):
|
||||||
already = {r["key"] for r in load_jsonl(args.already)}
|
already = {r["key"] for r in load_jsonl(args.already)}
|
||||||
|
|
@ -124,26 +131,27 @@ def main():
|
||||||
for r in rows:
|
for r in rows:
|
||||||
key = r["key"]
|
key = r["key"]
|
||||||
entry = {"key": key, "subject": r["subject"], "date": r.get("date"),
|
entry = {"key": key, "subject": r["subject"], "date": r.get("date"),
|
||||||
"source": r.get("source"), "pimg": r["pimg"], "label": None,
|
"source": r.get("source"), "img": r["img"], "pimg": r.get("pimg"), "label": None,
|
||||||
"role": "rej", "vol_mm3": 0, "maxp": None}
|
"role": "rej", "vol_mm3": 0, "maxp": None}
|
||||||
try:
|
try:
|
||||||
pimg_itk = sitk.ReadImage(r["pimg"])
|
img_itk = sitk.ReadImage(r["img"])
|
||||||
vol = read_nii_arr(r["pimg"]).astype("float32")
|
vol = read_nii_arr(r["img"]).astype("float32")
|
||||||
pt = tumor_probs_from_npz(npz[key])
|
pt = tumor_probs_from_npz(npz[key])
|
||||||
if pt.shape != vol.shape:
|
if pt.shape != vol.shape:
|
||||||
raise ValueError(f"prob shape {pt.shape} != image shape {vol.shape} for {key}")
|
raise ValueError(f"prob shape {pt.shape} != native image {vol.shape} for {key}")
|
||||||
entry["maxp"] = round(float(pt.max()), 4)
|
entry["maxp"] = round(float(pt.max()), 4)
|
||||||
got = pos_mask(pt, args.tau_pos, args.min_cc_frac, vol_lo, vol_hi)
|
voxel_vol_mm3 = float(np.prod(np.asarray(img_itk.GetSpacing(), dtype=np.float64)))
|
||||||
|
got = pos_mask(pt, args.tau_pos, args.min_cc_frac, vol_lo, vol_hi, voxel_vol_mm3)
|
||||||
if got is not None:
|
if got is not None:
|
||||||
mask, cc_frac, vol_mm3 = got
|
mask, cc_frac, vol_mm3 = got
|
||||||
lp = os.path.join(out, key + "_label.nii.gz")
|
lp = os.path.join(out, key + "_label.nii.gz")
|
||||||
write_label(mask, lp, pimg_itk)
|
write_label(mask, lp, img_itk)
|
||||||
entry.update({"role": "pos", "label": lp, "vol_mm3": vol_mm3, "cc_frac": round(cc_frac, 3)})
|
entry.update({"role": "pos", "label": lp, "vol_mm3": vol_mm3, "cc_frac": round(cc_frac, 3)})
|
||||||
else:
|
else:
|
||||||
frac = neg_frac_bg(pt, vol, args.tau_neg, args.neg_frac)
|
frac = neg_frac_bg(pt, vol, args.tau_neg, args.neg_frac)
|
||||||
if frac is not None:
|
if frac is not None:
|
||||||
lp = os.path.join(out, key + "_label.nii.gz")
|
lp = os.path.join(out, key + "_label.nii.gz")
|
||||||
write_arr(np.zeros(vol.shape, dtype="uint8"), lp, itk_img=pimg_itk)
|
write_arr(np.zeros(vol.shape, dtype="uint8"), lp, itk_img=img_itk)
|
||||||
entry.update({"role": "neg", "label": lp, "neg_conf": round(frac, 4)})
|
entry.update({"role": "neg", "label": lp, "neg_conf": round(frac, 4)})
|
||||||
except Exception as e: # noqa
|
except Exception as e: # noqa
|
||||||
entry["role"] = "error"
|
entry["role"] = "error"
|
||||||
|
|
@ -153,7 +161,7 @@ def main():
|
||||||
|
|
||||||
n_pos0 = sum(1 for x in rows_out if x["role"] == "pos")
|
n_pos0 = sum(1 for x in rows_out if x["role"] == "pos")
|
||||||
n_neg0 = sum(1 for x in rows_out if x["role"] == "neg")
|
n_neg0 = sum(1 for x in rows_out if x["role"] == "neg")
|
||||||
n_rej = consistency_filter(rows_out, out, args.max_rel_dist, args.vol_ratio)
|
n_rej = consistency_filter(rows_out, out, args.max_rel_dist, args.vol_ratio, image_key="img")
|
||||||
for x in rows_out:
|
for x in rows_out:
|
||||||
if x["role"] == "rej":
|
if x["role"] == "rej":
|
||||||
x["label"] = None
|
x["label"] = None
|
||||||
|
|
|
||||||
|
|
@ -1,9 +1,13 @@
|
||||||
"""Holdout test evaluation of an nnU-Net round model.
|
"""Holdout test evaluation of an nnU-Net round model.
|
||||||
|
|
||||||
Predicts the test volumes (multi-GPU sharding, --save_probabilities) and scores:
|
Predicts the test volumes (NATIVE inputs, multi-GPU sharding,
|
||||||
|
--save_probabilities). nnU-Net applies the plan's preprocessing at predict time
|
||||||
|
and returns both the tumor probability and the argmax seg on each case's native
|
||||||
|
grid, so we score directly there:
|
||||||
dice - tumor probability (softmax channel 1) thresholded at 0.5
|
dice - tumor probability (softmax channel 1) thresholded at 0.5
|
||||||
(same convention as scripts/08_eval.py)
|
(same convention as scripts/08_eval.py)
|
||||||
dice_hard - argmax segmentation written by nnUNetv2_predict
|
dice_hard - argmax segmentation written by nnUNetv2_predict
|
||||||
|
Ground truth is the row's native label.
|
||||||
Writes: <out> json + <out>.json -> per_row jsonl alongside.
|
Writes: <out> json + <out>.json -> per_row jsonl alongside.
|
||||||
|
|
||||||
Usage: python scripts_nnu/05_nnu_eval_test.py --rows data/manifests/split_test.jsonl \
|
Usage: python scripts_nnu/05_nnu_eval_test.py --rows data/manifests/split_test.jsonl \
|
||||||
|
|
@ -15,6 +19,7 @@ import sys
|
||||||
import json
|
import json
|
||||||
import shutil
|
import shutil
|
||||||
import subprocess
|
import subprocess
|
||||||
|
import SimpleITK as sitk
|
||||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||||
from src.common import load_jsonl, save_jsonl, read_nii_arr, dice
|
from src.common import load_jsonl, save_jsonl, read_nii_arr, dice
|
||||||
|
|
@ -35,10 +40,12 @@ def main():
|
||||||
tag = os.path.basename(os.path.normpath(args.out)) + "_predict"
|
tag = os.path.basename(os.path.normpath(args.out)) + "_predict"
|
||||||
in_dir = os.path.join(d("nnu/pred"), f"in_test_{tag}")
|
in_dir = os.path.join(d("nnu/pred"), f"in_test_{tag}")
|
||||||
out_dir = os.path.join(d("nnu/pred"), f"out_test_{tag}")
|
out_dir = os.path.join(d("nnu/pred"), f"out_test_{tag}")
|
||||||
|
if any("img" not in r or "label" not in r for r in rows):
|
||||||
|
raise SystemExit("split rows lack native img/label; rerun scripts/05_build_splits.py")
|
||||||
shutil.rmtree(in_dir, ignore_errors=True)
|
shutil.rmtree(in_dir, ignore_errors=True)
|
||||||
os.makedirs(in_dir, exist_ok=True)
|
os.makedirs(in_dir, exist_ok=True)
|
||||||
for r in rows:
|
for r in rows:
|
||||||
os.symlink(os.path.abspath(r["pimg"]), os.path.join(in_dir, r["key"] + "_0000.nii.gz"))
|
os.symlink(os.path.abspath(r["img"]), os.path.join(in_dir, r["key"] + "_0000.nii.gz"))
|
||||||
|
|
||||||
env = nnu_env()
|
env = nnu_env()
|
||||||
procs = []
|
procs = []
|
||||||
|
|
@ -80,17 +87,26 @@ def main():
|
||||||
per, dice_p, dice_h = [], [], []
|
per, dice_p, dice_h = [], [], []
|
||||||
for r in rows:
|
for r in rows:
|
||||||
key = r["key"]
|
key = r["key"]
|
||||||
lab = read_nii_arr(r.get("plabel") or r.get("label"))
|
|
||||||
try:
|
try:
|
||||||
|
img_itk = sitk.ReadImage(r["img"])
|
||||||
|
vol = read_nii_arr(r["img"]).astype("float32")
|
||||||
|
lab_path = r.get("label")
|
||||||
|
if not lab_path:
|
||||||
|
raise ValueError("row lacks native label; rerun scripts/05_build_splits.py")
|
||||||
|
lab = sitk.GetArrayFromImage(sitk.ReadImage(lab_path))
|
||||||
|
if lab.shape != vol.shape:
|
||||||
|
lab = sitk.GetArrayFromImage(sitk.Resample(sitk.ReadImage(lab_path), img_itk,
|
||||||
|
sitk.Transform(), sitk.sitkNearestNeighbor, 0.0))
|
||||||
|
lab = (lab > 0).astype("uint8")
|
||||||
pt = tumor_probs_from_npz(npz[key])
|
pt = tumor_probs_from_npz(npz[key])
|
||||||
if pt.shape != lab.shape:
|
if pt.shape != vol.shape:
|
||||||
raise ValueError(f"prob {pt.shape} vs label {lab.shape}")
|
raise ValueError(f"prob {pt.shape} vs native image {vol.shape}")
|
||||||
d1 = dice((pt >= 0.5).astype("uint8"), (lab > 0).astype("uint8"))
|
d1 = dice((pt >= 0.5).astype("uint8"), lab)
|
||||||
dice_p.append(d1)
|
dice_p.append(d1)
|
||||||
hard = read_nii_arr(segs[key])
|
hard = read_nii_arr(segs[key])
|
||||||
if hard.shape != lab.shape:
|
if hard.shape != lab.shape:
|
||||||
raise ValueError(f"hard seg {hard.shape} vs label {lab.shape}")
|
raise ValueError(f"hard seg {hard.shape} vs native label {lab.shape}")
|
||||||
d2 = dice((hard > 0).astype("uint8"), (lab > 0).astype("uint8"))
|
d2 = dice((hard > 0).astype("uint8"), lab)
|
||||||
dice_h.append(d2)
|
dice_h.append(d2)
|
||||||
per.append({"key": key, "dice": round(d1, 4), "dice_hard": round(d2, 4)})
|
per.append({"key": key, "dice": round(d1, 4), "dice_hard": round(d2, 4)})
|
||||||
except Exception as e: # noqa
|
except Exception as e: # noqa
|
||||||
|
|
|
||||||
|
|
@ -54,7 +54,10 @@ def main():
|
||||||
|
|
||||||
base = []
|
base = []
|
||||||
for f in (train_f, val_f):
|
for f in (train_f, val_f):
|
||||||
base += [{"key": r["key"], "pimg": r["pimg"], "label": r["plabel"]} for r in load_jsonl(f)]
|
for r in load_jsonl(f):
|
||||||
|
if "img" not in r or not r.get("label"):
|
||||||
|
raise SystemExit(f"{f} rows lack native img/label; rerun scripts/05_build_splits.py")
|
||||||
|
base.append({"key": r["key"], "img": r["img"], "label": r["label"]})
|
||||||
save_jsonl(base, os.path.join(pdir_root, "base_rows.jsonl"))
|
save_jsonl(base, os.path.join(pdir_root, "base_rows.jsonl"))
|
||||||
print(f"[nnu-orch] base dataset rows: {len(base)}", flush=True)
|
print(f"[nnu-orch] base dataset rows: {len(base)}", flush=True)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -12,7 +12,7 @@ import subprocess
|
||||||
|
|
||||||
ROOT = os.environ.get("LONGITUDINAL_ROOT", "/mnt/b4/xfr/git26/longitudinal")
|
ROOT = os.environ.get("LONGITUDINAL_ROOT", "/mnt/b4/xfr/git26/longitudinal")
|
||||||
sys.path.insert(0, ROOT)
|
sys.path.insert(0, ROOT)
|
||||||
from src.common import load_jsonl, save_jsonl, read_nii_arr # noqa: E402
|
from src.common import load_jsonl, save_jsonl, read_nii_arr, head_mask_from_image, largest_cc # noqa: E402
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import SimpleITK as sitk
|
import SimpleITK as sitk
|
||||||
|
|
@ -128,9 +128,43 @@ def predict_cmd(in_dir, out_dir, gpus, part, chk="checkpoint_best.pth", tta=True
|
||||||
|
|
||||||
|
|
||||||
# ---------------- raw dataset (nnU-Net format) ----------------
|
# ---------------- raw dataset (nnU-Net format) ----------------
|
||||||
|
#
|
||||||
|
# Pipeline B feeds nnU-Net the NATIVE volumes (no resampling / cropping /
|
||||||
|
# intensity normalization beforehand): nnUNetv2_plan_and_preprocess does all of
|
||||||
|
# that itself (nonzero-bbox crop, resample to the plan's target spacing,
|
||||||
|
# per-channel normalization, .b2nd storage).
|
||||||
|
|
||||||
|
def same_frame(a, b):
|
||||||
|
return (a.GetSize() == b.GetSize() and a.GetSpacing() == b.GetSpacing()
|
||||||
|
and a.GetOrigin() == b.GetOrigin() and a.GetDirection() == b.GetDirection())
|
||||||
|
|
||||||
|
|
||||||
|
def native_label_target(img_path, lab_path, link_path):
|
||||||
|
"""Label target for a raw-dataset row: 0/1 on the image's native grid.
|
||||||
|
|
||||||
|
Symlinks the label when it already sits on the image grid with values
|
||||||
|
{0, 1}; otherwise writes a converted copy to link_path (nearest-neighbor
|
||||||
|
warp to the image grid + binarization to 0/1). This is label fix-up only —
|
||||||
|
no common-grid resampling, cropping, or normalization.
|
||||||
|
"""
|
||||||
|
img_itk = sitk.ReadImage(img_path)
|
||||||
|
lab_itk = sitk.ReadImage(lab_path)
|
||||||
|
if same_frame(img_itk, lab_itk):
|
||||||
|
vals = set(np.unique(sitk.GetArrayFromImage(lab_itk)).tolist())
|
||||||
|
if vals <= {0, 1}:
|
||||||
|
os.symlink(os.path.abspath(lab_path), link_path)
|
||||||
|
return
|
||||||
|
warped = sitk.Resample(lab_itk, img_itk, sitk.Transform(), sitk.sitkNearestNeighbor, 0.0)
|
||||||
|
out = sitk.GetImageFromArray((sitk.GetArrayFromImage(warped) > 0).astype(np.uint8))
|
||||||
|
out.CopyInformation(img_itk)
|
||||||
|
os.makedirs(os.path.dirname(link_path), exist_ok=True)
|
||||||
|
sitk.WriteImage(out, link_path, True)
|
||||||
|
|
||||||
|
|
||||||
def make_raw_dataset(rows, channel="0000"):
|
def make_raw_dataset(rows, channel="0000"):
|
||||||
"""rows: list of {key, pimg, label}. Rebuilds imagesTr/labelsTr symlinks + dataset.json."""
|
"""rows: list of {key, img, label} — img/label are native (unprocessed) NIfTI
|
||||||
|
paths. Rebuilds imagesTr/labelsTr + dataset.json. Images are symlinked
|
||||||
|
as-is; labels via native_label_target. See section note."""
|
||||||
img_dir = os.path.join(raw_ds(), "imagesTr")
|
img_dir = os.path.join(raw_ds(), "imagesTr")
|
||||||
lab_dir = os.path.join(raw_ds(), "labelsTr")
|
lab_dir = os.path.join(raw_ds(), "labelsTr")
|
||||||
shutil.rmtree(img_dir, ignore_errors=True)
|
shutil.rmtree(img_dir, ignore_errors=True)
|
||||||
|
|
@ -139,7 +173,7 @@ def make_raw_dataset(rows, channel="0000"):
|
||||||
os.makedirs(lab_dir, exist_ok=True)
|
os.makedirs(lab_dir, exist_ok=True)
|
||||||
seen = set()
|
seen = set()
|
||||||
for r in rows:
|
for r in rows:
|
||||||
key, img, lab = r["key"], r["pimg"], r["label"]
|
key, img, lab = r["key"], r["img"], r["label"]
|
||||||
if key in seen:
|
if key in seen:
|
||||||
raise ValueError(f"duplicate case id in dataset: {key}")
|
raise ValueError(f"duplicate case id in dataset: {key}")
|
||||||
seen.add(key)
|
seen.add(key)
|
||||||
|
|
@ -148,7 +182,7 @@ def make_raw_dataset(rows, channel="0000"):
|
||||||
if not os.path.exists(lab):
|
if not os.path.exists(lab):
|
||||||
raise FileNotFoundError(f"label missing for {key}: {lab}")
|
raise FileNotFoundError(f"label missing for {key}: {lab}")
|
||||||
os.symlink(os.path.abspath(img), os.path.join(img_dir, f"{key}_{channel}.nii.gz"))
|
os.symlink(os.path.abspath(img), os.path.join(img_dir, f"{key}_{channel}.nii.gz"))
|
||||||
os.symlink(os.path.abspath(lab), os.path.join(lab_dir, f"{key}.nii.gz"))
|
native_label_target(img, lab, os.path.join(lab_dir, f"{key}.nii.gz"))
|
||||||
ds_json = {"channel_names": {"0": "t1c"},
|
ds_json = {"channel_names": {"0": "t1c"},
|
||||||
"labels": {"background": 0, "tumor": 1},
|
"labels": {"background": 0, "tumor": 1},
|
||||||
"numTraining": len(seen),
|
"numTraining": len(seen),
|
||||||
|
|
@ -178,6 +212,11 @@ def write_splits(train_keys, val_keys):
|
||||||
|
|
||||||
|
|
||||||
# ---------------- label writing / selection (same gates as scripts/06) ----------------
|
# ---------------- label writing / selection (same gates as scripts/06) ----------------
|
||||||
|
#
|
||||||
|
# Prediction outputs are ALREADY on each case's native grid: nnUNetv2_predict
|
||||||
|
# (--save_probabilities) resamples tumor probabilities and the argmax seg back
|
||||||
|
# to the original input shape. Gates below therefore run directly on the native
|
||||||
|
# volume; physical mm is recovered from the native voxel spacing.
|
||||||
|
|
||||||
def write_label(arr_u01, path, ref_itk):
|
def write_label(arr_u01, path, ref_itk):
|
||||||
img = sitk.GetImageFromArray(arr_u01.astype(np.uint8))
|
img = sitk.GetImageFromArray(arr_u01.astype(np.uint8))
|
||||||
|
|
@ -186,8 +225,11 @@ def write_label(arr_u01, path, ref_itk):
|
||||||
sitk.WriteImage(img, path, True)
|
sitk.WriteImage(img, path, True)
|
||||||
|
|
||||||
|
|
||||||
def pos_mask(p_tumor, tau_pos, min_cc_frac, vol_lo, vol_hi):
|
def pos_mask(p_tumor, tau_pos, min_cc_frac, vol_lo, vol_hi, voxel_vol_mm3=1.0):
|
||||||
"""Returns (mask, cc_frac, vol_mm3) if the positive gates pass, else None."""
|
"""Returns (mask, cc_frac, vol_mm3) if the positive gates pass, else None.
|
||||||
|
|
||||||
|
voxel_vol_mm3 is the native voxel volume in mm3 (Product of NIfTI spacing);
|
||||||
|
1.0 for 1mm processed volumes (Pipelines A/C)."""
|
||||||
m = p_tumor >= tau_pos
|
m = p_tumor >= tau_pos
|
||||||
if m.sum() == 0:
|
if m.sum() == 0:
|
||||||
return None
|
return None
|
||||||
|
|
@ -198,15 +240,26 @@ def pos_mask(p_tumor, tau_pos, min_cc_frac, vol_lo, vol_hi):
|
||||||
sizes = ndimage.sum(m, lab, range(1, n + 1))
|
sizes = ndimage.sum(m, lab, range(1, n + 1))
|
||||||
big = (lab == (int(np.argmax(sizes)) + 1)).astype(np.uint8)
|
big = (lab == (int(np.argmax(sizes)) + 1)).astype(np.uint8)
|
||||||
cc_frac = float(big.sum()) / float(m.sum())
|
cc_frac = float(big.sum()) / float(m.sum())
|
||||||
vol_mm3 = int(big.sum())
|
vol_mm3 = int(big.sum()) * voxel_vol_mm3
|
||||||
if cc_frac >= min_cc_frac and vol_lo <= vol_mm3 <= vol_hi:
|
if cc_frac >= min_cc_frac and vol_lo <= vol_mm3 <= vol_hi:
|
||||||
return big, cc_frac, vol_mm3
|
return big, cc_frac, vol_mm3
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _interior_mask(vol, min_interior=5000):
|
||||||
|
"""Head-interior mask. Processed [0,1] percentile-normalized volumes keep
|
||||||
|
the legacy `vol > 0.02` rule; native raw-intensity volumes use the robust
|
||||||
|
percentile / largest-CC head mask (see src/common.head_mask_from_image)."""
|
||||||
|
if float(vol.max()) <= 1.0 + 1e-6:
|
||||||
|
m = vol > 0.02
|
||||||
|
else:
|
||||||
|
m = head_mask_from_image(vol)
|
||||||
|
return m if m.sum() >= min_interior else None
|
||||||
|
|
||||||
|
|
||||||
def neg_frac_bg(p_tumor, vol, tau_neg, min_frac, min_interior=5000):
|
def neg_frac_bg(p_tumor, vol, tau_neg, min_frac, min_interior=5000):
|
||||||
interior = vol > 0.02
|
interior = _interior_mask(vol, min_interior)
|
||||||
if interior.sum() < min_interior:
|
if interior is None:
|
||||||
return None
|
return None
|
||||||
frac = float(((1.0 - p_tumor)[interior] >= tau_neg).mean())
|
frac = float(((1.0 - p_tumor)[interior] >= tau_neg).mean())
|
||||||
return frac if frac >= min_frac else None
|
return frac if frac >= min_frac else None
|
||||||
|
|
@ -225,31 +278,36 @@ def tumor_probs_from_npz(npz_path):
|
||||||
# The absolute patient-space resampling used by scripts/06 cannot work for this
|
# The absolute patient-space resampling used by scripts/06 cannot work for this
|
||||||
# data: each acquisition has its own scanner/patient coordinate system (table
|
# data: each acquisition has its own scanner/patient coordinate system (table
|
||||||
# offset, head orientation), so true labels from two visits of the same subject
|
# offset, head orientation), so true labels from two visits of the same subject
|
||||||
# resample to ~0 dice. Instead we gate on frame-invariant features:
|
# resample to ~0 dice. Instead we gate on frame-invariant features, computed in
|
||||||
# * tumor centroid relative to the head centroid (crop space, 1mm voxels)
|
# each case's own grid with physical mm (native spacing for Pipeline B, 1mm for
|
||||||
# * tumor volume ratio
|
# A/C):
|
||||||
|
# * tumor centroid relative to the head centroid (offset in mm)
|
||||||
|
# * tumor volume ratio (mm3)
|
||||||
# Two timepoints are compatible if they agree with at least one accepted
|
# Two timepoints are compatible if they agree with at least one accepted
|
||||||
# neighbor on both; greedy removal until stable (same scheme as scripts/06).
|
# neighbor on both; greedy removal until stable (same scheme as scripts/06).
|
||||||
|
|
||||||
def head_centroid(vol):
|
def head_centroid(vol):
|
||||||
m = vol > 0.02
|
m = _interior_mask(vol)
|
||||||
if m.sum() < 5000:
|
if m is None:
|
||||||
return None
|
return None
|
||||||
lab, n = ndimage.label(m)
|
m = largest_cc(m)
|
||||||
if n == 0:
|
if int(m.sum()) == 0:
|
||||||
return None
|
return None
|
||||||
sizes = ndimage.sum(m, lab, range(1, n + 1))
|
return np.array(ndimage.center_of_mass(m))
|
||||||
return np.array(ndimage.center_of_mass(lab == (int(np.argmax(sizes)) + 1)))
|
|
||||||
|
|
||||||
|
|
||||||
def rel_tumor_features(mask, vol):
|
def rel_tumor_features(mask, vol, spacing=(1.0, 1.0, 1.0)):
|
||||||
"""Returns (centroid_rel_to_head [z,y,x] in mm, vol_mm3) or None."""
|
"""Returns (centroid_rel_to_head [z,y,x] in mm, vol_mm3) or None.
|
||||||
|
|
||||||
|
spacing is the NIfTI voxel spacing in array order (z, y, x); (1, 1, 1) for
|
||||||
|
1mm processed volumes (Pipelines A/C)."""
|
||||||
if int(mask.sum()) == 0:
|
if int(mask.sum()) == 0:
|
||||||
return None
|
return None
|
||||||
|
spacing = np.asarray(spacing, dtype=np.float64)
|
||||||
tc = np.array(ndimage.center_of_mass(mask))
|
tc = np.array(ndimage.center_of_mass(mask))
|
||||||
hc = head_centroid(vol)
|
hc = head_centroid(vol)
|
||||||
rel = tc - hc if hc is not None else None
|
rel = (tc - hc) * spacing if hc is not None else None
|
||||||
return rel, int(mask.sum())
|
return rel, int(mask.sum()) * float(spacing.prod())
|
||||||
|
|
||||||
|
|
||||||
def _compatible(fa, fb, max_rel_dist_mm, vol_ratio_max):
|
def _compatible(fa, fb, max_rel_dist_mm, vol_ratio_max):
|
||||||
|
|
@ -264,10 +322,12 @@ def _compatible(fa, fb, max_rel_dist_mm, vol_ratio_max):
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
def consistency_filter(rows, out_dir, max_rel_dist_mm=40.0, vol_ratio_max=10.0):
|
def consistency_filter(rows, out_dir, max_rel_dist_mm=40.0, vol_ratio_max=10.0, image_key="pimg"):
|
||||||
"""Greedy removal of positive timepoints inconsistent with accepted neighbors.
|
"""Greedy removal of positive timepoints inconsistent with accepted neighbors.
|
||||||
|
|
||||||
masks from <out_dir>/<key>_label.nii.gz, volumes from row["pimg"]. Returns n_rejected.
|
masks from <out_dir>/<key>_label.nii.gz, volumes + native spacing from
|
||||||
|
row[image_key] ("pimg": 1mm processed, Pipelines A/C; "img": native,
|
||||||
|
Pipeline B). Returns n_rejected.
|
||||||
"""
|
"""
|
||||||
by_subj = {}
|
by_subj = {}
|
||||||
for r in rows:
|
for r in rows:
|
||||||
|
|
@ -281,8 +341,11 @@ def consistency_filter(rows, out_dir, max_rel_dist_mm=40.0, vol_ratio_max=10.0):
|
||||||
feats = {}
|
feats = {}
|
||||||
for i, r in enumerate(tps):
|
for i, r in enumerate(tps):
|
||||||
mask = read_nii_arr(os.path.join(out_dir, r["key"] + "_label.nii.gz")).astype(bool)
|
mask = read_nii_arr(os.path.join(out_dir, r["key"] + "_label.nii.gz")).astype(bool)
|
||||||
vol = read_nii_arr(r["pimg"]).astype("float32")
|
vol = read_nii_arr(r[image_key]).astype("float32")
|
||||||
feats[i] = rel_tumor_features(mask, vol)
|
if mask.shape != vol.shape:
|
||||||
|
raise ValueError(f"mask shape {mask.shape} != volume {r[image_key]} shape {vol.shape} for {r['key']}")
|
||||||
|
sp = np.asarray(sitk.ReadImage(r[image_key]).GetSpacing()[::-1], dtype=np.float64)
|
||||||
|
feats[i] = rel_tumor_features(mask, vol, sp)
|
||||||
accepted = list(range(len(tps)))
|
accepted = list(range(len(tps)))
|
||||||
while True:
|
while True:
|
||||||
changed = False
|
changed = False
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue