[FlashInfer v0.6.18] add FlashInfer CuTe DSL NVFP4 W4A16 mode (#35120)
This commit is contained in:
@@ -759,6 +759,11 @@ SGLang supports various environment variables that can be used to configure its
|
|||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Enable FlashInfer TRT-LLM or CuTe DSL v2 (no A2A or FlashInfer A2A) per-token FP32 activation scaling for serialized <code>modelopt_fp4</code> checkpoints; checkpoint activation scales are treated as <code>1</code></td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Enable FlashInfer TRT-LLM or CuTe DSL v2 (no A2A or FlashInfer A2A) per-token FP32 activation scaling for serialized <code>modelopt_fp4</code> checkpoints; checkpoint activation scales are treated as <code>1</code></td>
|
||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>false</code></td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>false</code></td>
|
||||||
</tr>
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_FLASHINFER_CUTEDSL_NVFP4_W4A16</code></td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Use BF16 activations and outputs with FlashInfer CuTe DSL NVFP4 weights on NVIDIA SM100-family GPUs (SM100/SM103). Serialized dense linear layers use this mode with <code>--fp4-gemm-backend flashinfer_cutedsl</code>. MoE layers use it with <code>--moe-runner-backend flashinfer_cutedsl</code>; the MoE path supports online weight quantization and serialized ModelOpt NVFP4 weights with either no A2A or FlashInfer A2A, and honors <code>SGLANG_FLASHINFER_MOE_FUSED_FINALIZE</code>.</td>
|
||||||
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>false</code></td>
|
||||||
|
</tr>
|
||||||
<tr>
|
<tr>
|
||||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_FLASHINFER_MOE_FUSED_FINALIZE</code></td>
|
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_FLASHINFER_MOE_FUSED_FINALIZE</code></td>
|
||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Use FlashInfer's fused atomic CUTLASS and CuTe DSL MoE finalize for best performance. Deterministic inference overrides this to <code>false</code>.</td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Use FlashInfer's fused atomic CUTLASS and CuTe DSL MoE finalize for best performance. Deterministic inference overrides this to <code>false</code>.</td>
|
||||||
|
|||||||
@@ -259,7 +259,17 @@ def handle_a2a_moe(server_args: Any):
|
|||||||
), "Flashinfer MoE A2A is only supported with dp_size == tp_size and --enable-dp-attention"
|
), "Flashinfer MoE A2A is only supported with dp_size == tp_size and --enable-dp-attention"
|
||||||
if cfg.deepep_mode != "auto":
|
if cfg.deepep_mode != "auto":
|
||||||
logger.warning("--deepep-mode is ignored for Flashinfer MoE A2A")
|
logger.warning("--deepep-mode is ignored for Flashinfer MoE A2A")
|
||||||
if not envs.SGLANG_MOE_NVFP4_DISPATCH.is_set() and (
|
use_cutedsl_w4a16 = (
|
||||||
|
resolved_view(server_args).moe_runner_backend == "flashinfer_cutedsl"
|
||||||
|
and envs.SGLANG_FLASHINFER_CUTEDSL_NVFP4_W4A16.get()
|
||||||
|
)
|
||||||
|
if use_cutedsl_w4a16:
|
||||||
|
if envs.SGLANG_MOE_NVFP4_DISPATCH.get():
|
||||||
|
raise ValueError(
|
||||||
|
"CuTe DSL NVFP4 W4A16 requires BF16 FlashInfer MoE "
|
||||||
|
"dispatch; unset SGLANG_MOE_NVFP4_DISPATCH."
|
||||||
|
)
|
||||||
|
elif not envs.SGLANG_MOE_NVFP4_DISPATCH.is_set() and (
|
||||||
resolved_view(server_args).quantization == "modelopt_fp4"
|
resolved_view(server_args).quantization == "modelopt_fp4"
|
||||||
or model_config_of(server_args).nvfp4_moe_meta is not None
|
or model_config_of(server_args).nvfp4_moe_meta is not None
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -945,6 +945,8 @@ class Envs:
|
|||||||
# Enable per-token FP32 activation scaling for serialized ModelOpt FP4 with
|
# Enable per-token FP32 activation scaling for serialized ModelOpt FP4 with
|
||||||
# FlashInfer TRT-LLM or CuTe DSL v2 MoE.
|
# FlashInfer TRT-LLM or CuTe DSL v2 MoE.
|
||||||
SGLANG_FLASHINFER_NVFP4_PER_TOKEN_ACTIVATION = EnvBool(False)
|
SGLANG_FLASHINFER_NVFP4_PER_TOKEN_ACTIVATION = EnvBool(False)
|
||||||
|
# Use BF16 activations with FlashInfer CuTe DSL NVFP4 dense and MoE weights.
|
||||||
|
SGLANG_FLASHINFER_CUTEDSL_NVFP4_W4A16 = EnvBool(False)
|
||||||
# Launch the TRT-LLM MoE grouped GEMMs with PDL only at or below this
|
# Launch the TRT-LLM MoE grouped GEMMs with PDL only at or below this
|
||||||
# token count.
|
# token count.
|
||||||
SGLANG_TRTLLM_MOE_PDL_MAX_TOKENS = EnvInt(8192)
|
SGLANG_TRTLLM_MOE_PDL_MAX_TOKENS = EnvInt(8192)
|
||||||
|
|||||||
@@ -1186,6 +1186,16 @@ def should_apply_lm_head_quant_method(lm_head, quant_method) -> bool:
|
|||||||
# carrying the draft model's stale ModelOpt quant_method. Only use the
|
# carrying the draft model's stale ModelOpt quant_method. Only use the
|
||||||
# ModelOpt lm_head kernel when the runtime quantization state matches it.
|
# ModelOpt lm_head kernel when the runtime quantization state matches it.
|
||||||
if method_name == "ModelOptFp4LinearMethod":
|
if method_name == "ModelOptFp4LinearMethod":
|
||||||
|
if quant_method.quant_mode == "w4a16":
|
||||||
|
return lm_head.weight.dtype == torch.uint8 and _has_lm_head_runtime_attrs(
|
||||||
|
lm_head,
|
||||||
|
(
|
||||||
|
"weight_scale_interleaved",
|
||||||
|
"alpha",
|
||||||
|
"input_size_per_partition",
|
||||||
|
"output_size_per_partition",
|
||||||
|
),
|
||||||
|
)
|
||||||
if lm_head.weight.dtype == torch.int32 and _has_lm_head_runtime_attrs(
|
if lm_head.weight.dtype == torch.int32 and _has_lm_head_runtime_attrs(
|
||||||
lm_head,
|
lm_head,
|
||||||
(
|
(
|
||||||
|
|||||||
@@ -257,7 +257,10 @@ def refresh_cutedsl_standard_scales_for_weight_update(
|
|||||||
w1_alpha, fc2_input_scale, w2_alpha, used_input_scale = (
|
w1_alpha, fc2_input_scale, w2_alpha, used_input_scale = (
|
||||||
resolve_cutedsl_standard_scales(layer)
|
resolve_cutedsl_standard_scales(layer)
|
||||||
)
|
)
|
||||||
if layer.quant_config.use_per_token_activation:
|
if (
|
||||||
|
layer.quant_config.use_per_token_activation
|
||||||
|
and layer._cutedsl_wrapper.quant_mode == "w4a4"
|
||||||
|
):
|
||||||
used_input_scale = _make_per_token_global_scale(used_input_scale)
|
used_input_scale = _make_per_token_global_scale(used_input_scale)
|
||||||
|
|
||||||
new_scales = (w1_alpha, fc2_input_scale, w2_alpha)
|
new_scales = (w1_alpha, fc2_input_scale, w2_alpha)
|
||||||
@@ -319,6 +322,8 @@ def ensure_cutedsl_wrapper(layer: torch.nn.Module) -> None:
|
|||||||
"Install with: pip install flashinfer"
|
"Install with: pip install flashinfer"
|
||||||
) from e
|
) from e
|
||||||
|
|
||||||
|
quant_mode = "w4a16" if envs.SGLANG_FLASHINFER_CUTEDSL_NVFP4_W4A16.get() else "w4a4"
|
||||||
|
|
||||||
assert layer.intermediate_size_per_partition > 0, (
|
assert layer.intermediate_size_per_partition > 0, (
|
||||||
f"CuteDSL MoE: intermediate_size_per_partition must be > 0, "
|
f"CuteDSL MoE: intermediate_size_per_partition must be > 0, "
|
||||||
f"got {layer.intermediate_size_per_partition}. Check EP/TP configuration."
|
f"got {layer.intermediate_size_per_partition}. Check EP/TP configuration."
|
||||||
@@ -360,12 +365,13 @@ def ensure_cutedsl_wrapper(layer: torch.nn.Module) -> None:
|
|||||||
activation_type=_cutedsl_wrapper_activation_type(
|
activation_type=_cutedsl_wrapper_activation_type(
|
||||||
layer.moe_runner_config.activation, ActivationType
|
layer.moe_runner_config.activation, ActivationType
|
||||||
),
|
),
|
||||||
|
quant_mode=quant_mode,
|
||||||
)
|
)
|
||||||
|
|
||||||
w1_alpha, fc2_input_scale, w2_alpha, used_input_scale = (
|
w1_alpha, fc2_input_scale, w2_alpha, used_input_scale = (
|
||||||
resolve_cutedsl_standard_scales(layer)
|
resolve_cutedsl_standard_scales(layer)
|
||||||
)
|
)
|
||||||
if layer.quant_config.use_per_token_activation:
|
if layer.quant_config.use_per_token_activation and quant_mode == "w4a4":
|
||||||
used_input_scale = _make_per_token_global_scale(used_input_scale)
|
used_input_scale = _make_per_token_global_scale(used_input_scale)
|
||||||
layer._cutedsl_scales = (w1_alpha, fc2_input_scale, w2_alpha)
|
layer._cutedsl_scales = (w1_alpha, fc2_input_scale, w2_alpha)
|
||||||
layer._cutedsl_input_scale = used_input_scale
|
layer._cutedsl_input_scale = used_input_scale
|
||||||
@@ -422,6 +428,9 @@ class CuteDslFp4MoeQuantInfo(MoeQuantInfo):
|
|||||||
# v2 only: quantize hidden states with per-token dynamic activation scales.
|
# v2 only: quantize hidden states with per-token dynamic activation scales.
|
||||||
use_per_token_activation: bool = False
|
use_per_token_activation: bool = False
|
||||||
|
|
||||||
|
# v2 only: FlashInfer CuTe DSL activation/weight quantization mode.
|
||||||
|
quant_mode: str = "w4a4"
|
||||||
|
|
||||||
# v1 only: SBO down-GEMM overlap args.
|
# v1 only: SBO down-GEMM overlap args.
|
||||||
down_gemm_overlap_args: Optional[DownGemmOverlapArgs] = None
|
down_gemm_overlap_args: Optional[DownGemmOverlapArgs] = None
|
||||||
|
|
||||||
@@ -461,6 +470,10 @@ def fused_experts_none_to_flashinfer_cutedsl_fp4(
|
|||||||
per_token_activation=True,
|
per_token_activation=True,
|
||||||
backend="cute-dsl",
|
backend="cute-dsl",
|
||||||
)
|
)
|
||||||
|
elif quant_info.quant_mode == "w4a16":
|
||||||
|
x_fp4 = hidden_states
|
||||||
|
x_sf = None
|
||||||
|
per_token_scale = None
|
||||||
else:
|
else:
|
||||||
x_fp4, x_sf = fp4_quantize(
|
x_fp4, x_sf = fp4_quantize(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
@@ -470,11 +483,12 @@ def fused_experts_none_to_flashinfer_cutedsl_fp4(
|
|||||||
)
|
)
|
||||||
per_token_scale = None
|
per_token_scale = None
|
||||||
|
|
||||||
seq_len, hidden_size = hidden_states.shape
|
if quant_info.quant_mode != "w4a16":
|
||||||
x_fp4 = x_fp4.reshape(seq_len, hidden_size // 2)
|
seq_len, hidden_size = hidden_states.shape
|
||||||
x_sf = x_sf.view(torch.float8_e4m3fn).reshape(
|
x_fp4 = x_fp4.reshape(seq_len, hidden_size // 2)
|
||||||
seq_len, hidden_size // _FP4_SF_VEC_SIZE
|
x_sf = x_sf.view(torch.float8_e4m3fn).reshape(
|
||||||
)
|
seq_len, hidden_size // _FP4_SF_VEC_SIZE
|
||||||
|
)
|
||||||
|
|
||||||
output = quant_info.wrapper.run(
|
output = quant_info.wrapper.run(
|
||||||
x=x_fp4,
|
x=x_fp4,
|
||||||
@@ -484,7 +498,9 @@ def fused_experts_none_to_flashinfer_cutedsl_fp4(
|
|||||||
w1_weight=quant_info.w13_weight,
|
w1_weight=quant_info.w13_weight,
|
||||||
w1_weight_sf=quant_info.w13_weight_sf,
|
w1_weight_sf=quant_info.w13_weight_sf,
|
||||||
w1_alpha=quant_info.w1_alpha,
|
w1_alpha=quant_info.w1_alpha,
|
||||||
fc2_input_scale=quant_info.a2_scale,
|
fc2_input_scale=(
|
||||||
|
None if quant_info.quant_mode == "w4a16" else quant_info.a2_scale
|
||||||
|
),
|
||||||
w2_weight=quant_info.w2_weight,
|
w2_weight=quant_info.w2_weight,
|
||||||
w2_weight_sf=quant_info.w2_weight_sf,
|
w2_weight_sf=quant_info.w2_weight_sf,
|
||||||
w2_alpha=quant_info.w2_alpha,
|
w2_alpha=quant_info.w2_alpha,
|
||||||
@@ -541,6 +557,9 @@ def fused_experts_flashinfer_to_flashinfer_cutedsl_fp4(
|
|||||||
# NVFP4 dispatch, inputs are already quantized.
|
# NVFP4 dispatch, inputs are already quantized.
|
||||||
x_fp4 = hidden_states
|
x_fp4 = hidden_states
|
||||||
per_token_scale = None
|
per_token_scale = None
|
||||||
|
elif quant_info.quant_mode == "w4a16":
|
||||||
|
x_fp4 = hidden_states
|
||||||
|
per_token_scale = None
|
||||||
else:
|
else:
|
||||||
if quant_info.use_per_token_activation:
|
if quant_info.use_per_token_activation:
|
||||||
from flashinfer import SfLayout, nvfp4_quantize
|
from flashinfer import SfLayout, nvfp4_quantize
|
||||||
@@ -575,7 +594,9 @@ def fused_experts_flashinfer_to_flashinfer_cutedsl_fp4(
|
|||||||
w1_weight=quant_info.w13_weight,
|
w1_weight=quant_info.w13_weight,
|
||||||
w1_weight_sf=quant_info.w13_weight_sf,
|
w1_weight_sf=quant_info.w13_weight_sf,
|
||||||
w1_alpha=quant_info.w1_alpha,
|
w1_alpha=quant_info.w1_alpha,
|
||||||
fc2_input_scale=quant_info.a2_scale,
|
fc2_input_scale=(
|
||||||
|
None if quant_info.quant_mode == "w4a16" else quant_info.a2_scale
|
||||||
|
),
|
||||||
w2_weight=quant_info.w2_weight,
|
w2_weight=quant_info.w2_weight,
|
||||||
w2_weight_sf=quant_info.w2_weight_sf,
|
w2_weight_sf=quant_info.w2_weight_sf,
|
||||||
w2_alpha=quant_info.w2_alpha,
|
w2_alpha=quant_info.w2_alpha,
|
||||||
|
|||||||
@@ -117,11 +117,12 @@ logger = logging.getLogger(__name__)
|
|||||||
def _sglang_fp4_gemm_fake(
|
def _sglang_fp4_gemm_fake(
|
||||||
input: torch.Tensor,
|
input: torch.Tensor,
|
||||||
weight: torch.Tensor,
|
weight: torch.Tensor,
|
||||||
input_sf: torch.Tensor,
|
input_sf: Optional[torch.Tensor],
|
||||||
weight_sf: torch.Tensor,
|
weight_sf: torch.Tensor,
|
||||||
alpha: torch.Tensor,
|
alpha: torch.Tensor,
|
||||||
out_dtype: torch.dtype,
|
out_dtype: torch.dtype,
|
||||||
out_features: int,
|
out_features: int,
|
||||||
|
quant_mode: str = "w4a4",
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
M = input.shape[-2]
|
M = input.shape[-2]
|
||||||
N = int(out_features)
|
N = int(out_features)
|
||||||
@@ -132,11 +133,12 @@ def _sglang_fp4_gemm_fake(
|
|||||||
def fp4_gemm(
|
def fp4_gemm(
|
||||||
input: torch.Tensor,
|
input: torch.Tensor,
|
||||||
weight: torch.Tensor,
|
weight: torch.Tensor,
|
||||||
input_sf: torch.Tensor,
|
input_sf: Optional[torch.Tensor],
|
||||||
weight_sf: torch.Tensor,
|
weight_sf: torch.Tensor,
|
||||||
alpha: torch.Tensor,
|
alpha: torch.Tensor,
|
||||||
out_dtype: torch.dtype,
|
out_dtype: torch.dtype,
|
||||||
out_features: int,
|
out_features: int,
|
||||||
|
quant_mode: str = "w4a4",
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
if not enable_flashinfer_fp4_gemm:
|
if not enable_flashinfer_fp4_gemm:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
@@ -145,9 +147,23 @@ def fp4_gemm(
|
|||||||
fp4_backend = get_fp4_gemm_runner_backend()
|
fp4_backend = get_fp4_gemm_runner_backend()
|
||||||
# Use the remapping logic to convert SGLang backend names to FlashInfer API names
|
# Use the remapping logic to convert SGLang backend names to FlashInfer API names
|
||||||
backend = fp4_backend.get_flashinfer_backend()
|
backend = fp4_backend.get_flashinfer_backend()
|
||||||
return flashinfer_fp4_gemm(
|
if quant_mode == "w4a4":
|
||||||
input, weight, input_sf, weight_sf, alpha, out_dtype, backend=backend
|
return flashinfer_fp4_gemm(
|
||||||
)
|
input, weight, input_sf, weight_sf, alpha, out_dtype, backend=backend
|
||||||
|
)
|
||||||
|
elif quant_mode == "w4a16":
|
||||||
|
from flashinfer import mm_bf16_fp4
|
||||||
|
|
||||||
|
return mm_bf16_fp4(
|
||||||
|
input,
|
||||||
|
weight,
|
||||||
|
weight_sf,
|
||||||
|
alpha,
|
||||||
|
backend=backend,
|
||||||
|
out_dtype=out_dtype,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unsupported FlashInfer FP4 GEMM quant mode: {quant_mode}")
|
||||||
|
|
||||||
|
|
||||||
if is_cuda() and (not get_platform().is_sm120) and (fp4_quantize is not None):
|
if is_cuda() and (not get_platform().is_sm120) and (fp4_quantize is not None):
|
||||||
@@ -1675,6 +1691,14 @@ class ModelOptFp4LinearMethod(LinearMethodBase):
|
|||||||
|
|
||||||
def __init__(self, quant_config: ModelOptFp4Config):
|
def __init__(self, quant_config: ModelOptFp4Config):
|
||||||
self.quant_config = quant_config
|
self.quant_config = quant_config
|
||||||
|
self.quant_mode = (
|
||||||
|
"w4a16"
|
||||||
|
if (
|
||||||
|
envs.SGLANG_FLASHINFER_CUTEDSL_NVFP4_W4A16.get()
|
||||||
|
and get_fp4_gemm_runner_backend().is_flashinfer_cutedsl()
|
||||||
|
)
|
||||||
|
else "w4a4"
|
||||||
|
)
|
||||||
|
|
||||||
def create_weights(
|
def create_weights(
|
||||||
self,
|
self,
|
||||||
@@ -1770,6 +1794,22 @@ class ModelOptFp4LinearMethod(LinearMethodBase):
|
|||||||
input_scale_2 = layer.input_scale.max().to(torch.float32)
|
input_scale_2 = layer.input_scale.max().to(torch.float32)
|
||||||
weight_scale_2 = layer.weight_scale_2.max().to(torch.float32)
|
weight_scale_2 = layer.weight_scale_2.max().to(torch.float32)
|
||||||
|
|
||||||
|
if self.quant_mode == "w4a16":
|
||||||
|
from flashinfer import prepare_bf16_fp4_weights
|
||||||
|
|
||||||
|
weight, weight_scale, alpha = prepare_bf16_fp4_weights(
|
||||||
|
layer.weight,
|
||||||
|
swizzle_blockscale(layer.weight_scale),
|
||||||
|
weight_scale_2.reshape(1),
|
||||||
|
backend=get_fp4_gemm_runner_backend().get_flashinfer_backend(),
|
||||||
|
)
|
||||||
|
copy_or_rebind_param(layer, "weight", weight)
|
||||||
|
copy_or_rebind_param(layer, "weight_scale_interleaved", weight_scale)
|
||||||
|
copy_or_rebind_param(layer, "alpha", alpha)
|
||||||
|
return
|
||||||
|
elif self.quant_mode != "w4a4":
|
||||||
|
raise ValueError(f"Unsupported FP4 GEMM quant mode: {self.quant_mode}")
|
||||||
|
|
||||||
# alpha / input_scale_inv stay as scalar Parameters. Aliasing them into
|
# alpha / input_scale_inv stay as scalar Parameters. Aliasing them into
|
||||||
# the [N_partitions] source slot breaks fused-QKV linears whose
|
# the [N_partitions] source slot breaks fused-QKV linears whose
|
||||||
# downstream kernels assume scalar input scale.
|
# downstream kernels assume scalar input scale.
|
||||||
@@ -1958,55 +1998,76 @@ class ModelOptFp4LinearMethod(LinearMethodBase):
|
|||||||
bias=bias,
|
bias=bias,
|
||||||
)
|
)
|
||||||
|
|
||||||
# `_accepts_prequantized_fp4` is the explicit opt-in so an accidental
|
if self.quant_mode == "w4a4":
|
||||||
# tuple from unrelated code can't silently bypass quantization.
|
# `_accepts_prequantized_fp4` is the explicit opt-in so an accidental
|
||||||
if getattr(layer, "_accepts_prequantized_fp4", False) and isinstance(x, tuple):
|
# tuple from unrelated code can't silently bypass quantization.
|
||||||
x_fp4, x_scale_interleaved = x
|
if getattr(layer, "_accepts_prequantized_fp4", False) and isinstance(
|
||||||
x_m = x_fp4.shape[0]
|
x, tuple
|
||||||
output_dtype = layer.params_dtype
|
):
|
||||||
else:
|
x_fp4, x_scale_interleaved = x
|
||||||
# NVFP4_AWQ: apply the per-input-channel pre_quant_scale.
|
x_m = x_fp4.shape[0]
|
||||||
|
output_dtype = layer.params_dtype
|
||||||
|
else:
|
||||||
|
# NVFP4_AWQ: apply the per-input-channel pre_quant_scale.
|
||||||
|
if self.quant_config.is_awq:
|
||||||
|
x = x * layer.pre_quant_scale
|
||||||
|
x_fp4, x_scale_interleaved = fp4_quantize(x, layer.input_scale_inv)
|
||||||
|
x_m, _ = x.shape
|
||||||
|
output_dtype = x.dtype
|
||||||
|
|
||||||
|
output_size = layer.output_size_per_partition
|
||||||
|
w_n, _ = layer.weight.shape
|
||||||
|
output_shape = [x_m, output_size]
|
||||||
|
|
||||||
|
assert x_fp4.dtype == torch.uint8
|
||||||
|
assert layer.weight.dtype == torch.uint8
|
||||||
|
assert layer.weight_scale_interleaved.dtype == torch.float8_e4m3fn
|
||||||
|
assert layer.alpha.dtype == torch.float32
|
||||||
|
|
||||||
|
# Pad activations to match weight K-dimension padding
|
||||||
|
weights_padding_cols = getattr(layer, "weights_padding_cols", 0)
|
||||||
|
x_fp4 = pad_nvfp4_activation_for_cutlass(x_fp4, weights_padding_cols)
|
||||||
|
|
||||||
|
w = layer.weight
|
||||||
|
w_scale_interleaved = layer.weight_scale_interleaved
|
||||||
|
if enable_flashinfer_fp4_gemm:
|
||||||
|
w = layer.weight.T
|
||||||
|
w_scale_interleaved = layer.weight_scale_interleaved.T
|
||||||
|
|
||||||
|
out = fp4_gemm(
|
||||||
|
x_fp4,
|
||||||
|
w,
|
||||||
|
x_scale_interleaved,
|
||||||
|
w_scale_interleaved,
|
||||||
|
layer.alpha,
|
||||||
|
output_dtype,
|
||||||
|
w_n,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Slice output to remove N-dimension padding
|
||||||
|
out = slice_nvfp4_output(out, output_size)
|
||||||
|
|
||||||
|
if bias is not None:
|
||||||
|
out = out + bias
|
||||||
|
return out.view(*output_shape)
|
||||||
|
elif self.quant_mode == "w4a16":
|
||||||
if self.quant_config.is_awq:
|
if self.quant_config.is_awq:
|
||||||
x = x * layer.pre_quant_scale
|
x = x * layer.pre_quant_scale
|
||||||
x_fp4, x_scale_interleaved = fp4_quantize(x, layer.input_scale_inv)
|
out = fp4_gemm(
|
||||||
x_m, _ = x.shape
|
x.reshape(-1, x.shape[-1]),
|
||||||
output_dtype = x.dtype
|
layer.weight,
|
||||||
|
None,
|
||||||
output_size = layer.output_size_per_partition
|
layer.weight_scale_interleaved,
|
||||||
w_n, _ = layer.weight.shape
|
layer.alpha,
|
||||||
output_shape = [x_m, output_size]
|
torch.bfloat16,
|
||||||
|
layer.output_size_per_partition,
|
||||||
assert x_fp4.dtype == torch.uint8
|
self.quant_mode,
|
||||||
assert layer.weight.dtype == torch.uint8
|
)
|
||||||
assert layer.weight_scale_interleaved.dtype == torch.float8_e4m3fn
|
if bias is not None:
|
||||||
assert layer.alpha.dtype == torch.float32
|
out = out + bias
|
||||||
|
return out.view(*x.shape[:-1], layer.output_size_per_partition)
|
||||||
# Pad activations to match weight K-dimension padding
|
else:
|
||||||
weights_padding_cols = getattr(layer, "weights_padding_cols", 0)
|
raise ValueError(f"Unsupported FP4 GEMM quant mode: {self.quant_mode}")
|
||||||
x_fp4 = pad_nvfp4_activation_for_cutlass(x_fp4, weights_padding_cols)
|
|
||||||
|
|
||||||
w = layer.weight
|
|
||||||
w_scale_interleaved = layer.weight_scale_interleaved
|
|
||||||
if enable_flashinfer_fp4_gemm:
|
|
||||||
w = layer.weight.T
|
|
||||||
w_scale_interleaved = layer.weight_scale_interleaved.T
|
|
||||||
|
|
||||||
out = fp4_gemm(
|
|
||||||
x_fp4,
|
|
||||||
w,
|
|
||||||
x_scale_interleaved,
|
|
||||||
w_scale_interleaved,
|
|
||||||
layer.alpha,
|
|
||||||
output_dtype,
|
|
||||||
w_n,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Slice output to remove N-dimension padding
|
|
||||||
out = slice_nvfp4_output(out, output_size)
|
|
||||||
|
|
||||||
if bias is not None:
|
|
||||||
out = out + bias
|
|
||||||
return out.view(*output_shape)
|
|
||||||
|
|
||||||
|
|
||||||
def deinterleave_w13(weight: torch.Tensor, *, up_first: bool = False) -> torch.Tensor:
|
def deinterleave_w13(weight: torch.Tensor, *, up_first: bool = False) -> torch.Tensor:
|
||||||
@@ -2498,9 +2559,15 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
|||||||
w13_input_scale = layer.w13_input_scale.max(dim=-1).values.to(torch.float32)
|
w13_input_scale = layer.w13_input_scale.max(dim=-1).values.to(torch.float32)
|
||||||
w2_input_scale = layer.w2_input_scale
|
w2_input_scale = layer.w2_input_scale
|
||||||
|
|
||||||
if self.quant_config.use_per_token_activation:
|
use_cutedsl_w4a16 = (
|
||||||
|
self._is_cutedsl_v2_standard
|
||||||
|
and envs.SGLANG_FLASHINFER_CUTEDSL_NVFP4_W4A16.get()
|
||||||
|
)
|
||||||
|
if self.quant_config.use_per_token_activation or use_cutedsl_w4a16:
|
||||||
# FlashInfer computes activation scales dynamically per token, so
|
# FlashInfer computes activation scales dynamically per token, so
|
||||||
# the static checkpoint activation scale is intentionally neutral.
|
# the static checkpoint activation scale is intentionally neutral.
|
||||||
|
# CuTe DSL W4A16 keeps activations in BF16, so its GEMM alphas must
|
||||||
|
# likewise contain only the NVFP4 weight decode scales.
|
||||||
w13_input_scale = torch.ones_like(w13_input_scale, dtype=torch.float32)
|
w13_input_scale = torch.ones_like(w13_input_scale, dtype=torch.float32)
|
||||||
w2_input_scale = torch.ones_like(w2_input_scale, dtype=torch.float32)
|
w2_input_scale = torch.ones_like(w2_input_scale, dtype=torch.float32)
|
||||||
|
|
||||||
@@ -2566,8 +2633,12 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
|||||||
copy_or_rebind_param(layer, "gemm1_beta", gemm1_beta)
|
copy_or_rebind_param(layer, "gemm1_beta", gemm1_beta)
|
||||||
|
|
||||||
# TODO: for flashinfer always do MOE_NVFP4_DISPATCH
|
# TODO: for flashinfer always do MOE_NVFP4_DISPATCH
|
||||||
use_dispatch_fp4 = not self.quant_config.use_per_token_activation and (
|
use_dispatch_fp4 = (
|
||||||
MOE_NVFP4_DISPATCH or should_use_flashinfer_cutlass_moe_fp4_allgather()
|
not self.quant_config.use_per_token_activation
|
||||||
|
and not use_cutedsl_w4a16
|
||||||
|
and (
|
||||||
|
MOE_NVFP4_DISPATCH or should_use_flashinfer_cutlass_moe_fp4_allgather()
|
||||||
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
layer.dispatcher.set_quant_config(
|
layer.dispatcher.set_quant_config(
|
||||||
@@ -2882,6 +2953,11 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if self._is_cutedsl_v1_deepep:
|
if self._is_cutedsl_v1_deepep:
|
||||||
|
if envs.SGLANG_FLASHINFER_CUTEDSL_NVFP4_W4A16.get():
|
||||||
|
raise ValueError(
|
||||||
|
"SGLANG_FLASHINFER_CUTEDSL_NVFP4_W4A16 does not support "
|
||||||
|
"the CuTe DSL v1 DeepEP masked MoE path."
|
||||||
|
)
|
||||||
# v1 path: DeepEP low-latency + flashinfer_cutedsl_moe_masked.
|
# v1 path: DeepEP low-latency + flashinfer_cutedsl_moe_masked.
|
||||||
# Weights are [Gate, Up] (non-interleaved) with swizzled blockscales.
|
# Weights are [Gate, Up] (non-interleaved) with swizzled blockscales.
|
||||||
quant_info = CuteDslFp4MoeQuantInfo(
|
quant_info = CuteDslFp4MoeQuantInfo(
|
||||||
@@ -2904,6 +2980,7 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
|||||||
# with [Up, Gate] interleaved weights and MMA blockscales.
|
# with [Up, Gate] interleaved weights and MMA blockscales.
|
||||||
ensure_cutedsl_wrapper(layer)
|
ensure_cutedsl_wrapper(layer)
|
||||||
w1_alpha, fc2_input_scale, w2_alpha = layer._cutedsl_scales
|
w1_alpha, fc2_input_scale, w2_alpha = layer._cutedsl_scales
|
||||||
|
quant_mode = layer._cutedsl_wrapper.quant_mode
|
||||||
quant_info = CuteDslFp4MoeQuantInfo(
|
quant_info = CuteDslFp4MoeQuantInfo(
|
||||||
w13_weight=layer.w13_weight,
|
w13_weight=layer.w13_weight,
|
||||||
w2_weight=layer.w2_weight,
|
w2_weight=layer.w2_weight,
|
||||||
@@ -2918,7 +2995,10 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
|||||||
a1_scale=layer._cutedsl_input_scale,
|
a1_scale=layer._cutedsl_input_scale,
|
||||||
a2_scale=fc2_input_scale,
|
a2_scale=fc2_input_scale,
|
||||||
wrapper=layer._cutedsl_wrapper,
|
wrapper=layer._cutedsl_wrapper,
|
||||||
use_per_token_activation=self.quant_config.use_per_token_activation,
|
use_per_token_activation=(
|
||||||
|
self.quant_config.use_per_token_activation and quant_mode == "w4a4"
|
||||||
|
),
|
||||||
|
quant_mode=quant_mode,
|
||||||
)
|
)
|
||||||
return self.runner.run(dispatch_output, quant_info)
|
return self.runner.run(dispatch_output, quant_info)
|
||||||
|
|
||||||
|
|||||||
@@ -773,6 +773,7 @@ class DeepseekV2MoE(nn.Module):
|
|||||||
self.shared_experts.gate_up_proj.quant_method,
|
self.shared_experts.gate_up_proj.quant_method,
|
||||||
ModelOptFp4LinearMethod,
|
ModelOptFp4LinearMethod,
|
||||||
)
|
)
|
||||||
|
and self.shared_experts.gate_up_proj.quant_method.quant_mode == "w4a4"
|
||||||
and isinstance(
|
and isinstance(
|
||||||
self.shared_experts.down_proj.quant_method,
|
self.shared_experts.down_proj.quant_method,
|
||||||
ModelOptFp4LinearMethod,
|
ModelOptFp4LinearMethod,
|
||||||
|
|||||||
@@ -182,6 +182,7 @@ def _maybe_enable_silu_fp4_quant_fusion(mlp: nn.Module) -> None:
|
|||||||
|
|
||||||
if not (
|
if not (
|
||||||
isinstance(mlp.gate_up_proj.quant_method, ModelOptFp4LinearMethod)
|
isinstance(mlp.gate_up_proj.quant_method, ModelOptFp4LinearMethod)
|
||||||
|
and mlp.gate_up_proj.quant_method.quant_mode == "w4a4"
|
||||||
and isinstance(mlp.down_proj.quant_method, ModelOptFp4LinearMethod)
|
and isinstance(mlp.down_proj.quant_method, ModelOptFp4LinearMethod)
|
||||||
):
|
):
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ from sglang.test.test_utils import (
|
|||||||
popen_launch_server,
|
popen_launch_server,
|
||||||
)
|
)
|
||||||
|
|
||||||
register_cuda_ci(est_time=800, stage="nightly", runner_config="4-gpu-b200")
|
register_cuda_ci(est_time=1200, stage="nightly", runner_config="4-gpu-b200")
|
||||||
|
|
||||||
|
|
||||||
class FlashinferNvFp4OnlineMoeBackendBase:
|
class FlashinferNvFp4OnlineMoeBackendBase:
|
||||||
@@ -124,5 +124,47 @@ class TestFlashinferCuteDSLMoeBackendNvFp4Online(
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class TestFlashinferCuteDSLMoeBackendNvFp4OnlineW4A16(
|
||||||
|
FlashinferNvFp4OnlineMoeBackendBase, CustomTestCase
|
||||||
|
):
|
||||||
|
backend = "flashinfer_cutedsl"
|
||||||
|
model = "nvidia/NVIDIA-Nemotron-3-Super-120B-A12B-FP8"
|
||||||
|
extra_args = [
|
||||||
|
"--reasoning-parser",
|
||||||
|
"nemotron_3",
|
||||||
|
"--tool-call-parser",
|
||||||
|
"qwen3_coder",
|
||||||
|
"--speculative-algorithm",
|
||||||
|
"EAGLE",
|
||||||
|
"--speculative-num-steps",
|
||||||
|
"3",
|
||||||
|
"--speculative-eagle-topk",
|
||||||
|
"1",
|
||||||
|
"--speculative-num-draft-tokens",
|
||||||
|
"4",
|
||||||
|
"--dp-size",
|
||||||
|
"4",
|
||||||
|
"--enable-dp-attention",
|
||||||
|
"--enable-dp-lm-head",
|
||||||
|
"--moe-a2a-backend",
|
||||||
|
"flashinfer",
|
||||||
|
"--cuda-graph-backend-prefill",
|
||||||
|
"disabled",
|
||||||
|
]
|
||||||
|
eval_args = {"max_tokens": 16000, "temperature": 1.0, "top_p": 0.95}
|
||||||
|
spec_accept_length_threshold = 2.5
|
||||||
|
extra_env = {
|
||||||
|
"FLASHINFER_NVFP4_4OVER6": "1",
|
||||||
|
"FLASHINFER_NVFP4_4OVER6_ERR_MODE": "MSE",
|
||||||
|
"FLASHINFER_NVFP4_4OVER6_ERR_USE_FAST_MATH": "1",
|
||||||
|
"FLASHINFER_NVFP4_4OVER6_E4M3_USE_256": "1",
|
||||||
|
"SGLANG_FLASHINFER_CUTEDSL_NVFP4_W4A16": "1",
|
||||||
|
"SGLANG_FLASHINFER_NVFP4_PER_TOKEN_ACTIVATION": "0",
|
||||||
|
"SGLANG_FLASHINFER_MOE_FUSED_FINALIZE": "1",
|
||||||
|
"SGLANG_MOE_NVFP4_DISPATCH": "0",
|
||||||
|
"SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "4096",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
|
||||||
register_cuda_ci(est_time=320, stage="extra-b", runner_config="4-gpu-b200")
|
register_cuda_ci(est_time=420, stage="extra-b", runner_config="4-gpu-b200")
|
||||||
|
|
||||||
import time
|
import time
|
||||||
import unittest
|
import unittest
|
||||||
@@ -269,5 +269,37 @@ class TestServerUpdateWeightsFromDiskNVFP4CuteDSL(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestServerUpdateWeightsFromDiskNVFP4W4A16CuteDSL(
|
||||||
|
UpdateWeightsFromDiskBase, CustomTestCase
|
||||||
|
):
|
||||||
|
model = "nvidia/Qwen3-30B-A3B-NVFP4"
|
||||||
|
decode_payload = {**UpdateWeightsFromDiskBase.decode_payload, "routed_dp_rank": 0}
|
||||||
|
launch_env = {
|
||||||
|
"SGLANG_FLASHINFER_CUTEDSL_NVFP4_W4A16": "1",
|
||||||
|
"SGLANG_FLASHINFER_NVFP4_PER_TOKEN_ACTIVATION": "0",
|
||||||
|
"SGLANG_MOE_NVFP4_DISPATCH": "0",
|
||||||
|
"SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "4096",
|
||||||
|
}
|
||||||
|
backend_test_suites = (
|
||||||
|
{
|
||||||
|
"name": "flashinfer_cutedsl_nvfp4_w4a16",
|
||||||
|
"other_args": (
|
||||||
|
"--tp-size",
|
||||||
|
"4",
|
||||||
|
"--dp-size",
|
||||||
|
"4",
|
||||||
|
"--enable-dp-attention",
|
||||||
|
"--ep-size",
|
||||||
|
"4",
|
||||||
|
"--fp4-gemm-backend",
|
||||||
|
"flashinfer_cutedsl",
|
||||||
|
"--moe-runner-backend",
|
||||||
|
"flashinfer_cutedsl",
|
||||||
|
"--enable-deterministic-inference",
|
||||||
|
),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -31,6 +31,7 @@ def test_flashinfer_prefill_returns_standard_combine_input():
|
|||||||
wrapper.run.return_value = expected_output
|
wrapper.run.return_value = expected_output
|
||||||
quant_info = SimpleNamespace(
|
quant_info = SimpleNamespace(
|
||||||
wrapper=wrapper,
|
wrapper=wrapper,
|
||||||
|
quant_mode="w4a4",
|
||||||
use_per_token_activation=False,
|
use_per_token_activation=False,
|
||||||
a1_scale=torch.tensor(1.0),
|
a1_scale=torch.tensor(1.0),
|
||||||
a2_scale=torch.tensor(1.0),
|
a2_scale=torch.tensor(1.0),
|
||||||
|
|||||||
@@ -962,6 +962,20 @@ class TestModelOptMixedPrecisionConfig(CustomTestCase):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_lm_head_guard_accepts_modelopt_fp4_cutedsl_w4a16_runtime_state(self):
|
||||||
|
lm_head = nn.Module()
|
||||||
|
lm_head.weight = nn.Parameter(
|
||||||
|
torch.empty(128, 1024, dtype=torch.uint8), requires_grad=False
|
||||||
|
)
|
||||||
|
lm_head.weight_scale_interleaved = nn.Parameter(torch.empty(1))
|
||||||
|
lm_head.alpha = nn.Parameter(torch.empty(1))
|
||||||
|
lm_head.input_size_per_partition = 2048
|
||||||
|
lm_head.output_size_per_partition = 128
|
||||||
|
quant_method = ModelOptFp4LinearMethod(ModelOptFp4Config())
|
||||||
|
quant_method.quant_mode = "w4a16"
|
||||||
|
|
||||||
|
self.assertTrue(should_apply_lm_head_quant_method(lm_head, quant_method))
|
||||||
|
|
||||||
def test_lm_head_guard_rejects_stale_modelopt_fp4_method_on_dense_head(self):
|
def test_lm_head_guard_rejects_stale_modelopt_fp4_method_on_dense_head(self):
|
||||||
lm_head = nn.Module()
|
lm_head = nn.Module()
|
||||||
lm_head.weight = nn.Parameter(torch.empty(128000, 2048))
|
lm_head.weight = nn.Parameter(torch.empty(128000, 2048))
|
||||||
|
|||||||
Reference in New Issue
Block a user