Download runtime/experimental/nvfp4/modules.py from ajh-code/Fish-Audio-S2-Pro-NVFP4-Balanced: direct link, hf CLI and curl.
- Browser
- Download file 70.3 kB
-
https://huggingface.co/ajh-code/Fish-Audio-S2-Pro-NVFP4-Balanced/resolve/main/runtime/experimental/nvfp4/modules.py
- Command line
-
hf download hf://ajh-code/Fish-Audio-S2-Pro-NVFP4-Balanced/runtime/experimental/nvfp4/modules.py
-
curl -L -o modules.py https://huggingface.co/ajh-code/Fish-Audio-S2-Pro-NVFP4-Balanced/resolve/main/runtime/experimental/nvfp4/modules.py
70.3 kB
| """Portable native NVFP4 modules for the S2-Pro NVFP4 V1 runtime.""" | |
| from __future__ import annotations | |
| import hashlib | |
| import json | |
| import os | |
| import sys | |
| from dataclasses import asdict, dataclass | |
| from functools import lru_cache | |
| from pathlib import Path | |
| from typing import Any | |
| import numpy as np | |
| import torch | |
| import torch.nn.functional as F | |
| from torch import nn | |
| from torch.nn.attention import SDPBackend, sdpa_kernel | |
| from fish_speech.models.text2semantic.llama import apply_rotary_emb | |
| ROOT = Path( | |
| os.environ.get("FISH_NVFP4_ROOT", Path(__file__).resolve().parents[3]) | |
| ).resolve() | |
| NATIVE_SOURCE_ROOT = Path( | |
| os.environ.get( | |
| "FISH_NVFP4_NATIVE_ROOT", | |
| ROOT / "runtime" / "native", | |
| ) | |
| ).resolve() | |
| BUILD_ROOT = Path( | |
| os.environ.get( | |
| "FISH_NVFP4_BUILD_ROOT", | |
| ROOT / "runtime-data" / "torch-extensions", | |
| ) | |
| ).resolve() | |
| COMFY_KITCHEN_ROOT = os.environ.get("COMFY_KITCHEN_ROOT") | |
| DIRECT_SOURCE_ROOT = NATIVE_SOURCE_ROOT / "direct_w4a4_m1" | |
| SILU_PRODUCT_SOURCE_ROOT = NATIVE_SOURCE_ROOT / "silu_product_nvfp4_m1" | |
| RMSNORM_SOURCE_ROOT = NATIVE_SOURCE_ROOT / "rmsnorm_nvfp4_m1" | |
| SMALLM_SOURCE_ROOT = NATIVE_SOURCE_ROOT / "smallm_gemv" | |
| DIRECT_SOURCE_FILES = ( | |
| "direct_w4a4_m1.cpp", | |
| "direct_w4a4_m1.cu", | |
| "direct_w4a4_m1.h", | |
| ) | |
| SILU_PRODUCT_SOURCE_FILES = ( | |
| "silu_product_nvfp4_m1.cpp", | |
| "silu_product_nvfp4_m1.cu", | |
| "silu_product_nvfp4_m1.h", | |
| ) | |
| RMSNORM_SOURCE_FILES = ( | |
| "rmsnorm_nvfp4_m1.cpp", | |
| "rmsnorm_nvfp4_m1.cu", | |
| "rmsnorm_nvfp4_m1.h", | |
| ) | |
| SMALLM_SOURCE_FILES = ( | |
| "smallm_gemv.cpp", | |
| "smallm_gemv.cu", | |
| "smallm_gemv.h", | |
| ) | |
| if COMFY_KITCHEN_ROOT and COMFY_KITCHEN_ROOT not in sys.path: | |
| sys.path.insert(0, COMFY_KITCHEN_ROOT) | |
| from comfy_kitchen.tensor import QuantizedTensor, TensorCoreNVFP4Layout | |
| class ConversionRecord: | |
| name: str | |
| in_features: int | |
| out_features: int | |
| parameters: int | |
| probe_cosine: float | |
| class FusedMLPConversionRecord: | |
| name: str | |
| in_features: int | |
| intermediate_features: int | |
| out_features: int | |
| parameters: int | |
| probe_cosine: float | |
| class FusedTransformerConversionRecord: | |
| name: str | |
| parameters: int | |
| wqkv_probe_cosine: float | |
| wo_probe_cosine: float | |
| mlp_probe_cosine: float | |
| def quantize_nvfp4( | |
| value: torch.Tensor, | |
| *, | |
| scale: torch.Tensor | float | None = None, | |
| ) -> QuantizedTensor: | |
| """Apply the pinned tensor-wide dynamic NVFP4 policy.""" | |
| return QuantizedTensor.from_float( | |
| value, | |
| "TensorCoreNVFP4Layout", | |
| scale=scale, | |
| ) | |
| def apply_nvfp4_rounding_checkpoint( | |
| model: nn.Module, | |
| checkpoint: Path | str, | |
| ) -> dict[str, Any]: | |
| """Apply a packed learned-rounding delta after ordinary model conversion.""" | |
| from safetensors.torch import load_file | |
| checkpoint_path = Path(checkpoint) | |
| if checkpoint_path.is_dir(): | |
| candidates = sorted(checkpoint_path.glob("*.safetensors")) | |
| if len(candidates) != 1: | |
| raise ValueError( | |
| f"Expected one safetensors file in {checkpoint_path}, found {len(candidates)}" | |
| ) | |
| checkpoint_path = candidates[0] | |
| tensors = load_file(str(checkpoint_path), device=str(next(model.parameters()).device)) | |
| qdata_keys = sorted(key for key in tensors if key.endswith(".qdata")) | |
| if not qdata_keys: | |
| raise ValueError(f"No packed qdata entries in {checkpoint_path}") | |
| records = [] | |
| for qdata_key in qdata_keys: | |
| module_name = qdata_key.removesuffix(".qdata") | |
| module = model.get_submodule(module_name) | |
| if not isinstance(module, NVFP4Linear): | |
| raise TypeError(f"Checkpoint target is not NVFP4Linear: {module_name}") | |
| block_key = module_name + ".block_scale" | |
| tensor_key = module_name + ".tensor_scale" | |
| if block_key not in tensors or tensor_key not in tensors: | |
| raise ValueError(f"Checkpoint lacks scales for {module_name}") | |
| checkpoint_qdata = tensors[qdata_key] | |
| checkpoint_block_scale = tensors[block_key] | |
| checkpoint_tensor_scale = tensors[tensor_key] | |
| if checkpoint_qdata.shape != module.qdata.shape: | |
| raise ValueError(f"qdata shape mismatch for {module_name}") | |
| if not torch.equal(checkpoint_block_scale, module.weight_block_scale): | |
| raise ValueError(f"Block scale mismatch for {module_name}") | |
| if not torch.equal(checkpoint_tensor_scale, module.weight_scale): | |
| raise ValueError(f"Tensor scale mismatch for {module_name}") | |
| changed_bytes = int((checkpoint_qdata != module.qdata).sum()) | |
| old_codes = torch.stack( | |
| (module.qdata >> 4, module.qdata & 0x0F), | |
| dim=-1, | |
| ) | |
| new_codes = torch.stack( | |
| (checkpoint_qdata >> 4, checkpoint_qdata & 0x0F), | |
| dim=-1, | |
| ) | |
| changed_weights = int((old_codes != new_codes).sum()) | |
| module.qdata.copy_(checkpoint_qdata) | |
| records.append( | |
| { | |
| "module": module_name, | |
| "changed_packed_bytes": changed_bytes, | |
| "changed_weights": changed_weights, | |
| "weights": module.out_features * module.in_features, | |
| } | |
| ) | |
| return { | |
| "checkpoint": str(checkpoint_path), | |
| "modules": len(records), | |
| "changed_packed_bytes": sum(row["changed_packed_bytes"] for row in records), | |
| "changed_weights": sum(row["changed_weights"] for row in records), | |
| "records": records, | |
| } | |
| def _hadamard_blocks(value: torch.Tensor, block_size: int) -> torch.Tensor: | |
| """Materialized orthonormal block Hadamard for quality prototypes.""" | |
| if block_size < 2 or block_size & (block_size - 1): | |
| raise ValueError(f"Hadamard block size must be a power of two: {block_size}") | |
| if value.shape[-1] % block_size: | |
| raise ValueError( | |
| f"Width {value.shape[-1]} is not divisible by block size {block_size}" | |
| ) | |
| original_dtype = value.dtype | |
| transformed = value.float().reshape(*value.shape[:-1], -1, block_size) | |
| stride = 1 | |
| while stride < block_size: | |
| pairs = transformed.reshape( | |
| *transformed.shape[:-1], | |
| block_size // (2 * stride), | |
| 2, | |
| stride, | |
| ) | |
| left = pairs[..., 0, :] | |
| right = pairs[..., 1, :] | |
| transformed = torch.stack((left + right, left - right), dim=-2).reshape( | |
| *transformed.shape | |
| ) | |
| stride *= 2 | |
| return (transformed.reshape_as(value) / block_size**0.5).to(original_dtype) | |
| def _direct_source_hash() -> str: | |
| digest = hashlib.sha256() | |
| for filename in DIRECT_SOURCE_FILES: | |
| digest.update((DIRECT_SOURCE_ROOT / filename).read_bytes()) | |
| return digest.hexdigest() | |
| def _source_hash(source_root: Path, filenames: tuple[str, ...]) -> str: | |
| digest = hashlib.sha256() | |
| for filename in filenames: | |
| digest.update((source_root / filename).read_bytes()) | |
| return digest.hexdigest() | |
| def load_direct_w4a4_m1_extension(*, verbose: bool = False) -> Any: | |
| """Build the pinned direct packed-NVFP4 M=1 primitive locally.""" | |
| from torch.utils.cpp_extension import load | |
| missing = [ | |
| filename | |
| for filename in DIRECT_SOURCE_FILES | |
| if not (DIRECT_SOURCE_ROOT / filename).is_file() | |
| ] | |
| if missing: | |
| raise RuntimeError(f"Missing pinned NVFP4 native sources: {missing}") | |
| build_root = BUILD_ROOT / "direct_w4a4_m1" | |
| build_root.mkdir(parents=True, exist_ok=True) | |
| os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "12.0") | |
| os.environ.setdefault("MAX_JOBS", "2") | |
| return load( | |
| name=f"s2_pro_direct_w4a4_m1_{_direct_source_hash()[:12]}", | |
| sources=[ | |
| str(DIRECT_SOURCE_ROOT / "direct_w4a4_m1.cpp"), | |
| str(DIRECT_SOURCE_ROOT / "direct_w4a4_m1.cu"), | |
| ], | |
| extra_cflags=["-O3"], | |
| extra_cuda_cflags=["-O3", "--use_fast_math", "-lineinfo"], | |
| extra_include_paths=[str(DIRECT_SOURCE_ROOT)], | |
| build_directory=str(build_root), | |
| with_cuda=True, | |
| verbose=verbose, | |
| is_python_module=True, | |
| ) | |
| def load_silu_product_nvfp4_m1_extension(*, verbose: bool = False) -> Any: | |
| """Build the pinned fused SiLU/product-to-NVFP4 M=1 primitive.""" | |
| from torch.utils.cpp_extension import load | |
| missing = [ | |
| filename | |
| for filename in SILU_PRODUCT_SOURCE_FILES | |
| if not (SILU_PRODUCT_SOURCE_ROOT / filename).is_file() | |
| ] | |
| if missing: | |
| raise RuntimeError(f"Missing pinned NVFP4 SiLU sources: {missing}") | |
| build_root = BUILD_ROOT / "silu_product_nvfp4_m1" | |
| build_root.mkdir(parents=True, exist_ok=True) | |
| os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "12.0") | |
| os.environ.setdefault("MAX_JOBS", "2") | |
| return load( | |
| name=( | |
| "s2_pro_silu_product_nvfp4_m1_" | |
| f"{_source_hash(SILU_PRODUCT_SOURCE_ROOT, SILU_PRODUCT_SOURCE_FILES)[:12]}" | |
| ), | |
| sources=[ | |
| str(SILU_PRODUCT_SOURCE_ROOT / "silu_product_nvfp4_m1.cpp"), | |
| str(SILU_PRODUCT_SOURCE_ROOT / "silu_product_nvfp4_m1.cu"), | |
| ], | |
| extra_cflags=["-O3"], | |
| extra_cuda_cflags=["-O3", "--use_fast_math", "-lineinfo"], | |
| extra_include_paths=[str(SILU_PRODUCT_SOURCE_ROOT)], | |
| build_directory=str(build_root), | |
| with_cuda=True, | |
| verbose=verbose, | |
| is_python_module=True, | |
| ) | |
| def load_rmsnorm_nvfp4_m1_extension(*, verbose: bool = False) -> Any: | |
| """Build the pinned fused Fish-compatible RMSNorm-to-NVFP4 M=1 primitive.""" | |
| from torch.utils.cpp_extension import load | |
| missing = [ | |
| filename | |
| for filename in RMSNORM_SOURCE_FILES | |
| if not (RMSNORM_SOURCE_ROOT / filename).is_file() | |
| ] | |
| if missing: | |
| raise RuntimeError(f"Missing pinned NVFP4 RMSNorm sources: {missing}") | |
| build_root = BUILD_ROOT / "rmsnorm_nvfp4_m1" | |
| build_root.mkdir(parents=True, exist_ok=True) | |
| os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "12.0") | |
| os.environ.setdefault("MAX_JOBS", "2") | |
| return load( | |
| name=( | |
| "s2_pro_rmsnorm_nvfp4_m1_" | |
| f"{_source_hash(RMSNORM_SOURCE_ROOT, RMSNORM_SOURCE_FILES)[:12]}" | |
| ), | |
| sources=[ | |
| str(RMSNORM_SOURCE_ROOT / "rmsnorm_nvfp4_m1.cpp"), | |
| str(RMSNORM_SOURCE_ROOT / "rmsnorm_nvfp4_m1.cu"), | |
| ], | |
| extra_cflags=["-O3"], | |
| extra_cuda_cflags=["-O3", "--use_fast_math", "-lineinfo"], | |
| extra_include_paths=[str(RMSNORM_SOURCE_ROOT)], | |
| build_directory=str(build_root), | |
| with_cuda=True, | |
| verbose=verbose, | |
| is_python_module=True, | |
| ) | |
| def load_smallm_nvfp4_extension(*, verbose: bool = False) -> Any: | |
| """Build the pinned packed-weight W4A16 small-M GEMV primitive.""" | |
| from torch.utils.cpp_extension import load | |
| missing = [ | |
| filename | |
| for filename in SMALLM_SOURCE_FILES | |
| if not (SMALLM_SOURCE_ROOT / filename).is_file() | |
| ] | |
| if missing: | |
| raise RuntimeError(f"Missing pinned NVFP4 small-M sources: {missing}") | |
| build_root = BUILD_ROOT / "smallm_gemv" | |
| build_root.mkdir(parents=True, exist_ok=True) | |
| os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "12.0") | |
| os.environ.setdefault("MAX_JOBS", "2") | |
| return load( | |
| name=( | |
| "s2_pro_smallm_nvfp4_" | |
| f"{_source_hash(SMALLM_SOURCE_ROOT, SMALLM_SOURCE_FILES)[:12]}" | |
| ), | |
| sources=[ | |
| str(SMALLM_SOURCE_ROOT / "smallm_gemv.cpp"), | |
| str(SMALLM_SOURCE_ROOT / "smallm_gemv.cu"), | |
| ], | |
| extra_cflags=["-O3"], | |
| extra_cuda_cflags=["-O3", "--use_fast_math", "-lineinfo"], | |
| extra_include_paths=[str(SMALLM_SOURCE_ROOT)], | |
| build_directory=str(build_root), | |
| with_cuda=True, | |
| verbose=verbose, | |
| is_python_module=True, | |
| ) | |
| class NVFP4Linear(nn.Module): | |
| """BF16-input linear using a native packed NVFP4 weight with BF16 output.""" | |
| def __init__( | |
| self, | |
| qdata: torch.Tensor, | |
| weight_block_scale: torch.Tensor, | |
| weight_scale: torch.Tensor, | |
| *, | |
| in_features: int, | |
| out_features: int, | |
| direct_m1: bool = False, | |
| w4a16_max_m: int = 0, | |
| correction_down: torch.Tensor | None = None, | |
| correction_up: torch.Tensor | None = None, | |
| sparse_correction_indices: torch.Tensor | None = None, | |
| sparse_correction_weight: torch.Tensor | None = None, | |
| input_hadamard_block_size: int = 0, | |
| ) -> None: | |
| super().__init__() | |
| if w4a16_max_m < 0: | |
| raise ValueError("w4a16_max_m must be nonnegative") | |
| if direct_m1 and w4a16_max_m: | |
| raise ValueError("Choose only one NVFP4 M=1 backend") | |
| self.in_features = int(in_features) | |
| self.out_features = int(out_features) | |
| self.direct_m1 = bool(direct_m1) | |
| self.w4a16_max_m = int(w4a16_max_m) | |
| if input_hadamard_block_size and ( | |
| input_hadamard_block_size < 2 | |
| or input_hadamard_block_size & (input_hadamard_block_size - 1) | |
| or self.in_features % input_hadamard_block_size | |
| ): | |
| raise ValueError( | |
| "Input Hadamard block size must be a power of two that divides " | |
| f"the input width, got {input_hadamard_block_size}" | |
| ) | |
| self.input_hadamard_block_size = int(input_hadamard_block_size) | |
| self.register_buffer("qdata", qdata) | |
| self.register_buffer("weight_block_scale", weight_block_scale) | |
| self.register_buffer("weight_scale", weight_scale) | |
| if (correction_down is None) != (correction_up is None): | |
| raise ValueError("Low-rank correction requires both down and up factors") | |
| if correction_down is not None: | |
| if correction_down.dim() != 2 or correction_up.dim() != 2: | |
| raise ValueError("Low-rank correction factors must be 2D") | |
| if correction_down.shape[1] != self.in_features: | |
| raise ValueError("Low-rank down factor input width does not match") | |
| if correction_up.shape != (self.out_features, correction_down.shape[0]): | |
| raise ValueError("Low-rank up factor shape does not match") | |
| self.register_buffer("correction_down", correction_down) | |
| self.register_buffer("correction_up", correction_up) | |
| if (sparse_correction_indices is None) != (sparse_correction_weight is None): | |
| raise ValueError( | |
| "Sparse correction requires both channel indices and a weight" | |
| ) | |
| if sparse_correction_indices is not None: | |
| if ( | |
| sparse_correction_indices.dim() != 1 | |
| or sparse_correction_indices.dtype != torch.int64 | |
| ): | |
| raise ValueError("Sparse correction indices must be 1D int64") | |
| if sparse_correction_weight.shape != ( | |
| self.out_features, | |
| sparse_correction_indices.numel(), | |
| ): | |
| raise ValueError("Sparse correction weight shape does not match") | |
| if sparse_correction_weight.dtype != torch.bfloat16: | |
| raise ValueError("Sparse correction weight must be BF16") | |
| if sparse_correction_indices.numel() and ( | |
| int(sparse_correction_indices.min()) < 0 | |
| or int(sparse_correction_indices.max()) >= self.in_features | |
| ): | |
| raise ValueError("Sparse correction channel index is out of range") | |
| self.register_buffer("sparse_correction_indices", sparse_correction_indices) | |
| self.register_buffer("sparse_correction_weight", sparse_correction_weight) | |
| def from_weight( | |
| cls, | |
| weight: torch.Tensor, | |
| *, | |
| direct_m1: bool = False, | |
| w4a16_max_m: int = 0, | |
| quantization_scale: torch.Tensor | float | None = None, | |
| correction_down: torch.Tensor | None = None, | |
| correction_up: torch.Tensor | None = None, | |
| sparse_correction_indices: torch.Tensor | None = None, | |
| sparse_correction_weight: torch.Tensor | None = None, | |
| input_hadamard_block_size: int = 0, | |
| ) -> "NVFP4Linear": | |
| if weight.device.type != "cuda": | |
| raise ValueError("Quantize S2-Pro weights after moving them to CUDA") | |
| if weight.dtype != torch.bfloat16 or weight.dim() != 2: | |
| raise ValueError(f"Expected a 2D BF16 source weight, got {weight.dtype} {weight.shape}") | |
| packed = quantize_nvfp4( | |
| weight.contiguous(), | |
| scale=quantization_scale, | |
| ) | |
| return cls( | |
| packed._qdata, | |
| packed._params.block_scale, | |
| packed._params.scale, | |
| in_features=weight.shape[1], | |
| out_features=weight.shape[0], | |
| direct_m1=direct_m1, | |
| w4a16_max_m=w4a16_max_m, | |
| correction_down=correction_down, | |
| correction_up=correction_up, | |
| sparse_correction_indices=sparse_correction_indices, | |
| sparse_correction_weight=sparse_correction_weight, | |
| input_hadamard_block_size=input_hadamard_block_size, | |
| ) | |
| def from_linear( | |
| cls, | |
| linear: nn.Linear, | |
| *, | |
| direct_m1: bool = False, | |
| w4a16_max_m: int = 0, | |
| ) -> "NVFP4Linear": | |
| if linear.bias is not None: | |
| raise ValueError("The initial S2-Pro NVFP4 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}") | |
| return cls.from_weight( | |
| linear.weight, | |
| direct_m1=direct_m1, | |
| w4a16_max_m=w4a16_max_m, | |
| ) | |
| def _weight_quantized_tensor(self) -> QuantizedTensor: | |
| params = TensorCoreNVFP4Layout.Params( | |
| scale=self.weight_scale, | |
| orig_dtype=torch.bfloat16, | |
| orig_shape=(self.out_features, self.in_features), | |
| block_scale=self.weight_block_scale, | |
| ) | |
| return QuantizedTensor(self.qdata, "TensorCoreNVFP4Layout", params) | |
| def _direct(self, activation: QuantizedTensor) -> torch.Tensor: | |
| params = activation._params | |
| return load_direct_w4a4_m1_extension().linear( | |
| activation._qdata, | |
| params.block_scale, | |
| params.scale, | |
| self.qdata, | |
| self.weight_block_scale, | |
| self.weight_scale, | |
| None, | |
| ) | |
| def project_packed_m1( | |
| self, | |
| qdata: torch.Tensor, | |
| block_scale: torch.Tensor, | |
| tensor_scale: torch.Tensor, | |
| ) -> torch.Tensor: | |
| """Project an already packed logical M=1 activation without wrappers.""" | |
| return load_direct_w4a4_m1_extension().linear( | |
| qdata, | |
| block_scale, | |
| tensor_scale, | |
| self.qdata, | |
| self.weight_block_scale, | |
| self.weight_scale, | |
| None, | |
| ) | |
| def project_quantized(self, activation: QuantizedTensor) -> torch.Tensor: | |
| logical_m = int(activation._params.orig_shape[0]) | |
| if self.direct_m1 and logical_m == 1: | |
| return self._direct(activation) | |
| return F.linear(activation, self._weight_quantized_tensor(), None) | |
| def forward(self, value: torch.Tensor) -> torch.Tensor: | |
| if value.shape[-1] != self.in_features: | |
| raise ValueError( | |
| f"Expected input width {self.in_features}, got {value.shape[-1]}" | |
| ) | |
| input_shape = tuple(value.shape) | |
| correction_input = value.reshape(-1, self.in_features).contiguous() | |
| flattened = correction_input | |
| if self.input_hadamard_block_size: | |
| flattened = _hadamard_blocks( | |
| flattened, | |
| self.input_hadamard_block_size, | |
| ).contiguous() | |
| if 0 < flattened.shape[0] <= self.w4a16_max_m: | |
| output = load_smallm_nvfp4_extension().linear( | |
| flattened, | |
| self.qdata, | |
| self.weight_block_scale, | |
| self.weight_scale, | |
| None, | |
| ) | |
| else: | |
| activation = quantize_nvfp4(flattened) | |
| output = self.project_quantized(activation) | |
| if self.correction_down is not None: | |
| correction = F.linear( | |
| F.linear(correction_input, self.correction_down), | |
| self.correction_up, | |
| ) | |
| output = output + correction | |
| if self.sparse_correction_indices is not None: | |
| selected = correction_input.index_select( | |
| 1, | |
| self.sparse_correction_indices, | |
| ) | |
| output = output + F.linear(selected, self.sparse_correction_weight) | |
| return output.reshape(*input_shape[:-1], self.out_features) | |
| def extra_repr(self) -> str: | |
| backend = ( | |
| f"w4a16_through_m{self.w4a16_max_m}+w4a4_above_threshold" | |
| if self.w4a16_max_m | |
| else ("direct_m1+tensorcore" if self.direct_m1 else "tensorcore") | |
| ) | |
| return ( | |
| f"in_features={self.in_features}, out_features={self.out_features}, " | |
| f"weight=NVFP4_E2M1, activation=dynamic_NVFP4, output=BF16, backend={backend}, " | |
| f"correction_rank={0 if self.correction_down is None else self.correction_down.shape[0]}, " | |
| f"sparse_correction_channels=" | |
| f"{0 if self.sparse_correction_indices is None else self.sparse_correction_indices.numel()}, " | |
| f"input_hadamard_block_size={self.input_hadamard_block_size}" | |
| ) | |
| class NVFP4FeedForward(nn.Module): | |
| """S2 SwiGLU with shared input packing and fused product packing at M=1.""" | |
| def __init__( | |
| self, | |
| w1: NVFP4Linear, | |
| w2: NVFP4Linear, | |
| w3: NVFP4Linear, | |
| ) -> None: | |
| super().__init__() | |
| self.w1 = w1 | |
| self.w2 = w2 | |
| self.w3 = w3 | |
| self.decode_backend = "w4a16" if w1.w4a16_max_m else "w4a4" | |
| def from_module( | |
| cls, | |
| module: nn.Module, | |
| *, | |
| decode_backend: str = "w4a4", | |
| w4a16_max_m: int = 1, | |
| ) -> "NVFP4FeedForward": | |
| if decode_backend not in {"w4a4", "w4a16"}: | |
| raise ValueError(f"Unsupported NVFP4 MLP decode backend: {decode_backend}") | |
| for name in ("w1", "w2", "w3"): | |
| if not isinstance(getattr(module, name, None), nn.Linear): | |
| raise TypeError(f"Expected BF16 FeedForward.{name} linear") | |
| result = cls( | |
| NVFP4Linear.from_linear( | |
| module.w1, | |
| direct_m1=decode_backend == "w4a4", | |
| w4a16_max_m=(w4a16_max_m if decode_backend == "w4a16" else 0), | |
| ), | |
| NVFP4Linear.from_linear( | |
| module.w2, | |
| direct_m1=decode_backend == "w4a4", | |
| w4a16_max_m=(w4a16_max_m if decode_backend == "w4a16" else 0), | |
| ), | |
| NVFP4Linear.from_linear( | |
| module.w3, | |
| direct_m1=decode_backend == "w4a4", | |
| w4a16_max_m=(w4a16_max_m if decode_backend == "w4a16" else 0), | |
| ), | |
| ) | |
| result.decode_backend = decode_backend | |
| return result | |
| def _forward_m1(self, flattened: torch.Tensor) -> torch.Tensor: | |
| if self.decode_backend == "w4a16": | |
| return self.w2(F.silu(self.w1(flattened)) * self.w3(flattened)) | |
| packer = load_silu_product_nvfp4_m1_extension() | |
| qdata, block_scale, tensor_scale = packer.quantize_input(flattened) | |
| return self.forward_packed_m1(qdata, block_scale, tensor_scale) | |
| def forward_packed_m1( | |
| self, | |
| qdata: torch.Tensor, | |
| block_scale: torch.Tensor, | |
| tensor_scale: torch.Tensor, | |
| ) -> torch.Tensor: | |
| """Consume an already packed normalized M=1 activation.""" | |
| packer = load_silu_product_nvfp4_m1_extension() | |
| gate = self.w1.project_packed_m1(qdata, block_scale, tensor_scale) | |
| up = self.w3.project_packed_m1(qdata, block_scale, tensor_scale) | |
| product_qdata, product_block_scale, product_tensor_scale = packer.quantize( | |
| gate, up | |
| ) | |
| return self.w2.project_packed_m1( | |
| product_qdata, | |
| product_block_scale, | |
| product_tensor_scale, | |
| ) | |
| def _forward_tensorcore(self, flattened: torch.Tensor) -> torch.Tensor: | |
| activation = quantize_nvfp4(flattened) | |
| gate = self.w1.project_quantized(activation) | |
| up = self.w3.project_quantized(activation) | |
| return self.w2(F.silu(gate) * up) | |
| def forward(self, value: torch.Tensor) -> torch.Tensor: | |
| input_shape = tuple(value.shape) | |
| flattened = value.reshape(-1, input_shape[-1]).contiguous() | |
| if ( | |
| self.decode_backend == "w4a16" | |
| and flattened.shape[0] <= self.w1.w4a16_max_m | |
| ): | |
| output = self.w2(F.silu(self.w1(flattened)) * self.w3(flattened)) | |
| elif flattened.shape[0] == 1: | |
| output = self._forward_m1(flattened) | |
| else: | |
| output = self._forward_tensorcore(flattened) | |
| return output.reshape(*input_shape[:-1], self.w2.out_features) | |
| def extra_repr(self) -> str: | |
| return ( | |
| f"decode_backend={self.decode_backend}, " | |
| f"fused_silu_product_pack_m1={self.decode_backend == 'w4a4'}" | |
| ) | |
| class NVFP4TransformerBlock(nn.Module): | |
| """Slow S2 block with fused norm/MLP packing for autoregressive M=1.""" | |
| def __init__(self, source: nn.Module) -> None: | |
| super().__init__() | |
| source.attention.wqkv = NVFP4Linear.from_linear( | |
| source.attention.wqkv, direct_m1=True | |
| ) | |
| source.attention.wo = NVFP4Linear.from_linear( | |
| source.attention.wo, direct_m1=True | |
| ) | |
| self.attention = source.attention | |
| self.feed_forward = NVFP4FeedForward.from_module(source.feed_forward) | |
| self.ffn_norm = source.ffn_norm | |
| self.attention_norm = source.attention_norm | |
| self.train(source.training) | |
| def _packed_norm( | |
| self, | |
| value: torch.Tensor, | |
| norm: nn.Module, | |
| ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: | |
| flattened = value.reshape(-1, value.shape[-1]).contiguous() | |
| return tuple( | |
| load_rmsnorm_nvfp4_m1_extension().quantize( | |
| flattened, | |
| norm.weight.contiguous(), | |
| float(norm.eps), | |
| ) | |
| ) | |
| def _attention_m1( | |
| self, | |
| x: torch.Tensor, | |
| packed_norm: tuple[torch.Tensor, torch.Tensor, torch.Tensor], | |
| freqs_cis: torch.Tensor, | |
| mask: torch.Tensor | None, | |
| input_pos: torch.Tensor | None, | |
| ) -> torch.Tensor: | |
| attention = self.attention | |
| bsz, seqlen, _ = x.shape | |
| qkv = attention.wqkv.project_packed_m1(*packed_norm) | |
| q_size = attention.n_head * attention.head_dim | |
| kv_size = attention.n_local_heads * attention.head_dim | |
| q, k, v = qkv.split([q_size, kv_size, kv_size], dim=-1) | |
| q = q.view(bsz, seqlen, attention.n_head, attention.head_dim) | |
| k = k.view(bsz, seqlen, attention.n_local_heads, attention.head_dim) | |
| v = v.view(bsz, seqlen, attention.n_local_heads, attention.head_dim) | |
| if attention.attention_qk_norm: | |
| q = attention.q_norm(q) | |
| k = attention.k_norm(k) | |
| q = apply_rotary_emb(q, freqs_cis) | |
| k = apply_rotary_emb(k, freqs_cis) | |
| q, k, v = (item.transpose(1, 2) for item in (q, k, v)) | |
| if attention.kv_cache is not None: | |
| k, v = attention.kv_cache.update(input_pos, k, v) | |
| repeat = attention.n_head // attention.n_local_heads | |
| k = k.repeat_interleave(repeat, dim=1) | |
| v = v.repeat_interleave(repeat, dim=1) | |
| if attention.use_sdpa: | |
| if mask is None: | |
| with sdpa_kernel(SDPBackend.FLASH_ATTENTION): | |
| y = F.scaled_dot_product_attention( | |
| q, | |
| k, | |
| v, | |
| dropout_p=attention.dropout if attention.training else 0.0, | |
| is_causal=True, | |
| ) | |
| else: | |
| y = F.scaled_dot_product_attention( | |
| q, | |
| k, | |
| v, | |
| attn_mask=mask, | |
| dropout_p=attention.dropout if attention.training else 0.0, | |
| ) | |
| else: | |
| y = attention.eq_scaled_dot_product_attention(q, k, v, attn_mask=mask) | |
| y = y.transpose(1, 2).contiguous().view(bsz, seqlen, q_size) | |
| y_flat = y.reshape(-1, q_size) | |
| output_pack = load_silu_product_nvfp4_m1_extension().quantize_input(y_flat) | |
| return attention.wo.project_packed_m1(*output_pack).reshape( | |
| bsz, seqlen, attention.dim | |
| ) | |
| def forward( | |
| self, | |
| x: torch.Tensor, | |
| freqs_cis: torch.Tensor, | |
| mask: torch.Tensor, | |
| input_pos: torch.Tensor | None = None, | |
| ) -> torch.Tensor: | |
| if x.reshape(-1, x.shape[-1]).shape[0] != 1: | |
| h = x + self.attention( | |
| self.attention_norm(x), freqs_cis, mask, input_pos | |
| ) | |
| return h + self.feed_forward(self.ffn_norm(h)) | |
| attention_pack = self._packed_norm(x, self.attention_norm) | |
| h = x + self._attention_m1(x, attention_pack, freqs_cis, mask, input_pos) | |
| ffn_pack = self._packed_norm(h, self.ffn_norm) | |
| feed_forward = self.feed_forward.forward_packed_m1(*ffn_pack) | |
| return h + feed_forward.reshape_as(h) | |
| def extra_repr(self) -> str: | |
| return "fused_rmsnorm_pack_m1=True, fused_mlp_pack_m1=True" | |
| 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_slow_attention(name: str, module: nn.Module) -> bool: | |
| return ( | |
| isinstance(module, nn.Linear) | |
| and name.startswith("layers.") | |
| and ".attention." in name | |
| and name.rsplit(".", 1)[-1] in {"wqkv", "wo"} | |
| ) | |
| def _selected_slow_mlp_block(name: str, module: nn.Module) -> bool: | |
| return ( | |
| name.startswith("layers.") | |
| and name.endswith(".feed_forward") | |
| and all(isinstance(getattr(module, role, None), nn.Linear) for role in ("w1", "w2", "w3")) | |
| ) | |
| def _selected_slow_transformer_block(name: str, module: nn.Module) -> bool: | |
| parts = name.split(".") | |
| return ( | |
| len(parts) == 2 | |
| and parts[0] == "layers" | |
| and parts[1].isdigit() | |
| and hasattr(module, "attention") | |
| and hasattr(module, "feed_forward") | |
| and hasattr(module, "attention_norm") | |
| and hasattr(module, "ffn_norm") | |
| ) | |
| def _packed_weight_bytes(model: nn.Module) -> int: | |
| return sum( | |
| int(module.qdata.numel() * module.qdata.element_size()) | |
| + int(module.weight_block_scale.numel() * module.weight_block_scale.element_size()) | |
| + int(module.weight_scale.numel() * module.weight_scale.element_size()) | |
| for module in model.modules() | |
| if isinstance(module, NVFP4Linear) | |
| ) | |
| def _resolve_workspace_path(value: str | Path, *, relative_to: Path) -> Path: | |
| path = Path(value) | |
| if path.is_absolute(): | |
| return path | |
| relative_candidate = relative_to / path | |
| if relative_candidate.exists(): | |
| return relative_candidate | |
| return ROOT / path | |
| def _clip_weight_blocks(weight: torch.Tensor, ratio: float) -> torch.Tensor: | |
| if ratio == 1.0: | |
| return weight | |
| if not 0 < ratio <= 1: | |
| raise ValueError(f"NVFP4 clip ratio must be in (0,1], got {ratio}") | |
| if weight.shape[1] % 128: | |
| raise ValueError(f"Input width must be divisible by 128, got {weight.shape}") | |
| blocks = weight.float().reshape(weight.shape[0], -1, 128) | |
| threshold = blocks.abs().amax(dim=-1, keepdim=True) * ratio | |
| return torch.clamp(blocks, min=-threshold, max=threshold).reshape_as(weight) | |
| def _prepare_activation_scaling( | |
| model: nn.Module, | |
| report_path: Path, | |
| selected_layers: set[int], | |
| *, | |
| fold_norms: bool = True, | |
| ) -> tuple[ | |
| dict[int, torch.Tensor], | |
| dict[int, float], | |
| dict[int, dict[str, float]], | |
| dict[int, dict[str, torch.Tensor]], | |
| dict[str, Any], | |
| ]: | |
| sweep = json.loads(report_path.read_text()) | |
| calibration_values = sweep.get("calibration_reports") | |
| if calibration_values is None: | |
| calibration_values = [sweep["calibration_report"]] | |
| calibration_report_paths = [ | |
| _resolve_workspace_path(value, relative_to=report_path.parent) | |
| for value in calibration_values | |
| ] | |
| calibrations = [] | |
| activation_sets = [] | |
| for calibration_report_path in calibration_report_paths: | |
| calibration = json.loads(calibration_report_path.read_text()) | |
| with np.load(calibration_report_path.parent / calibration["arrays"]) as values: | |
| activation_sets.append(values["activations"].copy()) | |
| calibrations.append(calibration) | |
| activations = np.concatenate(activation_sets, axis=1) | |
| rows = {int(row["layer"]): row for row in sweep["results"]} | |
| missing = sorted(selected_layers - rows.keys()) | |
| if missing: | |
| raise ValueError(f"Activation-scaling report lacks layers: {missing}") | |
| clamp = float(sweep["scale_clamp"]) | |
| activation_statistic = sweep.get("activation_statistic", "absmax") | |
| scales: dict[int, torch.Tensor] = {} | |
| clip_ratios: dict[int, float] = {} | |
| quantization_scales: dict[int, dict[str, float]] = {} | |
| corrections: dict[int, dict[str, torch.Tensor]] = {} | |
| correction_tensors = None | |
| correction_tensors_value = sweep.get("correction_tensors") | |
| if correction_tensors_value is not None: | |
| from safetensors.torch import load_file | |
| correction_path = _resolve_workspace_path( | |
| correction_tensors_value, | |
| relative_to=report_path.parent, | |
| ) | |
| correction_tensors = load_file( | |
| correction_path, | |
| device=str(model.layers[0].ffn_norm.weight.device), | |
| ) | |
| records = [] | |
| for layer_index in sorted(selected_layers): | |
| block = model.layers[layer_index] | |
| gate = block.feed_forward.w1.weight | |
| up = block.feed_forward.w3.weight | |
| row = rows[layer_index] | |
| improvement = float(row["improvement_fraction_vs_unscaled"]) | |
| explicit_channel_scale = row.get("channel_scale") | |
| if explicit_channel_scale is not None: | |
| if len(explicit_channel_scale) != gate.shape[1]: | |
| raise ValueError( | |
| f"Layer {layer_index} channel scale has " | |
| f"{len(explicit_channel_scale)} values, expected {gate.shape[1]}" | |
| ) | |
| scale = torch.tensor( | |
| explicit_channel_scale, | |
| device=gate.device, | |
| dtype=torch.float32, | |
| ) | |
| if not torch.isfinite(scale).all() or (scale <= 0).any(): | |
| raise ValueError( | |
| f"Layer {layer_index} channel scale must be finite and positive" | |
| ) | |
| alpha = None | |
| clip_ratio = float(row.get("best_clip_ratio", 1.0)) | |
| elif improvement <= 0: | |
| scale = torch.ones(gate.shape[1], device=gate.device, dtype=torch.float32) | |
| alpha = None | |
| clip_ratio = 1.0 | |
| else: | |
| alpha = float(row["best_alpha"]) | |
| clip_ratio = float(row.get("best_clip_ratio", 1.0)) | |
| activation = torch.from_numpy(activations[layer_index]).to( | |
| device=gate.device, | |
| dtype=torch.float32, | |
| ) | |
| if activation_statistic == "absmax": | |
| activation_scale = activation.abs().amax(dim=0) | |
| elif activation_statistic == "abs_p99": | |
| activation_scale = torch.quantile(activation.abs(), 0.99, dim=0) | |
| elif activation_statistic == "abs_p999": | |
| activation_scale = torch.quantile(activation.abs(), 0.999, dim=0) | |
| elif activation_statistic == "mean_abs": | |
| activation_scale = activation.abs().mean(dim=0) | |
| elif activation_statistic == "rms": | |
| activation_scale = activation.square().mean(dim=0).sqrt() | |
| else: | |
| raise ValueError( | |
| f"Unknown activation statistic {activation_statistic!r}" | |
| ) | |
| activation_scale = activation_scale.clamp_min(1e-6) | |
| weight_max = torch.maximum( | |
| gate.float().abs().amax(dim=0), | |
| up.float().abs().amax(dim=0), | |
| ).clamp_min(1e-6) | |
| scale = activation_scale.pow(alpha) / weight_max.pow(1.0 - alpha) | |
| scale = scale / torch.exp(torch.mean(torch.log(scale))) | |
| scale = scale.clamp(min=1.0 / clamp, max=clamp) | |
| role_quantization_scales = row.get("weight_tensor_scales") | |
| if role_quantization_scales is not None: | |
| parsed_role_scales = { | |
| role: float(role_quantization_scales[role]) | |
| for role in ("w1", "w3") | |
| } | |
| if any( | |
| not np.isfinite(value) or value <= 0 | |
| for value in parsed_role_scales.values() | |
| ): | |
| raise ValueError( | |
| f"Layer {layer_index} weight tensor scales must be finite and positive" | |
| ) | |
| quantization_scales[layer_index] = parsed_role_scales | |
| low_rank = row.get("low_rank_correction") | |
| if low_rank is not None: | |
| if correction_tensors is None: | |
| raise ValueError( | |
| f"Layer {layer_index} has a low-rank correction without tensors" | |
| ) | |
| down_key = low_rank["down_key"] | |
| up_keys = low_rank["up_keys"] | |
| try: | |
| correction = { | |
| "down": correction_tensors[down_key], | |
| "w1_up": correction_tensors[up_keys["w1"]], | |
| "w3_up": correction_tensors[up_keys["w3"]], | |
| } | |
| except KeyError as error: | |
| raise ValueError( | |
| f"Layer {layer_index} correction tensor is missing: {error}" | |
| ) from error | |
| if any(value.dtype != torch.bfloat16 for value in correction.values()): | |
| raise ValueError( | |
| f"Layer {layer_index} correction tensors must be BF16" | |
| ) | |
| corrections[layer_index] = correction | |
| sparse = row.get("sparse_channel_correction") | |
| if sparse is not None: | |
| if correction_tensors is None: | |
| raise ValueError( | |
| f"Layer {layer_index} has a sparse correction without tensors" | |
| ) | |
| try: | |
| sparse_correction = { | |
| "sparse_indices": correction_tensors[sparse["indices_key"]], | |
| "w1_sparse_weight": correction_tensors[ | |
| sparse["weight_keys"]["w1"] | |
| ], | |
| "w3_sparse_weight": correction_tensors[ | |
| sparse["weight_keys"]["w3"] | |
| ], | |
| } | |
| except KeyError as error: | |
| raise ValueError( | |
| f"Layer {layer_index} sparse correction tensor is missing: {error}" | |
| ) from error | |
| if sparse_correction["sparse_indices"].dtype != torch.int64: | |
| raise ValueError( | |
| f"Layer {layer_index} sparse correction indices must be int64" | |
| ) | |
| if any( | |
| sparse_correction[key].dtype != torch.bfloat16 | |
| for key in ("w1_sparse_weight", "w3_sparse_weight") | |
| ): | |
| raise ValueError( | |
| f"Layer {layer_index} sparse correction weights must be BF16" | |
| ) | |
| corrections.setdefault(layer_index, {}).update(sparse_correction) | |
| input_hadamard_block_size = int(row.get("input_hadamard_block_size", 0)) | |
| if input_hadamard_block_size and ( | |
| input_hadamard_block_size < 2 | |
| or input_hadamard_block_size & (input_hadamard_block_size - 1) | |
| or gate.shape[1] % input_hadamard_block_size | |
| ): | |
| raise ValueError( | |
| f"Layer {layer_index} has invalid input Hadamard block size " | |
| f"{input_hadamard_block_size}" | |
| ) | |
| if fold_norms: | |
| norm = block.ffn_norm.weight | |
| norm.data.copy_((norm.float() / scale).to(dtype=norm.dtype)) | |
| scales[layer_index] = scale | |
| clip_ratios[layer_index] = clip_ratio | |
| records.append( | |
| { | |
| "layer": layer_index, | |
| "alpha": alpha, | |
| "calibration_improvement_fraction": improvement, | |
| "clip_ratio": clip_ratio, | |
| "scale_mode": ( | |
| "explicit_channel" | |
| if explicit_channel_scale is not None | |
| else "activation_formula" | |
| ), | |
| "correction_rank": ( | |
| None if low_rank is None else int(low_rank["rank"]) | |
| ), | |
| "sparse_correction_channels": ( | |
| None if sparse is None else int(sparse["channels"]) | |
| ), | |
| "input_hadamard_block_size": input_hadamard_block_size, | |
| "scale_min": float(scale.min()), | |
| "scale_p50": float(torch.quantile(scale, 0.50)), | |
| "scale_max": float(scale.max()), | |
| } | |
| ) | |
| return scales, clip_ratios, quantization_scales, corrections, { | |
| "sweep_report": str(report_path), | |
| "calibration_reports": [str(path) for path in calibration_report_paths], | |
| "calibration_history_sha256": [ | |
| calibration["history_sha256"] for calibration in calibrations | |
| ], | |
| "scale_clamp": clamp, | |
| "activation_statistic": activation_statistic, | |
| "correction_tensors": correction_tensors_value, | |
| "layers": records, | |
| } | |
| def convert_s2_pro_nvfp4( | |
| model: nn.Module, | |
| *, | |
| policy: str = "slow_mlp", | |
| direct_m1: bool = False, | |
| w4a16_max_m: int = 1, | |
| nvfp4_layers: set[int] | None = None, | |
| activation_scaling_report: Path | None = None, | |
| probe_seed: int = 20260818, | |
| ) -> dict[str, Any]: | |
| """Replace selected S2-Pro projections without modifying the FP8 path.""" | |
| if w4a16_max_m < 1: | |
| raise ValueError("w4a16_max_m must be positive") | |
| selectors = { | |
| "slow_mlp": (_selected_slow_mlp, 108), | |
| "slow_transformer": (_selected_slow_transformer, 180), | |
| } | |
| if policy == "w4a16_slow_transformer": | |
| candidates = [ | |
| (name, module) | |
| for name, module in model.named_modules() | |
| if _selected_slow_transformer(name, module) | |
| ] | |
| if len(candidates) != 180: | |
| raise RuntimeError( | |
| f"Expected 180 w4a16_slow_transformer projections, found {len(candidates)}" | |
| ) | |
| generator = torch.Generator(device=candidates[0][1].weight.device) | |
| generator.manual_seed(probe_seed) | |
| records = [] | |
| for name, linear in candidates: | |
| parent_name, attribute = name.rsplit(".", 1) | |
| parent = model.get_submodule(parent_name) | |
| probe = torch.randn( | |
| 1, | |
| linear.in_features, | |
| dtype=torch.bfloat16, | |
| device=linear.weight.device, | |
| generator=generator, | |
| ) * 0.1 | |
| reference = F.linear(probe, linear.weight) | |
| replacement = NVFP4Linear.from_linear( | |
| linear, w4a16_max_m=w4a16_max_m | |
| ) | |
| actual = replacement(probe) | |
| records.append( | |
| ConversionRecord( | |
| name=name, | |
| in_features=linear.in_features, | |
| out_features=linear.out_features, | |
| parameters=linear.weight.numel(), | |
| probe_cosine=float( | |
| F.cosine_similarity( | |
| actual.float().flatten(), | |
| reference.float().flatten(), | |
| dim=0, | |
| ).item() | |
| ), | |
| ) | |
| ) | |
| 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] | |
| parameters = sum(record.parameters for record in records) | |
| return { | |
| "policy": policy, | |
| "backend": ( | |
| f"w4a16_through_m{w4a16_max_m}+w4a4_above_threshold" | |
| ), | |
| "w4a16_max_m": w4a16_max_m, | |
| "modules": len(records), | |
| "projections": len(records), | |
| "parameters": parameters, | |
| "theoretical_bf16_source_bytes": parameters * 2, | |
| "packed_weight_bytes": _packed_weight_bytes(model), | |
| "probe_cosine_min": min(cosines), | |
| "probe_cosine_mean": sum(cosines) / len(cosines), | |
| "probe_cosine_max": max(cosines), | |
| "records": serialized, | |
| } | |
| mixed_role_policies = { | |
| "w4a16_mlp_mxfp8_attention": ({"w1", "w2", "w3"}, None), | |
| "w4a16_gate_up_mxfp8_rest": ({"w1", "w3"}, None), | |
| "w4a16_down_mxfp8_rest": ({"w2"}, None), | |
| "w4a16_gate_up_middle6_mxfp8_rest": ({"w1", "w3"}, set(range(15, 21))), | |
| "w4a16_gate_up_middle12_mxfp8_rest": ({"w1", "w3"}, set(range(12, 24))), | |
| "w4a16_gate_up_middle18_mxfp8_rest": ({"w1", "w3"}, set(range(9, 27))), | |
| "w4a16_gate_up_middle24_mxfp8_rest": ({"w1", "w3"}, set(range(6, 30))), | |
| "w4a16_gate_up_middle30_mxfp8_rest": ({"w1", "w3"}, set(range(3, 33))), | |
| "w4a16_gate_up_custom_mxfp8_rest": ({"w1", "w3"}, "custom"), | |
| } | |
| if policy in mixed_role_policies: | |
| from experimental.fp8 import MXFP8Linear | |
| candidates = [ | |
| (name, module) | |
| for name, module in model.named_modules() | |
| if _selected_slow_transformer(name, module) | |
| ] | |
| if len(candidates) != 180: | |
| raise RuntimeError( | |
| f"Expected 180 {policy} projections, " | |
| f"found {len(candidates)}" | |
| ) | |
| nvfp4_roles, policy_layers = mixed_role_policies[policy] | |
| if policy_layers == "custom": | |
| if not nvfp4_layers: | |
| raise ValueError(f"{policy} requires at least one nvfp4 layer") | |
| invalid_layers = sorted(set(nvfp4_layers) - set(range(36))) | |
| if invalid_layers: | |
| raise ValueError(f"Invalid slow-transformer layers: {invalid_layers}") | |
| selected_nvfp4_layers = set(nvfp4_layers) | |
| else: | |
| if nvfp4_layers is not None: | |
| raise ValueError("nvfp4_layers is only valid with the custom policy") | |
| selected_nvfp4_layers = policy_layers | |
| activation_scales: dict[int, torch.Tensor] = {} | |
| activation_clip_ratios: dict[int, float] = {} | |
| activation_quantization_scales: dict[int, dict[str, float]] = {} | |
| activation_corrections: dict[int, dict[str, torch.Tensor]] = {} | |
| activation_hadamard_blocks: dict[int, int] = {} | |
| activation_scaling = None | |
| if activation_scaling_report is not None: | |
| if nvfp4_roles != {"w1", "w3"}: | |
| raise ValueError( | |
| "Activation scaling requires NVFP4 gate and up projections together" | |
| ) | |
| layers_to_scale = ( | |
| set(range(36)) | |
| if selected_nvfp4_layers is None | |
| else set(selected_nvfp4_layers) | |
| ) | |
| ( | |
| activation_scales, | |
| activation_clip_ratios, | |
| activation_quantization_scales, | |
| activation_corrections, | |
| activation_scaling, | |
| ) = _prepare_activation_scaling( | |
| model, | |
| Path(activation_scaling_report), | |
| layers_to_scale, | |
| ) | |
| activation_hadamard_blocks = { | |
| int(record["layer"]): int(record["input_hadamard_block_size"]) | |
| for record in activation_scaling["layers"] | |
| if int(record.get("input_hadamard_block_size", 0)) | |
| } | |
| generator = torch.Generator(device=candidates[0][1].weight.device) | |
| generator.manual_seed(probe_seed) | |
| records = [] | |
| nvfp4_parameters = 0 | |
| mxfp8_parameters = 0 | |
| for name, linear in candidates: | |
| parent_name, attribute = name.rsplit(".", 1) | |
| parent = model.get_submodule(parent_name) | |
| probe = torch.randn( | |
| 1, | |
| linear.in_features, | |
| dtype=torch.bfloat16, | |
| device=linear.weight.device, | |
| generator=generator, | |
| ) * 0.1 | |
| reference = F.linear(probe, linear.weight) | |
| role = name.rsplit(".", 1)[-1] | |
| layer = int(name.split(".", 2)[1]) | |
| selected_layer = ( | |
| selected_nvfp4_layers is None or layer in selected_nvfp4_layers | |
| ) | |
| if ".feed_forward." in name and role in nvfp4_roles and selected_layer: | |
| activation_scale = activation_scales.get(layer) | |
| if activation_scale is None: | |
| replacement = NVFP4Linear.from_linear( | |
| linear, w4a16_max_m=w4a16_max_m | |
| ) | |
| actual_probe = probe | |
| else: | |
| scaled_weight = ( | |
| linear.weight.float() * activation_scale | |
| ) | |
| scaled_weight = _clip_weight_blocks( | |
| scaled_weight, | |
| activation_clip_ratios.get(layer, 1.0), | |
| ) | |
| input_hadamard_block_size = activation_hadamard_blocks.get( | |
| layer, | |
| 0, | |
| ) | |
| if input_hadamard_block_size: | |
| scaled_weight = _hadamard_blocks( | |
| scaled_weight, | |
| input_hadamard_block_size, | |
| ) | |
| scaled_weight = scaled_weight.to(dtype=torch.bfloat16) | |
| replacement = NVFP4Linear.from_weight( | |
| scaled_weight, | |
| w4a16_max_m=w4a16_max_m, | |
| quantization_scale=activation_quantization_scales.get( | |
| layer, {} | |
| ).get(role), | |
| correction_down=activation_corrections.get( | |
| layer, {} | |
| ).get("down"), | |
| correction_up=activation_corrections.get( | |
| layer, {} | |
| ).get(f"{role}_up"), | |
| sparse_correction_indices=activation_corrections.get( | |
| layer, {} | |
| ).get("sparse_indices"), | |
| sparse_correction_weight=activation_corrections.get( | |
| layer, {} | |
| ).get(f"{role}_sparse_weight"), | |
| input_hadamard_block_size=input_hadamard_block_size, | |
| ) | |
| actual_probe = ( | |
| probe.float() / activation_scale | |
| ).to(dtype=torch.bfloat16) | |
| precision = ( | |
| f"nvfp4_w4a16_through_m{w4a16_max_m}" | |
| + ("_activation_scaled" if activation_scale is not None else "") | |
| + ( | |
| "_low_rank_corrected" | |
| if "down" in activation_corrections.get(layer, {}) | |
| else "" | |
| ) | |
| + ( | |
| "_sparse_channel_corrected" | |
| if "sparse_indices" in activation_corrections.get(layer, {}) | |
| else "" | |
| ) | |
| + ( | |
| f"_hadamard{activation_hadamard_blocks[layer]}" | |
| if layer in activation_hadamard_blocks | |
| else "" | |
| ) | |
| ) | |
| nvfp4_parameters += linear.weight.numel() | |
| else: | |
| replacement = MXFP8Linear.from_linear(linear) | |
| actual_probe = probe | |
| precision = "mxfp8_w8a8" | |
| mxfp8_parameters += linear.weight.numel() | |
| actual = replacement(actual_probe) | |
| record = ConversionRecord( | |
| name=name, | |
| in_features=linear.in_features, | |
| out_features=linear.out_features, | |
| parameters=linear.weight.numel(), | |
| probe_cosine=float( | |
| F.cosine_similarity( | |
| actual.float().flatten(), | |
| reference.float().flatten(), | |
| dim=0, | |
| ).item() | |
| ), | |
| ) | |
| records.append({**asdict(record), "precision": precision}) | |
| setattr(parent, attribute, replacement) | |
| torch.cuda.synchronize(candidates[0][1].weight.device) | |
| nvfp4_bytes = _packed_weight_bytes(model) | |
| mxfp8_bytes = sum( | |
| module.weight_fp8.numel() * module.weight_fp8.element_size() | |
| + module.weight_scale_storage.numel() | |
| * module.weight_scale_storage.element_size() | |
| for module in model.modules() | |
| if isinstance(module, MXFP8Linear) | |
| ) | |
| cosines = [record["probe_cosine"] for record in records] | |
| parameters = nvfp4_parameters + mxfp8_parameters | |
| correction_tensors = { | |
| value.data_ptr(): value | |
| for correction in activation_corrections.values() | |
| for value in correction.values() | |
| } | |
| correction_parameters = sum( | |
| value.numel() | |
| for value in correction_tensors.values() | |
| if value.is_floating_point() | |
| ) | |
| correction_bytes = sum( | |
| value.numel() * value.element_size() | |
| for value in correction_tensors.values() | |
| ) | |
| low_rank_tensors = { | |
| value.data_ptr(): value | |
| for correction in activation_corrections.values() | |
| for key, value in correction.items() | |
| if key == "down" or key.endswith("_up") | |
| } | |
| sparse_correction_tensors = { | |
| value.data_ptr(): value | |
| for correction in activation_corrections.values() | |
| for key, value in correction.items() | |
| if key == "sparse_indices" or key.endswith("_sparse_weight") | |
| } | |
| low_rank_correction_parameters = sum( | |
| value.numel() for value in low_rank_tensors.values() | |
| ) | |
| low_rank_correction_bytes = sum( | |
| value.numel() * value.element_size() | |
| for value in low_rank_tensors.values() | |
| ) | |
| sparse_correction_parameters = sum( | |
| value.numel() | |
| for value in sparse_correction_tensors.values() | |
| if value.is_floating_point() | |
| ) | |
| sparse_correction_bytes = sum( | |
| value.numel() * value.element_size() | |
| for value in sparse_correction_tensors.values() | |
| ) | |
| return { | |
| "policy": policy, | |
| "backend": "selective_nvfp4_weight_bf16_activation+mxfp8_rest", | |
| "w4a16_max_m": w4a16_max_m, | |
| "nvfp4_mlp_roles": sorted(nvfp4_roles), | |
| "nvfp4_layers": ( | |
| "all" | |
| if selected_nvfp4_layers is None | |
| else sorted(selected_nvfp4_layers) | |
| ), | |
| "activation_scaling": activation_scaling, | |
| "modules": len(records), | |
| "projections": len(records), | |
| "parameters": parameters, | |
| "nvfp4_parameters": nvfp4_parameters, | |
| "mxfp8_parameters": mxfp8_parameters, | |
| "low_rank_correction_parameters": low_rank_correction_parameters, | |
| "low_rank_correction_bytes": low_rank_correction_bytes, | |
| "sparse_correction_parameters": sparse_correction_parameters, | |
| "sparse_correction_bytes": sparse_correction_bytes, | |
| "correction_parameters": correction_parameters, | |
| "correction_bytes": correction_bytes, | |
| "theoretical_bf16_source_bytes": parameters * 2, | |
| "packed_weight_bytes": nvfp4_bytes + mxfp8_bytes, | |
| "nvfp4_packed_weight_bytes": nvfp4_bytes, | |
| "mxfp8_packed_weight_bytes": mxfp8_bytes, | |
| "probe_cosine_min": min(cosines), | |
| "probe_cosine_mean": sum(cosines) / len(cosines), | |
| "probe_cosine_max": max(cosines), | |
| "records": records, | |
| } | |
| if policy == "hybrid_slow_transformer": | |
| from experimental.fp8 import MXFP8Linear | |
| nvfp4 = convert_s2_pro_nvfp4( | |
| model, | |
| policy="slow_mlp_fused", | |
| probe_seed=probe_seed, | |
| ) | |
| candidates = [ | |
| (name, module) | |
| for name, module in model.named_modules() | |
| if _selected_slow_attention(name, module) | |
| ] | |
| if len(candidates) != 72: | |
| raise RuntimeError( | |
| f"Expected 72 hybrid attention projections, found {len(candidates)}" | |
| ) | |
| generator = torch.Generator(device=candidates[0][1].weight.device) | |
| generator.manual_seed(probe_seed + 1) | |
| attention_records = [] | |
| for name, linear in candidates: | |
| parent_name, attribute = name.rsplit(".", 1) | |
| parent = model.get_submodule(parent_name) | |
| probe = torch.randn( | |
| 1, | |
| linear.in_features, | |
| dtype=torch.bfloat16, | |
| device=linear.weight.device, | |
| generator=generator, | |
| ) * 0.1 | |
| reference = F.linear(probe, linear.weight) | |
| replacement = MXFP8Linear.from_linear(linear) | |
| actual = replacement(probe) | |
| attention_records.append( | |
| ConversionRecord( | |
| name=name, | |
| in_features=linear.in_features, | |
| out_features=linear.out_features, | |
| parameters=linear.weight.numel(), | |
| probe_cosine=float( | |
| F.cosine_similarity( | |
| actual.float().flatten(), | |
| reference.float().flatten(), | |
| dim=0, | |
| ).item() | |
| ), | |
| ) | |
| ) | |
| setattr(parent, attribute, replacement) | |
| torch.cuda.synchronize(candidates[0][1].weight.device) | |
| attention_parameters = sum(record.parameters for record in attention_records) | |
| attention_packed_bytes = sum( | |
| module.weight_fp8.numel() * module.weight_fp8.element_size() | |
| + module.weight_scale_storage.numel() | |
| * module.weight_scale_storage.element_size() | |
| for module in model.modules() | |
| if isinstance(module, MXFP8Linear) | |
| ) | |
| cosines = [ | |
| *[record["probe_cosine"] for record in nvfp4["records"]], | |
| *[record.probe_cosine for record in attention_records], | |
| ] | |
| return { | |
| "policy": policy, | |
| "backend": "fused_nvfp4_mlp+mxfp8_attention", | |
| "modules": nvfp4["modules"] + len(attention_records), | |
| "projections": nvfp4["projections"] + len(attention_records), | |
| "parameters": nvfp4["parameters"] + attention_parameters, | |
| "theoretical_bf16_source_bytes": ( | |
| nvfp4["theoretical_bf16_source_bytes"] | |
| + attention_parameters * 2 | |
| ), | |
| "packed_weight_bytes": ( | |
| nvfp4["packed_weight_bytes"] + attention_packed_bytes | |
| ), | |
| "nvfp4_mlp": nvfp4, | |
| "mxfp8_attention": { | |
| "modules": len(attention_records), | |
| "parameters": attention_parameters, | |
| "packed_weight_bytes": attention_packed_bytes, | |
| "records": [asdict(record) for record in attention_records], | |
| }, | |
| "probe_cosine_min": min(cosines), | |
| "probe_cosine_mean": sum(cosines) / len(cosines), | |
| "probe_cosine_max": max(cosines), | |
| } | |
| if policy == "slow_transformer_fused": | |
| candidates = [ | |
| (name, module) | |
| for name, module in model.named_modules() | |
| if _selected_slow_transformer_block(name, module) | |
| ] | |
| if len(candidates) != 36: | |
| raise RuntimeError( | |
| f"Expected 36 slow_transformer_fused blocks, found {len(candidates)}" | |
| ) | |
| generator = torch.Generator(device=candidates[0][1].attention.wqkv.weight.device) | |
| generator.manual_seed(probe_seed) | |
| records = [] | |
| for name, module in candidates: | |
| parent_name, attribute = name.rsplit(".", 1) | |
| parent = model.get_submodule(parent_name) | |
| dim = module.attention.wqkv.in_features | |
| probe = torch.randn( | |
| 1, | |
| dim, | |
| dtype=torch.bfloat16, | |
| device=module.attention.wqkv.weight.device, | |
| generator=generator, | |
| ) * 0.1 | |
| wo_probe = torch.randn( | |
| 1, | |
| module.attention.wo.in_features, | |
| dtype=torch.bfloat16, | |
| device=module.attention.wo.weight.device, | |
| generator=generator, | |
| ) * 0.1 | |
| reference_wqkv = module.attention.wqkv(probe) | |
| reference_wo = module.attention.wo(wo_probe) | |
| reference_mlp = module.feed_forward(probe) | |
| parameters = sum( | |
| linear.weight.numel() | |
| for linear in ( | |
| module.attention.wqkv, | |
| module.attention.wo, | |
| module.feed_forward.w1, | |
| module.feed_forward.w2, | |
| module.feed_forward.w3, | |
| ) | |
| ) | |
| replacement = NVFP4TransformerBlock(module) | |
| actual_wqkv = replacement.attention.wqkv(probe) | |
| actual_wo = replacement.attention.wo(wo_probe) | |
| actual_mlp = replacement.feed_forward(probe) | |
| records.append( | |
| FusedTransformerConversionRecord( | |
| name=name, | |
| parameters=parameters, | |
| wqkv_probe_cosine=float( | |
| F.cosine_similarity( | |
| actual_wqkv.float().flatten(), | |
| reference_wqkv.float().flatten(), | |
| dim=0, | |
| ).item() | |
| ), | |
| wo_probe_cosine=float( | |
| F.cosine_similarity( | |
| actual_wo.float().flatten(), | |
| reference_wo.float().flatten(), | |
| dim=0, | |
| ).item() | |
| ), | |
| mlp_probe_cosine=float( | |
| F.cosine_similarity( | |
| actual_mlp.float().flatten(), | |
| reference_mlp.float().flatten(), | |
| dim=0, | |
| ).item() | |
| ), | |
| ) | |
| ) | |
| setattr(parent, attribute, replacement) | |
| torch.cuda.synchronize(candidates[0][1].attention.wqkv.qdata.device) | |
| serialized = [asdict(record) for record in records] | |
| cosines = [ | |
| cosine | |
| for record in records | |
| for cosine in ( | |
| record.wqkv_probe_cosine, | |
| record.wo_probe_cosine, | |
| record.mlp_probe_cosine, | |
| ) | |
| ] | |
| parameters = sum(record.parameters for record in records) | |
| return { | |
| "policy": policy, | |
| "backend": "direct_fused_m1+tensorcore_prefill", | |
| "modules": len(records), | |
| "projections": len(records) * 5, | |
| "parameters": parameters, | |
| "theoretical_bf16_source_bytes": parameters * 2, | |
| "packed_weight_bytes": _packed_weight_bytes(model), | |
| "probe_cosine_min": min(cosines), | |
| "probe_cosine_mean": sum(cosines) / len(cosines), | |
| "probe_cosine_max": max(cosines), | |
| "records": serialized, | |
| } | |
| if policy == "slow_mlp_fused": | |
| candidates = [ | |
| (name, module) | |
| for name, module in model.named_modules() | |
| if _selected_slow_mlp_block(name, module) | |
| ] | |
| if len(candidates) != 36: | |
| raise RuntimeError( | |
| f"Expected 36 slow_mlp_fused blocks, found {len(candidates)}" | |
| ) | |
| generator = torch.Generator(device=candidates[0][1].w1.weight.device) | |
| generator.manual_seed(probe_seed) | |
| records = [] | |
| for name, module in candidates: | |
| parent_name, attribute = name.rsplit(".", 1) | |
| parent = model.get_submodule(parent_name) | |
| probe = torch.randn( | |
| 1, | |
| module.w1.in_features, | |
| dtype=torch.bfloat16, | |
| device=module.w1.weight.device, | |
| generator=generator, | |
| ) * 0.1 | |
| reference = module(probe) | |
| replacement = NVFP4FeedForward.from_module(module) | |
| actual = replacement(probe) | |
| records.append( | |
| FusedMLPConversionRecord( | |
| name=name, | |
| in_features=module.w1.in_features, | |
| intermediate_features=module.w1.out_features, | |
| out_features=module.w2.out_features, | |
| parameters=( | |
| module.w1.weight.numel() | |
| + module.w2.weight.numel() | |
| + module.w3.weight.numel() | |
| ), | |
| probe_cosine=float( | |
| F.cosine_similarity( | |
| actual.float().flatten(), | |
| reference.float().flatten(), | |
| dim=0, | |
| ).item() | |
| ), | |
| ) | |
| ) | |
| setattr(parent, attribute, replacement) | |
| torch.cuda.synchronize(candidates[0][1].w1.weight.device) | |
| serialized = [asdict(record) for record in records] | |
| cosines = [record.probe_cosine for record in records] | |
| parameters = sum(record.parameters for record in records) | |
| return { | |
| "policy": policy, | |
| "backend": "direct_fused_m1+tensorcore_prefill", | |
| "modules": len(records), | |
| "projections": len(records) * 3, | |
| "parameters": parameters, | |
| "theoretical_bf16_source_bytes": parameters * 2, | |
| "packed_weight_bytes": _packed_weight_bytes(model), | |
| "probe_cosine_min": min(cosines), | |
| "probe_cosine_mean": sum(cosines) / len(cosines), | |
| "probe_cosine_max": max(cosines), | |
| "records": serialized, | |
| } | |
| if policy not in selectors: | |
| raise ValueError(f"Unsupported initial NVFP4 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)}" | |
| ) | |
| generator = torch.Generator(device=candidates[0][1].weight.device) | |
| generator.manual_seed(probe_seed) | |
| records = [] | |
| for name, linear in candidates: | |
| parent_name, attribute = name.rsplit(".", 1) | |
| parent = model.get_submodule(parent_name) | |
| replacement = NVFP4Linear.from_linear(linear, direct_m1=direct_m1) | |
| 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, | |
| "backend": "direct_m1+tensorcore" if direct_m1 else "tensorcore", | |
| "modules": len(records), | |
| "parameters": sum(record.parameters for record in records), | |
| "theoretical_bf16_source_bytes": sum(record.parameters * 2 for record in records), | |
| "packed_weight_bytes": _packed_weight_bytes(model), | |
| "probe_cosine_min": min(cosines), | |
| "probe_cosine_mean": sum(cosines) / len(cosines), | |
| "probe_cosine_max": max(cosines), | |
| "records": serialized, | |
| } | |