ajh-code's picture
Add files using upload-large-folder tool
16f5171 verified
Raw
History Blame Contribute Delete
10.9 kB
"""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