generated from dellevin/template
22
This commit is contained in:
@@ -40,6 +40,7 @@ _tasks = {}
|
||||
DEFAULT_PARAMS = {
|
||||
'uvr_project_path': '', 'model_dir_mode': 'absolute',
|
||||
'demucs_model_dir': '', 'vr_model_dir': '', 'mdx_model_dir': '',
|
||||
'demucs_model_dir_mode': 'absolute', 'vr_model_dir_mode': 'absolute', 'mdx_model_dir_mode': 'absolute',
|
||||
'arch_type': 'Demucs',
|
||||
'save_format': 'wav', 'wav_type': 'PCM_16', 'mp3_bitrate': '320k',
|
||||
'is_gpu': True, 'device_set': 'Default',
|
||||
@@ -70,7 +71,7 @@ def _save_config(cfg):
|
||||
json.dump(cfg, f, ensure_ascii=False, indent=2)
|
||||
|
||||
|
||||
def _resolve_model_dir(raw_dir, cfg=None):
|
||||
def _resolve_model_dir(raw_dir, cfg=None, mode=None):
|
||||
"""将模型目录路径解析为绝对路径"""
|
||||
if not raw_dir:
|
||||
return ''
|
||||
@@ -79,7 +80,8 @@ def _resolve_model_dir(raw_dir, cfg=None):
|
||||
if cfg is None:
|
||||
cfg = _load_config()
|
||||
uvr_path = cfg.get('uvr_project_path', '')
|
||||
mode = cfg.get('model_dir_mode', 'absolute')
|
||||
if mode is None:
|
||||
mode = cfg.get('model_dir_mode', 'absolute')
|
||||
if mode == 'relative' and uvr_path:
|
||||
return os.path.normpath(os.path.join(uvr_path, raw_dir))
|
||||
# absolute 模式下如果输入的是相对路径,也尝试拼接 uvr_project_path
|
||||
@@ -92,13 +94,15 @@ def _resolve_model_dir(raw_dir, cfg=None):
|
||||
|
||||
def _get_model_dir_for_arch(arch_type, explicit_dir=None, cfg=None):
|
||||
"""获取指定架构的模型目录:优先显式路径,否则从配置自动推断"""
|
||||
if explicit_dir:
|
||||
return _resolve_model_dir(explicit_dir, cfg)
|
||||
if cfg is None:
|
||||
cfg = _load_config()
|
||||
mode_map = {'Demucs': 'demucs_model_dir_mode', 'VR Arc': 'vr_model_dir_mode', 'MDX-Net': 'mdx_model_dir_mode'}
|
||||
mode = cfg.get(mode_map.get(arch_type, ''), cfg.get('model_dir_mode', 'absolute'))
|
||||
if explicit_dir:
|
||||
return _resolve_model_dir(explicit_dir, cfg, mode=mode)
|
||||
key_map = {'Demucs': 'demucs_model_dir', 'VR Arc': 'vr_model_dir', 'MDX-Net': 'mdx_model_dir'}
|
||||
raw = cfg.get(key_map.get(arch_type, ''), '')
|
||||
return _resolve_model_dir(raw, cfg)
|
||||
return _resolve_model_dir(raw, cfg, mode=mode)
|
||||
|
||||
|
||||
def _get_uvr_paths():
|
||||
@@ -639,18 +643,47 @@ def _update_progress(task_id, step, inference_iterations=0):
|
||||
return
|
||||
progress = min(99, max(1, int((step + inference_iterations) * 100)))
|
||||
task['progress'] = progress
|
||||
import time as _t
|
||||
# 记录子进度(inference_iterations > 0 表示推理中的迭代进度)
|
||||
if inference_iterations > 0:
|
||||
last = task.get('_last_iter_log', 0)
|
||||
if inference_iterations - last >= 0.1:
|
||||
task['_last_iter_log'] = inference_iterations
|
||||
task['logs'].append('[%s] Inference iteration: %.0f%%' % (_t.strftime('%H:%M:%S'), inference_iterations * 100))
|
||||
if len(task['logs']) > 500:
|
||||
task['logs'] = task['logs'][-500:]
|
||||
else:
|
||||
# 主进度每 20% 记录一条
|
||||
last = task.get('_last_progress_log', 0)
|
||||
if progress - last >= 20:
|
||||
task['_last_progress_log'] = progress
|
||||
task['logs'].append('[%s] Progress: %d%%' % (_t.strftime('%H:%M:%S'), progress))
|
||||
if len(task['logs']) > 500:
|
||||
task['logs'] = task['logs'][-500:]
|
||||
|
||||
|
||||
def _make_process_data(task_id, model_data, audio_path, export_path):
|
||||
"""构造 process_data 字典"""
|
||||
base = os.path.splitext(os.path.basename(audio_path))[0]
|
||||
task = _tasks.get(task_id)
|
||||
def _write_console(*args, **kwargs):
|
||||
msg = ' '.join(str(a) for a in args)
|
||||
msg = msg.replace('\r\n', '\n').replace('\r', '\n').strip()
|
||||
if msg and task:
|
||||
import time as _t
|
||||
for line in msg.split('\n'):
|
||||
line = line.strip()
|
||||
if line:
|
||||
task['logs'].append('[%s] %s' % (_t.strftime('%H:%M:%S'), line))
|
||||
if len(task['logs']) > 500:
|
||||
task['logs'] = task['logs'][-500:]
|
||||
return {
|
||||
'model_data': model_data,
|
||||
'export_path': export_path,
|
||||
'audio_file_base': base,
|
||||
'audio_file': audio_path,
|
||||
'set_progress_bar': lambda step, it=0: _update_progress(task_id, step, it),
|
||||
'write_to_console': lambda *_, **__: None,
|
||||
'write_to_console': _write_console,
|
||||
'process_iteration': lambda: None,
|
||||
'cached_source_callback': lambda *_, **__: (None, None),
|
||||
'cached_model_source_holder': lambda *_, **__: None,
|
||||
@@ -663,6 +696,29 @@ def _make_process_data(task_id, model_data, audio_path, export_path):
|
||||
def _do_separate(task_id, model_path, process_method, audio_path, params, model_meta):
|
||||
"""后台线程执行分离"""
|
||||
task = _tasks[task_id]
|
||||
import time as _time
|
||||
|
||||
# 日志捕获
|
||||
class _LogCapture:
|
||||
def __init__(self, orig):
|
||||
self._orig = orig
|
||||
def write(self, msg):
|
||||
if msg and msg.strip():
|
||||
text = msg.replace('\r\n', '\n').replace('\r', '\n').strip()
|
||||
for line in text.split('\n'):
|
||||
line = line.strip()
|
||||
if line:
|
||||
task['logs'].append('[%s] %s' % (_time.strftime('%H:%M:%S'), line))
|
||||
if len(task['logs']) > 500:
|
||||
task['logs'] = task['logs'][-500:]
|
||||
self._orig.write(msg)
|
||||
def flush(self):
|
||||
self._orig.flush()
|
||||
|
||||
old_stdout, old_stderr = sys.stdout, sys.stderr
|
||||
sys.stdout = _LogCapture(old_stdout)
|
||||
sys.stderr = _LogCapture(old_stderr)
|
||||
|
||||
try:
|
||||
_ensure_uvr_imports()
|
||||
|
||||
@@ -727,6 +783,8 @@ def _do_separate(task_id, model_path, process_method, audio_path, params, model_
|
||||
task['error'] = str(e)
|
||||
traceback.print_exc()
|
||||
finally:
|
||||
sys.stdout = old_stdout
|
||||
sys.stderr = old_stderr
|
||||
try:
|
||||
os.remove(audio_path)
|
||||
except OSError:
|
||||
@@ -945,6 +1003,7 @@ def start_separate():
|
||||
'stems': [],
|
||||
'error': None,
|
||||
'device': '',
|
||||
'logs': [],
|
||||
}
|
||||
|
||||
threading.Thread(
|
||||
@@ -968,6 +1027,7 @@ def status(task_id):
|
||||
}
|
||||
if task.get('device'):
|
||||
resp['device'] = task['device']
|
||||
resp['logs'] = task.get('logs', [])[-200:]
|
||||
if task['status'] == 'done':
|
||||
resp['stems'] = task['stems']
|
||||
resp['count'] = len(task['stems'])
|
||||
@@ -976,6 +1036,14 @@ def status(task_id):
|
||||
return jsonify(resp)
|
||||
|
||||
|
||||
@bp.route('/clear-logs/<task_id>', methods=['POST'])
|
||||
def clear_logs(task_id):
|
||||
task = _tasks.get(task_id)
|
||||
if task:
|
||||
task['logs'] = []
|
||||
return jsonify({'success': True})
|
||||
|
||||
|
||||
@bp.route('/download/<task_id>/<stem>')
|
||||
def download_stem(task_id, stem):
|
||||
task = _tasks.get(task_id)
|
||||
|
||||
Reference in New Issue
Block a user