ajh-code's picture
Add files using upload-large-folder tool
16f5171 verified
Raw
History Blame
5.28 kB
#!/usr/bin/env python
"""Serve the self-contained English S2-Pro mixed NVFP4 V1 checkpoint."""
from __future__ import annotations
import argparse
import json
import multiprocessing
import os
import sys
from argparse import Namespace
from pathlib import Path
import torch
import uvicorn
from kui.asgi import JSONResponse
ROOT = Path(
os.environ.get("FISH_NVFP4_ROOT", Path(__file__).resolve().parents[1])
).resolve()
RUNTIME_ROOT = ROOT / "runtime"
if str(RUNTIME_ROOT) not in sys.path:
sys.path.insert(0, str(RUNTIME_ROOT))
from experimental.codec import (
load_compact_codec_model,
load_reference_audio_soundfile,
warm_reference_encoder,
)
from experimental.nvfp4 import launch_mixed_nvfp4_thread_safe_queue
DEFAULT_CHECKPOINT = ROOT
DEFAULT_WEB_UI = RUNTIME_ROOT / "web" / "index.html"
EXPECTED_FORMAT = "fish-s2-pro-project-local-nvfp4-mixed"
EXPECTED_POLICY = "w4a16_gate_up_middle30_mxfp8_rest"
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--checkpoint", type=Path, default=DEFAULT_CHECKPOINT)
parser.add_argument("--host", default="0.0.0.0")
parser.add_argument("--port", type=int, default=8080)
parser.add_argument("--device", default="cuda:0")
parser.add_argument("--cache-length", type=int, default=3072)
parser.add_argument("--max-text-length", type=int, default=0)
parser.add_argument("--api-key")
parser.add_argument("--verify-checksums", action="store_true")
return parser.parse_args()
def validate(args: argparse.Namespace) -> dict:
if not 1 <= args.port <= 65535:
raise SystemExit("--port must be between 1 and 65535")
metadata_path = args.checkpoint / "quantization.json"
codec_path = args.checkpoint / "codec.pth"
if not metadata_path.is_file() or not codec_path.is_file():
raise SystemExit(f"Incomplete NVFP4 checkpoint: {args.checkpoint}")
metadata = json.loads(metadata_path.read_text())
if metadata.get("format") != EXPECTED_FORMAT:
raise SystemExit(f"Unsupported checkpoint format: {metadata.get('format')}")
if metadata.get("policy") != EXPECTED_POLICY:
raise SystemExit(f"Unexpected NVFP4 policy: {metadata.get('policy')}")
if metadata.get("fresh_load_verification", {}).get("status") != "passed":
raise SystemExit("Checkpoint lacks a passed fresh-load verification")
if not torch.cuda.is_available() or torch.cuda.get_device_capability(args.device)[0] != 12:
raise SystemExit("The mixed NVFP4 kernels require an SM120 CUDA GPU")
return metadata
def main() -> int:
args = parse_args()
metadata = validate(args)
import tools.server.model_manager as model_manager_module
import tools.server.views as server_views
from fish_speech.inference_engine.reference_loader import ReferenceLoader
def checkpoint_queue_loader(checkpoint_path, device, precision, compile=False):
return launch_mixed_nvfp4_thread_safe_queue(
checkpoint_path,
device,
precision,
compile,
max_length=args.cache_length,
verify_checksums=args.verify_checksums,
)
def compact_codec_loader(config_name, checkpoint_path, device):
codec = load_compact_codec_model(
config_name,
checkpoint_path,
device,
torch.bfloat16,
offload_reference=True,
)
codec._reference_warmup_report = warm_reference_encoder(codec, device)
return codec
model_manager_module.launch_thread_safe_queue = checkpoint_queue_loader
model_manager_module.load_decoder_model = compact_codec_loader
ReferenceLoader.load_audio = staticmethod(load_reference_audio_soundfile)
server_views._WEBUI_HTML = DEFAULT_WEB_UI
@server_views.routes.http.get("/v1/model")
async def model_info():
return JSONResponse(
{
"model": "V1 · Fish Audio S2-Pro NVFP4 Balanced",
"checkpoint": str(args.checkpoint),
"checkpoint_status": metadata["status"],
"quantization": "60 native NVFP4 W4A16/W4A4 + 120 native MXFP8 W8A8",
"release": metadata.get("release"),
"policy": metadata["policy"],
"sampling": metadata["qualified_sampling"],
"fresh_load_verification": metadata["fresh_load_verification"],
"device": args.device,
"cache_length": args.cache_length,
}
)
from tools.api_server import API
upstream_args = Namespace(
mode="tts",
device=args.device,
half=False,
compile=False,
llama_checkpoint_path=str(args.checkpoint),
decoder_checkpoint_path=str(args.checkpoint / "codec.pth"),
decoder_config_name="modded_dac_vq",
max_text_length=args.max_text_length,
listen=f"{args.host}:{args.port}",
workers=1,
api_key=args.api_key,
)
multiprocessing.set_start_method("spawn", force=True)
app = API(args=upstream_args).app
uvicorn.run(app, host=args.host, port=args.port, workers=1, log_level="info")
return 0
if __name__ == "__main__":
raise SystemExit(main())