File size: 3,231 Bytes
93ffd19 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 | #!/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()
|