463 lines
17 KiB
Python
463 lines
17 KiB
Python
from __future__ import annotations
|
|
|
|
import io
|
|
import multiprocessing
|
|
import os
|
|
import threading
|
|
import time
|
|
import traceback
|
|
import uuid
|
|
from multiprocessing.connection import Connection
|
|
from typing import Any, Callable
|
|
|
|
|
|
MAX_NEW_TOKENS = 384
|
|
DEFAULT_SYNTHESIS_TIMEOUT_SECONDS = 120.0
|
|
DEFAULT_STARTUP_TIMEOUT_SECONDS = 480.0
|
|
# 启动失败后再次拉起 worker 的最小间隔:避免"启动超时→立即重启→再超时"的死循环
|
|
# 在系统高负载时持续加载 torch/CUDA,进一步加剧卡顿。
|
|
STARTUP_FAILURE_BACKOFF_SECONDS = 60.0
|
|
DEFAULT_CPU_THREADS = 4
|
|
DEFAULT_CPU_AFFINITY_COUNT = 8
|
|
DEFAULT_PROCESS_PRIORITY = "below_normal"
|
|
|
|
|
|
class FasterQwenWorkerError(RuntimeError):
|
|
pass
|
|
|
|
|
|
class FasterQwenWorkerTimeout(FasterQwenWorkerError):
|
|
pass
|
|
|
|
|
|
def _generation_kwargs(settings: dict[str, Any], text: str) -> dict[str, Any]:
|
|
return {
|
|
"text": text,
|
|
"language": str(settings.get("language") or "Chinese"),
|
|
"non_streaming_mode": bool(settings.get("non_streaming_mode", True)),
|
|
"max_new_tokens": MAX_NEW_TOKENS,
|
|
}
|
|
|
|
|
|
def _bounded_int(value: Any, default: int, *, minimum: int, maximum: int) -> int:
|
|
try:
|
|
parsed = int(value)
|
|
except (TypeError, ValueError):
|
|
parsed = default
|
|
return max(minimum, min(maximum, parsed))
|
|
|
|
|
|
def _configure_worker_environment(settings: dict[str, Any]) -> int:
|
|
cpu_threads = _bounded_int(
|
|
settings.get("cpu_threads"),
|
|
DEFAULT_CPU_THREADS,
|
|
minimum=1,
|
|
maximum=8,
|
|
)
|
|
thread_value = str(cpu_threads)
|
|
for name in (
|
|
"OMP_NUM_THREADS",
|
|
"MKL_NUM_THREADS",
|
|
"OPENBLAS_NUM_THREADS",
|
|
"NUMEXPR_NUM_THREADS",
|
|
"VECLIB_MAXIMUM_THREADS",
|
|
"BLIS_NUM_THREADS",
|
|
):
|
|
os.environ[name] = thread_value
|
|
os.environ["TOKENIZERS_PARALLELISM"] = "false"
|
|
return cpu_threads
|
|
|
|
|
|
def _configure_torch_threads(torch_module: Any, cpu_threads: int) -> None:
|
|
torch_module.set_num_threads(cpu_threads)
|
|
try:
|
|
torch_module.set_num_interop_threads(1)
|
|
except RuntimeError:
|
|
# PyTorch only allows setting interop threads before parallel work starts.
|
|
pass
|
|
|
|
|
|
def _apply_worker_process_limits(settings: dict[str, Any]) -> dict[str, Any]:
|
|
cpu_count = max(1, int(os.cpu_count() or 1))
|
|
affinity_count = _bounded_int(
|
|
settings.get("cpu_affinity_count"),
|
|
DEFAULT_CPU_AFFINITY_COUNT,
|
|
minimum=0,
|
|
maximum=min(cpu_count, 63),
|
|
)
|
|
priority = str(settings.get("process_priority") or DEFAULT_PROCESS_PRIORITY).strip().lower()
|
|
applied_affinity = 0
|
|
applied_priority = "default"
|
|
|
|
if os.name == "nt":
|
|
try:
|
|
import ctypes
|
|
from ctypes import wintypes
|
|
|
|
kernel32 = ctypes.WinDLL("kernel32", use_last_error=True)
|
|
kernel32.GetCurrentProcess.restype = wintypes.HANDLE
|
|
kernel32.SetPriorityClass.argtypes = [wintypes.HANDLE, wintypes.DWORD]
|
|
kernel32.SetPriorityClass.restype = wintypes.BOOL
|
|
kernel32.SetProcessAffinityMask.argtypes = [wintypes.HANDLE, ctypes.c_size_t]
|
|
kernel32.SetProcessAffinityMask.restype = wintypes.BOOL
|
|
process_handle = kernel32.GetCurrentProcess()
|
|
priority_classes = {
|
|
"idle": 0x00000040,
|
|
"below_normal": 0x00004000,
|
|
"normal": 0x00000020,
|
|
}
|
|
priority_class = priority_classes.get(priority, priority_classes[DEFAULT_PROCESS_PRIORITY])
|
|
if kernel32.SetPriorityClass(process_handle, priority_class):
|
|
applied_priority = priority if priority in priority_classes else DEFAULT_PROCESS_PRIORITY
|
|
if affinity_count > 0:
|
|
affinity_mask = (1 << affinity_count) - 1
|
|
if kernel32.SetProcessAffinityMask(process_handle, affinity_mask):
|
|
applied_affinity = affinity_count
|
|
except Exception:
|
|
pass
|
|
|
|
return {
|
|
"cpu_affinity_count": applied_affinity,
|
|
"process_priority": applied_priority,
|
|
}
|
|
|
|
|
|
def _load_runtime(settings: dict[str, Any]) -> dict[str, Any]:
|
|
device = str(settings.get("device") or "cuda")
|
|
if device == "cpu" and "CUDA_VISIBLE_DEVICES" not in os.environ:
|
|
os.environ["CUDA_VISIBLE_DEVICES"] = ""
|
|
|
|
cpu_threads = _configure_worker_environment(settings)
|
|
import torch
|
|
_configure_torch_threads(torch, cpu_threads)
|
|
from faster_qwen3_tts import FasterQwen3TTS
|
|
|
|
if device == "cpu":
|
|
torch.cuda.is_available = lambda: False
|
|
|
|
load_kwargs: dict[str, Any] = {}
|
|
if device == "cpu":
|
|
load_kwargs["device"] = "cpu"
|
|
model = FasterQwen3TTS.from_pretrained(
|
|
str(settings.get("model_name_or_path") or "Qwen/Qwen3-TTS-12Hz-0.6B-Base"),
|
|
**load_kwargs,
|
|
)
|
|
|
|
# voice_clone_prompt 不再预计算:当前 faster_qwen3_tts 的 FasterQwen3TTS 没有
|
|
# create_voice_clone_prompt 方法,预计算只会失败。改为每次合成时在
|
|
# _synthesize_wav 里直接传 ref_audio/ref_text/xvec_only 参数。
|
|
voice_clone_prompt = None
|
|
|
|
return {
|
|
"model": model,
|
|
"torch": torch,
|
|
"voice_clone_prompt": voice_clone_prompt,
|
|
"settings": settings,
|
|
"cpu_threads": cpu_threads,
|
|
}
|
|
|
|
|
|
def _synthesize_wav(runtime: dict[str, Any], text: str) -> bytes:
|
|
import soundfile as sf
|
|
|
|
model = runtime["model"]
|
|
torch = runtime["torch"]
|
|
settings = runtime["settings"]
|
|
safe_text = str(text or "").strip()[:80] or "欢迎来到直播间。"
|
|
kwargs = _generation_kwargs(settings, safe_text)
|
|
voice_clone_prompt = runtime.get("voice_clone_prompt")
|
|
if voice_clone_prompt is not None:
|
|
kwargs["voice_clone_prompt"] = voice_clone_prompt
|
|
else:
|
|
ref_audio = str(settings.get("ref_audio") or "")
|
|
if not ref_audio:
|
|
raise RuntimeError("Faster-Qwen3-TTS requires ref_audio")
|
|
kwargs.update({
|
|
"ref_audio": ref_audio,
|
|
"ref_text": str(settings.get("ref_text") or "") or None,
|
|
"xvec_only": bool(settings.get("xvec_only", True)),
|
|
"append_silence": bool(settings.get("append_silence", True)),
|
|
})
|
|
|
|
if not torch.cuda.is_available():
|
|
torch.backends.cudnn.enabled = False
|
|
with torch.inference_mode():
|
|
wavs, sample_rate = model.generate_voice_clone(**kwargs)
|
|
|
|
output = io.BytesIO()
|
|
audio = wavs[0]
|
|
if isinstance(audio, torch.Tensor):
|
|
audio = audio.cpu().numpy()
|
|
sf.write(output, audio, sample_rate, format="WAV")
|
|
return output.getvalue()
|
|
|
|
|
|
def faster_qwen_worker_main(connection: Connection, settings: dict[str, Any]) -> None:
|
|
try:
|
|
cpu_threads = _configure_worker_environment(settings)
|
|
process_limits = _apply_worker_process_limits(settings)
|
|
load_started = time.monotonic()
|
|
runtime = _load_runtime(settings)
|
|
load_ms = int((time.monotonic() - load_started) * 1000)
|
|
|
|
warmup_started = time.monotonic()
|
|
_synthesize_wav(runtime, "系统启动")
|
|
warmup_ms = int((time.monotonic() - warmup_started) * 1000)
|
|
connection.send({
|
|
"type": "ready",
|
|
"pid": os.getpid(),
|
|
"load_ms": load_ms,
|
|
"warmup_ms": warmup_ms,
|
|
"max_new_tokens": MAX_NEW_TOKENS,
|
|
"cpu_threads": cpu_threads,
|
|
**process_limits,
|
|
})
|
|
except BaseException as exc:
|
|
try:
|
|
connection.send({
|
|
"type": "startup_error",
|
|
"error_type": type(exc).__name__,
|
|
"error": str(exc),
|
|
"traceback": traceback.format_exc(),
|
|
})
|
|
except Exception:
|
|
pass
|
|
connection.close()
|
|
return
|
|
|
|
while True:
|
|
try:
|
|
message = connection.recv()
|
|
except (EOFError, OSError):
|
|
break
|
|
if not isinstance(message, dict):
|
|
continue
|
|
command = str(message.get("command") or "")
|
|
if command == "stop":
|
|
break
|
|
if command != "synthesize":
|
|
continue
|
|
|
|
request_id = str(message.get("request_id") or "")
|
|
started = time.monotonic()
|
|
try:
|
|
audio = _synthesize_wav(runtime, str(message.get("text") or ""))
|
|
connection.send({
|
|
"type": "result",
|
|
"request_id": request_id,
|
|
"audio": audio,
|
|
"duration_ms": int((time.monotonic() - started) * 1000),
|
|
"bytes": len(audio),
|
|
"max_new_tokens": MAX_NEW_TOKENS,
|
|
})
|
|
except BaseException as exc:
|
|
try:
|
|
connection.send({
|
|
"type": "error",
|
|
"request_id": request_id,
|
|
"error_type": type(exc).__name__,
|
|
"error": str(exc),
|
|
"traceback": traceback.format_exc(),
|
|
})
|
|
except Exception:
|
|
break
|
|
connection.close()
|
|
|
|
|
|
class FasterQwenWorkerClient:
|
|
def __init__(
|
|
self,
|
|
settings: dict[str, Any],
|
|
logger=None,
|
|
*,
|
|
synthesis_timeout_seconds: float = DEFAULT_SYNTHESIS_TIMEOUT_SECONDS,
|
|
startup_timeout_seconds: float = DEFAULT_STARTUP_TIMEOUT_SECONDS,
|
|
context=None,
|
|
process_target: Callable[..., None] | None = None,
|
|
):
|
|
self.settings = dict(settings)
|
|
self.logger = logger
|
|
self.synthesis_timeout_seconds = max(1.0, float(synthesis_timeout_seconds))
|
|
self.startup_timeout_seconds = max(10.0, float(startup_timeout_seconds))
|
|
self._context = context or multiprocessing.get_context("spawn")
|
|
self._process_target = process_target or faster_qwen_worker_main
|
|
self._lock = threading.RLock()
|
|
self._process = None
|
|
self._connection = None
|
|
self._worker_pid = 0
|
|
self._next_start_after = 0.0
|
|
|
|
@property
|
|
def worker_pid(self) -> int:
|
|
return int(self._worker_pid or 0)
|
|
|
|
def _log(self, level: str, message: str, *args) -> None:
|
|
if self.logger:
|
|
getattr(self.logger, level)(message, *args)
|
|
|
|
def _is_alive_locked(self) -> bool:
|
|
return bool(self._process is not None and self._process.is_alive())
|
|
|
|
def ensure_ready(self) -> dict[str, Any]:
|
|
with self._lock:
|
|
if self._is_alive_locked() and self._connection is not None:
|
|
return {"pid": self.worker_pid, "reused": True}
|
|
return self._start_worker_locked()
|
|
|
|
def _start_worker_locked(self) -> dict[str, Any]:
|
|
now = time.monotonic()
|
|
if now < self._next_start_after:
|
|
wait = int(self._next_start_after - now)
|
|
raise FasterQwenWorkerError(
|
|
f"TTS worker 启动退避中,距上次启动失败不足 {int(STARTUP_FAILURE_BACKOFF_SECONDS)} 秒,"
|
|
f"约 {wait} 秒后可重试"
|
|
)
|
|
self._terminate_worker_locked("replace_stale_worker", graceful=False)
|
|
parent_connection, child_connection = self._context.Pipe(duplex=True)
|
|
process = self._context.Process(
|
|
target=self._process_target,
|
|
args=(child_connection, self.settings),
|
|
name="FasterQwen3TTSWorker",
|
|
daemon=True,
|
|
)
|
|
process.start()
|
|
try:
|
|
child_connection.close()
|
|
except Exception:
|
|
pass
|
|
self._process = process
|
|
self._connection = parent_connection
|
|
self._worker_pid = int(getattr(process, "pid", 0) or 0)
|
|
self._log("info", "[FasterQwenTTS] worker 已启动, pid=%s,正在加载和预热", self.worker_pid)
|
|
|
|
if not parent_connection.poll(self.startup_timeout_seconds):
|
|
self._next_start_after = time.monotonic() + STARTUP_FAILURE_BACKOFF_SECONDS
|
|
self._terminate_worker_locked("startup_timeout", graceful=False)
|
|
raise FasterQwenWorkerError(
|
|
f"Faster-Qwen3-TTS worker startup exceeded {self.startup_timeout_seconds:.0f}s"
|
|
)
|
|
try:
|
|
message = parent_connection.recv()
|
|
except (EOFError, OSError) as exc:
|
|
self._next_start_after = time.monotonic() + STARTUP_FAILURE_BACKOFF_SECONDS
|
|
self._terminate_worker_locked("startup_connection_closed", graceful=False)
|
|
raise FasterQwenWorkerError("Faster-Qwen3-TTS worker exited during startup") from exc
|
|
if not isinstance(message, dict) or message.get("type") != "ready":
|
|
if isinstance(message, dict):
|
|
error = str(message.get("error") or message.get("error_type") or "unknown startup error")
|
|
else:
|
|
error = "invalid startup response"
|
|
self._next_start_after = time.monotonic() + STARTUP_FAILURE_BACKOFF_SECONDS
|
|
self._terminate_worker_locked("startup_error", graceful=False)
|
|
raise FasterQwenWorkerError(error)
|
|
self._next_start_after = 0.0
|
|
self._log(
|
|
"info",
|
|
"[FasterQwenTTS] worker 预热完成, pid=%s load=%sms warmup=%sms max_new_tokens=%s",
|
|
self.worker_pid,
|
|
message.get("load_ms"),
|
|
message.get("warmup_ms"),
|
|
message.get("max_new_tokens"),
|
|
)
|
|
self._log(
|
|
"info",
|
|
"[FasterQwenTTS] worker 资源限制: cpu_threads=%s affinity=%s priority=%s",
|
|
message.get("cpu_threads"),
|
|
message.get("cpu_affinity_count"),
|
|
message.get("process_priority"),
|
|
)
|
|
return message
|
|
|
|
def synthesize(self, text: str) -> tuple[bytes, dict[str, Any]]:
|
|
with self._lock:
|
|
self.ensure_ready()
|
|
request_id = uuid.uuid4().hex
|
|
connection = self._connection
|
|
try:
|
|
connection.send({
|
|
"command": "synthesize",
|
|
"request_id": request_id,
|
|
"text": str(text or ""),
|
|
})
|
|
except (BrokenPipeError, EOFError, OSError) as exc:
|
|
self._restart_after_failure_locked("send_failed")
|
|
raise FasterQwenWorkerError("Faster-Qwen3-TTS worker connection failed") from exc
|
|
|
|
if not connection.poll(self.synthesis_timeout_seconds):
|
|
self._log(
|
|
"error",
|
|
"[FasterQwenTTS] 单次合成超过 %.0f 秒,强制终止 worker pid=%s",
|
|
self.synthesis_timeout_seconds,
|
|
self.worker_pid,
|
|
)
|
|
restart_error = self._restart_after_failure_locked("synthesis_timeout")
|
|
suffix = f"; restart failed: {restart_error}" if restart_error else ""
|
|
raise FasterQwenWorkerTimeout(
|
|
f"Faster-Qwen3-TTS synthesis exceeded {self.synthesis_timeout_seconds:.0f}s{suffix}"
|
|
)
|
|
try:
|
|
message = connection.recv()
|
|
except (EOFError, OSError) as exc:
|
|
self._restart_after_failure_locked("worker_exited")
|
|
raise FasterQwenWorkerError("Faster-Qwen3-TTS worker exited during synthesis") from exc
|
|
|
|
if not isinstance(message, dict) or message.get("request_id") != request_id:
|
|
self._restart_after_failure_locked("invalid_response")
|
|
raise FasterQwenWorkerError("Faster-Qwen3-TTS worker returned an invalid response")
|
|
if message.get("type") == "error":
|
|
error = str(message.get("error") or message.get("error_type") or "synthesis failed")
|
|
lowered = error.lower()
|
|
if "cuda" in lowered or "out of memory" in lowered or "device-side" in lowered:
|
|
self._restart_after_failure_locked("cuda_error")
|
|
raise FasterQwenWorkerError(error)
|
|
if message.get("type") != "result" or not isinstance(message.get("audio"), bytes):
|
|
self._restart_after_failure_locked("invalid_result")
|
|
raise FasterQwenWorkerError("Faster-Qwen3-TTS worker returned no audio")
|
|
return message["audio"], message
|
|
|
|
def _restart_after_failure_locked(self, reason: str) -> str:
|
|
self._terminate_worker_locked(reason, graceful=False)
|
|
try:
|
|
self._start_worker_locked()
|
|
return ""
|
|
except Exception as exc:
|
|
self._log("error", "[FasterQwenTTS] worker 自动重启失败: %s", exc)
|
|
return str(exc)
|
|
|
|
def _terminate_worker_locked(self, reason: str, *, graceful: bool) -> None:
|
|
process = self._process
|
|
connection = self._connection
|
|
self._process = None
|
|
self._connection = None
|
|
self._worker_pid = 0
|
|
if process is None:
|
|
if connection is not None:
|
|
try:
|
|
connection.close()
|
|
except Exception:
|
|
pass
|
|
return
|
|
|
|
if graceful and process.is_alive() and connection is not None:
|
|
try:
|
|
connection.send({"command": "stop"})
|
|
process.join(timeout=3.0)
|
|
except Exception:
|
|
pass
|
|
if process.is_alive():
|
|
self._log("warning", "[FasterQwenTTS] 终止 worker, reason=%s pid=%s", reason, process.pid)
|
|
process.terminate()
|
|
process.join(timeout=10.0)
|
|
if process.is_alive():
|
|
process.kill()
|
|
process.join(timeout=5.0)
|
|
if connection is not None:
|
|
try:
|
|
connection.close()
|
|
except Exception:
|
|
pass
|
|
|
|
def close(self) -> None:
|
|
with self._lock:
|
|
self._terminate_worker_locked("shutdown", graceful=True)
|