refactor: load models asynchronously and add CUDA warmup for VoiceDesign and CustomVoice servers

This commit is contained in:
mARTin-B78 2026-05-26 23:07:51 +02:00
parent eef48d11f5
commit d04fa88853
2 changed files with 96 additions and 27 deletions

View File

@ -21,6 +21,7 @@ import uvicorn
from fastapi import FastAPI, HTTPException from fastapi import FastAPI, HTTPException
from fastapi.responses import JSONResponse, Response, StreamingResponse from fastapi.responses import JSONResponse, Response, StreamingResponse
from pydantic import BaseModel from pydantic import BaseModel
from contextlib import asynccontextmanager
sys.path.append("/app") sys.path.append("/app")
from faster_qwen3_tts.model import FasterQwen3TTS from faster_qwen3_tts.model import FasterQwen3TTS
@ -35,6 +36,51 @@ default_voice: str = None
SAMPLE_RATE = 24000 SAMPLE_RATE = 24000
DEFAULT_MAX_NEW_TOKENS = 2048 DEFAULT_MAX_NEW_TOKENS = 2048
_model_lock = threading.Lock() _model_lock = threading.Lock()
_load_model_kwargs = None
@asynccontextmanager
async def lifespan(app: FastAPI):
loop = asyncio.get_event_loop()
loop.run_in_executor(None, _do_load_and_warmup)
yield
def _do_load_and_warmup():
global tts_model, SAMPLE_RATE
import torch
args = _load_model_kwargs
try:
logger.info("Loading CustomVoice model %s ...", args.model)
model = FasterQwen3TTS.from_pretrained(
args.model,
device=args.device,
dtype=torch.bfloat16,
attn_implementation="sdpa",
max_seq_len=args.max_seq_len,
)
SAMPLE_RATE = model.sample_rate
logger.info("Model ready. Sample rate: %d Hz", SAMPLE_RATE)
# Warmup
logger.info("Warming up CUDA graphs (first request will be fast)...")
try:
for _ in model.generate_custom_voice_streaming(
text="Warmup.",
speaker="Ryan",
language="English"
):
pass
logger.info("CUDA warmup complete — server ready.")
except Exception as exc:
logger.warning("Warmup failed (non-fatal): %s", exc)
tts_model = model
except Exception as exc:
logger.error("Failed to load model: %s", exc)
app = FastAPI(lifespan=lifespan)
class SpeechRequest(BaseModel): class SpeechRequest(BaseModel):
@ -196,7 +242,7 @@ async def options_handler(path: str):
def main(): def main():
global tts_model, voices, default_voice, SAMPLE_RATE, DEFAULT_MAX_NEW_TOKENS global voices, default_voice, SAMPLE_RATE, DEFAULT_MAX_NEW_TOKENS, _load_model_kwargs
parser = argparse.ArgumentParser() parser = argparse.ArgumentParser()
parser.add_argument("--model", default="/models/Qwen3-TTS-CustomVoice") parser.add_argument("--model", default="/models/Qwen3-TTS-CustomVoice")
@ -207,25 +253,13 @@ def main():
parser.add_argument("--max-seq-len", type=int, default=2048) parser.add_argument("--max-seq-len", type=int, default=2048)
args = parser.parse_args() args = parser.parse_args()
DEFAULT_MAX_NEW_TOKENS = args.max_seq_len DEFAULT_MAX_NEW_TOKENS = args.max_seq_len
_load_model_kwargs = args
with open(args.voices) as f: with open(args.voices) as f:
voices = json.load(f) voices = json.load(f)
default_voice = next(iter(voices), None) default_voice = next(iter(voices), None)
_build_voice_list() _build_voice_list()
import torch
logger.info("Loading CustomVoice model %s ...", args.model)
tts_model = FasterQwen3TTS.from_pretrained(
args.model,
device=args.device,
dtype=torch.bfloat16,
attn_implementation="sdpa",
max_seq_len=args.max_seq_len,
)
SAMPLE_RATE = tts_model.sample_rate
logger.info("Model ready. Sample rate: %d Hz", SAMPLE_RATE)
uvicorn.run(app, host=args.host, port=args.port, log_level="info") uvicorn.run(app, host=args.host, port=args.port, log_level="info")

