Add tools
Browse files- tools/create_amos_ct_mr_manifest.py +196 -0
- tools/create_dataset_json.py +98 -0
- tools/eval.py +45 -0
- tools/export_source_memory.py +98 -0
- tools/inspect_dataset.py +25 -0
- tools/make_sfda_manifest.py +39 -0
- tools/remap_labels_in_manifest.py +89 -0
- tools/train.py +46 -0
tools/create_amos_ct_mr_manifest.py
ADDED
|
@@ -0,0 +1,196 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python
|
| 2 |
+
"""Create a SACFlow manifest for AMOS22 CT->MRI or MRI->CT.
|
| 3 |
+
|
| 4 |
+
Expected AMOS root:
|
| 5 |
+
dataset.json
|
| 6 |
+
labeled_data_meta_0000_0599.csv
|
| 7 |
+
imagesTr/ labelsTr/ imagesVa/ labelsVa/ imagesTs/ labelsTs/
|
| 8 |
+
|
| 9 |
+
This v2 script handles the common ambiguity in AMOS case numbering:
|
| 10 |
+
- zero_based: amos_0000..amos_0499 = CT, amos_0500..amos_0599 = MRI
|
| 11 |
+
- one_based: amos_0001..amos_0500 = CT, amos_0501..amos_0600 = MRI
|
| 12 |
+
- auto: infers zero_based vs one_based from the observed min/max IDs
|
| 13 |
+
|
| 14 |
+
It also tries metadata first when a modality column exists. Always verify printed counts.
|
| 15 |
+
"""
|
| 16 |
+
from __future__ import annotations
|
| 17 |
+
import argparse, csv, json, re
|
| 18 |
+
from pathlib import Path
|
| 19 |
+
from typing import Optional
|
| 20 |
+
|
| 21 |
+
NIFTI_EXTS = [".nii.gz", ".nii", ".mha", ".mhd"]
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
def strip_ext(p: Path) -> str:
|
| 25 |
+
name = p.name
|
| 26 |
+
for ext in NIFTI_EXTS:
|
| 27 |
+
if name.endswith(ext):
|
| 28 |
+
return name[:-len(ext)]
|
| 29 |
+
return p.stem
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def case_num(s: str) -> Optional[int]:
|
| 33 |
+
m = re.search(r"(\d+)$", str(s))
|
| 34 |
+
return int(m.group(1)) if m else None
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
def list_imgs(d: Path):
|
| 38 |
+
out = []
|
| 39 |
+
for ext in NIFTI_EXTS:
|
| 40 |
+
out += list(d.glob(f"*{ext}"))
|
| 41 |
+
return sorted(out)
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
def load_modality_map(meta_csv: Optional[Path]):
|
| 45 |
+
if not meta_csv or not meta_csv.exists():
|
| 46 |
+
return {}, None, None
|
| 47 |
+
with open(meta_csv, newline="") as f:
|
| 48 |
+
rows = list(csv.DictReader(f))
|
| 49 |
+
if not rows:
|
| 50 |
+
return {}, None, None
|
| 51 |
+
cols = list(rows[0].keys())
|
| 52 |
+
mod_col = None
|
| 53 |
+
# Look for a column whose values contain explicit modality strings.
|
| 54 |
+
for c in cols:
|
| 55 |
+
allvals = {str(r.get(c, "")).strip().lower() for r in rows}
|
| 56 |
+
if any(v in {"ct", "mri", "mr"} for v in allvals):
|
| 57 |
+
mod_col = c
|
| 58 |
+
break
|
| 59 |
+
id_cols = [c for c in cols if any(k in c.lower() for k in ["amos", "case", "id", "name", "file", "image"])]
|
| 60 |
+
if not id_cols:
|
| 61 |
+
id_cols = list(cols)
|
| 62 |
+
mapping = {}
|
| 63 |
+
if mod_col:
|
| 64 |
+
for r in rows:
|
| 65 |
+
mod = str(r.get(mod_col, "")).strip().lower()
|
| 66 |
+
if mod == "mr":
|
| 67 |
+
mod = "mri"
|
| 68 |
+
if mod not in {"ct", "mri"}:
|
| 69 |
+
continue
|
| 70 |
+
for c in id_cols:
|
| 71 |
+
v = str(r.get(c, "")).strip()
|
| 72 |
+
if not v:
|
| 73 |
+
continue
|
| 74 |
+
stem = strip_ext(Path(v))
|
| 75 |
+
keys = {v, Path(v).name, stem}
|
| 76 |
+
n = case_num(v)
|
| 77 |
+
if n is not None:
|
| 78 |
+
keys.update({str(n), f"{n:04d}", f"amos_{n:04d}"})
|
| 79 |
+
for k in keys:
|
| 80 |
+
mapping[str(k).lower()] = mod
|
| 81 |
+
return mapping, mod_col, id_cols
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
def infer_id_scheme(root: Path) -> str:
|
| 85 |
+
ids = []
|
| 86 |
+
for split in ["Tr", "Va", "Ts"]:
|
| 87 |
+
d = root / f"images{split}"
|
| 88 |
+
for img in list_imgs(d):
|
| 89 |
+
n = case_num(strip_ext(img))
|
| 90 |
+
if n is not None:
|
| 91 |
+
ids.append(n)
|
| 92 |
+
if not ids:
|
| 93 |
+
return "zero_based"
|
| 94 |
+
mn, mx = min(ids), max(ids)
|
| 95 |
+
# AMOS labeled release is 600 cases. If IDs run 1..600, use one_based.
|
| 96 |
+
# If IDs run 0..599, use zero_based.
|
| 97 |
+
if mn >= 1 and mx >= 600:
|
| 98 |
+
return "one_based"
|
| 99 |
+
return "zero_based"
|
| 100 |
+
|
| 101 |
+
|
| 102 |
+
def infer_modality(stem: str, mod_map: dict, id_scheme: str, ct_max_id: Optional[int], mri_min_id: Optional[int]):
|
| 103 |
+
keys = [stem.lower(), Path(stem).name.lower()]
|
| 104 |
+
n = case_num(stem)
|
| 105 |
+
if n is not None:
|
| 106 |
+
keys += [str(n), f"{n:04d}", f"amos_{n:04d}"]
|
| 107 |
+
for k in keys:
|
| 108 |
+
if k in mod_map:
|
| 109 |
+
return mod_map[k]
|
| 110 |
+
if n is None:
|
| 111 |
+
return None
|
| 112 |
+
if ct_max_id is not None:
|
| 113 |
+
return "ct" if n <= ct_max_id else "mri"
|
| 114 |
+
if mri_min_id is not None:
|
| 115 |
+
return "mri" if n >= mri_min_id else "ct"
|
| 116 |
+
if id_scheme == "one_based":
|
| 117 |
+
return "ct" if 1 <= n <= 500 else "mri"
|
| 118 |
+
# zero_based
|
| 119 |
+
return "ct" if 0 <= n <= 499 else "mri"
|
| 120 |
+
|
| 121 |
+
|
| 122 |
+
def make_items(root: Path, split: str, modality: str, mod_map: dict, prefix: str, id_scheme: str, ct_max_id: Optional[int], mri_min_id: Optional[int]):
|
| 123 |
+
img_dir = root / f"images{split}"
|
| 124 |
+
lab_dir = root / f"labels{split}"
|
| 125 |
+
items = []
|
| 126 |
+
if not img_dir.exists():
|
| 127 |
+
return items
|
| 128 |
+
for img in list_imgs(img_dir):
|
| 129 |
+
stem = strip_ext(img)
|
| 130 |
+
mod = infer_modality(stem, mod_map, id_scheme, ct_max_id, mri_min_id)
|
| 131 |
+
if mod != modality:
|
| 132 |
+
continue
|
| 133 |
+
lab = None
|
| 134 |
+
if lab_dir.exists():
|
| 135 |
+
for ext in NIFTI_EXTS:
|
| 136 |
+
cand = lab_dir / f"{stem}{ext}"
|
| 137 |
+
if cand.exists():
|
| 138 |
+
lab = cand
|
| 139 |
+
break
|
| 140 |
+
item = {"id": f"{prefix}_{split}_{stem}", "case_id": stem, "modality": modality, "image": str(img.resolve())}
|
| 141 |
+
if lab is not None:
|
| 142 |
+
item["label"] = str(lab.resolve())
|
| 143 |
+
items.append(item)
|
| 144 |
+
return sorted(items, key=lambda x: x["id"])
|
| 145 |
+
|
| 146 |
+
|
| 147 |
+
def main():
|
| 148 |
+
ap = argparse.ArgumentParser()
|
| 149 |
+
ap.add_argument("--amos-root", required=True)
|
| 150 |
+
ap.add_argument("--output", required=True)
|
| 151 |
+
ap.add_argument("--source", choices=["ct", "mri"], default="ct")
|
| 152 |
+
ap.add_argument("--target", choices=["ct", "mri"], default="mri")
|
| 153 |
+
ap.add_argument("--meta", default=None)
|
| 154 |
+
ap.add_argument("--name", default=None)
|
| 155 |
+
ap.add_argument("--id-scheme", choices=["auto", "zero_based", "one_based"], default="auto")
|
| 156 |
+
ap.add_argument("--ct-max-id", type=int, default=None, help="Override: IDs <= this are CT")
|
| 157 |
+
ap.add_argument("--mri-min-id", type=int, default=None, help="Override: IDs >= this are MRI")
|
| 158 |
+
args = ap.parse_args()
|
| 159 |
+
root = Path(args.amos_root).resolve()
|
| 160 |
+
meta = Path(args.meta).resolve() if args.meta else root / "labeled_data_meta_0000_0599.csv"
|
| 161 |
+
mod_map, mod_col, id_cols = load_modality_map(meta)
|
| 162 |
+
id_scheme = infer_id_scheme(root) if args.id_scheme == "auto" else args.id_scheme
|
| 163 |
+
print(f"Metadata: {meta if meta.exists() else 'not found'}")
|
| 164 |
+
print(f"Detected modality column: {mod_col}; candidate id columns: {id_cols}")
|
| 165 |
+
if not mod_map:
|
| 166 |
+
print(f"WARNING: no explicit modality map from metadata; using numeric ID scheme: {id_scheme}")
|
| 167 |
+
if args.ct_max_id is not None:
|
| 168 |
+
print(f"Override active: IDs <= {args.ct_max_id} treated as CT")
|
| 169 |
+
if args.mri_min_id is not None:
|
| 170 |
+
print(f"Override active: IDs >= {args.mri_min_id} treated as MRI")
|
| 171 |
+
manifest = {"name": args.name or f"amos_{args.source}2{args.target}"}
|
| 172 |
+
for split_name, code in [("train", "Tr"), ("val", "Va"), ("test", "Ts")]:
|
| 173 |
+
manifest[f"source_{split_name}"] = make_items(root, code, args.source, mod_map, "source", id_scheme, args.ct_max_id, args.mri_min_id)
|
| 174 |
+
manifest[f"target_{split_name}"] = make_items(root, code, args.target, mod_map, "target", id_scheme, args.ct_max_id, args.mri_min_id)
|
| 175 |
+
out = Path(args.output)
|
| 176 |
+
out.parent.mkdir(parents=True, exist_ok=True)
|
| 177 |
+
with open(out, "w") as f:
|
| 178 |
+
json.dump(manifest, f, indent=2)
|
| 179 |
+
print(f"Wrote {out}")
|
| 180 |
+
total_source = 0
|
| 181 |
+
total_target = 0
|
| 182 |
+
for k, v in manifest.items():
|
| 183 |
+
if isinstance(v, list):
|
| 184 |
+
with_labels = sum(1 for item in v if "label" in item)
|
| 185 |
+
print(f"{k}: {len(v)} items, {with_labels} labels")
|
| 186 |
+
if k.startswith("source_"):
|
| 187 |
+
total_source += len(v)
|
| 188 |
+
if k.startswith("target_"):
|
| 189 |
+
total_target += len(v)
|
| 190 |
+
print(f"TOTAL source ({args.source}): {total_source}")
|
| 191 |
+
print(f"TOTAL target ({args.target}): {total_target}")
|
| 192 |
+
print("\nIf using AMOS22 labeled 600 cases, expect about 500 CT and 100 MRI. If you see 499/101, rerun with --id-scheme one_based or --ct-max-id 500.")
|
| 193 |
+
|
| 194 |
+
|
| 195 |
+
if __name__ == "__main__":
|
| 196 |
+
main()
|
tools/create_dataset_json.py
ADDED
|
@@ -0,0 +1,98 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python
|
| 2 |
+
from __future__ import annotations
|
| 3 |
+
import argparse
|
| 4 |
+
import json
|
| 5 |
+
from pathlib import Path
|
| 6 |
+
import random
|
| 7 |
+
|
| 8 |
+
NIFTI_EXTS = [".nii.gz", ".nii", ".mha", ".mhd"]
|
| 9 |
+
|
| 10 |
+
def list_images(folder):
|
| 11 |
+
folder = Path(folder)
|
| 12 |
+
files = []
|
| 13 |
+
for ext in NIFTI_EXTS:
|
| 14 |
+
files.extend(folder.glob(f"*{ext}"))
|
| 15 |
+
return sorted(files)
|
| 16 |
+
|
| 17 |
+
def stem_nii(p: Path):
|
| 18 |
+
name = p.name
|
| 19 |
+
for ext in NIFTI_EXTS:
|
| 20 |
+
if name.endswith(ext):
|
| 21 |
+
return name[:-len(ext)]
|
| 22 |
+
return p.stem
|
| 23 |
+
|
| 24 |
+
def match_labels(images, labels_dir):
|
| 25 |
+
if labels_dir is None:
|
| 26 |
+
return {stem_nii(p): None for p in images}
|
| 27 |
+
labels = list_images(labels_dir)
|
| 28 |
+
lab_map = {stem_nii(p): p for p in labels}
|
| 29 |
+
out = {}
|
| 30 |
+
for img in images:
|
| 31 |
+
sid = stem_nii(img)
|
| 32 |
+
if sid in lab_map:
|
| 33 |
+
out[sid] = lab_map[sid]
|
| 34 |
+
else:
|
| 35 |
+
# fuzzy: strip common prefixes
|
| 36 |
+
cand = None
|
| 37 |
+
for k, v in lab_map.items():
|
| 38 |
+
if sid in k or k in sid:
|
| 39 |
+
cand = v; break
|
| 40 |
+
out[sid] = cand
|
| 41 |
+
return out
|
| 42 |
+
|
| 43 |
+
def make_items(images_dir, labels_dir=None, domain=""):
|
| 44 |
+
imgs = list_images(images_dir)
|
| 45 |
+
labels = match_labels(imgs, labels_dir)
|
| 46 |
+
items = []
|
| 47 |
+
for p in imgs:
|
| 48 |
+
sid = stem_nii(p)
|
| 49 |
+
item = {"id": f"{domain}_{sid}" if domain else sid, "image": str(p)}
|
| 50 |
+
if labels.get(sid):
|
| 51 |
+
item["label"] = str(labels[sid])
|
| 52 |
+
items.append(item)
|
| 53 |
+
return items
|
| 54 |
+
|
| 55 |
+
def split_items(items, val_fraction, test_fraction, seed):
|
| 56 |
+
items = list(items)
|
| 57 |
+
random.Random(seed).shuffle(items)
|
| 58 |
+
n = len(items)
|
| 59 |
+
n_test = int(round(n * test_fraction))
|
| 60 |
+
n_val = int(round(n * val_fraction))
|
| 61 |
+
test = items[:n_test]
|
| 62 |
+
val = items[n_test:n_test+n_val]
|
| 63 |
+
train = items[n_test+n_val:]
|
| 64 |
+
return train, val, test
|
| 65 |
+
|
| 66 |
+
def main():
|
| 67 |
+
ap = argparse.ArgumentParser()
|
| 68 |
+
ap.add_argument("--name", required=True)
|
| 69 |
+
ap.add_argument("--source-images", required=True)
|
| 70 |
+
ap.add_argument("--source-labels", default=None)
|
| 71 |
+
ap.add_argument("--target-images", required=True)
|
| 72 |
+
ap.add_argument("--target-labels", default=None)
|
| 73 |
+
ap.add_argument("--output", required=True)
|
| 74 |
+
ap.add_argument("--val-fraction", type=float, default=0.2)
|
| 75 |
+
ap.add_argument("--test-fraction", type=float, default=0.2)
|
| 76 |
+
ap.add_argument("--seed", type=int, default=1337)
|
| 77 |
+
args = ap.parse_args()
|
| 78 |
+
src = make_items(args.source_images, args.source_labels, "source")
|
| 79 |
+
tgt = make_items(args.target_images, args.target_labels, "target")
|
| 80 |
+
src_train, src_val, src_test = split_items(src, args.val_fraction, args.test_fraction, args.seed)
|
| 81 |
+
tgt_train, tgt_val, tgt_test = split_items(tgt, args.val_fraction, args.test_fraction, args.seed)
|
| 82 |
+
man = {
|
| 83 |
+
"name": args.name,
|
| 84 |
+
"source_train": src_train,
|
| 85 |
+
"source_val": src_val,
|
| 86 |
+
"source_test": src_test,
|
| 87 |
+
"target_train": tgt_train,
|
| 88 |
+
"target_val": tgt_val,
|
| 89 |
+
"target_test": tgt_test,
|
| 90 |
+
}
|
| 91 |
+
out = Path(args.output)
|
| 92 |
+
out.parent.mkdir(parents=True, exist_ok=True)
|
| 93 |
+
with open(out, "w") as f:
|
| 94 |
+
json.dump(man, f, indent=2)
|
| 95 |
+
print(f"Wrote {out}")
|
| 96 |
+
print({k: len(v) for k, v in man.items() if isinstance(v, list)})
|
| 97 |
+
if __name__ == "__main__":
|
| 98 |
+
main()
|
tools/eval.py
ADDED
|
@@ -0,0 +1,45 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python
|
| 2 |
+
from __future__ import annotations
|
| 3 |
+
import argparse
|
| 4 |
+
import json
|
| 5 |
+
from pathlib import Path
|
| 6 |
+
import torch
|
| 7 |
+
from sacflow.utils.config import load_yaml
|
| 8 |
+
from sacflow.utils.misc import seed_everything, ensure_dir
|
| 9 |
+
from sacflow.utils.distributed import init_distributed, cleanup, is_main_process
|
| 10 |
+
from sacflow.data.loader import build_loader
|
| 11 |
+
from sacflow.models.unet3d import build_model
|
| 12 |
+
from sacflow.engine.train_loop import evaluate
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
def main():
|
| 16 |
+
ap = argparse.ArgumentParser()
|
| 17 |
+
ap.add_argument("--config", required=True)
|
| 18 |
+
ap.add_argument("--checkpoint", default=None)
|
| 19 |
+
ap.add_argument("--split", default=None)
|
| 20 |
+
args = ap.parse_args()
|
| 21 |
+
cfg = load_yaml(args.config)
|
| 22 |
+
if args.checkpoint:
|
| 23 |
+
cfg.setdefault("eval", {})["checkpoint"] = args.checkpoint
|
| 24 |
+
if args.split:
|
| 25 |
+
cfg.setdefault("eval", {})["split"] = args.split
|
| 26 |
+
seed_everything(int(cfg.get("seed", 1337)))
|
| 27 |
+
device = init_distributed(cfg.get("distributed", {}).get("backend", "nccl"))
|
| 28 |
+
model = build_model(cfg).to(device)
|
| 29 |
+
ckpt_path = cfg.get("eval", {}).get("checkpoint") or cfg.get("train", {}).get("source_checkpoint")
|
| 30 |
+
if ckpt_path is None:
|
| 31 |
+
ckpt_path = str(Path(cfg["output_dir"]) / "checkpoints" / "best.pt")
|
| 32 |
+
ckpt = torch.load(ckpt_path, map_location="cpu")
|
| 33 |
+
model.load_state_dict(ckpt.get("model", ckpt), strict=False)
|
| 34 |
+
split = cfg.get("eval", {}).get("split", "target_test")
|
| 35 |
+
loader = build_loader(cfg, split=split, training=False, require_label=True)
|
| 36 |
+
metrics = evaluate(model, loader, cfg, device)
|
| 37 |
+
if is_main_process():
|
| 38 |
+
print(json.dumps(metrics, indent=2))
|
| 39 |
+
out = ensure_dir(Path(cfg["output_dir"]) / "eval")
|
| 40 |
+
with open(out / f"metrics_{split}.json", "w") as f:
|
| 41 |
+
json.dump(metrics, f, indent=2)
|
| 42 |
+
cleanup()
|
| 43 |
+
|
| 44 |
+
if __name__ == "__main__":
|
| 45 |
+
main()
|
tools/export_source_memory.py
ADDED
|
@@ -0,0 +1,98 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python
|
| 2 |
+
from __future__ import annotations
|
| 3 |
+
import argparse
|
| 4 |
+
from pathlib import Path
|
| 5 |
+
import copy
|
| 6 |
+
import torch
|
| 7 |
+
import torch.nn.functional as F
|
| 8 |
+
from tqdm import tqdm
|
| 9 |
+
from sacflow.utils.config import load_yaml
|
| 10 |
+
from sacflow.utils.misc import seed_everything, ensure_dir, move_to_device
|
| 11 |
+
from sacflow.utils.distributed import init_distributed, cleanup, is_main_process
|
| 12 |
+
from sacflow.data.loader import build_loader
|
| 13 |
+
from sacflow.models.unet3d import build_model
|
| 14 |
+
from sacflow.methods.task_space import centered_classifier_basis, project_task_and_residual
|
| 15 |
+
from sacflow.methods.source_memory import hard_onehot, class_moments, save_source_memory
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
def main():
|
| 19 |
+
ap = argparse.ArgumentParser()
|
| 20 |
+
ap.add_argument("--config", required=True)
|
| 21 |
+
ap.add_argument("--checkpoint", required=True)
|
| 22 |
+
ap.add_argument("--output", required=True)
|
| 23 |
+
ap.add_argument("--num-passes", type=int, default=3, help="Number of random-crop passes over source_train for memory estimation")
|
| 24 |
+
args = ap.parse_args()
|
| 25 |
+
cfg = load_yaml(args.config)
|
| 26 |
+
# Source memory should estimate source feature statistics, not augmentation noise.
|
| 27 |
+
cfg = copy.deepcopy(cfg)
|
| 28 |
+
cfg.setdefault("data", {}).setdefault("augmentation", {})
|
| 29 |
+
cfg["data"]["augmentation"] = {"random_flip": False, "random_intensity_shift": 0.0, "random_intensity_scale": 0.0}
|
| 30 |
+
seed_everything(int(cfg.get("seed", 1337)))
|
| 31 |
+
device = init_distributed(cfg.get("distributed", {}).get("backend", "nccl"))
|
| 32 |
+
model = build_model(cfg).to(device)
|
| 33 |
+
ckpt = torch.load(args.checkpoint, map_location="cpu")
|
| 34 |
+
model.load_state_dict(ckpt.get("model", ckpt), strict=False)
|
| 35 |
+
model.eval()
|
| 36 |
+
loader = build_loader(cfg, split="source_train", training=True, require_label=True)
|
| 37 |
+
W = model.final_classifier_weight().detach().to(device)
|
| 38 |
+
Q = centered_classifier_basis(W)
|
| 39 |
+
C = cfg["data"]["num_classes"]
|
| 40 |
+
mu_acc = None
|
| 41 |
+
var_acc = None
|
| 42 |
+
feat_mu_acc = None
|
| 43 |
+
feat_var_acc = None
|
| 44 |
+
count_acc = None
|
| 45 |
+
with torch.no_grad():
|
| 46 |
+
for pass_idx in range(max(1, args.num_passes)):
|
| 47 |
+
for batch in tqdm(loader, desc=f"source memory pass {pass_idx+1}/{max(1,args.num_passes)}", disable=not is_main_process()):
|
| 48 |
+
batch = move_to_device(batch, device)
|
| 49 |
+
logits, feats = model(batch["image"], return_features=True)
|
| 50 |
+
feat = feats["prelogit"]
|
| 51 |
+
task, residual = project_task_and_residual(feat, Q)
|
| 52 |
+
labels = batch["label"]
|
| 53 |
+
if labels.shape[-3:] != residual.shape[-3:]:
|
| 54 |
+
labels = F.interpolate(labels[:,None].float(), size=residual.shape[-3:], mode="nearest")[:,0].long()
|
| 55 |
+
probs = hard_onehot(labels, C)
|
| 56 |
+
mu, std, counts = class_moments(residual, probs)
|
| 57 |
+
fmu, fstd, _ = class_moments(feat, probs)
|
| 58 |
+
var = std.pow(2)
|
| 59 |
+
fvar = fstd.pow(2)
|
| 60 |
+
if mu_acc is None:
|
| 61 |
+
mu_acc = mu * counts[:,None]
|
| 62 |
+
var_acc = (var + mu.pow(2)) * counts[:,None]
|
| 63 |
+
feat_mu_acc = fmu * counts[:,None]
|
| 64 |
+
feat_var_acc = (fvar + fmu.pow(2)) * counts[:,None]
|
| 65 |
+
count_acc = counts
|
| 66 |
+
else:
|
| 67 |
+
mu_acc += mu * counts[:,None]
|
| 68 |
+
var_acc += (var + mu.pow(2)) * counts[:,None]
|
| 69 |
+
feat_mu_acc += fmu * counts[:,None]
|
| 70 |
+
feat_var_acc += (fvar + fmu.pow(2)) * counts[:,None]
|
| 71 |
+
count_acc += counts
|
| 72 |
+
count = count_acc.clamp_min(1e-6)
|
| 73 |
+
mu = mu_acc / count[:,None]
|
| 74 |
+
second = var_acc / count[:,None]
|
| 75 |
+
var = (second - mu.pow(2)).clamp_min(1e-6)
|
| 76 |
+
fmu = feat_mu_acc / count[:,None]
|
| 77 |
+
fsecond = feat_var_acc / count[:,None]
|
| 78 |
+
fvar = (fsecond - fmu.pow(2)).clamp_min(1e-6)
|
| 79 |
+
mem = {
|
| 80 |
+
"classifier_weight": W.detach().cpu(),
|
| 81 |
+
"task_basis_Q": Q.detach().cpu(),
|
| 82 |
+
"residual_mu": mu.detach().cpu(),
|
| 83 |
+
"residual_std": torch.sqrt(var).detach().cpu(),
|
| 84 |
+
"feature_mu": fmu.detach().cpu(),
|
| 85 |
+
"feature_std": torch.sqrt(fvar).detach().cpu(),
|
| 86 |
+
"class_counts": count.detach().cpu(),
|
| 87 |
+
"num_classes": C,
|
| 88 |
+
"feature_dim": W.shape[1],
|
| 89 |
+
"note": "Compact source memory. No raw source images stored. Estimated from source_train random crops with augmentation disabled.",
|
| 90 |
+
"num_passes": args.num_passes,
|
| 91 |
+
}
|
| 92 |
+
if is_main_process():
|
| 93 |
+
save_source_memory(args.output, mem)
|
| 94 |
+
print(f"Saved source memory to {args.output}")
|
| 95 |
+
cleanup()
|
| 96 |
+
|
| 97 |
+
if __name__ == "__main__":
|
| 98 |
+
main()
|
tools/inspect_dataset.py
ADDED
|
@@ -0,0 +1,25 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python
|
| 2 |
+
from __future__ import annotations
|
| 3 |
+
import argparse, json
|
| 4 |
+
from pathlib import Path
|
| 5 |
+
import nibabel as nib
|
| 6 |
+
import numpy as np
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
def main():
|
| 10 |
+
ap = argparse.ArgumentParser()
|
| 11 |
+
ap.add_argument("--manifest", required=True)
|
| 12 |
+
args = ap.parse_args()
|
| 13 |
+
man = json.load(open(args.manifest))
|
| 14 |
+
for split, items in man.items():
|
| 15 |
+
if not isinstance(items, list):
|
| 16 |
+
continue
|
| 17 |
+
print(split, len(items))
|
| 18 |
+
for it in items[:3]:
|
| 19 |
+
img = nib.load(it["image"])
|
| 20 |
+
print(" ", it.get("id"), img.shape, img.header.get_zooms()[:3], "label", bool(it.get("label")))
|
| 21 |
+
if it.get("label"):
|
| 22 |
+
lab = np.asanyarray(nib.load(it["label"]).dataobj)
|
| 23 |
+
print(" labels", np.unique(lab)[:20])
|
| 24 |
+
if __name__ == "__main__":
|
| 25 |
+
main()
|
tools/make_sfda_manifest.py
ADDED
|
@@ -0,0 +1,39 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python
|
| 2 |
+
"""Create an SFDA-safe manifest by dropping labels from specified target splits.
|
| 3 |
+
|
| 4 |
+
This prevents accidental target-label leakage during adaptation. Validation/test labels are
|
| 5 |
+
kept by default so metrics can be computed.
|
| 6 |
+
"""
|
| 7 |
+
from __future__ import annotations
|
| 8 |
+
import argparse
|
| 9 |
+
import json
|
| 10 |
+
from pathlib import Path
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
def main():
|
| 14 |
+
ap = argparse.ArgumentParser()
|
| 15 |
+
ap.add_argument("--input", required=True, help="Input manifest JSON")
|
| 16 |
+
ap.add_argument("--output", required=True, help="Output manifest JSON")
|
| 17 |
+
ap.add_argument("--drop-splits", nargs="+", default=["target_train"], help="Splits whose labels should be removed")
|
| 18 |
+
args = ap.parse_args()
|
| 19 |
+
with open(args.input, "r") as f:
|
| 20 |
+
man = json.load(f)
|
| 21 |
+
for split in args.drop_splits:
|
| 22 |
+
for item in man.get(split, []):
|
| 23 |
+
item.pop("label", None)
|
| 24 |
+
man.setdefault("notes", {})["sfda_safe"] = {
|
| 25 |
+
"dropped_label_splits": args.drop_splits,
|
| 26 |
+
"warning": "Target adaptation labels removed to avoid SFDA leakage."
|
| 27 |
+
}
|
| 28 |
+
out = Path(args.output)
|
| 29 |
+
out.parent.mkdir(parents=True, exist_ok=True)
|
| 30 |
+
with open(out, "w") as f:
|
| 31 |
+
json.dump(man, f, indent=2)
|
| 32 |
+
print(f"Wrote {out}")
|
| 33 |
+
for k, v in man.items():
|
| 34 |
+
if isinstance(v, list):
|
| 35 |
+
print(f"{k}: {len(v)} items, {sum('label' in x for x in v)} labels")
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
if __name__ == "__main__":
|
| 39 |
+
main()
|
tools/remap_labels_in_manifest.py
ADDED
|
@@ -0,0 +1,89 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python
|
| 2 |
+
"""Remap labels referenced by a SACFlow manifest and write a new manifest.
|
| 3 |
+
|
| 4 |
+
Useful for AMOS CT->MRI shared-label protocols. Example:
|
| 5 |
+
labels 14 and 15 -> background 0, num_classes=14.
|
| 6 |
+
"""
|
| 7 |
+
from __future__ import annotations
|
| 8 |
+
import argparse
|
| 9 |
+
import json
|
| 10 |
+
from pathlib import Path
|
| 11 |
+
import nibabel as nib
|
| 12 |
+
import numpy as np
|
| 13 |
+
from tqdm import tqdm
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
def parse_map(map_items):
|
| 17 |
+
out = {}
|
| 18 |
+
for item in map_items or []:
|
| 19 |
+
if ":" not in item:
|
| 20 |
+
raise ValueError(f"Bad --map entry {item}; expected old:new")
|
| 21 |
+
a, b = item.split(":", 1)
|
| 22 |
+
out[int(a)] = int(b)
|
| 23 |
+
return out
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
def main():
|
| 27 |
+
ap = argparse.ArgumentParser()
|
| 28 |
+
ap.add_argument("--input", required=True)
|
| 29 |
+
ap.add_argument("--output", required=True)
|
| 30 |
+
ap.add_argument("--label-root", required=True, help="Directory for remapped labels")
|
| 31 |
+
ap.add_argument("--map", nargs="*", default=[], help="Label mappings old:new, e.g. 14:0 15:0")
|
| 32 |
+
ap.add_argument("--num-classes", type=int, default=None)
|
| 33 |
+
ap.add_argument("--protocol-name", default="remapped")
|
| 34 |
+
ap.add_argument("--overwrite", action="store_true")
|
| 35 |
+
args = ap.parse_args()
|
| 36 |
+
mapping = parse_map(args.map)
|
| 37 |
+
with open(args.input, "r") as f:
|
| 38 |
+
man = json.load(f)
|
| 39 |
+
label_root = Path(args.label_root)
|
| 40 |
+
label_root.mkdir(parents=True, exist_ok=True)
|
| 41 |
+
seen = {}
|
| 42 |
+
|
| 43 |
+
def remap_one(path_str: str) -> str:
|
| 44 |
+
path = Path(path_str)
|
| 45 |
+
if path_str in seen:
|
| 46 |
+
return seen[path_str]
|
| 47 |
+
out = label_root / f"{path.parent.name}_{path.name}"
|
| 48 |
+
if out.exists() and not args.overwrite:
|
| 49 |
+
seen[path_str] = str(out.resolve())
|
| 50 |
+
return seen[path_str]
|
| 51 |
+
img = nib.load(str(path))
|
| 52 |
+
arr = img.get_fdata().astype(np.int16)
|
| 53 |
+
for old, new in mapping.items():
|
| 54 |
+
arr[arr == old] = new
|
| 55 |
+
if args.num_classes is not None:
|
| 56 |
+
arr[arr >= args.num_classes] = 0
|
| 57 |
+
arr[arr < 0] = 0
|
| 58 |
+
out_img = nib.Nifti1Image(arr.astype(np.uint8), img.affine, img.header)
|
| 59 |
+
out_img.set_data_dtype(np.uint8)
|
| 60 |
+
nib.save(out_img, str(out))
|
| 61 |
+
seen[path_str] = str(out.resolve())
|
| 62 |
+
return seen[path_str]
|
| 63 |
+
|
| 64 |
+
count = 0
|
| 65 |
+
for split, items in man.items():
|
| 66 |
+
if not isinstance(items, list):
|
| 67 |
+
continue
|
| 68 |
+
for item in tqdm(items, desc=f"remap {split}"):
|
| 69 |
+
if item.get("label"):
|
| 70 |
+
item["label"] = remap_one(item["label"])
|
| 71 |
+
count += 1
|
| 72 |
+
man["label_protocol"] = {
|
| 73 |
+
"name": args.protocol_name,
|
| 74 |
+
"mapping": {str(k): v for k, v in mapping.items()},
|
| 75 |
+
"num_classes": args.num_classes,
|
| 76 |
+
"description": "Labels remapped using tools/remap_labels_in_manifest.py",
|
| 77 |
+
}
|
| 78 |
+
out_path = Path(args.output)
|
| 79 |
+
out_path.parent.mkdir(parents=True, exist_ok=True)
|
| 80 |
+
with open(out_path, "w") as f:
|
| 81 |
+
json.dump(man, f, indent=2)
|
| 82 |
+
print(f"Wrote {out_path}; remapped {count} label references into {label_root}")
|
| 83 |
+
for split, items in man.items():
|
| 84 |
+
if isinstance(items, list):
|
| 85 |
+
print(f"{split}: {len(items)} items, {sum('label' in x for x in items)} labels")
|
| 86 |
+
|
| 87 |
+
|
| 88 |
+
if __name__ == "__main__":
|
| 89 |
+
main()
|
tools/train.py
ADDED
|
@@ -0,0 +1,46 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python
|
| 2 |
+
from __future__ import annotations
|
| 3 |
+
import argparse
|
| 4 |
+
from pathlib import Path
|
| 5 |
+
import yaml
|
| 6 |
+
from sacflow.utils.config import load_yaml, save_yaml, set_by_path
|
| 7 |
+
from sacflow.utils.misc import seed_everything, ensure_dir
|
| 8 |
+
from sacflow.utils.distributed import init_distributed, cleanup, is_main_process, barrier
|
| 9 |
+
from sacflow.utils.wandb_utils import init_wandb, wandb_finish
|
| 10 |
+
from sacflow.engine.train_loop import run_training
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
def main():
|
| 14 |
+
ap = argparse.ArgumentParser()
|
| 15 |
+
ap.add_argument("--config", required=True)
|
| 16 |
+
ap.add_argument("--name", default=None)
|
| 17 |
+
ap.add_argument("--resume", nargs="?", const="auto", default=None, help="Resume checkpoint path, or use --resume without a value for auto-detection.")
|
| 18 |
+
ap.add_argument("--opts", nargs="*", default=[], help="Override config values: key=value, e.g. train.epochs=220 optim.lr=1e-4")
|
| 19 |
+
args = ap.parse_args()
|
| 20 |
+
cfg = load_yaml(args.config)
|
| 21 |
+
for opt in args.opts:
|
| 22 |
+
if "=" not in opt:
|
| 23 |
+
raise ValueError(f"Invalid --opts entry {opt!r}; expected key=value")
|
| 24 |
+
key, value = opt.split("=", 1)
|
| 25 |
+
try:
|
| 26 |
+
parsed = yaml.safe_load(value)
|
| 27 |
+
except Exception:
|
| 28 |
+
parsed = value
|
| 29 |
+
set_by_path(cfg, key, parsed)
|
| 30 |
+
if args.resume is not None:
|
| 31 |
+
cfg.setdefault("train", {})["resume_checkpoint"] = args.resume
|
| 32 |
+
seed_everything(int(cfg.get("seed", 1337)))
|
| 33 |
+
device = init_distributed(cfg.get("distributed", {}).get("backend", "nccl"))
|
| 34 |
+
out_dir = ensure_dir(cfg["output_dir"])
|
| 35 |
+
if is_main_process():
|
| 36 |
+
save_yaml(cfg, out_dir / "resolved_config.yaml")
|
| 37 |
+
run = init_wandb(cfg, run_name=args.name or Path(cfg["output_dir"]).name)
|
| 38 |
+
try:
|
| 39 |
+
run_training(cfg, device, wandb_run=run)
|
| 40 |
+
finally:
|
| 41 |
+
barrier()
|
| 42 |
+
wandb_finish(run)
|
| 43 |
+
cleanup()
|
| 44 |
+
|
| 45 |
+
if __name__ == "__main__":
|
| 46 |
+
main()
|