单算子开发
本章节说明如何在分布式 kernel 内部使用普通 Triton 算子——tl.load / tl.store / tl.dot 等标准算子如何与 dl.* 分布式原语组合,以及这种组合有哪些需要注意的地方。
一、核心原则:远端指针就是普通指针
分布式 kernel 与普通 Triton kernel 之间只隔着一层:dl.symm_at() 返回的远端指针,在语法和语义上都是一个普通的 Triton 指针。
remote_ptr = dl.symm_at(peer_mem_ptr, target_rank) # 拿到远端地址
remote_ptrs = remote_ptr + offs_m[:, None] * stride_m + offs_n[None, :] * stride_n
tl.store(remote_ptrs, data, mask=mask) # 和写本地内存完全一样
所有普通 Triton 算子——指针算术、tl.load、tl.store、tl.atomic_add、tl.dot——对远端指针的用法与对本地指针完全一致。不需要特殊的"远程 load"算子,也不需要显式的数据搬运调用。
这是整个编程模型的基础:跨卡访问被表达为地址空间的扩展,而不是一组新的通信 API。
1.1 三条使用规则
规则 1:symm_at 只能作用在标量基址上
# 正确:先 symm_at 拿基址,再加偏移向量
remote_ptr = dl.symm_at(peer_mem_ptr, target_rank)
remote_ptrs = remote_ptr + offs[:, None] * stride
# 错误:symm_at 不接受 block pointer
remote_ptrs = dl.symm_at(peer_mem_ptr + offs[:, None] * stride, target_rank)
底层 distributed.symm_at 操作在 IR 层就断言了输入必须是标量指针。
规则 2:读本地对称内存有两种等价写法
local_ptr = dl.symm_at(peer_mem_ptr, rank) # 显式写法
local_ptr = peer_mem_ptr # 等价,推荐
symm_at(p, my_rank) 的计算结果就是 p 本身(本地基址减本地基址等于 0 偏移)。直接用裸指针更简洁,也少一条指令。
规则 3:前面有 dl.wait 时必须经过 dl.consume_token
token = dl.wait(sig_ptr, 1, "npu", "acquire")
data_ptr = dl.consume_token(data_ptr, token) # 不能省略
data = tl.load(data_ptr + offs, mask=mask)
详见第三节。
二、普通算子在分布式 kernel 中的用法
2.1 tl.dot
完全是标准写法,没有任何分布式相关的改动:
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for block_id_k in range(0, num_k_blocks):
a = tl.load(a_ptrs, mask=(offs_k[None, :] < K) & matmul_msk_am, other=0.0)
b = tl.load(b_ptrs, mask=(offs_k[:, None] < K) & msk_n, other=0.0)
accumulator += tl.dot(a, b)
c = accumulator.to(dtype)
唯一的分布式差异在于 a_ptrs 的基址是什么:
融合模式 |
|
说明 |
|---|---|---|
AllGather + GEMM |
|
数据已被所有 peer 推到本地对称内存 |
GEMM + ReduceScatter |
|
就是本地输入张量 |
累加器一律用 fp32,最后 .to(dtype) 转回。这与单卡 GEMM 的惯例相同——Cube Engine 的 MMA 累加类型就是 FP32。
重要影响:kernel 中是否含有
tl.dot,libshmem_device.barrier_all会隐式决定 launch 模式(cv0:1 还是 cv1:2)。
2.2 tl.atomic_add
归约类通信的核心算子。用于把从远端拉回的 partial 结果累加到本地输出:
c_mask = (offs_cm[:, None] < m_per_rank) & (offs_cn[None, :] < N)
c_temp = tl.load(remote_ptrs) # 远端读
tl.atomic_add(c_ptr + c_offs, c_temp, mask=c_mask) # 本地原子累加
atomic_add 的目标应当是本地内存,源才是远端。反过来(每张卡都往同一个远端地址原子加)会造成严重的写端竞争。
2.3 tl.load / tl.store
对远端指针的读写与本地完全一致:
写对称内存的 mask 和读源数据的 mask 通常不同。
因为源数据在输入张量的全局坐标系中,目标在对称内存 buffer 的坐标系中,两个坐标系的边界条件不同:
# 读源数据:用 A 的全局坐标系判断边界
a = tl.load(a_ptrs, mask=(comm_offs_k[None, :] < K) & comm_msk_m, other=0.0)
# 写对称内存:用 buffer 坐标系判断边界
tl.store(remote_ptrs, a, mask=(comm_offs_k[None, :] < K) & peermem_comm_msk_m)
什么时候可以省略 mask:当数据按 BLOCK 对齐写入 buffer、且 buffer 本身也按 BLOCK 分配时,越界不可能发生:
tl.store(peer_mem_ptrs, c) # GEMM 结果写 buffer,无需 mask
...
c_temp = tl.load(remote_ptrs) # 从 buffer 读,无需 mask
此时真实的输出边界靠最后写 C 时的 c_mask 保证。
三、dl.wait 与数据访问的衔接
3.1 为什么需要 consume_token
dl.wait 返回一个 token,必须通过 dl.consume_token(ptr, token) 把它"注入"到后续要访问的数据指针上:
token = dl.wait(sig_ptr, 1, "npu", "acquire")
data_ptr = dl.consume_token(data_ptr, token)
data = tl.load(data_ptr + offs, mask=mask)
漏掉 consume_token 的后果不是内存乱序,而是整条 dl.wait 被删除。
原因在 IR 层面:distributed.wait 的内存效果只有 Read,如果它的返回值无人使用,这条操作就是死代码,会被优化 pass 直接消除掉——同步就此彻底消失。
consume_token 在运行时不做任何事(实现就是把输入原样返回,零开销),它存在的唯一目的就是在 IR 上建立 wait 与数据访问之间的依赖边。
四、Ascend 专属语言扩展
分布式 kernel 中会用到几个 Ascend 平台的语言扩展,它们来自上游 Triton 的 Ascend 语言扩展模块:
from triton.language.extra.cann.extension import sub_vec_id, sub_vec_num
API |
返回 |
说明 |
|---|---|---|
|
本 AI Core 内的 Vector Core 序号 |
运行时值 |
|
每个 AI Core 的 Vector Core 数 |
constexpr,等于 |
Host 侧查询核数:
from triton.backends.ascend.driver import NPUUtils
aicore_num = NPUUtils().get_aicore_num()
aivec_num = NPUUtils().get_aivector_core_num()
不要硬编码核数——不同产品型号与变体的核数不同,一律运行时查询。
4.1 典型用法:门控通信段
subblock_idx = sub_vec_id()
if subblock_idx == 0:
# 通信段:只在一个 Vector Core 上发起搬运
remote_ptr = dl.symm_at(peer_mem_ptr, target_rank)
tl.store(remote_ptr + offs, data, mask=mask)
为什么只用一半的 Vector Core 做通信:一个 AI Core 有两个 Vector Core,其中一半参与通信就足以打满通信带宽,另一个留给计算更有价值。这不是浪费。
五、任务划分惯例
5.1 grid-stride 循环
分布式 kernel 中的所有 tile 循环一律采用 grid-stride 形式:
ncore = tl.num_programs(axis=0)
pid = tl.program_id(axis=0)
for task_idx in range(pid, total_tasks, ncore):
...
好处是任务数与核数解耦,尾块自动处理,不需要额外的边界判断。
5.2 多维任务号的线性化
当任务空间是多维的(如 rank × head × d-block),先线性化再解包:
for task_idx in range(pid, total_prod_tasks, ncore):
tmp = task_idx
rank_loop_id = tmp % rank_size
tmp //= rank_size
h_id = tmp % H
block_id_d = tmp // H
target_rank = (rank + rank_loop_id) % rank_size
注意最后一行:target_rank = (rank + rank_loop_id) % rank_size 而不是直接用 rank_loop_id。如果所有卡都按 0, 1, 2, ... 的顺序访问目标,任意时刻所有卡都在打同一个 peer,造成 incast 热点。加上自己的 rank 做旋转后,各卡在同一时刻打向不同的 peer,链路并发度最大化。
5.3 任务空间要先展平
一条从实践中总结的规则:不要只沿某一个维度分核。
# 错误:只沿 rank_size 维分核:world_size=2 时只有 2 个核在干活
for rank_loop_id in range(logical_core_id, rank_size, n_role_cores):
...
# 正确:展平后分核:所有核都有活干
total_work = num_blocks_s * total_prod_tasks
for wid in range(logical_core_id, total_work, n_role_cores):
global_id_s = wid // total_prod_tasks
task_idx = wid % total_prod_tasks
六、边界处理:用 mask 而非 padding
分布式 kernel 处理非对齐形状时一律用动态 mask,不做 padding:
二维边界 mask(最常用)
offs_s = global_id_s * COMM_BLOCK_S + tl.arange(0, COMM_BLOCK_S)
offs_d = block_id_d * COMM_BLOCK_D + tl.arange(0, COMM_BLOCK_D)
mask = (offs_s < S)[:, None] & (offs_d < D)[None, :]
行 mask 广播(当列维恒等于 BLOCK 时,省掉列判断)
offs_s = block_id * BLOCK_S + tl.arange(0, BLOCK_S)
row_mask = offs_s < sequence_length
io_mask = row_mask[:, None]