#!/usr/bin/env python 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: # fuzzy: strip common prefixes 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()