[Perf] Optimize DeepSeek-R1 w4afp8 glue kernels (#10027)
Co-authored-by: Fan Yin <1106310035@qq.com>
This commit is contained in:
@@ -10,15 +10,17 @@ from sgl_kernel import (
|
|||||||
silu_and_mul,
|
silu_and_mul,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
from sglang.srt.distributed import get_moe_expert_parallel_world_size
|
||||||
from sglang.srt.layers.moe.ep_moe.kernels import (
|
from sglang.srt.layers.moe.ep_moe.kernels import (
|
||||||
|
cutlass_w4_run_moe_ep_preproess,
|
||||||
deepep_ll_get_cutlass_w4a8_moe_mm_data,
|
deepep_ll_get_cutlass_w4a8_moe_mm_data,
|
||||||
deepep_permute_triton_kernel,
|
deepep_permute_triton_kernel,
|
||||||
deepep_post_reorder_triton_kernel,
|
deepep_post_reorder_triton_kernel,
|
||||||
deepep_run_moe_deep_preprocess,
|
deepep_run_moe_deep_preprocess,
|
||||||
post_reorder_triton_kernel_for_cutlass_moe,
|
post_reorder_for_cutlass_moe,
|
||||||
pre_reorder_triton_kernel_for_cutlass_moe,
|
pre_reorder_for_cutlass_moe,
|
||||||
run_moe_ep_preproess,
|
|
||||||
silu_and_mul_masked_post_per_tensor_quant_fwd,
|
silu_and_mul_masked_post_per_tensor_quant_fwd,
|
||||||
|
silu_mul_static_tensorwise_quant_for_cutlass_moe,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -44,6 +46,7 @@ def cutlass_w4a8_moe(
|
|||||||
a1_scale: Optional[torch.Tensor] = None,
|
a1_scale: Optional[torch.Tensor] = None,
|
||||||
a2_scale: Optional[torch.Tensor] = None,
|
a2_scale: Optional[torch.Tensor] = None,
|
||||||
apply_router_weight_on_input: bool = False,
|
apply_router_weight_on_input: bool = False,
|
||||||
|
routed_scaling_factor: float = 1.0,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
This function computes a w4a8-quantized Mixture of Experts (MoE) layer
|
This function computes a w4a8-quantized Mixture of Experts (MoE) layer
|
||||||
@@ -108,11 +111,11 @@ def cutlass_w4a8_moe(
|
|||||||
assert topk == 1, "apply_router_weight_on_input is only implemented for topk=1"
|
assert topk == 1, "apply_router_weight_on_input is only implemented for topk=1"
|
||||||
|
|
||||||
device = a.device
|
device = a.device
|
||||||
topk_ids = torch.where(topk_ids == -1, num_local_experts, topk_ids)
|
if get_moe_expert_parallel_world_size() > 1:
|
||||||
|
topk_ids = torch.where(topk_ids == -1, num_local_experts, topk_ids)
|
||||||
|
|
||||||
_, src2dst, _ = run_moe_ep_preproess(
|
src2dst = cutlass_w4_run_moe_ep_preproess(
|
||||||
topk_ids,
|
topk_ids,
|
||||||
num_local_experts,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
gateup_input = torch.empty(
|
gateup_input = torch.empty(
|
||||||
@@ -121,7 +124,7 @@ def cutlass_w4a8_moe(
|
|||||||
dtype=torch.float8_e4m3fn,
|
dtype=torch.float8_e4m3fn,
|
||||||
)
|
)
|
||||||
|
|
||||||
pre_reorder_triton_kernel_for_cutlass_moe[(m,)](
|
pre_reorder_for_cutlass_moe(
|
||||||
a,
|
a,
|
||||||
gateup_input,
|
gateup_input,
|
||||||
src2dst,
|
src2dst,
|
||||||
@@ -129,8 +132,8 @@ def cutlass_w4a8_moe(
|
|||||||
a1_scale,
|
a1_scale,
|
||||||
num_local_experts,
|
num_local_experts,
|
||||||
topk,
|
topk,
|
||||||
|
m,
|
||||||
k,
|
k,
|
||||||
BLOCK_SIZE=512,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# NOTE: a_map and c_map are not used in the get_cutlass_w4a8_moe_mm_data kernel,
|
# NOTE: a_map and c_map are not used in the get_cutlass_w4a8_moe_mm_data kernel,
|
||||||
@@ -151,7 +154,7 @@ def cutlass_w4a8_moe(
|
|||||||
)
|
)
|
||||||
|
|
||||||
c1 = torch.empty((m * topk, n * 2), device=device, dtype=torch.bfloat16)
|
c1 = torch.empty((m * topk, n * 2), device=device, dtype=torch.bfloat16)
|
||||||
c2 = torch.zeros((m * topk, k), device=device, dtype=torch.bfloat16)
|
c2 = torch.empty((m * topk, k), device=device, dtype=torch.bfloat16)
|
||||||
|
|
||||||
cutlass_w4a8_moe_mm(
|
cutlass_w4a8_moe_mm(
|
||||||
c1,
|
c1,
|
||||||
@@ -169,13 +172,12 @@ def cutlass_w4a8_moe(
|
|||||||
topk,
|
topk,
|
||||||
)
|
)
|
||||||
|
|
||||||
intermediate = torch.empty((m * topk, n), device=device, dtype=torch.bfloat16)
|
|
||||||
silu_and_mul(c1, intermediate)
|
|
||||||
|
|
||||||
intermediate_q = torch.empty(
|
intermediate_q = torch.empty(
|
||||||
intermediate.shape, dtype=torch.float8_e4m3fn, device=device
|
(m * topk, n), dtype=torch.float8_e4m3fn, device=device
|
||||||
|
)
|
||||||
|
silu_mul_static_tensorwise_quant_for_cutlass_moe(
|
||||||
|
c1, intermediate_q, a2_scale.float(), expert_offsets[-1:], m * topk, n
|
||||||
)
|
)
|
||||||
sgl_per_tensor_quant_fp8(intermediate, intermediate_q, a2_scale.float(), True)
|
|
||||||
|
|
||||||
cutlass_w4a8_moe_mm(
|
cutlass_w4a8_moe_mm(
|
||||||
c2,
|
c2,
|
||||||
@@ -194,16 +196,18 @@ def cutlass_w4a8_moe(
|
|||||||
)
|
)
|
||||||
|
|
||||||
output = torch.empty_like(a)
|
output = torch.empty_like(a)
|
||||||
post_reorder_triton_kernel_for_cutlass_moe[(m,)](
|
|
||||||
|
post_reorder_for_cutlass_moe(
|
||||||
c2,
|
c2,
|
||||||
output,
|
output,
|
||||||
src2dst,
|
src2dst,
|
||||||
topk_ids,
|
topk_ids,
|
||||||
topk_weights,
|
topk_weights,
|
||||||
topk,
|
|
||||||
num_local_experts,
|
num_local_experts,
|
||||||
|
topk,
|
||||||
|
m,
|
||||||
k,
|
k,
|
||||||
BLOCK_SIZE=512,
|
routed_scaling_factor,
|
||||||
)
|
)
|
||||||
return output
|
return output
|
||||||
|
|
||||||
|
|||||||
@@ -16,6 +16,60 @@ if _is_cuda:
|
|||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
|
|
||||||
|
|
||||||
|
def _get_launch_config_1d(device, numel):
|
||||||
|
MAX_THREADS_PER_BLOCK = 1024
|
||||||
|
MIN_THREADS_PER_BLOCK = 512
|
||||||
|
MAX_WAVES = 8 # empirical numbers
|
||||||
|
|
||||||
|
props = torch.cuda.get_device_properties(device)
|
||||||
|
sm_count = props.multi_processor_count
|
||||||
|
max_threads_per_sm = props.max_threads_per_multi_processor
|
||||||
|
max_num_blocks = sm_count * max_threads_per_sm // MAX_THREADS_PER_BLOCK
|
||||||
|
|
||||||
|
block_dim = MAX_THREADS_PER_BLOCK
|
||||||
|
|
||||||
|
def get_num_blocks(block_dim):
|
||||||
|
return triton.cdiv(numel, block_dim)
|
||||||
|
|
||||||
|
while (
|
||||||
|
block_dim > MIN_THREADS_PER_BLOCK
|
||||||
|
and get_num_blocks(block_dim // 2) <= max_num_blocks
|
||||||
|
):
|
||||||
|
block_dim = block_dim // 2
|
||||||
|
|
||||||
|
num_blocks = get_num_blocks(block_dim)
|
||||||
|
grid_dim = min(num_blocks, max_num_blocks * MAX_WAVES)
|
||||||
|
|
||||||
|
return (grid_dim,), block_dim
|
||||||
|
|
||||||
|
|
||||||
|
def _get_launch_config_2d(device, m, n):
|
||||||
|
MAX_THREADS_PER_BLOCK = 1024
|
||||||
|
MIN_THREADS_PER_BLOCK = 512
|
||||||
|
MAX_WAVES = 8 # empirical numbers
|
||||||
|
|
||||||
|
props = torch.cuda.get_device_properties(device)
|
||||||
|
sm_count = props.multi_processor_count
|
||||||
|
max_threads_per_sm = props.max_threads_per_multi_processor
|
||||||
|
max_num_blocks = sm_count * max_threads_per_sm // MAX_THREADS_PER_BLOCK
|
||||||
|
|
||||||
|
block_dim = MAX_THREADS_PER_BLOCK
|
||||||
|
|
||||||
|
def get_num_blocks(block_dim):
|
||||||
|
return m * triton.cdiv(n, block_dim)
|
||||||
|
|
||||||
|
while (
|
||||||
|
block_dim > MIN_THREADS_PER_BLOCK
|
||||||
|
and get_num_blocks(block_dim // 2) <= max_num_blocks
|
||||||
|
):
|
||||||
|
block_dim = block_dim // 2
|
||||||
|
|
||||||
|
grid_dim_x = triton.cdiv(n, block_dim)
|
||||||
|
grid_dim_y = max(min(m, max_num_blocks * MAX_WAVES // grid_dim_x), 1)
|
||||||
|
|
||||||
|
return (grid_dim_y, grid_dim_x), block_dim
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
def deepep_permute_triton_kernel(
|
def deepep_permute_triton_kernel(
|
||||||
input_ptr,
|
input_ptr,
|
||||||
@@ -142,25 +196,17 @@ def compute_seg_indptr_triton_kernel(reorder_topk_ids, seg_indptr, num_toks):
|
|||||||
tl.store(seg_indptr + expert_id_minus_1 + 1, target_location + 1)
|
tl.store(seg_indptr + expert_id_minus_1 + 1, target_location + 1)
|
||||||
|
|
||||||
|
|
||||||
def run_moe_ep_preproess(topk_ids: torch.Tensor, num_local_experts: int):
|
def cutlass_w4_run_moe_ep_preproess(topk_ids: torch.Tensor):
|
||||||
reorder_topk_ids, reorder_ids = torch.sort(topk_ids.view(-1), stable=True)
|
_, reorder_ids = torch.sort(topk_ids.view(-1), stable=True)
|
||||||
|
|
||||||
seg_indptr = torch.zeros(
|
|
||||||
num_local_experts + 1, device=topk_ids.device, dtype=torch.int64
|
|
||||||
)
|
|
||||||
src2dst = torch.empty(topk_ids.numel(), device=topk_ids.device, dtype=torch.int32)
|
|
||||||
|
|
||||||
compute_seg_indptr_triton_kernel[(num_local_experts,)](
|
|
||||||
reorder_topk_ids, seg_indptr, topk_ids.numel()
|
|
||||||
)
|
|
||||||
|
|
||||||
BLOCK_SIZE = 512
|
BLOCK_SIZE = 512
|
||||||
grid = (triton.cdiv(topk_ids.numel(), BLOCK_SIZE),)
|
grid = (triton.cdiv(topk_ids.numel(), BLOCK_SIZE),)
|
||||||
|
src2dst = torch.empty(topk_ids.numel(), device=topk_ids.device, dtype=torch.int32)
|
||||||
compute_src2dst_triton_kernel[grid](
|
compute_src2dst_triton_kernel[grid](
|
||||||
reorder_ids, src2dst, topk_ids.numel(), BLOCK_SIZE
|
reorder_ids, src2dst, topk_ids.numel(), BLOCK_SIZE
|
||||||
)
|
)
|
||||||
|
|
||||||
return reorder_topk_ids, src2dst, seg_indptr
|
return src2dst
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
@@ -172,36 +218,68 @@ def pre_reorder_triton_kernel_for_cutlass_moe(
|
|||||||
a1_scales_ptr,
|
a1_scales_ptr,
|
||||||
num_local_experts,
|
num_local_experts,
|
||||||
topk,
|
topk,
|
||||||
|
num_tokens,
|
||||||
hidden_size,
|
hidden_size,
|
||||||
BLOCK_SIZE: tl.constexpr,
|
BLOCK_SIZE: tl.constexpr,
|
||||||
|
NUM_STAGES: tl.constexpr,
|
||||||
):
|
):
|
||||||
OutDtype = gateup_input_ptr.dtype.element_ty
|
OutDtype = gateup_input_ptr.dtype.element_ty
|
||||||
|
|
||||||
src_idx_int32 = tl.program_id(0)
|
if a1_scales_ptr is not None:
|
||||||
src_idx = src_idx_int32.to(tl.int64)
|
a1_scale = 1.0 / tl.load(a1_scales_ptr)
|
||||||
src2dst_ptr = src2dst_ptr + src_idx * topk
|
else:
|
||||||
topk_ids_ptr = topk_ids_ptr + src_idx * topk
|
a1_scale = 1.0
|
||||||
src_ptr = input_ptr + src_idx * hidden_size
|
|
||||||
|
|
||||||
vec = tl.arange(0, BLOCK_SIZE)
|
offset = BLOCK_SIZE * tl.program_id(1) + tl.arange(0, BLOCK_SIZE)
|
||||||
|
mask = offset < hidden_size
|
||||||
|
|
||||||
for idx in range(topk):
|
start_src_idx = tl.program_id(0)
|
||||||
expert_id = tl.load(topk_ids_ptr + idx)
|
step = tl.num_programs(0)
|
||||||
if expert_id != num_local_experts:
|
|
||||||
if a1_scales_ptr is not None:
|
|
||||||
scale = 1.0 / tl.load(a1_scales_ptr)
|
|
||||||
else:
|
|
||||||
scale = 1.0
|
|
||||||
|
|
||||||
dst_idx_int32 = tl.load(src2dst_ptr + idx)
|
for src_idx_int32 in tl.range(
|
||||||
dst_idx = dst_idx_int32.to(tl.int64)
|
start_src_idx, num_tokens, step, num_stages=NUM_STAGES
|
||||||
dst_ptr = gateup_input_ptr + dst_idx * hidden_size
|
):
|
||||||
for start_offset in tl.range(0, hidden_size, BLOCK_SIZE):
|
src_idx = src_idx_int32.to(tl.int64)
|
||||||
offset = start_offset + vec
|
token_src2dst_ptr = src2dst_ptr + src_idx * topk
|
||||||
mask = offset < hidden_size
|
token_topk_ids_ptr = topk_ids_ptr + src_idx * topk
|
||||||
in_data = tl.load(src_ptr + offset, mask=mask).to(tl.float32)
|
|
||||||
out_data = (in_data * scale).to(OutDtype)
|
src_ptr_offs = input_ptr + src_idx * hidden_size + offset
|
||||||
tl.store(dst_ptr + offset, out_data, mask=mask)
|
dst_ptr_offs = gateup_input_ptr + offset
|
||||||
|
in_data = tl.load(src_ptr_offs, mask=mask).to(tl.float32)
|
||||||
|
out_data = (in_data * a1_scale).to(OutDtype)
|
||||||
|
for idx in range(topk):
|
||||||
|
expert_id = tl.load(token_topk_ids_ptr + idx)
|
||||||
|
if expert_id != num_local_experts:
|
||||||
|
dst_idx = tl.load(token_src2dst_ptr + idx)
|
||||||
|
tl.store(dst_ptr_offs + dst_idx * hidden_size, out_data, mask=mask)
|
||||||
|
|
||||||
|
|
||||||
|
def pre_reorder_for_cutlass_moe(
|
||||||
|
input,
|
||||||
|
gateup_input,
|
||||||
|
src2dst,
|
||||||
|
topk_ids,
|
||||||
|
a1_scales,
|
||||||
|
num_local_experts,
|
||||||
|
topk,
|
||||||
|
num_tokens,
|
||||||
|
hidden_size,
|
||||||
|
):
|
||||||
|
grid, block_dim = _get_launch_config_2d(input.device, num_tokens, hidden_size)
|
||||||
|
|
||||||
|
pre_reorder_triton_kernel_for_cutlass_moe[grid](
|
||||||
|
input_ptr=input,
|
||||||
|
gateup_input_ptr=gateup_input,
|
||||||
|
src2dst_ptr=src2dst,
|
||||||
|
topk_ids_ptr=topk_ids,
|
||||||
|
a1_scales_ptr=a1_scales,
|
||||||
|
num_local_experts=num_local_experts,
|
||||||
|
topk=topk,
|
||||||
|
num_tokens=num_tokens,
|
||||||
|
hidden_size=hidden_size,
|
||||||
|
BLOCK_SIZE=block_dim,
|
||||||
|
NUM_STAGES=3,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
# copy from https://github.com/ModelTC/lightllm/blob/a000ab69098654df4731f5b12587dd4e7f0a4f41/lightllm/common/fused_moe/moe_silu_and_mul_mix_quant_ep.py
|
# copy from https://github.com/ModelTC/lightllm/blob/a000ab69098654df4731f5b12587dd4e7f0a4f41/lightllm/common/fused_moe/moe_silu_and_mul_mix_quant_ep.py
|
||||||
@@ -351,6 +429,62 @@ def silu_and_mul_masked_post_quant_fwd(
|
|||||||
return
|
return
|
||||||
|
|
||||||
|
|
||||||
|
@triton.jit
|
||||||
|
def silu_mul_static_tensorwise_quant_triton_kernel_for_cutlass_moe(
|
||||||
|
input_ptr,
|
||||||
|
output_ptr,
|
||||||
|
scale_ptr,
|
||||||
|
num_tokens_tensor_ptr,
|
||||||
|
intermediate_size,
|
||||||
|
BLOCK_SIZE: tl.constexpr,
|
||||||
|
NUM_STAGES: tl.constexpr,
|
||||||
|
):
|
||||||
|
OutDtype = output_ptr.dtype.element_ty
|
||||||
|
|
||||||
|
num_tokens = tl.load(num_tokens_tensor_ptr)
|
||||||
|
numel = num_tokens * intermediate_size
|
||||||
|
gate_ptr = input_ptr
|
||||||
|
up_ptr = input_ptr + intermediate_size
|
||||||
|
scale = 1.0 / tl.load(scale_ptr)
|
||||||
|
|
||||||
|
start_idx = tl.program_id(0) * BLOCK_SIZE
|
||||||
|
step = tl.num_programs(0) * BLOCK_SIZE
|
||||||
|
|
||||||
|
for id in tl.range(start_idx, numel, step, num_stages=NUM_STAGES):
|
||||||
|
ids = id + tl.arange(0, BLOCK_SIZE)
|
||||||
|
token_ids = ids // intermediate_size
|
||||||
|
mask = ids < numel
|
||||||
|
|
||||||
|
offs = ids + token_ids * intermediate_size
|
||||||
|
gate = tl.load(gate_ptr + offs, mask=mask, other=0.0).to(tl.float32)
|
||||||
|
up = tl.load(up_ptr + offs, mask=mask, other=0.0).to(tl.float32)
|
||||||
|
output = gate / (1 + tl.exp(-gate)) * up * scale
|
||||||
|
tl.store(output_ptr + ids, output.to(OutDtype), mask=mask)
|
||||||
|
|
||||||
|
|
||||||
|
def silu_mul_static_tensorwise_quant_for_cutlass_moe(
|
||||||
|
input: torch.Tensor,
|
||||||
|
output: torch.Tensor,
|
||||||
|
scale: torch.Tensor,
|
||||||
|
num_tokens_tensor: torch.Tensor,
|
||||||
|
expected_num_tokens: int,
|
||||||
|
intermediate_size: int,
|
||||||
|
):
|
||||||
|
grid, block_dim = _get_launch_config_1d(
|
||||||
|
input.device, expected_num_tokens * intermediate_size
|
||||||
|
)
|
||||||
|
|
||||||
|
silu_mul_static_tensorwise_quant_triton_kernel_for_cutlass_moe[grid](
|
||||||
|
input_ptr=input,
|
||||||
|
output_ptr=output,
|
||||||
|
scale_ptr=scale,
|
||||||
|
num_tokens_tensor_ptr=num_tokens_tensor,
|
||||||
|
intermediate_size=intermediate_size,
|
||||||
|
BLOCK_SIZE=block_dim,
|
||||||
|
NUM_STAGES=3,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
def post_reorder_triton_kernel_for_cutlass_moe(
|
def post_reorder_triton_kernel_for_cutlass_moe(
|
||||||
down_output_ptr,
|
down_output_ptr,
|
||||||
@@ -358,38 +492,77 @@ def post_reorder_triton_kernel_for_cutlass_moe(
|
|||||||
src2dst_ptr,
|
src2dst_ptr,
|
||||||
topk_ids_ptr,
|
topk_ids_ptr,
|
||||||
topk_weights_ptr,
|
topk_weights_ptr,
|
||||||
topk,
|
|
||||||
num_local_experts,
|
num_local_experts,
|
||||||
|
topk,
|
||||||
|
num_tokens,
|
||||||
hidden_size,
|
hidden_size,
|
||||||
|
routed_scaling_factor: float,
|
||||||
BLOCK_SIZE: tl.constexpr,
|
BLOCK_SIZE: tl.constexpr,
|
||||||
|
NUM_STAGES: tl.constexpr,
|
||||||
):
|
):
|
||||||
InDtype = down_output_ptr.dtype.element_ty
|
OutDtype = output_ptr.dtype.element_ty
|
||||||
|
|
||||||
src_idx_int32 = tl.program_id(0)
|
offset = BLOCK_SIZE * tl.program_id(1) + tl.arange(0, BLOCK_SIZE)
|
||||||
src_idx = src_idx_int32.to(tl.int64)
|
mask = offset < hidden_size
|
||||||
src2dst_ptr = src2dst_ptr + src_idx * topk
|
|
||||||
topk_ids_ptr = topk_ids_ptr + src_idx * topk
|
|
||||||
topk_weights_ptr = topk_weights_ptr + src_idx * topk
|
|
||||||
|
|
||||||
store_ptr = output_ptr + src_idx * hidden_size
|
down_output_ptr_offs = down_output_ptr + offset
|
||||||
|
output_ptr_offs = output_ptr + offset
|
||||||
|
|
||||||
vec = tl.arange(0, BLOCK_SIZE)
|
start_src_idx = tl.program_id(0)
|
||||||
|
step = tl.num_programs(0)
|
||||||
|
|
||||||
for start_offset in tl.range(0, hidden_size, BLOCK_SIZE):
|
for src_idx_int32 in tl.range(
|
||||||
offset = start_offset + vec
|
start_src_idx, num_tokens, step, num_stages=NUM_STAGES
|
||||||
mask = offset < hidden_size
|
):
|
||||||
|
src_idx = src_idx_int32.to(tl.int64)
|
||||||
|
token_src2dst_ptr = src2dst_ptr + src_idx * topk
|
||||||
|
token_topk_ids_ptr = topk_ids_ptr + src_idx * topk
|
||||||
|
token_topk_weights_ptr = topk_weights_ptr + src_idx * topk
|
||||||
|
|
||||||
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(token_topk_ids_ptr + idx)
|
||||||
if expert_id != num_local_experts:
|
if expert_id != num_local_experts:
|
||||||
dst_idx_int32 = tl.load(src2dst_ptr + idx)
|
dst_idx_int32 = tl.load(token_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)
|
dst_idx = dst_idx
|
||||||
load_ptr = down_output_ptr + dst_idx * hidden_size
|
weight_scale = tl.load(token_topk_weights_ptr + idx).to(tl.float32)
|
||||||
in_data = tl.load(load_ptr + offset, mask=mask)
|
load_ptr_offs = down_output_ptr_offs + dst_idx * hidden_size
|
||||||
sum_vec += in_data * weigh_scale
|
in_data = tl.load(load_ptr_offs, mask=mask).to(tl.float32)
|
||||||
tl.store(store_ptr + offset, sum_vec, mask=mask)
|
sum_vec += in_data * weight_scale
|
||||||
|
sum_vec *= routed_scaling_factor
|
||||||
|
store_ptr_offs = output_ptr_offs + src_idx * hidden_size
|
||||||
|
tl.store(store_ptr_offs, sum_vec.to(OutDtype), mask=mask)
|
||||||
|
|
||||||
|
|
||||||
|
def post_reorder_for_cutlass_moe(
|
||||||
|
down_output,
|
||||||
|
output,
|
||||||
|
src2dst,
|
||||||
|
topk_ids,
|
||||||
|
topk_weights,
|
||||||
|
num_local_experts,
|
||||||
|
topk,
|
||||||
|
num_tokens,
|
||||||
|
hidden_size,
|
||||||
|
routed_scaling_factor: float,
|
||||||
|
):
|
||||||
|
grid, block_dim = _get_launch_config_2d(down_output.device, num_tokens, hidden_size)
|
||||||
|
|
||||||
|
post_reorder_triton_kernel_for_cutlass_moe[grid](
|
||||||
|
down_output_ptr=down_output,
|
||||||
|
output_ptr=output,
|
||||||
|
src2dst_ptr=src2dst,
|
||||||
|
topk_ids_ptr=topk_ids,
|
||||||
|
topk_weights_ptr=topk_weights,
|
||||||
|
num_local_experts=num_local_experts,
|
||||||
|
topk=topk,
|
||||||
|
num_tokens=num_tokens,
|
||||||
|
hidden_size=hidden_size,
|
||||||
|
routed_scaling_factor=routed_scaling_factor,
|
||||||
|
BLOCK_SIZE=block_dim,
|
||||||
|
NUM_STAGES=3,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
|
|||||||
@@ -270,17 +270,17 @@ class W4AFp8MoEMethod(FusedMoEMethodBase):
|
|||||||
layer.w2_weight_scale_inv = Parameter(w2_weight_scale, requires_grad=False)
|
layer.w2_weight_scale_inv = Parameter(w2_weight_scale, requires_grad=False)
|
||||||
|
|
||||||
# Process input scales
|
# Process input scales
|
||||||
w13_input_scale_max = layer.w13_input_scale.max().to(dtype).item()
|
w13_input_scale_max = layer.w13_input_scale.max().to(torch.float32).item()
|
||||||
new_w13_input_scale = torch.tensor(
|
new_w13_input_scale = torch.tensor(
|
||||||
[w13_input_scale_max],
|
[w13_input_scale_max],
|
||||||
dtype=dtype,
|
dtype=torch.float32,
|
||||||
device=device,
|
device=device,
|
||||||
)
|
)
|
||||||
layer.w13_input_scale = Parameter(new_w13_input_scale, requires_grad=False)
|
layer.w13_input_scale = Parameter(new_w13_input_scale, requires_grad=False)
|
||||||
|
|
||||||
w2_input_scale_max = layer.w2_input_scale.max().to(dtype).item()
|
w2_input_scale_max = layer.w2_input_scale.max().to(torch.float32).item()
|
||||||
new_w2_input_scale = torch.tensor(
|
new_w2_input_scale = torch.tensor(
|
||||||
[w2_input_scale_max], dtype=dtype, device=device
|
[w2_input_scale_max], dtype=torch.float32, device=device
|
||||||
)
|
)
|
||||||
layer.w2_input_scale = Parameter(new_w2_input_scale, requires_grad=False)
|
layer.w2_input_scale = Parameter(new_w2_input_scale, requires_grad=False)
|
||||||
|
|
||||||
@@ -324,9 +324,8 @@ class W4AFp8MoEMethod(FusedMoEMethodBase):
|
|||||||
self.problem_sizes2,
|
self.problem_sizes2,
|
||||||
layer.w13_input_scale,
|
layer.w13_input_scale,
|
||||||
layer.w2_input_scale,
|
layer.w2_input_scale,
|
||||||
|
routed_scaling_factor=self.moe_runner_config.routed_scaling_factor or 1.0,
|
||||||
)
|
)
|
||||||
if self.moe_runner_config.routed_scaling_factor is not None:
|
|
||||||
output *= self.moe_runner_config.routed_scaling_factor
|
|
||||||
return StandardCombineInput(hidden_states=output)
|
return StandardCombineInput(hidden_states=output)
|
||||||
|
|
||||||
def apply_deepep_ll(
|
def apply_deepep_ll(
|
||||||
|
|||||||
Reference in New Issue
Block a user