跳到正文
原文
Cohere Labs:官方研究博客·· 23 天前精选AI 评分68

Cohere 详解 North Mini Code 的 megakernel 推理引擎,单 H100 上比 vLLM 快 1.25–1.41 倍

Inside the megakernel serving engine for North Mini Code

AI 导读

Cohere 发布为 North Mini Code 构建的围绕 decode megakernel 的推理引擎,BF16 下单张 H100 端到端解码吞吐比 vLLM 快 1.25–1.41 倍,batch size 1 时达 292 tok/s(SoL 的 62%),代码已在 GitHub 开放。

推荐理由

原文给出手写 decode megakernel 的完整实现细节与真实吞吐数据,读者可以了解如何把现有内核组装成单个持久内核并接入真实服务。

正文 · AI 翻译

今天,Cohere 推出了一款面向 North Mini Code 的推理服务引擎,其核心是一个解码 megakernel:在单张 H100 上使用 BF16,端到端比 vLLM 快 1.25× - 1.41×。可在 GitHub 上探索该服务引擎背后的代码。

大多数 LLM 推理服务栈仍将每次前向传播视为一系列 kernel:启动 QKV,等待;启动 attention,等待;启动 MoE,等待。每次启动本身都没问题。问题在于其间的等待。在小批量下,GPU 在每个解码步骤中都有相当大一部分时间在等待,而非计算。

自回归解码,尤其是在较小批量下,本质上是受内存带宽限制的。对于每个解码步骤,我们从 HBM 中搬运很大一部分内存,而相对而言计算量较少。这意味着正确的问题是:我们能多有效地利用内存带宽,而非 flops。以 North Mini Code 为例,这是一个 30B 模型,每个 token 激活 3.3B 参数,在 BF16 下意味着每个解码步骤要流式传输 6.6 GB 权重,外加 8K 上下文下约 0.5 GB 的 KV cache。H100 通过 HBM 提供 3.35 TB/s 的带宽,使光速上限(SoL)约为 470 tok/s。vLLM 为该模型提供服务时速度为 185 tok/s,仅为 SoL 的 39%。

Megakernel 近来作为弥合这一差距的方法而备受关注:与其使用上百个小 kernel,不如将整个前向传播作为一个持久 kernel 运行。从 Hazy Research 的“Look Ma, No Bubbles!” 这一开创性工作开始(我们将在下文概述其设计),已有大量后续工作发布,实现了不同程度的加速。现有工作主要朝两个方向发展:自动生成 megakernel 的编译器,以及在批量大小为 1 时测量解码速度的独立演示。

我们更进一步。本文介绍的是我们所认为的首个围绕解码 megakernel 构建的完整服务系统。它支持真实服务器所需的一切:连续批处理、分页注意力以及不规则序列长度,全部封装在兼容 OpenAI 的端点之后,并支持工具调用。将 OpenCode 指向它,你就可以用它来编程。

在批量大小为 1 时,我们的 megakernel 达到 292 tok/s,即 SoL 的 62%——比 vLLM 快 1.58×。这一优势在各种批量大小以及高达 256K 的上下文下均保持,且没有可测量的精度损失。

图 1:批量大小为 1 时不同上下文长度的解码吞吐量。Megakernel 始终优于 vLLM。

我们还发现,megakernel 的编写难度远低于其名声所暗示的,因此我们附上了一份将你已有的 kernel 移植为一个 megakernel 的指南。我们的实现是一个单独的 CUDA 文件:没有编译器,没有新的编程范式,没有奇特的抽象——只有普通的 tiled GEMM 和普通的分页注意力,经过重构以适配单一的调用约定。

什么是 Megakernel?

一个 GPU 大致由 100–150 个独立的处理器组成,称为 SM(流式多处理器),它们全部运行同一个程序——一个 kernel——处理不同的数据片段。megakernel 是一个单一的常驻 kernel,运行整个前向传播:我们为每个 SM 精确启动一个线程块,它在整个解码步骤中保持常驻。每个块不是从驱动接收工作,而是读取一个 任务列表——一份它应执行的小块工作清单,由主机准备并放在全局内存中。数据依赖不再由 kernel 边界来编码,而是表示为全局内存中的显式计数器,任务完成时递增,需要输入时自旋等待。

结果是,调度单元从一个完整的操作缩小到一个操作的一个 tile,同步单元从整个 GPU 缩小到任务所依赖的特定生产者。

图 2:一个解码步骤,从主机到设备。主机不再为每个操作启动一个 kernel,而是将该步骤分解为任务——每个任务是一个操作的一个 tile——并以轮询方式将它们分配到各 SM 上,因此每个 SM 在全局内存中都有自己的任务列表。大多数任务的顺序由主机上的调度器决定(见调度器一节)。完整注意力和 MoE 是例外:它们的任务数量取决于实时序列长度和路由,因此它们进入任何 SM 都可以拉取的共享工作队列。我们将在后续章节中描述细节。

