#!/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())