[Quantization] Support Quark Dense + MoE FP8 & FP8 PTPC (#10485)
Co-authored-by: HAI <hixiao@gmail.com> Co-authored-by: kk <43161300+kkHuang-amd@users.noreply.github.com>
This commit is contained in:
@@ -604,158 +604,16 @@ def apply_fp8_linear(
|
||||
output_shape = [*input.shape[:-1], weight.shape[1]]
|
||||
|
||||
if compressed_tensor_quant:
|
||||
# cutlass_scaled_mm supports per tensor/channel W and per tensor/token A
|
||||
# for sgl-kernel fp8_scaled_mm, it support per channel W now
|
||||
# Maybe apply padding to output, see comment in __init__
|
||||
num_token_padding = output_padding
|
||||
if cutlass_fp8_supported and weight_scale.numel() == weight.shape[1]:
|
||||
qinput, x_scale = scaled_fp8_quant(
|
||||
input_2d,
|
||||
input_scale,
|
||||
use_per_token_if_dynamic=use_per_token_if_dynamic,
|
||||
)
|
||||
|
||||
# Fused GEMM_DQ
|
||||
if VLLM_AVAILABLE and use_vllm_cutlass_w8a8_fp8_kernel:
|
||||
# Fall back to vllm cutlass w8a8 fp8 kernel
|
||||
output = ops.cutlass_scaled_mm(
|
||||
qinput,
|
||||
weight,
|
||||
out_dtype=input.dtype,
|
||||
scale_a=x_scale,
|
||||
scale_b=weight_scale,
|
||||
bias=bias,
|
||||
)
|
||||
else:
|
||||
assert (
|
||||
weight_scale.numel() == weight.shape[1]
|
||||
), "cutlass w8a8 fp8 sgl-kernel only supports per-channel scale"
|
||||
|
||||
cutlass_compatible_b = (
|
||||
weight.shape[0] % 16 == 0 and weight.shape[1] % 16 == 0
|
||||
)
|
||||
if not cutlass_compatible_b or use_triton_w8a8_fp8_kernel:
|
||||
# Massage the input to be 2D
|
||||
qinput = qinput.view(-1, qinput.shape[-1])
|
||||
output = triton_scaled_mm(
|
||||
qinput, weight, x_scale, weight_scale, input.dtype, bias
|
||||
)
|
||||
else:
|
||||
output = fp8_scaled_mm(
|
||||
qinput,
|
||||
weight,
|
||||
x_scale,
|
||||
weight_scale,
|
||||
out_dtype=input.dtype,
|
||||
bias=bias,
|
||||
)
|
||||
return output.view(*output_shape)
|
||||
|
||||
# torch.scaled_mm supports per tensor weights + activations only
|
||||
# so fallback to naive if per channel or per token
|
||||
else:
|
||||
# Maybe apply padding to output, see comment in __init__
|
||||
qinput, x_scale = (
|
||||
scaled_fp8_quant(
|
||||
input_2d,
|
||||
input_scale,
|
||||
num_token_padding=output_padding,
|
||||
use_per_token_if_dynamic=use_per_token_if_dynamic,
|
||||
)
|
||||
if _is_cuda
|
||||
else ops.scaled_fp8_quant(
|
||||
input_2d,
|
||||
input_scale,
|
||||
num_token_padding=output_padding,
|
||||
use_per_token_if_dynamic=use_per_token_if_dynamic,
|
||||
)
|
||||
)
|
||||
|
||||
per_tensor_weights = weight_scale.numel() == 1
|
||||
per_tensor_activations = x_scale.numel() == 1
|
||||
|
||||
if per_tensor_weights and per_tensor_activations:
|
||||
# Fused GEMM_DQ
|
||||
output = torch._scaled_mm(
|
||||
qinput,
|
||||
weight,
|
||||
out_dtype=input.dtype,
|
||||
scale_a=x_scale,
|
||||
scale_b=weight_scale,
|
||||
bias=bias,
|
||||
)
|
||||
return _process_scaled_mm_output(output, input_2d.shape, output_shape)
|
||||
|
||||
elif (
|
||||
use_per_token_if_dynamic
|
||||
and not per_tensor_weights
|
||||
and not per_tensor_activations
|
||||
and (USE_ROWWISE_TORCH_SCALED_MM or _use_aiter)
|
||||
):
|
||||
# into this sector means use dynamic per-token-per-channel quant
|
||||
# per-token scale quant for input matrix, every row(one token) have one scale factor
|
||||
# per-channel scale quant for weight matrix, every col(one channel) have one scale factor
|
||||
if _use_aiter:
|
||||
# gemm_a8w8_bpreshuffle(XQ, WQ, x_scale, w_scale, dtype)
|
||||
# XQ -> input tensor, shape = (m, k)
|
||||
# WQ -> weight tensor, shape = (n, k), with preshuffe get better perf
|
||||
# x_scale -> input scale tensor, shape = (m, 1)
|
||||
# w_scale -> weight scale tensor, shape = (n ,1)
|
||||
# dtype -> output dtype
|
||||
output = gemm_a8w8_bpreshuffle(
|
||||
XQ=qinput,
|
||||
WQ=weight,
|
||||
x_scale=x_scale,
|
||||
w_scale=weight_scale,
|
||||
dtype=input.dtype,
|
||||
)
|
||||
if bias is not None:
|
||||
output += bias
|
||||
return _process_scaled_mm_output(
|
||||
output, input_2d.shape, [*input.shape[:-1], weight.shape[0]]
|
||||
)
|
||||
else:
|
||||
# For now validated on ROCm platform
|
||||
# fp8 rowwise scaling in torch._scaled_mm is introduced in
|
||||
# https://github.com/pytorch/pytorch/pull/144432 using hipBLASLt
|
||||
# and ROCm 6.3, which only exists in torch 2.7 and above.
|
||||
# For CUDA platform please validate if the
|
||||
# torch._scaled_mm support rowwise scaled GEMM
|
||||
# Fused GEMM_DQ Rowwise GEMM
|
||||
output = torch._scaled_mm(
|
||||
qinput,
|
||||
weight,
|
||||
out_dtype=input.dtype,
|
||||
scale_a=x_scale,
|
||||
scale_b=weight_scale.t(),
|
||||
bias=bias,
|
||||
)
|
||||
return _process_scaled_mm_output(
|
||||
output, input_2d.shape, output_shape
|
||||
)
|
||||
else:
|
||||
# Fallback for channelwise case, where we use unfused DQ
|
||||
# due to limitations with scaled_mm
|
||||
|
||||
# Symmetric quantized GEMM by definition computes the following:
|
||||
# C = (s_x * X) (s_w * W) + bias
|
||||
# This is equivalent to dequantizing the weights and activations
|
||||
# before applying a GEMM.
|
||||
#
|
||||
# In order to compute quantized operands, a quantized kernel
|
||||
# will rewrite the above like so:
|
||||
# C = s_w * s_x * (X * W) + bias
|
||||
#
|
||||
# For the scaled_mm fallback case, we break this down, since it
|
||||
# does not support s_w being a vector.
|
||||
return _apply_fallback_scaled_mm(
|
||||
qinput,
|
||||
weight,
|
||||
x_scale,
|
||||
weight_scale,
|
||||
input_2d.shape,
|
||||
output_shape,
|
||||
bias,
|
||||
input.dtype,
|
||||
)
|
||||
num_token_padding = None
|
||||
qinput, x_scale = scaled_fp8_quant(
|
||||
input_2d,
|
||||
input_scale,
|
||||
num_token_padding=num_token_padding,
|
||||
use_per_token_if_dynamic=use_per_token_if_dynamic,
|
||||
)
|
||||
else:
|
||||
# cutlass w8a8 fp8 sgl-kernel only supports per-token scale
|
||||
if input_scale is not None:
|
||||
@@ -783,53 +641,12 @@ def apply_fp8_linear(
|
||||
input_2d, group_size=input_2d.shape[1]
|
||||
)
|
||||
|
||||
if cutlass_fp8_supported:
|
||||
try:
|
||||
if VLLM_AVAILABLE and use_vllm_cutlass_w8a8_fp8_kernel:
|
||||
# Fall back to vllm cutlass w8a8 fp8 kernel
|
||||
output = ops.cutlass_scaled_mm(
|
||||
qinput,
|
||||
weight,
|
||||
out_dtype=input.dtype,
|
||||
scale_a=x_scale,
|
||||
scale_b=weight_scale,
|
||||
bias=bias,
|
||||
)
|
||||
else:
|
||||
assert (
|
||||
weight_scale.numel() == weight.shape[1]
|
||||
), "cutlass w8a8 fp8 sgl-kernel only supports per-channel scale"
|
||||
|
||||
cutlass_compatible_b = (
|
||||
weight.shape[0] % 16 == 0 and weight.shape[1] % 16 == 0
|
||||
)
|
||||
if not cutlass_compatible_b or use_triton_w8a8_fp8_kernel:
|
||||
# Massage the input to be 2D
|
||||
qinput = qinput.view(-1, qinput.shape[-1])
|
||||
output = triton_scaled_mm(
|
||||
qinput, weight, x_scale, weight_scale, input.dtype, bias
|
||||
)
|
||||
else:
|
||||
output = fp8_scaled_mm(
|
||||
qinput,
|
||||
weight,
|
||||
x_scale,
|
||||
weight_scale,
|
||||
out_dtype=input.dtype,
|
||||
bias=bias,
|
||||
)
|
||||
return output.view(*output_shape)
|
||||
except (ImportError, NameError, AttributeError):
|
||||
pass
|
||||
|
||||
# torch.scaled_mm supports per tensor weights + activations only
|
||||
# so fallback to naive if per channel or per token
|
||||
per_tensor_weights = weight_scale.numel() == 1
|
||||
per_tensor_activations = x_scale.numel() == 1
|
||||
|
||||
if per_tensor_weights and per_tensor_activations:
|
||||
# Fused GEMM_DQ
|
||||
output = torch._scaled_mm(
|
||||
if cutlass_fp8_supported and weight_scale.numel() == weight.shape[1]:
|
||||
# cutlass_scaled_mm supports per tensor/channel W and per tensor/token A
|
||||
# for sgl-kernel fp8_scaled_mm, it support per channel W now
|
||||
if VLLM_AVAILABLE and use_vllm_cutlass_w8a8_fp8_kernel:
|
||||
# Fall back to vllm cutlass w8a8 fp8 kernel
|
||||
output = ops.cutlass_scaled_mm(
|
||||
qinput,
|
||||
weight,
|
||||
out_dtype=input.dtype,
|
||||
@@ -837,33 +654,112 @@ def apply_fp8_linear(
|
||||
scale_b=weight_scale,
|
||||
bias=bias,
|
||||
)
|
||||
return _process_scaled_mm_output(output, input_2d.shape, output_shape)
|
||||
|
||||
else:
|
||||
# Fallback for channelwise case, where we use unfused DQ
|
||||
# due to limitations with scaled_mm
|
||||
cutlass_compatible_b = (
|
||||
weight.shape[0] % 16 == 0 and weight.shape[1] % 16 == 0
|
||||
)
|
||||
if not cutlass_compatible_b or use_triton_w8a8_fp8_kernel:
|
||||
# Massage the input to be 2D
|
||||
qinput = qinput.view(-1, qinput.shape[-1])
|
||||
output = triton_scaled_mm(
|
||||
qinput, weight, x_scale, weight_scale, input.dtype, bias
|
||||
)
|
||||
else:
|
||||
output = fp8_scaled_mm(
|
||||
qinput,
|
||||
weight,
|
||||
x_scale,
|
||||
weight_scale,
|
||||
out_dtype=input.dtype,
|
||||
bias=bias,
|
||||
)
|
||||
return output.view(*output_shape)
|
||||
|
||||
# Symmetric quantized GEMM by definition computes the following:
|
||||
# C = (s_x * X) (s_w * W) + bias
|
||||
# This is equivalent to dequantizing the weights and activations
|
||||
# before applying a GEMM.
|
||||
#
|
||||
# In order to compute quantized operands, a quantized kernel
|
||||
# will rewrite the above like so:
|
||||
# C = s_w * s_x * (X * W) + bias
|
||||
#
|
||||
# For the scaled_mm fallback case, we break this down, since it
|
||||
# does not support s_w being a vector.
|
||||
return _apply_fallback_scaled_mm(
|
||||
# torch.scaled_mm supports per tensor weights + activations only
|
||||
# so fallback to naive if per channel or per token
|
||||
per_tensor_weights = weight_scale.numel() == 1
|
||||
per_tensor_activations = x_scale.numel() == 1
|
||||
|
||||
if (
|
||||
use_per_token_if_dynamic
|
||||
and not per_tensor_weights
|
||||
and not per_tensor_activations
|
||||
and (USE_ROWWISE_TORCH_SCALED_MM or _use_aiter)
|
||||
):
|
||||
# into this sector means use dynamic per-token-per-channel quant
|
||||
# per-token scale quant for input matrix, every row(one token) have one scale factor
|
||||
# per-channel scale quant for weight matrix, every col(one channel) have one scale factor
|
||||
if _use_aiter:
|
||||
# gemm_a8w8_bpreshuffle(XQ, WQ, x_scale, w_scale, dtype)
|
||||
# XQ -> input tensor, shape = (m, k)
|
||||
# WQ -> weight tensor, shape = (n, k), with preshuffe get better perf
|
||||
# x_scale -> input scale tensor, shape = (m, 1)
|
||||
# w_scale -> weight scale tensor, shape = (n ,1)
|
||||
# dtype -> output dtype
|
||||
output = gemm_a8w8_bpreshuffle(
|
||||
XQ=qinput,
|
||||
WQ=weight.T,
|
||||
x_scale=x_scale,
|
||||
w_scale=weight_scale,
|
||||
dtype=input.dtype,
|
||||
)
|
||||
if bias is not None:
|
||||
output += bias
|
||||
return _process_scaled_mm_output(output, input_2d.shape, output_shape)
|
||||
else:
|
||||
# For now validated on ROCm platform
|
||||
# fp8 rowwise scaling in torch._scaled_mm is introduced in
|
||||
# https://github.com/pytorch/pytorch/pull/144432 using hipBLASLt
|
||||
# and ROCm 6.3, which only exists in torch 2.7 and above.
|
||||
# For CUDA platform please validate if the
|
||||
# torch._scaled_mm support rowwise scaled GEMM
|
||||
# Fused GEMM_DQ Rowwise GEMM
|
||||
output = torch._scaled_mm(
|
||||
qinput,
|
||||
weight,
|
||||
x_scale,
|
||||
weight_scale,
|
||||
input_2d.shape,
|
||||
output_shape,
|
||||
bias,
|
||||
input.dtype,
|
||||
out_dtype=input.dtype,
|
||||
scale_a=x_scale,
|
||||
scale_b=weight_scale.t(),
|
||||
bias=bias,
|
||||
)
|
||||
return _process_scaled_mm_output(output, input_2d.shape, output_shape)
|
||||
|
||||
if per_tensor_weights and per_tensor_activations:
|
||||
# Fused GEMM_DQ
|
||||
output = torch._scaled_mm(
|
||||
qinput,
|
||||
weight,
|
||||
out_dtype=input.dtype,
|
||||
scale_a=x_scale,
|
||||
scale_b=weight_scale,
|
||||
bias=bias,
|
||||
)
|
||||
return _process_scaled_mm_output(output, input_2d.shape, output_shape)
|
||||
|
||||
# Fallback for channelwise case, where we use unfused DQ
|
||||
# due to limitations with scaled_mm
|
||||
|
||||
# Symmetric quantized GEMM by definition computes the following:
|
||||
# C = (s_x * X) (s_w * W) + bias
|
||||
# This is equivalent to dequantizing the weights and activations
|
||||
# before applying a GEMM.
|
||||
#
|
||||
# In order to compute quantized operands, a quantized kernel
|
||||
# will rewrite the above like so:
|
||||
# C = s_w * s_x * (X * W) + bias
|
||||
#
|
||||
# For the scaled_mm fallback case, we break this down, since it
|
||||
# does not support s_w being a vector.
|
||||
return _apply_fallback_scaled_mm(
|
||||
qinput,
|
||||
weight,
|
||||
x_scale,
|
||||
weight_scale,
|
||||
input_2d.shape,
|
||||
output_shape,
|
||||
bias,
|
||||
input.dtype,
|
||||
)
|
||||
|
||||
|
||||
def can_auto_enable_marlin_fp8() -> bool:
|
||||
|
||||
Reference in New Issue
Block a user