#!/usr/bin/env python3 """Validate the self-contained Fish S2-Pro NVFP4/MXFP8 V1 release.""" from __future__ import annotations import hashlib import json import struct from pathlib import Path, PurePosixPath ROOT = Path(__file__).resolve().parent DTYPE_BYTES = { "BF16": 2, "F32": 4, "F8_E4M3": 1, "I32": 4, "U8": 1, } def fail(message: str) -> None: raise RuntimeError(message) def sha256(path: Path) -> str: digest = hashlib.sha256() with path.open("rb") as handle: for chunk in iter(lambda: handle.read(8 * 1024 * 1024), b""): digest.update(chunk) return digest.hexdigest() def load_json(relative: str) -> dict: path = ROOT / relative if not path.is_file(): fail(f"missing required JSON file: {relative}") try: return json.loads(path.read_text(encoding="utf-8")) except (OSError, json.JSONDecodeError) as error: fail(f"invalid JSON in {relative}: {error}") def safetensors_header(path: Path) -> dict: with path.open("rb") as handle: raw_length = handle.read(8) if len(raw_length) != 8: fail(f"truncated safetensors length: {path.name}") header_length = struct.unpack(" path.stat().st_size - 8: fail(f"invalid safetensors header length: {path.name}") try: return json.loads(handle.read(header_length)) except json.JSONDecodeError as error: fail(f"invalid safetensors header in {path.name}: {error}") def validate_metadata() -> dict: config = load_json("config.json") policy = config.get("fish_s2_quantization") if not isinstance(policy, dict): fail("config.json has no fish_s2_quantization dictionary") expected = { "format": "mixed_nvfp4_mxfp8", "profile": "balanced", "release": "v1", "policy": "w4a16_gate_up_middle30_mxfp8_rest", "nvfp4_modules": 60, "mxfp8_modules": 120, "runtime": "bundled", "hardware_family": "sm_120", } if policy != expected: fail(f"unexpected config quantization policy: {policy}") quantization = load_json("quantization.json") if quantization.get("format") != "fish-s2-pro-project-local-nvfp4-mixed": fail("unexpected quantization format") if quantization.get("policy") != "w4a16_gate_up_middle30_mxfp8_rest": fail("unexpected quantization policy") release = quantization.get("release", {}) if release.get("version") != "1.0" or release.get("xpo3_release") is not False: fail("release must identify V1 and explicitly remain outside XPO3") if release.get("repository_name") != "ajh-code/Fish-Audio-S2-Pro-NVFP4-Balanced": fail("release metadata has the wrong Hugging Face repository name") if release.get("self_contained_weights") is not True: fail("release does not declare self-contained weights") fresh_load = quantization.get("fresh_load_verification", {}) if fresh_load.get("status") != "passed": fail("checkpoint lacks passed fresh-load verification") conversion = quantization.get("conversion", {}) if len(conversion.get("records", [])) != 180: fail("quantization metadata must contain 180 projection records") if conversion.get("correction_parameters") != 0: fail("V1 loader does not accept correction-bearing checkpoints") if conversion.get("w4a16_max_m") != 1: fail("V1 must use the qualified M=1 W4A16 decode boundary") source = (ROOT / ".source").read_text(encoding="utf-8") for revision in ( "1de9996b6be38b745688de084d87a5633f714e4e", "e5e292632cb11e7a27b2b7487f58f612bc101e13", "a04c1b63b1a7a670840fb3e97a82c0dbe2a35ded", "7a03467b90d6feff6bd196928dfe156bd173f36e", ): if revision not in source: fail(f".source is missing pinned revision {revision}") license_text = (ROOT / "LICENSE.md").read_text(encoding="utf-8") if "FISH AUDIO RESEARCH LICENSE AGREEMENT" not in license_text: fail("LICENSE.md is not the Fish Audio Research License") notice = (ROOT / "Notice").read_text(encoding="utf-8") if "This model is licensed under the Fish Audio Research License" not in notice: fail("Notice lacks the required Fish Audio attribution") if "Built with Fish Audio" not in notice: fail("Notice lacks the required Built with Fish Audio statement") readme = (ROOT / "README.md").read_text(encoding="utf-8") for required in ( "Built with Fish Audio", "This is not yet an XPO3 release", "Commercial use requires a separate", ): if required not in readme: fail(f"README.md lacks required release statement: {required}") return quantization def validate_checkpoint(quantization: dict) -> tuple[int, int, int]: index = load_json("model.safetensors.index.json") weight_map = index.get("weight_map") if not isinstance(weight_map, dict) or not weight_map: fail("checkpoint index has no weight map") shard_names = sorted(set(weight_map.values())) if shard_names != [ "model-00001-of-00003.safetensors", "model-00002-of-00003.safetensors", "model-00003-of-00003.safetensors", ]: fail(f"unexpected checkpoint shards: {shard_names}") discovered: dict[str, str] = {} logical_bytes = 0 for shard_name in shard_names: shard_path = ROOT / shard_name if not shard_path.is_file(): fail(f"missing checkpoint shard: {shard_name}") header = safetensors_header(shard_path) for name, record in header.items(): if name == "__metadata__": continue if name in discovered: fail(f"duplicate tensor across shards: {name}") dtype = record.get("dtype") shape = record.get("shape") offsets = record.get("data_offsets") if dtype not in DTYPE_BYTES or not isinstance(shape, list): fail(f"unsupported tensor metadata for {name}") if ( not isinstance(offsets, list) or len(offsets) != 2 or not all(isinstance(value, int) for value in offsets) or offsets[0] < 0 or offsets[1] < offsets[0] ): fail(f"invalid data offsets for {name}") elements = 1 for dimension in shape: if not isinstance(dimension, int) or dimension < 0: fail(f"invalid shape for {name}") elements *= dimension tensor_bytes = elements * DTYPE_BYTES[dtype] if offsets[1] - offsets[0] != tensor_bytes: fail(f"tensor byte range mismatch for {name}") logical_bytes += tensor_bytes discovered[name] = shard_name if discovered != weight_map: missing = sorted(set(weight_map) - set(discovered)) extra = sorted(set(discovered) - set(weight_map)) fail(f"checkpoint index mismatch; missing={missing[:3]} extra={extra[:3]}") expected_size = int(index.get("metadata", {}).get("total_size", -1)) if logical_bytes != expected_size: fail(f"logical checkpoint size mismatch: {logical_bytes} != {expected_size}") if logical_bytes != quantization.get("state_payload_bytes"): fail("quantization state_payload_bytes does not match the checkpoint") counts = { "qdata": sum(name.endswith(".qdata") for name in weight_map), "weight_block_scale": sum( name.endswith(".weight_block_scale") for name in weight_map ), "weight_scale": sum(name.endswith(".weight_scale") for name in weight_map), "weight_fp8": sum(name.endswith(".weight_fp8") for name in weight_map), "weight_scale_storage": sum( name.endswith(".weight_scale_storage") for name in weight_map ), } if counts != { "qdata": 60, "weight_block_scale": 60, "weight_scale": 60, "weight_fp8": 120, "weight_scale_storage": 120, }: fail(f"unexpected packed tensor counts: {counts}") if counts != quantization.get("packed_tensor_counts"): fail("packed tensor counts disagree with quantization.json") return len(weight_map), len(shard_names), logical_bytes def validate_runtime() -> None: required = ( "codec.pth", "client.py", "launch.sh", "install.sh", "Dockerfile", "compose.yaml", "runtime/server.py", "runtime/web/index.html", "runtime/experimental/codec.py", "runtime/experimental/nvfp4/checkpoint.py", "runtime/experimental/nvfp4/modules.py", "runtime/experimental/fp8/modules.py", "runtime/native/smallm_gemv/smallm_gemv.cpp", "runtime/native/smallm_gemv/smallm_gemv.cu", "runtime/native/smallm_gemv/smallm_gemv.h", "runtime/native/LICENSE", "vendor/fish-speech/.project-root", "vendor/fish-speech/LICENSE", "vendor/fish-speech/pyproject.toml", "vendor/fish-speech/fish_speech/configs/modded_dac_vq.yaml", "vendor/fish-speech/tools/api_server.py", ) for relative in required: if not (ROOT / relative).is_file(): fail(f"missing required release file: {relative}") def validate_manifest() -> int: manifest = load_json("MANIFEST.json") records = manifest.get("files") if not isinstance(records, list) or not records: fail("MANIFEST.json has no file records") seen: set[str] = set() for record in records: relative = record.get("path") if not isinstance(relative, str): fail("manifest record has no path") pure = PurePosixPath(relative) if pure.is_absolute() or ".." in pure.parts or relative == "MANIFEST.json": fail(f"unsafe or recursive manifest path: {relative}") if relative in seen: fail(f"duplicate manifest path: {relative}") seen.add(relative) path = ROOT.joinpath(*pure.parts) if not path.is_file(): fail(f"manifest file is missing: {relative}") if path.stat().st_size != record.get("size"): fail(f"manifest size mismatch: {relative}") if sha256(path) != record.get("sha256"): fail(f"manifest hash mismatch: {relative}") return len(records) def main() -> None: quantization = validate_metadata() checkpoint_tensors, checkpoint_shards, logical_bytes = validate_checkpoint( quantization ) validate_runtime() manifest_files = validate_manifest() print( json.dumps( { "status": "pass", "release": "v1", "xpo3_release": False, "checkpoint_tensors": checkpoint_tensors, "checkpoint_shards": checkpoint_shards, "checkpoint_logical_bytes": logical_bytes, "nvfp4_projections": 60, "mxfp8_projections": 120, "manifest_files": manifest_files, "self_contained_weights": True, }, indent=2, ) ) if __name__ == "__main__": main()