tinygrad 请求链路中间件开发:从 UOp 图到 PatternMatcher 与 TinyJit 完整指南
【免费下载链接】tinygradYou like pytorch? You like micrograd? You love tinygrad! ❤️项目地址: https://gitcode.com/GitHub_Trending/tiny/tinygrad
tinygrad 的架构像一个洋葱模型:Tensor 前端、UOp 图、调度器、lowering 引擎、runtime,每一层都能被拦截和替换。想在"请求"经过框架时插入自定义逻辑?调度器里可以挂 PatternMatcher 规则,JIT 可以缓存整个执行子图,UOp 改写也可以挂任意 pattern。这套机制就是 tinygrad 的"中间件层",本文带你走完整个链路。
🧅 为什么说 tinygrad 天生适合插中间件
前面提到 docs/developer/developer.md 把框架拆成四块:前端、调度器、lowering 引擎、运行时。前端只是 UOp 图的语法糖,Tensor 操作会被拆成一个个小节点;调度器把大图切成一个个 kernel 调用;lowering 把每个 kernel 编译成可执行代码;runtime 负责把代码派发到具体设备。
这套结构意味着每一段代码都是"可拦截的":前端到图之间有改写阶段,图到 kernel 之间有调度阶段,kernel 到设备之间有 runtime 阶段。其他常见框架的图叠在 cuDNN、cuBLAS 等固定 kernel 库上,你只能选;tinygrad 直接生成 kernel,任何一段计算都可以被你的代码改写成别的形状。
🕸️ UOp 图:理解 tinygrad 请求的载体
每个节点是一个四元组(op, dtype, src, arg),定义见 tinygrad/uop/:
m = UOp(Ops.MUL, dtypes.float, src=(a, b), arg=None)op决定这个节点做什么src是输入节点的列表,指向上游arg携带参数,例如 kernel 名、常量值
节点分 base 和 view 两类:base 真正占用 buffer,view 只是切出来的视图,不额外占内存。这张图不可变,所以"改写"就等于"匹配一个子图,替换成另一个子图"——这正是中间件最擅长干的事。
🎯 PatternMatcher:tinygrad 的核心中间件
PatternMatcher 定义在 tinygrad/uop/ops.py,本质是一张规则表。每条规则三件事:一个要 match 的子图、一个用来 replace 的子图、可选的 guard 条件。框架反复扫全图,每匹配一次就替换一次,直到没有变化为止。
from tinygrad.uop.ops import PatternMatcher, UOp, Ops, dtypes pm = PatternMatcher([ (UOp(Ops.MUL, src=(UOp.variable("a"), UOp.const(1.0))), lambda ctx, a: UOp(Ops.ADD, src=(a, a))), ]) new_graph = graph_rewrite(sink, pm, name="my_middleware")上面这条规则做的事情:a * 1.0改写成a + a。你可以把自定义规则塞进调度阶段或 lowering 阶段,就像在中间件链里加一个环节;框架原有的 beam 搜索优化、内存规划本身也是同一套 PatternMatcher 规则。
⚡ TinyJit:把整个请求打包成一次回放
TinyJit 是请求级中间件:包一层函数,让它跑三次就稳定下来。
from tinygrad.engine.jit import TinyJit @TinyJit def train_step(x, target): loss = model(x).sparse_categorical_crossentropy(target) return loss.backward()第一次调用走原函数,当作热身;第二次捕获这次执行产生的全部 kernel 调用,合并成一个大 LINEAR,交给调度器重新做内存规划、再编译,支持 graph 的设备还会把多个 kernel 打包成一次 graph launch;从第三次开始就不再走 Python 前端了,只替换输入 buffer,直接执行已经编译好的东西。JIT=0关闭,JIT=2升级到更激进的 graph 批量。
🔍 用 VIZ 与 DEBUG 观察请求被谁改写了
打开两个环境变量就能看到整条链路的执行轨迹:
VIZ=1打开可视化窗口,每一次graph_rewrite的中间快照都能回放,PatternMatcher 支持 trace 参数追踪每条规则命中几次DEBUG=2在终端打印调度器的拆分结果和 JIT 捕获了多少 kernel
跑一遍 test/test_tiny.py 或 examples/beautiful_mnist.py,你能完整看到一次"请求"在每个中间件阶段经历了什么。
⚠️ 新手最容易踩的四个坑
- JIT 输入必须是真实 buffer。传视图或切片进去会抛
JIT inputs must be real buffers; use .clone(),先.clone()再传。 - 同一个张量不能当两个输入。会直接抛
duplicate inputs to JIT。 - TinyJit 里不能再套 TinyJit。捕获阶段会抛
RuntimeError,嵌套结构要拆层。 - PatternMatcher 规则要写成幂等的,否则会反复触发,框架有迭代上限保护,但你的规则最好在
replace后不再匹配自己。
这四个坑都有对应的异常类型和测试用例,报错信息通常能直接指向问题,建议读一遍 test/test_uops.py 和 test/null/test_pattern_matcher.py。
下一步做什么
三件事按顺序做就行:
- 用
DEBUG=2跑一遍 examples/beautiful_mnist.py,先看清每个请求在哪一层被拦截。 - 打开
VIZ=1,在浏览器里逐帧看 UOp 图的改写过程,找到"原来这条规则在这里生效"的直观感受。 - 挑一个你感兴趣的优化阶段(tinygrad/codegen/ 里按 simplify、opt、decomp、late 分层),写一条自己的 PatternMatcher 规则,跑通现有测试。
做到第三步,tinygrad 就从"你用的框架"变成"你能改的框架"。
【免费下载链接】tinygradYou like pytorch? You like micrograd? You love tinygrad! ❤️项目地址: https://gitcode.com/GitHub_Trending/tiny/tinygrad
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考