View File

@ -20,6 +20,7 @@ import uvicorn
from fastapi import FastAPI, HTTPException from fastapi import FastAPI, HTTPException
from fastapi.responses import Response, StreamingResponse, JSONResponse from fastapi.responses import Response, StreamingResponse, JSONResponse
from pydantic import BaseModel from pydantic import BaseModel
from contextlib import asynccontextmanager
sys.path.append("/app") sys.path.append("/app")
from faster_qwen3_tts.model import FasterQwen3TTS from faster_qwen3_tts.model import FasterQwen3TTS
@ -34,6 +35,51 @@ default_voice: str = None
SAMPLE_RATE = 24000 SAMPLE_RATE = 24000
DEFAULT_MAX_NEW_TOKENS = 2048 DEFAULT_MAX_NEW_TOKENS = 2048
_model_lock = threading.Lock() _model_lock = threading.Lock()
_load_model_kwargs = None
@asynccontextmanager
async def lifespan(app: FastAPI):
loop = asyncio.get_event_loop()
loop.run_in_executor(None, _do_load_and_warmup)
yield
def _do_load_and_warmup():
global tts_model, SAMPLE_RATE
import torch
args = _load_model_kwargs
try:
logger.info("Loading VoiceDesign model %s", args.model)
model = FasterQwen3TTS.from_pretrained(
args.model,
device=args.device,
dtype=torch.bfloat16,
attn_implementation="sdpa",
max_seq_len=args.max_seq_len,
)
SAMPLE_RATE = model.sample_rate
logger.info("Model ready. Sample rate: %d Hz", SAMPLE_RATE)
# Warmup
logger.info("Warming up CUDA graphs (first request will be fast)...")
try:
for _ in model.generate_voice_design_streaming(
text="Warmup.",
instruct="Warmup.",
language="English"
):
pass
logger.info("CUDA warmup complete — server ready.")
except Exception as exc:
logger.warning("Warmup failed (non-fatal): %s", exc)
tts_model = model
except Exception as exc:
logger.error("Failed to load model: %s", exc)
app = FastAPI(lifespan=lifespan)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@ -208,7 +254,7 @@ async def options_handler(path: str):
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
def main(): def main():
global tts_model, voices, default_voice, SAMPLE_RATE, DEFAULT_MAX_NEW_TOKENS global voices, default_voice, SAMPLE_RATE, DEFAULT_MAX_NEW_TOKENS, _load_model_kwargs
parser = argparse.ArgumentParser() parser = argparse.ArgumentParser()
parser.add_argument("--model", default="/models/Qwen3-TTS-VoiceDesign") parser.add_argument("--model", default="/models/Qwen3-TTS-VoiceDesign")
@ -219,24 +265,13 @@ def main():
parser.add_argument("--max-seq-len", type=int, default=2048) parser.add_argument("--max-seq-len", type=int, default=2048)
args = parser.parse_args() args = parser.parse_args()
DEFAULT_MAX_NEW_TOKENS = args.max_seq_len DEFAULT_MAX_NEW_TOKENS = args.max_seq_len
_load_model_kwargs = args
with open(args.voices) as f: with open(args.voices) as f:
voices = json.load(f) voices = json.load(f)
default_voice = next(iter(voices), None) default_voice = next(iter(voices), None)
_build_voice_list() _build_voice_list()
import torch
logger.info("Loading VoiceDesign model %s", args.model)
tts_model = FasterQwen3TTS.from_pretrained(
args.model,
device=args.device,
dtype=torch.bfloat16,
attn_implementation="sdpa",
max_seq_len=args.max_seq_len,
)
SAMPLE_RATE = tts_model.sample_rate
logger.info("Model ready. Sample rate: %d Hz", SAMPLE_RATE)
uvicorn.run(app, host=args.host, port=args.port, log_level="info") uvicorn.run(app, host=args.host, port=args.port, log_level="info")