作者:WS、ZCX、YS、PZL from DeepLink Group @ Shanghai AI Lab

概述

Grouped MatMul 的难点不仅在矩阵乘本身,还在于如何根据 Device Tensor 中的分组信息组织不规则任务。本文以 M 轴分组为例,分析预计算映射在 Ascend 910B 上引入约 1 ms 额外开销的原因,并通过 24 个常驻 program 将不同 Group 的输出 tile 组成连续任务流,减少 Host 同步、Kernel Launch 和组尾 Core 空闲;随后使用对角线 Swizzle 改善访存竞争与 L2 局部性。主测试集上,该实现整体接近 Torch 融合算子,最佳测试点约快 3.4%。

Triton 系列往期课程:

Triton第一课:不写 CUDA,也能实现高性能 GPU Kernel

Triton第二课:实现 LayerNorm 算子的六种方式

Triton第三课:Grouped GEMM 算子优化思路与实战技巧

Triton第四课:基于对称内存融合 AllGather 与 MatMul,提速1.56倍

Grouped MatMul 的计算语义

问题定义

先从普通 MatMul 开始。普通 MatMul 只有一套左右矩阵:

[M, K] @ [K, N] -> [M, N]

M 轴分组的 Grouped MatMul 则把左矩阵的 M 行切成 G 个连续区域,每个区域使用不同的权重。它的输入为:

  • 激活矩阵 A,形状为 [M, K]

  • 权重矩阵 B,形状为 [G, K, N]

  • 分组大小 group_size,形状为 [G],且位于 Device;

  • 输出矩阵 C,形状为 [M, N]

举个简单的例子:

group_size = [128, 256, 64]
G = 3
M = 128 + 256 + 64 = 448

这次计算会被拆成三段:

Group 0:C[  0:128, :] = A[  0:128, :] @ B[0, :, :]
Group 1:C[128:384, :] = A[128:384, :] @ B[1, :, :]
Group 2:C[384:448, :] = A[384:448, :] @ B[2, :, :]

写成通用公式,就是:

start_g = sum(group_size[:g])
end_g   = start_g + group_size[g]
C[start_g:end_g, :] = A[start_g:end_g, :] @ B[g, :, :]

在这里插入图片描述

如图所示,左矩阵 A 沿 M 轴被切成多个 Group,每个 Group 分别与 B 中对应的 [K, N] 权重相乘,结果仍按原来的行序写入 C

这里有一个关键约束:同一个计算块不能跨越 Group 边界。 如果它前半部分属于 Group 1、后半部分属于 Group 2,就需要在一个 tile 内处理两套权重,显著增加 Kernel 的控制流与地址计算复杂度。

Host、Device 与 Device Tensor

在这篇文章中:

  • Host 指 CPU,以及运行 Python 和 PyTorch 调度代码的进程;

  • Device 指执行计算的 Ascend NPU;

  • Device Tensor 指数据存放在 NPU 显存中的 Tensor。

例如:

group_size = torch.tensor([128, 256, 64], device="npu")

Python 知道这个 Tensor 的形状和类型,但三个具体数值保存在 NPU 上。如果 Python 调用 `.item()` 读取其中一个值,CPU 必须先等待 NPU 完成相关任务,再把数值传回 CPU。这个等待过程就是 Device-to-Host 同步

Host 侧逐组计算的开销

最直接的实现是在 Python 中逐组读取大小并调用 MatMul:

start = 0
for g in range(num_groups):
    group_m = group_size[g].item()
    C[start:start + group_m] = A[start:start + group_m] @ B[g]
    start += group_m

语义没有问题,但性能上有两个麻烦:

  • .item() 会让 CPU 等待 NPU,并把 group_m 传回 Host;

  • 每个 Group 单独执行 MatMul,会产生多次 Kernel Launch,小 Group 越多,调度开销越明显。

因此,更合适的实现是让 Host 只启动一次 Kernel,后续工作全部留在 NPU 上:读取 group_size、判断 Group 边界、派发任务并完成矩阵乘。

从输出 tile 理解任务数量

