Add .gitignore, AGENTS.md, scripts directory, and src directory to initialize the repository.
211 lines
No EOL
8.4 KiB
Python
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() |