TensorPlay AI
打开导航
博客

TPX:计算与梯度如何解耦

自动微分不必把梯度状态塞进每一次计算。本文沿着一次前向与反向传播,说明 TPX 如何记录、排序、执行并验证梯度。

先回答:为什么把自动微分放在 TPX?

P10 的职责是计算张量值,TPX 的职责是记录值之间的依赖并计算梯度。如果把 requires_grad、grad 和反向节点直接塞进 P10,每一次纯数值计算都会携带训练状态,硬件和计算层也会被迫理解微分。

TPX 选择组合 p10::Tensor:需要梯度时使用 tpx::Tensor 包住计算核心,不需要梯度时直接调用 P10。这个边界让学习者可以分别阅读“数值如何算”和“梯度如何来”。

  • P10:创建张量、执行算子、选择后端 kernel。
  • TPX:保存微分状态、构建 DAG、调度 GradFn。
  • 两层之间:通过张量值和算子调用连接,而不是互相污染实现。

按需追踪:只为需要的计算建立图

当 requires_grad 为 true 时,TPX 记录输入、输出、GradFn 与算子参数,并把每个结果连接到产生它的节点。requires_grad 为 false 时,TPX 不创建这段反向状态,前向路径可以保持接近纯 P10。

这不是简单地开关一个布尔值。它决定了哪些 Tensor 需要保存上下文、哪些节点会进入 DAG,以及 backward() 最终能够沿哪条路径回溯。

x = tp.ones((2, 2), requires_grad=True)
y = x.matmul(w).relu()

TPX records: matmul -> relu -> y
P10 executes: tensor kernels for matmul and relu

前向计算与反向调度分离

前向过程中,TPX 只需要把算子关系记录下来;调用 backward() 后,它先从目标节点开始做拓扑排序,再按反向顺序调度每个节点的 GradFn。

GradFn 决定某个算子的局部梯度规则,但真正的张量加法、乘法和累加仍由 P10 执行。于是自动微分可以替换规则,计算后端也可以替换 kernel,两者不需要互相复制。

loss.backward()
  -> topologically sort the recorded DAG
  -> call each node's GradFn
  -> accumulate gradients with P10 operations
  -> expose x.grad and parameter gradients

用一条具体链路检查梯度

以 y = x * w 为例,前向节点保存 x、w 和输出 y;反向时,乘法的 GradFn 根据上游梯度分别计算 dy/dx = w 与 dy/dw = x,再用 P10 完成对应的张量运算。

因为局部规则是显式节点,数值错误可以被定位到某一个 GradFn、某一次累加或某个后端 kernel,而不是只能从最终 loss 猜原因。

  • 检查每个节点保存的输入、输出和算子参数。
  • 检查反向拓扑顺序是否符合依赖关系。
  • 用有限差分或小规模手算结果校验局部梯度。
  • 分别测试启用与关闭追踪时的前向数值。

扩展边界:新规则、新硬件与新实验

自定义 GradFn 不需要修改 P10 的张量核心;新硬件也不必重新实现整个自动微分系统。研究者可以先替换一个局部梯度规则,再单独观察它对前向、反向和内存的影响。

这正是 TPX 的工程价值:它不承诺自动解决所有训练性能问题,而是把微分过程变成可以阅读、验证和重新组合的实验对象。

Ask DeepWiki