# -*- coding: utf-8 -*- # Copyright (c) Alibaba, Inc. and its affiliates. import argparse import datetime import importlib import os import random import sys import gradio as gr import scepter from scepter.modules.utils.config import Config from scepter.modules.utils.file_system import FS from scepter.modules.utils.logger import get_logger, init_logger from scepter.modules.utils.import_utils import get_dirname if os.path.exists('__init__.py'): package_name = 'scepter_ext' spec = importlib.util.spec_from_file_location(package_name, '__init__.py') package = importlib.util.module_from_spec(spec) sys.modules[package_name] = package spec.loader.exec_module(package) def prepare(config): if 'FILE_SYSTEM' in config: for fs_info in config['FILE_SYSTEM']: FS.init_fs_client(fs_info) if 'LOG_FILE' in config: logger = get_logger() tid = '{0:%Y%m%d%H%M%S%f}'.format(datetime.datetime.now()) + ''.join( [str(random.randint(1, 10)) for i in range(3)]) init_logger(logger, log_file=config['LOG_FILE'].format(tid)) class TabManager(): def __init__(self): pass if __name__ == '__main__': parser = argparse.ArgumentParser() parser.add_argument('--cfg', dest='config', type=str, default=os.path.join( os.path.dirname(get_dirname()), 'scepter/methods/studio/scepter_ui.yaml')) parser.add_argument('--debug', dest='debug', action='store_true', help='Switch debug mode.') parser.add_argument( '--host', dest='host', default=None, help='The host of Gradio, default is set in ui config.') parser.add_argument('--port', dest='port', default=None, help='The port of Gradio.') parser.add_argument('--language', dest='language', choices=['en', 'zh'], default='en', help='Now we only support english(en) and chinese(zh)') parser.add_argument('--tab', dest='tab', choices=['all', 'chatbot'], default='all', help='The tabs will be launched, ' 'set [all] to use all tools and set [chatbot] to use chatbot only.') parser.add_argument('--model', dest='model', default=None, help='Example: sd*|flux') args = parser.parse_args() if not os.path.exists(args.config): print( f"{args.config} doesn't exist, find this file in {os.path.dirname(get_dirname())}" ) args.config = os.path.join(os.path.dirname(get_dirname()), args.config) assert os.path.exists(args.config) config = Config(load=True, cfg_file=args.config) prepare(config) tab_manager = TabManager() interfaces = [] for info in config['INTERFACE']: name = info.get('NAME_EN', '') if args.language == 'en' else info['NAME'] ifid = info['IFID'] if not FS.exists(info['CONFIG']): info['CONFIG'] = os.path.join(os.path.dirname(scepter.dirname), info['CONFIG']) if not FS.exists(info['CONFIG']): raise f"{info['CONFIG']} doesn't exist." interface = None model = args.model if args.model else None if ifid == 'home' and args.tab in ["all", ifid]: from scepter.studio.home.home import HomeUI interface = HomeUI(info['CONFIG'], is_debug=args.debug, language=args.language, root_work_dir=config.WORK_DIR) print('init home page success!') if ifid == 'preprocess' and args.tab in ["all", ifid]: from scepter.studio.preprocess.preprocess import PreprocessUI interface = PreprocessUI(info['CONFIG'], is_debug=args.debug, language=args.language, root_work_dir=config.WORK_DIR) print('init preprocess success!') if ifid == 'self_train' and args.tab in ["all", ifid]: from scepter.studio.self_train.self_train import SelfTrainUI interface = SelfTrainUI(info['CONFIG'], is_debug=args.debug, language=args.language, root_work_dir=config.WORK_DIR) print('init self-train success!') if ifid == 'tuner_manager' and args.tab in ["all", ifid]: from scepter.studio.tuner_manager.tuner_manager import TunerManagerUI interface = TunerManagerUI(info['CONFIG'], is_debug=args.debug, language=args.language, root_work_dir=config.WORK_DIR) print('init tuner-manager success!') if ifid == 'inference' and args.tab in ["all", ifid]: from scepter.studio.inference.inference import InferenceUI interface = InferenceUI(info['CONFIG'], is_debug=args.debug, language=args.language, model=model, root_work_dir=config.WORK_DIR) print('init inference success!') if ifid == 'chatbot' and args.tab in ["all", ifid]: from scepter.studio.chatbot.chatbot import ChatBotUI interface = ChatBotUI(info['CONFIG'], root_work_dir=config.WORK_DIR) print('init chatbot success!') if ifid == '': pass # TODO: Add New Features if interface: interfaces.append((interface, name, ifid)) setattr(tab_manager, ifid, interface) css = """ .upload_zone { height: 100px; } """ with gr.Blocks(css=css) as demo: if 'BANNER' in config: gr.HTML(config.BANNER) else: gr.Markdown( f"

{config.get('TITLE', 'scepter studio')}

" ) setattr(tab_manager, 'user_name', gr.Text(value='admin', visible=False, show_label=False)) with gr.Tabs(elem_id='tabs') as tabs: setattr(tab_manager, 'tabs', tabs) for interface, label, ifid in interfaces: with gr.TabItem(label, id=ifid, elem_id=f'tab_{ifid}'): interface.create_ui() for interface, label, ifid in interfaces: interface.set_callbacks(tab_manager) auth_info = {} if config.have('AUTH_INFO'): for auth_user in config.AUTH_INFO: auth_info[auth_user.USER] = auth_user.PASSWD def check_auth(user_name, password): if user_name in auth_info: return auth_info[user_name] == password else: return False def init_value(req: gr.Request): print(req.username, 'have login') return gr.Text(value=req.username, visible=False) if len(auth_info) > 0: demo.load(init_value, outputs=[tab_manager.user_name]) allowed_paths = [config['WORK_DIR']] allowed_paths.extend(list(set([fs_cfg['TEMP_DIR'] for fs_cfg in config['FILE_SYSTEM']])) if 'FILE_SYSTEM' in config else []) demo.queue(status_update_rate=1).launch( server_name=args.host if args.host else config['HOST'], server_port=int(args.port) if args.port else config['PORT'], root_path=config['ROOT'], show_error=True, debug=True, auth=check_auth if len(auth_info) > 0 else None, allowed_paths=allowed_paths)