[DSv4] Loading Time Weight Dequant (#27867)

Co-authored-by: Peng Zhang <aniz1905@gmail.com>
This commit is contained in:
laixin
2026-07-07 09:54:34 +08:00
committed by GitHub
co-authored by Peng Zhang
parent 3cbb7568bd
commit 6279805962
5 changed files with 148 additions and 3 deletions
+6 -1
View File
@@ -325,7 +325,10 @@ class ModelConfig:
self.is_fp4_experts: bool = False
if is_deepseek_v4(self.hf_config):
self.is_fp4_experts = envs.SGLANG_DSV4_FP4_EXPERTS.get()
if not envs.SGLANG_DSV4_FP4_EXPERTS.is_set():
if (
not envs.SGLANG_DSV4_FP4_EXPERTS.is_set()
or envs.SGLANG_DSV4_FP4_DEQUANT.is_set()
):
from sglang.srt.configs.deepseek_v4 import try_detect_fp4_experts
detected = try_detect_fp4_experts(self.model_path)
@@ -335,6 +338,8 @@ class ModelConfig:
"Auto-detected DSV4 routed-expert layout: is_fp4_experts=%s",
self.is_fp4_experts,
)
if envs.SGLANG_DSV4_FP4_DEQUANT.get():
envs.SGLANG_DSV4_FP4_DEQUANT.set(self.is_fp4_experts is not None)
# HF config.json inherits topk_group=4 from the V3 template, but
# DSV4 trains with no group limiting (sqrtsoftplus + full-expert
+1
View File
@@ -822,6 +822,7 @@ class Envs:
# Set False when using FP4-to-FP8 converted DeepSeek V4 checkpoint.
SGLANG_DSV4_FP4_EXPERTS = EnvBool(True)
SGLANG_DSV4_FP4_DEQUANT = EnvBool(False)
# Default reasoning_effort for dsv4 chat encoder when request doesn't set it.
# Accepts "", "max", "high" (empty string means unset); other values filtered to None.
SGLANG_DSV4_REASONING_EFFORT = EnvStr("")
@@ -144,6 +144,73 @@ ACTIVATION_SCHEMES = ["static", "dynamic"]
logger = logging.getLogger(__name__)
DSV4_DEQUANT_FP4_TABLE = torch.tensor(
[
0.0,
0.5,
1.0,
1.5,
2.0,
3.0,
4.0,
6.0,
0.0,
-0.5,
-1.0,
-1.5,
-2.0,
-3.0,
-4.0,
-6.0,
],
dtype=torch.float32,
)
def cast_e2m1fn_to_e4m3fn(
x: torch.Tensor, scale: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Casts a tensor from e2m1fn to e4m3fn losslessly.
"""
assert x.dtype == torch.int8
assert x.ndim == 2
out_dim, in_dim = x.size()
in_dim *= 2
fp8_block_size = 128
fp4_block_size = 32
assert in_dim % fp8_block_size == 0 and out_dim % fp8_block_size == 0
assert scale.size(0) == out_dim and scale.size(1) == in_dim // fp4_block_size
x = x.view(torch.uint8)
low = x & 0x0F
high = (x >> 4) & 0x0F
table = DSV4_DEQUANT_FP4_TABLE.to(x.device)
x = torch.stack([table[low.long()], table[high.long()]], dim=-1).flatten(2)
# max_fp4 (6.0) * MAX_OFFSET must fit in e4m3fn (max 448)
# 6.0 * 2^6 = 384 < 448; 6.0 * 2^7 = 768 > 448; so MAX_OFFSET_BITS = 6
MAX_OFFSET_BITS = 6
bOut = out_dim // fp8_block_size
bIn = in_dim // fp8_block_size
# bOut, bIn, 128, 128
x = x.view(bOut, fp8_block_size, bIn, fp8_block_size).transpose(1, 2)
# bOut, bIn, 128*4
scale = scale.float().view(bOut, fp8_block_size, bIn, -1).transpose(1, 2).flatten(2)
## bOut, bIn, 1
scale_max_offset_bits = scale.amax(dim=-1, keepdim=True) / (2**MAX_OFFSET_BITS)
# bOut, bIn, 128*4
offset = scale / scale_max_offset_bits
# bOut, bIn, 128, 128
offset = offset.unflatten(-1, (fp8_block_size, -1)).repeat_interleave(
fp4_block_size, dim=-1
)
x = (x * offset).transpose(1, 2).reshape(out_dim, in_dim)
return x.to(torch.float8_e4m3fn), scale_max_offset_bits.squeeze(-1).to(
torch.float8_e8m0fnu
)
class Fp8Config(QuantizationConfig):
"""Config class for FP8."""
@@ -162,6 +229,7 @@ class Fp8Config(QuantizationConfig):
# DSV4 mxfp4-packed (True) vs converted FP8 (False); injected by
# model_loader from ModelConfig. Default False off the DSV4 path.
self.is_fp4_experts = is_fp4_experts
self.dequant_fp4_to_fp8 = False
self.is_checkpoint_fp8_serialized = is_checkpoint_fp8_serialized
if is_checkpoint_fp8_serialized:
log_info_on_rank0(logger, "Detected fp8 checkpoint.")
@@ -282,6 +350,12 @@ class Fp8Config(QuantizationConfig):
fp8_method = Fp8MoEMethod(self)
if self.is_fp4_experts and self.dequant_fp4_to_fp8:
assert (
get_moe_runner_backend().is_auto()
), f"{get_moe_runner_backend()} is not compatible with SGLANG_DSV4_FP4_DEQUANT=1"
return fp8_method
if self.is_fp4_experts and get_moe_runner_backend().is_marlin():
from sglang.srt.layers.quantization.mxfp4_marlin_moe import (
Mxfp4MarlinMoEMethod,
@@ -860,6 +934,7 @@ class Fp8MoEMethod(FusedMoEMethodBase):
self.use_mxfp8 or self.quant_config.weight_block_size is not None
)
self.is_fp4_expert = self.quant_config.is_fp4_experts
self.dequant_fp4_to_fp8 = self.quant_config.dequant_fp4_to_fp8
self.with_bias = False
if get_moe_runner_backend().is_cutlass():
assert (
@@ -1321,6 +1396,26 @@ class Fp8MoEMethod(FusedMoEMethodBase):
# Check if MoE will actually use DeepGEMM runner
will_use_deepgemm = self.is_deepgemm_moe_runner_backend_enabled()
if self.is_fp4_expert and self.dequant_fp4_to_fp8:
for weight_param, scale_param in [
(layer.w13_weight, layer.w13_weight_scale_inv),
(layer.w2_weight, layer.w2_weight_scale_inv),
]:
num_experts = weight_param.shape[0]
new_weights = []
new_scales = []
for e in range(num_experts):
w, s = cast_e2m1fn_to_e4m3fn(
weight_param.data[e], scale_param.data[e]
)
new_weights.append(w)
new_scales.append(s)
weight_param.data = torch.stack(new_weights)
scale_param.data = torch.stack(new_scales).float()
scale_param.format_ue8m0 = False
self.is_fp4_expert = False
logger.warning_once("Dequantized FP4 expert weights to FP8.")
if self.is_fp4_expert:
if get_moe_runner_backend().is_marlin():
layer.w13_weight.data = layer.w13_weight.data.view(torch.int8)
+1
View File
@@ -247,6 +247,7 @@ def _get_quantization_config(
if isinstance(quant_config, Fp8Config):
quant_config.is_fp4_experts = model_config.is_fp4_experts
quant_config.dequant_fp4_to_fp8 = envs.SGLANG_DSV4_FP4_DEQUANT.get()
# Handle hybrid NVFP4 moe (nvidia/DeepSeek-V4-Pro-NVFP4)
nvfp4_meta = model_config.nvfp4_moe_meta
if nvfp4_meta is not None: