pytorch后端解读¶
目标:理解一个 op 从 Python 调用到后端 kernel 的完整过程,以及
eager/Graph两种执行模式下的 stream 与同步语义。
从简单计算到算子¶
- Python 层:
torch.add(a, b)走到Tensor方法或 ATen 函数绑定,实质是进入c10::Dispatcher。 - schema:每个 op 都有声明式 schema(如
aten::add.Tensor(Tensor self, Tensor other, Scalar alpha=1) -> Tensor),它约定了重载解析、dtype 提升等契约。 - dispatcher:按 dispatch key 选中实现;
at::native里是参考实现,后端只需覆盖自己关心的 op,其余交给 fallback。 - OpPreparation:多数 op 不自己造输出,而是先算出输出的 shape / dtype / device(借助
TensorIterator、类型提升规则等),再用at::empty之类完成分配。 - structured kernel:把“元信息推理”与“计算”分离。后端只实现计算部分,元信息复用通用逻辑;也可注册 meta 函数来支持
FakeTensor/ meta 设备推理。 - 后端最先要打通的三个基础 op:
aten::empty、aten::empty_strided、aten::_copy_from——其余算子大多建立在它们之上。
Text Only
torch.add(a, b)
-> Dispatcher(按 dispatch key 选实现)
-> OpPreparation:计算输出 shape / dtype / device
-> at::empty -> 后端 allocator 分配
-> 后端 kernel 写入输出
TODO
- 挑一个有代表性的算子(如
add或softmax),把 schema、meta 推理、后端实现逐层拆开记录。 - 记录 dtype 提升与隐式类型转换发生在哪一层。
'Stream' 'Event'机制实现¶
c10::Stream:轻量句柄,内含StreamId+DeviceIndex+ 设备类型,不持有资源,可自由拷贝。c10::Event:跨 stream 的完成标记,带EventFlag与状态查询 / 同步接口(如synchronize、query)。- 当前 stream 的查询与设置:CPU / CUDA 各有
getCurrentCUDAStream()这类接口;自定义后端必须提供等价实现(社区常见命名形如getCurrentPrivateUse1Stream,以所用版本为准),否则 guard 和 kernel 都不知道该往哪条流排队。 - guard:
c10::DeviceGuard/c10::OptionalDeviceGuard/impl::InlineDeviceGuard在作用域内切换“当前设备 + 当前 stream”,退出时还原。这是后端必须支持的接口,否则大量框架代码无法工作。 - 延迟释放:
recordStream把内存与 stream 绑定,保证该 stream 上后续工作完成前内存不被复用;跨 stream 使用时需显式 event 同步。 - 典型同步点:H2D 拷贝后要让计算流等待拷贝流;D2H 读回结果前要等待计算完成。同步粒度选错是“结果偶发错误”的高频原因。
TODO
- 整理“当前 stream / 默认 stream / 自定义 stream”三者的行为差异。
- 记录一次跨 stream 同步缺失导致错误结果的实际案例。
'eager' 'Graph'机制实现¶
- eager 模式:Python 调用即执行,autograd 在运行时动态记录反向图。灵活,但 Python 开销大、kernel launch 密集。
- 图捕获的前置条件:后端需要支持 stream 语义、allocator 的
recordStream、以及可被捕获的内存池(graph / private mempool),否则捕获会失败或产生悬垂指针。 torch.compile链路(概念层):- Dynamo:把 Python 字节码转成 FX graph;遇到无法追踪的代码就产生 graph break,把图拆成多段。
- AOTAutograd:为前向 / 反向生成带 autograd 信息的图,并做图变换。
- Inductor:把图降级为可执行代码 / kernel(默认生成 Triton 或 C++)。
FakeTensor/ meta device:在“只跑元信息、不碰真实内存”的 fake 模式下完成 shape / dtype 推理与图变换,是编译链路能跑通的关键支撑。- 后端接入点:提供自定义 compiler backend 并注册,把 Inductor 产出或自己的图编译结果接到设备 runtime。
TODO
- 补
torch.compile(backend=...)的最小自定义后端示例。 - 记录 graph break 的实际触发点,以及它如何影响后端的编译粒度。