# 单算子开发 本章节说明如何在分布式 kernel 内部使用普通 Triton 算子——`tl.load` / `tl.store` / `tl.dot` 等标准算子如何与 `dl.*` 分布式原语组合,以及这种组合有哪些需要注意的地方。 --- ## 一、核心原则:远端指针就是普通指针 分布式 kernel 与普通 Triton kernel 之间只隔着一层:**`dl.symm_at()` 返回的远端指针,在语法和语义上都是一个普通的 Triton 指针。** ```python 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` 只能作用在标量基址上** ```python # 正确:先 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:读本地对称内存有两种等价写法** ```python 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`** ```python 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` 完全是标准写法,没有任何分布式相关的改动: ```python 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` 的基址是什么**: | 融合模式 | `a_ptrs` 基址 | 说明 | | --- | --- | --- | | AllGather + GEMM | `peer_mem_ptr` | 数据已被所有 peer 推到本地对称内存 | | GEMM + ReduceScatter | `a_ptr` | 就是本地输入张量 | 累加器一律用 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 结果累加到本地输出: ```python 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 的坐标系中,两个坐标系的边界条件不同: ```python # 读源数据:用 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 分配时,越界不可能发生: ```python 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)` 把它"注入"到后续要访问的数据指针上: ```python 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 语言扩展模块: ```python from triton.language.extra.cann.extension import sub_vec_id, sub_vec_num ``` | API | 返回 | 说明 | | --- | --- | --- | | `sub_vec_id()` | 本 AI Core 内的 Vector Core 序号 | 运行时值 | | `sub_vec_num()` | 每个 AI Core 的 Vector Core 数 | constexpr,等于 `get_aivector_core_num() // get_aicore_num()` | Host 侧查询核数: ```python from triton.backends.ascend.driver import NPUUtils aicore_num = NPUUtils().get_aicore_num() aivec_num = NPUUtils().get_aivector_core_num() ``` **不要硬编码核数**——不同产品型号与变体的核数不同,一律运行时查询。 ### 4.1 典型用法:门控通信段 ```python 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 形式: ```python 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),先线性化再解包: ```python 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 任务空间要先展平 一条从实践中总结的规则:**不要只沿某一个维度分核。** ```python # 错误:只沿 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(最常用)** ```python 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 时,省掉列判断) ```python offs_s = block_id * BLOCK_S + tl.arange(0, BLOCK_S) row_mask = offs_s < sequence_length io_mask = row_mask[:, None] ```