[MoE] dedup triton_kernels backend quant-arg asserts and fill weight dtype guard (#28689)
This commit is contained in:
@@ -19,6 +19,7 @@ from triton_kernels.matmul_ogs import (
|
|||||||
)
|
)
|
||||||
from triton_kernels.numerics import InFlexData
|
from triton_kernels.numerics import InFlexData
|
||||||
from triton_kernels.swiglu import swiglu_fn
|
from triton_kernels.swiglu import swiglu_fn
|
||||||
|
from triton_kernels.tensor import FP4
|
||||||
|
|
||||||
from sglang.srt.utils import is_cuda
|
from sglang.srt.utils import is_cuda
|
||||||
|
|
||||||
@@ -32,6 +33,26 @@ if TYPE_CHECKING:
|
|||||||
from sglang.srt.layers.moe.topk import TopKOutput
|
from sglang.srt.layers.moe.topk import TopKOutput
|
||||||
|
|
||||||
|
|
||||||
|
def _assert_unsupported_quant_args(
|
||||||
|
use_fp8_w8a8: bool,
|
||||||
|
per_channel_quant: bool,
|
||||||
|
expert_map: Optional[torch.Tensor],
|
||||||
|
w1_scale: Optional[torch.Tensor],
|
||||||
|
w2_scale: Optional[torch.Tensor],
|
||||||
|
a1_scale: Optional[torch.Tensor],
|
||||||
|
a2_scale: Optional[torch.Tensor],
|
||||||
|
block_shape: Optional[list[int]],
|
||||||
|
) -> None:
|
||||||
|
assert use_fp8_w8a8 is False, "use_fp8_w8a8 is not supported"
|
||||||
|
assert per_channel_quant is False, "per_channel_quant is not supported"
|
||||||
|
assert expert_map is None, "expert_map is not supported"
|
||||||
|
assert w1_scale is None, "w1_scale is not supported"
|
||||||
|
assert w2_scale is None, "w2_scale is not supported"
|
||||||
|
assert a1_scale is None, "a1_scale is not supported"
|
||||||
|
assert a2_scale is None, "a2_scale is not supported"
|
||||||
|
assert block_shape is None, "block_shape is not supported"
|
||||||
|
|
||||||
|
|
||||||
def quantize(w, dtype, dev, **opt):
|
def quantize(w, dtype, dev, **opt):
|
||||||
if dtype == "bf16":
|
if dtype == "bf16":
|
||||||
return w.to(torch.bfloat16), InFlexData()
|
return w.to(torch.bfloat16), InFlexData()
|
||||||
@@ -105,14 +126,16 @@ def triton_kernel_fused_experts(
|
|||||||
block_shape: Optional[list[int]] = None,
|
block_shape: Optional[list[int]] = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
|
|
||||||
assert use_fp8_w8a8 is False, "use_fp8_w8a8 is not supported"
|
_assert_unsupported_quant_args(
|
||||||
assert per_channel_quant is False, "per_channel_quant is not supported"
|
use_fp8_w8a8,
|
||||||
assert expert_map is None, "expert_map is not supported"
|
per_channel_quant,
|
||||||
assert w1_scale is None, "w1_scale is not supported"
|
expert_map,
|
||||||
assert w2_scale is None, "w2_scale is not supported"
|
w1_scale,
|
||||||
assert a1_scale is None, "a1_scale is not supported"
|
w2_scale,
|
||||||
assert a2_scale is None, "a2_scale is not supported"
|
a1_scale,
|
||||||
assert block_shape is None, "block_shape is not supported"
|
a2_scale,
|
||||||
|
block_shape,
|
||||||
|
)
|
||||||
|
|
||||||
# type check
|
# type check
|
||||||
assert hidden_states.dtype == torch.bfloat16, "hidden_states must be bfloat16"
|
assert hidden_states.dtype == torch.bfloat16, "hidden_states must be bfloat16"
|
||||||
@@ -253,21 +276,24 @@ def triton_kernel_fused_experts_with_bias(
|
|||||||
gemm1_alpha: Optional[float] = None,
|
gemm1_alpha: Optional[float] = None,
|
||||||
gemm1_clamp_limit: Optional[float] = None,
|
gemm1_clamp_limit: Optional[float] = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
assert use_fp8_w8a8 is False, "use_fp8_w8a8 is not supported"
|
_assert_unsupported_quant_args(
|
||||||
assert per_channel_quant is False, "per_channel_quant is not supported"
|
use_fp8_w8a8,
|
||||||
assert expert_map is None, "expert_map is not supported"
|
per_channel_quant,
|
||||||
assert w1_scale is None, "w1_scale is not supported"
|
expert_map,
|
||||||
assert w2_scale is None, "w2_scale is not supported"
|
w1_scale,
|
||||||
assert a1_scale is None, "a1_scale is not supported"
|
w2_scale,
|
||||||
assert a2_scale is None, "a2_scale is not supported"
|
a1_scale,
|
||||||
assert block_shape is None, "block_shape is not supported"
|
a2_scale,
|
||||||
|
block_shape,
|
||||||
|
)
|
||||||
|
|
||||||
# type check
|
# type check
|
||||||
assert hidden_states.dtype == torch.bfloat16, "hidden_states must be bfloat16"
|
assert hidden_states.dtype == torch.bfloat16, "hidden_states must be bfloat16"
|
||||||
for w in (w1, w2):
|
for w in (w1, w2):
|
||||||
# TODO assert bf16 or mxfp4
|
assert w.dtype in (
|
||||||
# assert (w.dtype == torch.bfloat16) or check-is-mxfp4, f"w must be bfloat16 or mxfp4 {w1.dtype=}"
|
torch.bfloat16,
|
||||||
pass
|
FP4,
|
||||||
|
), f"w must be bfloat16 or mxfp4 (FP4), got {w.dtype}"
|
||||||
|
|
||||||
# Shape check
|
# Shape check
|
||||||
assert hidden_states.ndim == 2, "hidden_states must be 2D"
|
assert hidden_states.ndim == 2, "hidden_states must be 2D"
|
||||||
|
|||||||
Reference in New Issue
Block a user