[MXFP8] Use FlashInfer CUTLASS for dense GEMM on SM120, delete Triton path (#33208)
Co-authored-by: Brayden Zhong <brayden@radixark.ai>
This commit is contained in:
co-authored by
Brayden Zhong
parent
699fcdc936
commit
b3ee679467
@@ -172,7 +172,6 @@ _TRITON_KERNELS = [
|
|||||||
("fp8_kernel", "sglang_per_token_quant_fp8"),
|
("fp8_kernel", "sglang_per_token_quant_fp8"),
|
||||||
("fp8_kernel", "static_quant_fp8"),
|
("fp8_kernel", "static_quant_fp8"),
|
||||||
("fp8_kernel", "w8a8_block_fp8_matmul"),
|
("fp8_kernel", "w8a8_block_fp8_matmul"),
|
||||||
("fp8_kernel", "mxfp8_block_scaled_matmul_triton"),
|
|
||||||
("fp8_kernel", "per_tensor_quant_mla_fp8"),
|
("fp8_kernel", "per_tensor_quant_mla_fp8"),
|
||||||
("fp8_kernel", "per_token_group_quant_mla_deep_gemm_masked_fp8"),
|
("fp8_kernel", "per_token_group_quant_mla_deep_gemm_masked_fp8"),
|
||||||
("fp8_kernel", "per_token_group_quant_fp8_hopper_moe_mn_major"),
|
("fp8_kernel", "per_token_group_quant_fp8_hopper_moe_mn_major"),
|
||||||
|
|||||||
@@ -23,11 +23,6 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
|
|
||||||
try:
|
|
||||||
from triton.tools.tensor_descriptor import TensorDescriptor
|
|
||||||
except:
|
|
||||||
pass
|
|
||||||
|
|
||||||
from sglang.kernels.jit.utils import is_arch_support_pdl
|
from sglang.kernels.jit.utils import is_arch_support_pdl
|
||||||
from sglang.kernels.ops.quantization.fp8_utils import fp8_dtype_to_triton
|
from sglang.kernels.ops.quantization.fp8_utils import fp8_dtype_to_triton
|
||||||
from sglang.srt.layers import deep_gemm_wrapper
|
from sglang.srt.layers import deep_gemm_wrapper
|
||||||
@@ -40,8 +35,6 @@ from sglang.srt.utils import (
|
|||||||
is_cuda,
|
is_cuda,
|
||||||
is_hip,
|
is_hip,
|
||||||
is_musa,
|
is_musa,
|
||||||
is_sm100_supported,
|
|
||||||
is_sm120_supported,
|
|
||||||
log_info_on_rank0,
|
log_info_on_rank0,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.custom_op import register_custom_op
|
from sglang.srt.utils.custom_op import register_custom_op
|
||||||
@@ -51,8 +44,6 @@ _is_hip = is_hip()
|
|||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
_is_cpu = is_cpu()
|
_is_cpu = is_cpu()
|
||||||
_is_musa = is_musa()
|
_is_musa = is_musa()
|
||||||
_is_sm100_supported = is_sm100_supported()
|
|
||||||
_is_sm120_supported = is_sm120_supported()
|
|
||||||
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
||||||
|
|
||||||
if _is_cuda or _is_musa:
|
if _is_cuda or _is_musa:
|
||||||
@@ -1474,207 +1465,6 @@ def w8a8_block_fp8_matmul(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
# Copied and adapted from https://github.com/triton-lang/triton/blob/main/python/tutorials/10-block-scaled-matmul.py
|
|
||||||
@triton.jit
|
|
||||||
def _mxfp8_block_scaled_matmul_kernel( #
|
|
||||||
a_desc, #
|
|
||||||
a_scale_desc, #
|
|
||||||
b_desc, #
|
|
||||||
b_scale_desc, #
|
|
||||||
c_desc, #
|
|
||||||
M: tl.constexpr, #
|
|
||||||
N: tl.constexpr, #
|
|
||||||
K: tl.constexpr, #
|
|
||||||
output_type: tl.constexpr, #
|
|
||||||
BLOCK_M: tl.constexpr, #
|
|
||||||
BLOCK_N: tl.constexpr, #
|
|
||||||
BLOCK_K: tl.constexpr, #
|
|
||||||
rep_m: tl.constexpr, #
|
|
||||||
rep_n: tl.constexpr, #
|
|
||||||
rep_k: tl.constexpr, #
|
|
||||||
NUM_STAGES: tl.constexpr, #
|
|
||||||
): #
|
|
||||||
if output_type == 0:
|
|
||||||
output_dtype = tl.float32
|
|
||||||
elif output_type == 1:
|
|
||||||
output_dtype = tl.float16
|
|
||||||
elif output_type == 2:
|
|
||||||
output_dtype = tl.bfloat16
|
|
||||||
|
|
||||||
pid = tl.program_id(axis=0)
|
|
||||||
num_pid_m = tl.cdiv(M, BLOCK_M)
|
|
||||||
pid_m = pid % num_pid_m
|
|
||||||
pid_n = pid // num_pid_m
|
|
||||||
offs_am = pid_m * BLOCK_M
|
|
||||||
offs_bn = pid_n * BLOCK_N
|
|
||||||
offs_k_a = 0
|
|
||||||
offs_k_b = 0
|
|
||||||
offs_scale_m = pid_m * rep_m
|
|
||||||
offs_scale_n = pid_n * rep_n
|
|
||||||
offs_scale_k = 0
|
|
||||||
|
|
||||||
VEC_SIZE: tl.constexpr = 32
|
|
||||||
|
|
||||||
accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
|
|
||||||
for k in tl.range(0, tl.cdiv(K, BLOCK_K), num_stages=NUM_STAGES):
|
|
||||||
a = a_desc.load([offs_am, offs_k_a])
|
|
||||||
b = b_desc.load([offs_bn, offs_k_b])
|
|
||||||
scale_a = a_scale_desc.load([0, offs_scale_m, offs_scale_k, 0, 0])
|
|
||||||
scale_b = b_scale_desc.load([0, offs_scale_n, offs_scale_k, 0, 0])
|
|
||||||
|
|
||||||
scale_a = (
|
|
||||||
scale_a.reshape(rep_m, rep_k, 32, 4, 4)
|
|
||||||
.trans(0, 3, 2, 1, 4)
|
|
||||||
.reshape(BLOCK_M, BLOCK_K // VEC_SIZE)
|
|
||||||
)
|
|
||||||
scale_b = (
|
|
||||||
scale_b.reshape(rep_n, rep_k, 32, 4, 4)
|
|
||||||
.trans(0, 3, 2, 1, 4)
|
|
||||||
.reshape(BLOCK_N, BLOCK_K // VEC_SIZE)
|
|
||||||
)
|
|
||||||
|
|
||||||
accumulator = tl.dot_scaled(
|
|
||||||
a, scale_a, "e4m3", b.T, scale_b, "e4m3", accumulator
|
|
||||||
)
|
|
||||||
|
|
||||||
offs_k_a += BLOCK_K
|
|
||||||
offs_k_b += BLOCK_K
|
|
||||||
offs_scale_k += rep_k
|
|
||||||
|
|
||||||
c_desc.store([offs_am, offs_bn], accumulator.to(output_dtype))
|
|
||||||
|
|
||||||
|
|
||||||
# Copied and adapted from https://github.com/triton-lang/triton/blob/main/python/tutorials/10-block-scaled-matmul.py
|
|
||||||
def mxfp8_block_scaled_matmul_triton(
|
|
||||||
a: torch.Tensor,
|
|
||||||
a_scale: torch.Tensor,
|
|
||||||
b: torch.Tensor,
|
|
||||||
b_scale: torch.Tensor,
|
|
||||||
output_dtype: torch.dtype,
|
|
||||||
*,
|
|
||||||
block_m: int = 128,
|
|
||||||
block_n: int = 256,
|
|
||||||
block_k: int = 128,
|
|
||||||
num_stages: Optional[int] = None,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
"""Block-scaled matmul for MXFP8 using Triton dot_scaled.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
num_stages: Number of pipeline stages. If None, auto-selects based on GPU:
|
|
||||||
SM120: 1, SM100: 4.
|
|
||||||
"""
|
|
||||||
if num_stages is None:
|
|
||||||
num_stages = 1 if _is_sm120_supported else (4 if _is_sm100_supported else 1)
|
|
||||||
M, K = a.shape
|
|
||||||
N, K_b = b.shape
|
|
||||||
assert K == K_b
|
|
||||||
|
|
||||||
if output_dtype == torch.float32:
|
|
||||||
output_type = 0
|
|
||||||
elif output_dtype == torch.float16:
|
|
||||||
output_type = 1
|
|
||||||
elif output_dtype == torch.bfloat16:
|
|
||||||
output_type = 2
|
|
||||||
else:
|
|
||||||
raise ValueError(f"Unsupported output dtype: {output_dtype}")
|
|
||||||
|
|
||||||
rep_m = block_m // 128
|
|
||||||
rep_n = block_n // 128
|
|
||||||
rep_k = block_k // 32 // 4
|
|
||||||
|
|
||||||
a_desc = TensorDescriptor.from_tensor(a, [block_m, block_k])
|
|
||||||
b_desc = TensorDescriptor.from_tensor(b, [block_n, block_k])
|
|
||||||
|
|
||||||
scale_block_shape = [1, rep_m, rep_k, 2, 256]
|
|
||||||
a_scale_desc = TensorDescriptor.from_tensor(a_scale, block_shape=scale_block_shape)
|
|
||||||
scale_block_shape = [1, rep_n, rep_k, 2, 256]
|
|
||||||
b_scale_desc = TensorDescriptor.from_tensor(b_scale, block_shape=scale_block_shape)
|
|
||||||
|
|
||||||
output = torch.empty((M, N), dtype=output_dtype, device=a.device)
|
|
||||||
c_desc = TensorDescriptor.from_tensor(output, [block_m, block_n])
|
|
||||||
|
|
||||||
grid = (triton.cdiv(M, block_m) * triton.cdiv(N, block_n), 1)
|
|
||||||
_mxfp8_block_scaled_matmul_kernel[grid](
|
|
||||||
a_desc,
|
|
||||||
a_scale_desc,
|
|
||||||
b_desc,
|
|
||||||
b_scale_desc,
|
|
||||||
c_desc,
|
|
||||||
M,
|
|
||||||
N,
|
|
||||||
K,
|
|
||||||
output_type,
|
|
||||||
block_m,
|
|
||||||
block_n,
|
|
||||||
block_k,
|
|
||||||
rep_m,
|
|
||||||
rep_n,
|
|
||||||
rep_k,
|
|
||||||
num_stages,
|
|
||||||
)
|
|
||||||
return output
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
|
||||||
def _pack_mxfp8_scales_kernel(
|
|
||||||
scale_ptr,
|
|
||||||
out_ptr,
|
|
||||||
M: tl.constexpr,
|
|
||||||
K_GROUPS: tl.constexpr,
|
|
||||||
SCALE_K: tl.constexpr,
|
|
||||||
TOTAL: tl.constexpr,
|
|
||||||
BLOCK: tl.constexpr,
|
|
||||||
):
|
|
||||||
offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
|
|
||||||
mask = offs < TOTAL
|
|
||||||
|
|
||||||
idx256 = offs % 256
|
|
||||||
tmp = offs // 256
|
|
||||||
two = tmp % 2
|
|
||||||
tmp = tmp // 2
|
|
||||||
scale_k = tmp % SCALE_K
|
|
||||||
scale_m = tmp // SCALE_K
|
|
||||||
|
|
||||||
within = two * 256 + idx256
|
|
||||||
row_inner_32 = within // 16
|
|
||||||
rem = within - row_inner_32 * 16
|
|
||||||
row_outer_4 = rem // 4
|
|
||||||
k_inner_4 = rem - row_outer_4 * 4
|
|
||||||
|
|
||||||
row = scale_m * 128 + row_outer_4 * 32 + row_inner_32
|
|
||||||
col = scale_k * 4 + k_inner_4
|
|
||||||
value = tl.load(scale_ptr + row * K_GROUPS + col, mask & (row < M), other=127)
|
|
||||||
tl.store(out_ptr + offs, value, mask)
|
|
||||||
|
|
||||||
|
|
||||||
def pack_mxfp8_scales_triton(scale_u8: torch.Tensor) -> torch.Tensor:
|
|
||||||
assert scale_u8.dim() == 2, f"Expected 2D scale tensor, got {scale_u8.dim()}D"
|
|
||||||
scale_u8 = scale_u8.contiguous()
|
|
||||||
m, k_groups = scale_u8.shape
|
|
||||||
assert (
|
|
||||||
k_groups % 4 == 0
|
|
||||||
), f"{k_groups=} must be divisible by 4 (K must be multiple of 128)"
|
|
||||||
|
|
||||||
scale_m = triton.cdiv(m, 128)
|
|
||||||
scale_k = k_groups // 4
|
|
||||||
out = torch.empty(
|
|
||||||
(1, scale_m, scale_k, 2, 256), dtype=scale_u8.dtype, device=scale_u8.device
|
|
||||||
)
|
|
||||||
total = out.numel()
|
|
||||||
block = 1024
|
|
||||||
grid = (triton.cdiv(total, block),)
|
|
||||||
_pack_mxfp8_scales_kernel[grid](
|
|
||||||
scale_u8,
|
|
||||||
out,
|
|
||||||
m,
|
|
||||||
k_groups,
|
|
||||||
scale_k,
|
|
||||||
total,
|
|
||||||
BLOCK=block,
|
|
||||||
)
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
def _per_tensor_quant_mla_fp8_stage1(
|
def _per_tensor_quant_mla_fp8_stage1(
|
||||||
x_ptr,
|
x_ptr,
|
||||||
|
|||||||
@@ -60,11 +60,11 @@ from sglang.srt.layers.quantization.fp8_utils import (
|
|||||||
deepgemm_w8a8_block_fp8_linear_with_fallback,
|
deepgemm_w8a8_block_fp8_linear_with_fallback,
|
||||||
dispatch_w8a8_block_fp8_linear,
|
dispatch_w8a8_block_fp8_linear,
|
||||||
dispatch_w8a8_mxfp8_linear,
|
dispatch_w8a8_mxfp8_linear,
|
||||||
get_fp8_gemm_runner_backend,
|
|
||||||
input_to_float8,
|
input_to_float8,
|
||||||
mxfp8_group_quantize,
|
mxfp8_group_quantize,
|
||||||
normalize_e4m3fn_to_e4m3fnuz,
|
normalize_e4m3fn_to_e4m3fnuz,
|
||||||
requant_block_scale_ue8m0_for_deepgemm,
|
requant_block_scale_ue8m0_for_deepgemm,
|
||||||
|
resolve_mxfp8_dense_gemm_backend,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.quantization.kv_cache import BaseKVCacheMethod
|
from sglang.srt.layers.quantization.kv_cache import BaseKVCacheMethod
|
||||||
from sglang.srt.layers.quantization.marlin_utils_fp8 import prepare_fp8_layer_for_marlin
|
from sglang.srt.layers.quantization.marlin_utils_fp8 import prepare_fp8_layer_for_marlin
|
||||||
@@ -84,8 +84,10 @@ from sglang.srt.runtime_context import get_parallel
|
|||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
get_bool_env_var,
|
get_bool_env_var,
|
||||||
|
is_blackwell_supported,
|
||||||
is_cpu,
|
is_cpu,
|
||||||
is_cuda,
|
is_cuda,
|
||||||
|
is_flashinfer_available,
|
||||||
is_gfx95_supported,
|
is_gfx95_supported,
|
||||||
is_hip,
|
is_hip,
|
||||||
is_musa,
|
is_musa,
|
||||||
@@ -470,7 +472,9 @@ class Fp8LinearMethod(LinearMethodBase):
|
|||||||
self.weight_block_size = self.quant_config.weight_block_size
|
self.weight_block_size = self.quant_config.weight_block_size
|
||||||
self.w8a8_block_fp8_linear = None
|
self.w8a8_block_fp8_linear = None
|
||||||
self.w8a8_mxfp8_linear = None
|
self.w8a8_mxfp8_linear = None
|
||||||
|
self.mxfp8_dense_backend = None
|
||||||
if self.use_mxfp8 and not self.convert_mxfp8_to_block:
|
if self.use_mxfp8 and not self.convert_mxfp8_to_block:
|
||||||
|
self.mxfp8_dense_backend = resolve_mxfp8_dense_gemm_backend()
|
||||||
self.w8a8_mxfp8_linear = dispatch_w8a8_mxfp8_linear()
|
self.w8a8_mxfp8_linear = dispatch_w8a8_mxfp8_linear()
|
||||||
else:
|
else:
|
||||||
self.w8a8_block_fp8_linear = dispatch_w8a8_block_fp8_linear()
|
self.w8a8_block_fp8_linear = dispatch_w8a8_block_fp8_linear()
|
||||||
@@ -730,7 +734,7 @@ class Fp8LinearMethod(LinearMethodBase):
|
|||||||
if not self.use_mxfp8:
|
if not self.use_mxfp8:
|
||||||
return
|
return
|
||||||
|
|
||||||
backend = get_fp8_gemm_runner_backend()
|
backend = self.mxfp8_dense_backend
|
||||||
if backend.is_flashinfer_trtllm():
|
if backend.is_flashinfer_trtllm():
|
||||||
from flashinfer import shuffle_matrix_a, shuffle_matrix_sf_a
|
from flashinfer import shuffle_matrix_a, shuffle_matrix_sf_a
|
||||||
|
|
||||||
@@ -783,13 +787,28 @@ class Fp8LinearMethod(LinearMethodBase):
|
|||||||
"weight_scale_inv_swizzled",
|
"weight_scale_inv_swizzled",
|
||||||
block_scale_interleave(scale_u8.contiguous()).contiguous(),
|
block_scale_interleave(scale_u8.contiguous()).contiguous(),
|
||||||
)
|
)
|
||||||
elif get_fp8_gemm_runner_backend().is_deep_gemm():
|
elif backend.is_deep_gemm():
|
||||||
from sglang.srt.layers.deep_gemm_wrapper.configurer import (
|
from sglang.srt.layers.deep_gemm_wrapper.configurer import (
|
||||||
DEEPGEMM_SCALE_UE8M0,
|
DEEPGEMM_SCALE_UE8M0,
|
||||||
)
|
)
|
||||||
|
|
||||||
n, k = layer.weight.shape
|
n, k = layer.weight.shape
|
||||||
scale_u8 = layer.weight_scale_inv.data
|
scale_u8 = layer.weight_scale_inv.data
|
||||||
|
layer.weight_scale_inv_swizzled = None
|
||||||
|
if n % 64 != 0 or k % 128 != 0:
|
||||||
|
if not (is_blackwell_supported() and is_flashinfer_available()):
|
||||||
|
raise RuntimeError(
|
||||||
|
f"--fp8-gemm-backend=deep_gemm cannot serve MXFP8 weight shape "
|
||||||
|
f"({n}, {k}) (needs N % 64 == 0 and K % 128 == 0), and this "
|
||||||
|
"device has no FlashInfer MXFP8 fallback kernel."
|
||||||
|
)
|
||||||
|
from flashinfer import block_scale_interleave
|
||||||
|
|
||||||
|
copy_or_rebind_param(
|
||||||
|
layer,
|
||||||
|
"weight_scale_inv_swizzled",
|
||||||
|
block_scale_interleave(scale_u8.contiguous()).contiguous(),
|
||||||
|
)
|
||||||
scale_fp32 = (
|
scale_fp32 = (
|
||||||
(scale_u8.contiguous().view(-1).to(torch.int32) << 23)
|
(scale_u8.contiguous().view(-1).to(torch.int32) << 23)
|
||||||
.view(torch.float32)
|
.view(torch.float32)
|
||||||
@@ -807,9 +826,6 @@ class Fp8LinearMethod(LinearMethodBase):
|
|||||||
else:
|
else:
|
||||||
scale_packed = scale_fp32
|
scale_packed = scale_fp32
|
||||||
copy_or_rebind_param(layer, "weight_scale_inv_deepgemm", scale_packed)
|
copy_or_rebind_param(layer, "weight_scale_inv_deepgemm", scale_packed)
|
||||||
else:
|
|
||||||
# Triton path consumes canonical 2D UE8M0 uint8 scales directly.
|
|
||||||
return
|
|
||||||
|
|
||||||
def _quantize_mxfp8_weights(self, layer: Module) -> None:
|
def _quantize_mxfp8_weights(self, layer: Module) -> None:
|
||||||
weight = layer.weight.data
|
weight = layer.weight.data
|
||||||
@@ -959,32 +975,15 @@ class Fp8LinearMethod(LinearMethodBase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if self.use_mxfp8:
|
if self.use_mxfp8:
|
||||||
backend = get_fp8_gemm_runner_backend()
|
backend = self.mxfp8_dense_backend
|
||||||
|
extra_kwargs = {}
|
||||||
if backend.is_flashinfer_cutlass():
|
if backend.is_flashinfer_cutlass():
|
||||||
weight_scale = layer.weight_scale_inv_swizzled
|
weight_scale = layer.weight_scale_inv_swizzled
|
||||||
elif backend.is_flashinfer_trtllm():
|
elif backend.is_flashinfer_trtllm():
|
||||||
weight_scale = layer.weight_scale_inv_shuffled
|
weight_scale = layer.weight_scale_inv_shuffled
|
||||||
elif get_fp8_gemm_runner_backend().is_deep_gemm():
|
elif backend.is_deep_gemm():
|
||||||
weight_scale = getattr(
|
weight_scale = layer.weight_scale_inv_deepgemm
|
||||||
layer, "weight_scale_inv_deepgemm", layer.weight_scale_inv
|
extra_kwargs["weight_scale_swizzled"] = layer.weight_scale_inv_swizzled
|
||||||
)
|
|
||||||
if isinstance(x, tuple):
|
|
||||||
return self.w8a8_mxfp8_linear(
|
|
||||||
input=x[0],
|
|
||||||
weight=layer.weight,
|
|
||||||
weight_scale=weight_scale,
|
|
||||||
input_scale=x[1],
|
|
||||||
bias=bias,
|
|
||||||
weight_scale_fallback=layer.weight_scale_inv,
|
|
||||||
)
|
|
||||||
return self.w8a8_mxfp8_linear(
|
|
||||||
input=x,
|
|
||||||
weight=layer.weight,
|
|
||||||
weight_scale=weight_scale,
|
|
||||||
input_scale=None,
|
|
||||||
bias=bias,
|
|
||||||
weight_scale_fallback=layer.weight_scale_inv,
|
|
||||||
)
|
|
||||||
else:
|
else:
|
||||||
weight_scale = layer.weight_scale_inv
|
weight_scale = layer.weight_scale_inv
|
||||||
if isinstance(x, tuple):
|
if isinstance(x, tuple):
|
||||||
@@ -994,6 +993,7 @@ class Fp8LinearMethod(LinearMethodBase):
|
|||||||
weight_scale=weight_scale,
|
weight_scale=weight_scale,
|
||||||
input_scale=x[1],
|
input_scale=x[1],
|
||||||
bias=bias,
|
bias=bias,
|
||||||
|
**extra_kwargs,
|
||||||
)
|
)
|
||||||
return self.w8a8_mxfp8_linear(
|
return self.w8a8_mxfp8_linear(
|
||||||
input=x,
|
input=x,
|
||||||
@@ -1001,6 +1001,7 @@ class Fp8LinearMethod(LinearMethodBase):
|
|||||||
weight_scale=weight_scale,
|
weight_scale=weight_scale,
|
||||||
input_scale=None,
|
input_scale=None,
|
||||||
bias=bias,
|
bias=bias,
|
||||||
|
**extra_kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.block_quant:
|
if self.block_quant:
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from functools import lru_cache
|
from functools import lru_cache, partial
|
||||||
from typing import TYPE_CHECKING, Callable, List, Optional, Tuple, Union
|
from typing import TYPE_CHECKING, Callable, List, Optional, Tuple, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -24,8 +24,6 @@ from sglang.kernels.ops.quantization.fp8_kernel import (
|
|||||||
fp8_max,
|
fp8_max,
|
||||||
fp8_min,
|
fp8_min,
|
||||||
is_fp8_fnuz,
|
is_fp8_fnuz,
|
||||||
mxfp8_block_scaled_matmul_triton,
|
|
||||||
pack_mxfp8_scales_triton,
|
|
||||||
per_token_group_quant_fp8,
|
per_token_group_quant_fp8,
|
||||||
scaled_fp8_quant,
|
scaled_fp8_quant,
|
||||||
sglang_per_token_quant_fp8,
|
sglang_per_token_quant_fp8,
|
||||||
@@ -60,7 +58,6 @@ _is_hip = is_hip()
|
|||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
_is_fp8_fnuz = is_fp8_fnuz()
|
_is_fp8_fnuz = is_fp8_fnuz()
|
||||||
_is_sm100_supported = is_sm100_supported()
|
_is_sm100_supported = is_sm100_supported()
|
||||||
_is_sm120_supported = is_sm120_supported()
|
|
||||||
_is_gfx95_supported = is_gfx95_supported()
|
_is_gfx95_supported = is_gfx95_supported()
|
||||||
_is_musa = is_musa()
|
_is_musa = is_musa()
|
||||||
|
|
||||||
@@ -308,6 +305,32 @@ class Fp8GemmRunnerBackend(Enum):
|
|||||||
return self == Fp8GemmRunnerBackend.AITER
|
return self == Fp8GemmRunnerBackend.AITER
|
||||||
|
|
||||||
|
|
||||||
|
class Mxfp8DenseGemmBackend(Enum):
|
||||||
|
"""Enum for MXFP8 dense linear backend selection, resolved separately from
|
||||||
|
`Fp8GemmRunnerBackend`."""
|
||||||
|
|
||||||
|
FLASHINFER_CUTLASS = "flashinfer_cutlass"
|
||||||
|
FLASHINFER_TRTLLM = "flashinfer_trtllm"
|
||||||
|
DEEP_GEMM = "deep_gemm"
|
||||||
|
GFX95_DOT_SCALED = "gfx95_dot_scaled"
|
||||||
|
UNSUPPORTED = "unsupported"
|
||||||
|
|
||||||
|
def is_flashinfer_cutlass(self) -> bool:
|
||||||
|
return self == Mxfp8DenseGemmBackend.FLASHINFER_CUTLASS
|
||||||
|
|
||||||
|
def is_flashinfer_trtllm(self) -> bool:
|
||||||
|
return self == Mxfp8DenseGemmBackend.FLASHINFER_TRTLLM
|
||||||
|
|
||||||
|
def is_deep_gemm(self) -> bool:
|
||||||
|
return self == Mxfp8DenseGemmBackend.DEEP_GEMM
|
||||||
|
|
||||||
|
def is_gfx95_dot_scaled(self) -> bool:
|
||||||
|
return self == Mxfp8DenseGemmBackend.GFX95_DOT_SCALED
|
||||||
|
|
||||||
|
def is_unsupported(self) -> bool:
|
||||||
|
return self == Mxfp8DenseGemmBackend.UNSUPPORTED
|
||||||
|
|
||||||
|
|
||||||
FP8_GEMM_RUNNER_BACKEND: Fp8GemmRunnerBackend | None = None
|
FP8_GEMM_RUNNER_BACKEND: Fp8GemmRunnerBackend | None = None
|
||||||
|
|
||||||
|
|
||||||
@@ -495,21 +518,64 @@ def dispatch_w8a8_block_fp8_linear() -> Callable:
|
|||||||
return _dispatch_auto_backend()
|
return _dispatch_auto_backend()
|
||||||
|
|
||||||
|
|
||||||
def dispatch_w8a8_mxfp8_linear() -> Callable:
|
def resolve_mxfp8_dense_gemm_backend() -> Mxfp8DenseGemmBackend:
|
||||||
|
"""Pick the MXFP8 dense linear backend, honoring `--fp8-gemm-backend` only when it
|
||||||
|
names a backend that owns an MXFP8 dense kernel."""
|
||||||
backend = get_fp8_gemm_runner_backend()
|
backend = get_fp8_gemm_runner_backend()
|
||||||
|
|
||||||
|
if backend.is_flashinfer_trtllm():
|
||||||
|
if not (_is_sm100_supported and is_flashinfer_available()):
|
||||||
|
raise RuntimeError(
|
||||||
|
"MXFP8 dense GEMM requested via --fp8-gemm-backend=flashinfer_trtllm, "
|
||||||
|
"but that kernel requires SM100/SM103 GPUs and FlashInfer."
|
||||||
|
)
|
||||||
|
return Mxfp8DenseGemmBackend.FLASHINFER_TRTLLM
|
||||||
|
|
||||||
|
if backend.is_deep_gemm():
|
||||||
|
if not deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM:
|
||||||
|
raise RuntimeError(
|
||||||
|
"MXFP8 dense GEMM requested via --fp8-gemm-backend=deep_gemm, but "
|
||||||
|
"DeepGEMM is not available (package missing or "
|
||||||
|
"SGLANG_ENABLE_JIT_DEEPGEMM=0)."
|
||||||
|
)
|
||||||
|
return Mxfp8DenseGemmBackend.DEEP_GEMM
|
||||||
|
|
||||||
|
if _is_hip and _is_gfx95_supported:
|
||||||
|
return Mxfp8DenseGemmBackend.GFX95_DOT_SCALED
|
||||||
|
|
||||||
|
if is_blackwell_supported() and is_flashinfer_available():
|
||||||
|
return Mxfp8DenseGemmBackend.FLASHINFER_CUTLASS
|
||||||
|
|
||||||
|
if deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM:
|
||||||
|
return Mxfp8DenseGemmBackend.DEEP_GEMM
|
||||||
|
|
||||||
|
return Mxfp8DenseGemmBackend.UNSUPPORTED
|
||||||
|
|
||||||
|
|
||||||
|
def _unsupported_mxfp8_linear(*args, **kwargs) -> torch.Tensor:
|
||||||
|
raise RuntimeError(
|
||||||
|
"No MXFP8 dense GEMM kernel is available on this device. MXFP8 dense linear "
|
||||||
|
"requires Blackwell (SM100/SM103/SM110/SM120) with FlashInfer, Hopper (SM90) "
|
||||||
|
"with DeepGEMM, or ROCm gfx95."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def dispatch_w8a8_mxfp8_linear() -> Callable:
|
||||||
|
backend = resolve_mxfp8_dense_gemm_backend()
|
||||||
if backend.is_deep_gemm():
|
if backend.is_deep_gemm():
|
||||||
return _deepgemm_w8a8_mxfp8_linear_with_fallback
|
return _deepgemm_w8a8_mxfp8_linear_with_fallback
|
||||||
elif backend.is_flashinfer_cutlass() or backend.is_flashinfer_trtllm():
|
elif backend.is_flashinfer_trtllm():
|
||||||
return flashinfer_mxfp8_blockscaled_linear
|
return partial(flashinfer_mxfp8_blockscaled_linear, backend="trtllm")
|
||||||
elif backend.is_triton():
|
elif backend.is_flashinfer_cutlass():
|
||||||
return triton_mxfp8_blockscaled_linear
|
return partial(flashinfer_mxfp8_blockscaled_linear, backend="cutlass")
|
||||||
elif _is_hip and _is_gfx95_supported:
|
elif backend.is_unsupported():
|
||||||
|
return _unsupported_mxfp8_linear
|
||||||
|
|
||||||
from sglang.kernels.ops.quantization.mxfp8_amd_gfx95 import (
|
from sglang.kernels.ops.quantization.mxfp8_amd_gfx95 import (
|
||||||
dot_scaled_mxfp8_blockscaled_linear,
|
dot_scaled_mxfp8_blockscaled_linear,
|
||||||
)
|
)
|
||||||
|
|
||||||
return dot_scaled_mxfp8_blockscaled_linear
|
return dot_scaled_mxfp8_blockscaled_linear
|
||||||
return triton_mxfp8_blockscaled_linear
|
|
||||||
|
|
||||||
|
|
||||||
def _deepgemm_w8a8_mxfp8_linear_with_fallback(
|
def _deepgemm_w8a8_mxfp8_linear_with_fallback(
|
||||||
@@ -518,7 +584,7 @@ def _deepgemm_w8a8_mxfp8_linear_with_fallback(
|
|||||||
weight_scale: torch.Tensor,
|
weight_scale: torch.Tensor,
|
||||||
input_scale: Optional[torch.Tensor] = None,
|
input_scale: Optional[torch.Tensor] = None,
|
||||||
bias: Optional[torch.Tensor] = None,
|
bias: Optional[torch.Tensor] = None,
|
||||||
weight_scale_fallback: Optional[torch.Tensor] = None,
|
weight_scale_swizzled: Optional[torch.Tensor] = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
from sglang.kernels.ops.quantization.fp8_kernel import (
|
from sglang.kernels.ops.quantization.fp8_kernel import (
|
||||||
sglang_per_token_group_quant_fp8,
|
sglang_per_token_group_quant_fp8,
|
||||||
@@ -532,8 +598,18 @@ def _deepgemm_w8a8_mxfp8_linear_with_fallback(
|
|||||||
dtype_supported = output_dtype == torch.bfloat16
|
dtype_supported = output_dtype == torch.bfloat16
|
||||||
|
|
||||||
if not (shape_supported and dtype_supported):
|
if not (shape_supported and dtype_supported):
|
||||||
return triton_mxfp8_blockscaled_linear(
|
if weight_scale_swizzled is None:
|
||||||
input, weight, weight_scale_fallback, input_scale, bias
|
raise RuntimeError(
|
||||||
|
f"DeepGEMM cannot serve this MXFP8 GEMM ({shape_supported=}, "
|
||||||
|
f"{dtype_supported=}) and no FlashInfer fallback scale was prepared "
|
||||||
|
"for this layer. Re-run with --fp8-gemm-backend=flashinfer_cutlass."
|
||||||
|
)
|
||||||
|
return flashinfer_mxfp8_blockscaled_linear(
|
||||||
|
input=input,
|
||||||
|
weight=weight,
|
||||||
|
weight_scale=weight_scale_swizzled,
|
||||||
|
input_scale=input_scale,
|
||||||
|
bias=bias,
|
||||||
)
|
)
|
||||||
|
|
||||||
input_2d = input.view(-1, input.shape[-1])
|
input_2d = input.view(-1, input.shape[-1])
|
||||||
@@ -1132,197 +1208,6 @@ def mxfp8_group_quantize(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
|||||||
return q_input.contiguous(), scale_u8.contiguous()
|
return q_input.contiguous(), scale_u8.contiguous()
|
||||||
|
|
||||||
|
|
||||||
def _pack_mxfp8_scales(scale_u8: torch.Tensor) -> torch.Tensor:
|
|
||||||
if (
|
|
||||||
_is_hip
|
|
||||||
and _is_gfx95_supported
|
|
||||||
and scale_u8.is_cuda
|
|
||||||
and scale_u8.shape[0] % 128 != 0
|
|
||||||
):
|
|
||||||
return pack_mxfp8_scales_triton(scale_u8)
|
|
||||||
|
|
||||||
# Pack (M, K//32) UE8M0 scales into the layout expected by tl.dot_scaled.
|
|
||||||
assert scale_u8.dim() == 2, f"Expected 2D scale tensor, got {scale_u8.dim()}D"
|
|
||||||
scale_u8 = scale_u8.contiguous()
|
|
||||||
m, k_groups = scale_u8.shape
|
|
||||||
assert (
|
|
||||||
k_groups % 4 == 0
|
|
||||||
), f"{k_groups=} must be divisible by 4 (K must be multiple of 128)"
|
|
||||||
|
|
||||||
scale_m = ceil_div(m, 128)
|
|
||||||
if m % 128 != 0:
|
|
||||||
pad_rows = scale_m * 128 - m
|
|
||||||
pad = torch.full(
|
|
||||||
(pad_rows, k_groups),
|
|
||||||
127,
|
|
||||||
dtype=scale_u8.dtype,
|
|
||||||
device=scale_u8.device,
|
|
||||||
)
|
|
||||||
scale_u8 = torch.cat([scale_u8, pad], dim=0)
|
|
||||||
|
|
||||||
scale_k = k_groups // 4
|
|
||||||
scale_u8 = scale_u8.view(scale_m, 128, scale_k, 4)
|
|
||||||
scale_u8 = scale_u8.view(scale_m, 4, 32, scale_k, 4)
|
|
||||||
packed = scale_u8.permute(0, 3, 2, 1, 4).contiguous()
|
|
||||||
return packed.view(1, scale_m, scale_k, 2, 256)
|
|
||||||
|
|
||||||
|
|
||||||
@register_custom_op(
|
|
||||||
op_name="triton_mxfp8_block_scaled_matmul",
|
|
||||||
mutates_args=[],
|
|
||||||
fake_impl=lambda a, a_scale, b, b_scale, output_dtype, block_m=128, block_n=256, block_k=128, num_stages=None: ( # noqa: E501
|
|
||||||
a.new_empty((a.shape[0], b.shape[0]), dtype=output_dtype)
|
|
||||||
),
|
|
||||||
)
|
|
||||||
def triton_mxfp8_block_scaled_matmul(
|
|
||||||
a: torch.Tensor,
|
|
||||||
a_scale: torch.Tensor,
|
|
||||||
b: torch.Tensor,
|
|
||||||
b_scale: torch.Tensor,
|
|
||||||
output_dtype: torch.dtype,
|
|
||||||
*,
|
|
||||||
block_m: int = 128,
|
|
||||||
block_n: int = 256,
|
|
||||||
block_k: int = 128,
|
|
||||||
num_stages: Optional[int] = None,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
"""Opaque custom op wrapper to prevent Dynamo tracing Triton grid math."""
|
|
||||||
return mxfp8_block_scaled_matmul_triton(
|
|
||||||
a,
|
|
||||||
a_scale,
|
|
||||||
b,
|
|
||||||
b_scale,
|
|
||||||
output_dtype=output_dtype,
|
|
||||||
block_m=block_m,
|
|
||||||
block_n=block_n,
|
|
||||||
block_k=block_k,
|
|
||||||
num_stages=num_stages,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _raw_triton_mxfp8_blockscaled_linear(
|
|
||||||
input: torch.Tensor,
|
|
||||||
weight: torch.Tensor,
|
|
||||||
weight_scale: torch.Tensor,
|
|
||||||
input_scale: Optional[torch.Tensor] = None,
|
|
||||||
bias: Optional[torch.Tensor] = None,
|
|
||||||
output_dtype: Optional[torch.dtype] = None,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
if not (
|
|
||||||
(_is_cuda and (_is_sm100_supported or _is_sm120_supported))
|
|
||||||
or (_is_hip and _is_gfx95_supported)
|
|
||||||
):
|
|
||||||
raise RuntimeError(
|
|
||||||
"MXFP8 dense linear requires Blackwell GPUs (SM100/SM120) or ROCm gfx95."
|
|
||||||
)
|
|
||||||
|
|
||||||
input_2d = input.view(-1, input.shape[-1]).contiguous()
|
|
||||||
output_shape = [*input.shape[:-1], weight.shape[0]]
|
|
||||||
|
|
||||||
block_m = 128
|
|
||||||
block_n = 256 if weight.shape[0] % 256 == 0 else 128
|
|
||||||
block_k = 128
|
|
||||||
|
|
||||||
m, k = input_2d.shape
|
|
||||||
n, k_w = weight.shape
|
|
||||||
assert k == k_w, f"{k=} does not match {k_w=}"
|
|
||||||
assert k % 128 == 0, f"{k=} must be divisible by 128 for MXFP8"
|
|
||||||
assert n % block_n == 0, f"{n=} must be divisible by {block_n}"
|
|
||||||
assert weight.dtype == torch.float8_e4m3fn, "MXFP8 weight must be FP8 E4M3."
|
|
||||||
assert weight_scale.dtype == torch.uint8, "MXFP8 weight_scale must be UE8M0 uint8."
|
|
||||||
assert weight_scale.dim() in (
|
|
||||||
2,
|
|
||||||
5,
|
|
||||||
), (
|
|
||||||
"MXFP8 weight_scale must be canonical 2D or packed 5D, "
|
|
||||||
f"got {weight_scale.dim()}D."
|
|
||||||
)
|
|
||||||
|
|
||||||
if input_scale is None:
|
|
||||||
q_input, x_scale_u8 = mxfp8_group_quantize(input_2d)
|
|
||||||
else:
|
|
||||||
q_input = input_2d
|
|
||||||
x_scale_u8 = input_scale
|
|
||||||
assert x_scale_u8.dtype == torch.uint8, "MXFP8 input_scale must be UE8M0 uint8."
|
|
||||||
assert x_scale_u8.shape == (m, k // 32)
|
|
||||||
|
|
||||||
if output_dtype is None:
|
|
||||||
if input_2d.dtype in (torch.float16, torch.bfloat16, torch.float32):
|
|
||||||
output_dtype = input_2d.dtype
|
|
||||||
else:
|
|
||||||
output_dtype = torch.bfloat16
|
|
||||||
|
|
||||||
if m % block_m != 0:
|
|
||||||
pad_rows = ceil_div(m, block_m) * block_m - m
|
|
||||||
q_input = torch.cat(
|
|
||||||
[
|
|
||||||
q_input,
|
|
||||||
torch.zeros((pad_rows, k), device=q_input.device, dtype=q_input.dtype),
|
|
||||||
],
|
|
||||||
dim=0,
|
|
||||||
)
|
|
||||||
pad_scale = torch.full(
|
|
||||||
(pad_rows, k // 32),
|
|
||||||
127,
|
|
||||||
device=x_scale_u8.device,
|
|
||||||
dtype=x_scale_u8.dtype,
|
|
||||||
)
|
|
||||||
x_scale_u8 = torch.cat([x_scale_u8, pad_scale], dim=0)
|
|
||||||
|
|
||||||
a_scale_packed = _pack_mxfp8_scales(x_scale_u8)
|
|
||||||
b_scale_packed = (
|
|
||||||
weight_scale.contiguous()
|
|
||||||
if weight_scale.dim() == 5
|
|
||||||
else _pack_mxfp8_scales(weight_scale)
|
|
||||||
)
|
|
||||||
|
|
||||||
num_stages = 1 if _is_sm120_supported else (4 if _is_sm100_supported else 1)
|
|
||||||
output = triton_mxfp8_block_scaled_matmul(
|
|
||||||
q_input,
|
|
||||||
a_scale_packed,
|
|
||||||
weight.contiguous(),
|
|
||||||
b_scale_packed,
|
|
||||||
output_dtype=output_dtype,
|
|
||||||
block_m=block_m,
|
|
||||||
block_n=block_n,
|
|
||||||
block_k=block_k,
|
|
||||||
num_stages=num_stages,
|
|
||||||
)
|
|
||||||
output = output[:m, :]
|
|
||||||
if bias is not None:
|
|
||||||
output += bias
|
|
||||||
return output.to(dtype=output_dtype).view(*output_shape)
|
|
||||||
|
|
||||||
|
|
||||||
@register_custom_op(
|
|
||||||
op_name="triton_mxfp8_blockscaled_linear",
|
|
||||||
mutates_args=[],
|
|
||||||
fake_impl=lambda input, weight, weight_scale, input_scale=None, bias=None, output_dtype=None: (
|
|
||||||
input.new_empty(
|
|
||||||
(*input.shape[:-1], weight.shape[0]),
|
|
||||||
dtype=(output_dtype if output_dtype is not None else input.dtype),
|
|
||||||
)
|
|
||||||
),
|
|
||||||
)
|
|
||||||
def triton_mxfp8_blockscaled_linear(
|
|
||||||
input: torch.Tensor,
|
|
||||||
weight: torch.Tensor,
|
|
||||||
weight_scale: torch.Tensor,
|
|
||||||
input_scale: Optional[torch.Tensor] = None,
|
|
||||||
bias: Optional[torch.Tensor] = None,
|
|
||||||
output_dtype: Optional[torch.dtype] = None,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
"""Opaque custom-op wrapper to prevent Dynamo guards on MXFP8 padding branches."""
|
|
||||||
return _raw_triton_mxfp8_blockscaled_linear(
|
|
||||||
input=input,
|
|
||||||
weight=weight,
|
|
||||||
weight_scale=weight_scale,
|
|
||||||
input_scale=input_scale,
|
|
||||||
bias=bias,
|
|
||||||
output_dtype=output_dtype,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def flashinfer_mxfp8_blockscaled_linear(
|
def flashinfer_mxfp8_blockscaled_linear(
|
||||||
input: torch.Tensor,
|
input: torch.Tensor,
|
||||||
weight: torch.Tensor,
|
weight: torch.Tensor,
|
||||||
@@ -1330,13 +1215,15 @@ def flashinfer_mxfp8_blockscaled_linear(
|
|||||||
input_scale: Optional[torch.Tensor] = None,
|
input_scale: Optional[torch.Tensor] = None,
|
||||||
bias: Optional[torch.Tensor] = None,
|
bias: Optional[torch.Tensor] = None,
|
||||||
output_dtype: Optional[torch.dtype] = None,
|
output_dtype: Optional[torch.dtype] = None,
|
||||||
|
backend: str = "cutlass",
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""MXFP8 dense linear via FlashInfer mm_mxfp8."""
|
"""MXFP8 dense linear via FlashInfer mm_mxfp8. `weight_scale` must be the layout
|
||||||
input_2d = input.view(-1, input.shape[-1]).contiguous()
|
the backend expects, prepared at load time."""
|
||||||
|
input_2d = input.view(-1, input.shape[-1])
|
||||||
output_shape = [*input.shape[:-1], weight.shape[0]]
|
output_shape = [*input.shape[:-1], weight.shape[0]]
|
||||||
|
|
||||||
m, k = input_2d.shape
|
k = input_2d.shape[1]
|
||||||
n, k_w = weight.shape
|
k_w = weight.shape[1]
|
||||||
if k != k_w:
|
if k != k_w:
|
||||||
raise ValueError(f"Input K={k} does not match weight K={k_w}.")
|
raise ValueError(f"Input K={k} does not match weight K={k_w}.")
|
||||||
if k % 32 != 0:
|
if k % 32 != 0:
|
||||||
@@ -1350,7 +1237,7 @@ def flashinfer_mxfp8_blockscaled_linear(
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
q_input = input_2d
|
q_input = input_2d
|
||||||
x_scale_u8 = input_scale.contiguous()
|
x_scale_u8 = input_scale
|
||||||
|
|
||||||
if output_dtype is None:
|
if output_dtype is None:
|
||||||
if input_2d.dtype in (torch.float16, torch.bfloat16, torch.float32):
|
if input_2d.dtype in (torch.float16, torch.bfloat16, torch.float32):
|
||||||
@@ -1358,34 +1245,19 @@ def flashinfer_mxfp8_blockscaled_linear(
|
|||||||
else:
|
else:
|
||||||
output_dtype = torch.bfloat16
|
output_dtype = torch.bfloat16
|
||||||
|
|
||||||
# Ensure transposed tensors are contiguous for FlashInfer's internal runner.
|
if backend == "trtllm":
|
||||||
weight_t = weight.contiguous().t()
|
weight_scale_t = weight_scale.view(-1)
|
||||||
|
else:
|
||||||
|
weight_scale_t = weight_scale.t() if weight_scale.ndim == 2 else weight_scale
|
||||||
|
|
||||||
if get_fp8_gemm_runner_backend().is_flashinfer_trtllm():
|
|
||||||
weight_scale_t = weight_scale.contiguous().view(-1)
|
|
||||||
output = flashinfer_mm_mxfp8(
|
output = flashinfer_mm_mxfp8(
|
||||||
q_input,
|
q_input,
|
||||||
weight_t,
|
weight.t(),
|
||||||
x_scale_u8,
|
x_scale_u8,
|
||||||
weight_scale_t,
|
weight_scale_t,
|
||||||
out_dtype=output_dtype,
|
out_dtype=output_dtype,
|
||||||
use_8x4_sf_layout=False,
|
use_8x4_sf_layout=False,
|
||||||
backend="trtllm",
|
backend=backend,
|
||||||
)
|
|
||||||
elif get_fp8_gemm_runner_backend().is_flashinfer_cutlass():
|
|
||||||
weight_scale_t = (
|
|
||||||
weight_scale.contiguous().t()
|
|
||||||
if weight_scale.ndim == 2
|
|
||||||
else weight_scale.contiguous()
|
|
||||||
)
|
|
||||||
output = flashinfer_mm_mxfp8(
|
|
||||||
q_input,
|
|
||||||
weight_t,
|
|
||||||
x_scale_u8,
|
|
||||||
weight_scale_t,
|
|
||||||
out_dtype=output_dtype,
|
|
||||||
use_8x4_sf_layout=False,
|
|
||||||
backend="cutlass",
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if bias is not None:
|
if bias is not None:
|
||||||
|
|||||||
@@ -17,9 +17,8 @@ from sglang.srt.layers.moe.topk import TopKConfig, select_experts
|
|||||||
from sglang.srt.layers.quantization.fp8_utils import (
|
from sglang.srt.layers.quantization.fp8_utils import (
|
||||||
input_to_float8,
|
input_to_float8,
|
||||||
mxfp8_group_quantize,
|
mxfp8_group_quantize,
|
||||||
triton_mxfp8_blockscaled_linear,
|
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import is_sm100_supported, is_sm120_supported
|
from sglang.srt.utils import is_blackwell_supported, is_flashinfer_available
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
_is_cuda = torch.cuda.is_available() and torch.version.cuda
|
_is_cuda = torch.cuda.is_available() and torch.version.cuda
|
||||||
@@ -452,11 +451,20 @@ class TestMXFP8DenseLinear(CustomTestCase):
|
|||||||
def setUpClass(cls):
|
def setUpClass(cls):
|
||||||
if not torch.cuda.is_available():
|
if not torch.cuda.is_available():
|
||||||
raise unittest.SkipTest("CUDA is not available")
|
raise unittest.SkipTest("CUDA is not available")
|
||||||
if not (is_sm100_supported() or is_sm120_supported()):
|
if not (is_blackwell_supported() and is_flashinfer_available()):
|
||||||
raise unittest.SkipTest("MXFP8 requires Blackwell (SM100/SM120)")
|
raise unittest.SkipTest(
|
||||||
|
"MXFP8 dense linear requires Blackwell + FlashInfer"
|
||||||
|
)
|
||||||
torch.set_default_device("cuda")
|
torch.set_default_device("cuda")
|
||||||
|
|
||||||
def _mxfp8_dense_linear(self, M, NK, dtype, seed):
|
def _mxfp8_dense_linear(self, M, NK, dtype, seed):
|
||||||
|
from flashinfer import block_scale_interleave
|
||||||
|
|
||||||
|
from sglang.srt.layers.quantization.fp8_utils import (
|
||||||
|
flashinfer_mxfp8_blockscaled_linear,
|
||||||
|
flashinfer_mxfp8_quantize,
|
||||||
|
)
|
||||||
|
|
||||||
N, K = NK
|
N, K = NK
|
||||||
torch.manual_seed(seed)
|
torch.manual_seed(seed)
|
||||||
|
|
||||||
@@ -465,23 +473,29 @@ class TestMXFP8DenseLinear(CustomTestCase):
|
|||||||
|
|
||||||
weight_fp32 = torch.randn((N, K), dtype=torch.float32) / 4
|
weight_fp32 = torch.randn((N, K), dtype=torch.float32) / 4
|
||||||
weight_q, weight_scale_u8 = mxfp8_group_quantize(weight_fp32)
|
weight_q, weight_scale_u8 = mxfp8_group_quantize(weight_fp32)
|
||||||
|
weight_scale_swizzled = block_scale_interleave(weight_scale_u8).contiguous()
|
||||||
|
|
||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
q_input, input_scale_u8 = mxfp8_group_quantize(input_fp16.to(torch.float32))
|
q_input, input_scale_swizzled = flashinfer_mxfp8_quantize(
|
||||||
a_dq = _mxfp8_group_dequant(q_input, input_scale_u8)
|
input_fp16, is_sf_swizzled_layout=True, alignment=32
|
||||||
|
)
|
||||||
|
_, input_scale_linear = flashinfer_mxfp8_quantize(
|
||||||
|
input_fp16, is_sf_swizzled_layout=False, alignment=32
|
||||||
|
)
|
||||||
|
a_dq = _mxfp8_group_dequant(q_input, input_scale_linear.view(M, K // 32))
|
||||||
b_dq = _mxfp8_group_dequant(weight_q, weight_scale_u8)
|
b_dq = _mxfp8_group_dequant(weight_q, weight_scale_u8)
|
||||||
ref_out = torch.matmul(a_dq, b_dq.t()).to(dtype)
|
ref_out = torch.matmul(a_dq, b_dq.t()).to(dtype)
|
||||||
|
|
||||||
out = triton_mxfp8_blockscaled_linear(
|
out = flashinfer_mxfp8_blockscaled_linear(
|
||||||
input=input_fp16,
|
input=input_fp16,
|
||||||
weight=weight_q,
|
weight=weight_q,
|
||||||
weight_scale=weight_scale_u8,
|
weight_scale=weight_scale_swizzled,
|
||||||
)
|
)
|
||||||
out_prequant = triton_mxfp8_blockscaled_linear(
|
out_prequant = flashinfer_mxfp8_blockscaled_linear(
|
||||||
input=q_input,
|
input=q_input,
|
||||||
weight=weight_q,
|
weight=weight_q,
|
||||||
weight_scale=weight_scale_u8,
|
weight_scale=weight_scale_swizzled,
|
||||||
input_scale=input_scale_u8,
|
input_scale=input_scale_swizzled,
|
||||||
output_dtype=dtype,
|
output_dtype=dtype,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -48,14 +48,14 @@ def test_fused_equals_separate(T, N1, N2, K):
|
|||||||
x = torch.randn(T, K, dtype=torch.bfloat16, device=dev)
|
x = torch.randn(T, K, dtype=torch.bfloat16, device=dev)
|
||||||
|
|
||||||
s1p, s2p = _pack_weight_scale(s1), _pack_weight_scale(s2)
|
s1p, s2p = _pack_weight_scale(s1), _pack_weight_scale(s2)
|
||||||
out1 = mxfp8_linear(x, w1, s1p, weight_scale_fallback=s1)
|
out1 = mxfp8_linear(x, w1, s1p)
|
||||||
out2 = mxfp8_linear(x, w2, s2p, weight_scale_fallback=s2)
|
out2 = mxfp8_linear(x, w2, s2p)
|
||||||
ref = torch.cat([out1, out2], dim=-1)
|
ref = torch.cat([out1, out2], dim=-1)
|
||||||
|
|
||||||
w = torch.cat([w1, w2], dim=0).contiguous()
|
w = torch.cat([w1, w2], dim=0).contiguous()
|
||||||
s = torch.cat([s1, s2], dim=0).contiguous()
|
s = torch.cat([s1, s2], dim=0).contiguous()
|
||||||
sp = _pack_weight_scale(s)
|
sp = _pack_weight_scale(s)
|
||||||
fused = mxfp8_linear(x, w, sp, weight_scale_fallback=s)
|
fused = mxfp8_linear(x, w, sp)
|
||||||
|
|
||||||
assert fused.shape == ref.shape
|
assert fused.shape == ref.shape
|
||||||
assert torch.equal(
|
assert torch.equal(
|
||||||
|
|||||||
Reference in New Issue
Block a user