Fix FP8 Triton dtype selection on A100 (#31340)
Co-authored-by: Daniel Afrimi <dafrimi@aws-dfw-cs-001-login-01.cm.cluster>
This commit is contained in:
co-authored by
Daniel Afrimi
parent
3079157175
commit
7de8792758
@@ -29,6 +29,7 @@ except:
|
|||||||
pass
|
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.srt.layers import deep_gemm_wrapper
|
from sglang.srt.layers import deep_gemm_wrapper
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
ceil_align,
|
ceil_align,
|
||||||
@@ -834,6 +835,7 @@ def _static_quant_fp8(
|
|||||||
fp8_max,
|
fp8_max,
|
||||||
# Meta-parameters
|
# Meta-parameters
|
||||||
BLOCK: tl.constexpr,
|
BLOCK: tl.constexpr,
|
||||||
|
FP8_DTYPE: tl.constexpr,
|
||||||
REPEAT_SCALE: tl.constexpr,
|
REPEAT_SCALE: tl.constexpr,
|
||||||
USE_PDL: tl.constexpr = False,
|
USE_PDL: tl.constexpr = False,
|
||||||
):
|
):
|
||||||
@@ -862,9 +864,9 @@ def _static_quant_fp8(
|
|||||||
tl.extra.cuda.gdc_launch_dependents()
|
tl.extra.cuda.gdc_launch_dependents()
|
||||||
|
|
||||||
y_s_inv = 1.0 / y_s
|
y_s_inv = 1.0 / y_s
|
||||||
y_q = tl.clamp(y * y_s_inv, fp8_min, fp8_max).to(y_q_ptr.dtype.element_ty)
|
y_q = tl.clamp(y * y_s_inv, fp8_min, fp8_max).to(FP8_DTYPE)
|
||||||
|
|
||||||
tl.store(y_q_ptr + cols, y_q, mask=mask)
|
tl.store(y_q_ptr + cols, y_q.to(tl.uint8, bitcast=True), mask=mask)
|
||||||
if REPEAT_SCALE:
|
if REPEAT_SCALE:
|
||||||
tl.store(y_s_repeat_ptr, y_s)
|
tl.store(y_s_repeat_ptr, y_s)
|
||||||
|
|
||||||
@@ -910,7 +912,7 @@ def static_quant_fp8(
|
|||||||
pdl_kwargs = {"USE_PDL": True, "launch_pdl": True} if is_arch_support_pdl() else {}
|
pdl_kwargs = {"USE_PDL": True, "launch_pdl": True} if is_arch_support_pdl() else {}
|
||||||
_static_quant_fp8[(M,)](
|
_static_quant_fp8[(M,)](
|
||||||
x,
|
x,
|
||||||
x_q,
|
x_q.view(torch.uint8),
|
||||||
x_s,
|
x_s,
|
||||||
x_s_repeat,
|
x_s_repeat,
|
||||||
N,
|
N,
|
||||||
@@ -918,6 +920,7 @@ def static_quant_fp8(
|
|||||||
fp8_min=fp8_min,
|
fp8_min=fp8_min,
|
||||||
fp8_max=fp8_max,
|
fp8_max=fp8_max,
|
||||||
BLOCK=BLOCK,
|
BLOCK=BLOCK,
|
||||||
|
FP8_DTYPE=fp8_dtype_to_triton(fp8_dtype),
|
||||||
REPEAT_SCALE=repeat_scale,
|
REPEAT_SCALE=repeat_scale,
|
||||||
num_warps=num_warps,
|
num_warps=num_warps,
|
||||||
num_stages=num_stages,
|
num_stages=num_stages,
|
||||||
@@ -1706,6 +1709,7 @@ def _per_tensor_quant_mla_fp8_stage2(
|
|||||||
x_stride_s,
|
x_stride_s,
|
||||||
fp8_min,
|
fp8_min,
|
||||||
fp8_max,
|
fp8_max,
|
||||||
|
FP8_DTYPE: tl.constexpr,
|
||||||
BLOCK_SIZE: tl.constexpr,
|
BLOCK_SIZE: tl.constexpr,
|
||||||
):
|
):
|
||||||
seq_id = tl.program_id(0)
|
seq_id = tl.program_id(0)
|
||||||
@@ -1720,8 +1724,8 @@ def _per_tensor_quant_mla_fp8_stage2(
|
|||||||
x_q_ptr += head_id * num_seq * head_size + seq_id * head_size
|
x_q_ptr += head_id * num_seq * head_size + seq_id * head_size
|
||||||
|
|
||||||
x = tl.load(x_ptr + offset, mask=mask, other=0.0).to(tl.float32)
|
x = tl.load(x_ptr + offset, mask=mask, other=0.0).to(tl.float32)
|
||||||
x_q = tl.clamp(x * x_s_inv, fp8_min, fp8_max).to(x_q_ptr.dtype.element_ty)
|
x_q = tl.clamp(x * x_s_inv, fp8_min, fp8_max).to(FP8_DTYPE)
|
||||||
tl.store(x_q_ptr + offset, x_q, mask=mask)
|
tl.store(x_q_ptr + offset, x_q.to(tl.uint8, bitcast=True), mask=mask)
|
||||||
|
|
||||||
|
|
||||||
def per_tensor_quant_mla_fp8(
|
def per_tensor_quant_mla_fp8(
|
||||||
@@ -1757,13 +1761,14 @@ def per_tensor_quant_mla_fp8(
|
|||||||
_per_tensor_quant_mla_fp8_stage2[grid](
|
_per_tensor_quant_mla_fp8_stage2[grid](
|
||||||
x,
|
x,
|
||||||
x_s_out,
|
x_s_out,
|
||||||
x_q,
|
x_q.view(torch.uint8),
|
||||||
num_seq,
|
num_seq,
|
||||||
head_size,
|
head_size,
|
||||||
x.stride(0),
|
x.stride(0),
|
||||||
x.stride(1),
|
x.stride(1),
|
||||||
fp8_min,
|
fp8_min,
|
||||||
fp8_max,
|
fp8_max,
|
||||||
|
fp8_dtype_to_triton(fp8_dtype),
|
||||||
BLOCK_SIZE,
|
BLOCK_SIZE,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1786,6 +1791,7 @@ def _per_token_group_quant_mla_deep_gemm_masked_fp8(
|
|||||||
eps,
|
eps,
|
||||||
fp8_min,
|
fp8_min,
|
||||||
fp8_max,
|
fp8_max,
|
||||||
|
FP8_DTYPE: tl.constexpr,
|
||||||
NUM_GROUP: tl.constexpr,
|
NUM_GROUP: tl.constexpr,
|
||||||
BLOCK: tl.constexpr,
|
BLOCK: tl.constexpr,
|
||||||
):
|
):
|
||||||
@@ -1814,9 +1820,13 @@ def _per_token_group_quant_mla_deep_gemm_masked_fp8(
|
|||||||
)
|
)
|
||||||
_absmax = tl.maximum(tl.max(tl.abs(y)), eps)
|
_absmax = tl.maximum(tl.max(tl.abs(y)), eps)
|
||||||
y_s = _absmax / fp8_max
|
y_s = _absmax / fp8_max
|
||||||
y_q = tl.clamp(y / y_s, fp8_min, fp8_max).to(y_q_ptr.dtype.element_ty)
|
y_q = tl.clamp(y / y_s, fp8_min, fp8_max).to(FP8_DTYPE)
|
||||||
|
|
||||||
tl.store(y_q_ptr + gid * group_size + cols, y_q, mask=mask)
|
tl.store(
|
||||||
|
y_q_ptr + gid * group_size + cols,
|
||||||
|
y_q.to(tl.uint8, bitcast=True),
|
||||||
|
mask=mask,
|
||||||
|
)
|
||||||
tl.store(y_s_ptr + gid * y_s_stride_g, y_s)
|
tl.store(y_s_ptr + gid * y_s_stride_g, y_s)
|
||||||
|
|
||||||
|
|
||||||
@@ -1846,7 +1856,7 @@ def per_token_group_quant_mla_deep_gemm_masked_fp8(
|
|||||||
|
|
||||||
_per_token_group_quant_mla_deep_gemm_masked_fp8[grid](
|
_per_token_group_quant_mla_deep_gemm_masked_fp8[grid](
|
||||||
x,
|
x,
|
||||||
x_q,
|
x_q.view(torch.uint8),
|
||||||
x_s,
|
x_s,
|
||||||
masked_m,
|
masked_m,
|
||||||
group_size,
|
group_size,
|
||||||
@@ -1859,6 +1869,7 @@ def per_token_group_quant_mla_deep_gemm_masked_fp8(
|
|||||||
eps,
|
eps,
|
||||||
-fp8_max,
|
-fp8_max,
|
||||||
fp8_max,
|
fp8_max,
|
||||||
|
fp8_dtype_to_triton(dtype),
|
||||||
num_tiles_k,
|
num_tiles_k,
|
||||||
BLOCK_SIZE,
|
BLOCK_SIZE,
|
||||||
)
|
)
|
||||||
@@ -2028,6 +2039,7 @@ def _per_token_group_quant_fp8_hopper_moe_mn_major(
|
|||||||
K: tl.constexpr,
|
K: tl.constexpr,
|
||||||
BLOCK_K: tl.constexpr,
|
BLOCK_K: tl.constexpr,
|
||||||
M_ALIGNMENT: tl.constexpr,
|
M_ALIGNMENT: tl.constexpr,
|
||||||
|
FP8_DTYPE: tl.constexpr,
|
||||||
BLOCK_M: tl.constexpr, # tune
|
BLOCK_M: tl.constexpr, # tune
|
||||||
):
|
):
|
||||||
k_offset = tl.program_id(0)
|
k_offset = tl.program_id(0)
|
||||||
@@ -2047,13 +2059,13 @@ def _per_token_group_quant_fp8_hopper_moe_mn_major(
|
|||||||
inp = tl.load(a_ptrs, mask=a_mask).to(tl.float32) # [BLOCK_M, BLOCK_K]
|
inp = tl.load(a_ptrs, mask=a_mask).to(tl.float32) # [BLOCK_M, BLOCK_K]
|
||||||
inp_amax = tl.max(tl.abs(inp), axis=1) # [BLOCK_M,]
|
inp_amax = tl.max(tl.abs(inp), axis=1) # [BLOCK_M,]
|
||||||
inp_amax = tl.clamp(inp_amax, min=1e-4, max=float("inf"))
|
inp_amax = tl.clamp(inp_amax, min=1e-4, max=float("inf"))
|
||||||
inp_fp8 = (inp * (448.0 / inp_amax[:, None])).to(tl.float8e4nv)
|
inp_fp8 = (inp * (448.0 / inp_amax[:, None])).to(FP8_DTYPE)
|
||||||
|
|
||||||
# Store fp8
|
# Store fp8
|
||||||
a_fp8_ptrs = (
|
a_fp8_ptrs = (
|
||||||
a_fp8 + current_expert_offset * K + coord_m[:, None] * K + coord_k[None, :]
|
a_fp8 + current_expert_offset * K + coord_m[:, None] * K + coord_k[None, :]
|
||||||
)
|
)
|
||||||
tl.store(a_fp8_ptrs, inp_fp8, mask=a_mask)
|
tl.store(a_fp8_ptrs, inp_fp8.to(tl.uint8, bitcast=True), mask=a_mask)
|
||||||
|
|
||||||
# Store sfa
|
# Store sfa
|
||||||
k = tl.cdiv(K, BLOCK_K)
|
k = tl.cdiv(K, BLOCK_K)
|
||||||
@@ -2092,11 +2104,12 @@ def per_token_group_quant_fp8_hopper_moe_mn_major(
|
|||||||
A,
|
A,
|
||||||
expert_offsets,
|
expert_offsets,
|
||||||
problem_sizes,
|
problem_sizes,
|
||||||
a_q,
|
a_q.view(torch.uint8),
|
||||||
sfa,
|
sfa,
|
||||||
K,
|
K,
|
||||||
group_size,
|
group_size,
|
||||||
expert_tokens_alignment,
|
expert_tokens_alignment,
|
||||||
|
fp8_dtype_to_triton(fp8_dtype),
|
||||||
)
|
)
|
||||||
return a_q, sfa
|
return a_q, sfa
|
||||||
|
|
||||||
|
|||||||
@@ -26,6 +26,8 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
|
|
||||||
|
from sglang.kernels.ops.quantization.fp8_utils import fp8_dtype_to_triton
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
def _fp8_quantize_kernel(
|
def _fp8_quantize_kernel(
|
||||||
@@ -52,9 +54,10 @@ def _fp8_quantize_kernel(
|
|||||||
x = tl.load(x_ptr + x_off, mask=m_mask[:, None])
|
x = tl.load(x_ptr + x_off, mask=m_mask[:, None])
|
||||||
|
|
||||||
x_fp8 = (x.to(tl.float32) * scale_inv).to(FP8_DTYPE)
|
x_fp8 = (x.to(tl.float32) * scale_inv).to(FP8_DTYPE)
|
||||||
|
x_fp8_bytes = x_fp8.to(tl.uint8, bitcast=True)
|
||||||
|
|
||||||
out_off = m_idx[:, None] * out_row_stride + n_idx[None, :]
|
out_off = m_idx[:, None] * out_row_stride + n_idx[None, :]
|
||||||
tl.store(out_ptr + out_off, x_fp8, mask=m_mask[:, None])
|
tl.store(out_ptr + out_off, x_fp8_bytes, mask=m_mask[:, None])
|
||||||
|
|
||||||
if ENABLE_PDL:
|
if ENABLE_PDL:
|
||||||
tl.extra.cuda.gdc_launch_dependents()
|
tl.extra.cuda.gdc_launch_dependents()
|
||||||
@@ -123,7 +126,7 @@ def fp8_quantize(
|
|||||||
out_M, _, out_row_stride = _flatten_to_2d(out)
|
out_M, _, out_row_stride = _flatten_to_2d(out)
|
||||||
assert out_M == M
|
assert out_M == M
|
||||||
|
|
||||||
fp8_dtype_const = tl.float8e4nv if fp8_dtype is torch.float8_e4m3fn else tl.float8e5
|
fp8_dtype_const = fp8_dtype_to_triton(fp8_dtype)
|
||||||
|
|
||||||
if M <= 2048:
|
if M <= 2048:
|
||||||
block_m = 4
|
block_m = 4
|
||||||
@@ -141,7 +144,7 @@ def fp8_quantize(
|
|||||||
|
|
||||||
_fp8_quantize_kernel[grid](
|
_fp8_quantize_kernel[grid](
|
||||||
x,
|
x,
|
||||||
out,
|
out.view(torch.uint8),
|
||||||
scale_inv,
|
scale_inv,
|
||||||
M,
|
M,
|
||||||
x_row_stride,
|
x_row_stride,
|
||||||
|
|||||||
@@ -0,0 +1,50 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Optional, Tuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import triton.language as tl
|
||||||
|
|
||||||
|
from sglang.kernels.jit.utils import (
|
||||||
|
get_jit_cuda_arch,
|
||||||
|
is_hip_runtime,
|
||||||
|
is_musa_runtime,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def cuda_capability_uses_fp8_e4b15(cuda_capability: Tuple[int, int]) -> bool:
|
||||||
|
"""Triton names E4M3 as fp8e4b15 on CUDA architectures before SM89."""
|
||||||
|
return cuda_capability < (8, 9)
|
||||||
|
|
||||||
|
|
||||||
|
def use_fp8_e4b15_for_e4m3fn(
|
||||||
|
device: Optional[int] = None,
|
||||||
|
cuda_capability: Optional[Tuple[int, int]] = None,
|
||||||
|
) -> bool:
|
||||||
|
if cuda_capability is None:
|
||||||
|
if is_hip_runtime() or is_musa_runtime() or not torch.cuda.is_available():
|
||||||
|
return False
|
||||||
|
if device is None:
|
||||||
|
arch = get_jit_cuda_arch()
|
||||||
|
cuda_capability = (arch.major, arch.minor)
|
||||||
|
else:
|
||||||
|
cuda_capability = torch.cuda.get_device_capability(device)
|
||||||
|
|
||||||
|
return cuda_capability_uses_fp8_e4b15(cuda_capability)
|
||||||
|
|
||||||
|
|
||||||
|
def fp8_dtype_to_triton(
|
||||||
|
fp8_dtype: torch.dtype,
|
||||||
|
*,
|
||||||
|
device: Optional[int] = None,
|
||||||
|
cuda_capability: Optional[Tuple[int, int]] = None,
|
||||||
|
):
|
||||||
|
if fp8_dtype == torch.float8_e4m3fn:
|
||||||
|
if use_fp8_e4b15_for_e4m3fn(device, cuda_capability):
|
||||||
|
return tl.float8e4b15
|
||||||
|
return tl.float8e4nv
|
||||||
|
if fp8_dtype == torch.float8_e4m3fnuz:
|
||||||
|
return tl.float8e4b8
|
||||||
|
if fp8_dtype == torch.float8_e5m2:
|
||||||
|
return tl.float8e5
|
||||||
|
raise ValueError(f"Unsupported FP8 dtype: {fp8_dtype}")
|
||||||
@@ -38,6 +38,7 @@ from sglang.srt.layers.quantization.fp8 import Fp8Config
|
|||||||
from sglang.srt.layers.quantization.fp8_utils import (
|
from sglang.srt.layers.quantization.fp8_utils import (
|
||||||
apply_fp8_linear,
|
apply_fp8_linear,
|
||||||
apply_fp8_linear_bmm_flashinfer,
|
apply_fp8_linear_bmm_flashinfer,
|
||||||
|
can_auto_enable_marlin_fp8,
|
||||||
cutlass_fp8_supported,
|
cutlass_fp8_supported,
|
||||||
is_blackwell_supported,
|
is_blackwell_supported,
|
||||||
)
|
)
|
||||||
@@ -47,6 +48,9 @@ from sglang.srt.layers.quantization.marlin_utils_fp4 import (
|
|||||||
prepare_moe_nvfp4_layer_for_marlin,
|
prepare_moe_nvfp4_layer_for_marlin,
|
||||||
prepare_nvfp4_layer_for_marlin,
|
prepare_nvfp4_layer_for_marlin,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.layers.quantization.marlin_utils_fp8 import (
|
||||||
|
prepare_fp8_layer_for_marlin,
|
||||||
|
)
|
||||||
from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod
|
from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod
|
||||||
from sglang.srt.layers.quantization.utils import (
|
from sglang.srt.layers.quantization.utils import (
|
||||||
convert_to_channelwise,
|
convert_to_channelwise,
|
||||||
@@ -502,6 +506,11 @@ class ModelOptFp8LinearMethod(LinearMethodBase):
|
|||||||
self.enable_flashinfer_bmm = (
|
self.enable_flashinfer_bmm = (
|
||||||
is_sm100_supported() or is_sm120_supported()
|
is_sm100_supported() or is_sm120_supported()
|
||||||
) and is_flashinfer_available()
|
) and is_flashinfer_available()
|
||||||
|
self.use_marlin = False
|
||||||
|
if is_cuda():
|
||||||
|
self.use_marlin = (
|
||||||
|
envs.SGLANG_FORCE_FP8_MARLIN.get() or can_auto_enable_marlin_fp8()
|
||||||
|
)
|
||||||
|
|
||||||
def create_weights(
|
def create_weights(
|
||||||
self,
|
self,
|
||||||
@@ -526,6 +535,7 @@ class ModelOptFp8LinearMethod(LinearMethodBase):
|
|||||||
layer.logical_widths = output_partition_sizes
|
layer.logical_widths = output_partition_sizes
|
||||||
layer.input_size_per_partition = input_size_per_partition
|
layer.input_size_per_partition = input_size_per_partition
|
||||||
layer.output_size_per_partition = output_size_per_partition
|
layer.output_size_per_partition = output_size_per_partition
|
||||||
|
layer.orig_dtype = params_dtype
|
||||||
|
|
||||||
# Register weight
|
# Register weight
|
||||||
layer.register_parameter(
|
layer.register_parameter(
|
||||||
@@ -563,6 +573,10 @@ class ModelOptFp8LinearMethod(LinearMethodBase):
|
|||||||
max_w_scale = convert_to_channelwise(max_w_scale, layer.logical_widths)
|
max_w_scale = convert_to_channelwise(max_w_scale, layer.logical_widths)
|
||||||
layer.weight_scale = Parameter(max_w_scale, requires_grad=False)
|
layer.weight_scale = Parameter(max_w_scale, requires_grad=False)
|
||||||
layer.input_scale = Parameter(layer.input_scale.max(), requires_grad=False)
|
layer.input_scale = Parameter(layer.input_scale.max(), requires_grad=False)
|
||||||
|
if self.use_marlin:
|
||||||
|
prepare_fp8_layer_for_marlin(layer)
|
||||||
|
# Marlin uses FP8 weights with unquantized activations.
|
||||||
|
del layer.input_scale
|
||||||
|
|
||||||
def apply(
|
def apply(
|
||||||
self,
|
self,
|
||||||
@@ -571,6 +585,16 @@ class ModelOptFp8LinearMethod(LinearMethodBase):
|
|||||||
bias: Optional[torch.Tensor] = None,
|
bias: Optional[torch.Tensor] = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""Applies FP8 linear transformation."""
|
"""Applies FP8 linear transformation."""
|
||||||
|
if self.use_marlin:
|
||||||
|
return torch.ops.sglang.apply_fp8_marlin_linear(
|
||||||
|
input=x,
|
||||||
|
weight=layer.weight,
|
||||||
|
weight_scale=layer.weight_scale,
|
||||||
|
workspace=layer.workspace,
|
||||||
|
size_n=layer.output_size_per_partition,
|
||||||
|
size_k=layer.input_size_per_partition,
|
||||||
|
bias=bias,
|
||||||
|
)
|
||||||
if self.enable_flashinfer_bmm and layer.input_scale is not None:
|
if self.enable_flashinfer_bmm and layer.input_scale is not None:
|
||||||
return apply_fp8_linear_bmm_flashinfer(
|
return apply_fp8_linear_bmm_flashinfer(
|
||||||
input=x,
|
input=x,
|
||||||
|
|||||||
Reference in New Issue
Block a user