vignette/scripts/higgs-tts-server.py

194 lines
7.5 KiB
Python

# -*- coding: utf-8 -*-
"""권리 안전한 Vignette P1용 Higgs Audio v3 상주 TTS 서버.
Higgs 모델은 한 번만 GPU에 올리고, 저장소의 무참조 synthetic seed만 화자
reference로 사용한다. 실존 인물/성우 음성은 읽지 않는다. 모델 라이선스 때문에
127.0.0.1 로컬 개발 환경에서만 실행한다.
"""
from __future__ import annotations
import argparse
import importlib
import io
import json
import sys
import time
from http.server import BaseHTTPRequestHandler, HTTPServer
from pathlib import Path
from typing import Any
import soundfile as sf
import torch
REPO_ROOT = Path(__file__).resolve().parents[1]
DEFAULT_COMFY_ROOT = Path(r"C:\Users\encep\Tools\ComfyUI_windows_portable")
DEFAULT_REFERENCE_DIR = (
REPO_ROOT / "docs" / "voice-art" / "p1-seoyeon-higgs-v3-20260627"
)
MODEL_ID = "higgs-audio-v3-tts-4b"
MAX_TEXT_CHARS = 2_000
def _load_reference(reference_dir: Path) -> tuple[dict[str, Any], str]:
manifest = json.loads((reference_dir / "manifest.json").read_text(encoding="utf-8"))
seed = manifest["synthetic_seed"]
reference_path = reference_dir / seed["wav"]
samples, sample_rate = sf.read(reference_path, dtype="float32")
if samples.ndim == 2:
samples = samples.mean(axis=1)
reference = {
"waveform": torch.from_numpy(samples[None, None, :]).float(),
"sample_rate": int(sample_rate),
}
return reference, str(seed["text"])
def _load_bundle(comfy_root: Path):
node_root = comfy_root / "ComfyUI" / "custom_nodes"
if not node_root.is_dir():
raise RuntimeError(f"ComfyUI custom_nodes 경로가 없습니다: {node_root}")
sys.path.insert(0, str(node_root))
loader = importlib.import_module("Higgs_v3-TTS-ComfyUI.loader")
native = importlib.import_module("Higgs_v3-TTS-ComfyUI.native")
choices = list(loader.get_model_choices())
choice = MODEL_ID if MODEL_ID in choices else next(
(item for item in choices if "higgs" in item.casefold()), None
)
if choice is None:
raise RuntimeError("설치된 Higgs Audio v3 TTS 모델을 찾지 못했습니다.")
bundle = loader.load_higgs_bundle(
model_choice=choice,
dtype_name="auto",
device_name="auto",
attention="auto",
download_if_missing=False,
)
return native, bundle, choice
class HiggsRuntime:
def __init__(self, comfy_root: Path, reference_dir: Path) -> None:
started = time.perf_counter()
self.native, self.bundle, self.model_choice = _load_bundle(comfy_root)
self.reference_audio, self.reference_text = _load_reference(reference_dir)
self.loaded_seconds = round(time.perf_counter() - started, 3)
def synthesize_wav(self, text: str) -> bytes:
generated = self.native.generate_higgs_audio(
self.bundle,
text=text,
reference_audio=self.reference_audio,
reference_audio_path="",
reference_text=self.reference_text,
max_new_tokens=2048,
temperature=0.8,
top_p=0.95,
top_k=50,
seed=0,
trim_reference_audio=True,
silence_threshold_db=-42.0,
max_reference_seconds=12.0,
progress_callback=None,
)
waveform = generated["waveform"]
if not isinstance(waveform, torch.Tensor):
waveform = torch.as_tensor(waveform)
data = waveform.detach().float().cpu()
if data.ndim == 3:
data = data[0]
if data.ndim == 2:
data = data.numpy().T
else:
data = data.numpy()
output = io.BytesIO()
sf.write(output, data, int(generated["sample_rate"]), format="WAV")
return output.getvalue()
def make_handler(runtime: HiggsRuntime):
class Handler(BaseHTTPRequestHandler):
server_version = "VignetteHiggsTTS/1.0"
def log_message(self, format: str, *args: object) -> None:
print(f"[higgs-tts] {self.address_string()} {format % args}", flush=True)
def _send_json(self, status: int, body: dict[str, Any]) -> None:
payload = json.dumps(body, ensure_ascii=False).encode("utf-8")
self.send_response(status)
self.send_header("Content-Type", "application/json; charset=utf-8")
self.send_header("Content-Length", str(len(payload)))
self.send_header("Cache-Control", "no-store")
self.end_headers()
self.wfile.write(payload)
def do_GET(self) -> None: # noqa: N802
if self.path != "/health":
self._send_json(404, {"detail": "not found"})
return
self._send_json(
200,
{
"status": "ok",
"model": MODEL_ID,
"model_choice": runtime.model_choice,
"loaded_seconds": runtime.loaded_seconds,
"reference_policy": "synthetic-seed-only",
},
)
def do_POST(self) -> None: # noqa: N802
if self.path != "/tts":
self._send_json(404, {"detail": "not found"})
return
try:
content_length = int(self.headers.get("Content-Length", "0"))
if content_length <= 0 or content_length > 32_768:
raise ValueError("invalid content length")
body = json.loads(self.rfile.read(content_length).decode("utf-8"))
text = str(body.get("text") or "").strip()
if not text:
raise ValueError("text is required")
if len(text) > MAX_TEXT_CHARS:
raise ValueError(f"text exceeds {MAX_TEXT_CHARS} characters")
started = time.perf_counter()
audio = runtime.synthesize_wav(text)
elapsed = time.perf_counter() - started
self.send_response(200)
self.send_header("Content-Type", "audio/wav")
self.send_header("Content-Length", str(len(audio)))
self.send_header("Cache-Control", "no-store")
self.send_header("X-Higgs-Generation-Seconds", f"{elapsed:.3f}")
self.end_headers()
self.wfile.write(audio)
except (ValueError, json.JSONDecodeError) as exc:
self._send_json(400, {"detail": str(exc)})
except Exception as exc: # 모델 오류는 본문에 민감정보 없이 타입만 노출
print(f"[higgs-tts] generation failed: {exc!r}", flush=True)
self._send_json(500, {"detail": type(exc).__name__})
return Handler
def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument("--host", default="127.0.0.1", choices=["127.0.0.1", "localhost"])
parser.add_argument("--port", type=int, default=9881)
parser.add_argument("--comfy-root", type=Path, default=DEFAULT_COMFY_ROOT)
parser.add_argument("--reference-dir", type=Path, default=DEFAULT_REFERENCE_DIR)
args = parser.parse_args()
print("[higgs-tts] Higgs 모델과 synthetic seed를 로드합니다...", flush=True)
runtime = HiggsRuntime(args.comfy_root, args.reference_dir)
print(
f"[higgs-tts] ready model={runtime.model_choice} load={runtime.loaded_seconds}s "
f"url=http://{args.host}:{args.port}",
flush=True,
)
HTTPServer((args.host, args.port), make_handler(runtime)).serve_forever()
return 0
if __name__ == "__main__":
raise SystemExit(main())