File size: 3,323 Bytes
bed2cee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
100
101
import argparse
import os
import random
import sys


AUDIO_EXTS = {".wav", ".flac", ".mp3"}


def main():
    parser = argparse.ArgumentParser(
        description="Split MUSAN audio files into train/test/dev (8:1:1) per subdirectory."
    )
    parser.add_argument("root_dir", help="MUSAN root directory")
    parser.add_argument("output_dir", help="Output directory for {tr,tt,cv}.wavlist")
    parser.add_argument(
        "--seed", type=int, default=42, help="Random seed for reproducible splits"
    )
    args = parser.parse_args()

    root_dir = os.path.abspath(args.root_dir)
    output_dir = os.path.abspath(args.output_dir)
    os.makedirs(output_dir, exist_ok=True)

    # Collect audio files grouped by their immediate parent directory
    # key: (parent_dir_abs_path), value: list of absolute audio file paths
    dir_map = {}
    for dirpath, _, filenames in os.walk(root_dir):
        audio_files = [
            os.path.join(dirpath, f)
            for f in filenames
            if os.path.splitext(f)[1].lower() in AUDIO_EXTS
        ]
        if audio_files:
            # Group by the immediate parent directory so we split per-subfolder
            dir_map[dirpath] = audio_files

    if not dir_map:
        print("No audio files found in the given root directory.", file=sys.stderr)
        sys.exit(1)

    random.seed(args.seed)

    train_files = []
    test_files = []
    dev_files = []

    for d, files in sorted(dir_map.items()):
        n = len(files)
        files_sorted = sorted(files)
        random.shuffle(files_sorted)

        if n < 10:
            rel = os.path.relpath(d, root_dir)
            print(
                f"WARNING: '{rel}' has only {n} audio file(s) (< 10). "
                "Prioritizing 1 for test, 1 for dev, rest for train.",
                file=sys.stderr,
            )
            # Priority: 1 test, 1 dev, rest train
            tst = files_sorted[:1] if n >= 1 else []
            dev = files_sorted[1:2] if n >= 2 else []
            trn = files_sorted[2:]
            test_files.extend(tst)
            dev_files.extend(dev)
            train_files.extend(trn)
        else:
            n_test = max(1, round(n * 0.1))
            n_dev = max(1, round(n * 0.1))
            # Adjust to ensure n_test + n_dev < n, prioritize test first if overflow
            if n_test + n_dev >= n:
                n_dev = max(0, n - n_test - 1)
                if n_dev == 0:
                    n_test = n - 1
            n_train = n - n_test - n_dev

            test_files.extend(files_sorted[:n_test])
            dev_files.extend(files_sorted[n_test:n_test + n_dev])
            train_files.extend(files_sorted[n_test + n_dev:])

    # Sort for deterministic output
    train_files.sort()
    test_files.sort()
    dev_files.sort()

    with open(os.path.join(output_dir, "tr.wavlist"), "w") as f:
        f.write("\n".join(train_files) + "\n")

    with open(os.path.join(output_dir, "tt.wavlist"), "w") as f:
        f.write("\n".join(test_files) + "\n")

    with open(os.path.join(output_dir, "cv.wavlist"), "w") as f:
        f.write("\n".join(dev_files) + "\n")

    print(f"Train: {len(train_files)}, Test: {len(test_files)}, Dev: {len(dev_files)}")
    print(f"Files written to {output_dir}/{{tr,tt,cv}}.wavlist")


if __name__ == "__main__":
    main()