"""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: /rows.jsonl, /.pseudo.nii.gz, /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, DATA, 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) # columns = x, y, z axis directions cv = np.array(m["crop_vox"]) # crop starts in (z, y, x) array coords o = origin + cv[2] * R[:, 0] + cv[1] * R[:, 1] + cv[0] * 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(DATA, "vols.json")): vstats = json.load(open(os.path.join(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()