加速从何而来

推理引擎通常为每个操作启动一个 kernel——RMSNorm、QKV、注意力、MoE 等等——大部分优化工作都花在让每个 kernel 尽可能快上。这对训练和预填充很有效,因为那里的工作负载是大型计算受限的 GEMM,每个 kernel 都有足够的工作来饱和 SM。

解码则相反:它受内存带宽和延迟限制。一个解码步骤主要是低算术强度的 GEMV,因此其速度取决于权重从 HBM 流式传输到共享内存的速度。任何阻止权重移动的因素都是时间损失,而每个操作一个 kernel 的方法有多个地方会让权重停止移动。这些停顿加起来,占了典型推理引擎未使用的 61% 带宽中的大部分。

最简单的收益来自减少启动和同步开销。在两个连续的 kernel 之间,每个 SM 都必须完成,任何 SM 才能开始下一个,并且驱动必须分派下一个网格。对于一个由每层数十个小 kernel 组成的解码步骤,这些间隙会累积起来。megakernel 每个步骤支付一次这种成本,而不是每个操作支付一次。我们列出另外三个对该模型更重要的收益,大致按影响排序:

1. 减少波量化

假设一个 kernel 有 200 个 tile 的工作要做,而 GPU 有 132 个 SM。前 132 个 tile 并行运行;剩余的 68 个在第二波运行,同时 64 个 SM 空闲。该 kernel 需要两波的时间来完成 1.5 波的工作,而 kernel 越小,这种取整就越糟糕。这不是我们通过更均匀地划分工作就能解决的问题。GEMM tile 形状受矩阵维度和 kernel 设计的约束,因此总 tile 数很少恰好是 SM 数量的整数倍。

图 3:波量化会降低 GPU 利用率。

在 megakernel 中,没有需要向上取整的边界:只要某个 tile 的输入就绪,它就能在任意空闲的 SM 上启动。North Mini Code 从中获益比大多数架构更多,因为它使用了并行 transformer 层:注意力和 MoE 前馈网络都基于同一份归一化输入计算,并且只在层末通过一个融合的残差相加 + RMSNorm 重新汇合,因此注意力和 MoE 都不需要对方的输出。

图 4:North Mini Code 使用的并行 transformer 层。

借助 megakernel,我们可以用就绪的工作“回填”空闲的 SM。并行 transformer 层让我们能够以更激进的方式进行回填:只要有可能,我们就确定性地把可能已就绪可运行的任务放到空闲的 SM 上。放置的细节在任务调度器一节中描述。在下图中,我们对比了由传统服务栈运行的 MoE 解码层和由 megakernel 运行的 MoE 解码层。

图 5:一个 MoE 解码层,相同的操作、相同的顺序,以两种方式调度。上:每个操作一个 kernel。由于波量化和硬件抖动,注意力不会同时完成。kernel 边界处的屏障让每个 SM 都等待最后一个(阴影部分)。同样的模式在每个 kernel 边界处重复。下:megakernel。每个 SM 完成自己的注意力工作后就开始一个 MoE tile,因此同一时间段内承载的是工作而非等待。为便于看清任务,图中画了 16 个 SM 和每个操作几十个 tile;真实 kernel 有 132 个 SM 和数千个任务。时长仅为示意,时间轴是相对的。

2. 消除虚假依赖

即使工作负载完全相同,SM 也不总是在同一时间完成工作。kernel 边界是一个全网格屏障,因此最慢的 SM 决定了所有 SM 的节奏。例如,如果注意力被拆分到 4 个 key/value 组,而其中一组提前完成,那个 SM 就会空闲等待另外三组赶上,即使它下一个操作实际需要的数据已经在内存中。细粒度屏障消除了这种虚假依赖:给定 KV 组的 O-proj 在该组的注意力输出落地后立即开始。类似地,只要对应专家的 up projection 任务完成,MoE down projection 任务就可以开始,而无需等待所有专家的 up projection。

3. 权重预取

权重是不可变的——它们完全不依赖本步骤的激活值。因此,一个任务可以在其激活依赖被满足之前就开始把权重 tile 从 HBM 流式加载到共享内存,而 kernel 边界会禁止这一点。我们在 router 和 QKV 投影上最激进地使用这一点,它们在上一层 O-proj 的尾部、RMSNorm 甚至还没运行时就预取权重,以利用原本会被闲置的带宽。


我们的起点

我们的设计借鉴了前面提到的 Hazy Research 那篇开创性文章的大量洞见,该文章把 Llama-3.2-1B 的前向传播融合进单个 kernel,在 batch size 为 1 时达到 H100 内存带宽的 78%,而 vLLM 和 SGLang 大约只有其一半。他们的三个想法对我们至关重要:

  • GPU 上的“任务解释器”模式。 每个 SM 遍历一份在主机上准备并在前向传播中复用的任务描述符列表。一个控制 warp 读取描述符,并将任务分派给各种设备端函数,每个函数实现一种操作。我们在图 2 和图 6 中描述了该想法的实现。
  • 基于计数器的同步。 依赖屏障是全局内存中的普通整数,在每一步之前清零。任务完成时递增一个,开始前自旋等待一个。我们在“屏障”部分和图 7 中详细描述了我们的实现。
  • 跨任务边界的重叠。 一个任务可以在同一 SM 上的前一个任务仍在存储结果时就开始加载其权重。

我们做了哪些不同

我们的实现与他们的 megakernel 在几个方面有所不同。我们的 GEMM 实现严重依赖张量核心指令,即使批大小为 1 也是如此,因为我们发现在我们的场景中 wgmma 比 CUDA 核心略快,并且减少了寄存器压力。

另一个区别是我们不使用共享内存分页来实现权重预取。我们最初尝试了共享内存分页,它允许内存加载在前一个任务释放其缓冲区之前开始。但在实践中,簿记复杂,引入了持续的错误来源,并且开销很高,超过了收益。相反,每个操作码都获得自己的 warp 专用流水线,其共享内存布局在编译时静态确定,我们从两个更廉价的地方获得重叠。

在同一类型的连续 GEMM 任务之间。 一个 MoE GEMM 任务遍历一个瓦片列表,流水线在整个列表中携带其阶段相位,而不是在每个瓦片边界排空并重新填充流水线。在 MoE GEMM 内部,下一个瓦片的权重已经在传输中,而当前瓦片仍在张量核心或尾声阶段。因为它是具有相同共享内存布局的相同操作,共享内存分页简化为一个简单的多阶段流水线。这种预取与持久分组 GEMM 内核的精神相同。

在 GEMM 流水线内部。 权重是不可变的,因此生产者 warp 在跨 SM 等待激活值之前发出其权重瓦片加载;在本文后面展示的伪代码(列表 1)中,prefetch_weight_tiles 位于 wait_input_bars 之上。MoE 下投影就是这样一个例子:其专家权重开始从 HBM 流式传输,而 up/gate 仍在计算 down 将消费的隐藏状态。因此,一个因输入而阻塞的任务仍在移动字节。另一个例子是在 RMS 范数完成之前预取 QKV 和路由器。

我们还在调度上投入了大量精力。大多数任务遵循主机构建的静态调度,我们可以精确调整;注意力和 MoE 使用本地工作窃取来平衡来自连续批处理和路由的运行时相关工作。我们将在下面回到这一点。

所有操作共用一个 ABI

整个 megakernel 可以看作是由一个通用调用约定拼接起来的多个更小的 kernel——每个更小的 kernel 都必须严格用 3 个 warp 组(每个 warp 组有 4 个 warp)来实现,其中包含 8 个 consumer warp、1 个 controller warp、1 个 producer warp 和 1 个 storer warp,每个 warp 都有各自的寄存器需求。此外,每个更小的 kernel 都必须从一个固定大小的任务描述符中读取它的“参数”。我们觉得这种约定类似于编译器和操作系统中的应用程序二进制接口(ABI)。在整篇文章中,我们会用 “ABI” 这个词来描述这一约定。

让我们以 GEMM 为例来理解 ABI 是如何工作的。

图 6:一个 GEMM tile 如何变成 SM 上的工作。主机将输出划分为多个 tile,并将每个 tile 写成一个 32 位 int32 描述符。在设备上,一个 controller warp 读取下一个描述符,并根据其 opcode 进行切换——QKV、attention、O-proj,或另外十四个中的任意一个——这样同一个 SM 就能连续运行不同的操作。2×4 网格仅作示意;batch 1 下真实的 QKV 是 80 个 tile。

让它仍然可以手工编写的原因在于,每一项操作——不只是 GEMM——都遵守这一 ABI:任务是什么、它以什么形状运行,以及它如何发出完成信号。一旦所有操作都遵循同一套 ABI,将它们组装成一个 megakernel 就变得非常可控。本节剩余部分将拆解这三件事。

任务

megakernel 并没有发明新的操作。它采用通常的 decode 图——QKV、attention、O-proj、router、MoE、residual/RMSNorm、LM head——并将其降低为一个由小型分块 任务 组成的列表。十六个 opcode 覆盖了整个 decode 步骤:

稠密 GEMM

QKV_PROJ、O_PROJ、FFN_UPGATE_ACT、FFN_DOWN、LM_HEAD

一次矩阵乘法的一个输出 tile,可选地是 split-K 归约的一个切片。QKV_PROJ 还会在其 epilogue 中在寄存器内应用 RoPE;FFN_UPGATE_ACT 在其 epilogue 中融合了 SiLU 与乘法。

Attention

