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