跳到主要内容

数据截至 (上游 commit 118e812b316f)

05 · 设备、运行时与 TinyJit

这一章讲什么: 编译产物落地的最后一段——设备怎么抽象、显存怎么管、kernel 怎么发射,以及 @TinyJit 怎么把「调度+编译」的固定开销从每步训练里抹掉。


1. 它要解决的小问题

到第 4 章结束,我们手上是一个个编译好的 PROGRAM(源码+二进制)。剩下三件事:

  1. 设备差异:Metal 用 ObjC API、CUDA 用驱动 API、AMD/NV 干脆自己写用户态驱动——怎么用一个界面装下它们;
  2. 显存管理:中间 buffer 每秒成百上千次分配释放,不能每次都真找驱动要;
  3. 重复开销:训练每步的图都一样,调度+编译不能每步重来——这就是 JIT 存在的理由。

2. 设备抽象四层

2.1 一张分层图

_Device(单例,按名注册/懒加载)
│ Device["METAL"] → import tinygrad.runtime.ops_metal

Compiled(一个物理设备)
│ 捆四样:Allocator(内存) / renderers(代码生成)
│ runtime(Program 加载发射) / graph(图执行,可空)

Buffer ──► Allocator(LRU 缓存真实分配)

Program(编译好的 kernel,__call__ 即发射)

2.2 _Device:按名懒加载

Device 是全局单例 _Device(tinygrad/device.py:15-56)。Device["METAL"] 第一次被取下标时,按名字动态 importlib.import_module("tinygrad.runtime.ops_metal") 并找到 MetalDevice 类实例化(tinygrad/device.py:27-37)——没装 CUDA 的机器不会因为 import tinygrad 而炸。默认设备按 ALL_DEVICES(tinygrad/device.py:14)顺序探测第一个能打开的(_select_device,tinygrad/device.py:47-55)。

2.3 Compiled:设备对象

Compiled.__init__(tinygrad/device.py:352-356)把四样东西捆在一起,以后端 MetalDevice 为例(tinygrad/runtime/ops_metal.py:31-50):

成员Metal 实例职责
allocatorMetalAllocator显存分配/释放/拷贝
renderers[MetalRenderer]第 4 章的渲染器(可多选)
runtime_tMetalProgram把 BINARY 加载成可发射的 Program
graphMetalGraph设备图执行(虚拟化 Metal 上置 None,tinygrad/runtime/ops_metal.py:50)

2.4 Buffer 与 LRUAllocator

Buffer(tinygrad/device.py:102-217)的逻辑要点:

  • 惰性分配:构造只记元数据,ensure_allocated/allocate 才真分配(tinygrad/device.py:143-162),并累计 GlobalCounters.mem_used;
  • 视图零拷贝:Buffer.view(tinygrad/device.py:215-217)造一个 base 相同、带 offset 的新 Buffer——这就是调度器 arena 子分配和 SHRINK 的物理基础;
  • 拷贝统一走 LINEAR:copy_from 也包成一个单 CALL 的 LINEAR 扔给执行引擎(tinygrad/device.py:207-214),没有第二条拷贝路径。

LRUAllocator(tinygrad/device.py:250-270)在真实分配器外面加一层 (size, options) → 空闲块 缓存:free 不还真,alloc 先查缓存;真不够了(MemoryError)才 free_cache 清仓重试(tinygrad/device.py:260-263)。训练循环里同尺寸 buffer 的分配因此接近零成本。

2.5 后端家族速览

tinygrad/runtime/ops_*.py 一个文件一个后端,成熟度分几档:

后端特点
自研用户态驱动ops_nv.pyops_amd.pyHCQ 路线:直接 ioctl 建队列/填命令,不经过官方运行时;docs/developer/hc.md 专述
官方 APIops_metal.pyops_cuda.pyops_hip.pyops_cl.pyops_webgpu.pyops_qcom.pyops_dsp.py经厂商 API;编译各自用厂商工具链或内嵌 clang
CPUops_cpu.pyclang/LLVM jit;NUM_CPU_THREADS 控制线程
参照/虚拟ops_python.py(纯 Python 解释执行 IR)、ops_null.pyops_disk.py(mmap 文件即 buffer)、ops_npy.py测试与权重加载用