Triton 不会让一个 program 计算完整输出矩阵,而是把输出切成许多大小为 [BLOCK_M, BLOCK_N] 的 tile。一个 program 负责一个 tile。

假设某个 Group 有 300 行,同时:

group_size[g] = 300
N = 1000
BLOCK_M = 128
BLOCK_N = 256

M 方向需要 3 块,N 方向需要 4 块:

num_block_m[g] = ceil(300 / 128) = 3
num_block_n    = ceil(1000 / 256) = 4
tasks[g]       = 3 * 4 = 12

也就是说,这个 Group 会产生 12 个独立任务。最后一个 M tile 只有 44 行有效数据,最后一个 N tile 只有 232 列有效数据,其余位置通过 Mask 屏蔽。

这里没有出现 K,是因为 K 是每个 tile 内部的归约维度。K 会决定一个 tile 要循环累加多少次,但不会改变输出一共有多少个 tile。

预计算分组映射

本节先介绍一种最直观的实现:提前算好 tile 到 Group 的映射表,让主 Kernel 直接查表。但它会在路径上引入约 1ms 的前处理开销,这也是后续第 3 节核心优化中要解决的。

构建 tile 到 Group 的映射

主 Kernel 最想知道的是:“我现在计算的这个 tile 属于哪个 Group?”最直接的办法,是提前准备一张对照表 m_indices_pad。表中第 i 个元素,记录第 i 个 M tile 的 Group ID。

这样做的好处是,主 Kernel 不需要临时扫描 group_size,只需读取:

group_id = tl.load(m_indices_pad + pid_m)

这张映射表的生成过程如下。

对第 g 组,先将 M 维逻辑补齐到 BLOCK_M 的整数倍:

