跳到主要内容

数据截至 (上游 commit 118e812b316f)

02 · UOp 与图重写引擎

这一章讲什么: tinygrad 的中枢神经系统。读懂 UOp 这一种数据结构和 graph_rewrite 这一种变换,你就拿到了读全库任何一处的钥匙——因为求导、融合、优化、代码生成全是它们的组合。


1. 它要解决的小问题

一个编译器要经过十几轮变换:化简、求导、融合、切 kernel、展开循环、生成代码。每轮变换如果都发明一种 IR 和一套遍历代码,框架就爆炸了。

tinygrad 的回答很激进:IR 只有一种,变换只有一种

  • IR:UOp——一棵不可变的、带 op 码的表达式树;
  • 变换:graph_rewrite(图, PatternMatcher(规则表))——模式匹配重写。

这正是它能把全套功能压进两万来行的根本原因。


2. UOp:唯一的节点

2.1 五个字段

UOp 是一个不可变 dataclass(tinygrad/uop/ops.py:237-242):

@dataclass(eq=False, slots=True)
class UOp(RandMixin, metaclass=UOpMetaClass):
op:Ops # 这个节点是什么操作(枚举,见 §3)
dtype:DType = dtypes.void # 结果的 dtype(void 表示不返回值)
src:tuple[UOp, ...] = tuple() # 子节点(输入)
arg:Any = None # 随 op 而定的附加参数(shape、axis、常量值…)
tag:Any = None # 临时标记(调度/跟踪用)

不同生命周期阶段的图,只是 Ops 的不同子集:张量图里有 RESHAPE/EXPAND/REDUCE,kernel AST 里有 RANGE/INDEX/LOAD/STORE,渲染前的线性列表里有 SPECIAL/IF/END。一种结构贯穿到底。

2.2 全局 interning:同构子图即同一对象

UOp 不能直接 UOp(...) 裸造——元类 UOpMetaClass.__call__(tinygrad/uop/ops.py:194-215)先用 (op, dtype, src, arg, tag) 查全局缓存 ucache(tinygrad/uop/ops.py:199-200),命中就返回已有对象。这就是 hash-consing:

  • 相等 = 指针相等,dict/set 操作全是 O(1);
  • UOp.key(tinygrad/uop/ops.py:268-270)给出整棵子树的内容哈希(sha256),调度缓存、编译缓存、BEAM 缓存都拿它当键。

dtype 也不是白给的:dtype_from_uop(tinygrad/uop/ops.py:118-192)按 op 推导(例如 SIN/LOG2 自动提升到最小浮点、WHERE 取分支的最小上界),构造时自动填。

2.3 shape 在图上,不在 Tensor 上

_shape 是一个递归属性(tinygrad/uop/ops.py:330-465),按 op 逐个定义 shape 规则:movement op 做 shape 变换,REDUCE 砍掉前 arg[1] 个轴(tinygrad/uop/ops.py:445-449),CONST 是标量。

值得单独点出的是:广播规则长在 _shape(tinygrad/uop/ops.py:455-462)——任何 Broadcastable 节点的 shape 是所有输入 shape 右对齐广播的结果(_broadcast_shape,tinygrad/uop/ops.py:74-82)。所以 Tensor 根本不用存 shape。

2.4 遍历就是拓扑排序

toposort(tinygrad/uop/ops.py:297-309)是显式栈实现的迭代版(避免递归爆栈),返回「依赖在前」的节点序;backward_slice(tinygrad/uop/ops.py:279-283)就是去掉自己的 toposort。全库所有「遍历这张图」的需求都收敛到这一个函数。


3. Ops 全集地图

Ops 是一个 IntEnum(tinygrad/uop/__init__.py:13),文件头注释提醒:枚举顺序决定 toposort 的顺序(tinygrad/uop/__init__.py:12)。按用途分群(注释里的分组即源码里的分组):

