docs(readme): update project documentation and directory structure
Update README.md to include detailed project overview, environment requirements, data sources, and a comprehensive directory layout. Add nnU-Net pipeline documentation and directory descriptions. Update .gitignore to exclude nnU-Net specific directories and add new scripts directory for nnU-Net pipeline. Add initial nnU-Net pipeline scripts.
This commit is contained in:
parent
a491ae0523
commit
8c813db209
11 changed files with 1125 additions and 0 deletions
2
.gitignore
vendored
2
.gitignore
vendored
|
|
@ -2,4 +2,6 @@ data/
|
|||
runs/
|
||||
results/
|
||||
logs/
|
||||
nnu/
|
||||
runs_nnu/
|
||||
__pycache__/
|
||||
|
|
|
|||
206
README.md
206
README.md
|
|
@ -0,0 +1,206 @@
|
|||
# Longitudinal T1c Brain-Tumor Segmentation
|
||||
|
||||
Longitudinal (repeated-measures) analysis of brain T1 post-contrast (T1c) MRI with
|
||||
tumor segmentation, plus an **iterative pseudo-labeling** study that leverages a
|
||||
large unlabeled pool from the same subjects/series to improve a held-out,
|
||||
patient-level test Dice.
|
||||
|
||||
All volumes are preprocessed to a uniform 1 mm isotropic, head-cropped,
|
||||
percentile-normalized form so that tumor counts are directly comparable in mm³
|
||||
across sources and timepoints.
|
||||
|
||||
## Environment
|
||||
|
||||
Conda env `longitudinal` (Python 3.14, torch 2.14 +cu126). Activate before any
|
||||
command:
|
||||
|
||||
```bash
|
||||
source /opt/conda/etc/profile.d/conda.sh && conda activate longitudinal
|
||||
```
|
||||
|
||||
GPU (CUDA 12.6) is available; multi-GPU jobs use `torchrun` (in-house) or
|
||||
nnU-Net's own DDP (`-num_gpus`). Run scripts from the repo root.
|
||||
|
||||
## Data sources
|
||||
|
||||
| Source | What | Notes |
|
||||
|---|---|---|
|
||||
| `ntuh` | Labeled T1c + tumor segmentation | Native MR, deduped per acquisition |
|
||||
| `m6` | Labeled (GTV warped from CT) + large unlabeled pool | GTV registered CT→T1c |
|
||||
| `lee` | Longitudinal T1c scanned as JPG + DICOM txt | Volumes reconstructed from slices |
|
||||
|
||||
## Directory layout
|
||||
|
||||
```
|
||||
data/
|
||||
manifests/ jsonl row tables (see Pipeline A)
|
||||
proc/ <key>.nii.gz (1mm, cropped, normalized)
|
||||
proc/<key>_label.nii.gz
|
||||
procmeta/<key>.json geometry (origin/direction/crop_vox) + normalization
|
||||
vols.json labeled tumor volume stats (mm3 percentiles)
|
||||
pseudo/roundK/ in-house pseudo-labels (rows.jsonl, masks, summary.json)
|
||||
src/ U-Net, dataset, losses, training/eval helpers
|
||||
scripts/ Pipeline A (in-house 3D U-Net)
|
||||
scripts_nnu/ Pipeline B (nnU-Net)
|
||||
runs/roundK/ in-house checkpoints (best.pt, final.pt, state.pt)
|
||||
runs_nnu/roundK/ nnU-Net checkpoint snapshots (best_nnu.pth)
|
||||
nnu/ nnU-Net raw / preprocessed / results trees
|
||||
results/ evaluation JSON tables + plots
|
||||
logs/ run logs
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Pipeline A — In-house 3D U-Net (`scripts/`)
|
||||
|
||||
3-class-capable but used as binary (background / tumor) `Unet3D`, 96³ patches,
|
||||
DDP, per-sample weighted loss.
|
||||
|
||||
| # | Script | Purpose |
|
||||
|---|---|---|
|
||||
| 01 | `01_build_ntuh_manifest.py` | Build NTUH2022G4 labeled T1c + seg manifest |
|
||||
| 02 | `02_build_m6_dataset.py` | Build M6-2025 manifests (GTV CT→T1c registration/warp) |
|
||||
| 03 | `03_scan_lee_t1c.py` | Scan lee for brain T1c series → raw + selected manifests |
|
||||
| 04 | `04_reconstruct_lee.py` | Reconstruct lee T1c niftis from JPG slices + txt metadata |
|
||||
| 05 | `05_build_splits.py` | Patient-level train/val/test splits + unlabeled pool + volume stats |
|
||||
| 06 | `06_pseudo_label.py` | Pseudo-label the pool with a round model (sliding-window + TTA) |
|
||||
| 07 | `07_train.py` | DDP training entrypoint (labeled + weighted pseudo rows) |
|
||||
| 08 | `08_eval.py` | Held-out test evaluation (Dice from probability map) |
|
||||
| 09 | `09_run_iterative.py` | Orchestrates rounds 0…K and produces the summary table/plot |
|
||||
| — | `preprocess.py` | Crop/normalize source volumes → `data/proc/` |
|
||||
| — | `scan_procs.py` | Flag corrupt/non-finite processed volumes |
|
||||
| — | `test_dataloader.py` | Smoke-test the training dataloader |
|
||||
|
||||
Pseudo-label gates (per volume): **pos** when `p_tumor ≥ tau_pos`, largest-CC
|
||||
fraction ≥ `min_cc_frac`, and volume within the labeled p2–p98 range; **neg** when
|
||||
≥ `neg_frac` of interior voxels have `p_bg ≥ tau_neg`; then a per-subject
|
||||
longitudinal consistency filter.
|
||||
|
||||
Run the full study:
|
||||
|
||||
```bash
|
||||
python scripts/09_run_iterative.py --rounds 4 --gpus 3
|
||||
```
|
||||
|
||||
Individual stages are standalone, e.g.:
|
||||
|
||||
```bash
|
||||
torchrun --standalone --nproc_per_node 3 scripts/07_train.py \
|
||||
--rows data/manifests/split_train.jsonl --val data/manifests/split_val.jsonl \
|
||||
--epochs 40 --lr 3e-4 --batch 3 --ckpt-dir runs/round0
|
||||
python scripts/08_eval.py --rows data/manifests/split_test.jsonl --ckpt runs/round0/best.pt
|
||||
```
|
||||
|
||||
### Known issue in Pipeline A (longitudinal consistency)
|
||||
|
||||
`06_pseudo_label.py`'s `grid_info()`/`dice_a_on_b()` resample timepoints into a
|
||||
shared physical space using `procmeta` origin/direction/crop_vox. This is
|
||||
**not reliable across separate acquisitions**:
|
||||
|
||||
- Each scan has its own scanner/patient coordinate frame (table offset + head
|
||||
pose), so two visits of the same subject do not overlap in absolute space.
|
||||
Benchmarking against labeled multi-timepoint patients gives median cross-visit
|
||||
true-label Dice ≈ 0, meaning the consistency gate tends to over-reject.
|
||||
- `crop_vox` is stored as `[z, y, x]` array-axis starts, while `direction` is
|
||||
ordered `(x_dir, y_dir, z_dir)`; the pairings in `grid_info` mix these up.
|
||||
|
||||
This is noted for the record; Pipeline B below replaces the gate with a
|
||||
frame-invariant one. Pipeline A was left unchanged.
|
||||
|
||||
---
|
||||
|
||||
## Pipeline B — nnU-Net iterative pseudo-labeling (`scripts_nnu/`)
|
||||
|
||||
The same study driven by **nnU-Net v2** (installed build `nnunetv2` 2.8.1) as the
|
||||
segmentation backbone, keeping the identical patient-level splits, unlabeled
|
||||
pool, and evaluation protocol so results are directly comparable to Pipeline A.
|
||||
|
||||
> This installed nnU-Net is a modern fork: preprocessed data is `.b2nd`
|
||||
> (blosc2), checkpoints store `network_weights` (not `model`), and
|
||||
> `--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
|
||||
> docs.
|
||||
|
||||
Dataset: `Dataset210_NTUH_T1C_PL`, single channel `t1c`, 2 classes
|
||||
(background=0, tumor=1), `3d_fullres` only, `nnUNetPlans`, fold 0.
|
||||
|
||||
| # | Script | Purpose |
|
||||
|---|---|---|
|
||||
| 01 | `01_nnu_prepare_dataset.py` | Build/rebuild the raw dataset (symlinked `imagesTr`/`labelsTr` + `dataset.json`) from a rows jsonl |
|
||||
| 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 |
|
||||
| 04 | `04_nnu_pseudo_label.py` | Multi-GPU `nnUNetv2_predict` on the pool + selection gates |
|
||||
| 05 | `05_nnu_eval_test.py` | Held-out test evaluation (probability Dice@0.5 + hard-seg Dice) |
|
||||
| 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 |
|
||||
| — | `nnu_common.py` | Shared paths, env, dataset/split helpers, selection + consistency |
|
||||
|
||||
### Round semantics
|
||||
|
||||
- **Round 0:** train on the labeled patient-level train split (nnU-Net internal
|
||||
validation = the held-out val split, patient-disjoint via `splits_final.json`).
|
||||
Pseudo-label the remaining pool → `data/pseudo_nnu/round1`.
|
||||
- **Round k (k ≥ 1):** dataset grows with all accepted pseudo-labels from rounds
|
||||
1…k (positives + zero-mask negatives, de-duplicated by key); re-plan/preprocess;
|
||||
**warm-start** training from round k−1's best checkpoint at a lower LR; predict
|
||||
the *remaining* pool; evaluate on test.
|
||||
|
||||
Because this build has no incremental preprocessing, each round re-plans and
|
||||
re-preprocesses the whole (growing) dataset; the pool shrinks each round since
|
||||
accepted cases are excluded from re-prediction.
|
||||
|
||||
### Custom trainer
|
||||
|
||||
`NTUHLPLTrainer` (resolved through the `nnUNet_extTrainer` env var) adds:
|
||||
|
||||
- `NNU_PL_EPOCHS` / `NNU_PL_LR` — epoch count and initial PolyLR (set per round).
|
||||
- **Full-weight warm start** (`NNU_PL_WARMSTART`): loads *all* weights including
|
||||
the segmentation head in `on_train_start`. The CLI `-pretrained_weights` flag
|
||||
deliberately skips `.seg_layers.` keys, which would silently re-initialize the
|
||||
head and break round-to-round fine-tuning — hence this hook.
|
||||
|
||||
### Longitudinal consistency (frame-invariant)
|
||||
|
||||
Instead of resampling into absolute patient space (unreliable, see Pipeline A
|
||||
note), two timepoints are compatible if they agree with at least one accepted
|
||||
neighbor on both:
|
||||
|
||||
- tumor centroid offset **relative to the head centroid** ≤ `--max-rel-dist` (default 40 mm), and
|
||||
- tumor **volume ratio** ≤ `--vol-ratio` (default 10×)
|
||||
|
||||
same greedy-removal scheme as Pipeline A, but robust to scanner/pose differences
|
||||
across visits.
|
||||
|
||||
### Running
|
||||
|
||||
Full study from the repo root:
|
||||
|
||||
```bash
|
||||
python scripts_nnu/06_nnu_run_iterative.py --rounds 4 --gpus 3
|
||||
```
|
||||
|
||||
Defaults: baseline 250 epochs @ 1e-2; warm-started rounds 75 epochs @ 1e-3.
|
||||
Pseudo-label gates match Pipeline A (`--tau-pos 0.95`, `--min-cc-frac 0.2`,
|
||||
`--neg-frac 0.9`, volume p2–p98). Optional flags: `--no-tta`,
|
||||
`--no-neg-pseudo`, and the gate overrides. Outputs to `results/round{k}_test_nnu.json`,
|
||||
`results/iterative_table_nnu.jsonl`, `results/iterative_dice_nnu.png`; per-round
|
||||
checkpoints snapshotted to `runs_nnu/round{k}/best_nnu.pth`.
|
||||
|
||||
Individual stages:
|
||||
|
||||
```bash
|
||||
python scripts_nnu/01_nnu_prepare_dataset.py --rows <rows.jsonl>
|
||||
python scripts_nnu/02_nnu_plan_preprocess.py --val data/manifests/split_val.jsonl
|
||||
python scripts_nnu/03_nnu_train.py --gpus 3 --epochs 250 --lr 1e-2
|
||||
python scripts_nnu/04_nnu_pseudo_label.py --pool data/manifests/unlabeled_pool.jsonl \
|
||||
--out data/pseudo_nnu/round1 --gpus 3
|
||||
python scripts_nnu/05_nnu_eval_test.py --rows data/manifests/split_test.jsonl \
|
||||
--out results/round0_test_nnu.json --gpus 3
|
||||
```
|
||||
|
||||
### Verified
|
||||
|
||||
End-to-end smoke-tested on a scratch 4-case dataset: dataset build → planning →
|
||||
`splits_final.json` → training 1 epoch → warm-start → 2-way sharded prediction →
|
||||
selection gates → test eval. The consistency filter's keep / reject /
|
||||
single-timepoint paths are unit-tested with synthetic volumes.
|
||||
27
scripts_nnu/01_nnu_prepare_dataset.py
Normal file
27
scripts_nnu/01_nnu_prepare_dataset.py
Normal file
|
|
@ -0,0 +1,27 @@
|
|||
"""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.
|
||||
Rebuilds imagesTr/labelsTr as symlinks and writes dataset.json.
|
||||
|
||||
Usage: python scripts_nnu/01_nnu_prepare_dataset.py --rows <rows.jsonl>
|
||||
"""
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
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__))))
|
||||
from src.common import load_jsonl
|
||||
from nnu_common import make_raw_dataset, raw_ds
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--rows", required=True, help="jsonl with {key, pimg, label}")
|
||||
args = ap.parse_args()
|
||||
rows = load_jsonl(args.rows)
|
||||
n = make_raw_dataset(rows)
|
||||
print(f"[nnu:prepare] {raw_ds()}: {n} cases", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
33
scripts_nnu/02_nnu_plan_preprocess.py
Normal file
33
scripts_nnu/02_nnu_plan_preprocess.py
Normal file
|
|
@ -0,0 +1,33 @@
|
|||
"""Plan + preprocess the nnU-Net dataset (always --clean), then write the
|
||||
subject-level splits_final.json: val = split_val cases, train = the rest
|
||||
(all base + pseudo cases), same 5-fold content so -f 0 is the one we use.
|
||||
|
||||
Usage: python scripts_nnu/02_nnu_plan_preprocess.py --val data/manifests/split_val.jsonl [--npp 8]
|
||||
"""
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
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__))))
|
||||
from src.common import load_jsonl
|
||||
from nnu_common import run, nnu_env, plan_preprocess_cmd, dataset_case_keys, write_splits, preproc_ds
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--val", required=True, help="jsonl of internal-validation cases (subject-disjoint)")
|
||||
ap.add_argument("--npp", type=int, default=8, help="preprocess processes")
|
||||
ap.add_argument("--log", default=None)
|
||||
args = ap.parse_args()
|
||||
log = args.log or os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "logs", "nnu_plan.log")
|
||||
run(plan_preprocess_cmd(args.npp), log, env=nnu_env())
|
||||
val_keys = [r["key"] for r in load_jsonl(args.val)]
|
||||
all_keys = dataset_case_keys()
|
||||
assert all_keys, "raw dataset is empty; run 01_nnu_prepare_dataset.py first"
|
||||
assert set(val_keys) <= set(all_keys), f"val keys not in dataset: {set(val_keys) - set(all_keys)}"
|
||||
n_tr, n_va = write_splits(all_keys, val_keys)
|
||||
print(f"[nnu:plan] splits_final.json written to {preproc_ds()}: train={n_tr} val={n_va}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
32
scripts_nnu/03_nnu_train.py
Normal file
32
scripts_nnu/03_nnu_train.py
Normal file
|
|
@ -0,0 +1,32 @@
|
|||
"""Train the nnU-Net model for one pseudo-labeling round.
|
||||
|
||||
Epochs / initial lr / warm-start checkpoint are passed to NTUHLPLTrainer via
|
||||
NNU_PL_* env vars (see scripts_nnu/trainers/ntuh_pl_trainer.py).
|
||||
|
||||
Usage:
|
||||
nnUNetv2_train wrapper: python scripts_nnu/03_nnu_train.py \
|
||||
--gpus 3 --epochs 250 --lr 1e-2 [--warmstart runs_nnu/round0/best_nnu.pth]
|
||||
"""
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
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__))))
|
||||
from nnu_common import run, nnu_env, train_cmd, best_ckpt
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--gpus", type=int, default=3)
|
||||
ap.add_argument("--epochs", type=int, default=250)
|
||||
ap.add_argument("--lr", type=float, default=1e-2)
|
||||
ap.add_argument("--warmstart", default=None, help="checkpoint for full-weight warm start")
|
||||
ap.add_argument("--log", default=None)
|
||||
args = ap.parse_args()
|
||||
log = args.log or os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "logs", "nnu_train.log")
|
||||
run(train_cmd(args.gpus), log, env=nnu_env(epoch=args.epochs, lr=args.lr, warmstart=args.warmstart))
|
||||
print(f"[nnu:train] done; best checkpoint at {best_ckpt()}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
193
scripts_nnu/04_nnu_pseudo_label.py
Normal file
193
scripts_nnu/04_nnu_pseudo_label.py
Normal file
|
|
@ -0,0 +1,193 @@
|
|||
"""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.
|
||||
2. Runs one nnUNetv2_predict per GPU (-num_parts/-part_id), --save_probabilities.
|
||||
3. Selection gates (same as scripts/06_pseudo_label.py):
|
||||
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)
|
||||
neg: >= neg_frac of interior (vol > 0.02) voxels have p_bg >= tau_neg
|
||||
4. Per-subject longitudinal consistency filter over accepted positive timepoints
|
||||
(head-relative tumor centroid distance + volume ratio; absolute patient-space
|
||||
grids are not comparable across acquisitions).
|
||||
Writes <out>/rows.jsonl, <out>/<key>_label.nii.gz, <out>/accepted.jsonl, <out>/summary.json.
|
||||
|
||||
Usage:
|
||||
python scripts_nnu/04_nnu_pseudo_label.py --pool data/manifests/unlabeled_pool.jsonl \
|
||||
--out data/pseudo_nnu/round1 --gpus 3
|
||||
"""
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
import json
|
||||
import shutil
|
||||
import numpy as np
|
||||
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__))))
|
||||
from src.common import load_jsonl, save_jsonl, read_nii_arr, write_arr
|
||||
import SimpleITK as sitk
|
||||
from nnu_common import (ROOT, d, nnu_env, predict_cmd, tumor_probs_from_npz,
|
||||
write_label, pos_mask, neg_frac_bg, consistency_filter,
|
||||
load_voxel_stats)
|
||||
|
||||
|
||||
def build_input_folder(rows, in_dir):
|
||||
shutil.rmtree(in_dir, ignore_errors=True)
|
||||
os.makedirs(in_dir, exist_ok=True)
|
||||
for r in rows:
|
||||
os.symlink(os.path.abspath(r["pimg"]), os.path.join(in_dir, r["key"] + "_0000.nii.gz"))
|
||||
|
||||
|
||||
def find_outputs(out_dir, ext):
|
||||
found = {}
|
||||
for part in os.listdir(out_dir):
|
||||
pd = os.path.join(out_dir, part)
|
||||
if not os.path.isdir(pd):
|
||||
continue
|
||||
for fn in os.listdir(pd):
|
||||
if fn.endswith(ext):
|
||||
found[fn[:-len(ext)]] = os.path.join(pd, fn)
|
||||
return found
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--pool", required=True)
|
||||
ap.add_argument("--out", required=True)
|
||||
ap.add_argument("--already", default=None, help="jsonl of case keys already added to the dataset")
|
||||
ap.add_argument("--gpus", type=int, default=3)
|
||||
ap.add_argument("--chk", default="checkpoint_best.pth")
|
||||
ap.add_argument("--no-tta", action="store_true")
|
||||
ap.add_argument("--npp", type=int, default=2, help="predict subprocesses per GPU")
|
||||
ap.add_argument("--tau-pos", type=float, default=0.95)
|
||||
ap.add_argument("--tau-neg", type=float, default=0.98)
|
||||
ap.add_argument("--neg-frac", type=float, default=0.90)
|
||||
ap.add_argument("--vol-qp", type=float, nargs=2, default=[2, 98])
|
||||
ap.add_argument("--min-cc-frac", type=float, default=0.2)
|
||||
ap.add_argument("--max-rel-dist", type=float, default=40.0,
|
||||
help="max mm of head-relative tumor centroid shift between compatible timepoints")
|
||||
ap.add_argument("--vol-ratio", type=float, default=10.0,
|
||||
help="max tumor volume ratio between compatible timepoints")
|
||||
ap.add_argument("--skip-predict", action="store_true", help="reuse existing prediction outputs")
|
||||
ap.add_argument("--keep-pred", action="store_true")
|
||||
args = ap.parse_args()
|
||||
|
||||
outname = os.path.basename(os.path.normpath(args.out))
|
||||
out = d(args.out)
|
||||
pool_rows = load_jsonl(args.pool)
|
||||
already = set()
|
||||
if args.already and os.path.exists(args.already):
|
||||
already = {r["key"] for r in load_jsonl(args.already)}
|
||||
rows = [r for r in pool_rows if r["key"] not in already]
|
||||
print(f"[nnu:pseudo:{outname}] pool={len(pool_rows)} already_in_dataset={len(already)} to_predict={len(rows)}", flush=True)
|
||||
|
||||
vstats = load_voxel_stats()
|
||||
vol_lo = vstats.get(f"p{args.vol_qp[0]:.0f}", 1.0)
|
||||
vol_hi = vstats.get(f"p{args.vol_qp[1]:.0f}", 50000.0)
|
||||
|
||||
tag = outname
|
||||
in_dir = os.path.join(d("nnu/pred"), f"in_pool_{tag}")
|
||||
out_dir = os.path.join(d("nnu/pred"), f"out_pool_{tag}")
|
||||
if rows:
|
||||
build_input_folder(rows, in_dir)
|
||||
if not args.skip_predict:
|
||||
import subprocess
|
||||
os.makedirs(out_dir, exist_ok=True)
|
||||
env = nnu_env()
|
||||
procs = []
|
||||
for g in range(args.gpus):
|
||||
e = dict(env)
|
||||
e["CUDA_VISIBLE_DEVICES"] = str(g)
|
||||
logf = os.path.join(d("logs/nnu"), f"pred_{tag}_part{g}.log")
|
||||
with open(logf, "wb") as lf:
|
||||
procs.append(subprocess.Popen(predict_cmd(in_dir, os.path.join(out_dir, f"part{g}"),
|
||||
args.gpus, g, chk=args.chk,
|
||||
tta=not args.no_tta, npp=args.npp),
|
||||
stdout=lf, stderr=subprocess.STDOUT, env=e, cwd=ROOT))
|
||||
for g, p in enumerate(procs):
|
||||
rc = p.wait()
|
||||
if rc != 0:
|
||||
print(f"[nnu:pseudo:{outname}] part{g} failed rc={rc}; retrying once (skips finished cases)", flush=True)
|
||||
e = dict(env)
|
||||
e["CUDA_VISIBLE_DEVICES"] = str(g)
|
||||
logf = os.path.join(d("logs/nnu"), f"pred_{tag}_part{g}_retry.log")
|
||||
with open(logf, "wb") as lf:
|
||||
p2 = subprocess.Popen(predict_cmd(in_dir, os.path.join(out_dir, f"part{g}"),
|
||||
args.gpus, g, chk=args.chk,
|
||||
tta=not args.no_tta, npp=args.npp),
|
||||
stdout=lf, stderr=subprocess.STDOUT, env=e, cwd=ROOT)
|
||||
rc2 = p2.wait()
|
||||
if rc2 != 0:
|
||||
raise RuntimeError(f"prediction part{g} failed twice; see {logf}")
|
||||
|
||||
npz = find_outputs(out_dir, ".npz") if os.path.isdir(out_dir) else {}
|
||||
rows_out, n_err = [], 0
|
||||
for r in rows:
|
||||
key = r["key"]
|
||||
entry = {"key": key, "subject": r["subject"], "date": r.get("date"),
|
||||
"source": r.get("source"), "pimg": r["pimg"], "label": None,
|
||||
"role": "rej", "vol_mm3": 0, "maxp": None}
|
||||
try:
|
||||
pimg_itk = sitk.ReadImage(r["pimg"])
|
||||
vol = read_nii_arr(r["pimg"]).astype("float32")
|
||||
pt = tumor_probs_from_npz(npz[key])
|
||||
if pt.shape != vol.shape:
|
||||
raise ValueError(f"prob shape {pt.shape} != image shape {vol.shape} for {key}")
|
||||
entry["maxp"] = round(float(pt.max()), 4)
|
||||
got = pos_mask(pt, args.tau_pos, args.min_cc_frac, vol_lo, vol_hi)
|
||||
if got is not None:
|
||||
mask, cc_frac, vol_mm3 = got
|
||||
lp = os.path.join(out, key + "_label.nii.gz")
|
||||
write_label(mask, lp, pimg_itk)
|
||||
entry.update({"role": "pos", "label": lp, "vol_mm3": vol_mm3, "cc_frac": round(cc_frac, 3)})
|
||||
else:
|
||||
frac = neg_frac_bg(pt, vol, args.tau_neg, args.neg_frac)
|
||||
if frac is not None:
|
||||
lp = os.path.join(out, key + "_label.nii.gz")
|
||||
write_arr(np.zeros(vol.shape, dtype="uint8"), lp, itk_img=pimg_itk)
|
||||
entry.update({"role": "neg", "label": lp, "neg_conf": round(frac, 4)})
|
||||
except Exception as e: # noqa
|
||||
entry["role"] = "error"
|
||||
n_err += 1
|
||||
print(f"[nnu:pseudo:{outname}] {key} ERR {e!r}", flush=True)
|
||||
rows_out.append(entry)
|
||||
|
||||
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_rej = consistency_filter(rows_out, out, args.max_rel_dist, args.vol_ratio)
|
||||
for x in rows_out:
|
||||
if x["role"] == "rej":
|
||||
x["label"] = None
|
||||
accepted = [x for x in rows_out if x["role"] in ("pos", "neg")]
|
||||
save_jsonl(rows_out, os.path.join(out, "rows.jsonl"))
|
||||
save_jsonl(accepted, os.path.join(out, "accepted.jsonl"))
|
||||
posv = [x["vol_mm3"] for x in rows_out if x["role"] == "pos"]
|
||||
summ = {
|
||||
"n_pool_predicted": len(rows_out),
|
||||
"n_predicted": len(npz),
|
||||
"n_pos": sum(1 for x in rows_out if x["role"] == "pos"),
|
||||
"n_neg": sum(1 for x in rows_out if x["role"] == "neg"),
|
||||
"n_pos_before_consistency": n_pos0,
|
||||
"n_rejected_consistency": n_rej,
|
||||
"n_other_rej": sum(1 for x in rows_out if x["role"] == "rej"),
|
||||
"n_error": n_err,
|
||||
"pos_vol_mm3": {"med": float(np.median(posv)) if posv else 0,
|
||||
"p5": float(float(np.percentile(posv, 5))) if posv else 0,
|
||||
"p95": float(float(np.percentile(posv, 95))) if posv else 0},
|
||||
"tau_pos": args.tau_pos, "tau_neg": args.tau_neg, "neg_frac": args.neg_frac,
|
||||
"vol_range": [vol_lo, vol_hi], "max_rel_dist_mm": args.max_rel_dist,
|
||||
"vol_ratio": args.vol_ratio, "chk": args.chk, "tta": not args.no_tta,
|
||||
}
|
||||
with open(os.path.join(out, "summary.json"), "w") as f:
|
||||
json.dump(summ, f, indent=1)
|
||||
print(f"[nnu:pseudo] round {outname}: {json.dumps(summ)}", flush=True)
|
||||
|
||||
if not args.keep_pred and n_err == 0:
|
||||
shutil.rmtree(in_dir, ignore_errors=True)
|
||||
shutil.rmtree(out_dir, ignore_errors=True)
|
||||
elif n_err > 0:
|
||||
print(f"[nnu:pseudo:{outname}] kept prediction outputs ({in_dir}, {out_dir}) due to {n_err} errors; "
|
||||
f"rerun with --skip-predict --keep-pred after fixing", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
113
scripts_nnu/05_nnu_eval_test.py
Normal file
113
scripts_nnu/05_nnu_eval_test.py
Normal file
|
|
@ -0,0 +1,113 @@
|
|||
"""Holdout test evaluation of an nnU-Net round model.
|
||||
|
||||
Predicts the test volumes (multi-GPU sharding, --save_probabilities) and scores:
|
||||
dice - tumor probability (softmax channel 1) thresholded at 0.5
|
||||
(same convention as scripts/08_eval.py)
|
||||
dice_hard - argmax segmentation written by nnUNetv2_predict
|
||||
Writes: <out> json + <out>.json -> per_row jsonl alongside.
|
||||
|
||||
Usage: python scripts_nnu/05_nnu_eval_test.py --rows data/manifests/split_test.jsonl \
|
||||
--out results/round0_test_nnu.json --gpus 3
|
||||
"""
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
import json
|
||||
import shutil
|
||||
import subprocess
|
||||
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__))))
|
||||
from src.common import load_jsonl, save_jsonl, read_nii_arr, dice
|
||||
from nnu_common import ROOT, d, nnu_env, predict_cmd, tumor_probs_from_npz
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--rows", required=True)
|
||||
ap.add_argument("--out", required=True)
|
||||
ap.add_argument("--gpus", type=int, default=3)
|
||||
ap.add_argument("--chk", default="checkpoint_best.pth")
|
||||
ap.add_argument("--no-tta", action="store_true")
|
||||
ap.add_argument("--npp", type=int, default=2)
|
||||
args = ap.parse_args()
|
||||
|
||||
rows = load_jsonl(args.rows)
|
||||
tag = os.path.basename(os.path.normpath(args.out)) + "_predict"
|
||||
in_dir = os.path.join(d("nnu/pred"), f"in_test_{tag}")
|
||||
out_dir = os.path.join(d("nnu/pred"), f"out_test_{tag}")
|
||||
shutil.rmtree(in_dir, ignore_errors=True)
|
||||
os.makedirs(in_dir, exist_ok=True)
|
||||
for r in rows:
|
||||
os.symlink(os.path.abspath(r["pimg"]), os.path.join(in_dir, r["key"] + "_0000.nii.gz"))
|
||||
|
||||
env = nnu_env()
|
||||
procs = []
|
||||
for g in range(args.gpus):
|
||||
e = dict(env)
|
||||
e["CUDA_VISIBLE_DEVICES"] = str(g)
|
||||
logf = os.path.join(d("logs/nnu"), f"pred_test_{tag}_part{g}.log")
|
||||
with open(logf, "wb") as lf:
|
||||
procs.append(subprocess.Popen(predict_cmd(in_dir, os.path.join(out_dir, f"part{g}"),
|
||||
args.gpus, g, chk=args.chk,
|
||||
tta=not args.no_tta, npp=args.npp),
|
||||
stdout=lf, stderr=subprocess.STDOUT, env=e, cwd=ROOT))
|
||||
for g, p in enumerate(procs):
|
||||
rc = p.wait()
|
||||
if rc != 0:
|
||||
print(f"[nnu:eval] test predict part{g} failed rc={rc}; retrying once", flush=True)
|
||||
e = dict(env)
|
||||
e["CUDA_VISIBLE_DEVICES"] = str(g)
|
||||
logf = os.path.join(d("logs/nnu"), f"pred_test_{tag}_part{g}_retry.log")
|
||||
with open(logf, "wb") as lf:
|
||||
p2 = subprocess.Popen(predict_cmd(in_dir, os.path.join(out_dir, f"part{g}"),
|
||||
args.gpus, g, chk=args.chk,
|
||||
tta=not args.no_tta, npp=args.npp),
|
||||
stdout=lf, stderr=subprocess.STDOUT, env=e, cwd=ROOT)
|
||||
if p2.wait() != 0:
|
||||
raise RuntimeError(f"test prediction part{g} failed twice; see {logf}")
|
||||
|
||||
npz, segs = {}, {}
|
||||
for part in os.listdir(out_dir):
|
||||
pd = os.path.join(out_dir, part)
|
||||
if not os.path.isdir(pd):
|
||||
continue
|
||||
for fn in os.listdir(pd):
|
||||
if fn.endswith(".npz"):
|
||||
npz[fn[:-4]] = os.path.join(pd, fn)
|
||||
elif fn.endswith(".nii.gz"):
|
||||
segs[fn[:-7]] = os.path.join(pd, fn)
|
||||
|
||||
per, dice_p, dice_h = [], [], []
|
||||
for r in rows:
|
||||
key = r["key"]
|
||||
lab = read_nii_arr(r.get("plabel") or r.get("label"))
|
||||
try:
|
||||
pt = tumor_probs_from_npz(npz[key])
|
||||
if pt.shape != lab.shape:
|
||||
raise ValueError(f"prob {pt.shape} vs label {lab.shape}")
|
||||
d1 = dice((pt >= 0.5).astype("uint8"), (lab > 0).astype("uint8"))
|
||||
dice_p.append(d1)
|
||||
hard = read_nii_arr(segs[key])
|
||||
if hard.shape != lab.shape:
|
||||
raise ValueError(f"hard seg {hard.shape} vs label {lab.shape}")
|
||||
d2 = dice((hard > 0).astype("uint8"), (lab > 0).astype("uint8"))
|
||||
dice_h.append(d2)
|
||||
per.append({"key": key, "dice": round(d1, 4), "dice_hard": round(d2, 4)})
|
||||
except Exception as e: # noqa
|
||||
print(f"[nnu:eval] {key} ERR {e!r}", flush=True)
|
||||
|
||||
res = {"ckpt_chron": args.chk, "n": len(per),
|
||||
"dice": float(sum(dice_p) / len(dice_p)) if dice_p else 0.0,
|
||||
"dice_hard": float(sum(dice_h) / len(dice_h)) if dice_h else 0.0,
|
||||
"n_pred": len(npz), "per_row": per}
|
||||
os.makedirs(os.path.dirname(os.path.abspath(args.out)), exist_ok=True)
|
||||
with open(args.out, "w") as f:
|
||||
json.dump(res, f, indent=1)
|
||||
save_jsonl(per, args.out.replace(".json", "_per_row.jsonl"))
|
||||
print(json.dumps({k: res[k] for k in ("n", "dice", "dice_hard", "n_pred")}, indent=1), flush=True)
|
||||
print("saved", args.out, flush=True)
|
||||
shutil.rmtree(in_dir, ignore_errors=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
166
scripts_nnu/06_nnu_run_iterative.py
Normal file
166
scripts_nnu/06_nnu_run_iterative.py
Normal file
|
|
@ -0,0 +1,166 @@
|
|||
"""Orchestrates the nnU-Net iterative pseudo-labeling study (parallel to scripts/09).
|
||||
|
||||
Round 0: train nnU-Net on the labeled patient-level train split (internal val =
|
||||
held-out patient-level val split, subject-disjoint).
|
||||
Round k: dataset grows with accepted pseudo-labels (pos + neg) from rounds 1..k;
|
||||
re-plan/preprocess; warm-start training (full weights) from round k-1
|
||||
best checkpoint at a lower initial lr; then pseudo-label the remaining
|
||||
unlabeled pool (gates identical to scripts/06) and evaluate the model on
|
||||
the held-out patient-level test split.
|
||||
|
||||
Usage: python scripts_nnu/06_nnu_run_iterative.py [--rounds 4] [--gpus 3]
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
import json
|
||||
import shutil
|
||||
import argparse
|
||||
import subprocess
|
||||
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__))))
|
||||
from src.common import load_jsonl, save_jsonl
|
||||
from nnu_common import ROOT, d, best_ckpt
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--rounds", type=int, default=4)
|
||||
ap.add_argument("--gpus", type=int, default=3)
|
||||
ap.add_argument("--base-epochs", type=int, default=250)
|
||||
ap.add_argument("--base-lr", type=float, default=1e-2)
|
||||
ap.add_argument("--round-epochs", type=int, default=75)
|
||||
ap.add_argument("--round-lr", type=float, default=1e-3)
|
||||
ap.add_argument("--no-neg-pseudo", action="store_true", help="do not add negative pseudo cases to training")
|
||||
ap.add_argument("--no-tta", action="store_true", help="disable mirroring TTA in inference")
|
||||
ap.add_argument("--tau-pos", type=float, default=0.95)
|
||||
ap.add_argument("--tau-neg", type=float, default=0.98)
|
||||
ap.add_argument("--neg-frac", type=float, default=0.90)
|
||||
ap.add_argument("--vol-qp", type=float, nargs=2, default=[2, 98])
|
||||
ap.add_argument("--min-cc-frac", type=float, default=0.2)
|
||||
ap.add_argument("--max-rel-dist", type=float, default=40.0)
|
||||
ap.add_argument("--vol-ratio", type=float, default=10.0)
|
||||
args = ap.parse_args()
|
||||
|
||||
man = os.path.join(ROOT, "data/manifests")
|
||||
train_f = os.path.join(man, "split_train.jsonl")
|
||||
val_f = os.path.join(man, "split_val.jsonl")
|
||||
test_f = os.path.join(man, "split_test.jsonl")
|
||||
pool_f = os.path.join(man, "unlabeled_pool.jsonl")
|
||||
results_dir = d("results")
|
||||
logdir = d("logs")
|
||||
main_log = os.path.join(logdir, "nnu_iterative.log")
|
||||
pdir_root = d("data/pseudo_nnu")
|
||||
added_f = os.path.join(pdir_root, "added_keys.jsonl")
|
||||
|
||||
base = []
|
||||
for f in (train_f, val_f):
|
||||
base += [{"key": r["key"], "pimg": r["pimg"], "label": r["plabel"]} for r in load_jsonl(f)]
|
||||
save_jsonl(base, os.path.join(pdir_root, "base_rows.jsonl"))
|
||||
print(f"[nnu-orch] base dataset rows: {len(base)}", flush=True)
|
||||
|
||||
def py(script, extra):
|
||||
print(f"[nnu-orch] $ python {script} {extra}", flush=True)
|
||||
with open(main_log, "a") as f:
|
||||
f.write(f"$ python {script} {extra}\n")
|
||||
subprocess.run([sys.executable, f"{ROOT}/scripts_nnu/{script}", *extra.split()],
|
||||
stdout=f, stderr=subprocess.STDOUT, check=True, cwd=ROOT)
|
||||
|
||||
def prepare(rows):
|
||||
f = os.path.join(pdir_root, "current_rows.jsonl")
|
||||
save_jsonl(rows, f)
|
||||
py("01_nnu_prepare_dataset.py", f"--rows {f}")
|
||||
|
||||
def plan():
|
||||
py("02_nnu_plan_preprocess.py", f"--val {val_f}")
|
||||
|
||||
def train(round, epochs, lr, warmstart):
|
||||
extra = f"--gpus {args.gpus} --epochs {epochs} --lr {lr} --log {os.path.join(logdir, f'nnu_train_r{round}.log')}"
|
||||
if warmstart:
|
||||
extra += f" --warmstart {warmstart}"
|
||||
py("03_nnu_train.py", extra)
|
||||
dst = os.path.join(ROOT, "runs_nnu", f"round{round}", "best_nnu.pth")
|
||||
os.makedirs(os.path.dirname(dst), exist_ok=True)
|
||||
shutil.copy2(best_ckpt(), dst)
|
||||
return dst
|
||||
|
||||
def predict_pool(round, added):
|
||||
extra = (f"--pool {pool_f} --out {pdir_root}/round{round} --gpus {args.gpus} "
|
||||
f"--tau-pos {args.tau_pos} --tau-neg {args.tau_neg} --neg-frac {args.neg_frac} "
|
||||
f"--vol-qp {args.vol_qp[0]} {args.vol_qp[1]} --min-cc-frac {args.min_cc_frac} "
|
||||
f"--max-rel-dist {args.max_rel_dist} --vol-ratio {args.vol_ratio}")
|
||||
if args.no_tta:
|
||||
extra += " --no-tta"
|
||||
if added:
|
||||
extra += f" --already {added_f}"
|
||||
py("04_nnu_pseudo_label.py", extra)
|
||||
return os.path.join(pdir_root, f"round{round}", "accepted.jsonl")
|
||||
|
||||
def eval_test(round, out_name):
|
||||
extra = (f"--rows {test_f} --out {results_dir}/{out_name} --gpus {args.gpus}")
|
||||
if args.no_tta:
|
||||
extra += " --no-tta"
|
||||
py("05_nnu_eval_test.py", extra)
|
||||
return os.path.join(results_dir, out_name)
|
||||
|
||||
# ---- round 0: baseline ----
|
||||
prepare(base)
|
||||
plan()
|
||||
train(0, args.base_epochs, args.base_lr, None)
|
||||
accepted_f = [predict_pool(1, None)]
|
||||
save_jsonl(load_jsonl(accepted_f[0]), added_f)
|
||||
eval_test(0, "round0_test_nnu.json")
|
||||
|
||||
for k in range(1, args.rounds + 1):
|
||||
rows = list(base)
|
||||
for f in accepted_f:
|
||||
for r in load_jsonl(f):
|
||||
if args.no_neg_pseudo and r["role"] == "neg":
|
||||
continue
|
||||
if all(r["key"] != x["key"] for x in rows):
|
||||
rows.append(r)
|
||||
prepare(rows)
|
||||
plan()
|
||||
train(k, args.round_epochs, args.round_lr, os.path.join(ROOT, "runs_nnu", f"round{k-1}", "best_nnu.pth"))
|
||||
accepted_f.append(predict_pool(k + 1, added_f))
|
||||
save_jsonl([r for f in accepted_f for r in load_jsonl(f)], added_f)
|
||||
eval_test(k, f"round{k}_test_nnu.json")
|
||||
|
||||
# ---- report (same layout as scripts/09) ----
|
||||
table = []
|
||||
for k in range(args.rounds + 1):
|
||||
resf = os.path.join(results_dir, f"round{k}_test_nnu.json")
|
||||
if not os.path.exists(resf):
|
||||
continue
|
||||
r = json.load(open(resf))
|
||||
row = {"round": k, "test_dice": round(r["dice"], 4), "test_dice_hard": round(r["dice_hard"], 4),
|
||||
"n_test": r["n"], "ckpt": os.path.join(ROOT, "runs_nnu", f"round{k}", "best_nnu.pth")}
|
||||
pf = os.path.join(pdir_root, f"round{k}", "summary.json")
|
||||
if k > 0 and os.path.exists(pf):
|
||||
s = json.load(open(pf))
|
||||
row.update({"n_pos": s["n_pos"], "n_neg": s["n_neg"],
|
||||
"n_rej_cons": s["n_rejected_consistency"], "n_error": s["n_error"],
|
||||
"pos_vol_med_mm3": s["pos_vol_mm3"]["med"]})
|
||||
table.append(row)
|
||||
save_jsonl(table, os.path.join(results_dir, "iterative_table_nnu.jsonl"))
|
||||
print(json.dumps(table, indent=1))
|
||||
try:
|
||||
import matplotlib
|
||||
matplotlib.use("Agg")
|
||||
import matplotlib.pyplot as plt
|
||||
rs = [t["round"] for t in table]
|
||||
ds = [t["test_dice"] for t in table]
|
||||
plt.figure(figsize=(6, 4))
|
||||
plt.plot(rs, ds, "o-")
|
||||
plt.xlabel("pseudo-labeling round")
|
||||
plt.ylabel("holdout tumor Dice (nnU-Net)")
|
||||
for x, y in zip(rs, ds):
|
||||
plt.annotate(f"{y:.3f}", (x, y), textcoords="offset points", xytext=(0, 8), fontsize=8)
|
||||
plt.grid(alpha=0.3)
|
||||
plt.tight_layout()
|
||||
plt.savefig(os.path.join(results_dir, "iterative_dice_nnu.png"), dpi=150)
|
||||
except Exception as e: # noqa
|
||||
print("plot failed:", repr(e))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
307
scripts_nnu/nnu_common.py
Normal file
307
scripts_nnu/nnu_common.py
Normal file
|
|
@ -0,0 +1,307 @@
|
|||
"""Shared helpers for the nnU-Net iterative pseudo-labeling pipeline (scripts_nnu/).
|
||||
|
||||
Runs against the installed nnunetv2 2.8.1 build (checkpoint key `network_weights`,
|
||||
probabilities saved as <case>.npz channel-first, folds in splits_final.json).
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
import json
|
||||
import shutil
|
||||
import subprocess
|
||||
|
||||
ROOT = os.environ.get("LONGITUDINAL_ROOT", "/mnt/b4/xfr/git26/longitudinal")
|
||||
sys.path.insert(0, ROOT)
|
||||
from src.common import load_jsonl, save_jsonl, read_nii_arr # noqa: E402
|
||||
|
||||
import numpy as np
|
||||
import SimpleITK as sitk
|
||||
from scipy import ndimage
|
||||
|
||||
DS_ID = 210
|
||||
DS_NAME = f"Dataset{DS_ID:03d}_NTUH_T1C_PL"
|
||||
TRAINER = "NTUHLPLTrainer"
|
||||
CONFIG = "3d_fullres"
|
||||
FOLD = 0
|
||||
PLANS = "nnUNetPlans"
|
||||
|
||||
SCRIPTS_NNU = os.path.join(ROOT, "scripts_nnu")
|
||||
TRAINERS_DIR = os.path.join(SCRIPTS_NNU, "trainers")
|
||||
|
||||
|
||||
def nnu_root():
|
||||
return d(os.path.join("nnu"))
|
||||
|
||||
|
||||
def raw_ds():
|
||||
return os.path.join(nnu_root(), "raw", DS_NAME)
|
||||
|
||||
|
||||
def preproc_ds():
|
||||
return os.path.join(nnu_root(), "preprocessed", DS_NAME)
|
||||
|
||||
|
||||
def results_root():
|
||||
return os.path.join(nnu_root(), "results")
|
||||
|
||||
|
||||
def model_folder():
|
||||
return os.path.join(results_root(), DS_NAME, f"{TRAINER}__{PLANS}__{CONFIG}")
|
||||
|
||||
|
||||
def fold_dir():
|
||||
return os.path.join(model_folder(), f"fold_{FOLD}")
|
||||
|
||||
|
||||
def best_ckpt():
|
||||
return os.path.join(fold_dir(), "checkpoint_best.pth")
|
||||
|
||||
|
||||
def final_ckpt():
|
||||
return os.path.join(fold_dir(), "checkpoint_final.pth")
|
||||
|
||||
|
||||
def d(name):
|
||||
p = os.path.join(ROOT, name)
|
||||
os.makedirs(p, exist_ok=True)
|
||||
return p
|
||||
|
||||
|
||||
def nnu_env(epoch=None, lr=None, warmstart=None):
|
||||
env = dict(os.environ)
|
||||
env["nnUNet_raw"] = os.path.join(nnu_root(), "raw")
|
||||
env["nnUNet_preprocessed"] = os.path.join(nnu_root(), "preprocessed")
|
||||
env["nnUNet_results"] = results_root()
|
||||
env["nnUNet_extTrainer"] = TRAINERS_DIR
|
||||
env.pop("NNU_PL_EPOCHS", None)
|
||||
env.pop("NNU_PL_LR", None)
|
||||
env.pop("NNU_PL_WARMSTART", None)
|
||||
env.pop("CUDA_VISIBLE_DEVICES", None)
|
||||
if epoch is not None:
|
||||
env["NNU_PL_EPOCHS"] = str(epoch)
|
||||
if lr is not None:
|
||||
env["NNU_PL_LR"] = repr(float(lr))
|
||||
if warmstart:
|
||||
env["NNU_PL_WARMSTART"] = str(warmstart)
|
||||
return env
|
||||
|
||||
|
||||
def run(cmd, log, env=None, retries=2):
|
||||
os.makedirs(os.path.dirname(os.path.abspath(log)), exist_ok=True)
|
||||
cmd_s = " ".join(str(c) for c in cmd)
|
||||
for attempt in range(retries + 1):
|
||||
try:
|
||||
with open(log, "a") as f:
|
||||
f.write(f"$ (attempt {attempt + 1}) " + cmd_s + "\n")
|
||||
f.flush()
|
||||
print(f"$ (attempt {attempt + 1}) " + cmd_s, flush=True)
|
||||
subprocess.run(cmd, cwd=ROOT, stdout=f, stderr=subprocess.STDOUT,
|
||||
env=env, check=True)
|
||||
return
|
||||
except subprocess.CalledProcessError:
|
||||
if attempt == retries:
|
||||
raise
|
||||
print(f"[nnu] command failed, retrying in 60s: {cmd_s}", flush=True)
|
||||
time.sleep(60)
|
||||
|
||||
|
||||
def train_cmd(gpus):
|
||||
return ["nnUNetv2_train", str(DS_ID), CONFIG, str(FOLD), "-tr", TRAINER,
|
||||
"-p", PLANS, "-num_gpus", str(gpus)]
|
||||
|
||||
|
||||
def plan_preprocess_cmd(npp=8):
|
||||
return ["nnUNetv2_plan_and_preprocess", "-d", str(DS_ID), "-c", CONFIG,
|
||||
"-np", str(npp), "--clean", "--no_pbar"]
|
||||
|
||||
|
||||
def predict_cmd(in_dir, out_dir, gpus, part, chk="checkpoint_best.pth", tta=True,
|
||||
npp=2, nps=2):
|
||||
cmd = ["nnUNetv2_predict", "-i", str(in_dir), "-o", str(out_dir), "-d", str(DS_ID),
|
||||
"-c", CONFIG, "-tr", TRAINER, "-p", PLANS, "-f", str(FOLD), "-chk", chk,
|
||||
"--save_probabilities", "--continue_prediction", "--disable_progress_bar",
|
||||
"-num_parts", str(gpus), "-part_id", str(part),
|
||||
"-npp", str(npp), "-nps", str(nps)]
|
||||
if not tta:
|
||||
cmd.append("--disable_tta")
|
||||
return cmd
|
||||
|
||||
|
||||
# ---------------- raw dataset (nnU-Net format) ----------------
|
||||
|
||||
def make_raw_dataset(rows, channel="0000"):
|
||||
"""rows: list of {key, pimg, label}. Rebuilds imagesTr/labelsTr symlinks + dataset.json."""
|
||||
img_dir = os.path.join(raw_ds(), "imagesTr")
|
||||
lab_dir = os.path.join(raw_ds(), "labelsTr")
|
||||
shutil.rmtree(img_dir, ignore_errors=True)
|
||||
shutil.rmtree(lab_dir, ignore_errors=True)
|
||||
os.makedirs(img_dir, exist_ok=True)
|
||||
os.makedirs(lab_dir, exist_ok=True)
|
||||
seen = set()
|
||||
for r in rows:
|
||||
key, img, lab = r["key"], r["pimg"], r["label"]
|
||||
if key in seen:
|
||||
raise ValueError(f"duplicate case id in dataset: {key}")
|
||||
seen.add(key)
|
||||
if not os.path.exists(img):
|
||||
raise FileNotFoundError(f"image missing for {key}: {img}")
|
||||
if not os.path.exists(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(lab), os.path.join(lab_dir, f"{key}.nii.gz"))
|
||||
ds_json = {"channel_names": {"0": "t1c"},
|
||||
"labels": {"background": 0, "tumor": 1},
|
||||
"numTraining": len(seen),
|
||||
"file_ending": ".nii.gz"}
|
||||
with open(os.path.join(raw_ds(), "dataset.json"), "w") as f:
|
||||
json.dump(ds_json, f, indent=1)
|
||||
return len(seen)
|
||||
|
||||
|
||||
def dataset_case_keys():
|
||||
img_dir = os.path.join(raw_ds(), "imagesTr")
|
||||
out = []
|
||||
if os.path.isdir(img_dir):
|
||||
for fn in sorted(os.listdir(img_dir)):
|
||||
if fn.endswith("_0000.nii.gz"):
|
||||
out.append(fn[:-len("_0000.nii.gz")])
|
||||
return out
|
||||
|
||||
|
||||
def write_splits(train_keys, val_keys):
|
||||
val_keys = sorted(set(val_keys))
|
||||
train_keys = sorted(set(k for k in train_keys if k not in val_keys))
|
||||
folds = [{"train": train_keys, "val": val_keys}] * 5
|
||||
with open(os.path.join(preproc_ds(), "splits_final.json"), "w") as f:
|
||||
json.dump(folds, f)
|
||||
return len(train_keys), len(val_keys)
|
||||
|
||||
|
||||
# ---------------- label writing / selection (same gates as scripts/06) ----------------
|
||||
|
||||
def write_label(arr_u01, path, ref_itk):
|
||||
img = sitk.GetImageFromArray(arr_u01.astype(np.uint8))
|
||||
img.CopyInformation(ref_itk)
|
||||
os.makedirs(os.path.dirname(path), exist_ok=True)
|
||||
sitk.WriteImage(img, path, True)
|
||||
|
||||
|
||||
def pos_mask(p_tumor, tau_pos, min_cc_frac, vol_lo, vol_hi):
|
||||
"""Returns (mask, cc_frac, vol_mm3) if the positive gates pass, else None."""
|
||||
m = p_tumor >= tau_pos
|
||||
if m.sum() == 0:
|
||||
return None
|
||||
m = ndimage.median_filter(m, size=(3, 3, 3))
|
||||
if m.sum() == 0:
|
||||
return None
|
||||
lab, n = ndimage.label(m)
|
||||
sizes = ndimage.sum(m, lab, range(1, n + 1))
|
||||
big = (lab == (int(np.argmax(sizes)) + 1)).astype(np.uint8)
|
||||
cc_frac = float(big.sum()) / float(m.sum())
|
||||
vol_mm3 = int(big.sum())
|
||||
if cc_frac >= min_cc_frac and vol_lo <= vol_mm3 <= vol_hi:
|
||||
return big, cc_frac, vol_mm3
|
||||
return None
|
||||
|
||||
|
||||
def neg_frac_bg(p_tumor, vol, tau_neg, min_frac, min_interior=5000):
|
||||
interior = vol > 0.02
|
||||
if interior.sum() < min_interior:
|
||||
return None
|
||||
frac = float(((1.0 - p_tumor)[interior] >= tau_neg).mean())
|
||||
return frac if frac >= min_frac else None
|
||||
|
||||
|
||||
def tumor_probs_from_npz(npz_path):
|
||||
"""<case>.npz probabilities: channel-first (C,H,W,D), C=2 (bg, tumor)."""
|
||||
probs = np.load(npz_path, allow_pickle=False)["probabilities"]
|
||||
if probs.shape[0] != 2:
|
||||
raise ValueError(f"unexpected probability channels {probs.shape} in {npz_path}")
|
||||
return probs[1]
|
||||
|
||||
|
||||
# ---------------- longitudinal consistency (head-relative) ----------------
|
||||
#
|
||||
# The absolute patient-space resampling used by scripts/06 cannot work for this
|
||||
# data: each acquisition has its own scanner/patient coordinate system (table
|
||||
# offset, head orientation), so true labels from two visits of the same subject
|
||||
# resample to ~0 dice. Instead we gate on frame-invariant features:
|
||||
# * tumor centroid relative to the head centroid (crop space, 1mm voxels)
|
||||
# * tumor volume ratio
|
||||
# Two timepoints are compatible if they agree with at least one accepted
|
||||
# neighbor on both; greedy removal until stable (same scheme as scripts/06).
|
||||
|
||||
def head_centroid(vol):
|
||||
m = vol > 0.02
|
||||
if m.sum() < 5000:
|
||||
return None
|
||||
lab, n = ndimage.label(m)
|
||||
if n == 0:
|
||||
return None
|
||||
sizes = ndimage.sum(m, lab, range(1, n + 1))
|
||||
return np.array(ndimage.center_of_mass(lab == (int(np.argmax(sizes)) + 1)))
|
||||
|
||||
|
||||
def rel_tumor_features(mask, vol):
|
||||
"""Returns (centroid_rel_to_head [z,y,x] in mm, vol_mm3) or None."""
|
||||
if int(mask.sum()) == 0:
|
||||
return None
|
||||
tc = np.array(ndimage.center_of_mass(mask))
|
||||
hc = head_centroid(vol)
|
||||
rel = tc - hc if hc is not None else None
|
||||
return rel, int(mask.sum())
|
||||
|
||||
|
||||
def _compatible(fa, fb, max_rel_dist_mm, vol_ratio_max):
|
||||
ra, va = fa
|
||||
rb, vb = fb
|
||||
if ra is not None and rb is not None and np.linalg.norm(ra - rb) > max_rel_dist_mm:
|
||||
return False
|
||||
if va > 0 and vb > 0:
|
||||
ratio = max(va / vb, vb / va)
|
||||
if ratio > vol_ratio_max:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def consistency_filter(rows, out_dir, max_rel_dist_mm=40.0, vol_ratio_max=10.0):
|
||||
"""Greedy removal of positive timepoints inconsistent with accepted neighbors.
|
||||
|
||||
masks from <out_dir>/<key>_label.nii.gz, volumes from row["pimg"]. Returns n_rejected.
|
||||
"""
|
||||
by_subj = {}
|
||||
for r in rows:
|
||||
if r["role"] == "pos":
|
||||
by_subj.setdefault(r["subject"], []).append(r)
|
||||
rejected = 0
|
||||
for subj, tps in by_subj.items():
|
||||
if len(tps) < 2:
|
||||
continue
|
||||
tps.sort(key=lambda x: (x.get("date") or ""))
|
||||
feats = {}
|
||||
for i, r in enumerate(tps):
|
||||
mask = read_nii_arr(os.path.join(out_dir, r["key"] + "_label.nii.gz")).astype(bool)
|
||||
vol = read_nii_arr(r["pimg"]).astype("float32")
|
||||
feats[i] = rel_tumor_features(mask, vol)
|
||||
accepted = list(range(len(tps)))
|
||||
while True:
|
||||
changed = False
|
||||
for ai in list(accepted):
|
||||
neigh = [b for b in (ai - 1, ai + 1) if b in accepted]
|
||||
if not neigh or feats[ai] is None:
|
||||
continue
|
||||
if not any(_compatible(feats[ai], feats[b], max_rel_dist_mm, vol_ratio_max)
|
||||
for b in neigh if feats[b] is not None):
|
||||
tps[ai]["role"] = "rejected"
|
||||
accepted.remove(ai)
|
||||
rejected += 1
|
||||
changed = True
|
||||
break
|
||||
if not changed:
|
||||
break
|
||||
return rejected
|
||||
|
||||
|
||||
def load_voxel_stats():
|
||||
p = os.path.join(ROOT, "data/vols.json")
|
||||
return json.load(open(p)) if os.path.exists(p) else {}
|
||||
0
scripts_nnu/trainers/__init__.py
Normal file
0
scripts_nnu/trainers/__init__.py
Normal file
46
scripts_nnu/trainers/ntuh_pl_trainer.py
Normal file
46
scripts_nnu/trainers/ntuh_pl_trainer.py
Normal file
|
|
@ -0,0 +1,46 @@
|
|||
import os
|
||||
import torch
|
||||
|
||||
from nnunetv2.training.nnUNetTrainer.nnUNetTrainer import nnUNetTrainer
|
||||
|
||||
|
||||
class NTUHLPLTrainer(nnUNetTrainer):
|
||||
"""nnU-Net trainer for iterative pseudo-labeling on NTUH T1c tumor segmentation.
|
||||
|
||||
Env knobs (read in __init__, applied before initialize() builds optimizer/LR):
|
||||
NNU_PL_EPOCHS total epochs of this run (default 250)
|
||||
NNU_PL_LR initial lr, PolyLR decay (default 1e-2)
|
||||
|
||||
Warm start (round-to-round fine-tuning):
|
||||
NNU_PL_WARMSTART path to an nnU-Net checkpoint. Its FULL network weights,
|
||||
segmentation head included, are loaded in on_train_start(). The CLI
|
||||
-pretrained_weights flag deliberately skips `.seg_layers.` keys, which is
|
||||
wrong for iterative pseudo-labeling, hence this hook.
|
||||
"""
|
||||
|
||||
def __init__(self, plans: dict, configuration: str, fold: int, dataset_json: dict,
|
||||
device: torch.device = torch.device("cuda")):
|
||||
super().__init__(plans, configuration, fold, dataset_json, device)
|
||||
self.num_epochs = int(float(os.environ.get("NNU_PL_EPOCHS", "250")))
|
||||
self.initial_lr = float(os.environ.get("NNU_PL_LR", "1e-2"))
|
||||
self._pl_warmstart_done = False
|
||||
|
||||
def on_train_start(self):
|
||||
super().on_train_start()
|
||||
if self._pl_warmstart_done:
|
||||
return
|
||||
self._pl_warmstart_done = True
|
||||
ws = os.environ.get("NNU_PL_WARMSTART")
|
||||
if not ws:
|
||||
return
|
||||
mod = self.network
|
||||
if hasattr(mod, "module"):
|
||||
mod = mod.module
|
||||
if hasattr(mod, "_orig_mod"):
|
||||
mod = mod._orig_mod
|
||||
ckpt = torch.load(ws, map_location=self.device, weights_only=False)
|
||||
w = ckpt["network_weights"]
|
||||
mod.load_state_dict(w, strict=True)
|
||||
torch.cuda.empty_cache()
|
||||
print(f"[NTUHLPLTrainer] warm-started full network weights from {ws}", flush=True)
|
||||
self.print_to_log_file(f"warm-started full network weights from {ws}")
|
||||
Loading…
Reference in a new issue