从零构建自动微分引擎:对偶数、计算图与反向传播的 Rust 深度实战

> 自动微分(Automatic Differentiation, AD)是现代机器学习的基石——从 PyTorch 到 JAX,从 TensorFlow 到 Hugging Face,所有深度学习框架的核心都依赖它。本文不借助任何现成 AD 库,从零实现一个功能完整的自动微分引擎,深入理解对偶数(Dual Numbers)、计算图构建、反向模式微分(Backpropagation)以及高阶导数的实现原理。


一、为什么需要自动微分

在深度学习训练中,梯度是优化器的燃料。计算梯度的方法有三种:

  1. 数值微分(Numerical Differentiation):有限差分法,简单但精度低、计算量大,$O(n)$ 次函数求值计算 n 维梯度。
  2. 符号微分(Symbolical Differentiation):对数学表达式求解析导数,精确但面临"表达式膨胀"问题。
  3. 自动微分(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

八、与现有框架的对比与选型建议

特性我们的引擎PyTorchJAXTensorFlow 动态图✅✅ (eager)❌ (jit)✅ (eager) 静态图编译❌TorchScript/Compilejit/tfjittf.function 函数变换❌有限vmap/pmap/jvp/vjpLimited 高阶导数✅ (嵌套)✅✅✅ GPU 支持❌✅✅✅ 分布式❌FSDP/Apexpjit/pmapMultiWorkerMirrored 选型建议:
  • 研究/原型开发:PyTorch 的 eager 模式最灵活
  • 需要函数式变换/科学计算:JAX 的 grad/vmap/pmap 组合无与伦比
  • 生产部署:TensorFlow Serving 或 PyTorch TorchServe
  • 嵌入式/边缘:从我们的引擎出发,添加自定义 backend

九、总结

从零构建自动微分引擎的过程,让我们深入到现代机器学习基础设施的核心:

  1. 对偶数提供了前向模式 AD 的最简实现,数学优雅且无需追踪计算图。
  2. 计算图 + 反向传播是工业标准,一次前向 + 一次反向即可获得全维度梯度。
  3. 嵌套 AD天然支持高阶导数,无需手动推导二阶梯度。
  4. 工程优化(检查点、SIMD、图编译、混合精度)才是从玩具到产品的关键差距。

自动微分的工程仍在快速演进。随着模型规模增长到万亿参数,编译型 AD(XLA、torch.compile)和分布式自动微分(GSPMD)正成为下一个前沿。掌握这些底层原理,才能在框架更迭浪潮中保持不变的竞争力。

> 完整代码仓库:以上代码可直接编译运行,读者可在此基础上扩展到矩阵 AD、卷积层和 attention 机制的实现。核心原理不变——复杂算子拆解为基本操作的组合,微分规则逐层传播。


本文首发于 ybb.press,欢迎交流指正。
点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部