跳转至

LowerHostTensorCollectives Pass

概览

LowerHostTensorCollectives 将 host orchestrator 中的 pld.tensor.allreducepld.tensor.barrierpld.tensor.broadcastpld.tensor.reduce_scatterpld.tensor.allgatherpld.tensor.all_to_allpld.tensor.all_to_all_v 调用改写为编译器内部的 builtin chip dispatch。它在 MaterializeCommDomainScopes 之后运行, 因此 window 绑定的 data tensor 和用户显式传入或编译器合成的 signal tensor 已经带有 WindowBuffer 反向引用,并属于推断出的通信域。

该 pass 不修改非 host 函数。InCore allreduce 仍然走 LowerCompositeOps

Pipeline 位置

... -> SynthesizeAllReduceSignals -> MaterializeCommDomainScopes -> LowerHostTensorCollectives -> MaterializeDistTensorCtx -> Simplify(最终) -> MaterializeRuntimeScopes

最终的 Simplify 位于本 pass 之后,用于继续折叠生成的循环边界或常量表达式, 随后再插入 runtime scopes。

行为

对于 host orchestrator 中的调用:

data = pld.tensor.allreduce(data, signal, op=pld.ReduceOp.Sum)
data = pld.tensor.allreduce(data, signal, op=pld.ReduceOp.Sum, core_num=4)
data = pld.tensor.allreduce(data, signal, op=pld.ReduceOp.Sum, mode="ring")
signal = pld.tensor.barrier(signal)
data = pld.tensor.broadcast(data, signal, root=0)
data = pld.tensor.reduce_scatter(data, signal, op=pld.ReduceOp.Sum)
data = pld.tensor.allgather(stage, data, signal)
data = pld.tensor.all_to_all(stage, data, signal)
data = pld.tensor.all_to_all_v(input, target, signal, send_counts, recv_counts)

pld.tensor.allreduce 根据 mode kwarg 进行分发:默认 mode="mesh" 会 lower 到 builtin.tensor.allreduce,而 mode="ring" 会 lower 到 builtin.tensor.allreduce_ring。其他取值将作为用户错误被拒绝。

对于 allgather / all_to_all / all_to_all_vstage/input(TPUT 源) 与 data/target(结果)必须是两个不同的 window。allgatherstage 只保存本 rank 的单个分片,形状为 [1, SIZE]all_to_allstage 每行 携带一个按目的地划分的分片,形状为 [NR, SIZE]all_to_all_vinput 每个目的地携带一个 MAX_RECV 行容量块,形状为 [NR*MAX_RECV, SIZE]all_to_all / all_to_all_v 两种情况下 data/target 都是 peer 推入的 结果窗口。all_to_all_v 还额外要求 send_counts(在这一层是窗口绑定的, 仅本地使用)和 recv_counts(窗口绑定,通过 pld.system.notify 跨 rank 发布)——五个窗口参数都必须位于同一个 CommDomainScopeStmt 中,并且必须 两两互不相同(任意一对发生别名都是跨进程竞争,无论是 data 与 data、 data 与 control,还是 control 与 control 之间)。

本 pass 会为每个参与设备生成对应的 builtin.tensor.* 调用(如 builtin.tensor.allreducebuiltin.tensor.allreduce_ringbuiltin.tensor.barrierbuiltin.tensor.broadcastbuiltin.tensor.reduce_scatterbuiltin.tensor.allgatherbuiltin.tensor.all_to_allbuiltin.tensor.all_to_all_v)。若外层 comm-domain scope 带有显式 device 列表,则生成 SeqStmts;否则生成顺序 for r in pld.system.world_size() 循环。

每个生成的 builtin call 携带来源 pld.tensor.* 调用中 collective 特定的 参数和 kwarg 属性。窗口绑定的 INOUT tensor 原样传递;标量 kwarg 值 (oprootdtype,以及 mesh AllReduce 的 core_num)转发给 builtin。 all_to_all_vMAX_RECV 不是 lowering 时的属性:HOST 内核在入口把它推导为 target.shape[0] / nranks(运行时通信域大小),因此不再需要按 MAX_RECV 进行代码生成的 variant 混入,块布局也始终与实际运行的设备数一致。

