This commit is contained in:
DelLevin-Home
2026-06-17 10:43:33 +08:00
parent db033c661f
commit 05e733e872
18 changed files with 3399 additions and 254 deletions

View File

@@ -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)