LLM System: 通信计算融合算子实现 01 - sm80GEMM + ReduceScatter

flux实现:sm80 gemm+rs 看flux的论文来说,原理并不难,修改的是ffn的第二个gemm的epi阶段,让epi阶段的写回操作变成写入reducescatter的通信buffer,相当于做了一次零拷贝。 几个关键的点:1. 如何用cutlass自定义epi阶段。2. 用了什么ptx。3. 清零? cutlass自定义epilogue–sm80 常规的epi,写回gmem,这里改成写在scatter-aware-memory(我自己起的名字)。cutlass封装自定义epi的逻辑抽象是:输出tile存到哪里,以及做什么运算,这分别是两个模版类: using EVT_D = decltype(this->evt_d(kparams)); using StoreD = decltype(this->evt_store_d(kparams)); using EVT = cutlass::epilogue::threadblock::Sm80EVT<StoreD, EVT_D>; StoreD就是规定了怎么存,EVTD就是规定了怎么算。 EVT全称epilogue visitor tree就是定义了一个epi阶段的计算-访存操作流,用树状图的形式保存epi阶段的一堆操作(做成图应该是为了方便接入nvcc?不懂为啥非得在这强调是tree)一个树节点(操作)是一个visitor。 custom_evt_d() EVT_Compute0 = alpha * accumulator 代码里对应 VisitorCompute<cutlass::multiplies, …>,两个输入是 VisitorScalarBroadcast 和 VisitorAccFetch。第一个负责广播 alpha,第二个v从 GEMM accumulator 里拿当前 tile 的结果,然后 multiply EVT_Compute1 = beta * C + EVT_Compute0 = beta * C + alpha * accumulator 代码里对应 VisitorAuxLoadGemmk 先把 C/bias 读进 epilogue,然后 VisitorCompute<cutlass::multiply_add, …> 做 beta * C + alpha * acc。因为CUTLASS 2.x 没有 SM90 那种 SrcFetch 替代物,所以这里用 VisitorAuxLoadGemmk 来读 C。 ...

July 5, 2026 · 4 min

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