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