代表成员出现阶段
定义/特殊BUFFERPARAMSPECIAL(GPU 维度)、VARIABLE贯穿
非渲染节点NOOPSINKAFTERGROUPTUPLE/GETTUPLEFUNCTION/CALLPROGRAM/LINEAR/SOURCE/BINARY调度与编译产物
load/storeINDEXLOADSTORESHRINKkernel AST
数学ADD/MUL/...EXP2/LOG2/SIN/...WHEREWMMA(TensorCore)、MULACC全程
控制流RANGEENDIF/ENDIFBARRIERCONSTkernel AST
只在张量图CONTIGUOUSDETACHCOPYRESHAPE/PERMUTE/EXPAND/PAD/FLIPUNSHARDREDUCE/ALLREDUCE调度前
模式编译 IRPYLITERALUPat 编译(见 §5)

GroupOp(tinygrad/uop/__init__.py:106-137)再给出语义分组:ALUMovementCommutative(可交换,匹配时尝试子节点排列)、Reduce 等——重写规则大量使用这些集合,而不是逐个列枚举。


4. UPat:给 UOp 写的「正则表达式」

4.1 直觉

重写一条规则要回答两个问题:「什么样的节点要换」(模式)和「换成什么」(替换函数)。UPat(tinygrad/uop/ops.py:1337)就是模式的 DSL:

# 示意,非源码 —— 风格与 symbolic.py 里的真实规则一致
# 规则:「0 加任何数 = 那个数」
(UPat(Ops.ADD, src=(UPat.cvar("c"), UPat.var("x"))), lambda c, x: x if c.val == 0 else None)
  • UPat(Ops.ADD, src=(...)):匹配 op 为 ADD、且子结构匹配;
  • UPat.var("x") / UPat.cvar("c"):捕获任意节点 / 常量节点,按名字传给替换函数;
  • 返回 None 表示「这条规则不适用,继续试下一条」。

UPat.match(tinygrad/uop/ops.py:1424-1441)是回溯匹配器;对可交换算子(GroupOp.Commutative),UPat.__init__ 会把 src 展开成全部排列再试(tinygrad/uop/ops.py:1352-1353)。

4.2 PatternMatcher:规则的容器与索引

PatternMatcher(tinygrad/uop/ops.py:1474-1498)按「根的 op」把规则分桶(pdict),重写一个节点时只看同桶规则。两个提速设计:

  • early reject:每条模式预计算「子节点里必须出现哪些 op」(tinygrad/uop/ops.py:1363-1366),目标节点没有就直接跳过(tinygrad/uop/ops.py:1494-1496);
  • 模式编译成 Python:UPat 本身也是 UOp 表达的谓词树,upat_compile 把它编译成一段 Python 源码再 eval(tinygrad/uop/upat.py:9-12 的文件头注释说明了这件事),比逐层解释 match 快得多;UPAT_COMPILE=0 时退回解释执行(tinygrad/uop/ops.py:1467-1472)。

4.3 graph_rewrite:驱动器

引擎主体是 RewriteContext.unified_rewrite(tinygrad/uop/ops.py:1694-1760),graph_rewrite(tinygrad/uop/ops.py:1763-1765)是它的入口。工作流程:

sink(整张图)
│ 显式栈,自底向上走

① bottom_up(bpm)时:先在当前节点连跑到不动点(带环检测)
② 子节点都重写完后重建节点(src 变了才新建,interning 免费去重)
③ top_down(pm)时:对重建节点试一次规则


每个旧节点 → 新节点 记进 replace 表,返回替换后的根

要点:

  • 一次性重写整棵树,重复子树只处理一次(replace 表去重);
  • bottom_up=True 用于「先局部化简再向上传播」的场景(如 symbolic),默认先建后匹配;
  • FUNCTION/CALL 的函数体默认不进重写范围(tinygrad/uop/ops.py:530),需要时传 enter_calls=True;
  • 重写出环(规则来回踢皮球)会被栈上限和不动点环检测拦下报错(tinygrad/uop/ops.py:1699:1700)——规则集的终止性是规则作者的责任,引擎只负责报警

4.4 symbolic:最大的规则库

tinygrad/uop/symbolic.py 是全库最大的规则集,按强度分三档:

名字位置内容
symbolic_simpletinygrad/uop/symbolic.py:106常量折叠、恒等式、Invalid(死代码)传播
symbolictinygrad/uop/symbolic.py:234上一档 + 交换律重排、结合律、因式分解
symtinygrad/uop/symbolic.py:451上一档 + valid(边界条件)化简,调度和 codegen 里最常用