3. 执行引擎:run_linear

run_linear(tinygrad/engine/realize.py:320-324)三步:

  1. compile_linear(tinygrad/engine/realize.py:310-316):可选 CPU 交叉验证、应用 BEAM、lower_and_compile(见下)、optimize_local_size;
  2. lower_and_compile(tinygrad/engine/realize.py:268-293):收集 LINEAR 里还没编译的 kernel,按缓存键去重,用进程池并行编译(get_worker_pool,tinygrad/engine/worker.py——worker 禁设备、禁 Ctrl-C 的 SpawnProcess),编好把 PROGRAM 换进 CALL;
  3. pm_exec 逐 CALL 派发(tinygrad/engine/realize.py:299-306)。

派发按 CALL 的 ast 类型分派,这是「一切皆 UOp」在执行层的体现:

ast执行器干什么
PROGRAMexec_kernel(tinygrad/engine/realize.py:186-196)绑 buffer、算发射维度、runtime(...) 发射
COPYexec_copy(tinygrad/engine/realize.py:174-184)优先 _transfer(同厂设备直连)/copy_from_disk,退回经 host 倒手
CUSTOM_FUNCTION "graph"exec_graph设备图整体重放
CUSTOM_FUNCTION "encdec"exec_encdec视频编解码(HEVC)
CUSTOM_FUNCTION "hcq" / "validate"exec_hcq / exec_validateHCQ 命令批 / CPU 交叉验证

多设备(MultiBuffer)时 unwrap_multi(tinygrad/engine/realize.py:165-172)逐设备拆开发射,设备号以 _device_num 注入 var_vals。


4. TinyJit:捕获一次,重放千次

4.1 三段状态机

_TinyJit.__call__(tinygrad/engine/jit.py:235-281)用计数器 cnt 分三段:

cnt行为
0原样跑(并收集输入签名),相当于 baseline
1捕获:capturing 列表挂上自己(tinygrad/engine/jit.py:244-252),这次执行里所有 create_linear_with_vars 只记录 LINEAR 不执行(tinygrad/schedule/__init__.py:204-207)
≥2重放:核对输入签名后,直接跑 CapturedJit

4.2 捕获时做了什么

捕获到的所有 LINEAR 拼成一个大的,然后 jit_lower(tinygrad/engine/jit.py:65-75):

  1. 输入 buffer 全部换成 PARAM(tinygrad/engine/jit.py:70)——重放时按槽位喂新 buffer;
  2. 全量 memory_plan_rewrite——JIT 下的内存规划比 eager 更激进,因为整个调用序列已知;
  3. compile_linear 编译全部 kernel;
  4. graph_split_rewrite(tinygrad/engine/jit.py:32-59):把连续可图的 kernel 段打包成 CUSTOM_FUNCTION "graph"(如 MetalGraph/CUDA Graph),一次 API 调用提交几十上百个 kernel。

产物 CapturedJit(tinygrad/engine/jit.py:163-191)持有:参数化 LINEAR、输入签名(形状变量、dtype、device)、返回结构。重放时 _copy_input 会复制「会被写」的输入,避免污染调用方的 buffer(tinygrad/engine/jit.py:178-182)。

4.3 原理演示

# 示意,非源码 —— TinyJit 的用户体验
@TinyJit
def train_step(x, y):
return model(x).sparse_categorical_crossentropy(y).backward().realize()

train_step(x, y) # 第 1 次:正常跑,记录输入签名
train_step(x, y) # 第 2 次:捕获全部 kernel,编译,连成图
train_step(x, y) # 第 3 次起:换数据(同形状),直接重放,几乎零 Python 开销

