refactor apply_w8a8_block_fp8_linear in fp (#6545)
This commit is contained in:
@@ -1,5 +1,6 @@
|
||||
import os
|
||||
from typing import List, Optional, Tuple
|
||||
from curses import flash
|
||||
from typing import Callable, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
@@ -21,7 +22,8 @@ from sglang.srt.layers.quantization.fp8_kernel import (
|
||||
scaled_fp8_quant,
|
||||
sglang_per_token_quant_fp8,
|
||||
static_quant_fp8,
|
||||
w8a8_block_fp8_matmul,
|
||||
w8a8_block_fp8_matmul_deepgemm,
|
||||
w8a8_block_fp8_matmul_triton,
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
get_bool_env_var,
|
||||
@@ -134,7 +136,20 @@ if ENABLE_FLASHINFER_GEMM:
|
||||
from flashinfer.gemm import gemm_fp8_nt_groupwise
|
||||
|
||||
|
||||
def apply_w8a8_block_fp8_linear(
|
||||
def dispatch_w8a8_block_fp8_linear() -> Callable:
|
||||
if ENABLE_FLASHINFER_GEMM:
|
||||
return flashinfer_gemm_w8a8_block_fp8_linear
|
||||
elif CUTLASS_BLOCK_FP8_SUPPORTED:
|
||||
return cutlass_w8a8_block_fp8_linear_with_fallback
|
||||
elif _is_hip and use_aiter_moe:
|
||||
return aiter_w8a8_block_fp8_linear
|
||||
elif _ENABLE_JIT_DEEPGEMM:
|
||||
return deepgemm_w8a8_block_fp8_linear_with_fallback
|
||||
else:
|
||||
return triton_w8a8_block_fp8_linear
|
||||
|
||||
|
||||
def flashinfer_gemm_w8a8_block_fp8_linear(
|
||||
input: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
block_size: List[int],
|
||||
@@ -143,58 +158,148 @@ def apply_w8a8_block_fp8_linear(
|
||||
bias: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
assert input_scale is None
|
||||
# View input as 2D matrix for fp8 methods
|
||||
|
||||
input_2d = input.view(-1, input.shape[-1])
|
||||
output_shape = [*input.shape[:-1], weight.shape[0]]
|
||||
# TODO: add more robust shape check here
|
||||
shape_supported_by_cutlass = (
|
||||
weight.shape[0] % 128 == 0 and weight.shape[1] % 128 == 0
|
||||
|
||||
q_input, x_scale = sglang_per_token_group_quant_fp8(
|
||||
input_2d, block_size[1], column_major_scales=False
|
||||
)
|
||||
|
||||
x_scale_input = x_scale.T.contiguous()
|
||||
weight_scale_input = weight_scale.T.contiguous()
|
||||
|
||||
output = gemm_fp8_nt_groupwise(
|
||||
q_input, weight, x_scale_input, weight_scale_input, out_dtype=input_2d.dtype
|
||||
)
|
||||
if ENABLE_FLASHINFER_GEMM:
|
||||
q_input, x_scale = sglang_per_token_group_quant_fp8(
|
||||
input_2d, block_size[1], column_major_scales=False
|
||||
)
|
||||
x_scale_input = x_scale.T.contiguous()
|
||||
weight_scale_input = weight_scale.T.contiguous()
|
||||
output = gemm_fp8_nt_groupwise(
|
||||
q_input, weight, x_scale_input, weight_scale_input, out_dtype=input.dtype
|
||||
)
|
||||
elif CUTLASS_BLOCK_FP8_SUPPORTED and shape_supported_by_cutlass:
|
||||
q_input, x_scale = per_token_group_quant_fp8(
|
||||
input_2d, block_size[1], column_major_scales=True
|
||||
)
|
||||
output = fp8_blockwise_scaled_mm(
|
||||
q_input, weight.T, x_scale, weight_scale.T, out_dtype=input.dtype
|
||||
)
|
||||
elif _is_hip and use_aiter_moe:
|
||||
q_input, x_scale = per_token_group_quant_fp8(
|
||||
input_2d, block_size[1], column_major_scales=False
|
||||
)
|
||||
output = torch.zeros(
|
||||
[q_input.shape[0], weight.shape[0]],
|
||||
dtype=input.dtype,
|
||||
device=q_input.device,
|
||||
)
|
||||
gemm_a8w8_blockscale(q_input, weight, x_scale, weight_scale, output)
|
||||
else:
|
||||
if _ENABLE_JIT_DEEPGEMM:
|
||||
q_input, x_scale = sglang_per_token_group_quant_fp8(
|
||||
input_2d,
|
||||
block_size[1],
|
||||
column_major_scales=True,
|
||||
scale_tma_aligned=True,
|
||||
)
|
||||
else:
|
||||
q_input, x_scale = per_token_group_quant_fp8(
|
||||
input_2d, block_size[1], column_major_scales=False
|
||||
)
|
||||
output = w8a8_block_fp8_matmul(
|
||||
q_input, weight, x_scale, weight_scale, block_size, output_dtype=input.dtype
|
||||
)
|
||||
|
||||
if bias is not None:
|
||||
output = output + bias
|
||||
return output.to(dtype=input.dtype).view(*output_shape)
|
||||
output += bias
|
||||
|
||||
return output.to(dtype=input_2d.dtype).view(*output_shape)
|
||||
|
||||
|
||||
def cutlass_w8a8_block_fp8_linear_with_fallback(
|
||||
input: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
block_size: List[int],
|
||||
weight_scale: torch.Tensor,
|
||||
input_scale: Optional[torch.Tensor] = None,
|
||||
bias: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
assert input_scale is None
|
||||
|
||||
# TODO: add more robust shape check here
|
||||
shape_supported = weight.shape[0] % 128 == 0 and weight.shape[1] % 128 == 0
|
||||
|
||||
if not shape_supported:
|
||||
# fallback to triton
|
||||
return triton_w8a8_block_fp8_linear(
|
||||
input, weight, block_size, weight_scale, input_scale, bias
|
||||
)
|
||||
|
||||
input_2d = input.view(-1, input.shape[-1])
|
||||
output_shape = [*input.shape[:-1], weight.shape[0]]
|
||||
|
||||
q_input, x_scale = per_token_group_quant_fp8(
|
||||
input_2d, block_size[1], column_major_scales=True
|
||||
)
|
||||
output = fp8_blockwise_scaled_mm(
|
||||
q_input, weight.T, x_scale, weight_scale.T, out_dtype=input_2d.dtype
|
||||
)
|
||||
if bias is not None:
|
||||
output += bias
|
||||
return output.to(dtype=input_2d.dtype).view(*output_shape)
|
||||
|
||||
|
||||
def deepgemm_w8a8_block_fp8_linear_with_fallback(
|
||||
input: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
block_size: List[int],
|
||||
weight_scale: torch.Tensor,
|
||||
input_scale: Optional[torch.Tensor] = None,
|
||||
bias: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
assert input_scale is None
|
||||
|
||||
output_dtype = input.dtype
|
||||
dtype_supported = output_dtype == torch.bfloat16
|
||||
|
||||
# TODO: add more robust shape check here
|
||||
shape_supported = weight.shape[0] % 128 == 0 and weight.shape[1] % 128 == 0
|
||||
|
||||
if not (shape_supported and dtype_supported):
|
||||
# fall back to triton
|
||||
return triton_w8a8_block_fp8_linear(
|
||||
input, weight, block_size, weight_scale, input_scale, bias
|
||||
)
|
||||
|
||||
input_2d = input.view(-1, input.shape[-1])
|
||||
output_shape = [*input.shape[:-1], weight.shape[0]]
|
||||
|
||||
q_input, x_scale = sglang_per_token_group_quant_fp8(
|
||||
input_2d,
|
||||
block_size[1],
|
||||
column_major_scales=True,
|
||||
scale_tma_aligned=True,
|
||||
)
|
||||
output = w8a8_block_fp8_matmul_deepgemm(
|
||||
q_input, weight, x_scale, weight_scale, block_size, output_dtype=output_dtype
|
||||
)
|
||||
if bias is not None:
|
||||
output += bias
|
||||
return output.to(dtype=output_dtype).view(*output_shape)
|
||||
|
||||
|
||||
def aiter_w8a8_block_fp8_linear(
|
||||
input: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
block_size: List[int],
|
||||
weight_scale: torch.Tensor,
|
||||
input_scale: Optional[torch.Tensor] = None,
|
||||
bias: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
assert input_scale is None
|
||||
input_2d = input.view(-1, input.shape[-1])
|
||||
output_shape = [*input.shape[:-1], weight.shape[0]]
|
||||
|
||||
q_input, x_scale = per_token_group_quant_fp8(
|
||||
input_2d, block_size[1], column_major_scales=False
|
||||
)
|
||||
output = torch.zeros(
|
||||
[q_input.shape[0], weight.shape[0]],
|
||||
dtype=input_2d.dtype,
|
||||
device=q_input.device,
|
||||
)
|
||||
gemm_a8w8_blockscale(q_input, weight, x_scale, weight_scale, output)
|
||||
|
||||
if bias is not None:
|
||||
output += bias
|
||||
|
||||
return output.to(dtype=input_2d.dtype).view(*output_shape)
|
||||
|
||||
|
||||
def triton_w8a8_block_fp8_linear(
|
||||
input: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
block_size: List[int],
|
||||
weight_scale: torch.Tensor,
|
||||
input_scale: Optional[torch.Tensor] = None,
|
||||
bias: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
assert input_scale is None
|
||||
input_2d = input.view(-1, input.shape[-1])
|
||||
output_shape = [*input.shape[:-1], weight.shape[0]]
|
||||
|
||||
q_input, x_scale = per_token_group_quant_fp8(
|
||||
input_2d, block_size[1], column_major_scales=False
|
||||
)
|
||||
output = w8a8_block_fp8_matmul_triton(
|
||||
q_input, weight, x_scale, weight_scale, block_size, output_dtype=input_2d.dtype
|
||||
)
|
||||
if bias is not None:
|
||||
output += bias
|
||||
return output.to(dtype=input_2d.dtype).view(*output_shape)
|
||||
|
||||
|
||||
def input_to_float8(
|
||||
|
||||
Reference in New Issue
Block a user