ajh-code's picture
Add files using upload-large-folder tool
16f5171 verified
Raw
History Blame Contribute Delete
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
@dataclass
class ConversionRecord:
name: str
in_features: int
out_features: int
parameters: int
probe_cosine: float
@dataclass
class FusedMLPConversionRecord:
name: str
in_features: int
intermediate_features: int
out_features: int
parameters: int
probe_cosine: float
@dataclass
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,
)
@torch.inference_mode()
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()
@lru_cache(maxsize=1)
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,
)
@lru_cache(maxsize=1)
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,
)
@lru_cache(maxsize=1)
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,
)
@lru_cache(maxsize=1)
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)
@classmethod
@torch.inference_mode()
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,
)
@classmethod
@torch.inference_mode()
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"
@classmethod
@torch.inference_mode()
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)
@torch.inference_mode()
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,
}
@torch.inference_mode()
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,
}