[MoE] Support BF16 standard A2A with DeepGEMM runner (#26473)
Co-authored-by: popsiclexu <zhenxue.xu@mthreads.com>
This commit is contained in:
+1
-1
@@ -12,7 +12,7 @@ ARG DEEPEP_COMMIT=9af0e0d0e74f3577af1979c9b9e1ac2cad0104ee
|
|||||||
ARG BUILD_AND_DOWNLOAD_PARALLEL=8
|
ARG BUILD_AND_DOWNLOAD_PARALLEL=8
|
||||||
ARG SGL_KERNEL_VERSION=0.4.3
|
ARG SGL_KERNEL_VERSION=0.4.3
|
||||||
ARG SGL_VERSION
|
ARG SGL_VERSION
|
||||||
ARG SGL_DEEP_GEMM_VERSION=0.1.0
|
ARG SGL_DEEP_GEMM_VERSION=0.1.1
|
||||||
ARG USE_LATEST_SGLANG=0
|
ARG USE_LATEST_SGLANG=0
|
||||||
ARG GDRCOPY_VERSION=2.5.1
|
ARG GDRCOPY_VERSION=2.5.1
|
||||||
ARG PIP_DEFAULT_INDEX
|
ARG PIP_DEFAULT_INDEX
|
||||||
|
|||||||
@@ -60,7 +60,7 @@ dependencies = [
|
|||||||
"sentencepiece",
|
"sentencepiece",
|
||||||
"setproctitle",
|
"setproctitle",
|
||||||
"flash-attn-4>=4.0.0b9",
|
"flash-attn-4>=4.0.0b9",
|
||||||
"sgl-deep-gemm==0.1.0",
|
"sgl-deep-gemm==0.1.1",
|
||||||
"sglang-kernel==0.4.3",
|
"sglang-kernel==0.4.3",
|
||||||
"soundfile==0.13.1",
|
"soundfile==0.13.1",
|
||||||
"tiktoken",
|
"tiktoken",
|
||||||
|
|||||||
@@ -349,7 +349,7 @@ class _BF16GroupedContWarmupExecutor(_BaseWarmupExecutor):
|
|||||||
self.a[:m],
|
self.a[:m],
|
||||||
self.b,
|
self.b,
|
||||||
self.out[:m],
|
self.out[:m],
|
||||||
m_indices=self.m_indices[:m],
|
self.m_indices[:m],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -475,10 +475,12 @@ def _silu_and_mul_kernel(
|
|||||||
input_ptr_offs + token_index * stride_input_1 + size_n,
|
input_ptr_offs + token_index * stride_input_1 + size_n,
|
||||||
mask=offs_in_d < size_n,
|
mask=offs_in_d < size_n,
|
||||||
other=0.0,
|
other=0.0,
|
||||||
)
|
).to(tl.float32)
|
||||||
gate = gate / (1 + tl.exp(-gate))
|
gate = gate / (1 + tl.exp(-gate))
|
||||||
gate = gate.to(input_ptr.dtype.element_ty)
|
|
||||||
gate_up = up * gate
|
gate_up = up * gate
|
||||||
|
# Compute SiLU in fp32 for better precision, then cast back to the
|
||||||
|
# input dtype.
|
||||||
|
gate_up = gate_up.to(input_ptr.dtype.element_ty)
|
||||||
tl.store(
|
tl.store(
|
||||||
output_ptr_offs + token_index * stride_output_1,
|
output_ptr_offs + token_index * stride_output_1,
|
||||||
gate_up,
|
gate_up,
|
||||||
@@ -702,17 +704,19 @@ def post_reorder_triton_kernel(
|
|||||||
offset = start_offset + vec
|
offset = start_offset + vec
|
||||||
mask = offset < hidden_size
|
mask = offset < hidden_size
|
||||||
|
|
||||||
sum_vec = tl.zeros([BLOCK_SIZE], dtype=InDtype)
|
sum_vec = tl.zeros([BLOCK_SIZE], dtype=tl.float32)
|
||||||
for idx in range(topk):
|
for idx in range(topk):
|
||||||
expert_id = tl.load(topk_ids_ptr + idx)
|
expert_id = tl.load(topk_ids_ptr + idx)
|
||||||
if expert_id > 0:
|
if expert_id >= 0:
|
||||||
dst_idx_int32 = tl.load(src2dst_ptr + idx)
|
dst_idx_int32 = tl.load(src2dst_ptr + idx)
|
||||||
dst_idx = dst_idx_int32.to(tl.int64)
|
dst_idx = dst_idx_int32.to(tl.int64)
|
||||||
weigh_scale = tl.load(topk_weights_ptr + idx).to(InDtype)
|
weigh_scale = tl.load(topk_weights_ptr + idx).to(tl.float32)
|
||||||
load_ptr = down_output_ptr + dst_idx * hidden_size
|
load_ptr = down_output_ptr + dst_idx * hidden_size
|
||||||
in_data = tl.load(load_ptr + offset, mask=mask)
|
# accumulate expert outputs in fp32 for better precision
|
||||||
|
# before casting to the final output dtype.
|
||||||
|
in_data = tl.load(load_ptr + offset, mask=mask).to(tl.float32)
|
||||||
sum_vec += in_data * weigh_scale
|
sum_vec += in_data * weigh_scale
|
||||||
tl.store(store_ptr + offset, sum_vec, mask=mask)
|
tl.store(store_ptr + offset, sum_vec.to(InDtype), mask=mask)
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
@@ -1136,6 +1140,7 @@ def fill_gateup_input_triton_kernel(
|
|||||||
hidden_size,
|
hidden_size,
|
||||||
scale_size,
|
scale_size,
|
||||||
BLOCK_SIZE: tl.constexpr,
|
BLOCK_SIZE: tl.constexpr,
|
||||||
|
IS_FP8: tl.constexpr,
|
||||||
):
|
):
|
||||||
|
|
||||||
src_idx_int32 = tl.program_id(0)
|
src_idx_int32 = tl.program_id(0)
|
||||||
@@ -1143,7 +1148,8 @@ def fill_gateup_input_triton_kernel(
|
|||||||
src2dst_ptr = src2dst_ptr + src_idx * topk
|
src2dst_ptr = src2dst_ptr + src_idx * topk
|
||||||
topk_ids_ptr = topk_ids_ptr + src_idx * topk
|
topk_ids_ptr = topk_ids_ptr + src_idx * topk
|
||||||
src_ptr = input_ptr + src_idx * hidden_size
|
src_ptr = input_ptr + src_idx * hidden_size
|
||||||
scale_src_ptr = scale_ptr + src_idx * scale_size
|
if IS_FP8:
|
||||||
|
scale_src_ptr = scale_ptr + src_idx * scale_size
|
||||||
|
|
||||||
vec = tl.arange(0, BLOCK_SIZE)
|
vec = tl.arange(0, BLOCK_SIZE)
|
||||||
for idx in range(topk):
|
for idx in range(topk):
|
||||||
@@ -1157,12 +1163,14 @@ def fill_gateup_input_triton_kernel(
|
|||||||
mask = offset < hidden_size
|
mask = offset < hidden_size
|
||||||
in_data = tl.load(src_ptr + offset, mask=mask)
|
in_data = tl.load(src_ptr + offset, mask=mask)
|
||||||
tl.store(dst_ptr + offset, in_data, mask=mask)
|
tl.store(dst_ptr + offset, in_data, mask=mask)
|
||||||
scale_dst_ptr = gateup_input_scale_ptr + dst_idx * scale_size
|
|
||||||
for start_offset in tl.range(0, scale_size, BLOCK_SIZE):
|
if IS_FP8:
|
||||||
offset = start_offset + vec
|
scale_dst_ptr = gateup_input_scale_ptr + dst_idx * scale_size
|
||||||
mask = offset < scale_size
|
for start_offset in tl.range(0, scale_size, BLOCK_SIZE):
|
||||||
in_scale = tl.load(scale_src_ptr + offset, mask=mask)
|
offset = start_offset + vec
|
||||||
tl.store(scale_dst_ptr + offset, in_scale, mask=mask)
|
mask = offset < scale_size
|
||||||
|
in_scale = tl.load(scale_src_ptr + offset, mask=mask)
|
||||||
|
tl.store(scale_dst_ptr + offset, in_scale, mask=mask)
|
||||||
|
|
||||||
|
|
||||||
def moe_ep_deepgemm_preprocess(
|
def moe_ep_deepgemm_preprocess(
|
||||||
@@ -1210,15 +1218,19 @@ def moe_ep_deepgemm_preprocess(
|
|||||||
block_shape = [128, 128]
|
block_shape = [128, 128]
|
||||||
assert len(block_shape) == 2
|
assert len(block_shape) == 2
|
||||||
block_n, block_k = block_shape[0], block_shape[1]
|
block_n, block_k = block_shape[0], block_shape[1]
|
||||||
|
is_fp8 = output_dtype == torch.float8_e4m3fn
|
||||||
|
if is_fp8:
|
||||||
|
# TODO: fuse this with the preprocess
|
||||||
|
hidden_states, scale = per_token_group_quant_fp8(hidden_states, block_k)
|
||||||
|
|
||||||
# TODO: fuse this with the preprocess
|
gateup_input_scale = torch.empty(
|
||||||
hidden_states, scale = per_token_group_quant_fp8(hidden_states, block_k)
|
(gateup_input.size(0), gateup_input.size(1), scale.size(1)),
|
||||||
|
device=hidden_states.device,
|
||||||
gateup_input_scale = torch.empty(
|
dtype=scale.dtype,
|
||||||
(gateup_input.size(0), gateup_input.size(1), scale.size(1)),
|
)
|
||||||
device=hidden_states.device,
|
else:
|
||||||
dtype=scale.dtype,
|
scale = None
|
||||||
)
|
gateup_input_scale = None
|
||||||
|
|
||||||
fill_gateup_input_triton_kernel[(hidden_states.shape[0],)](
|
fill_gateup_input_triton_kernel[(hidden_states.shape[0],)](
|
||||||
hidden_states,
|
hidden_states,
|
||||||
@@ -1229,8 +1241,9 @@ def moe_ep_deepgemm_preprocess(
|
|||||||
topk_ids,
|
topk_ids,
|
||||||
top_k,
|
top_k,
|
||||||
hidden_states.size(1),
|
hidden_states.size(1),
|
||||||
scale.size(1),
|
scale.size(1) if is_fp8 else 0,
|
||||||
BLOCK_SIZE=1024,
|
BLOCK_SIZE=1024,
|
||||||
|
IS_FP8=is_fp8,
|
||||||
)
|
)
|
||||||
|
|
||||||
return (
|
return (
|
||||||
|
|||||||
@@ -585,6 +585,11 @@ def pre_permute_standard_to_deep_gemm(
|
|||||||
topk_weights, topk_ids = topk_weights, topk_ids
|
topk_weights, topk_ids = topk_weights, topk_ids
|
||||||
|
|
||||||
# PreReorder
|
# PreReorder
|
||||||
|
output_dtype = (
|
||||||
|
torch.bfloat16
|
||||||
|
if quant_info.w13_weight.dtype == torch.bfloat16
|
||||||
|
else torch.float8_e4m3fn
|
||||||
|
)
|
||||||
masked_m, expected_m, src2dst, hidden_states, hidden_states_scale = (
|
masked_m, expected_m, src2dst, hidden_states, hidden_states_scale = (
|
||||||
moe_ep_deepgemm_preprocess(
|
moe_ep_deepgemm_preprocess(
|
||||||
topk_ids,
|
topk_ids,
|
||||||
@@ -592,6 +597,7 @@ def pre_permute_standard_to_deep_gemm(
|
|||||||
hidden_states,
|
hidden_states,
|
||||||
runner_config.top_k,
|
runner_config.top_k,
|
||||||
quant_info.block_shape,
|
quant_info.block_shape,
|
||||||
|
output_dtype=output_dtype,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user