ATTN_DECODE、ATTN_COMBINE、ATTN_DRAIN

ATTN_DECODE 在部分 KV 切片上计算 attention;ATTN_COMBINE 将这些部分结果合并为真正的输出。ATTN_DRAIN 仅用于 full attention 层。它是对动态生成的工作队列的认领者——参见后面的调度章节。

MoE 路由

ROUTER_GEMM、ROUTER_TOPK、ROUTE_FINALIZE、MOE_GATHER

为 128 个专家打分,计算每个 token 的 top-k 专家,并生成 MoE GEMM 工作队列。

MoE GEMM

MOE_UPGATE_ACT_DRAIN、MOE_DOWN_DRAIN、MOE_COMBINE

专家 FFN,作为由 ROUTE_FINALIZE 生成的工作队列的认领者。MOE_COMBINE 将每个 token 的 8 个专家输出按其 router 分数加权求和。

所有 GEMM opcode 共享同一条流水线。 O-proj、router、稠密 FFN、LM head 以及两个 MoE 专家 GEMM 都运行相同的 warp 专用主体:producer 将权重和激活 tile 流入共享内存的 stage ring,consumer 在 tensor core 上累加,storer 写出输出 tile 并在一个 barrier 上到达。每个 opcode 变化的是少量逐操作细节——涉及哪些张量、等待哪个 barrier 以及发出哪个 barrier 信号、epilogue 是否融合 SiLU 与乘法、store 是否为 split-K 归约。QKV 就是同一条流水线,只是在写回 HBM 之前在寄存器内应用了 RoPE。MoE drain 也是同一条流水线,只是选择下一个 tile 的方式不同:它们从队列中认领工作,而不是从描述符中读取坐标。

因此,添加新的 GEMM 任务只是一个小改动,而不是编写一个新内核。我们不需要重写加载/计算/存储循环,只需填写这些细节。

任务描述符

每个任务描述符是 32 个 int32 字段,由主机端编码器写入每个 SM 的缓冲区。字段 0 是操作码;其余字段是操作特定的:哪一层、哪个输出 tile、哪个 split-K 切片、要遍历多少个 K-tile,以及——最重要的部分——要等待哪个屏障、等待的计数是多少,以及完成时向哪个屏障发信号。

设备端很直接。每个 block 一个 warp 充当控制器:它将即将到来的描述符预取到一个小型共享内存环形缓冲区中,这样 SM 就不会因等待任务描述符而停顿。其他 warp 从该环形缓冲区中取出任务并执行。

线程块形状

每个任务,无论操作码是什么,都在相同的线程块形状下运行。12 个 warp,组织为 3 个 warpgroup:

WG0 warp 0

控制器

将即将到来的任务描述符预取到共享内存环形缓冲区中

WG0 warp 1

生产者

发出 TMA 加载、初始化信号量、等待输入屏障

WG0 warp 2

存储者

发出 TMA 存储、向输出屏障发信号

WG0 warp 3

空闲

未使用

大部分内容遵循独立 Hopper GEMM 的标准生产者/消费者 warp 特化模式。megakernel 的贡献在于,每个操作码——包括 attention 和 RMSNorm——都使用相同的形状,因此一个刚完成 QKV tile 的 SM 可以接着运行一个 attention 切片,而无需改变其 warp 数量或每个 warp 的角色。

这些角色是编译期标签,而不是运行时分支。每个操作体都基于角色进行模板化,因此 if constexpr (role == PRODUCER) 会在编译期删除不可达代码。这让我们可以在一个函数中编写生产者、消费者和存储者,就像编写普通 GEMM 一样,而每个 warp 只保留它实际运行的路径。

控制器从不进入该函数。它有自己的循环,其唯一职责是将下一个描述符预取到一个小型共享内存环形缓冲区中,并在槽位就绪时通知工作线程。一个小问题是,工作线程之间不能使用 __syncthreads() 进行同步——那会等待独立于工作线程运行的控制器 warp——因此它们通过一个命名屏障(worker_sync)在彼此之间汇合,该屏障排除了控制器。

以下是 megakernel 中实现的 GEMM 任务概览。如果你写过 warp 特化的 Hopper 内核,你可能会觉得这个结构很熟悉。MoE drain 就是同一个任务,在从队列中认领的 tile 上循环调用。

def gemm_task(task):  # compute Y = X @ W
    if constexpr (role == producer):
        prefetch_weight_tiles(task)      # optional, before the wait
        wait_input_bars(task)            # cross-SM: my activations ready?
        for k in k_tiles(task):
            tma_load_A_B(k)              # async copy HBM -> shared memory
            signal_stage_ready(k)        # intra-block: stage k is loaded
    elif constexpr (role == consumer):
        for k in k_tiles(task):
            wait_stage_ready(k)
            wgmma(k)                     # tensor-core MMA
        epilogue_to_smem()
    elif constexpr (role == storer):
        tma_store_Y()                    # async copy shared memory -> HBM
        arrive_output_bar(task)          # cross-SM: my tile is visible
    worker_sync()                        # all worker warps; excluding controller

