CoolFace
Apppublic

lenML/ChatTTS-Forge

sourceHugging Faceagpl-3.0updated 2y agoView on Hugging Face
301likes
models_setup.py75 linesDownload Raw Back to modules
1import argparse2import logging3 4from modules import generate_audio5from modules.devices import devices6from modules.Enhancer.ResembleEnhance import load_enhancer7from modules.models import load_chat_tts8from modules.utils import env9 10 11def setup_model_args(parser: argparse.ArgumentParser):12    parser.add_argument("--compile", action="store_true", help="Enable model compile")13    parser.add_argument(14        "--no_half",15        action="store_true",16        help="Disalbe half precision for model inference",17    )18    parser.add_argument(19        "--off_tqdm",20        action="store_true",21        help="Disable tqdm progress bar",22    )23    parser.add_argument(24        "--device_id",25        type=str,26        help="Select the default CUDA device to use (export CUDA_VISIBLE_DEVICES=0,1,etc might be needed before)",27        default=None,28    )29    parser.add_argument(30        "--use_cpu",31        nargs="+",32        help="use CPU as torch device for specified modules",33        default=[],34        type=str.lower,35        choices=["all", "chattts", "enhancer", "trainer"],36    )37    parser.add_argument(38        "--lru_size",39        type=int,40        default=64,41        help="Set the size of the request cache pool, set it to 0 will disable lru_cache",42    )43    parser.add_argument(44        "--debug_generate",45        action="store_true",46        help="Enable debug mode for audio generation",47    )48    parser.add_argument(49        "--preload_models",50        action="store_true",51        help="Preload all models at startup",52    )53 54 55def process_model_args(args: argparse.Namespace):56    lru_size = env.get_and_update_env(args, "lru_size", 64, int)57    compile = env.get_and_update_env(args, "compile", False, bool)58    device_id = env.get_and_update_env(args, "device_id", None, str)59    use_cpu = env.get_and_update_env(args, "use_cpu", [], list)60    no_half = env.get_and_update_env(args, "no_half", False, bool)61    off_tqdm = env.get_and_update_env(args, "off_tqdm", False, bool)62    debug_generate = env.get_and_update_env(args, "debug_generate", False, bool)63    preload_models = env.get_and_update_env(args, "preload_models", False, bool)64 65    generate_audio.setup_lru_cache()66    devices.reset_device()67    devices.first_time_calculation()68 69    if debug_generate:70        generate_audio.logger.setLevel(logging.DEBUG)71 72    if preload_models:73        load_chat_tts()74        load_enhancer()75