| |
| """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()) |
|
|