从零构建自动微分引擎:对偶数、计算图与反向传播的 Rust 深度实战
> 自动微分(Automatic Differentiation, AD)是现代机器学习的基石——从 PyTorch 到 JAX,从 TensorFlow 到 Hugging Face,所有深度学习框架的核心都依赖它。本文不借助任何现成 AD 库,从零实现一个功能完整的自动微分引擎,深入理解对偶数(Dual Numbers)、计算图构建、反向模式微分(Backpropagation)以及高阶导数的实现原理。
一、为什么需要自动微分
在深度学习训练中,梯度是优化器的燃料。计算梯度的方法有三种:
- 数值微分(Numerical Differentiation):有限差分法,简单但精度低、计算量大,$O(n)$ 次函数求值计算 n 维梯度。
- 符号微分(Symbolical Differentiation):对数学表达式求解析导数,精确但面临"表达式膨胀"问题。
- 自动微分(Automatic Differentiation):精确、高效,$O(1)$ 次函数求值计算梯度,是现代框架的唯一选择。
自动微分不是数值近似,也不是符号推导。它通过链式法则系统地追踪每个基本操作的导数,精确地在任意可微点计算导数。
二、数学基础:对偶数与微分
对偶数(Dual Numbers)是自动微分的数学根基。形式为 $a + b\epsilon$,其中 $\epsilon^2 = 0$(类似于虚数单位 $i^2 = -1$,但对偶数的平方为零)。
对偶数的核心性质:若 $f$ 在 $a$ 处可微,则 $$f(a + b\epsilon) = f(a) + f'(a) \cdot b\epsilon$$
将 $b=1$ 代入,即得 $f(a + \epsilon) = f(a) + f'(a)\epsilon$。这意味着:函数在对偶数上的计算结果天然包含导数值。
// 对偶数定义
#[derive(Debug, Clone, Copy)]
struct Dual {
val: f64, // 实部:函数值
deriv: f64, // 对偶部:导数值
}
impl Dual {
fn new(val: f64, deriv: f64) -> Self {
Self { val, deriv }
}
fn constant(val: f64) -> Self {
Self { val, deriv: 0.0 }
}
fn variable(val: f64) -> Self {
Self { val, deriv: 1.0 } // 自变量的导数为 1
}
}
实现基本运算,导数部分遵循链式法则:
use std::ops::{Add, Sub, Mul, Div, Neg};
impl Add for Dual {
type Output = Self;
fn add(self, rhs: Self) -> Self::Output {
Self {
val: self.val + rhs.val,
deriv: self.deriv + rhs.deriv, // (f+g)' = f' + g'
}
}
}
impl Sub for Dual {
type Output = Self;
fn sub(self, rhs: Self) -> Self::Output {
Self {
val: self.val - rhs.val,
deriv: self.deriv - rhs.deriv,
}
}
}
impl Mul for Dual {
type Output = Self;
fn mul(self, rhs: Self) -> Self::Output {
Self {
val: self.val * rhs.val,
// 乘法法则: (fg)' = f'g + fg'
deriv: self.deriv * rhs.val + self.val * rhs.deriv,
}
}
}
impl Div for Dual {
type Output = Self;
fn div(self, rhs: Self) -> Self::Output {
Self {
val: self.val / rhs.val,
// 除法法则: (f/g)' = (f'g - fg') / g²
deriv: (self.deriv * rhs.val - self.val * rhs.deriv) / (rhs.val * rhs.val),
}
}
}
impl Neg for Dual {
type Output = Self;
fn neg(self) -> Self::Output {
Self {
val: -self.val,
deriv: -self.deriv,
}
}
}
指数函数、三角函数和对数函数的对偶数扩展:
impl Dual {
fn exp(self) -> Self {
let val = self.val.exp();
Self {
val,
deriv: val * self.deriv, // d/dx e^x = e^x, 所以导数 = e^x * x'
}
}
fn ln(self) -> Self {
Self {
val: self.val.ln(),
deriv: self.deriv / self.val, // d/dx ln(x) = 1/x
}
}
fn sin(self) -> Self {
Self {
val: self.val.sin(),
deriv: self.val.cos() * self.deriv, // d/dx sin(x) = cos(x)
}
}
fn cos(self) -> Self {
Self {
val: self.val.cos(),
deriv: -self.val.sin() * self.deriv, // d/dx cos(x) = -sin(x)
}
}
fn powf(self, n: f64) -> Self {
Self {
val: self.val.powf(n),
deriv: n * self.val.powf(n - 1.0) * self.deriv,
}
}
fn sqrt(self) -> Self {
let val = self.val.sqrt();
Self {
val,
deriv: self.deriv / (2.0 * val),
}
}
fn relu(self) -> Self {
Self {
val: if self.val > 0.0 { self.val } else { 0.0 },
deriv: if self.val > 0.0 { self.deriv } else { 0.0 },
}
}
fn sigmoid(self) -> Self {
let val = 1.0 / (1.0 + (-self.val).exp());
Self {
val,
deriv: val * (1.0 - val) * self.deriv, // σ'(x) = σ(x)(1-σ(x))
}
}
fn tanh(self) -> Self {
let val = self.val.tanh();
Self {
val,
deriv: (1.0 - val * val) * self.deriv,
}
}
}
测试我们的对偶数引擎:
fn main() {
// 计算 f(x) = x² + 3x 在 x=2 处的值与导数
let x = Dual::variable(2.0);
let f = x * x + x * 3.0;
println!("f(2.0) = {}", f.val); // 输出: 10.0
println!("f'(2.0) = {}", f.deriv); // 输出: 7.0
// 验证: f'(x) = 2x + 3 = 2*2 + 3 = 7 ✓
// 测试复杂函数 g(x) = sin(x²) 在 x=1.5
let x = Dual::variable(1.5);
let g = (x * x).sin();
println!("g(1.5) = {:.6}", g.val); // sin(2.25) ≈ 0.778073
println!("g'(1.5) = {:.6}", g.deriv); // cos(2.25)*3.0 ≈ -2.05756
}
三、计算图与前向模式 AD
对偶数方法优雅但有一个局限:当函数有多个输入时($f: \mathbb{R}^n \to \mathbb{R}$),需要对每个输入独立设置扰动,计算 n 次前向传播才能得到完整梯度。这在高维场景下效率低下。
为此引入计算图(Computation Graph),配合反向模式自动微分(Reverse-Mode AD),这正是现代深度学习框架的核心机制。
3.1 计算图构建
将计算视为一个有向无环图(DAG):叶子节点是输入变量或常量,中间节点是基本运算,输出节点是最终结果。每个节点存储:
- 本地值(前向传播时计算)
- 本地偏导数(反向传播时使用)
// 计算图中的节点 ID
type NodeId = usize;
// 操作类型
#[derive(Clone)]
enum Op {
// 一元运算
Neg, Exp, Ln, Sin, Cos, Tanh, Sigmoid, Relu, Sqrt,
// 二元运算
Add, Sub, Mul, Div, Pow(f64),
// Reward/Indicator
Constant(f64),
}
// 节点存储
struct Node {
op: Op,
inputs: Vec<NodeId>, // 父节点
output: f64, // 前向传播计算的值
grad: f64, // 反向传播的梯度
}
// 计算图
struct ComputationGraph {
nodes: Vec<Node>,
}
impl ComputationGraph {
fn new() -> Self {
Self { nodes: Vec::new() }
}
fn add_node(&mut self, op: Op, inputs: Vec<NodeId>) -> NodeId {
let id = self.nodes.len();
self.nodes.push(Node {
op,
inputs,
output: 0.0,
grad: 0.0,
});
id
}
// 创建常量节点
fn constant(&mut self, val: f64) -> NodeId {
self.add_node(Op::Constant(val), vec![])
}
// 创建变量节点(占位符,前向时填充值)
fn variable(&mut self, val: f64) -> NodeId {
self.add_node(Op::Constant(val), vec![])
}
// 一元运算
fn unary_op(&mut self, op: Op, input: NodeId) -> NodeId {
self.add_node(op, vec![input])
}
// 二元运算
fn binary_op(&mut self, op: Op, left: NodeId, right: NodeId) -> NodeId {
self.add_node(op, vec![left, right])
}
// 便捷方法
fn add(&mut self, a: NodeId, b: NodeId) -> NodeId {
self.binary_op(Op::Add, a, b)
}
fn sub(&mut self, a: NodeId, b: NodeId) -> NodeId {
self.binary_op(Op::Sub, a, b)
}
fn mul(&mut self, a: NodeId, b: NodeId) -> NodeId {
self.binary_op(Op::Mul, a, b)
}
fn div(&mut self, a: NodeId, b: NodeId) -> NodeId {
self.binary_op(Op::Div, a, b)
}
fn neg(&mut self, a: NodeId) -> NodeId {
self.unary_op(Op::Neg, a)
}
fn sin(&mut self, a: NodeId) -> NodeId {
self.unary_op(Op::Sin, a)
}
fn cos(&mut self, a: NodeId) -> NodeId {
self.unary_op(Op::Cos, a)
}
fn exp(&mut self, a: NodeId) -> NodeId {
self.unary_op(Op::Exp, a)
}
fn sigmoid(&mut self, a: NodeId) -> NodeId {
self.unary_op(Op::Sigmoid, a)
}
fn tanh(&mut self, a: NodeId) -> NodeId {
self.unary_op(Op::Tanh, a)
}
fn relu(&mut self, a: NodeId) -> NodeId {
self.unary_op(Op::Relu, a)
}
fn ln(&mut self, a: NodeId) -> NodeId {
self.unary_op(Op::Ln, a)
}
fn sqrt(&mut self, a: NodeId) -> NodeId {
self.unary_op(Op::Sqrt, a)
}
fn powf(&mut self, a: NodeId, n: f64) -> NodeId {
self.binary_op(Op::Pow(n), a, self.constant(n))
}
}
3.2 前向传播
impl ComputationGraph {
fn forward(&mut self, output_id: NodeId) -> f64 {
// 递归计算所有依赖节点
self.eval_node(output_id);
self.nodes[output_id].output
}
fn eval_node(&mut self, id: NodeId) -> f64 {
let node = &self.nodes[id];
// 如果已经计算过(此处简化为不重复计算,实际需标记机制)
if node.output != 0.0 || matches!(node.op, Op::Constant(0.0)) {
// 更完善的实现需要 visited 标记
}
let inputs: Vec<f64> = node.inputs.iter().map(|&i| self.eval_node(i)).collect();
let result = match &node.op {
Op::Constant(v) => *v,
Op::Neg => -inputs[0],
Op::Add => inputs[0] + inputs[1],
Op::Sub => inputs[0] - inputs[1],
Op::Mul => inputs[0] * inputs[1],
Op::Div => inputs[0] / inputs[1],
Op::Exp => inputs[0].exp(),
Op::Ln => inputs[0].ln(),
Op::Sin => inputs[0].sin(),
Op::Cos => inputs[0].cos(),
Op::Tanh => inputs[0].tanh(),
Op::Sigmoid => 1.0 / (1.0 + (-inputs[0]).exp()),
Op::Relu => if inputs[0] > 0.0 { inputs[0] } else { 0.0 },
Op::Sqrt => inputs[0].sqrt(),
Op::Pow(n) => inputs[0].powf(*n),
};
self.nodes[id].output = result;
result
}
}
3.3 反向传播
反向传播是自动微分的精髓。从输出节点开始,沿着计算图反向传播梯度。对于一个节点 $y = f(x_1, x_2, ..., x_n)$,若已知 $\frac{dL}{dy}$(上游传入的梯度),则对每个输入 $x_i$: $$\frac{dL}{dx_i} = \frac{dL}{dy} \cdot \frac{\partial y}{\partial x_i}$$
impl ComputationGraph {
fn backward(&mut self, output_id: NodeId) {
// 初始化所有梯度为 0
for node in self.nodes.iter_mut() {
node.grad = 0.0;
}
// 输出节点的梯度设为 1(dL/dL = 1,或输出是标量损失)
self.nodes[output_id].grad = 1.0;
// 反向遍历节点(按创建顺序的逆序)
for id in (0..=output_id).rev() {
let node = self.nodes[id].clone(); // 避免借用冲突
let grad = node.grad;
match &node.op {
Op::Add => {
// d/dx (x+y) = 1, d/dy (x+y) = 1
self.nodes[node.inputs[0]].grad += grad * 1.0;
self.nodes[node.inputs[1]].grad += grad * 1.0;
}
Op::Sub => {
self.nodes[node.inputs[0]].grad += grad * 1.0;
self.nodes[node.inputs[1]].grad += grad * (-1.0);
}
Op::Mul => {
let x = self.nodes[node.inputs[0]].output;
let y = self.nodes[node.inputs[1]].output;
// d/dx (xy) = y, d/dy (xy) = x
self.nodes[node.inputs[0]].grad += grad * y;
self.nodes[node.inputs[1]].grad += grad * x;
}
Op::Div => {
let x = self.nodes[node.inputs[0]].output;
let y = self.nodes[node.inputs[1]].output;
// d/dx (x/y) = 1/y, d/dy (x/y) = -x/y²
self.nodes[node.inputs[0]].grad += grad / y;
self.nodes[node.inputs[1]].grad += grad * (-x / (y * y));
}
Op::Neg => {
self.nodes[node.inputs[0]].grad += grad * (-1.0);
}
Op::Exp => {
let val = self.nodes[node.inputs[0]].output;
self.nodes[node.inputs[0]].grad += grad * val.exp();
}
Op::Ln => {
let x = self.nodes[node.inputs[0]].output;
self.nodes[node.inputs[0]].grad += grad / x;
}
Op::Sin => {
let x = self.nodes[node.inputs[0]].output;
self.nodes[node.inputs[0]].grad += grad * x.cos();
}
Op::Cos => {
let x = self.nodes[node.inputs[0]].output;
self.nodes[node.inputs[0]].grad += grad * (-x.sin());
}
Op::Tanh => {
let val = self.nodes[node.inputs[0]].output;
self.nodes[node.inputs[0]].grad += grad * (1.0 - val * val);
}
Op::Sigmoid => {
let val = self.nodes[node.inputs[0]].output;
let sigma = 1.0 / (1.0 + (-val).exp());
self.nodes[node.inputs[0]].grad += grad * sigma * (1.0 - sigma);
}
Op::Relu => {
let x = self.nodes[node.inputs[0]].output;
if x > 0.0 {
self.nodes[node.inputs[0]].grad += grad;
}
}
Op::Sqrt => {
let x = self.nodes[node.inputs[0]].output;
self.nodes[node.inputs[0]].grad += grad / (2.0 * x.sqrt());
}
Op::Pow(n) => {
let x = self.nodes[node.inputs[0]].output;
self.nodes[node.inputs[0]].grad += grad * n * x.powf(n - 1.0);
}
Op::Constant(_) => {
// 常量节点无梯度
}
}
}
}
// 获取某个变量的梯度
fn gradient(&self, var_id: NodeId) -> f64 {
self.nodes[var_id].grad
}
}
四、实战:训练 XOR 网络
XOR 问题是神经网络的"hello world"——它是非线性不可分的,需要隐藏层才能解决。我们用自动微分从零训练一个 MLP:
use rand::Rng;
struct XORNetwork {
graph: ComputationGraph,
// 权重节点 ID
w1: NodeId, w2: NodeId, w3: NodeId, w4: NodeId, // 隐藏层权重
b1: NodeId, b2: NodeId, // 隐藏层偏置
w5: NodeId, w6: NodeId, // 输出层权重
b3: NodeId, // 输出层偏置
}
impl XORNetwork {
fn new() -> Self {
let mut g = ComputationGraph::new();
// 初始化权重(小的随机值)
let mut rng = rand::thread_rng();
let init = |g: &mut ComputationGraph| g.variable(rng.gen::<f64>() * 2.0 - 1.0);
XORNetwork {
w1: init(&mut g), w2: init(&mut g), w3: init(&mut g), w4: init(&mut g),
b1: init(&mut g), b2: init(&mut g),
w5: init(&mut g), w6: init(&mut g),
b3: init(&mut g),
graph: g,
}
}
fn forward(&mut self, x: NodeId, y: NodeId) -> NodeId {
// 隐藏层: h1 = sigmoid(w1*x + w2*y + b1)
// h2 = sigmoid(w3*x + w4*y + b2)
let g = &mut self.graph;
let h1_in = g.add(g.add(g.mul(self.w1, x), g.mul(self.w2, y)), self.b1);
let h1 = g.sigmoid(h1_in);
let h2_in = g.add(g.add(g.mul(self.w3, x), g.mul(self.w4, y)), self.b2);
let h2 = g.sigmoid(h2_in);
// 输出层: out = sigmoid(w5*h1 + w6*h2 + b3)
let out_in = g.add(g.add(g.mul(self.w5, h1), g.mul(self.w6, h2)), self.b3);
let out = g.sigmoid(out_in);
out
}
// 前向传播 + 计算损失(均方误差)
fn train_step(&mut self, x_val: f64, y_val: f64, target: f64, lr: f64) -> f64 {
let g = &mut self.graph;
// 重置计算图(简化实现:重新创建)
// 实际应用中需要动态图或每次重建
let x = g.variable(x_val);
let target_node = g.variable(target);
let pred = self.forward(x, x); // 简化:第二个输入复用 x
// 损失 = (pred - target)² 使用 MSE
let diff = g.sub(pred, target_node);
let loss = g.mul(diff, diff);
// 前向传播
self.graph.forward(loss);
let loss_val = self.graph.nodes[loss].output;
// 反向传播
self.graph.backward(loss);
// 更新权重
for &wid in &[self.w1, self.w2, self.w3, self.w4, self.b1, self.b2, self.w5, self.w6, self.b3] {
let grad = self.graph.nodes[wid].grad;
let val = self.graph.nodes[wid].output;
self.graph.nodes[wid].output = val - lr * grad;
}
loss_val
}
}
fn main() {
let mut net = XORNetwork::new();
let learning_rate = 0.5;
// XOR 真值表
let data = vec![
(0.0, 0.0, 0.0),
(0.0, 1.0, 1.0),
(1.0, 0.0, 1.0),
(1.0, 1.0, 0.0),
];
// 训练 10000 轮
for epoch in 0..10000 {
let mut total_loss = 0.0;
for &(x, y, t) in &data {
let loss = net.train_step(x, y, t, learning_rate);
total_loss += loss;
}
if epoch % 1000 == 0 {
println!("Epoch {} | Loss: {:.6}", epoch, total_loss);
}
}
// 测试结果
println!("\n=== XOR 测试结果 ===");
for &(x, y, t) in &data {
let loss = net.train_step(x, y, t, 0.0); // lr=0, 不更新
let output = 1.0 / (1.0 + (-loss).exp());
println!("XOR({}, {}) = {:.4} (期望: {})", x, y, output, t);
}
}
五、高阶导数与 Hessian 计算
自动微分的另一个诱人特性是:高阶导数可以通过嵌套 AD 自动获得。Hessian 矩阵的特征值分析在优化理论中非常重要:判断鞍点、分析收敛率、构建牛顿法。
// 对偶数嵌套:f64 + f64ε₁ + f64ε₂ + f64ε₁ε₂
// 用于同时计算函数值、一阶导和二阶导
#[derive(Debug, Clone, Copy)]
struct HyperDual {
val: f64,
grad1: f64, // 对 x1 的一阶导
grad2: f64, // 对 x2 的一阶导
hess: f64, // 二阶混合导 ∂²/(∂x₁∂x₂)
}
impl HyperDual {
fn constant(val: f64) -> Self {
Self { val, grad1: 0.0, grad2: 0.0, hess: 0.0 }
}
// 对 x1 的变量
fn var1(val: f64) -> Self {
Self { val, grad1: 1.0, grad2: 0.0, hess: 0.0 }
}
// 对 x2 的变量
fn var2(val: f64) -> Self {
Self { val, grad1: 0.0, grad2: 1.0, hess: 0.0 }
}
}
impl Add for HyperDual {
type Output = Self;
fn add(self, rhs: Self) -> Self::Output {
Self {
val: self.val + rhs.val,
grad1: self.grad1 + rhs.grad1,
grad2: self.grad2 + rhs.grad2,
hess: self.hess + rhs.hess,
}
}
}
impl Mul for HyperDual {
type Output = Self;
fn mul(self, rhs: Self) -> Self::Output {
Self {
val: self.val * rhs.val,
grad1: self.grad1 * rhs.val + self.val * rhs.grad1,
grad2: self.grad2 * rhs.val + self.val * rhs.grad2,
// ∂²(fg)/(∂x₁∂x₂) = ∂/∂x₁(∂(fg)/∂x₂)
// = ∂/∂x₁(f'₂g + fg'₂) = f'₁'₂g + f'₁g'₂ + f'₂g'₁ + fg'₁'₂
hess: self.hess * rhs.val
+ self.grad1 * rhs.grad2
+ self.grad2 * rhs.grad1
+ self.val * rhs.hess,
}
}
}
// 使用 HyperDual 计算 Hessian 的示例
fn compute_hessian() {
// f(x,y) = x²y + sin(x) + y³
let x = HyperDual::var1(2.0);
let y = HyperDual::var2(3.0);
let f = x * x * y + x.sin() + y * y * y;
println!("f(2,3) = {}", f.val); // 4*3 + sin(2) + 27 = 39.909
println!("∂f/∂x = {}", f.grad1); // 2*2*3 + cos(2) = 12 - 0.416 = 11.584
println!("∂f/∂y = {}", f.grad2); // 4 + 27 = 31
println!("∂²f/∂x∂y = {}", f.hess); // ∂/∂y(2xy + cos(x)) = 2x = 4
}
六、生产级 AD 框架的设计考量
从零实现 AD 引擎让我们理解了原理,但要构建像 PyTorch 这样的生产级框架还需考虑大量工程细节:
6.1 内存管理与梯度检查点
对于深层网络(如 ResNet-152 或大型 Transformer),存储所有中间激活值会耗尽 GPU 显存。梯度检查点(Gradient Checking/Rematerialization) 通过在前向传播时保存部分检查点,在反向时重新计算中间激活来节省内存。
实现思路:- 定义 checkpoint 函数:前向时只存储输入和函数,丢弃中间值
- 反向时从检查点重新执行前向,再执行反向
- 空间复杂度从 $O(L)$ 降为 $O(\sqrt{L})$,代价是约 33% 的额外计算
6.2 自动向量化的批量梯度
现代 AD 引擎自动识别批量维度,通过广播语义高效计算 batch 梯度。关键洞察:如果损失是 batch 中每个样本损失的均值,则对权重梯度只需一次 backward pass。
6.3 JVP 与 VJP 的对偶
- JVP(Jacobian-Vector Product):前向模式,适合 tall 矩阵(输出维度 >> 输入维度)
- VJP(Vector-Jacobian Product):反向模式,适合 wide 矩阵(输入维度 >> 输出维度),这是深度学习唯一需要的场景
JAX 的高明之处:jax.jvp 和 jax.vjp 作为原语可以自由组合。
6.4 不可微函数的处理
// 直通估计器(Straight-Through Estimator)
// 前向:二值化,反向:直通
fn straight_through_quantize(x: f64) -> (f64, impl Fn(f64) -> f64) {
let quantized = if x > 0.0 { 1.0 } else { -1.0 };
// 反向梯度 = 若 |x| < 1 则为 1,否则为 0(类似 clip)
let grad_fn = move |grad: f64| {
if x.abs() <= 1.0 { grad } else { 0.0 }
};
(quantized, grad_fn)
}
6.5 自定义梯度(Custom Gradients)
use std::ops::{Add, Sub, Mul, Div, Neg};
impl Add for Dual {
type Output = Self;
fn add(self, rhs: Self) -> Self::Output {
Self {
val: self.val + rhs.val,
deriv: self.deriv + rhs.deriv, // (f+g)' = f' + g'
}
}
}
impl Sub for Dual {
type Output = Self;
fn sub(self, rhs: Self) -> Self::Output {
Self {
val: self.val - rhs.val,
deriv: self.deriv - rhs.deriv,
}
}
}
impl Mul for Dual {
type Output = Self;
fn mul(self, rhs: Self) -> Self::Output {
Self {
val: self.val * rhs.val,
// 乘法法则: (fg)' = f'g + fg'
deriv: self.deriv * rhs.val + self.val * rhs.deriv,
}
}
}
impl Div for Dual {
type Output = Self;
fn div(self, rhs: Self) -> Self::Output {
Self {
val: self.val / rhs.val,
// 除法法则: (f/g)' = (f'g - fg') / g²
deriv: (self.deriv * rhs.val - self.val * rhs.deriv) / (rhs.val * rhs.val),
}
}
}
impl Neg for Dual {
type Output = Self;
fn neg(self) -> Self::Output {
Self {
val: -self.val,
deriv: -self.deriv,
}
}
}0
七、性能优化与工程实践
7.1 使用 SIMD 加速标量运算
对于需要逐元素操作的 AD(如大规模向量化),可以使用 Rust 的 std::simd(nightly)或 packed_simd 库并行处理多个计算图节点。
7.2 拓扑排序优化
在反向传播中,节点的求值顺序必须保证拓扑序。实际框架使用 Kahn 算法 在图构建完成后进行一次排序,避免递归求值带来的开销。
7.3 延迟求值与图编译
TensorFlow 2.x 的 tf.function 和 PyTorch 的 torch.compile 将 Python 代码编译为静态计算图:
- 消除 Python 解释器开销
- 实现算子融合(kernel fusion):合并
matmul + bias + relu为单次 kernel launch - 预先分配内存:静态图可一次性计算所有 tensor 的大小,避免运行时分配
7.4 混合精度训练
use std::ops::{Add, Sub, Mul, Div, Neg};
impl Add for Dual {
type Output = Self;
fn add(self, rhs: Self) -> Self::Output {
Self {
val: self.val + rhs.val,
deriv: self.deriv + rhs.deriv, // (f+g)' = f' + g'
}
}
}
impl Sub for Dual {
type Output = Self;
fn sub(self, rhs: Self) -> Self::Output {
Self {
val: self.val - rhs.val,
deriv: self.deriv - rhs.deriv,
}
}
}
impl Mul for Dual {
type Output = Self;
fn mul(self, rhs: Self) -> Self::Output {
Self {
val: self.val * rhs.val,
// 乘法法则: (fg)' = f'g + fg'
deriv: self.deriv * rhs.val + self.val * rhs.deriv,
}
}
}
impl Div for Dual {
type Output = Self;
fn div(self, rhs: Self) -> Self::Output {
Self {
val: self.val / rhs.val,
// 除法法则: (f/g)' = (f'g - fg') / g²
deriv: (self.deriv * rhs.val - self.val * rhs.deriv) / (rhs.val * rhs.val),
}
}
}
impl Neg for Dual {
type Output = Self;
fn neg(self) -> Self::Output {
Self {
val: -self.val,
deriv: -self.deriv,
}
}
}1
八、与现有框架的对比与选型建议
- 研究/原型开发:PyTorch 的 eager 模式最灵活
- 需要函数式变换/科学计算:JAX 的
grad/vmap/pmap组合无与伦比 - 生产部署:TensorFlow Serving 或 PyTorch TorchServe
- 嵌入式/边缘:从我们的引擎出发,添加自定义 backend
九、总结
从零构建自动微分引擎的过程,让我们深入到现代机器学习基础设施的核心:
- 对偶数提供了前向模式 AD 的最简实现,数学优雅且无需追踪计算图。
- 计算图 + 反向传播是工业标准,一次前向 + 一次反向即可获得全维度梯度。
- 嵌套 AD天然支持高阶导数,无需手动推导二阶梯度。
- 工程优化(检查点、SIMD、图编译、混合精度)才是从玩具到产品的关键差距。
自动微分的工程仍在快速演进。随着模型规模增长到万亿参数,编译型 AD(XLA、torch.compile)和分布式自动微分(GSPMD)正成为下一个前沿。掌握这些底层原理,才能在框架更迭浪潮中保持不变的竞争力。
> 完整代码仓库:以上代码可直接编译运行,读者可在此基础上扩展到矩阵 AD、卷积层和 attention 机制的实现。核心原理不变——复杂算子拆解为基本操作的组合,微分规则逐层传播。
本文首发于 ybb.press,欢迎交流指正。

发表评论 取消回复