Download runtime/experimental/nvfp4/server.py from ajh-code/Fish-Audio-S2-Pro-NVFP4-Balanced: direct link, hf CLI and curl.
- Browser
- Download file 3.79 kB
-
https://huggingface.co/ajh-code/Fish-Audio-S2-Pro-NVFP4-Balanced/resolve/main/runtime/experimental/nvfp4/server.py
- Command line
-
hf download hf://ajh-code/Fish-Audio-S2-Pro-NVFP4-Balanced/runtime/experimental/nvfp4/server.py
-
curl -L -o server.py https://huggingface.co/ajh-code/Fish-Audio-S2-Pro-NVFP4-Balanced/resolve/main/runtime/experimental/nvfp4/server.py
3.79 kB
| """Threaded Fish Speech queue for standalone mixed NVFP4 checkpoints.""" | |
| from __future__ import annotations | |
| import queue | |
| import threading | |
| import traceback | |
| from pathlib import Path | |
| import torch | |
| from loguru import logger | |
| from fish_speech.models.text2semantic.inference import ( | |
| GenerateRequest, | |
| WrappedGenerateResponse, | |
| decode_one_token_ar, | |
| generate_long, | |
| ) | |
| from experimental.fp8 import MXFP8Linear | |
| from .checkpoint import load_mixed_nvfp4_checkpoint | |
| from .modules import NVFP4Linear | |
| def launch_mixed_nvfp4_thread_safe_queue( | |
| checkpoint_path: str | Path, | |
| device: str | torch.device, | |
| precision: torch.dtype, | |
| compile: bool = False, | |
| *, | |
| max_length: int = 3072, | |
| verify_checksums: bool = False, | |
| ) -> queue.Queue: | |
| """Load the frozen mixed checkpoint once and serialize generation.""" | |
| if precision is not torch.bfloat16: | |
| raise ValueError("The mixed NVFP4 service requires BF16 exclusions") | |
| if compile: | |
| raise ValueError("torch.compile is not qualified for this checkpoint") | |
| if max_length < 1024: | |
| raise ValueError("NVFP4 API cache length must be at least 1024") | |
| input_queue: queue.Queue = queue.Queue() | |
| init_event = threading.Event() | |
| init_error: list[BaseException] = [] | |
| def worker() -> None: | |
| try: | |
| model = load_mixed_nvfp4_checkpoint( | |
| checkpoint_path, | |
| device=device, | |
| max_length=max_length, | |
| verify_checksums=verify_checksums, | |
| ) | |
| with torch.device(device): | |
| model.setup_caches( | |
| max_batch_size=1, | |
| max_seq_len=model.config.max_seq_len, | |
| dtype=torch.bfloat16, | |
| ) | |
| model._cache_setup_done = True | |
| nvfp4_count = sum( | |
| isinstance(module, NVFP4Linear) for module in model.modules() | |
| ) | |
| mxfp8_count = sum( | |
| isinstance(module, MXFP8Linear) for module in model.modules() | |
| ) | |
| if (nvfp4_count, mxfp8_count) != (60, 120): | |
| raise RuntimeError( | |
| f"Expected 60/120 NVFP4/MXFP8 modules, got {nvfp4_count}/{mxfp8_count}" | |
| ) | |
| logger.info( | |
| "Loaded standalone English NVFP4 checkpoint with " | |
| f"{nvfp4_count} NVFP4 and {mxfp8_count} MXFP8 projections on {device}" | |
| ) | |
| except BaseException as error: | |
| init_error.append(error) | |
| logger.error(traceback.format_exc()) | |
| init_event.set() | |
| return | |
| init_event.set() | |
| while True: | |
| item: GenerateRequest | None = input_queue.get() | |
| if item is None: | |
| break | |
| response_queue = item.response_queue | |
| try: | |
| for chunk in generate_long( | |
| model=model, | |
| decode_one_token=decode_one_token_ar, | |
| **item.request, | |
| ): | |
| response_queue.put( | |
| WrappedGenerateResponse(status="success", response=chunk) | |
| ) | |
| except Exception as error: | |
| logger.error(traceback.format_exc()) | |
| response_queue.put( | |
| WrappedGenerateResponse(status="error", response=error) | |
| ) | |
| finally: | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| threading.Thread( | |
| target=worker, | |
| name="fish-s2-pro-mixed-nvfp4-worker", | |
| daemon=True, | |
| ).start() | |
| init_event.wait() | |
| if init_error: | |
| raise RuntimeError("Mixed NVFP4 model worker failed to initialize") from init_error[0] | |
| return input_queue | |