Swizzle 辅助接口
gemm_swizzle2d_Nz
将 GEMM 迭代序号转换为输出矩阵的二维 tile 坐标。
tile 按列分组,相邻组反向遍历行,形成蛇形调度。最后一组不足
swizzle_offset 列时会自动处理。
签名
gemm_swizzle2d_Nz(
iter_id,
data_row_shape,
data_col_shape,
tile_row_shape,
tile_col_shape,
swizzle_offset=7,
)
参数
参数 |
含义 |
|---|---|
|
当前迭代序号。范围为 |
|
输出矩阵的行数,通常为 |
|
输出矩阵的列数,通常为 |
|
tile 行数,通常为 |
|
tile 列数,通常为 |
|
每组包含的列 tile 数量,默认为 |
返回值
返回 (data_row_idx, data_col_idx),即 tile 的行索引和列索引。
调度顺序
例如,3 行、10 列的 tile 网格在 swizzle_offset=7 时:前 7 列按
0 -> 1 -> 2 遍历行,后 3 列按 2 -> 1 -> 0 遍历。
使用场景
用于分布式融合 kernel 的 GEMM 阶段。
将相邻列的 tile 分组,有助于提高矩阵数据的缓存局部性。
蛇形遍历让相邻分组在同一行衔接,减少 tile 坐标的大跨度跳转。
返回的
(block_id_m, block_id_n)可直接用于计算 A、B、C 的分块地址。
注意事项
shape、tile 和
swizzle_offset参数必须大于0。接口只调整调度顺序,不改变数据布局。
示例
import triton
import triton.language as tl
from triton_dist.language.extra.ascend.algorithm import gemm_swizzle2d_Nz
@triton.jit
def gemm_tile_kernel(c_ptr, M, N, stride_cm, stride_cn,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr):
pid = tl.program_id(0)
ncore = tl.num_programs(0)
num_tiles = tl.cdiv(M, BLOCK_M) * tl.cdiv(N, BLOCK_N)
for iter_id in range(pid, num_tiles, ncore):
block_m, block_n = gemm_swizzle2d_Nz(
iter_id, M, N, BLOCK_M, BLOCK_N,
)
# 使用 offs_m 和 offs_n 计算 A、B、C 的分块地址。
offs_m = block_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = block_n * BLOCK_N + tl.arange(0, BLOCK_N)
dist_swizzle2d_Nz
将通信迭代序号转换为数据 tile、目标 rank 和尾块大小。
不同 tile 会从不同 rank 开始通信,避免多个 AI Core 集中访问同一 rank。
签名
dist_swizzle2d_Nz(
iter_id,
rank_size,
data_row_shape,
data_col_shape,
tile_row_shape,
tile_col_shape,
comm_npu_split=1,
)
参数
参数 |
含义 |
|---|---|
|
当前迭代序号。范围为 |
|
参与通信的 rank 数量。 |
|
数据行数。 |
|
数据列数。 |
|
每个 tile 的最大行数。 |
|
每个 tile 的最大列数。 |
|
每组交错调度的 rank 数量,默认为 |
返回值
返回 (data_row_idx, data_col_idx, rank_idx, comm_row_size, comm_col_size):
data_row_idx、data_col_idx:tile 的行、列索引。rank_idx:目标 rank。comm_row_size、comm_col_size:当前 tile 的实际行数和列数。
使用场景
用于 AllGather-GEMM、GEMM-ReduceScatter 和 GEMM-AllReduce 等融合 kernel。
轮转不同 tile 的目标 rank,有助于减少通信热点和 rank 间的访问竞争。
注意事项
comm_npu_split必须满足rank_size % comm_npu_split == 0。接口只调整调度顺序,不改变数据布局。
示例
import triton
import triton.language as tl
from triton_dist.language.extra.ascend.algorithm import dist_swizzle2d_Nz
@triton.jit
def communication_kernel(rank_size, rows, cols,
TILE_M: tl.constexpr, TILE_N: tl.constexpr):
pid = tl.program_id(0)
ncore = tl.num_programs(0)
num_iters = rank_size * tl.cdiv(rows, TILE_M) * tl.cdiv(cols, TILE_N)
for iter_id in range(pid, num_iters, ncore):
# target_rank 用于选择远端 rank;tile_rows 和 tile_cols 用于生成尾块 mask。
tile_m, tile_n, target_rank, tile_rows, tile_cols = dist_swizzle2d_Nz(
iter_id, rank_size, rows, cols, TILE_M, TILE_N,
)