def moe_drain_task(task):
    while True:
        tile_id = atomicAdd(n_claimed_tiles, 1)
        if tile_id >= len(moe_workqueue):
            break
        tile = moe_workqueue[tile_id]
        gemm_task(tile)

def controller_loop(tasks):
    for i in range(len(tasks)):
        slot = ring[i % RING]
        if i >= RING:
            wait(slot.done)              # workers finished with this slot
        slot.task = load(tasks[i])   # prefetch the next descriptor
        arrive(slot.ready)               # signal workers: you can start

def worker_loop():
    for i in range(len(tasks)):
        slot = ring[i % RING]
        wait(slot.ready)                 # descriptor is in shared memory
        gemm_task(slot.task)             # or switch(opcode) onto another body
        if storer:
            arrive(slot.done)            # slot is free for the next prefetch

清单 1:我们演示如何在 megakernel 中实现 GEMM 任务。整体结构遵循 Hopper 上标准的 warp 特化 GEMM。

if constexpr 分支是上面的编译期标签,而不是运行时 switch。独立内核中不会出现的两行是 wait_input_bars 和 arrive_output_bar——描述符中携带的跨 SM 计数器。位于等待之上的 prefetch_weight_tiles 是另一个可选的 megakernel 专属操作。

控制器始终比工作线程领先一个描述符,而从不加入它们的屏障。Attention、路由和 MoE drain 填充相同的生产者/消费者/存储者槽位;只有 worker_loop 中的 switch (opcode) 会改变。将你已有的内核放入 megakernel 的摩擦就是那两个屏障调用加上这个包装循环,而不是新的编程模型。


屏障

受 Hazy Research 启发,我们的屏障实现为全局内存中的计数器。

// wait: spin until enough upstream tasks have arrived
while (*(volatile const uint32_t*)bar < target) {
    __nanosleep(20);
}
__threadfence();

// arrive: publish my tile, then signal downstream tasks
fence.proxy.async;     // make async (TMA) stores visible first
__threadfence();       // make data computed by this SM visible to others
atomicAdd(bar, 1);

清单 2:我们将屏障实现为全局内存中的计数器。等待的线程进行自旋等待,直到依赖得到满足。

这就是整个依赖机制。一个任务等待单个 count,而不是等待某个特定的上游任务,这使得发出信号和等待依赖的时间复杂度均为 O(1),无论扇入和扇出如何。这些栅栏至关重要,可确保跨 SM 之间不存在静默导致数据损坏的数据竞争。

你如何编写一个 megakernel?


megakernel 听起来像是完全重写。实际上,ABI 将工作限制在可控范围内:一旦某个操作遵循通用的 threadblock 形状和描述符格式,它就会像其他任务一样接入 megakernel。性能关键的逻辑可以从现有 kernel 中复用。

对于从现有 kernel 库开始的人,我们会给出这样的建议:

  1. 从已经具备独立竞争力的 GEMM 和 attention kernel 开始。megakernel 消除的是操作之间的粘合代码;它不会让慢操作变快。我们的实现使用普通的 warp 专用流水线,在没有 TMA 多播或乒乓调度的情况下,单独就能与 cuBLAS 和 FlashAttention-3 匹敌。
  2. 让 kernel 适配 ABI。恰好使用 8 个消费者 warp,最多 3 个生产者/存储 warp。对于标准的 Hopper warp 专用 kernel,这主要意味着将生产者和消费者代码拆分为单独的函数,以便按 warpgroup 的寄存器限制生效。
  3. 添加表达其数据依赖的屏障。一个 GEMM tile 或一个 KV 组成为一个任务。下游操作需要小心地等待屏障,以确保数据完整性。
  4. 将任务描述符写入全局内存任务列表调整任务顺序和放置以实现最佳吞吐量。

困难的部分是第 3 步。计数屏障在 N 个线程到达时释放,而不检查这些线程是哪些,因此一个到达次数错误的操作码可能让 warp 漂移到不同的任务上,并在很久之后死锁。屏障记账需要明确的不变量和仔细的测试。

描述符格式、块形状和计数器协议构成了整个集成契约。这使得添加新操作成为一项边界清晰的工程任务,而不是一个新的 kernel 架构项目。


任务调度器:将任务映射到 SM

在我们实现 megakernel 本身之后,我们需要确定如何将任务分配给每个 SM,以及这些任务应按什么顺序执行。这种分配被称为任务调度。任务调度器接收任务图,并输出每个 SM 按顺序执行的精确任务列表。事实上,我们在设计调度方面有很大的自由度。细粒度屏障仅通过确保任务在其输入就绪后才开始来保证 megakernel 的正确性。只要调度不会导致循环等待,我们就可以随意打乱任务。虽然不同的调度产生相同的输出,但它们会以相当大的方式改变吞吐量。

