flashinfer.comm.moe_a2a_get_workspace_size_per_rank

flashinfer.comm.moe_a2a_get_workspace_size_per_rank(ep_size: int, max_num_tokens: int, total_dispatch_payload_size_per_token: int, combine_payload_size_per_token: int)

获取 MoeAlltoAll 操作每个 rank 的工作区大小。

参数:
  • ep_size – 总专家并行大小

  • max_num_tokens – 所有 rank 上的最大 token 数量

  • total_dispatch_payload_size_per_token – 分发阶段每个 token 的 payload 大小。这应该是所有 payload 的总和。

  • combine_payload_size_per_token – 组合阶段每个 token 的 payload 大小。

返回值:

每个 rank 的工作区大小,单位为字节

返回值类型:

workspace_size_per_rank