194 lines
7.5 KiB
Python
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())
|