我们的调度器大多是静态的,仅对少数操作使用动态工作窃取。

静态部分:波次顺序和放置。主机为每一层构建命名波次(qkv、router、attn、moe_up_drain、oproj、rmsnorm、……),选择一个顺序,将它们展平为一个任务列表,并将任务 k 分配给 SM k mod 132。当前的顺序将 router 和 route setup 放在 attention 之前:

qkv → router → top-k → route setup → MoE gather → attention → MoE up/down → O-proj → RMSNorm

这种排序让 MoE 分支获得先机,从而可以将计算密集度较低的路由与更重的 QKV GEMM 一起调度,以更好地利用 SM。除了这个简单的图景之外,调度的影响很难被单独分离出来。改变一个 wave 会改变其依赖项何时就绪、哪些任务共享同一个 SM,以及 HBM 带宽如何在任务之间分配。这些选择会在整个层中相互作用。

我们的实验表明调度确实重要。我们运行了两个 8K 输入长度、均匀路由的消融实验,仅改变 wave 顺序。

交错式。在 attention 开始后,MoE up/down 在 attention combine 之前运行。其理由是让两个分支的就绪工作都保持可用,从而减少 SM 等待时间:

qkv → router → top-k → route setup → MoE gather
→ attention → MoE up/down → attention combine → MoE combine → O-proj → RMSNorm

Attention 优先。在任何 MoE GEMM 之前先启动 attention combine 和 O-proj,同时仍允许 router 提前开始

qkv → router → top-k → route setup → MoE gather
  → attention → attention combine → O-proj → MoE up/down → MoE combine → RMSNorm

1

291

282

3%

236

19%

2

423

406

4%

364

14%

4

553

531

4%

509

8%

所有这些调度都产生正确的输出。我们还尝试了依赖亲和性放置:将任务放在它所依赖的生产者所在的同一个 SM 上,希望复用缓存并减少跨 SM 交接。这使内核变慢了 1–2%,因此发布版本使用朴素的轮询放置。表格显示了实测效果;关于为何某种排序或放置会胜出的完整因果模型仍未解决。

动态部分:针对可变大小工作的本地工作窃取。完整 attention 会为每个活跃请求读取长度差异极大的 KV,而 router 决定每个 MoE expert 接收多少 token。这些 tile 数量只有在 step 期间才会变得已知。

这种拆分是刻意的。静态调度固定了大部分任务顺序和放置,为我们提供了调优性能所需的控制——正如上面的 wave 顺序消融所示。本地工作窃取只处理主机无法提前知道的负载不均衡。静态任务列表预留了固定数量的小型 claimer 任务;一个 ATTN_DRAIN 或 MOE_*_DRAIN claimer 会原子地从其阶段的共享队列中窃取下一个条目并执行它,直到队列为空。Attention claimer 窃取 attention tile,MoE claimer 窃取 MoE tile。claimer 的数量控制并行度,而当前 attention 队列反映批次中的活跃请求。这使静态调度保持稳定,并在无需每次批次变化时重建整个列表的情况下平衡可变工作。

下面的任务图将静态 wave、细粒度依赖和本地认领阶段放在一个视图中:

图 7:内核视角下的一个 MoE 解码层。Attention 被画成四个 KV head,因为 O-proj 的 split-K 等于 KV head 数量,所以每个 O-proj split 恰好依赖一个 head 的 attention 输出——一个就绪的 head 可以在另一个 head 仍处于 QKV 时进入 O-proj。虚线卡片是上文描述的本地认领阶段。

我们尝试过的其他调度器

在项目早期,当我们还以稠密模型为目标时,我们构建了几个更复杂的调度器。其中一个是贪心的、拓扑感知的调度器,它根据每个任务的内存读取来估算其开销。我们还尝试了一种暴力搜索,从数百个随机候选中选出最佳调度。在那个工作负载上,两者的吞吐量都比普通的轮询调度高出约 10%。MoE 改变了问题:路由使得其工作量的大小和位置都是动态的,因此那些针对稠密模型的调度方案无法干净地迁移过来。我们最终选择了更简单的轮询调度,配合调优后的波次顺序,再加上上文所述的本地工作窃取。尽管如此,如果简单的轮询调度被证明不足以应对在 Blackwell GPU 上运行的大内核,或 NVFP4、FP8 等更窄的量化类型,我们计划重新审视那些更复杂的调度器。

内核周围的推理引擎

服务器有两个长期存活的主机侧线程。一个 Python 线程充当控制平面:它接收请求、运行预填充,并管理 KV 容量和批次。一个原生 C++ 线程负责解码,以实现最大性能。这两个线程以交错方式运行:Python 会暂停 C++ 以进行预填充或更改活动批次大小。

