| """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) |
| |
| |
| |
| 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]}") |
|
|
| |
| 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 |
|
|