若用户代码使用赋值形式,pass 会在生成的 builtin 调用之后追加 <result> = <original expr>,保留 public API 的 rebind 语义。

打印形式

builtin.tensor.* 算子在注册表中标记为 internal_only(内部专用):没有任何 DSL 包装器可以拼写它们,面向用户的算子构造路径也会按名字拒绝它们。但 Python printer 仍然需要打印它们,因此把它们放在 pl.builtin 命名空间下 —— 与它给任何 非 pld 注册算子加 pl. 前缀的规则一致:

for r_1 in pl.range(pl.const(0, pl.INT64), pld.system.world_size(), pl.const(1, pl.INT64)):
    pl.builtin.tensor.allreduce(
        data, signal, op=0, dtype=pl.FP32, core_num=1,
        attrs={"op": 0, "dtype": pl.FP32, "core_num": 1, "device": r_1,
               "arg_directions": [pl.adir.inout, pl.adir.inout]},
    )

Parser 能读回这种拼写(ast_parser._parse_builtin_op),因此 lowering 产生的 dispatch 可以完成 print -> parse 往返。它是仅供机器使用(machine-only)的表面, 限定在真正注册于 builtin. 下的名字,并且只接受 printer 能写出的形式: devicearg_directions 两个 attr 是必需的,因为 orchestration codegen 会在 内部检查(internal check)后读回它们。手写调用若省略这些 attr,会在 parse 阶段 按用户错误拒绝,而不是在 codegen 阶段表现为编译器 bug 诊断。用户代码请改写复合 形式 pld.tensor.*

注意:整程序的 assert_structural_equal 往返仍然被上一个 pass 阻断 —— MaterializeCommDomainScopes 会合成 CommDomainScopeStmt(打印为前导注释)以及 DistributedTensorType 上的 WindowBuffer 反向引用(完全不打印),二者都没有可解析回来的 DSL 表面。

检查

该 pass 要求两个参数都是已经 materialize 的 DistributedTensorType view,并且位于同一个 CommDomainScopeStmt 中。host allreduce builtin 支持 FP16、FP32 上的 ReduceOp.SumMaxMinProd,并支持任意正元素数量。它按 256 个 元素分块,并把 FP16 和 FP32 的 ragged load 范围都对齐到 32 字节,不改变逻辑 tensor shape。 signal 必须是 INT32 tensor,形状可以是 rank-1 [world_size] 或 rank-2 [world_size, signal_stride];当参与设备数静态可知时,signal 的静态容量必须足够。 由于 signal 由 pld.window 产生,它天然是 packed 的,builtin 按扁平 row-major 数组索引它。

mesh allreduce 为每个启动的 AIV block 分配一条 signal lane:rank-1 signal 仅在 core_num == 1 时有效;rank-2 signal 要求第二维是常量且 signal_stride >= core_num(允许更宽的 stride,因此显式 signal 可以带有多余 lane)。core_num 还必须不超过所配置 backend 的 AIV 核数——该 builtin 以 standalone AIV kernel 提交并设置了 require_sync_start,超额的 launch 永远无法 被准入,表现为挂死而非报错。未配置 backend 时(纯 IR 测试)跳过该检查。 多核仅支持 mesh:mode="ring" 要求 core_num == 1

Ring allreduce(mode="ring") 的 signal 为 rank-2,形状为 [2 * (NR - 1) + 1, NR],其 shape[0] 在 signal 两个维度均为编译期常量时必须恰好等于 2 * (NR - 1) + 1;仅 shape[0] 静态可知时则至少为 2 * (NR - 1) + 1(两个维度均为动态时无静态检查)。 当参与设备数静态可知时,signal 的静态容量必须足够。ring allreduce 会在运行时把每个 rank 的 src 划分为 NR 个均衡、可能不齐(ragged)的 chunk——chunk r 覆盖 [floor(numel*r/NR), floor(numel*(r+1)/NR))——因此 numel(src) 不必能被 NR 整除,host-ring 的 src 形状也不必静态已知。每次 TPUT 传输都通过 staging tile 的 valid shape 精确收窄到对应 chunk 的范围,因此不齐(ragged)与动态输入均得到完整支持。