这种交接之所以存在,是因为大内核运行的是一个提前构建好的任务列表:瓦片数量、屏障索引以及瓦片所属的槽位都已预先确定。因此,Python 在接纳、退役或重塑批次之前会先停住解码,然后 C++ 再以与新批次大小和上下文相匹配的调度恢复运行。下图展示了双方各自拥有的状态以及调度如何变化:

图 8:在一个步骤开始时,C++ 循环会清除大内核将使用的每个暂存缓冲区和屏障计数器。然后它取最长活跃请求的上下文,为当前批次大小选择一个预构建的调度。该调度固定了 GEMM 瓦片和其他静态工作;当相邻上下文桶的静态任务描述符缓冲区完全相同时,它们会复用同一个调度。全注意力是例外:C++ 根据活跃长度计算每个请求的拆分数量,将其上传,GPU 在排空注意力队列时读取它。

Python 和 C++ 轮流执行

C++ 拥有连续解码循环:它更新 KV 位置、清除每步状态、选择调度、启动大内核、采样,并流式输出生成的 token。只要有活跃请求,它就会重复该循环。Python 拥有所有会改变批次的操作:接纳请求、运行其预填充、分配或驱逐槽位,以及切换到更小的批次大小。

park / resume 箭头就是这两个所有者之间的交接:

# Python control plane
while server_is_running:
    request = wait_for_new_request()

    decode_service.request_pause()       # C++ finishes its current decode step.
    decode_service.wait_until_parked()   # GPU is now idle; C++ will not read batch state.

    prefill(request)                     # Ordinary prefill kernels populate its KV pages.
    decode_service.admit_or_evict(request)
    decode_service.switch_batch_size_if_needed()

    decode_service.resume()              # C++ snapshots the new state and keeps decoding.
// C++ decode service
while (!stop) {
    if (pause_requested || active_requests == 0) {
        signal_parked();                  // Python may now mutate the batch.
        wait_for_resume_or_admission();
        continue;
    }

    update_kv_page_tables();
    clear_scratch_and_barriers();
    select_schedule(max_live_context());
    launch_megakernel();
    sample_and_stream_tokens();
    retire_finished_requests();
}

停住是可变的批次状态的所有权边界。在 C++ 等待期间,Python 会更改每个请求的元数据、token 缓冲区、KV 页表,以及可能为新批次大小准备的指针。C++ 只在恢复后才对该状态进行快照,然后在下一个解码步骤中拥有它。这可以防止大内核观察到半更新的批次。当前的权衡很简单:预填充会暂停解码,因此引擎目前还不会在 GPU 上混合预填充和解码。当没有活跃请求时,C++ 会自行停住,Python 在接纳请求后将其唤醒。

从外部看,这是一个兼容 OpenAI 的服务器,支持流式传输、前缀缓存和工具调用。

已知限制

这些是当前实现的限制,而非设计的限制。我们计划填补它们。

  • 不混合 prefill 和 decode。一个 prefill 会暂停所有 decode 请求。对于大量非常短请求的工作负载,这会略微损害性能。
  • 最大 batch size 为 8。这是配置限制,而非 megakernel 的架构限制。服务器可以在需要时支持更大的 batch size,但我们尚未针对更大的 batch size 调优性能。
  • MK 仅用于 decode。Prefill 作为普通 PyTorch kernel 运行。

性能

Decode 吞吐量

设置。所有数据均在单张 H100(132 个 SM)上运行 North Mini Code 测得。基线为 vLLM v0.24,使用 FA3 attention 后端和 Triton MoE 后端,禁用 prefill。两个引擎都针对合成的 KV cache 进行 decode,以便将 decode 计算与 prefill 隔离开来比较。我们测量 1K 输出 token 上的 decode 吞吐量。

我们报告两种基准测试设置,它们之间的差异颇具启发性。

  • 真实 checkpoint。两个引擎都使用真实的 North Mini Code 权重,因此 MoE 专家分布是模型真实的、相关的路由。
  • 均匀路由。vLLM 使用模拟的均匀随机路由运行;megakernel 使用随机权重运行,从而产生近似均匀的专家分布。

在 8K 上下文下,按每个 batch size 归一化到 vLLM:

图 9:我们比较了 8K 上下文下 megakernel 与 vLLM 的 decode 吞吐量。Megakernel 在真实专家分布下获得更大的加速。

首先,对于均匀路由,当流水线气泡最多时加速最大。在 batch size 为 1 时,每个请求的工作量最少。因此,GPU 时间的很大一部分花在流水线气泡和 kernel 边界之间。随着 batch size 增大,流水线气泡占总时间的比例下降,差距缩小。

