CoolFace
Apppublic

codemo/fish-speech-1

sourceHugging Facecc-by-nc-sa-4.0updated 2y agoView on Hugging Face
0likes
create_train_split.py84 linesDownload Raw Back to vqgan
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