codemo/fish-speech-1
0
1import math2from pathlib import Path3from random import Random4 5import click6from loguru import logger7from pydub import AudioSegment8from tqdm import tqdm9 10from tools.file import AUDIO_EXTENSIONS, list_files, load_filelist11 12 13@click.command()14@click.argument("root", type=click.Path(exists=True, path_type=Path))15@click.option("--val-ratio", type=float, default=None)16@click.option("--val-count", type=int, default=None)17@click.option("--filelist", default=None, type=Path)18@click.option("--min-duration", default=None, type=float)19@click.option("--max-duration", default=None, type=float)20def main(root, val_ratio, val_count, filelist, min_duration, max_duration):21 if filelist:22 files = [i[0] for i in load_filelist(filelist)]23 else:24 files = list_files(root, AUDIO_EXTENSIONS, recursive=True, sort=True)25 26 if min_duration is None and max_duration is None:27 filtered_files = list(map(str, [file.relative_to(root) for file in files]))28 else:29 filtered_files = []30 for file in tqdm(files):31 try:32 audio = AudioSegment.from_file(str(file))33 duration = len(audio) / 1000.034 35 if min_duration is not None and duration < min_duration:36 logger.info(37 f"Skipping {file} due to duration {duration:.2f} < {min_duration:.2f}"38 )39 continue40 41 if max_duration is not None and duration > max_duration:42 logger.info(43 f"Skipping {file} due to duration {duration:.2f} > {max_duration:.2f}"44 )45 continue46 47 filtered_files.append(str(file.relative_to(root)))48 except Exception as e:49 logger.info(f"Error processing {file}: {e}")50 51 logger.info(52 f"Found {len(files)} files, remaining {len(filtered_files)} files after filtering"53 )54 55 Random(42).shuffle(filtered_files)56 57 if val_count is None and val_ratio is None:58 logger.info("Validation ratio and count not specified, using min(20%, 100)")59 val_size = min(100, math.ceil(len(filtered_files) * 0.2))60 elif val_count is not None and val_ratio is not None:61 logger.error("Cannot specify both val_count and val_ratio")62 return63 elif val_count is not None:64 if val_count < 1 or val_count > len(filtered_files):65 logger.error("val_count must be between 1 and number of files")66 return67 val_size = val_count68 else:69 val_size = math.ceil(len(filtered_files) * val_ratio)70 71 logger.info(f"Using {val_size} files for validation")72 73 with open(root / "vq_train_filelist.txt", "w", encoding="utf-8") as f:74 f.write("\n".join(filtered_files[val_size:]))75 76 with open(root / "vq_val_filelist.txt", "w", encoding="utf-8") as f:77 f.write("\n".join(filtered_files[:val_size]))78 79 logger.info("Done")80 81 82if __name__ == "__main__":83 main()84 