"""Model-local MXFP8 projection modules for S2-Pro inference research.""" from __future__ import annotations import hashlib import json from dataclasses import asdict, dataclass from pathlib import Path import torch import torch.nn.functional as F from torch import nn from safetensors.torch import load_file import fish_scales_ops as fso from fish_speech.models.text2semantic.llama import ( BaseModelArgs, DualARTransformer, precompute_freqs_cis, ) from fish_speech.tokenizer import FishTokenizer @dataclass class ConversionRecord: name: str in_features: int out_features: int parameters: int probe_cosine: float class MXFP8Linear(nn.Module): """BF16-input linear using native 1x32 MXFP8 activation/weight GEMM.""" def __init__( self, weight_fp8: torch.Tensor, weight_scale_storage: torch.Tensor, *, in_features: int, out_features: int, ) -> None: super().__init__() self.in_features = in_features self.out_features = out_features self.register_buffer("weight_fp8", weight_fp8) # The SM120 kernel consumes K-major scales. Store their physical # layout as a contiguous [K-block, N] tensor so safetensors can # serialize it canonically; transpose restores the required view. self.register_buffer("weight_scale_storage", weight_scale_storage) @classmethod @torch.inference_mode() def from_linear(cls, linear: nn.Linear) -> "MXFP8Linear": if linear.bias is not None: raise ValueError("The initial S2-Pro MXFP8 path supports bias-free linears") if linear.weight.device.type != "cuda": raise ValueError("Quantize S2-Pro linears after moving them to CUDA") if linear.weight.dtype != torch.bfloat16: raise ValueError(f"Expected BF16 source weight, got {linear.weight.dtype}") weight_fp8, weight_scale = fso.gemm.quantize_1x32_fp8(linear.weight) return cls( weight_fp8, weight_scale.t().contiguous(), in_features=linear.in_features, out_features=linear.out_features, ) def forward(self, x: torch.Tensor) -> torch.Tensor: if x.shape[-1] != self.in_features: raise ValueError( f"Expected input width {self.in_features}, got {x.shape[-1]}" ) prefix = x.shape[:-1] x_2d = x.reshape(-1, self.in_features).contiguous() x_fp8, x_scale = fso.gemm.quantize_1x32_fp8(x_2d) output = fso.gemm.linear_mxfp8( x_fp8, self.weight_fp8, x_scale, self.weight_scale_storage.t(), ) return output.reshape(*prefix, self.out_features) def extra_repr(self) -> str: return ( f"in_features={self.in_features}, out_features={self.out_features}, " "weight=MXFP8_1x32, activation=dynamic_MXFP8_1x32, output=BF16" ) def _selected_slow_mlp(name: str, module: nn.Module) -> bool: return ( isinstance(module, nn.Linear) and name.startswith("layers.") and ".feed_forward." in name and name.rsplit(".", 1)[-1] in {"w1", "w2", "w3"} ) def _selected_slow_transformer(name: str, module: nn.Module) -> bool: return isinstance(module, nn.Linear) and name.startswith("layers.") def _selected_fast_transformer(name: str, module: nn.Module) -> bool: return isinstance(module, nn.Linear) and name.startswith("fast_layers.") def _selected_all_transformers(name: str, module: nn.Module) -> bool: return _selected_slow_transformer(name, module) or _selected_fast_transformer( name, module ) @torch.inference_mode() def convert_s2_pro_mxfp8( model: nn.Module, *, policy: str = "slow_mlp", probe_seed: int = 20260817, ) -> dict: """Replace selected S2-Pro projections without patching global linear APIs.""" selectors = { "slow_mlp": (_selected_slow_mlp, 108), "slow_transformer": (_selected_slow_transformer, 180), "fast_transformer": (_selected_fast_transformer, 20), "all_transformers": (_selected_all_transformers, 200), } if policy not in selectors: raise ValueError(f"Unsupported initial MXFP8 policy: {policy}") selector, expected_modules = selectors[policy] candidates = [ (name, module) for name, module in model.named_modules() if selector(name, module) ] if len(candidates) != expected_modules: raise RuntimeError( f"Expected {expected_modules} {policy} projections, found {len(candidates)}" ) records = [] generator = torch.Generator(device=candidates[0][1].weight.device) generator.manual_seed(probe_seed) for name, linear in candidates: parent_name, attribute = name.rsplit(".", 1) parent = model.get_submodule(parent_name) replacement = MXFP8Linear.from_linear(linear) probe = torch.randn( 1, linear.in_features, dtype=torch.bfloat16, device=linear.weight.device, generator=generator, ) * 0.1 reference = F.linear(probe, linear.weight) actual = replacement(probe) probe_cosine = float( F.cosine_similarity( actual.float().flatten(), reference.float().flatten(), dim=0 ).item() ) records.append( ConversionRecord( name=name, in_features=linear.in_features, out_features=linear.out_features, parameters=linear.weight.numel(), probe_cosine=probe_cosine, ) ) setattr(parent, attribute, replacement) torch.cuda.synchronize(candidates[0][1].weight.device) serialized = [asdict(record) for record in records] cosines = [record.probe_cosine for record in records] return { "policy": policy, "modules": len(records), "parameters": sum(record.parameters for record in records), "theoretical_bf16_source_bytes": sum( record.parameters * 2 for record in records ), "probe_cosine_min": min(cosines), "probe_cosine_mean": sum(cosines) / len(cosines), "probe_cosine_max": max(cosines), "records": serialized, } def _sha256(path: Path) -> str: digest = hashlib.sha256() with path.open("rb") as handle: while chunk := handle.read(8 * 1024 * 1024): digest.update(chunk) return digest.hexdigest() def _install_empty_mxfp8_modules(model: nn.Module, policy: str) -> int: selectors = { "slow_transformer": (_selected_slow_transformer, 180), } if policy not in selectors: raise ValueError(f"Unsupported artifact policy: {policy}") selector, expected = selectors[policy] names = [name for name, module in model.named_modules() if selector(name, module)] if len(names) != expected: raise RuntimeError(f"Expected {expected} artifact modules, found {len(names)}") for name in names: linear = model.get_submodule(name) parent_name, attribute = name.rsplit(".", 1) parent = model.get_submodule(parent_name) replacement = MXFP8Linear( torch.empty( linear.out_features, linear.in_features, dtype=torch.float8_e4m3fn, device="meta", ), torch.empty( linear.in_features // 128, linear.out_features, dtype=torch.int32, device="meta", ), in_features=linear.in_features, out_features=linear.out_features, ) setattr(parent, attribute, replacement) return len(names) @torch.inference_mode() def load_mxfp8_checkpoint( path: str | Path, *, device: str | torch.device = "cuda:0", max_length: int = 4096, verify_checksums: bool = False, ) -> DualARTransformer: """Load the canonical checkpoint without materializing BF16 FP8 sources.""" path = Path(path) metadata = json.loads((path / "quantization.json").read_text()) if metadata["format"] != "fish-s2-pro-project-local-mxfp8": raise ValueError(f"Unsupported checkpoint format: {metadata['format']}") if torch.cuda.get_device_capability(device)[0] != 12: raise RuntimeError("This MXFP8 artifact currently requires sm_120") if verify_checksums: for filename, record in metadata["checksums"].items(): file_path = path / filename if file_path.stat().st_size != record["bytes"]: raise RuntimeError(f"Size mismatch for {filename}") if _sha256(file_path) != record["sha256"]: raise RuntimeError(f"SHA256 mismatch for {filename}") config = BaseModelArgs.from_pretrained(str(path)) config.max_seq_len = max_length with torch.device("meta"): model = DualARTransformer(config) model.tokenizer = FishTokenizer.from_pretrained(path) _install_empty_mxfp8_modules(model, metadata["policy"]) index_path = path / "model.safetensors.index.json" if index_path.is_file(): index = json.loads(index_path.read_text()) shard_names = sorted(set(index["weight_map"].values())) else: shard_names = ["model.safetensors"] expected_keys = set(model.state_dict()) loaded_keys: set[str] = set() for shard_name in shard_names: shard = load_file(path / shard_name, device="cpu") unexpected = set(shard) - expected_keys if unexpected: raise RuntimeError( f"Unexpected checkpoint tensors in {shard_name}: {sorted(unexpected)[:5]}" ) model.load_state_dict(shard, strict=False, assign=True) loaded_keys.update(shard) missing = expected_keys - loaded_keys if missing: raise RuntimeError(f"Missing checkpoint tensors: {sorted(missing)[:5]}") # These buffers are intentionally non-persistent and were meta tensors. model.freqs_cis = precompute_freqs_cis( config.max_seq_len, config.head_dim, config.rope_base, ) model.causal_mask = torch.tril( torch.ones(config.max_seq_len, config.max_seq_len, dtype=torch.bool) ) model.fast_freqs_cis = precompute_freqs_cis( config.num_codebooks, config.fast_head_dim, config.rope_base, ) model = model.to(device=device).eval() model.fixed_temperature = torch.tensor(0.7, device=device, dtype=torch.float) model.fixed_top_p = torch.tensor(0.7, device=device, dtype=torch.float) model.fixed_repetition_penalty = torch.tensor(1.5, device=device, dtype=torch.float) model._cache_setup_done = False return model