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