Fixing MXFP8 online quantization pipeline (#31510)
This commit is contained in:
@@ -703,9 +703,11 @@ def fused_experts_none_to_flashinfer_trtllm_fp8(
|
|||||||
|
|
||||||
if quant_info.use_mxfp8:
|
if quant_info.use_mxfp8:
|
||||||
assert quant_info.weight_block_k == 32
|
assert quant_info.weight_block_k == 32
|
||||||
from flashinfer import mxfp8_quantize
|
from sglang.srt.layers.quantization.fp8_utils import (
|
||||||
|
flashinfer_mxfp8_quantize,
|
||||||
|
)
|
||||||
|
|
||||||
a_q, a_sf = mxfp8_quantize(hidden_states, False, backend="cute-dsl")
|
a_q, a_sf = flashinfer_mxfp8_quantize(hidden_states, False)
|
||||||
# FlashInfer TRT-LLM MxFP8 expects token-major activation scales:
|
# FlashInfer TRT-LLM MxFP8 expects token-major activation scales:
|
||||||
# [num_tokens, hidden_size // 32] (no transpose).
|
# [num_tokens, hidden_size // 32] (no transpose).
|
||||||
a_sf_t = a_sf.view(torch.uint8).reshape(hidden_states.shape[0], -1)
|
a_sf_t = a_sf.view(torch.uint8).reshape(hidden_states.shape[0], -1)
|
||||||
|
|||||||
@@ -1798,10 +1798,12 @@ class Fp8MoEMethod(FusedMoEMethodBase):
|
|||||||
weight = weight.contiguous()
|
weight = weight.contiguous()
|
||||||
num_experts, m, k = weight.shape
|
num_experts, m, k = weight.shape
|
||||||
assert k % 32 == 0, f"{k=} must be divisible by 32 for MXFP8"
|
assert k % 32 == 0, f"{k=} must be divisible by 32 for MXFP8"
|
||||||
from flashinfer import mxfp8_quantize
|
from sglang.srt.layers.quantization.fp8_utils import (
|
||||||
|
flashinfer_mxfp8_quantize,
|
||||||
|
)
|
||||||
|
|
||||||
weight_flat = weight.view(-1, k).contiguous()
|
weight_flat = weight.view(-1, k).contiguous()
|
||||||
qweight, scale = mxfp8_quantize(weight_flat, False)
|
qweight, scale = flashinfer_mxfp8_quantize(weight_flat, False)
|
||||||
scale_u8 = (
|
scale_u8 = (
|
||||||
scale.view(torch.uint8).contiguous().view(num_experts, m, k // 32)
|
scale.view(torch.uint8).contiguous().view(num_experts, m, k // 32)
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -393,6 +393,7 @@ if is_blackwell_supported() and is_flashinfer_available():
|
|||||||
input: torch.Tensor,
|
input: torch.Tensor,
|
||||||
_is_sf_swizzled_layout: bool = True,
|
_is_sf_swizzled_layout: bool = True,
|
||||||
alignment: int = 32,
|
alignment: int = 32,
|
||||||
|
backend: str = "cute-dsl",
|
||||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
# Fake mode only needs dtypes and output rank to propagate compile graph.
|
# Fake mode only needs dtypes and output rank to propagate compile graph.
|
||||||
# The scale tensor shape is not consumed before the following fake mm op.
|
# The scale tensor shape is not consumed before the following fake mm op.
|
||||||
@@ -412,12 +413,18 @@ if is_blackwell_supported() and is_flashinfer_available():
|
|||||||
input: torch.Tensor,
|
input: torch.Tensor,
|
||||||
is_sf_swizzled_layout: bool = True,
|
is_sf_swizzled_layout: bool = True,
|
||||||
alignment: int = 32,
|
alignment: int = 32,
|
||||||
|
backend: str = "cute-dsl",
|
||||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
return _raw_flashinfer_mxfp8_quantize(
|
return _raw_flashinfer_mxfp8_quantize(
|
||||||
input,
|
input,
|
||||||
is_sf_swizzled_layout=is_sf_swizzled_layout,
|
is_sf_swizzled_layout=is_sf_swizzled_layout,
|
||||||
alignment=alignment,
|
alignment=alignment,
|
||||||
sf_swizzle_layout=SfLayout.layout_128x4,
|
backend=backend,
|
||||||
|
sf_swizzle_layout=(
|
||||||
|
SfLayout.layout_128x4
|
||||||
|
if is_sf_swizzled_layout
|
||||||
|
else SfLayout.layout_linear
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
@register_custom_op(
|
@register_custom_op(
|
||||||
|
|||||||
@@ -73,7 +73,6 @@ has_triton_kernels = is_triton_kernels_available()
|
|||||||
|
|
||||||
if is_flashinfer_available():
|
if is_flashinfer_available():
|
||||||
from flashinfer import (
|
from flashinfer import (
|
||||||
mxfp8_quantize,
|
|
||||||
nvfp4_block_scale_interleave,
|
nvfp4_block_scale_interleave,
|
||||||
trtllm_fp4_block_scale_moe,
|
trtllm_fp4_block_scale_moe,
|
||||||
)
|
)
|
||||||
@@ -1168,7 +1167,13 @@ class Mxfp4MoEMethod(FusedMoEMethodBase):
|
|||||||
value=0.0,
|
value=0.0,
|
||||||
)
|
)
|
||||||
elif self.flashinfer_mxfp4_moe_precision == "default":
|
elif self.flashinfer_mxfp4_moe_precision == "default":
|
||||||
x_quant, x_scale = mxfp8_quantize(x, False, alignment=self.hidden_size)
|
from sglang.srt.layers.quantization.fp8_utils import (
|
||||||
|
flashinfer_mxfp8_quantize,
|
||||||
|
)
|
||||||
|
|
||||||
|
x_quant, x_scale = flashinfer_mxfp8_quantize(
|
||||||
|
x, False, alignment=self.hidden_size
|
||||||
|
)
|
||||||
x_scale = x_scale.view(torch.float8_e4m3fn).reshape(*x.shape[:-1], -1)
|
x_scale = x_scale.view(torch.float8_e4m3fn).reshape(*x.shape[:-1], -1)
|
||||||
else:
|
else:
|
||||||
raise NotImplementedError()
|
raise NotImplementedError()
|
||||||
|
|||||||
@@ -25,7 +25,7 @@ from sglang.srt.utils.common import is_sm100_supported, next_power_of_2
|
|||||||
_MXFP8_QUANTIZE_BACKEND = "cute-dsl" if is_sm100_supported() else "cuda"
|
_MXFP8_QUANTIZE_BACKEND = "cute-dsl" if is_sm100_supported() else "cuda"
|
||||||
|
|
||||||
if is_flashinfer_available():
|
if is_flashinfer_available():
|
||||||
from flashinfer import mxfp8_quantize, shuffle_matrix_a, shuffle_matrix_sf_a
|
from flashinfer import shuffle_matrix_a, shuffle_matrix_sf_a
|
||||||
from flashinfer.fp4_quantization import block_scale_interleave
|
from flashinfer.fp4_quantization import block_scale_interleave
|
||||||
from flashinfer.fused_moe import trtllm_fp4_block_scale_routed_moe
|
from flashinfer.fused_moe import trtllm_fp4_block_scale_routed_moe
|
||||||
from flashinfer.fused_moe.core import (
|
from flashinfer.fused_moe.core import (
|
||||||
@@ -303,7 +303,11 @@ class Mxfp4FlashinferTrtllmMoEMethod:
|
|||||||
value=0.0,
|
value=0.0,
|
||||||
)
|
)
|
||||||
elif precision == "default":
|
elif precision == "default":
|
||||||
x_quant, x_scale = mxfp8_quantize(
|
from sglang.srt.layers.quantization.fp8_utils import (
|
||||||
|
flashinfer_mxfp8_quantize,
|
||||||
|
)
|
||||||
|
|
||||||
|
x_quant, x_scale = flashinfer_mxfp8_quantize(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
False,
|
False,
|
||||||
alignment=hidden_size,
|
alignment=hidden_size,
|
||||||
|
|||||||
@@ -1786,21 +1786,15 @@ class DeepseekV2AttentionMLA(
|
|||||||
and self.fused_qkv_a_proj_with_mqa.quant_method.quant_config.get_name()
|
and self.fused_qkv_a_proj_with_mqa.quant_method.quant_config.get_name()
|
||||||
in {"awq", "awq_marlin", "moe_wna16"}
|
in {"awq", "awq_marlin", "moe_wna16"}
|
||||||
)
|
)
|
||||||
self.use_min_latency_fused_a_gemm = (
|
self._use_min_latency_fused_a_gemm: bool | None = None
|
||||||
self.has_fused_proj
|
|
||||||
and not get_exec().deterministic.enable_deterministic_inference
|
|
||||||
and not self.is_packed_weight
|
|
||||||
and fused_a_gemm_weight_eligible(self.fused_qkv_a_proj_with_mqa)
|
|
||||||
)
|
|
||||||
self.fused_a_gemm_backend = "auto"
|
self.fused_a_gemm_backend = "auto"
|
||||||
|
|
||||||
self.has_q_b_proj = hasattr(self, "q_b_proj")
|
self.has_q_b_proj = hasattr(self, "q_b_proj")
|
||||||
q_b_proj_verified_shapes = {(2048, 2048), (4096, 2048)}
|
q_b_proj_verified_shapes = {(2048, 2048), (4096, 2048)}
|
||||||
self.use_min_latency_q_b_gemm = (
|
self._q_b_proj_verified_shape = self.has_q_b_proj and (
|
||||||
self.has_q_b_proj
|
tuple(self.q_b_proj.weight.shape) in q_b_proj_verified_shapes
|
||||||
and tuple(self.q_b_proj.weight.shape) in q_b_proj_verified_shapes
|
|
||||||
and fused_a_gemm_weight_eligible(self.q_b_proj)
|
|
||||||
)
|
)
|
||||||
|
self._use_min_latency_q_b_gemm: bool | None = None
|
||||||
|
|
||||||
self.init_mha_forward()
|
self.init_mha_forward()
|
||||||
self.init_mla_forward()
|
self.init_mla_forward()
|
||||||
@@ -2031,7 +2025,14 @@ class DeepseekV2AttentionMLA(
|
|||||||
self, hidden_states: torch.Tensor, forward_batch: ForwardBatch
|
self, hidden_states: torch.Tensor, forward_batch: ForwardBatch
|
||||||
):
|
):
|
||||||
assert self.q_lora_rank is not None
|
assert self.q_lora_rank is not None
|
||||||
if self.use_min_latency_fused_a_gemm:
|
if self._use_min_latency_fused_a_gemm is None:
|
||||||
|
self._use_min_latency_fused_a_gemm = (
|
||||||
|
self.has_fused_proj
|
||||||
|
and not get_exec().deterministic.enable_deterministic_inference
|
||||||
|
and not self.is_packed_weight
|
||||||
|
and fused_a_gemm_weight_eligible(self.fused_qkv_a_proj_with_mqa)
|
||||||
|
)
|
||||||
|
if self._use_min_latency_fused_a_gemm:
|
||||||
return linear_with_fused_a_gemm(
|
return linear_with_fused_a_gemm(
|
||||||
self.fused_qkv_a_proj_with_mqa,
|
self.fused_qkv_a_proj_with_mqa,
|
||||||
hidden_states,
|
hidden_states,
|
||||||
@@ -2040,7 +2041,12 @@ class DeepseekV2AttentionMLA(
|
|||||||
return self.fused_qkv_a_proj_with_mqa(hidden_states)[0]
|
return self.fused_qkv_a_proj_with_mqa(hidden_states)[0]
|
||||||
|
|
||||||
def q_b_proj_forward(self, q_lora: torch.Tensor) -> torch.Tensor:
|
def q_b_proj_forward(self, q_lora: torch.Tensor) -> torch.Tensor:
|
||||||
if self.use_min_latency_q_b_gemm:
|
if self._use_min_latency_q_b_gemm is None:
|
||||||
|
self._use_min_latency_q_b_gemm = (
|
||||||
|
self._q_b_proj_verified_shape
|
||||||
|
and fused_a_gemm_weight_eligible(self.q_b_proj)
|
||||||
|
)
|
||||||
|
if self._use_min_latency_q_b_gemm:
|
||||||
q = linear_with_fused_a_gemm(
|
q = linear_with_fused_a_gemm(
|
||||||
self.q_b_proj, q_lora, backend=self.fused_a_gemm_backend
|
self.q_b_proj, q_lora, backend=self.fused_a_gemm_backend
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -123,7 +123,7 @@ class Eagle3MLADecoderLayer(nn.Module):
|
|||||||
)
|
)
|
||||||
# Recompute fused-proj-dependent flags so they reflect the new input dim.
|
# Recompute fused-proj-dependent flags so they reflect the new input dim.
|
||||||
attn.has_fused_proj = True
|
attn.has_fused_proj = True
|
||||||
attn.use_min_latency_fused_a_gemm = False
|
attn._use_min_latency_fused_a_gemm = False
|
||||||
quant_method = getattr(attn.fused_qkv_a_proj_with_mqa, "quant_method", None)
|
quant_method = getattr(attn.fused_qkv_a_proj_with_mqa, "quant_method", None)
|
||||||
attn.is_packed_weight = (
|
attn.is_packed_weight = (
|
||||||
quant_method is not None
|
quant_method is not None
|
||||||
|
|||||||
Reference in New Issue
Block a user