Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
31 changes: 31 additions & 0 deletions tests/ut/_310p/fused_moe/test_experts_selector_310.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,3 +52,34 @@ def test_select_experts(self, global_num_experts):

assert topk_weights.shape == (8, 2)
assert topk_ids.shape == (8, 2)

def test_select_experts_chunks_large_token_batch(self):
num_tokens = 2050
hidden_states = torch.randn(num_tokens, 16)
router_logits = torch.randn(num_tokens, 8)

def mock_gating(logits, k):
return (
torch.ones(logits.shape[0], k),
torch.zeros(logits.shape[0], k, dtype=torch.int32),
None,
)

with patch(
"torch_npu.npu_moe_gating_top_k_softmax",
side_effect=mock_gating,
) as mock_npu:
topk_weights, topk_ids = select_experts(
hidden_states=hidden_states,
router_logits=router_logits,
top_k=2,
use_grouped_topk=False,
renormalize=True,
custom_routing_function=None,
scoring_func="softmax",
)

assert [call.args[0].shape[0] for call in mock_npu.call_args_list] == [1024, 1024, 2]
assert topk_weights.shape == (num_tokens, 2)
assert topk_ids.shape == (num_tokens, 2)
assert torch.all(topk_weights == 0.5)
11 changes: 10 additions & 1 deletion vllm_ascend/_310p/fused_moe/experts_selector.py
Original file line number Diff line number Diff line change
Expand Up @@ -58,7 +58,16 @@ def select_experts(
topk_ids: selected expert IDs of shape (num_tokens, top_k).
"""
if scoring_func == "softmax" and not use_grouped_topk and custom_routing_function is None:
topk_weights, topk_ids, _ = torch_npu.npu_moe_gating_top_k_softmax(router_logits, k=top_k)
# 310P returns invalid routing results when this op receives more than 1024 tokens.
if router_logits.shape[0] > 1024:
topk_results = [
torch_npu.npu_moe_gating_top_k_softmax(router_logits_chunk, k=top_k)
for router_logits_chunk in router_logits.split(1024, dim=0)
]
topk_weights = torch.cat([result[0] for result in topk_results], dim=0)
topk_ids = torch.cat([result[1] for result in topk_results], dim=0)
else:
topk_weights, topk_ids, _ = torch_npu.npu_moe_gating_top_k_softmax(router_logits, k=top_k)
Comment thread
Tflowers-0129 marked this conversation as resolved.
topk_weights = _renormalize_topk_weights(topk_weights, renormalize)
else:
topk_weights, topk_ids = _native_select_experts(
Expand Down
Loading