4.4 约束与坑(全是显式检查)

  • 捕获期禁止碰数据:Tensor._buffer 在 capturing 时直接抛 JitError(tinygrad/tensor.py:462-466)——值会被烘进图里;
  • 输入必须是已 realize 的真 buffer 且不能重复:_prepare_jit_inputs(tinygrad/engine/jit.py:201-212)两处显式 raise;
  • 返回值必须是 Tensor(或容器):_check_no_non_tensor_return(tinygrad/engine/jit.py:80-85);
  • 签名变就报错:形状变量集合/dtype/device 变了会 JitError(tinygrad/engine/jit.py:279-281)——动态 shape 要靠 Variable 显式声明,靠蒙是蒙不过的;
  • 嵌套 JIT 不支持(tinygrad/engine/jit.py:245-246)。

5. 多设备:shard 与 allreduce

Tensor.shard(devices, axis)(tinygrad/tensor.py:557-571)把张量标成多设备:图里出现 UNSHARD/MSTACK/MSELECT 节点。调度前 multi_pm(tinygrad/schedule/multi.py:284)把它们展开成逐设备的子图(广播复制、按 shard 切 movement op、跨设备拷贝显式化),之后走同一套 rangeify;梯度侧的 shard 归约由 pm_gradientUNSHARD 规则处理(tinygrad/mixin/gradient.py:82)。数据并行要的梯度 allreduce 由 tinygrad/schedule/allreduce.py 生成专门的归约函数。

examples/beautiful_mnist_multigpu.py 是这条路线的最小完整示例。


6. 关键细节与坑(全章汇总)

  • ALLOW_DEVICE_USAGE 是全局安全闸。 建图期(@function 内、worker 进程、BEAM 子任务)设备访问被禁,Device.__getitem__ 直接 assert(tinygrad/device.py:25);DISK/NPY/PYTHON 永远放行——它们不是真设备。
  • Buffer 生命周期挂在 UOp 上。 UOp.__del__ 给 BUFFER 节点减引用(tinygrad/uop/ops.py:245-249),Python GC 顺序因此直接决定显存释放时机;jit.py/tensor.py 里散落的 disable_gc/suppress_finalizing 是在跟这件事搏斗。
  • 运行时也有两级缓存:编译产物 runtime_cache(tinygrad/engine/realize.py:131-136)+ 磁盘级 compile_cached(tinygrad/device.py:307-312),重启进程也不用重编译。
  • DISK tensor 的写是特例。 assign 对 DISK 走 copy_from 直写,绕过调度(tinygrad/tensor.py:446-449 的 TODO 注释自认是 hack)。
  • HCQ2 是进行中的下一代执行路径。 getenv("HCQ2") 时才启用(tinygrad/engine/realize.py:308:315);本 commit 默认关,主线仍是逐 CALL 派发 (inferred:从开关位置推断)。

7. 代码地图

主题文件路径符号名
设备注册/选择tinygrad/device.py_DeviceDevice.__getitem___select_deviceALL_DEVICES
设备对象tinygrad/device.pyCompiledCompiled._select_renderer
显存tinygrad/device.pyBufferBuffer.viewAllocatorLRUAllocator
编译器外壳tinygrad/device.pyCompiler.compile_cached
执行引擎tinygrad/engine/realize.pyrun_linearcompile_linearlower_and_compilepm_execexec_kernelexec_copy
并行编译tinygrad/engine/worker.pyget_worker_pool
JITtinygrad/engine/jit.pyTinyJit_TinyJit.__call__CapturedJitjit_lowerprune_linear
设备图tinygrad/engine/jit.pytinygrad/runtime/graph/graph_split_rewriteGraphRunner
多设备tinygrad/tensor.pytinygrad/schedule/multi.pyTensor.shardmulti_pm
后端实例tinygrad/runtime/ops_metal.pyMetalDeviceMetalAllocatorMetalProgram