| |
| from __future__ import annotations |
| import argparse |
| import json |
| from pathlib import Path |
| import random |
|
|
| NIFTI_EXTS = [".nii.gz", ".nii", ".mha", ".mhd"] |
|
|
| def list_images(folder): |
| folder = Path(folder) |
| files = [] |
| for ext in NIFTI_EXTS: |
| files.extend(folder.glob(f"*{ext}")) |
| return sorted(files) |
|
|
| def stem_nii(p: Path): |
| name = p.name |
| for ext in NIFTI_EXTS: |
| if name.endswith(ext): |
| return name[:-len(ext)] |
| return p.stem |
|
|
| def match_labels(images, labels_dir): |
| if labels_dir is None: |
| return {stem_nii(p): None for p in images} |
| labels = list_images(labels_dir) |
| lab_map = {stem_nii(p): p for p in labels} |
| out = {} |
| for img in images: |
| sid = stem_nii(img) |
| if sid in lab_map: |
| out[sid] = lab_map[sid] |
| else: |
| |
| cand = None |
| for k, v in lab_map.items(): |
| if sid in k or k in sid: |
| cand = v; break |
| out[sid] = cand |
| return out |
|
|
| def make_items(images_dir, labels_dir=None, domain=""): |
| imgs = list_images(images_dir) |
| labels = match_labels(imgs, labels_dir) |
| items = [] |
| for p in imgs: |
| sid = stem_nii(p) |
| item = {"id": f"{domain}_{sid}" if domain else sid, "image": str(p)} |
| if labels.get(sid): |
| item["label"] = str(labels[sid]) |
| items.append(item) |
| return items |
|
|
| def split_items(items, val_fraction, test_fraction, seed): |
| items = list(items) |
| random.Random(seed).shuffle(items) |
| n = len(items) |
| n_test = int(round(n * test_fraction)) |
| n_val = int(round(n * val_fraction)) |
| test = items[:n_test] |
| val = items[n_test:n_test+n_val] |
| train = items[n_test+n_val:] |
| return train, val, test |
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--name", required=True) |
| ap.add_argument("--source-images", required=True) |
| ap.add_argument("--source-labels", default=None) |
| ap.add_argument("--target-images", required=True) |
| ap.add_argument("--target-labels", default=None) |
| ap.add_argument("--output", required=True) |
| ap.add_argument("--val-fraction", type=float, default=0.2) |
| ap.add_argument("--test-fraction", type=float, default=0.2) |
| ap.add_argument("--seed", type=int, default=1337) |
| args = ap.parse_args() |
| src = make_items(args.source_images, args.source_labels, "source") |
| tgt = make_items(args.target_images, args.target_labels, "target") |
| src_train, src_val, src_test = split_items(src, args.val_fraction, args.test_fraction, args.seed) |
| tgt_train, tgt_val, tgt_test = split_items(tgt, args.val_fraction, args.test_fraction, args.seed) |
| man = { |
| "name": args.name, |
| "source_train": src_train, |
| "source_val": src_val, |
| "source_test": src_test, |
| "target_train": tgt_train, |
| "target_val": tgt_val, |
| "target_test": tgt_test, |
| } |
| out = Path(args.output) |
| out.parent.mkdir(parents=True, exist_ok=True) |
| with open(out, "w") as f: |
| json.dump(man, f, indent=2) |
| print(f"Wrote {out}") |
| print({k: len(v) for k, v in man.items() if isinstance(v, list)}) |
| if __name__ == "__main__": |
| main() |
|
|