跳转到内容

TASO — 让机器自己发现深度学习图重写规则

待复核

TASO 是 2019 年 SOSP 上 Stanford 发的论文,提出一个想法:深度学习编译器里的”图重写规则”,不应该让专家一条条手写,而应该让机器自动枚举、再用数学证明它对,最后让搜索算法选用哪几条

日常类比:以前 XLA / TensorRT 的工程师像中医,靠经验背几百条”这两味药可以替换那一味”的口诀,背错了病人就死。TASO 的思路是把这件事工业化——让机器把所有可能的”两味药等于一味药”的组合枚举出来,每一条都让计算机证明”在任何病人身上都等价”,然后给医生一本有 700 多条经过数学验证的处方手册。

它的核心论点:图级优化的瓶颈不是搜索算法,而是规则集是否够大、够对

不理解这篇论文,下面这些事都没法解释:

  • 为什么后来的图优化器敢把规则数量做大——因为”先生成、再机器证明”比纯手写更敢扩
  • 为什么 MLIR 一类栈会强调 pattern 描述语言 + 校验——同一类”可检查的重写”思路
  • 为什么”自动机器学习编译器”说得通——规则正确性可证明时,规则集才能放大一个数量级
  • 为什么 xla-compiler / TensorRT 仍要维护庞大 pattern 库——手写路线的规模化成本,正是 TASO 要工业化的痛点

TASO 的设计是三个解耦的阶段

  1. 枚举(Generator):把基本算子(conv / matmul / add / split / concat 等)当原子,枚举所有节点数不超过 4 的子图。这一步不管对不对,纯粹”先把候选写满”。

  2. 验证(Verifier):从一组代数公理(结合律、分配律、线性性、矩阵乘的关联性等)出发,自动判断”子图 A 和子图 B 是否在任意输入下输出相同”。等价的两个子图就成为一条候选重写规则

  3. 搜索(Optimizer):把验证过的约 743 条规则喂给一个 cost-based 回溯搜索器。搜索器以”整图执行总时间”为目标函数,反复尝试用规则替换图中的子结构,挑选最快的重写序列。

类比:第一步像把所有”可能的拼图块”印出来;第二步像让验厂师挨个检查”这块和那块是不是真的能互换”;第三步像让机器人在拼图盘上反复换块,直到拼出最快的那种摆法。

案例 1:一条 TASO 自动发现的规则

Section titled “案例 1:一条 TASO 自动发现的规则”

下面是论文里的一条典型规则(手写库里没有这一条):

Conv(x, w1) + Conv(x, w2) 等价于 Conv(x, concat(w1, w2)) 的某个切片

逐部分解释:

  • 左边:用 x 分别和两个卷积核 w1w2 做卷积,再相加——两次访存、两次卷积
  • 右边:先把 w1w2 在通道维度拼起来,再做一次卷积——一次访存、一次卷积
  • 验证器靠卷积对加法的分配律证明两边等价
  • 搜索器在 GPU 上比较时间,发现右边在大 batch 下快 2 倍

工程师没写这条规则,TASO 自己枚举出来并证明了。这种”两次小卷积合一次大卷积”在传统库里要靠人观察 ResNet 才会想到,TASO 把它升级成对任意网络都成立的通用模式。

给 TASO 一张 BERT 一类的计算图,搜索循环可以想成下面这段伪代码:

queue = [原图]
while queue 非空:
G = 取出当前估时最短的图
for 每条已验证规则 R, 每个可匹配位置:
G2 = 用 R 的右边替换该位置
若 cost(G2) 更低 → 把 G2 放回 queue # 变慢就丢弃
输出 queue 里估时最低的图

逐部分解释:匹配像在图里找”能套进规则左边”的拼图块;替换后用 cost model(各算子实测耗时之和)估整图快慢;变快才留下。论文里整次搜索通常 不到 10 分钟;对新架构相对当时 SOTA 框架最高约 2.8×(cuDNN 后端约 1.3–2.8×),不是单点对比 XLA 的固定倍数。

