[NVIDIA] Fix SM107 MXFP8 activation prep (#35405)

Signed-off-by: Sahithi Chigurupati <chigurupati.sahithi@gmail.com>
Co-authored-by: Mohammad Miadh Angkad <176301910+mmangkad@users.noreply.github.com>
This commit is contained in:
Sahithi Chigurupati
2026-08-23 18:17:03 +08:00
committed by GitHub
co-authored by Mohammad Miadh Angkad
parent 27aa48bca1
commit 44db041700
4 changed files with 216 additions and 50 deletions
@@ -1,16 +1,39 @@
"""CPU unit tests for MXFP4 conversion and MXFP8 fake-output metadata."""
import unittest
import torch
from sglang.srt.layers.quantization.fp8_utils import (
_fake_flashinfer_mxfp8_quantize,
quantize_block_fp8_weight_to_mxfp4,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=4, suite="base-a-test-cpu")
class TestFp8UtilsMxfp4(unittest.TestCase):
class TestFp8UtilsMxfp4(CustomTestCase):
def test_fake_flashinfer_mxfp8_quantize_linear_scale_shape(self):
"""The fake op must flatten leading dimensions and preserve scale groups."""
input = torch.empty((2, 3, 96), dtype=torch.bfloat16)
quantized, scale = _fake_flashinfer_mxfp8_quantize(input, False, alignment=128)
self.assertEqual(quantized.shape, torch.Size([6, 128]))
self.assertEqual(quantized.dtype, torch.float8_e4m3fn)
self.assertEqual(scale.shape, torch.Size([24]))
self.assertEqual(scale.dtype, torch.uint8)
def test_fake_flashinfer_mxfp8_quantize_swizzled_scale_shape(self):
input = torch.empty((3, 64), dtype=torch.bfloat16)
quantized, scale = _fake_flashinfer_mxfp8_quantize(input, True, alignment=64)
self.assertEqual(quantized.shape, torch.Size([3, 64]))
self.assertEqual(scale.shape, torch.Size([512]))
def test_quantize_block_fp8_weight_to_mxfp4_shapes_and_dtype(self):
fp8_weight = (
torch.linspace(-2.0, 2.0, 32 * 32, dtype=torch.float32)
@@ -0,0 +1,125 @@
"""CPU unit tests for MXFP8 activation-preparation dispatch."""
import importlib
import unittest
from unittest.mock import patch
import torch
from sglang.srt.layers.quantization.mxfp4 import (
_prepare_flashinfer_mxfp8_activations,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
per_token_group_quant_module = importlib.import_module(
"sglang.kernels.ops.quantization.per_token_group_quant"
)
class TestMxfp4FlashinferActivationPrep(CustomTestCase):
def test_sm107_handoff_miss_uses_flashinfer_quantizer(self):
x = torch.randn(3, 64, dtype=torch.bfloat16)
x_quant = torch.empty(3, 64, dtype=torch.float8_e4m3fn)
x_scale = torch.arange(6, dtype=torch.uint8).reshape(3, 2)
with patch(
"sglang.srt.layers.moe.route_quant_handoff.take", return_value=None
) as take, patch(
"sglang.srt.layers.quantization.mxfp4._is_sm107_supported",
return_value=True,
), patch(
"sglang.srt.layers.quantization.fp8_utils.flashinfer_mxfp8_quantize",
return_value=(x_quant, x_scale),
create=True,
) as quantize:
actual_x, packed_topk, actual_quant, actual_scale = (
_prepare_flashinfer_mxfp8_activations(x, 64)
)
take.assert_called_once_with(x)
quantize.assert_called_once_with(x, False, alignment=64)
self.assertIs(actual_x, x)
self.assertIsNone(packed_topk)
self.assertIs(actual_quant, x_quant)
self.assertTrue(torch.equal(actual_scale.view(torch.uint8), x_scale))
def test_other_sm10x_handoff_miss_keeps_triton_quantizer(self):
x = torch.randn(3, 64, dtype=torch.bfloat16)
x_quant = torch.empty(3, 64, dtype=torch.float8_e4m3fn)
x_scale = torch.arange(6, dtype=torch.uint8).reshape(3, 2)
with patch(
"sglang.srt.layers.moe.route_quant_handoff.take", return_value=None
), patch(
"sglang.srt.layers.quantization.mxfp4._is_sm107_supported",
return_value=False,
), patch.object(
per_token_group_quant_module,
"per_token_group_quant",
return_value=(x_quant, x_scale),
) as quantize, patch(
"sglang.srt.layers.quantization.fp8_utils.flashinfer_mxfp8_quantize",
create=True,
) as flashinfer_quantize:
actual_x, packed_topk, actual_quant, actual_scale = (
_prepare_flashinfer_mxfp8_activations(x, 64)
)
quantize.assert_called_once_with(x, group_size=32, scale_ue8m0=True)
flashinfer_quantize.assert_not_called()
self.assertIs(actual_x, x)
self.assertIsNone(packed_topk)
self.assertIs(actual_quant, x_quant)
self.assertTrue(torch.equal(actual_scale.view(torch.uint8), x_scale))
def test_padded_input_keeps_flashinfer_quantizer(self):
"""A group-aligned input must use hidden-size-aligned quantization."""
x = torch.randn(3, 96, dtype=torch.bfloat16)
x_quant = torch.empty(3, 128, dtype=torch.float8_e4m3fn)
x_scale = torch.arange(12, dtype=torch.uint8)
with patch(
"sglang.srt.layers.quantization.fp8_utils.flashinfer_mxfp8_quantize",
return_value=(x_quant, x_scale),
create=True,
) as quantize, patch("sglang.srt.layers.moe.route_quant_handoff.take") as take:
actual_x, packed_topk, actual_quant, actual_scale = (
_prepare_flashinfer_mxfp8_activations(x, 128)
)
take.assert_not_called()
quantize.assert_called_once_with(x, False, alignment=128)
self.assertIs(actual_x, x)
self.assertIsNone(packed_topk)
self.assertIs(actual_quant, x_quant)
self.assertEqual(actual_scale.shape, torch.Size([3, 4]))
def test_kimi_handoff_skips_flashinfer_quantizer(self):
x = torch.randn(2, 64, dtype=torch.bfloat16)
packed_topk = torch.zeros(2, 4, dtype=torch.int32)
x_quant = torch.empty(2, 64, dtype=torch.float8_e4m3fn)
x_scale = torch.arange(4, dtype=torch.uint8).reshape(2, 2)
with patch(
"sglang.srt.layers.moe.route_quant_handoff.take",
return_value=(packed_topk, x_quant, x_scale),
), patch(
"sglang.srt.layers.quantization.fp8_utils.flashinfer_mxfp8_quantize",
create=True,
) as quantize:
actual_x, actual_packed, actual_quant, actual_scale = (
_prepare_flashinfer_mxfp8_activations(x, 64)
)
quantize.assert_not_called()
self.assertIs(actual_x, x)
self.assertIs(actual_packed, packed_topk)
self.assertIs(actual_quant, x_quant)
self.assertTrue(torch.equal(actual_scale.view(torch.uint8), x_scale))
if __name__ == "__main__":
unittest.main()