数据截至 (上游 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)。按用途分群(注释里的分组即源码里的分组):
| 群 | 代表成员 | 出现阶段 |
|---|---|---|
| 定义/特殊 | BUFFER、PARAM、SPECIAL(GPU 维度)、VARIABLE 系 | 贯穿 |
| 非渲染节点 | NOOP、SINK、AFTER、GROUP、TUPLE/GETTUPLE、FUNCTION/CALL、PROGRAM/LINEAR/SOURCE/BINARY | 调度与编译产物 |
| load/store | INDEX、LOAD、STORE、SHRINK | kernel AST |
| 数学 | ADD/MUL/...、EXP2/LOG2/SIN/...、WHERE、WMMA(TensorCore)、MULACC | 全程 |
| 控制流 | RANGE、END、IF/ENDIF、BARRIER、CONST | kernel AST |
| 只在张量图 | CONTIGUOUS、DETACH、COPY、RESHAPE/PERMUTE/EXPAND/PAD/FLIP、UNSHARD、REDUCE/ALLREDUCE | 调度前 |
| 模式编译 IR | PYLITERAL | UPat 编译(见 §5) |
GroupOp(tinygrad/uop/__init__.py:106-137)再给出语义分组:ALU、Movement、Commutative(可交换,匹配时尝试子节点排列)、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_simple | tinygrad/uop/symbolic.py:106 | 常量折叠、恒等式、Invalid(死代码)传播 |
symbolic | tinygrad/uop/symbolic.py:234 | 上一档 + 交换律重排、结合律、因式分解 |
sym | tinygrad/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(如buffers、all_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.py | UOp、UOpMetaClass、UOp.key、UOp.replace |
| op 枚举与分组 | tinygrad/uop/__init__.py | Ops、GroupOp |
| shape/dtype 推导 | tinygrad/uop/ops.py | UOp._shape、dtype_from_uop、_broadcast_shape |
| 遍历 | tinygrad/uop/ops.py | UOp.toposort、UOp.topovisit、backward_slice |
| 模式语言 | tinygrad/uop/ops.py | UPat、UPat.var、UPat.cvar、UPat.match |
| 规则容器 | tinygrad/uop/ops.py | PatternMatcher、early_reject |
| 模式编译 | tinygrad/uop/upat.py | upat_compile |
| 重写引擎 | tinygrad/uop/ops.py | RewriteContext、graph_rewrite、unified_rewrite |
| 化简规则库 | tinygrad/uop/symbolic.py、tinygrad/uop/divandmod.py | symbolic_simple、symbolic、sym、div_and_mod_symbolic |
| 运行时校验 | tinygrad/uop/spec.py | type_verify、spec_full |