sathiiii commited on
Commit
93ffd19
·
verified ·
1 Parent(s): e7dc055

Add tools

Browse files
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()