update v0.0.4

This commit is contained in:
LouieStark
2024-03-31 13:08:41 +08:00
parent 35aada8ce8
commit bf53829530
106 changed files with 6927 additions and 889 deletions
+7
View File
@@ -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))
+9
View File
@@ -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')
+23 -7
View File
@@ -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
View File
@@ -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)