[SP] Fix runtime_max_tokens_per_rank for sequence parallelism (#25685)
Co-authored-by: Ming Yang <minos.future@gmail.com> Co-authored-by: Yinghai Lu <yinghai@meta.com>
This commit is contained in:
co-authored by
Ming Yang
Yinghai Lu
parent
6f892047ec
commit
878e6b8886
@@ -212,11 +212,16 @@ class FlashinferDispatcher(BaseDispatcher):
|
||||
payloads.append(topk_ids)
|
||||
payloads.append(topk_weights)
|
||||
|
||||
self.runtime_max_tokens_per_rank = (
|
||||
max(get_dp_global_num_tokens())
|
||||
if get_dp_global_num_tokens() is not None
|
||||
else x.shape[0]
|
||||
)
|
||||
dp_global = get_dp_global_num_tokens()
|
||||
if dp_global is not None and len(dp_global) > 1:
|
||||
# DP attention: multiple DP ranks with different token counts.
|
||||
# Use the max across ranks so the A2A workspace fits the fattest.
|
||||
self.runtime_max_tokens_per_rank = max(dp_global)
|
||||
else:
|
||||
# dp_size=1 or SP: use the actual input tensor size (post-scatter
|
||||
# in SP mode, full batch otherwise). Avoids the pre-scatter
|
||||
# scheduler count which can exceed the workspace cap.
|
||||
self.runtime_max_tokens_per_rank = x.shape[0]
|
||||
recv_tensors = self.moe_a2a.dispatch(
|
||||
self.dummy_topk_ids_current_rank if self.has_dummy_token else topk_ids,
|
||||
payloads,
|
||||
|
||||
Reference in New Issue
Block a user