gemm和alltoall通算融合

1. 总体思想 这次做的是单机八卡 H200、NVLink、Ulysses CP 下的 GEMM 和 AllToAll 融合。先把 forward 写清楚: A2A → QKV projection → QK → PV → A2A 需要接起来的主边界有两个。输入侧是 A2A→QKV projection,通信先把各个 peer 的输入 tile 搬到本地最终布局,GEMM 拿到一块就算一块。输出侧是 batched PV→A2A,每个本地 head 都有一组独立的 P×V,GEMM 算完一个 tile,通信 CTA 立刻把它送到目标 rank 的最终 Ulysses 布局。 如果 GEMM 和 NCCL 顺序执行,端到端时间接近两段时间相加。这里把通信 CTA 和 GEMM CTA 放进同一个 cooperative persistent grid,两种 CTA 常驻在不同的 SM 上,用 tile 级 ready epoch 接力。A2A→GEMM 由通信生产、GEMM 消费;GEMM→A2A 交换生产消费关系。这样首批 tile 到达后就能启动计算,前面的 tile 也可以在后续 GEMM 还在跑时发出去。 ...

August 17, 2026 · 3 min