From d6a066131cad33a7429f6c1f5a6332d44df9e753 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Sun, 9 Aug 2026 23:37:54 -0700 Subject: [PATCH] [Feature] Support NVFP4 token embedding in ModelOpt mixed-precision checkpoints (#34222) Co-authored-by: Brayden Zhong --- .../srt/layers/quantization/modelopt_quant.py | 116 +++++++++++++++- test/registered/quant/test_nvfp4_embedding.py | 124 ++++++++++++++++++ 2 files changed, 239 insertions(+), 1 deletion(-) create mode 100755 test/registered/quant/test_nvfp4_embedding.py diff --git a/python/sglang/srt/layers/quantization/modelopt_quant.py b/python/sglang/srt/layers/quantization/modelopt_quant.py index 6422d82d1..10cb2b894 100755 --- a/python/sglang/srt/layers/quantization/modelopt_quant.py +++ b/python/sglang/srt/layers/quantization/modelopt_quant.py @@ -626,6 +626,106 @@ class ModelOptFp8KVCacheMethod(BaseKVCacheMethod): 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): """Configuration for ModelOpt MIXED_PRECISION checkpoints.""" @@ -823,7 +923,10 @@ class ModelOptMixedPrecisionConfig(ModelOptQuantConfig): ) -> Optional[QuantizeMethodBase]: from sglang.srt.layers.linear import LinearBase 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) @@ -842,6 +945,17 @@ class ModelOptMixedPrecisionConfig(ModelOptQuantConfig): return ModelOptNvFp4A16LinearMethod(self.nvfp4a16_config) 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): return ModelOptFp8KVCacheMethod(self.fp8_config) diff --git a/test/registered/quant/test_nvfp4_embedding.py b/test/registered/quant/test_nvfp4_embedding.py new file mode 100755 index 000000000..2716aa4eb --- /dev/null +++ b/test/registered/quant/test_nvfp4_embedding.py @@ -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)