不验证会怎样?看这条错误规则:

Reshape(x, [a, b]) + Reshape(x, [b, a]) 假装等于 2 * Reshape(x, [a, b])

形状不一样,不能加,但模式匹配引擎可能因为算子名匹配就替换了。手写库里这种 bug 修过很多次。TASO 的验证器(用 Z3 这类定理证明器,像验算老师拿公理本逐条核对)会立刻拒绝——因为代数公理推不出两边等价。

  1. 枚举爆炸:节点数从 4 加到 5,子图候选数量暴涨。论文工程上限定到 4 节点,务实但也是局限。
  2. 公理库不闭合:代数公理是手写的。新加 LayerNorm 却没补公理,相关规则会被拒——自动化只覆盖”规则发现”,没消掉人工。
  3. cost model 不准:GPU 受 launch / cache / 带宽影响;论文用算子实测建 cost,否则会选出”理论快、实际慢”的图。
  4. 数值精度变形:浮点结合律不严格成立;TASO 证的是代数等价,不保证比特等价,训练时可能漂 loss。

适用

  • 深度学习推理编译器(XLA / TensorRT / IREE / TVM Relay)的图级优化阶段
  • 需要把”规则正确性”当一等公民、规则集要快速扩张但不能引入 bug 的工业场景

不适用

  • 规则空间太大无法枚举(LLM 训练里超大融合子图,节点数轻松破百)
  • 算子语义难用代数公理刻画(随机算子、控制流密集代码)
  • 运行时热路径——搜索通常要数分钟,适合编译期预优化

把上面串成 TASO 优化一张模型图的全过程:

  1. 离线生成:枚举 ≤4 算子子图,先得到约 28744 条候选,剪枝后剩约 743 条(论文约 5 分钟量级,不是十小时)
  2. 离线验证:用约 43 条算子公理,在 Z3 里核对这 743 条是否代数等价(剪枝后的集合再验证,不是”通过率 60%”)
  3. 在线搜索:对输入图做 cost-based 回溯,尝试规则替换;论文实验里发现优化图通常 不到 10 分钟
  4. 输出:把等价更快的图交给下游 backend(如 cuDNN / TVM)做 codegen

用户侧只需丢进模型,编译期几分钟后拿到一张更快的等价图。

  • 2018 年:Zhihao Jia 的前作 MetaFlow 已把”图重写 + 回溯搜索”结合,但规则仍手写
  • 2019 年:与 Oded Padon 等合作加入公理验证与自动枚举——TASO 诞生;相对当时 TensorFlow / TensorRT 等框架,新架构上最高约 2.8×(继承并扩展 MetaFlow 的搜索,而非简单”比 MetaFlow 快 2.8×”)
  • 2020 年起:同一路线延伸到 FlexFlow / Unity 分布式,影响了 alpa-2022 的并行策略搜索
  • 2022 年:PyTorch 2.0 Inductor 讨论里常把 TASO 当作”少手写 fusion pattern”的参照
  • 2024 年前后:Mirage 等系统把”枚举 + 验证 + 搜索”推到更大粒度的 kernel 融合
  1. 解耦比聪明更重要:规则发现 / 验证 / 使用拆开,每段可独立替换——这是工程贡献
  2. 形式化验证能省钱:七百多条规则靠机器证明,比人工 review 便宜
  3. 瓶颈常在搜索空间:没换搜索算法,把候选规则扩到约 743 条,效果就上去了
  4. 公理是源真相:加新算子先扩公理再扩规则;昂贵枚举离线一次,便宜搜索在线按模型跑
  • xla-compiler —— XLA 图级 fusion / pattern 与 TASO 思路可对照
  • tvm-2018 —— TVM 重 schedule,TASO 重图重写;两层正交
  • alpa-2022 —— 把”自动搜索”从图重写扩到并行策略
  • pytorch —— Inductor 一代常引用”少手写 fusion”叙事
  • mlir —— PDL + 验证器延续”生成/描述 + 校验”模式