czyhust's picture
Add files using upload-large-folder tool
bed2cee verified
Raw
History Blame Contribute Delete
3.32 kB
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()