m_per_group_padding = triton.cdiv(size_per_group, BLOCK_M) * BLOCK_M
m_pad = m_per_group_padding.sum()
repeats = (m_per_group_padding // BLOCK_M).to(torch.int32)

首先把每个 Group 的行数在逻辑上补齐到 BLOCK_M 的整数倍。注意,这里的 Padding 是为了分配 tile,并不要求真的创建一份更大的输入张量。第 g 组对应的 tile 数为:

tiles_m[g] = ceil(group_size[g] / BLOCK_M)

随后通过一个类似 repeat_interleave 的小 Kernel,把 Group ID 重复相应次数:

@triton.jit
def repeat_interleave_kernel(
    group_ptr,
    repeats_ptr,
    repeat_cum_ptr,
    output_ptr,
):
    pid = tl.program_id(axis=0)
    repeat = tl.load(repeats_ptr + pid)
    start = tl.load(repeat_cum_ptr + pid) - repeat
    group = tl.load(group_ptr + pid)

    for r in range(repeat):
        tl.store(output_ptr + start + r, group)

配套的边界信息为:

group_end = size_per_group.cumsum(0)
group_start = group_end - size_per_group

m_indices_pad = torch.empty(
    m_pad // BLOCK_M,
    device=size_per_group.device,
    dtype=torch.int64,
)

在这里插入图片描述

上图中不同颜色表示不同 Group。Group 2 即使只有少量有效行,也会独占 pid_m_3,tile 内多出来的位置用 Mask 屏蔽;Group 3 从下一个 tile 重新开始。这样,一个 pid_m 始终只对应一套 Group 边界和一块权重。

预计算后,每个 M 轴 tile 都只属于一个 Group,不会出现同一 tile 横跨两套权重的情况。主 Kernel 可通过 m_indices_pad[pid_m] 直接得到 Group ID,再结合 group_startgroup_end 生成 Mask。

预计算方案的额外开销

预计算让主 Kernel 变得更简单,但端到端路径增加了 Padding 计算、除法、前缀和、内存分配以及映射表写入。在 Ascend 910B(AArch64 Host)测试环境中,这条预处理链路耗时约 1 ms。如果主体 MatMul 本身只有 1~2 ms,这笔固定开销就已经不可忽略。

此外,映射表还会增加一次全局内存写入和主 Kernel 中的一次全局内存读取。它用带宽换控制流,是否划算取决于 Group 数量、单组大小以及主体 GEMM 的计算量。

这说明,判断一次优化是否有效,不能只看主 Kernel 是否变简单,还需要把前处理、内存分配和额外 Kernel Launch 计入端到端耗时

因此,下一步将 Group 查询放回主 Kernel,并通过跨 Group 调度减少 AI Core 空闲。

核心优化:跨 Group 连续派发

连续任务流

前面已经知道,每个 Group 会产生若干 tile 任务。如果每个 Group 都从 Core 0 重新分配,当任务数不是 24 的整数倍时,最后一波总会有部分 Core 空闲。小 Group 越多,这种浪费出现得越频繁。

更高效的方式,是将所有 Group 的 tile 在逻辑上首尾相接。Group 0 的任务发完后,不再等待所有 Core 在 Group 边界重新对齐;空闲 Core 可以直接处理 Group 1 的任务。

实现上只启动 24 个 Triton program。第 c 个 program 负责逻辑编号满足 task_id % 24 == c 的任务,并在 Kernel 内循环处理后续 tile。
在这里插入图片描述

上图直观展示了两种派发方式:上半部分中,Group 0 只有 16 个任务,因此 Core 16~23 处于空闲状态;下半部分中,这 8 个 Core 直接处理 Group 1 的前 8 个任务。矩阵乘的计算内容没有变化,变化的只是任务顺序。

这种写法属于静态 Persistent Kernel 调度:不需要原子计数器,也不需要提前生成 tile -> group 映射表。

调度器实现

Kernel 中有两个与分组直接相关的参数:

  • num_groups 表示一共有多少个 Group;

  • group_size_ptr 指向 NPU 上的 group_size 数组,数组中保存每个 Group 的真实行数。

例如 group_size=[128, 256, 64] 时,num_groups=3。Kernel 通过 tl.load(group_size_ptr + group_idx) 依次读取 128、256 和 64。num_groups 被声明为 tl.constexpr,编译器可以据此优化 Group 循环;group_size 的具体数值仍然在运行时从 Device 读取。

下面给出完整的调度骨架。为突出任务分配逻辑,省略 MatMul 的 Load、Dot 和 Store:

@triton.jit
def m_grouped_gemm_kernel(
    a_ptr,
    b_ptr,
    c_ptr,
    group_size_ptr,
    num_groups: tl.constexpr,
    N: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_THRESHOLD: tl.constexpr,
):
    core_idx = tl.program_id(axis=0)
    total_cores = tl.num_programs(axis=0)
    num_block_n = tl.cdiv(N, BLOCK_N)

    # 下一组应该从哪个 Core 开始接着分配
    last_count = 0
    group_start = 0

    for group_idx in range(num_groups):
        group_m = tl.load(group_size_ptr + group_idx)
        group_end = group_start + group_m
        num_block_m = tl.cdiv(group_m, BLOCK_M)

        # 当前 Group 的逻辑任务区间为 [last_count, cur_count)
        cur_count = last_count + num_block_m * num_block_n
        cur_block = tl.where(
            core_idx >= last_count,
            core_idx,
            core_idx + total_cores,
        )

        while cur_block < cur_count:
            local_task = cur_block - last_count
            task_m_idx, task_n_idx = grouped_launch_npu(
                local_task,
                num_block_m,
                num_block_n,
                BLOCK_THRESHOLD,
            )

            # A 的有效行范围为 [group_start, group_end)
            # B 的第 0 维索引为 group_idx
            # 此处执行 BLOCK_M x BLOCK_N 输出 tile 的 MatMul
            # ...

            cur_block += total_cores

        last_count = cur_count % total_cores
        group_start = group_end

这段循环可以拆成五个步骤理解。

第一步:找到当前 Group 的行范围。

group_m = tl.load(group_size_ptr + group_idx)
group_end = group_start + group_m

如果 group_start=128group_m=256,当前组处理的就是 A[128:384, :]C[128:384, :],并使用 B[group_idx]

第二步:计算当前 Group 有多少个任务。

num_block_m = tl.cdiv(group_m, BLOCK_M)
cur_count = last_count + num_block_m * num_block_n

num_block_m * num_block_n 是当前 Group 的 tile 数。cur_count 不是任务数量,而是当前 Group 在连续任务流中的右边界。

第三步:每个 Core 找到自己的第一个任务。

cur_block = tl.where(
    core_idx >= last_count,
    core_idx,
    core_idx + total_cores,
)

如果新 Group 从位置 16 开始,那么 Core 16~23 直接领取 16~23;Core 0~15 需要先绕过一轮,因此领取 24~39。这样任务就从上一组结束的位置继续,而不是重新从 Core 0 开始。

第四步:把一维任务编号变成二维坐标。

local_task = cur_block - last_count
task_m_idx, task_n_idx = grouped_launch_npu(...)

local_task 是当前 Group 内从 0 开始的编号。grouped_launch_npu 再把它转换成输出矩阵中的 (task_m_idx, task_n_idx)。后面讨论的水平分核和对角线 Swizzle,就发生在这一步。

第五步:当前 Core 继续领取间隔为 24 的任务。

cur_block += total_cores

例如 Core 16 会依次尝试任务 16、40、64……,直到超过当前 Group 的 cur_count。Group 结束后:

last_count = cur_count % total_cores
group_start = group_end

last_count 把没有填满的一波任务交给下一个 Group 继续使用;group_start 则移动到 A 和 C 的下一段真实行。注意,group_start 按真实行数推进,不能按 Padding 后的行数推进。

Host 侧根据设备属性设置 Grid:

import triton.runtime.driver as driver

device = torch.npu.current_device()
num_cores = driver.active.utils.get_device_properties(device)["num_aicore"]

m_grouped_gemm_kernel[(num_cores,)](
    A,
    B,
    C,
    group_size,
    num_groups,
    N,
    BLOCK_M=BLOCK_M,
    BLOCK_N=BLOCK_N,
    BLOCK_THRESHOLD=BLOCK_THRESHOLD,
)

grid = num_cores 表示只启动一波 program,再让每个 program 在 Kernel 内循环处理多个 tile。这里容易混淆:Grid 变大不会增加 Kernel Launch 次数,但会增加 program 数量和调度波次。

用 16 + 25 个 tile 理解“无缝衔接”

现在用图中的数字完整走一遍。假设设备有 24 个 AI Core:

  • Group 0 有 16 个 tile;

  • Group 1 有 25 个 tile。

处理 Group 0 时,Core 0~15 各计算一个 tile,Core 16~23 在该组没有任务。此时:

last_count = 16 % 24 = 16

进入 Group 1 后,逻辑任务区间为 [16, 41)。Core 16~23 直接领取 Group 1 的 0~7 号 tile;已经完成 Group 0 的 Core 0~15 则领取 Group 1 的 8~23 号 tile。由于 Group 1 一共有 25 个 tile,Core 16 还会以 24 为步长继续处理 24 号 tile。

Core 0  ~ 15:Group 0 的 tile 0  ~ 15 → Group 1 的 tile 8  ~ 23
Core 16      :Group 1 的 tile 0、24
Core 17 ~ 23:Group 1 的 tile 1  ~ 7

看起来任务编号发生了一次旋转,其实规则很简单:第 core_idx 个 Core 始终处理满足 global_task_id % 24 == core_idx 的任务。这样既不会遗漏,也不会重复,还能让 Group 0 空出来的 Core 提前开始 Group 1。

正确性边界

跨 Group 调度最容易在边界位置出错,实现时需要检查以下五点:

  • group_start 必须按真实 group_size 累加,不能按 Padding 后的长度推进;

  • 最后一个 M tile 需要使用 row < group_end 的 Mask,避免读到下一组的激活;

  • B 指针必须使用当前 group_idx,不能由 Padding 后的全局 pid_m 直接推导;

  • 空 Group 应产生 0 个 tile,同时保持 group_start 更新正确;

  • group_size.sum() 应等于 M,否则需要在 Host 侧校验,或明确约定多余行的处理方式。

L2 Cache 优化:调整 tile 执行顺序

跨 Group 连续派发改善了 AI Core 利用率,但 tile 的执行顺序仍会影响 L2 Cache 命中率和并发访存冲突。即使 tile 数量完全相同,不同的编号映射也可能产生明显的性能差异。

对单个 Group:

num_tiles = ceil(M_g / BLOCK_M) * ceil(N / BLOCK_N)

调度器拿到的是一维编号 local_task,而输出矩阵是二维的,所以必须把它映射为 (task_m_idx, task_n_idx)如何完成这个映射,就是分核策略。
在这里插入图片描述

左图按行优先顺序编号,连续任务先沿 N 轴走;右图重新排列编号,让任务在一个较小的二维区域内展开。两种方法最终计算完全相同的 C,区别只是各 Core 在什么时间读取哪一块 A 和 B。

水平分核

水平分核采用最直接的行优先顺序:

task_m_idx = local_task // num_block_n
task_n_idx = local_task % num_block_n

在这里插入图片描述

这段代码很好理解,但矩阵较大时容易出现两个问题:

  • 多个 Core 同时读取相同的 A tile,访问压力集中,可能加剧片上存储体冲突或共享访存通路竞争;

  • 完成一整行输出需要遍历大量 B tile。若 B 的工作集超过 L2 容量,下一行再次使用这些数据时容易发生缓存失效和重复搬运。

普通 MatMul 中的 grouped ordering 或 Swizzle,本质上都是通过调整 tile 顺序,提高已进入 L2 的 A、B 数据复用率。Grouped MatMul 还需要满足一个额外限制:重排只能发生在当前 Group 内,不能把 tile 映射到相邻 Group。

8 × 8 对角线 Swizzle

一种更适合当前场景的做法,是把输出 tile 划分为 8 × 8 的小区域,再沿对角线顺序遍历。先用图看编号方式:
在这里插入图片描述

0~7 分布在第一条对角线上,8~15 移到相邻对角线,直到覆盖 64 个 tile。这样,连续 Core 不会全部集中在同一 M 行或 N 列上。

在这里插入图片描述

以 24 个 AI Core 同时执行为例,对角线重排把访问分散到不同的 M、N tile,主要希望获得两项收益:

  • 降低多个 Core 对同一片上 Bank 或访存路径的集中竞争;

  • 在较小的二维窗口内复用 A、B tile,降低工作集超过 L2 容量后反复换入换出的概率。

需要强调的是,Swizzle 不会减少计算量,只会调整计算顺序。8 × 8 也不是固定最优值,需要与 BLOCK_MBLOCK_N 一起调优。对于小矩阵或狭长矩阵,额外索引计算可能抵消缓存收益,因此代码用 BLOCK_THRESHOLD 在简单映射和对角线映射之间切换。

面向 Ascend 910B 的实现细节

AI Core 与片上存储层级

在本文测试配置中,Ascend 910B 有 24 个 AI Core。可以先做一个简化理解:Cube 负责矩阵乘累加,Vector 负责逐元素和归约,Scalar 负责循环、判断和地址计算;数据则要经过 UB、L1、L0A、L0B、L0C 等片上存储层级。

Grouped MatMul 的主体计算运行在 Cube 上,Group 扫描、边界判断和地址计算则会占用 Scalar/Vector 资源。控制逻辑过重时,Cube 可能等待数据或地址计算。因此,除了理论 FLOPs,还需要关注 Cube 是否能够持续获得输入数据。

Grid 与 AI Core 利用率

本文从 grid = num_aicore 开始调优,也就是启动 24 个 program,让每个 program 在 Kernel 内循环处理多个 tile。这样做有三个直接目的:

  • 避免逐 Group 启动 Kernel;

  • 避免 Grid 远大于 Core 数量时形成过多调度波次;

  • 通过跨 Group 连续派发,填补单组尾部不足一个完整波次的气泡。

“Grid 越小越好”并不是普适结论。 如果单个 program 循环过长,或者不同 tile 的耗时差异较大,一波 program 也可能出现负载不均。实践中可以从 1× num_aicore 开始,再测试 2× num_aicore

UB 空间与软件流水

910B 单核 UB 的物理容量为 192 KB,但编译器临时变量和 Double Buffer 也会占用空间,因此不能将全部容量用于 A、B、C tile。

调参时需要联合考虑:

BLOCK_M × BLOCK_K:A tile
BLOCK_K × BLOCK_N:B tile
BLOCK_M × BLOCK_N:累加或输出 tile

Block 太小时,Cube 利用率与连续搬运效率可能不足;Block 太大时,片上空间压力增大,还可能降低并发度。因此需要联合搜索 BLOCK_M/N/K,而不是单独增大某一个维度。

片上空间允许时,可以通过 num_stages 开启软件流水。目标是在 Cube 计算当前 K tile 时,提前加载下一个 K tile:

for k in range(0, K, BLOCK_K):
    load A tile, B tile
    compute dot
store C tile
时间片 迭代 0 迭代 1 迭代 2 迭代 3
0 Load
1 Compute Load
2 Store/Next Compute Load
3 Store/Next Compute Load
4 Store/Next Compute
5 Store

这里的表格只是帮助理解流水。实际 GEMM 通常在 K 循环中重叠“下一块 A/B 的加载”和“当前块的 Dot”,C 会在累加结束后统一写回。

提高访存连续性

同样读取 64 KB 数据,连续大块搬运通常比多次零散搬运更高效。访存优化可以按以下顺序进行:

  • 优先使用块指针。 明确提供张量的 `shape`、`stride`、`offset` 和 `order`,便于编译器识别二维块布局并选择更合适的搬运指令;

  • 增大连续维度的 Block。 在资源允许的前提下,优先扩大物理连续维度的搬运长度,提高有效带宽;

  • 补充地址属性提示。 在条件确实成立时,使用 `multiple_of`、`max_contiguous` 和 `max_constancy` 帮助编译器证明对齐与连续性;

  • 大段读取、片上切片。 对多次使用的小块非连续数据,可先搬运较大的连续区域,再用编译器扩展的 `slice` 原语在片上拆分;

  • 片上组合、连续写回。 将零散结果在片上组织后再一次性写入 Global Memory,减少小粒度 Store。

控制 Padding 比例

例如有效长度是 96,而硬件按 128 处理,多出来的 32 个位置就是 Padding。它能换来更规则的指令,也会带来无效搬运和计算。

Grouped MatMul 中尤其要区分两类 Padding:

  • K/N 维 Padding 可能帮助 Cube 使用规则 tile;

  • M 维的组尾 Padding 只能用于单组内部,Mask 不能越过 group_end,否则会错误地使用下一组的 A 数据和当前组的 B 权重

因此,选择 BLOCK_M 时要看 Group 大小分布。如果大量 Group 都很小,大 BLOCK_M 看起来规整,实际上可能大部分时间都在算 Padding。

编译提示与常量特化

在只允许 K 维 Padding 的前提下,可以向 DLCompiler 补充 Dot 相关提示:

dl.compile_hint(a, "dot_pad_only_k")
dl.compile_hint(b, "dot_pad_only_k")

测试中这条提示带来约 0.5% 的性能提升。它更像最后的精修,而不是主要优化手段。只有真实布局满足提示语义时才能使用。

tl.constexpr 可以帮助编译器做常量传播、循环展开和地址化简,适合 BLOCK_M/N/K、Swizzle 尺寸和相对稳定的 num_groups。频繁变化的参数不宜设成 constexpr,否则可能产生较多编译版本。

Cube 与 Vector 的配合

当 Kernel 同时包含矩阵乘和较重的逐元素计算时,可以尝试让 Cube 与 Vector 形成流水。但 C:V = 1:2 不应被视为固定代码结构,实际收益取决于两类任务能否并行,以及数据搬运是否支持这种并行。

如果还要融合 tl.sumtl.max 或激活函数,也要关注操作发生在哪个轴上。有时换一个布局能让归约更快,但布局转换本身也有成本,仍然要看端到端耗时。

性能结果

加速比按 Torch 耗时 / Triton 耗时 计算。结果大于 1,表示 Triton 更快;小于 1,表示 Triton 更慢。例如 1.034× 表示 Triton 约快 3.4%。

Group 数 Total M N K Triton(ms) Torch 融合算子(ms) Triton 加速比
128 655360 4096 4096 81.0825 82.0766 1.012×
128 327680 512 512 0.91926 0.95038 1.034×
128 327680 1536 2048 8.18756 8.08373 0.987×
128 327680 3072 4096 31.4788 31.1182 0.989×
128 327680 4096 1536 16.3721 16.0584 0.981×

这组数据中,Triton 已经接近 Torch 融合算子:最佳测试点约快 3.4%,最大回退约 1.9%。这里不建议计算一个简单平均值,因为不同 N/K 会改变计算强度、B 的工作集和最优 Block Shape。对算子开发者来说,最值得关注的是哪些 shape 仍然回退,以及它们是否需要单独的配置。

下面再列出六组 BF16 测试,Group 数均为 G=8。本文主体讨论 M 分组;表中的 K 分组表示分组发生在 K 方向,只作为补充结果。trans_b 表示 B 是否采用转置读取路径。这些数据用于展示 Grouped MatMul 对输入场景的敏感性,不宜直接计算平均加速比。

加速比仍按 Torch 耗时 / Triton 耗时 计算;“耗时变化”以 Torch 为基准,负数表示 Triton 缩短了耗时,正数表示 Triton 更慢。

序号
分组维度 Group 数 G M N K trans_b Triton(ms) Torch(ms) 加速比 Triton 耗时变化
1 K 8 1536 2048 20480 0.86707 2.49452 2.877× -65.2%
2 K 8 2048 768 20480 0.88558 7.36940 8.322× -88.0%
3 M 8 20480 2048 768 True 0.43832 0.35507 0.810× +23.4%
4 M 8 20480 1536 2048 True 0.83816 1.55567 1.856× -46.1%
5 M 8 20480 1536 2048 False 1.60312 1.46243 0.912× +9.6%
6 M 8 20480 2048 768 False 0.36734 0.29680 0.808× +23.8%

六组结果的差异较为明显:两个 K 分组测试分别获得 2.877×8.322× 加速;四个 M 分组测试中,只有 `trans_b=True、N=1536、K=2048` 的测试点获得 1.856× 加速,其余三个测试点为 0.808×~0.912×,性能仍有回退。这表明当前实现对分组维度、B 的布局以及 M/N/K 组合较为敏感,一套配置还不能覆盖所有场景。

总结

回到开头的问题:Grouped MatMul 难的不是 tl.dot,而是如何处理不规则的 Group。

整个优化过程分为三步:首先通过预计算映射表明确 Group 与 tile 的关系;发现约 1 ms 的前处理开销后,改为让 24 个 program 直接读取 Group 大小,并跨 Group 连续领取 tile;在改善 Core 利用率之后,再通过对角线 Swizzle 优化 L2 局部性。

核心结论可以归纳为四点:

  • Group 边界留在 Device 侧处理,避免 Host 同步;

  • 把不同 Group 的 tile 连成任务流,减少 Core 空闲;

  • 先让 Core 吃满,再讨论 Swizzle 和缓存复用;

  • 评估优化时看端到端耗时,不只看主 Kernel。

如果你喜欢我们的内容,欢迎赞同、收藏、关注我们!
也欢迎在评论区分享你的想法和实践。

参考文献

  1. Triton 官方教程:Group GEMM

  2. Triton 官方教程:Matrix Multiplication

  3. 昇腾开源文档:矩阵乘法(Matrix Multiplication)

Logo

鲲鹏昇腾开发者社区是面向全社会开放的“联接全球计算开发者,聚合华为+生态”的社区,内容涵盖鲲鹏、昇腾资源,帮助开发者快速获取所需的知识、经验、软件、工具、算力,支撑开发者易学、好用、成功,成为核心开发者。

更多推荐