support fused_moe_triton and moe_sum_all_reduce kernel fusion[reduce … (#19672)
Co-authored-by: undefined <zhouchen.arrebol@jd.com>
This commit is contained in:
@@ -47,6 +47,8 @@ _use_sgl_xpu = use_intel_xpu_backend()
|
|||||||
|
|
||||||
if _is_cuda:
|
if _is_cuda:
|
||||||
from sgl_kernel import gelu_and_mul, moe_sum_reduce, silu_and_mul
|
from sgl_kernel import gelu_and_mul, moe_sum_reduce, silu_and_mul
|
||||||
|
|
||||||
|
from sglang.srt.server_args import get_global_server_args
|
||||||
elif _is_cpu and _is_cpu_amx_available:
|
elif _is_cpu and _is_cpu_amx_available:
|
||||||
pass
|
pass
|
||||||
elif _is_hip:
|
elif _is_hip:
|
||||||
@@ -466,6 +468,14 @@ def fused_experts_impl(
|
|||||||
curr_topk_ids = topk_ids[begin_chunk_idx:end_chunk_idx]
|
curr_topk_ids = topk_ids[begin_chunk_idx:end_chunk_idx]
|
||||||
curr_topk_weights = topk_weights[begin_chunk_idx:end_chunk_idx]
|
curr_topk_weights = topk_weights[begin_chunk_idx:end_chunk_idx]
|
||||||
|
|
||||||
|
use_fused_moe_sum_all_reduce = (
|
||||||
|
get_global_server_args().enable_fused_moe_sum_all_reduce
|
||||||
|
and (not no_combine)
|
||||||
|
and (curr_topk_ids.shape[1] > 2)
|
||||||
|
and (not use_int8_w8a16)
|
||||||
|
and (not use_int4_w4a16)
|
||||||
|
)
|
||||||
|
|
||||||
sorted_token_ids, expert_ids, num_tokens_post_padded = moe_align_block_size(
|
sorted_token_ids, expert_ids, num_tokens_post_padded = moe_align_block_size(
|
||||||
curr_topk_ids, config["BLOCK_SIZE_M"], E
|
curr_topk_ids, config["BLOCK_SIZE_M"], E
|
||||||
)
|
)
|
||||||
@@ -569,14 +579,23 @@ def fused_experts_impl(
|
|||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported activation: {activation=}, with {is_gated=}")
|
raise ValueError(f"Unsupported activation: {activation=}, with {is_gated=}")
|
||||||
|
|
||||||
|
out_slice = None
|
||||||
|
if use_fused_moe_sum_all_reduce:
|
||||||
|
out_slice = out_hidden_states[begin_chunk_idx:end_chunk_idx]
|
||||||
|
out_slice.zero_()
|
||||||
|
|
||||||
invoke_fused_moe_kernel(
|
invoke_fused_moe_kernel(
|
||||||
intermediate_cache2,
|
intermediate_cache2,
|
||||||
w2,
|
w2,
|
||||||
b2,
|
b2,
|
||||||
(
|
(
|
||||||
|
out_slice
|
||||||
|
if use_fused_moe_sum_all_reduce
|
||||||
|
else (
|
||||||
intermediate_cache3
|
intermediate_cache3
|
||||||
if not no_combine and topk_ids.shape[1] != 1
|
if not no_combine and topk_ids.shape[1] != 1
|
||||||
else out_hidden_states[begin_chunk_idx:end_chunk_idx].unsqueeze(0)
|
else out_hidden_states[begin_chunk_idx:end_chunk_idx].unsqueeze(0)
|
||||||
|
)
|
||||||
),
|
),
|
||||||
a2_scale,
|
a2_scale,
|
||||||
w2_scale,
|
w2_scale,
|
||||||
@@ -599,6 +618,8 @@ def fused_experts_impl(
|
|||||||
a_use_tma=down_moe_use_tma,
|
a_use_tma=down_moe_use_tma,
|
||||||
b_use_tma=down_moe_use_tma,
|
b_use_tma=down_moe_use_tma,
|
||||||
filter_expert=filter_expert,
|
filter_expert=filter_expert,
|
||||||
|
fuse_sum_all_reduce=use_fused_moe_sum_all_reduce,
|
||||||
|
router_topk=curr_topk_ids.shape[1],
|
||||||
)
|
)
|
||||||
|
|
||||||
if routed_scaling_factor is None:
|
if routed_scaling_factor is None:
|
||||||
@@ -607,7 +628,13 @@ def fused_experts_impl(
|
|||||||
if no_combine:
|
if no_combine:
|
||||||
pass
|
pass
|
||||||
elif _is_cuda:
|
elif _is_cuda:
|
||||||
if topk_ids.shape[1] == 1 and routed_scaling_factor == 1.0:
|
if use_fused_moe_sum_all_reduce:
|
||||||
|
if routed_scaling_factor is None:
|
||||||
|
routed_scaling_factor = 1.0
|
||||||
|
if routed_scaling_factor != 1.0:
|
||||||
|
assert out_slice is not None
|
||||||
|
out_slice.mul_(routed_scaling_factor)
|
||||||
|
elif topk_ids.shape[1] == 1 and routed_scaling_factor == 1.0:
|
||||||
pass # we write directly into out_hidden_states
|
pass # we write directly into out_hidden_states
|
||||||
elif topk_ids.shape[1] == 2 and routed_scaling_factor == 1.0:
|
elif topk_ids.shape[1] == 2 and routed_scaling_factor == 1.0:
|
||||||
torch.add(
|
torch.add(
|
||||||
|
|||||||
@@ -377,6 +377,8 @@ def fused_moe_kernel(
|
|||||||
c_sorted: tl.constexpr,
|
c_sorted: tl.constexpr,
|
||||||
filter_expert: tl.constexpr,
|
filter_expert: tl.constexpr,
|
||||||
swap_ab: tl.constexpr,
|
swap_ab: tl.constexpr,
|
||||||
|
FUSE_SUM_ALL_REDUCE: tl.constexpr,
|
||||||
|
ROUTER_TOPK: tl.constexpr,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Implements the fused computation for a Mixture of Experts (MOE) using
|
Implements the fused computation for a Mixture of Experts (MOE) using
|
||||||
@@ -601,12 +603,25 @@ def fused_moe_kernel(
|
|||||||
# -----------------------------------------------------------
|
# -----------------------------------------------------------
|
||||||
# Write back the block of the output
|
# Write back the block of the output
|
||||||
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
|
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
|
||||||
|
|
||||||
|
if FUSE_SUM_ALL_REDUCE:
|
||||||
|
offs_token_out = offs_token // ROUTER_TOPK
|
||||||
|
c_ptrs = (
|
||||||
|
c_ptr + stride_cm * offs_token_out[:, None] + stride_cn * offs_cn[None, :]
|
||||||
|
)
|
||||||
|
c_mask = token_mask[:, None] & (offs_cn[None, :] < N)
|
||||||
|
tl.atomic_add(c_ptrs, accumulator, mask=c_mask)
|
||||||
|
else:
|
||||||
if c_sorted:
|
if c_sorted:
|
||||||
c_ptrs = (
|
c_ptrs = (
|
||||||
c_ptr + stride_cm * offs_token_id[:, None] + stride_cn * offs_cn[None, :]
|
c_ptr
|
||||||
|
+ stride_cm * offs_token_id[:, None]
|
||||||
|
+ stride_cn * offs_cn[None, :]
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
c_ptrs = c_ptr + stride_cm * offs_token[:, None] + stride_cn * offs_cn[None, :]
|
c_ptrs = (
|
||||||
|
c_ptr + stride_cm * offs_token[:, None] + stride_cn * offs_cn[None, :]
|
||||||
|
)
|
||||||
c_mask = token_mask[:, None] & (offs_cn[None, :] < N)
|
c_mask = token_mask[:, None] & (offs_cn[None, :] < N)
|
||||||
tl.store(c_ptrs, accumulator, mask=c_mask)
|
tl.store(c_ptrs, accumulator, mask=c_mask)
|
||||||
|
|
||||||
@@ -700,6 +715,8 @@ def invoke_fused_moe_kernel(
|
|||||||
b_use_tma: bool = False,
|
b_use_tma: bool = False,
|
||||||
c_sorted: bool = False,
|
c_sorted: bool = False,
|
||||||
filter_expert: bool = True,
|
filter_expert: bool = True,
|
||||||
|
fuse_sum_all_reduce: bool = False,
|
||||||
|
router_topk: int = 1,
|
||||||
) -> None:
|
) -> None:
|
||||||
assert topk_weights.stride(1) == 1
|
assert topk_weights.stride(1) == 1
|
||||||
assert sorted_token_ids.stride(0) == 1
|
assert sorted_token_ids.stride(0) == 1
|
||||||
@@ -767,11 +784,17 @@ def invoke_fused_moe_kernel(
|
|||||||
else:
|
else:
|
||||||
even_Ks = False
|
even_Ks = False
|
||||||
|
|
||||||
|
if fuse_sum_all_reduce:
|
||||||
|
assert not c_sorted, "fuse_sum_all_reduce only supports c_sorted=False"
|
||||||
|
|
||||||
if (
|
if (
|
||||||
(use_int8_w8a16 or use_int4_w4a16)
|
(use_int8_w8a16 or use_int4_w4a16)
|
||||||
and block_shape is not None
|
and block_shape is not None
|
||||||
and block_shape[1] > 0
|
and block_shape[1] > 0
|
||||||
):
|
):
|
||||||
|
assert (
|
||||||
|
not fuse_sum_all_reduce
|
||||||
|
), "fuse_sum_all_reduce is not supported for GPTQ/AWQ kernels"
|
||||||
assert B_scale is not None and B_scale.ndim == 3
|
assert B_scale is not None and B_scale.ndim == 3
|
||||||
assert B_zp is None or B_zp.ndim == 3
|
assert B_zp is None or B_zp.ndim == 3
|
||||||
assert bias is None
|
assert bias is None
|
||||||
@@ -878,6 +901,8 @@ def invoke_fused_moe_kernel(
|
|||||||
c_sorted=c_sorted,
|
c_sorted=c_sorted,
|
||||||
filter_expert=filter_expert,
|
filter_expert=filter_expert,
|
||||||
swap_ab=swap_ab,
|
swap_ab=swap_ab,
|
||||||
|
FUSE_SUM_ALL_REDUCE=fuse_sum_all_reduce,
|
||||||
|
ROUTER_TOPK=router_topk,
|
||||||
**config,
|
**config,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -646,6 +646,7 @@ class ServerArgs:
|
|||||||
nsa_prefill_cp_mode: str = "round-robin-split"
|
nsa_prefill_cp_mode: str = "round-robin-split"
|
||||||
enable_fused_qk_norm_rope: bool = False
|
enable_fused_qk_norm_rope: bool = False
|
||||||
enable_precise_embedding_interpolation: bool = False
|
enable_precise_embedding_interpolation: bool = False
|
||||||
|
enable_fused_moe_sum_all_reduce: bool = False
|
||||||
|
|
||||||
# Dynamic batch tokenizer
|
# Dynamic batch tokenizer
|
||||||
enable_dynamic_batch_tokenizer: bool = False
|
enable_dynamic_batch_tokenizer: bool = False
|
||||||
@@ -5032,6 +5033,11 @@ class ServerArgs:
|
|||||||
action="store_true",
|
action="store_true",
|
||||||
help="Enable corner alignment for resize of embeddings grid to ensure more accurate(but slower) evaluation of interpolated embedding values.",
|
help="Enable corner alignment for resize of embeddings grid to ensure more accurate(but slower) evaluation of interpolated embedding values.",
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--enable-fused-moe-sum-all-reduce",
|
||||||
|
action="store_true",
|
||||||
|
help="Enable fused moe triton and sum all reduce.",
|
||||||
|
)
|
||||||
|
|
||||||
# Dynamic batch tokenizer
|
# Dynamic batch tokenizer
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
|
|||||||
Reference in New Issue
Block a user