Ring allreduce 目前仅支持 ReduceOp.Sumdtype=FP32ReduceOp.MaxReduceOp.MinReduceOp.Prod 以及 FP16mode="ring" 下尚未支持。Ring allreduce 最多支持 16 个参与设备 (world_size <= 16)。

builtin.tensor.allreduce_ring 内核采用推送(push)模型:数据搬运使用 pto::comm::TPUT(远端写)——reduce-scatter 阶段通过 TPUT<AtomicAdd> 将部分和 累加到右邻居的 slot,allgather 阶段用非原子 TPUT 转发每个已归约的 chunk, 与树内 allgather / all_to_all host builtin 保持一致。顺序保证为每次传输前后 pipe_barrier(PIPE_ALL),并在每次 TNOTIFY 前加 dsb(DSB_DDR)(而非 pto.fence.barrier_all,后者不会排空 MTE DMA 流水线)。跨 rank 同步使用 O(1) 的 NeighborBarrier(只通知/等待左右两个 ring 邻居)——在 NPU 上安全是因为 TPUT 写流水线保证数据先于信号可见,而旧的拉取(pull)模型(TLOAD/TSTORE)不具备该保证。

ring 内核是自清零(self-clearing)的:最终屏障之后,尾声(epilogue)把每个已使用的 barrier 行恢复为 0(对每个带信用的单元做本地 TNOTIFY(-1, AtomicAdd)—— NeighborBarrier 每轮是两个邻居单元,RoundBarrier 回退路径是所有 P−1 个单元), 因此单个 signal buffer 可以像其它 host builtin(#2279)一样在连续调用间复用。 nranks == 2 时两个邻居收敛到同一个 peer,该唯一信用单元携带两个 +1,用两个 −1 恢复。

HOST collective 的所有 window 操作数——data 与 signal 都是如此——必须解析为两两不同的 WindowBuffer 分配。同一个 alloc_window_buffer 上的两个 pld.window() view 在 in-kernel TPUT/notify 下是跨进程数据竞争:data 对 data 是 reduce 覆盖, data 对 control 是 notify/count 写入与内核读取竞争,control 对 control 是 notify 与 count 发布竞争。LowerHostTensorCollectives 在生成 builtin dispatch 之前会拒绝任何别名对。

当参与设备数静态可知时,还会额外校验 broadcastroot kwarg:在显式静态设备子集上, 它必须满足 root < participating device count。完全动态的 "all device" 域在编译期无法 校验(那里没有可用的设备数)——与 signal 容量检查存在相同的已记录限制。

signal 可复用(对自清理的 host builtin 而言):这些 kernel 会在每次调用后 自清理屏障 cell(信用屏障尾声),因此一个合成或用户分配的 signal 可以支撑 任意数量的连续调用或循环迭代,无需重新分配。

all_to_all_v 的单次使用 Set(1)/wait≥1 信号无法在 host_orchfor/while 循环中复用——本 pass 之前紧邻运行的 MaterializeCommDomainScopes 会提前 拒绝这种情况(与 LowerCompositeOps 在 InCore 路径上强制的限制相同)。在显式 静态 device 子集上,all_to_all_v 的 signal shape[0] 必须与子集大小 精确相等(而非其他 collective 所要求的 >=),因为 MAX_RECV 是由 target.shape[0] / signal.shape[0] 推导得出的,signal 过度分配会导致 静默的错误降级。

Pass 属性

字段 取值
required {IRProperty::CommDomainScopesMaterialized}
produced {IRProperty::CommDomainScopesMaterialized}
invalidated {}

参考