diff --git a/docs/cookbook/autoregressive/DeepSeek/DeepSeek-V4.mdx b/docs/cookbook/autoregressive/DeepSeek/DeepSeek-V4.mdx index 47d1f5c47..f19b39be2 100644 --- a/docs/cookbook/autoregressive/DeepSeek/DeepSeek-V4.mdx +++ b/docs/cookbook/autoregressive/DeepSeek/DeepSeek-V4.mdx @@ -205,6 +205,18 @@ For the original Flash and Pro checkpoints: - `high-throughput`: MTP disabled — at saturation the verify step costs more than it saves. - MTP runs on the v2 speculative path. +**Shared experts fusion (Blackwell, flashinfer_mxfp4)** + +On the Blackwell fp4 recipes (`--moe-runner-backend flashinfer_mxfp4`), the shared expert runs as a separate FP8 MLP on an alternate stream by default. Adding: + +```bash Command +--enforce-shared-experts-fusion +``` + +routes it as one extra MXFP4 expert through the same trtllm-gen MoE kernel, so the whole MoE runs on a single stream (~4 fewer kernel launches and 2 fewer stream syncs per MoE layer). The shared expert is requantized from FP8 to MXFP4 at load time. Measured on GB200 tp4: gsm8k and AIME25 accuracy on par with the unfused baseline; Mean TTFT -13% to -21% and P99 ITL -15% to -53% at QPS 1-8 with neutral throughput. + +Only for deployments without expert parallelism (e.g. the single-node low-latency recipes): with `moe_ep_size > 1` the flag is rejected at startup, unless the DeepEP/MegaMOE per-rank shared-slot path is in use. + **Compressed attention state dtype** DeepSeek-V4 uses hybrid compressed attention for long-context efficiency. `SGLANG_DSV4_COMPRESS_STATE_DTYPE` controls the dtype of the C4 / C128 compressed attention state pools. Supported values are `float32` / `fp32` (default: `float32`) and `bfloat16` / `bf16`. For BF16 on the offline compression path: diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 495d932d1..74eba43ef 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -91,7 +91,10 @@ from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.moe import get_moe_a2a_backend, should_use_dp_reduce_scatterv from sglang.srt.layers.moe.fused_moe_triton import FusedMoE -from sglang.srt.layers.moe.utils import is_shared_experts_fusion_disabled +from sglang.srt.layers.moe.utils import ( + is_shared_experts_fusion_disabled, + uses_per_rank_fused_shared_slots, +) from sglang.srt.layers.quantization.fp8_utils import ( view_aiter_fused_rms_transposed_fp8_scale, ) @@ -3303,6 +3306,13 @@ class DeepseekV4ForCausalLM(nn.Module): "routed experts, so they cannot be fused into the quantized " "routed-expert path." ) + if get_parallel().moe_ep_size > 1 and not uses_per_rank_fused_shared_slots(): + return ( + "Expert parallelism keeps only a slice of the routed experts on " + "each rank, so the fused shared expert cannot be appended to the " + "routed weight tensor (only DeepEP/MegaMOE per-rank shared slots " + "support fusion under EP)." + ) if not get_exec().moe.enforce_shared_experts_fusion: return "Config does not support fused shared expert(s)." if hf_config.n_shared_experts != 1: diff --git a/test/registered/unit/models/test_deepseek_v4_mxfp4_shared_expert_requant.py b/test/registered/unit/models/test_deepseek_v4_mxfp4_shared_expert_requant.py new file mode 100644 index 000000000..1de679c0e --- /dev/null +++ b/test/registered/unit/models/test_deepseek_v4_mxfp4_shared_expert_requant.py @@ -0,0 +1,108 @@ +import unittest + +import torch + +from sglang.srt.layers.quantization.fp8_utils import 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") + +_E2M1_LUT = torch.tensor([0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0]) + + +def _dequant_mxfp4(packed: torch.Tensor, scales: torch.Tensor) -> torch.Tensor: + """Reference dequant of the MXFP4 layout the trtllm-gen kernel (and the DSV4 + checkpoint's routed experts) consume: low nibble = even column, sign in code + bit 3, one e8m0 scale per 32 in-row elements.""" + as_u8 = packed.view(torch.uint8) + codes_lo = (as_u8 & 0x0F).long() + codes_hi = (as_u8 >> 4).long() + + def decode(codes): + magnitudes = _E2M1_LUT[codes & 0x7] # bits 2:0 index the magnitude table + return torch.where(codes >= 8, -magnitudes, magnitudes) + + vals = torch.stack([decode(codes_lo), decode(codes_hi)], dim=-1) + vals = vals.reshape(packed.shape[0], packed.shape[1] * 2) + exponents = scales.view(torch.uint8).float() - 127.0 + return vals * torch.pow(2.0, exponents).repeat_interleave(32, dim=-1) + + +class TestQuantizeBlockFp8WeightToMxfp4(CustomTestCase): + """Pins the wire format of the load-time FP8→MXFP4 shared-expert requant + (FusedMoE._maybe_load_fp8_shared_expert_as_fp4): the produced bytes land in + the same expert tensors as the checkpoint's MXFP4-packed routed experts, so + nibble order / sign bit / e8m0 bias are a contract with the trtllm-gen + kernel, not an implementation detail.""" + + def _requant(self, weight_bf16, block=(128, 128)): + rows, cols = weight_bf16.shape + scale_rows = (rows + block[0] - 1) // block[0] + scale_cols = (cols + block[1] - 1) // block[1] + # Identity block scale: fp8 payload == dequantized value. + fp8_scale = torch.ones(scale_rows, scale_cols, dtype=torch.float32) + fp8_weight = weight_bf16.to(torch.float8_e4m3fn) + return quantize_block_fp8_weight_to_mxfp4(fp8_weight, fp8_scale, list(block)) + + def test_packing_layout_contract(self): + w = torch.zeros(1, 32, dtype=torch.bfloat16) + w[0, 0] = 0.5 # code 1 + w[0, 1] = -3.0 # magnitude idx 5, sign bit -> code 13 + w[0, 2] = 6.0 # code 7 + w[0, 3] = 1.5 # code 3 + packed, scales = self._requant(w) + + self.assertEqual(packed.dtype, torch.int8) + self.assertEqual(packed.shape, (1, 16)) + self.assertEqual(scales.dtype, torch.float8_e8m0fnu) + self.assertEqual(scales.shape, (1, 1)) + # group amax 6.0 -> scale 2**0 -> biased e8m0 exponent 127 + self.assertEqual(scales.view(torch.uint8)[0, 0].item(), 127) + as_u8 = packed.view(torch.uint8) + # byte 0 = code(0.5) | code(-3.0) << 4 ; low nibble is the even column + self.assertEqual(as_u8[0, 0].item(), 0x01 | (0x0D << 4)) + # byte 1 = code(6.0) | code(1.5) << 4 + self.assertEqual(as_u8[0, 1].item(), 0x07 | (0x03 << 4)) + # Zero padding is not asserted byte-exactly: the quantizer encodes 0.0 + # as -0.0 (code 8), which the kernel decodes back to zero. Check the + # padding dequantizes to zero without pinning its sign bit. + deq = _dequant_mxfp4(packed, scales) + self.assertTrue((deq[0, 4:] == 0).all()) + + def test_roundtrip_error_is_mxfp4_sized(self): + torch.manual_seed(0) + w = (torch.randn(64, 128) * 0.1).to(torch.bfloat16) + packed, scales = self._requant(w) + deq = _dequant_mxfp4(packed, scales) + # The fp8 cast itself is lossy; compare against what was actually quantized. + ref = w.to(torch.float8_e4m3fn).float() + rel_err = (deq - ref).norm() / ref.norm() + self.assertLess(rel_err.item(), 0.15) + + def test_exactly_representable_values_roundtrip(self): + w = torch.zeros(1, 64, dtype=torch.bfloat16) + w[0, :32] = 4.0 # amax 4 -> quantizes exactly + w[0, 32:] = 48.0 # amax 48 -> 6 * 2**3, exact + packed, scales = self._requant(w) + deq = _dequant_mxfp4(packed, scales) + torch.testing.assert_close(deq, w.float()) + + def test_block_scale_is_applied(self): + torch.manual_seed(1) + w = (torch.randn(16, 64) * 0.1).to(torch.bfloat16) + fp8_weight = w.to(torch.float8_e4m3fn) + fp8_scale = torch.full((1, 1), 2.0, dtype=torch.float32) + packed, scales = quantize_block_fp8_weight_to_mxfp4( + fp8_weight, fp8_scale, [128, 128] + ) + expected = fp8_weight.float() * 2.0 + deq = _dequant_mxfp4(packed, scales) + rel_err = (deq - expected).norm() / expected.norm() + self.assertLess(rel_err.item(), 0.15) + self.assertEqual(packed.shape, (16, 32)) + self.assertEqual(scales.shape, (16, 2)) + + +if __name__ == "__main__": + unittest.main()