[Feature] Support NVFP4 token embedding in ModelOpt mixed-precision checkpoints (#34222)
Co-authored-by: Brayden Zhong <b8zhong@uwaterloo.ca>
This commit is contained in:
co-authored by
Brayden Zhong
parent
e226bb711c
commit
d6a066131c
@@ -626,6 +626,106 @@ class ModelOptFp8KVCacheMethod(BaseKVCacheMethod):
|
|||||||
super().__init__(quant_config)
|
super().__init__(quant_config)
|
||||||
|
|
||||||
|
|
||||||
|
# E2M1 code -> value, indexed by the 4-bit code (sign << 3 | magnitude).
|
||||||
|
_E2M1_LUT = (0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0)
|
||||||
|
|
||||||
|
|
||||||
|
class ModelOptNvFp4EmbeddingMethod(QuantizeMethodBase):
|
||||||
|
"""NVFP4 token embedding, dequantized on gather."""
|
||||||
|
|
||||||
|
def __init__(self, quant_config: ModelOptFp4Config):
|
||||||
|
self.quant_config = quant_config
|
||||||
|
self.params_dtype = torch.bfloat16
|
||||||
|
|
||||||
|
def create_weights(
|
||||||
|
self,
|
||||||
|
layer: torch.nn.Module,
|
||||||
|
input_size_per_partition: int,
|
||||||
|
output_partition_sizes: List[int],
|
||||||
|
input_size: int,
|
||||||
|
output_size: int,
|
||||||
|
params_dtype: torch.dtype,
|
||||||
|
**extra_weight_attrs,
|
||||||
|
):
|
||||||
|
self.params_dtype = params_dtype
|
||||||
|
group_size = self.quant_config.group_size
|
||||||
|
if input_size_per_partition % group_size != 0:
|
||||||
|
raise ValueError(
|
||||||
|
f"NVFP4 embedding needs embedding_dim divisible by {group_size}, "
|
||||||
|
f"got {input_size_per_partition}."
|
||||||
|
)
|
||||||
|
num_rows = sum(output_partition_sizes)
|
||||||
|
weight_loader = extra_weight_attrs.get("weight_loader")
|
||||||
|
|
||||||
|
weight = ModelWeightParameter(
|
||||||
|
data=torch.empty(
|
||||||
|
num_rows, input_size_per_partition // 2, dtype=torch.uint8
|
||||||
|
),
|
||||||
|
input_dim=1,
|
||||||
|
output_dim=0,
|
||||||
|
weight_loader=weight_loader,
|
||||||
|
)
|
||||||
|
layer.register_parameter("weight", weight)
|
||||||
|
|
||||||
|
weight_scale = ModelWeightParameter(
|
||||||
|
data=torch.empty(
|
||||||
|
num_rows,
|
||||||
|
input_size_per_partition // group_size,
|
||||||
|
dtype=torch.float8_e4m3fn,
|
||||||
|
),
|
||||||
|
input_dim=1,
|
||||||
|
output_dim=0,
|
||||||
|
weight_loader=weight_loader,
|
||||||
|
)
|
||||||
|
layer.register_parameter("weight_scale", weight_scale)
|
||||||
|
|
||||||
|
weight_scale_2 = Parameter(
|
||||||
|
torch.empty(1, dtype=torch.float32), requires_grad=False
|
||||||
|
)
|
||||||
|
set_weight_attrs(
|
||||||
|
weight_scale_2,
|
||||||
|
{"weight_loader": lambda p, w: p.data.copy_(w.reshape(p.shape).float())},
|
||||||
|
)
|
||||||
|
layer.register_parameter("weight_scale_2", weight_scale_2)
|
||||||
|
|
||||||
|
# A buffer; CUDA graph capture rejects host->device copies.
|
||||||
|
layer.register_buffer(
|
||||||
|
"e2m1_lut",
|
||||||
|
torch.tensor(_E2M1_LUT, dtype=torch.float32),
|
||||||
|
persistent=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def apply(self, *args, **kwargs):
|
||||||
|
raise NotImplementedError(
|
||||||
|
"NVFP4 embedding is gather-only. Reaching here means a tied lm_head "
|
||||||
|
"is sharing this module; exclude the embedding from NVFP4 in the "
|
||||||
|
"quantization recipe to serve such a checkpoint."
|
||||||
|
)
|
||||||
|
|
||||||
|
def embedding(self, layer: torch.nn.Module, input_: torch.Tensor) -> torch.Tensor:
|
||||||
|
index_shape = input_.shape
|
||||||
|
flat = input_.reshape(-1)
|
||||||
|
packed = layer.weight[flat] # [T, H/2] uint8
|
||||||
|
scale = layer.weight_scale[flat] # [T, H/16] e4m3
|
||||||
|
rows, half = packed.shape
|
||||||
|
hidden = half * 2
|
||||||
|
|
||||||
|
codes = packed.new_empty((rows, hidden))
|
||||||
|
codes[:, 0::2] = packed & 0x0F
|
||||||
|
codes[:, 1::2] = packed >> 4
|
||||||
|
|
||||||
|
mag = layer.e2m1_lut[(codes & 0x7).long()]
|
||||||
|
vals = torch.where(codes & 0x8 != 0, -mag, mag)
|
||||||
|
|
||||||
|
group_size = self.quant_config.group_size
|
||||||
|
eff = scale.float() * layer.weight_scale_2.float()
|
||||||
|
out = vals.view(rows, hidden // group_size, group_size) * eff.unsqueeze(-1)
|
||||||
|
return out.view(*index_shape, hidden).to(self.params_dtype)
|
||||||
|
|
||||||
|
|
||||||
class ModelOptMixedPrecisionConfig(ModelOptQuantConfig):
|
class ModelOptMixedPrecisionConfig(ModelOptQuantConfig):
|
||||||
"""Configuration for ModelOpt MIXED_PRECISION checkpoints."""
|
"""Configuration for ModelOpt MIXED_PRECISION checkpoints."""
|
||||||
|
|
||||||
@@ -823,7 +923,10 @@ class ModelOptMixedPrecisionConfig(ModelOptQuantConfig):
|
|||||||
) -> Optional[QuantizeMethodBase]:
|
) -> Optional[QuantizeMethodBase]:
|
||||||
from sglang.srt.layers.linear import LinearBase
|
from sglang.srt.layers.linear import LinearBase
|
||||||
from sglang.srt.layers.moe.fused_moe_triton import FusedMoE
|
from sglang.srt.layers.moe.fused_moe_triton import FusedMoE
|
||||||
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
|
from sglang.srt.layers.vocab_parallel_embedding import (
|
||||||
|
ParallelLMHead,
|
||||||
|
VocabParallelEmbedding,
|
||||||
|
)
|
||||||
|
|
||||||
quant_algo = self._resolve_quant_algo(prefix)
|
quant_algo = self._resolve_quant_algo(prefix)
|
||||||
|
|
||||||
@@ -842,6 +945,17 @@ class ModelOptMixedPrecisionConfig(ModelOptQuantConfig):
|
|||||||
return ModelOptNvFp4A16LinearMethod(self.nvfp4a16_config)
|
return ModelOptNvFp4A16LinearMethod(self.nvfp4a16_config)
|
||||||
return UnquantizedLinearMethod()
|
return UnquantizedLinearMethod()
|
||||||
|
|
||||||
|
# Must stay after the ParallelLMHead branch: ParallelLMHead subclasses
|
||||||
|
# VocabParallelEmbedding, and a tied lm_head IS the embedding module.
|
||||||
|
if isinstance(layer, VocabParallelEmbedding):
|
||||||
|
if is_layer_skipped(
|
||||||
|
prefix, self.exclude_modules, self.packed_modules_mapping
|
||||||
|
) or self.is_layer_excluded(prefix):
|
||||||
|
return None
|
||||||
|
if quant_algo == "NVFP4":
|
||||||
|
return ModelOptNvFp4EmbeddingMethod(self.nvfp4_config)
|
||||||
|
return None
|
||||||
|
|
||||||
if self.kv_cache_quant_algo and isinstance(layer, RadixAttention):
|
if self.kv_cache_quant_algo and isinstance(layer, RadixAttention):
|
||||||
return ModelOptFp8KVCacheMethod(self.fp8_config)
|
return ModelOptFp8KVCacheMethod(self.fp8_config)
|
||||||
|
|
||||||
|
|||||||
+124
@@ -0,0 +1,124 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.layers.quantization.modelopt_quant import (
|
||||||
|
ModelOptFp4Config,
|
||||||
|
ModelOptNvFp4EmbeddingMethod,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
GROUP_SIZE = 16
|
||||||
|
|
||||||
|
# Written out independently of the implementation: the E2M1 code points in
|
||||||
|
# magnitude order, so index == the 3-bit magnitude code.
|
||||||
|
_REFERENCE_E2M1 = [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0]
|
||||||
|
|
||||||
|
|
||||||
|
def reference_dequant(
|
||||||
|
packed: torch.Tensor, block_scale: torch.Tensor, global_scale: float
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Comparison oracle. Kept as a plain per-element loop on purpose: a
|
||||||
|
vectorized rewrite would mirror the code under test."""
|
||||||
|
rows, half = packed.shape
|
||||||
|
hidden = half * 2
|
||||||
|
out = torch.zeros(rows, hidden, dtype=torch.float32)
|
||||||
|
for r in range(rows):
|
||||||
|
for c in range(hidden):
|
||||||
|
byte = int(packed[r, c // 2])
|
||||||
|
code = (byte & 0x0F) if c % 2 == 0 else (byte >> 4)
|
||||||
|
magnitude = _REFERENCE_E2M1[code & 0x7]
|
||||||
|
value = -magnitude if code & 0x8 else magnitude
|
||||||
|
scale = float(block_scale[r, c // GROUP_SIZE]) * global_scale
|
||||||
|
out[r, c] = value * scale
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def build_layer(method, vocab_size: int, hidden_size: int) -> torch.nn.Module:
|
||||||
|
"""Materialize through create_weights, then fill as a checkpoint would."""
|
||||||
|
layer = torch.nn.Module()
|
||||||
|
method.create_weights(
|
||||||
|
layer,
|
||||||
|
input_size_per_partition=hidden_size,
|
||||||
|
output_partition_sizes=[vocab_size],
|
||||||
|
input_size=hidden_size,
|
||||||
|
output_size=vocab_size,
|
||||||
|
params_dtype=torch.bfloat16,
|
||||||
|
)
|
||||||
|
|
||||||
|
generator = torch.Generator().manual_seed(0)
|
||||||
|
layer.weight.data.copy_(
|
||||||
|
torch.randint(
|
||||||
|
0,
|
||||||
|
256,
|
||||||
|
(vocab_size, hidden_size // 2),
|
||||||
|
dtype=torch.uint8,
|
||||||
|
generator=generator,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
# Keep the block scales in a range e4m3 represents exactly.
|
||||||
|
layer.weight_scale.data.copy_(
|
||||||
|
torch.randint(
|
||||||
|
1,
|
||||||
|
8,
|
||||||
|
(vocab_size, hidden_size // GROUP_SIZE),
|
||||||
|
dtype=torch.int32,
|
||||||
|
generator=generator,
|
||||||
|
).to(torch.float8_e4m3fn)
|
||||||
|
)
|
||||||
|
layer.weight_scale_2.data.fill_(0.125)
|
||||||
|
return layer
|
||||||
|
|
||||||
|
|
||||||
|
class TestNvFp4Embedding(CustomTestCase):
|
||||||
|
def setUp(self):
|
||||||
|
self.method = ModelOptNvFp4EmbeddingMethod(
|
||||||
|
ModelOptFp4Config(
|
||||||
|
is_checkpoint_nvfp4_serialized=True, group_size=GROUP_SIZE
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_matches_reference_dequant(self):
|
||||||
|
vocab_size, hidden_size = 24, 64
|
||||||
|
layer = build_layer(self.method, vocab_size, hidden_size)
|
||||||
|
self.assertEqual(tuple(layer.weight.shape), (vocab_size, hidden_size // 2))
|
||||||
|
self.assertEqual(
|
||||||
|
tuple(layer.weight_scale.shape), (vocab_size, hidden_size // GROUP_SIZE)
|
||||||
|
)
|
||||||
|
|
||||||
|
ids = torch.tensor([[0, 5, 5], [23, 11, 0]])
|
||||||
|
got = self.method.embedding(layer, ids)
|
||||||
|
expected = reference_dequant(
|
||||||
|
layer.weight[ids.reshape(-1)],
|
||||||
|
layer.weight_scale[ids.reshape(-1)].float(),
|
||||||
|
float(layer.weight_scale_2),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(tuple(got.shape), (2, 3, hidden_size))
|
||||||
|
self.assertEqual(got.dtype, torch.bfloat16)
|
||||||
|
torch.testing.assert_close(
|
||||||
|
got.reshape(-1, hidden_size).float(),
|
||||||
|
expected.to(torch.bfloat16).float(),
|
||||||
|
rtol=0,
|
||||||
|
atol=0,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_hidden_size_must_divide_group_size(self):
|
||||||
|
with self.assertRaisesRegex(ValueError, "divisible by 16"):
|
||||||
|
self.method.create_weights(
|
||||||
|
torch.nn.Module(),
|
||||||
|
input_size_per_partition=40,
|
||||||
|
output_partition_sizes=[8],
|
||||||
|
input_size=40,
|
||||||
|
output_size=8,
|
||||||
|
params_dtype=torch.bfloat16,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main(verbosity=2)
|
||||||
Reference in New Issue
Block a user