其次,加速取决于专家分布。在 batch size 为 8 时,megakernel 在真实专家分布下快 1.32 倍,在均匀路由下快 1.14 倍。真实请求经常选择相同的专家,使活跃专家集合变得稀疏。此时 MoE 的总工作量减少,因此流水线气泡在步骤中占据更大份额;这些正是 megakernel 可以消除的间隙。均匀路由将 token 分散到许多专家,留下更多 MoE 工作和更少可恢复的气泡。因此,合成均匀路由低估了 megakernel 在真实流量上的加速,我们在此将其作为更难的情况报告。

这一优势在不同上下文长度下同样成立。我们在下面展示 BS=4 和 BS=8。

图 10:Megakernel 的加速在不同 batch size 和上下文长度下均成立。

端到端服务

微基准测试隔离了 decode。此测试以 batch size 8 运行完整服务器:真实 prompt 通过 API 进入,prefill 运行,请求在不同时间完成,服务器持续填充 batch。每个引擎生成的 token 数量略有不同,因此仅比较总墙钟时间并不公平。我们同时报告墙钟时间和平均 decode 吞吐量。

端到端平均 decode 吞吐量提升为 1.25 倍 - 1.41 倍,小于仅 decode 的情况,原因有二:prefill 仍使用普通 PyTorch kernel 并暂停 decode;专家路由逐 batch 变化,且服务期间活跃 batch size 不断变化,因此 megakernel 的加速并非恒定。

基准测试 引擎 总时间(秒) 生成的 token 数 平均 decode 吞吐量(tok/s) 加速比
AIME 2025 MK 335 313,205 935 1.41×
vLLM 495 327,296 661
GPQA MK 3,837 3,021,691 787 1.25×
vLLM 4,386 2,768,267 631
MMLU-Pro(CS 子集) MK 1,302 1,235,345 948 1.33×
vLLM 1,756 1,252,448 713
SciCode MK 8,712 6,193,859 711 1.37×
vLLM 11,006 6,166,407 560
LiveCodeBench v6 MK 6,141 4,933,421 803 1.28×
vLLM 5,428 3,394,164 625

准确率

我们还验证了服务路径能够保持模型质量。下表在以下基准上对比了 megakernel 服务器与 vLLM 基线:

SciCode

38.9% ± 1.6%

38.2%

我们报告了 7 次运行的平均分和标准差。我们的 megakernel 得分与 vLLM 对应版本接近,证实了我们的 kernel 是准确的。

我们的收获

Megakernel 起步很简单。 我们不需要编译器或新的编程模型。普通的 GEMM 和 attention kernel 已经完成了困难的数学计算;将它们包装在一个共享的调用约定中——相同的 warp 角色、任务格式和屏障协议——就足以手工组装出一个 megakernel,一次一个操作。

Megakernel 在真实流量下对 MoE 模型获得更多加速。 真实请求经常命中相同的专家,因此 MoE 的工作是稀疏的,为 megakernel 留出更多空闲时间来填充其他就绪工作。

Megakernel 能很好地集成到服务器中。 连续批处理和分页注意力都能与 megakernel 很好地配合。解码作为单个持久 kernel 留在 GPU 上,由 C++ 主机循环和单独的 Python 线程控制引擎。

在 kernel 变快之后,调度就是剩下的加速空间。 哪个 SM 运行哪个 tile,以及以什么顺序运行,是 kernel 调优之外的另一个性能旋钮。

下一步

第一个问题是这个设计能在解码之外延伸多远。Prefill 的形态非常不同,混合 prefill/decode 批次会要求调度器在不中断延迟敏感的解码的情况下保持两种工作负载同时推进。

我们还在为 RTX Blackwell(例如 RTX Pro 6000 和 RTX 50 系列)构建 megakernel,FP8 和 FP4 量化已在路线图中。这是对 ABI 和调度设计在不同 GPU 架构和更低精度 kernel 下能保留多少的有用测试。我们计划很快发布一个 RTX megakernel。

之后,我们希望将同样的方法带到数据中心 Blackwell 和多 GPU 推理中。张量并行和专家并行引入了跨设备集合通信,这增加了一个新的同步边界,并可能带来更多需要由 megakernel 回收的流水线气泡。

开始使用

进一步了解 North Mini Code——Cohere 的首个智能体编码模型——并在 GitHub 上探索 megakernel 背后的代码。

致谢

我们感谢 Bharat Venkitesh 在整个工作中提供的技术支持。我们感谢来自 Nvidia 的 Stephen Jones、Brian Pharris、Vinod Grover、Frederic Bastien、Hua Huang 和 Disha Mehra 提供的各种富有洞见的讨论。我们感谢 Zewen Shen 在评估和数值准确性方面的一些有益讨论。该设计建立在 Hazy Research 的 Look Ma, No Bubbles! 的想法之上,并使用了 ThunderKittens tile 原语。还要感谢 FlashAttention、PyTorch、Transformers 和 vLLM 的作者和维护者。

来源:Cohere Labs:官方研究博客 · cohere.com