KaiSpace
tech

从 GQA 例子读懂 Mirage 的 muGraph 搜索与后端生成链路

本文只讨论 Mirage(OSDI 25'),不讨论 python/mirage/mpk 下的 MPK 子系统(这个更偏向于runtime scheduling)。请注意,这两个项目几乎没有联系,我并不清楚为什么要放在同一个repo下。(当然也有可能是我没看出来联系)

文章按真实执行顺序展开:先用 Python 构造 high-level graph,再进入 Cython/C++ search 生成候选 muGraph,随后 verifier 判断语义等价,最后由 CUDA 或 Triton backend 尝试 lower、compile、profile,并选出最快的 runnable graph。

image.png

1. 从 GQA benchmark 进入

benchmark/group_query_attention.py 是一个适合读 Mirage 主线的入口。它构造的 high-level graph 很短:

graph = mi.new_kernel_graph()
Q = graph.new_input(dims=(2 * batch_size, 256, 64), dtype=mi.float16)
K = graph.new_input(dims=(2 * batch_size, 64, 4096), dtype=mi.float16)
V = graph.new_input(dims=(2 * batch_size, 4096, 64), dtype=mi.float16)
A = graph.matmul(Q, K)
E = graph.exp(A)
S = graph.reduction(E, 2)
D = graph.div(E, S)
O = graph.matmul(D, V)
graph.mark_output(O)

这条 graph 对应的 attention 形态是:

QK -> exp -> reduction -> div -> PV

默认 batch_size=1 时,Q(2, 256, 64)K(2, 64, 4096)V(2, 4096, 64)。当前代码注释把它解释成 batch=1kv_heads=2q_per_kv*head_dim=256mirage_head_dim=64

这里有一个版本差异要先说明:这个 benchmark 第 40 行调用了

graph.superoptimize(config="attention", previous_checkpoint=filename, ...)

但当前 checkout 中 python/mirage/kernel.pyKNGraph.superoptimize 签名没有 previous_checkpoint 参数。因此,这个文件适合作为代码阅读入口,不能在不检查版本差异的情况下当成当前 checkout 的直接复现实验脚本。

2. Python 建图阶段:KNGraph 还不是 kernel

Python 侧核心对象是 KNGraphmi.new_kernel_graph() 返回一个 kernel-level graph wrapper,它内部持有 cygraph,也就是 Cython 包装过的 C++ graph 对象。

KNGraph 上的 matmul()exp()silu()div()reduction() 等方法基本是薄包装:Python 方法把调用转发给 self.cygraph,返回新的 DTensor 元数据对象。这个阶段只是在扩展计算图,没有生成 CUDA,也没有运行 kernel。

new_input() 会把 input 的 shape、stride、dtype 写进 graph 元数据。默认没有传 stride 时,Python 侧构造 row-major contiguous stride。后面 CUDA transpiler 会按这些 stride 特化地址计算,所以 compile() 阶段会检查运行时传入 tensor 的 shape 和 stride 必须与 graph 中的 DTensor 元数据一致。

fuse_tensors()narrow()split() 这类接口也属于 graph mutation。它们会修改 cygraph 中的 tensor 组织方式,但仍然不是后端 tile schedule,也不是 CUDA codegen。

2.1 KNGraphTBGraph 的层级分工

Mirage 的关键抽象是 multi-level graph。当前代码和文档里经常出现两类 tensor:

层级操作对象主要决定例子
KNGraphDTensor外层 kernel graph 结构matmul/add/reduction/customized
KNGraphDTensorop 输入连接matmul(A, B) 还是 matmul(C, D)
KNGraphDTensor是否生成 KN_CUSTOMIZED_OP把一段计算换成 graph-defined kernel
TBGraphSTensorCUDA block/tile 映射grid_dimblock_dim
TBGraphSTensorinput 如何读 tileinput_mapforloop_dim
TBGraphSTensorblock 内部循环forloop_range
TBGraphSTensorblock 内部计算TB_MATMULTB_REDUCTIONTB_FORLOOP_ACCUM
TBGraphSTensoroutput 如何写回output_map
TBGraphSTensor累加和后处理accumulator、epilogue

input_map 指定外层 DTensor 的维度如何映射到 blockIdx.x/y/z,从而决定当前 block 读哪块 tile。forloop_dim 指定这个 TB input 在 threadblock 内部沿哪个 tensor 维度做 for-loop 迭代读取或累加。当前实现把 forloop_dim 放在 TB input 上;从抽象设计角度看,它也可以被理解成一种和 input loader / accumulator 相关的 op-level 属性。

KNGraph 描述“有哪些 kernel-level tensor 和 kernel-level op”。它本身不会做细粒度 tile 切分。细粒度 tile/block 语义出现在 TBGraph 里,通常嵌在 kernel graph 的 KN_CUSTOMIZED_OP 内部。

image.png

Mirage 文档也支持这个边界:kernel graph 中的 tensor 在 device memory;block graph 中的中间 tensor 主要在 shared memory。kernel graph 的 graph-defined operator 由 lower-level block graph 定义语义和行为。

3. superoptimize 的总控逻辑

image.png

KNGraph.superoptimize() 是 Mirage 的高层优化入口。当前 checkout 中它的默认参数包括:

backend: str = "cuda"
warmup_iters: int = 16
profile_iters: int = 1000
use_graph_dataset: bool = True
use_cached_graphs: bool = True
is_formal_verified: bool = False

它的真实顺序是:

1. 查 graph_dataset
   如果已有同一 input graph + search 参数 + backend 的 optimized graph,直接返回。

2. 准备 search checkpoint
   use_cached_graphs=True 时,用 high-level graph hash 生成
   mirage_cached_mugraphs_<hash>.json。

3. 调用 search(...)
   Cython 把 Python 参数传给 C++ search_c.cc;
   C++ search 生成一批语义等价候选 CyKNGraph。

4. 进入 backend
   CUDA backend 编译/profile 候选;
   Triton backend 生成 Triton 文件/profile 候选。

5. 返回最快的 runnable graph

这里有两层缓存,作用不同。

use_graph_dataset=True 查的是进程内的最终 optimized graph。命中后,superoptimize() 直接返回已经选好的 graph。它更符合“我只需要 best graph”的需求。当前 graph_dataset.py 是一个内存字典,不是磁盘数据库。

use_cached_graphs=True 查的是 search 阶段产出的候选 muGraph JSON。命中后可以跳过 C++ search,但后续仍然要让 backend compile/profile 候选,再选最快者。因此它只省 search 时间,不省后端筛选时间。

4. Search 阶段:从 Python 参数到候选 muGraph

这一节应该放在 backend compile path 之前,因为实际执行中 search 先产生候选 graph,backend 后面才处理这些候选。

4.1 Python 到 C++ search 的桥

superoptimize() 没有命中 graph_dataset 时,会调用 Cython 侧的 search()

graph.superoptimize(...)
  -> search(...)
  -> python/mirage/_cython/core.pyx: def search(...)
  -> cython_search(...)
  -> src/search/search_c.cc: mirage::search_c::cython_search(...)

core.pyx 的职责是把 Python list/tuple 形式的搜索配置搬运成 C++ vector,然后调用 cython_search()

  • imaps 转成 vector<MInt3>
  • omaps 转成 vector<MInt3>
  • griddims 转成 vector<MDim3>
  • fmaps 转成 vector<int>
  • franges 转成 vector<int>

当前 checkout 里 blockdims 是特殊点:core.pyx 明确写了

assert blockdims is None, "TODO: support blockdims"

因此虽然 C++ GeneratorConfigblock_dim_to_explore,但从当前 Python superoptimize()blockdims 不是有效调参入口。

另一个容易误解的参数是 max_num_new_graphs。Cython 默认值是 1024,并用一个固定大小的 C 数组接收返回 graph。search_c.ccmax_num_graphs 的作用是保护 caller-provided output array,限制最多拷贝多少个结果回 Python。它不是 KernelGraphGenerator 内部搜索的 early stop 条件。

4.2 search config 控制什么

Mirage search 枚举候选 KNGraph/TBGraph 结构。用户传入的 config 会限制这个搜索空间。

imap 描述 block grid 维度如何映射到 input tensor 的数据维度。举例:

imap = {0, 1, -1}

可以理解为:

input tensor 的第 0 维由 blockIdx.x 切分
input tensor 的第 1 维由 blockIdx.y 切分
第 3 个 grid 维度不映射到这个 input tensor 的数据维度

Mirage 文档把不映射到数据维度的情况称为 replica dimension,即同一份输入在对应 grid 维度上被复制给多个 block 使用。

omap 描述 block graph 的局部 output 如何拼回 kernel-level output DTensor。和 imap 不同,omap 需要把不同 block 的输出写到 disjoint device memory 区域,否则多个 block 会写同一片输出。

grid_dim_to_explore 控制 CUDA grid 形状,也就是 x/y/z 方向各有多少个 block。

block_dim_to_explore 控制一个 CUDA block 内部的 thread 排布。不过如上所述,当前 Python wrapper 没有支持从 superoptimize()blockdims

fmap_to_explore 描述 threadblock graph 内部 for-loop 沿哪个 tensor 维度迭代/累加。

frange_to_explore 描述 for-loop 迭代多少次。它直接影响每次迭代处理的 tile 范围,以及 block graph 是否需要 accumulation。

4.3 search_c.cc 的入口职责

src/search/search_c.cc 是 Cython 到 C++ search 的入口。它有两个职责:

  1. 如果 filename 存在,把它当作 saved muGraph list 读入,跳过搜索。
  2. 如果 filename 不存在,构造 KernelGraphGenerator,运行搜索,并把结果写到 checkpoint 文件。

默认 search config 来自 GeneratorConfig::get_default_config()。其中默认 verifier 是 PROBABILISTIC_VERIFIER。当 Python 传 config="attention" 时,search_c.cc 会调用 enable_attention_specific_optimization(),这会打开 attention-specific flag,把 max_num_threadblock_graphs 设为 2,并把 TB_FORLOOP_ACCUM_REDTOX_LD_SUM_OP 加入可探索 TB op 集合。

4.4 search.cc 如何递归生成候选 graph

真正的递归搜索在 src/search/search.cc。核心函数是 KernelGraphGenerator::generate_next_operator()。它维护一个 SearchContext,而 context 有两个层级:

LV_KERNEL: 给外层 kernel graph 追加 op
LV_THREADBLOCK: 填充一个 KN_CUSTOMIZED_OP 内部的 threadblock graph

当搜索处在 kernel level 时,它会枚举普通 kernel op,例如 KN_MATMUL_OPKN_EXP_OPKN_DIV_OP,也可能枚举 KN_CUSTOMIZED_OP。如果选择了 KN_CUSTOMIZED_OP,搜索会进入 threadblock level,继续填充内部 TBGraph,包括 input mapping、for-loop、block-level op、output mapping 等。

每次尝试增加 op 前,search 会用 abstract expression 做便宜剪枝。当前代码在 check_abstract_expr() 中调用 subexpr_to_final_expr(expr),这里用 egg 做 sub-expression 检查。这个检查不是最终正确性证明,而是在昂贵 verifier 之前剪掉明显不可能贡献到目标输出的候选。代码中用egg和下列规则进行验证筛选。

image.png

最终判断发生在 KernelGraphGenerator::verify()。候选 graph 暂时没有显式 output marker;verifier 会检查候选 graph 的 frontier tensor 是否能按某种输出排列匹配目标 high-level graph。验证成功后,search 临时 mark output,把完整 graph 加入 generated_graphs,然后移除 output marker,让递归搜索继续。

4.5 ProbabilisticVerifier 和 FormalVerifier

ProbabilisticVerifier 是默认 verifier。它不是用真实用户输入运行 graph,而是为 input tensor 初始化 fingerprint,再按每个 op 的 fingerprint 规则传播,最后比较候选 output fingerprint 和原始 graph output fingerprint。它速度快,但本质是 probabilistic equivalence check,有理论上的碰撞或误判可能。

FormalVerifier 通过 is_formal_verified=True 启用。当前实现不是直接走 Z3,而是把 graph 转成 S-expression 风格的代数表达式,再调用 Rust egg e-graph rewrite 做等价检查。它会显式表达 partitionreplicatecombinereducepartial_sum 等 muGraph 语义。它更接近符号等价检查,但能力受表达式生成和 rewrite rules 覆盖范围限制,速度也比 fingerprint 路径更重。

两者都返回 OutputMatch,而不是简单 bool。原因是候选 graph 的 frontier outputs 顺序不一定和原始 graph 的 outputs 顺序一致,所以 verifier 会枚举输出排列,找到一个合法对应关系。

4.6 generate_next_symbolic_operator 是什么

search.cc 里还有 generate_next_symbolic_operator()。它是 concrete search 的 symbolic 版本。这个在作者的新论文PRISM里面有提到。

concrete search 直接枚举具体 graph、具体 mapping 和具体 loop range。symbolic search 会先构造带符号维度的 graph template,再调用 instantiate 逻辑给符号变量赋值。代码里 instantiate_symbolic_graph() 填充这些参数,最后会把 symbolic graph 转成 concrete kernel graph,并再次调用 verify()

当前普通 Python superoptimize() 路径调用的是 generate_kernel_graphs(),也就是 concrete search。generate_kernel_graphs_symbolic() 是独立入口,不是这条主线的默认路径。另外,在最新repo(2026/06/26)中并未看到完整的generate_next_symbolic_operator实现,似乎有参数没有完整填充,仍是不可用状态。

5. Backend 阶段:search-valid 之后才 compile/profile

Search 返回的是一批候选 graph。它们已经通过 verifier,但还不是可运行 kernel。CUDA/Triton backend 的任务是:尝试 lower 候选、编译或生成运行文件、实际运行 profile,并选最快者。

5.1 最重要的边界:search-valid 不等于 runnable kernel

Mirage 的 search verifier 检查的是语义等价:候选 muGraph 是否和用户给出的 high-level graph 计算同一个结果。

CUDA/Triton 后端检查的是另一件事:这个候选能不能在目标后端 lower、compile、run,并且性能是否值得选中。

search-valid muGraph != runnable CUDA kernel

在 CUDA backend 下,一个候选通过 verifier 后仍然可能失败:

  • generated shared memory 超过目标 GPU 上限,compile() 会把 graph 标成 invalid。
  • Hopper lowering 要求 block thread count 与 producer/consumer warp-group 配置兼容;不兼容时可能返回 CUDA_T_CONFIG_ERROR
  • generated CUDA code 可能被 nvcc 编译失败,Python 侧会记录 "CUDA compilation error"

include/mirage/transpiler/error_types.h 里定义了:

CUDA_T_SUCCESS = 0
CUDA_T_INSUFFICIENT_SMEM = 1
CUDA_T_LAYOUT_ERROR = 2
CUDA_T_CONFIG_ERROR = 3

src/transpiler/transpiler_tb_hopper.cc 里也能看到 Hopper-specific 约束。对于带 for-loop 的 threadblock graph,thread 数需要匹配 selected producer+consumer warp-group 数量;否则这个 graph 即使 search-valid,也不能按 Hopper lowering 的执行模型生成合法 kernel。

image.png

5.2 CUDA backend 的调用链

CUDA 分支从 KNGraph.superoptimize(..., backend="cuda") 的 backend 部分开始。此时 search 已经结束,all_graphs 是候选 muGraph 列表。

当前代码路径是:

KNGraph.superoptimize(..., backend="cuda")
  -> 对每个候选 g.compile(async_=True, ...)
     Hopper/Blackwell: sweep pipeline_stages 与 num_warp_groups
     Ampere/older: 使用默认编译配置
  -> KNGraph.compile(...)
  -> generate_cuda_program(...)
  -> python/mirage/_cython/core.pyx: generate_cuda_program(...)
  -> Cython 声明的 transpiler::transpile(...)
  -> src/transpiler/transpile.cc: transpiler::transpile(...)

compile() 先调用 C++ transpiler 生成 CUDA 源码。返回结果包含:

  • generated CUDA code
  • runtime buffer size
  • max shared memory size
  • profiler buffer size
  • output tensor allocation directives

随后 Python 侧把 generated CUDA code 和 kernel.py 里的 HARD_CODE launcher 拼在一起,写入临时目录的 test.cu,再用本机 nvcc 编译成 Python extension module。导入成功后,KNGraph.run 会指向扩展模块里的 launch()

CUDA backend 会真的运行候选 graph 来选最快者。默认每个可运行候选都会执行 16 次 warmup,再执行 1000 次 profile iteration,最后用 CUDA event 统计平均耗时:

perf = elapsed_time / profile_iters

这个数字很大,所以搜索候选很多时,实际 profile 成本会非常高。

Hopper/Blackwell 分支还有一个额外细节:superoptimize() 会尝试 pipeline_stages in [2, 3, 4]num_warp_groups in [2, 3, 4] 的组合。Cython 侧会把 num_warp_groups 转成 num_producer_wgs=1num_consumer_wgs=num_warp_groups-1。这个实现似乎和tilelang的后端类似,即一种确定的one-pass算法进行单层的consumer-producer规划,不会识别多层A->B->C->D这种pipeline。

5.3 backend优化

论文中提到了三个backend优化路径:

Layout optimization
是什么:决定每个 DTensor/STensor 在内存里怎么排列,比如哪个维度连续、stride 怎么设、shared memory 是否 swizzle。
怎么分析:把 layout 选择建成 Z3 optimization 问题,用约束保证合法,用 cost model 偏好 cp.async、vectorized copy、ldmatrix、低 shared-memory
开销等更快布局。

Operator scheduling
是什么:决定 threadblock graph 里的 operator 按什么顺序执行、哪里插同步、哪些 op 可以 fuse 成 chain。
怎么分析:基于 TB graph 的依赖关系、tensor producer/consumer关系、fusion 规则和同步需求,构造 pre-loop / loop / post-loop 的执行 schedule。

Memory planning
是什么:决定中间 DTensor/STensor 放在哪段 global/shared buffer,尽量复用内存、降低峰值占用。
怎么分析:根据 schedule 推导每个 tensor 的生命周期,然后用分配算法把生命周期不重叠的 tensor 放到同一段内存。

5.4 Triton backend 的调用链

Triton 分支独立于 CUDA/Hopper transpiler。它同样发生在 search 之后,输入也是候选 muGraph 列表。

当前路径是:

KNGraph.superoptimize(..., backend="triton")
  -> profile_and_select_best_graph(...)
  -> TritonProfiler.profile_graphs(...)
  -> _generate_profile_code_file(...)
  -> generate_triton_program(...)
  -> python/mirage/_cython/core.pyx: generate_triton_program(...)
  -> src/triton_transpiler/transpile.cc: triton_transpiler::transpile(...)

TritonProfiler 会为每个候选 graph 生成临时 Python 文件。这个文件里包含 generated Triton code 和 profile harness。默认同样是 16 次 warmup、1000 次 profile iteration。profiler 会运行这些临时文件,记录平均耗时,选择最快的 graph。

CUDA backend 和 Triton backend 的共同点是:它们都会对 search 返回的候选进行实际 profiling。不同点是:CUDA backend 通过 C++ CUDA transpiler 生成 .cu 并用 nvcc 编成临时 Python extension;Triton backend 生成 Python/Triton 文件并运行它。

6. 后端结构

从代码结构看,Mirage 的 CUDA backend 更像一个受控的 domain-specific CUDA emitter,而不是一个通用编译器后端。

这个判断来自几个具体事实:

  • CUDA transpiler 从 KNGraph 生成 CUDA 程序,入口是 src/transpiler/transpile.cctranspile()
  • kernel-level op 会在 transpiler_kn.cc 里被逐类 lowering。
  • KN_CUSTOMIZED_OP 会分发到 threadblock graph lowering。
  • Hopper 路径在 transpiler_tb_hopper.cc 里显式处理 warp-group、TMA、swizzle、shared-memory planning、pipeline stage 等。
  • 支持哪些 op、哪些 fusion、哪些 lowering pattern,很大程度上由 Mirage 自己的 operator enum、runtime helper 和 transpiler 分支决定。

这不如“把所有东西都交给通用 IR 和通用 scheduler”抽象,但它有直接好处:search 生成的结构、verifier 的假设、后端 lowering 的语义之间更可控。对 Mirage 这种会搜索 graph-level 和 threadblock-level 结构的系统来说,控制变量少是实际工程优势。

6.1 我对于其后端的想法

TVM 可以作为 Mirage backend 的一个可选 lowering/tuning 后端,而且可以通过手写 TensorIR schedule 保留细粒度排布;前提是 Mirage 先定义精确的 TBGraph -> TensorIR 语义映射,并明确哪些 schedule 维度可以交给 TVM 调,哪些必须由 Mirage search 结果固定。

7. 一句话总结

Python 写 high-level KNGraph,
search 枚举语义等价的 multi-level muGraph,
verifier 只保证 search-valid,
CUDA/Triton transpiler 再决定候选是否可 lower、可编译、可运行,
最后通过实际 profiling 选出最快 graph。

8. 一些代码的缺失

image.png

  1. 论文中画出了三个层级的graph,但是似乎代码中只有KNGraph和TBGraph两层。

  2. benchmark/group_query_attention.pydemo/demo_group_query_attention.py 无法正确运行,故目前无法确定其实际性能。后续会对其进行实测,TBC。

9. 推荐代码阅读路径

  1. mirage/benchmark/group_query_attention.py
  2. mirage/python/mirage/kernel.py
  3. mirage/python/mirage/_cython/core.pyx
  4. mirage/src/search/search_c.cc
  5. mirage/src/search/search.cc
  6. mirage/src/search/config.* and dim_strategy.*
  7. mirage/src/search/verification/probabilistic_verifier.cc
  8. mirage/src/transpiler/transpiler_kn.cc
  9. mirage/src/transpiler/transpiler_tb_hopper.cc
  10. mirage/include/mirage/transpiler/error_types.h

看完这些可以理清大致逻辑。

Comments

No comments yet.