除法/取模的专门化简单独放在 tinygrad/uop/divandmod.py(div_and_mod_symbolic),它是索引算术化简的主力——reshape 合并、边界消除都靠它。


5. 原理演示:亲手写一个微型 graph_rewrite

# 示意,非源码 —— 最小版模式重写,体会 UPat + PatternMatcher 在做什么
rules = [
(("add", 0, "x"), lambda x: x), # add(0, x) -> x
(("add", "x", "x"), lambda x: ("mul", 2, x)), # add(x, x) -> mul(2, x)
]

def rewrite(node):
if not isinstance(node, tuple): return node
node = (node[0],) + tuple(rewrite(s) for s in node[1:]) # 先重写子树(自底向上)
for (pat_op, *pat_args), fxn in rules:
if node[0] == pat_op and all(p == a or isinstance(p, str) for p, a in zip(pat_args, node[1:])):
return fxn(*[a for p, a in zip(pat_args, node[1:]) if isinstance(p, str)])
return node

assert rewrite(("add", 0, ("add", "x", "x"))) == ("mul", 2, "x")

真实引擎多了 interning、op 分桶、环检测和进入/跳过 CALL 的控制,但骨架就是「自底向上重建 + 每节点试规则」。


6. 关键细节与坑

  • 不可变 = 随意共享。 UOp 从不原地修改,改就是 replace 造新节点(tinygrad/uop/ops.py:254-258);配合 interning,重建成本极低,但「想给节点挂点临时状态」就得用 tag 或外部 WeakKeyDictionary(如 buffersall_metadata,tinygrad/uop/ops.py:218-219)。
  • 递归属性要防 Python 递归深度。 _shape/_ranges 这类递归属性不用 cached_property,而用自定义的 recursive_property(tinygrad/uop/ops.py:222-231)——先 toposort 再自底向上填,深图不炸栈。
  • resolve 的默认值是 True。 判断一个符号化布尔表达式时拿不定主意就当成真(tinygrad/uop/ops.py:58-62),这影响 shape 检查等处的行为——读规则时看到 resolve(x, default=False) 要意识到作者在刻意反转。
  • SPEC 是运行时类型检查器。 打开 SPEC 环境变量后,每次造 UOp 都会跑 spec_full 校验(tinygrad/uop/ops.py:206-214),调度和 codegen 的关键节点也有 type_verify(如 tinygrad/schedule/rangeify.py:399-402)。开发 tinygrad 本体时它是一等公民。
  • VIZ 是调试重写的显微镜。 VIZ=1 时每次 graph_rewrite 都被记录(TrackedGraphRewrite,tinygrad/uop/ops.py:1510-1515),python -m tinygrad.viz.cli 能逐步回放任何一条规则把图改成了什么样。
  • 规则顺序有意义。 同一个 PatternMatcher 里靠前的规则先匹配;pm1 + pm2 是拼接(tinygrad/uop/ops.py:1489-1490),组合规则库时顺序就是语义。

7. 代码地图

主题文件路径符号名
节点定义tinygrad/uop/ops.pyUOpUOpMetaClassUOp.keyUOp.replace
op 枚举与分组tinygrad/uop/__init__.pyOpsGroupOp
shape/dtype 推导tinygrad/uop/ops.pyUOp._shapedtype_from_uop_broadcast_shape
遍历tinygrad/uop/ops.pyUOp.toposortUOp.topovisitbackward_slice
模式语言tinygrad/uop/ops.pyUPatUPat.varUPat.cvarUPat.match
规则容器tinygrad/uop/ops.pyPatternMatcherearly_reject
模式编译tinygrad/uop/upat.pyupat_compile
重写引擎tinygrad/uop/ops.pyRewriteContextgraph_rewriteunified_rewrite
化简规则库tinygrad/uop/symbolic.pytinygrad/uop/divandmod.pysymbolic_simplesymbolicsymdiv_and_mod_symbolic
运行时校验tinygrad/uop/spec.pytype_verifyspec_full