longitudinal/scripts/06_pseudo_label.py
Furen Xiao b6fa62a763 feat: initial project structure
Add .gitignore, AGENTS.md, scripts directory, and src directory to initialize the repository.
2026-09-25 16:00:37 +08:00

211 lines
No EOL
8.4 KiB
Python

"""Iterative pseudo-labeling round.
torchrun --standalone --nproc_per_node=3 scripts/06_pseudo_label.py \
--ckpt runs/round0/best.pt --unlabeled data/manifests/unlabeled_pool.jsonl --out data/pseudo/round1
Per volume: sliding-window tumor probabilities (TTA). Selection:
pos: p_tumor >= tau_pos, largest-CC fraction >= min_cc_frac, volume within [vol_lo, vol_hi]
neg: >= neg_frac of interior voxels have p_bg >= tau_neg
Then a per-subject longitudinal consistency filter over accepted positive timepoints.
Writes: <out>/rows.jsonl, <out>/<key>.pseudo.nii.gz, <out>/summary.json
"""
import os
import sys
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import json
import argparse
import numpy as np
import SimpleITK as sitk
from scipy import ndimage
from src.common import ROOT, d, load_jsonl, save_jsonl, read_nii_arr
from src import training
def grid_info(key):
p = os.path.join(d("data/procmeta"), key + ".json")
m = json.load(open(p))
origin = np.array(m["origin"])
R = np.array(m["direction"]).reshape(3, 3) # row_dir, col_dir, slice_dir
cv = np.array(m["crop_vox"])
o = origin + cv[0] * R[0] + cv[1] * R[1] + cv[2] * R[2]
return o, R
def grid_itk(shape, key):
o, R = grid_info(key)
img = sitk.GetImageFromArray(np.zeros(shape, np.uint8))
img.SetSpacing((1.0, 1.0, 1.0))
img.SetOrigin(tuple(float(x) for x in o))
img.SetDirection(tuple(float(x) for x in R.flatten()))
return img
def load_pseudo(outdir, key):
return read_nii_arr(os.path.join(outdir, key + ".pseudo.nii.gz")).astype(bool)
def dice_a_on_b(keyA, arrA, keyB, arrB):
a_itk = sitk.GetImageFromArray(arrA.astype(np.uint8) % 255)
a_itk.CopyInformation(grid_itk(arrA.shape, keyA))
ref = grid_itk(arrB.shape, keyB)
am = sitk.GetArrayFromImage(sitk.Resample(a_itk, ref, sitk.Transform(), sitk.sitkNearestNeighbor, 0)).astype(bool)
bm = arrB
inter = (am & bm).sum()
return float(2 * inter / max(int(am.sum()) + int(bm.sum()), 1))
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--ckpt", required=True)
ap.add_argument("--unlabeled", required=True)
ap.add_argument("--out", required=True)
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("--cons-dice", type=float, default=0.30)
ap.add_argument("--win", type=int, default=96)
ap.add_argument("--step", type=int, default=64)
args = ap.parse_args()
import torch
rank = int(os.environ.get("RANK", 0))
world = int(os.environ.get("WORLD_SIZE", 1))
local_rank = int(os.environ.get("LOCAL_RANK", 0))
if world > 1:
import torch.distributed as dist
dist.init_process_group("nccl")
torch.cuda.set_device(local_rank)
device = f"cuda:{local_rank}"
outname = os.path.basename(os.path.normpath(args.out))
out = d(os.path.join("data/pseudo", outname))
rows = load_jsonl(args.unlabeled)
done = set()
part = os.path.join(out, f"part{rank}.jsonl")
if os.path.exists(part):
for r in load_jsonl(part):
done.add(r["key"])
rows = [r for r in rows if r["key"] not in done]
shard = rows[rank::world]
vstats = {}
if os.path.exists(os.path.join(ROOT, "data/vols.json")):
vstats = json.load(open(os.path.join(ROOT, "data/vols.json")))
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)
if rank == 0:
print(f"[pseudo:rank0] pool={len(rows)} shard={len(shard)} vol_range=[{vol_lo:.0f},{vol_hi:.0f}]mm3", flush=True)
ckpt = torch.load(args.ckpt, map_location=device, weights_only=True)
model = training.build_model(device=device)
model.load_state_dict(ckpt["model"])
model.eval()
buf = []
for i, r in enumerate(shard, 1):
key = r["key"]
try:
vol = read_nii_arr(r["pimg"]).astype(np.float32)
p = training.sliding_window_probs(model, vol, device, args.win, args.step, tta=True)
out_row = {"key": key, "subject": r["subject"], "date": r.get("date"),
"source": r.get("source"), "pimg": r["pimg"], "plabel": None,
"role": "rej", "vol_mm3": 0, "maxp": round(float(p.max()), 4)}
m = p >= args.tau_pos
if m.sum() > 0:
m = ndimage.median_filter(m, size=(3, 3, 3))
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 >= args.min_cc_frac and vol_lo <= vol_mm3 <= vol_hi:
out_row.update({"role": "pos", "vol_mm3": vol_mm3, "cc_frac": round(cc_frac, 3)})
sitk.WriteImage(sitk.GetImageFromArray(big), os.path.join(out, key + ".pseudo.nii.gz"), True)
interior = vol > 0.02
if out_row["role"] != "pos" and interior.sum() > 5000:
frac = float(((1.0 - p)[interior] >= args.tau_neg).mean())
if frac >= args.neg_frac:
out_row.update({"role": "neg", "neg_conf": round(frac, 4)})
buf.append(out_row)
except Exception as e: # noqa
print(f"[pseudo:rank{rank}] {key} ERR {e!r}", flush=True)
if i % 20 == 0:
with open(part, "a") as f:
for b in buf:
f.write(json.dumps(b) + "\n")
buf = []
print(f"[pseudo:rank{rank}] {i}/{len(shard)}", flush=True)
if buf:
with open(part, "a") as f:
for b in buf:
f.write(json.dumps(b) + "\n")
buf = []
if world > 1:
dist.barrier()
if rank != 0:
dist.destroy_process_group() if world > 1 else None
return
merged = []
for k in range(world):
p = os.path.join(out, f"part{k}.jsonl")
if os.path.exists(p):
merged.extend(load_jsonl(p))
# longitudinal consistency filter across accepted positive timepoints
by_subj = {}
for r in merged:
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 ""))
accepted = list(range(len(tps)))
while True:
changed = False
for ai in list(accepted):
neigh = [ai - 1, ai + 1]
neigh = [b for b in neigh if b in accepted]
if not neigh:
continue
ds = []
arrA = load_pseudo(out, tps[ai]["key"]).astype(np.uint8) * 255
for bi in neigh:
arrB = load_pseudo(out, tps[bi]["key"]).astype(np.uint8) * 255
ds.append(dice_a_on_b(tps[ai]["key"], arrA, tps[bi]["key"], arrB))
if max(ds) < args.cons_dice:
tps[ai]["role"] = "rejected"
accepted.remove(ai)
rejected += 1
changed = True
break
if not changed:
break
for r in merged:
if r["role"] == "pos":
r["plabel"] = os.path.join(out, r["key"] + ".pseudo.nii.gz")
save_jsonl(merged, os.path.join(out, "rows.jsonl"))
posv = [r["vol_mm3"] for r in merged if r["role"] == "pos"]
summ = {
"n_pool": len(merged),
"n_pos": sum(1 for r in merged if r["role"] == "pos"),
"n_neg": sum(1 for r in merged if r["role"] == "neg"),
"n_rejected_consistency": rejected,
"n_other_rej": sum(1 for r in merged if r["role"] == "rej"),
"pos_vol_mm3": {"med": float(np.median(posv)) if posv else 0,
"p5": float(np.percentile(posv, 5)) if posv else 0,
"p95": float(np.percentile(posv, 95)) if posv else 0},
"tau_pos": args.tau_pos, "vol_range": [vol_lo, vol_hi],
}
with open(os.path.join(out, "summary.json"), "w") as f:
json.dump(summ, f, indent=1)
print(f"[pseudo] round {outname}: {json.dumps(summ)}", flush=True)
if world > 1:
dist.destroy_process_group()
if __name__ == "__main__":
main()