update v0.0.4
This commit is contained in:
@@ -1,12 +1,19 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import argparse
|
||||
import importlib
|
||||
import os
|
||||
import sys
|
||||
|
||||
#
|
||||
from scepter.modules.utils.registry import REGISTRY_LIST
|
||||
|
||||
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)
|
||||
sys.path.insert(0, os.path.abspath(os.curdir))
|
||||
|
||||
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import argparse
|
||||
import importlib
|
||||
import os
|
||||
import sys
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -16,6 +18,13 @@ from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import FS
|
||||
from scepter.modules.utils.logger import get_logger
|
||||
|
||||
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 run_task(cfg):
|
||||
std_logger = get_logger(name='scepter')
|
||||
|
||||
@@ -1,12 +1,22 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import argparse
|
||||
import importlib
|
||||
import os
|
||||
import sys
|
||||
|
||||
from scepter.modules.solver.registry import SOLVERS
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.logger import get_logger
|
||||
|
||||
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 run_task(cfg):
|
||||
std_logger = get_logger(name='scepter')
|
||||
@@ -18,19 +28,21 @@ def run_task(cfg):
|
||||
|
||||
def update_config(cfg):
|
||||
if hasattr(cfg.args, 'learning_rate') and cfg.args.learning_rate:
|
||||
print(
|
||||
f'learning_rate change from {cfg.SOLVER.OPTIMIZER.LEARNING_RATE} to {cfg.args.learning_rate}'
|
||||
)
|
||||
if cfg.SOLVER.OPTIMIZER.get('LEARNING_RATE', None) is not None:
|
||||
print(
|
||||
f'learning_rate change from {cfg.SOLVER.OPTIMIZER.LEARNING_RATE} to {cfg.args.learning_rate}'
|
||||
)
|
||||
cfg.SOLVER.OPTIMIZER.LEARNING_RATE = float(cfg.args.learning_rate)
|
||||
if hasattr(cfg.args, 'max_steps') and cfg.args.max_steps:
|
||||
print(
|
||||
f'max_steps change from {cfg.SOLVER.MAX_STEPS} to {cfg.args.max_steps}'
|
||||
)
|
||||
if cfg.SOLVER.get('MAX_STEPS', None) is not None:
|
||||
print(
|
||||
f'max_steps change from {cfg.SOLVER.MAX_STEPS} to {cfg.args.max_steps}'
|
||||
)
|
||||
cfg.SOLVER.MAX_STEPS = int(cfg.args.max_steps)
|
||||
return cfg
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
def run():
|
||||
parser = argparse.ArgumentParser(description='Argparser for Scepter:\n')
|
||||
parser.add_argument('--learning_rate',
|
||||
dest='learning_rate',
|
||||
@@ -44,3 +56,7 @@ if __name__ == '__main__':
|
||||
cfg = Config(load=True, parser_ins=parser)
|
||||
cfg = update_config(cfg)
|
||||
we.init_env(cfg, logger=None, fn=run_task)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
run()
|
||||
|
||||
+48
-2
@@ -2,8 +2,10 @@
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import argparse
|
||||
import datetime
|
||||
import importlib
|
||||
import os
|
||||
import random
|
||||
import sys
|
||||
|
||||
import gradio as gr
|
||||
|
||||
@@ -12,6 +14,13 @@ 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
|
||||
|
||||
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:
|
||||
@@ -57,6 +66,13 @@ if __name__ == '__main__':
|
||||
default='en',
|
||||
help='Now we only support english(en) and chinese(zh)')
|
||||
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(scepter.dirname)}"
|
||||
)
|
||||
args.config = os.path.join(os.path.dirname(scepter.dirname),
|
||||
args.config)
|
||||
assert os.path.exists(args.config)
|
||||
config = Config(load=True, cfg_file=args.config)
|
||||
prepare(config)
|
||||
|
||||
@@ -74,24 +90,34 @@ if __name__ == '__main__':
|
||||
interface = None
|
||||
if ifid == 'home':
|
||||
from scepter.studio.home.home import HomeUI
|
||||
|
||||
interface = HomeUI(info['CONFIG'],
|
||||
is_debug=args.debug,
|
||||
language=args.language,
|
||||
root_work_dir=config.WORK_DIR)
|
||||
if ifid == 'preprocess':
|
||||
from scepter.studio.preprocess.preprocess import PreprocessUI
|
||||
|
||||
interface = PreprocessUI(info['CONFIG'],
|
||||
is_debug=args.debug,
|
||||
language=args.language,
|
||||
root_work_dir=config.WORK_DIR)
|
||||
if ifid == 'self_train':
|
||||
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)
|
||||
if ifid == 'tuner_manager':
|
||||
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)
|
||||
if ifid == 'inference':
|
||||
from scepter.studio.inference.inference import InferenceUI
|
||||
|
||||
interface = InferenceUI(info['CONFIG'],
|
||||
is_debug=args.debug,
|
||||
language=args.language,
|
||||
@@ -109,6 +135,8 @@ if __name__ == '__main__':
|
||||
gr.Markdown(
|
||||
f"<h2><center>{config.get('TITLE', 'scepter studio')}</center></h2>"
|
||||
)
|
||||
setattr(tab_manager, 'user_name',
|
||||
gr.Text(value='', visible=False, show_label=False))
|
||||
with gr.Tabs(elem_id='tabs') as tabs:
|
||||
setattr(tab_manager, 'tabs', tabs)
|
||||
for interface, label, ifid in interfaces:
|
||||
@@ -116,11 +144,29 @@ if __name__ == '__main__':
|
||||
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])
|
||||
|
||||
demo.queue(status_update_rate=1).launch(
|
||||
server_name=args.host if args.host else config['HOST'],
|
||||
server_port=args.port if args.port else config['PORT'],
|
||||
server_port=int(args.port) if args.port else config['PORT'],
|
||||
root_path=config['ROOT'],
|
||||
show_error=True,
|
||||
debug=True,
|
||||
enable_queue=True)
|
||||
enable_queue=True,
|
||||
auth=check_auth if len(auth_info) > 0 else None)
|
||||
|
||||
Reference in New Issue
Block a user