Skip to content

Commit 0d00131

Browse files
authored
moe_align_block_size op debug (#1240)
* fix(moe): Initialize output tensors and correct padding logic in moe_align_block_size_triton. * add moe_align_block_size benchmark & test * fix format * fix format * add comment
1 parent 7252103 commit 0d00131

4 files changed

Lines changed: 130 additions & 1 deletion

File tree

benchmark/test_special_perf.py

Lines changed: 69 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -561,3 +561,72 @@ def torch_op(input_tensor, output_tensor):
561561
)
562562
bench.set_gems(gems_op)
563563
bench.run()
564+
565+
566+
try:
567+
import os
568+
569+
os.environ["VLLM_CONFIGURE_LOGGING"] = "0"
570+
import vllm._custom_ops as vllm_ops
571+
572+
HAS_VLLM = True
573+
except ImportError:
574+
HAS_VLLM = False
575+
576+
577+
@pytest.mark.moe_align_block_size
578+
@pytest.mark.skipif(not HAS_VLLM, reason="vllm not installed")
579+
def test_perf_moe_align_block_size():
580+
def moe_align_block_size_input_fn(shape, dtype, device):
581+
# ------------ parameters ------------
582+
num_experts = shape[0]
583+
block_size = shape[1]
584+
dtype = torch.int32
585+
topk_ids = torch.randint(0, num_experts, (3, 4), dtype=dtype, device=device)
586+
max_num_tokens_padded = topk_ids.numel() + num_experts * (block_size - 1)
587+
588+
# padded_num_experts in vllm._custom_ops.moe_align_block_size
589+
# must be less than 1024
590+
if max_num_tokens_padded >= 1024:
591+
return
592+
593+
sorted_ids = torch.empty((max_num_tokens_padded,), dtype=dtype, device=device)
594+
max_num_m_blocks = max_num_tokens_padded // block_size
595+
expert_ids = torch.empty((max_num_m_blocks,), dtype=dtype, device=device)
596+
num_tokens_post_pad = torch.empty(1, dtype=dtype, device=device)
597+
598+
yield (
599+
topk_ids,
600+
num_experts,
601+
block_size,
602+
sorted_ids,
603+
expert_ids,
604+
num_tokens_post_pad,
605+
)
606+
607+
class MoeAlignBlockSizeBenchmark(GenericBenchmark2DOnly):
608+
def set_more_shapes(self):
609+
return [
610+
(16, 8),
611+
(16, 16),
612+
(16, 32),
613+
(32, 8),
614+
(32, 16),
615+
(32, 32),
616+
(64, 8),
617+
(64, 16),
618+
(128, 8),
619+
]
620+
621+
gems_op = flag_gems.moe_align_block_size_triton
622+
bench = MoeAlignBlockSizeBenchmark(
623+
op_name="moe_align_block_size_triton",
624+
input_fn=moe_align_block_size_input_fn,
625+
torch_op=vllm_ops.moe_align_block_size,
626+
dtypes=[
627+
torch.int32,
628+
],
629+
)
630+
631+
bench.set_gems(gems_op)
632+
bench.run()

src/flag_gems/fused/moe_align_block_size.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -107,6 +107,11 @@ def moe_align_block_size_triton(
107107
num_tokens_post_pad: torch.Tensor,
108108
) -> None:
109109
numel = topk_ids.numel()
110+
# The tensor needs to be padded before calculating IDs,
111+
# to prevent out-of-bounds address access.
112+
sorted_token_ids.fill_(numel)
113+
expert_ids.fill_(0)
114+
110115
grid = (num_experts,)
111116
tokens_cnts = torch.zeros(
112117
(num_experts + 1, num_experts), dtype=torch.int32, device=topk_ids.device

tests/test_DSA/test_bin_topk.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -121,7 +121,7 @@ def debug_topk_results(actual, expected, inputs, test_name=""):
121121
print(f" Actual indices: {sorted(actual_set)[:m]}...") # Only show first 10
122122
print(f" Expected indices: {sorted(expected_set)[:m]}...")
123123
print(
124-
f" Intersection: {len(intersection)}/{len(expected_set)} = {len(intersection)/len(expected_set):.4f}"
124+
f" Intersection: {len(intersection)}/{len(expected_set)} = {len(intersection) / len(expected_set):.4f}"
125125
)
126126

127127
# Check quality of actually selected values

tests/test_special_ops.py

Lines changed: 55 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1389,3 +1389,58 @@ def test_moe_sum(shape, dtype):
13891389
with flag_gems.use_gems():
13901390
flag_gems.moe_sum(inp1, res_out)
13911391
gems_assert_close(res_out, ref_out, dtype)
1392+
1393+
1394+
try:
1395+
import vllm._custom_ops as vllm_ops
1396+
1397+
HAS_VLLM = True
1398+
except ImportError:
1399+
HAS_VLLM = False
1400+
1401+
1402+
# ref: https://github.com/vllm-project/vllm/blob/main/tests/kernels/moe/test_moe.py
1403+
@pytest.mark.moe_align_block_size
1404+
@pytest.mark.parametrize("num_experts", [32, 256, 512])
1405+
@pytest.mark.parametrize("block_size", [8, 16, 32])
1406+
@pytest.mark.skipif(not HAS_VLLM, reason="vllm not installed")
1407+
def test_accuracy_moe_align_block_size(
1408+
num_experts,
1409+
block_size,
1410+
):
1411+
# ------------ parameters ------------
1412+
dtype = torch.int32
1413+
topk_ids = torch.randint(0, num_experts, (3, 4), dtype=dtype, device="cuda")
1414+
max_num_tokens_padded = topk_ids.numel() + num_experts * (block_size - 1)
1415+
sorted_ids = torch.empty((max_num_tokens_padded,), dtype=dtype, device="cuda")
1416+
max_num_m_blocks = max_num_tokens_padded // block_size
1417+
expert_ids = torch.empty((max_num_m_blocks,), dtype=dtype, device="cuda")
1418+
num_tokens_post_pad = torch.empty(1, dtype=dtype, device="cuda")
1419+
1420+
topk_ids_vllm = topk_ids.clone()
1421+
sorted_ids_vllm = sorted_ids.clone()
1422+
expert_ids_vllm = expert_ids.clone()
1423+
num_tokens_post_pad_vllm = num_tokens_post_pad.clone()
1424+
1425+
flag_gems.moe_align_block_size_triton(
1426+
topk_ids=topk_ids,
1427+
num_experts=num_experts,
1428+
block_size=block_size,
1429+
sorted_token_ids=sorted_ids,
1430+
expert_ids=expert_ids,
1431+
num_tokens_post_pad=num_tokens_post_pad,
1432+
)
1433+
1434+
vllm_ops.moe_align_block_size(
1435+
topk_ids=topk_ids_vllm,
1436+
num_experts=num_experts,
1437+
block_size=block_size,
1438+
sorted_token_ids=sorted_ids_vllm,
1439+
experts_ids=expert_ids_vllm,
1440+
num_tokens_post_pad=num_tokens_post_pad_vllm,
1441+
)
1442+
1443+
torch.cuda.synchronize()
1444+
gems_assert_close(sorted_ids, sorted_ids_vllm, dtype=dtype)
1445+
gems_assert_close(expert_ids, expert_ids_vllm, dtype=dtype)
1446+
gems_assert_close(num_tokens_post_pad, num_tokens_post_pad_vllm, dtype=dtype)

0 commit comments

Comments
 (0)