File size: 5,276 Bytes
16f5171 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 | #!/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())
|