lenML/ChatTTS-Forge
301
1import logging2import os3import sys4 5from modules.ffmpeg_env import setup_ffmpeg_path6 7try:8 setup_ffmpeg_path()9 # NOTE: 因为 logger 都是在模块中初始化,所以这个 config 必须在最前面10 logging.basicConfig(11 level=os.getenv("LOG_LEVEL", "INFO"),12 format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",13 )14except BaseException:15 pass16 17import argparse18 19from modules import config20from modules.api.api_setup import process_api_args, setup_api_args21from modules.api.app_config import app_description, app_title, app_version22from modules.gradio_dcls_fix import dcls_patch23from modules.models_setup import process_model_args, setup_model_args24from modules.utils.env import get_and_update_env25from modules.utils.ignore_warn import ignore_useless_warnings26from modules.utils.torch_opt import configure_torch_optimizations27from modules.webui import webui_config28from modules.webui.app import create_interface, webui_init29 30import subprocess31 32subprocess.run(33 "pip install flash-attn --no-build-isolation",34 env={"FLASH_ATTENTION_SKIP_CUDA_BUILD": "TRUE"},35 shell=True,36)37 38dcls_patch()39ignore_useless_warnings()40 41 42def setup_webui_args(parser: argparse.ArgumentParser):43 parser.add_argument("--server_name", type=str, help="server name")44 parser.add_argument("--server_port", type=int, help="server port")45 parser.add_argument(46 "--share", action="store_true", help="share the gradio interface"47 )48 parser.add_argument("--debug", action="store_true", help="enable debug mode")49 parser.add_argument("--auth", type=str, help="username:password for authentication")50 parser.add_argument(51 "--tts_max_len",52 type=int,53 help="Max length of text for TTS",54 )55 parser.add_argument(56 "--ssml_max_len",57 type=int,58 help="Max length of text for SSML",59 )60 parser.add_argument(61 "--max_batch_size",62 type=int,63 help="Max batch size for TTS",64 )65 # webui_Experimental66 parser.add_argument(67 "--webui_experimental",68 action="store_true",69 help="Enable webui_experimental features",70 )71 parser.add_argument(72 "--language",73 type=str,74 help="Set the default language for the webui",75 )76 parser.add_argument(77 "--api",78 action="store_true",79 help="use api=True to launch the API together with the webui (run launch.py for only API server)",80 )81 82 83def process_webui_args(args):84 server_name = get_and_update_env(args, "server_name", "0.0.0.0", str)85 server_port = get_and_update_env(args, "server_port", 7860, int)86 share = get_and_update_env(args, "share", False, bool)87 debug = get_and_update_env(args, "debug", False, bool)88 auth = get_and_update_env(args, "auth", None, str)89 language = get_and_update_env(args, "language", "zh-CN", str)90 api = get_and_update_env(args, "api", False, bool)91 92 webui_config.experimental = get_and_update_env(93 args, "webui_experimental", False, bool94 )95 webui_config.tts_max = get_and_update_env(args, "tts_max_len", 1000, int)96 webui_config.ssml_max = get_and_update_env(args, "ssml_max_len", 5000, int)97 webui_config.max_batch_size = get_and_update_env(args, "max_batch_size", 8, int)98 99 webui_config.experimental = get_and_update_env(100 args, "webui_experimental", False, bool101 )102 webui_config.tts_max = get_and_update_env(args, "tts_max_len", 1000, int)103 webui_config.ssml_max = get_and_update_env(args, "ssml_max_len", 5000, int)104 webui_config.max_batch_size = get_and_update_env(args, "max_batch_size", 8, int)105 106 configure_torch_optimizations()107 webui_init()108 demo = create_interface()109 110 if auth:111 auth = tuple(auth.split(":"))112 113 app, local_url, share_url = demo.queue().launch(114 server_name=server_name,115 server_port=server_port,116 share=share,117 debug=debug,118 auth=auth,119 show_api=False,120 prevent_thread_lock=True,121 inbrowser=sys.platform == "win32",122 app_kwargs={123 "title": app_title,124 "description": app_description,125 "version": app_version,126 "redoc_url": (127 None128 if api is False129 else None if config.runtime_env_vars.no_docs else "/redoc"130 ),131 "docs_url": (132 None133 if api is False134 else None if config.runtime_env_vars.no_docs else "/docs"135 ),136 },137 )138 # gradio uses a very open CORS policy via app.user_middleware, which makes it possible for139 # an attacker to trick the user into opening a malicious HTML page, which makes a request to the140 # running web ui and do whatever the attacker wants, including installing an extension and141 # running its code. We disable this here. Suggested by RyotaK.142 app.user_middleware = [143 x for x in app.user_middleware if x.cls.__name__ != "CustomCORSMiddleware"144 ]145 146 if api:147 process_api_args(args, app)148 149 demo.block_thread()150 151 152if __name__ == "__main__":153 import dotenv154 155 dotenv.load_dotenv(156 dotenv_path=os.getenv("ENV_FILE", ".env.webui"),157 )158 159 parser = argparse.ArgumentParser(description="Gradio App")160 161 setup_webui_args(parser)162 setup_model_args(parser)163 setup_api_args(parser)164 165 args = parser.parse_args()166 167 process_model_args(args)168 process_webui_args(args)169 