You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用Rust的ndarray实现支持自动微分的异构泛型计算图?

异构泛型自动微分计算图的Rust实现方案

问题背景

我想在Rust中实现支持自动微分的计算图,选用了类NumPy的ndarray crate。常规的图结构实现(比如带引用的Vec<Node>或引用计数指针)在这里行不通,因为ndarray::Array<T, D>拥有自身数据,且T(如f32、Complex<f32>)和D(维度)都是泛型参数。我需要构建一个各节点可拥有不同T、D泛型的计算图,以此支持FFT这类实数转复数的操作。

现有尝试代码

use std::{cell::RefCell, ops::Add, rc::Rc};
use ndarray::prelude::*;
use num_traits::One;

pub enum OpType {
    Add,
    MatVecMul,
    Leaf
}

pub struct Expr<'t, T, D> {
    index: usize,
    tape: &'t Tape,
    output: Rc<Array<T, D>>
}

// 因Array的泛型特性,Node必须携带泛型参数
pub struct Node<T, D> {
    value: Rc<Array<T, D>>,
    op: OpType,
    deps: [usize; 2],
}

// 当前问题核心:无法为该trait添加返回泛型Array的方法(违反对象安全)
pub trait TapeNode {

}

impl <T, D> TapeNode for Node<T, D> {

}

pub struct Tape {
    nodes: RefCell<Vec<Box<dyn TapeNode>>>
}

impl Tape {
    pub fn new() -> Self {
        Tape {
            nodes: RefCell::new(Vec::new())
        }
    }

    pub fn push_leaf<'t, T, D>(&'t self, value: Array<T, D>) -> Expr<'t, T, D>
    where
        T: 'static,
        D: 'static
    {
        let mut nodes = self.nodes.borrow_mut();
        let value_ref = Rc::from(value);
        let new_node = Node {
            value: value_ref.clone(),
            op: OpType::Leaf,
            deps: [0, 0],
        };

        let len = nodes.len();
        nodes.push(Box::from(new_node));

        Expr { index: len, tape: &self, output: value_ref.clone() }
    }

    pub fn push_1<'t, T, D>(&'t self, value: Array<T, D>, id_0: usize, op: OpType) -> Expr<'t, T, D> 
    where
        T: 'static,
        D: 'static
    {
        let mut nodes = self.nodes.borrow_mut();
        let value_ref = Rc::from(value);
        let new_node = Node {
            value: value_ref.clone(),
            op,
            deps: [id_0, 0],
        };

        let len = nodes.len();
        nodes.push(Box::from(new_node));

        Expr { index: len, tape: &self, output: value_ref.clone() }
    }

    pub fn push_2<'t, T, D>(&'t self, value: Array<T, D>, id_0: usize, id_1: usize, op: OpType) -> Expr<'t, T, D> 
    where
        T: 'static,
        D: 'static
    {
        let mut nodes = self.nodes.borrow_mut();
        let value_ref = Rc::from(value);
        let new_node = Node {
            value: value_ref.clone(),
            op,
            deps: [id_0, id_1],
        };

        let len = nodes.len();
        nodes.push(Box::from(new_node));

        Expr { index: len, tape: &self, output: value_ref.clone() }
    }
}

impl <'t, T, D> Add for Expr<'t, T, D> 
where
    T: Add<T, Output = T> + Clone + One + 'static,
    D: Dimension + 'static
{
    type Output = Self;
    fn add(self, rhs: Self) -> Self::Output {
        let added = self.output.as_ref() + rhs.output.as_ref();
        self.tape.push_2(added, self.index, rhs.index, OpType::Add)
    }
}

/*
反向传播时会出现问题:需要反向遍历磁带并反复应用链式法则,
这要求能获取每个节点的Array,但当前方案无法实现。
*/

核心问题

当前代码用Tape存储动态分发的TapeNode trait对象实现了基础图结构,但反向传播时需要获取每个节点的Array。若给TapeNode添加返回泛型Array<T, D>的方法,会违反Rust的对象安全规则(泛型方法无法用于动态分发的trait对象),导致整个体系失效。

解决方案思路

方案一:改进现有动态分发方案

1. 用枚举封装所有节点类型

放弃纯动态分发,定义一个AnyNode枚举,枚举所有你需要支持的Node<T, D>组合:

use ndarray::{Ix1, Ix2};
use num_complex::Complex;

pub enum AnyNode {
    LeafF32(Node<f32, Ix1>),
    AddF32(Node<f32, Ix1>),
    ComplexFFT(Node<Complex<f32>, Ix2>),
    // 根据需求添加更多类型组合
}

然后将Tape的nodes字段改为RefCell<Vec<AnyNode>>,反向传播时通过模式匹配获取对应类型的Array:

// 反向传播示例
for node in self.nodes.borrow_mut().iter().rev() {
    match node {
        AnyNode::LeafF32(n) => {
            // 处理f32一维叶子节点的梯度
            let val = &n.value;
            // ...
        }
        AnyNode::ComplexFFT(n) => {
            // 处理复数二维FFT节点的梯度
            let val = &n.value;
            // ...
        }
        // 其他节点类型分支
    }
}

优点是实现简单,类型安全;缺点是需要预先枚举所有可能的T和D组合,扩展性有限。

2. 类型擦除+向下转型

给TapeNode添加返回&dyn Any的方法,将节点值擦除为Any类型,反向传播时根据操作类型向下转型:

use std::any::Any;

pub trait TapeNode {
    fn value_as_any(&self) -> &dyn Any;
    fn op(&self) -> OpType;
    fn deps(&self) -> [usize; 2];
}

impl<T: 'static, D: 'static> TapeNode for Node<T, D> {
    fn value_as_any(&self) -> &dyn Any {
        &self.value
    }

    fn op(&self) -> OpType {
        self.op.clone()
    }

    fn deps(&self) -> [usize; 2] {
        self.deps
    }
}

反向传播时,根据OpType判断节点可能的类型,再用downcast_ref获取具体的Array:

for node in self.nodes.borrow_mut().iter().rev() {
    match node.op() {
        OpType::Add => {
            // 尝试转型为f32一维数组
            if let Some(val) = node.value_as_any().downcast_ref::<Rc<Array<f32, Ix1>>>() {
                // 处理加法节点的反向传播
            }
            // 尝试转型为复数二维数组
            else if let Some(val) = node.value_as_any().downcast_ref::<Rc<Array<Complex<f32>, Ix2>>>() {
                // 处理复数加法节点的反向传播
            }
            // 其他可能的类型分支
        }
        OpType::Leaf => {
            // 处理叶子节点的梯度
        }
        // 其他操作类型分支
    }
}

优点是保留了动态分发的灵活性;缺点是需要手动处理所有可能的类型分支,代码量较大,且运行时可能出现转型失败的情况(需确保OpType与节点类型严格对应)。

方案二:参考neuronika的混合模式实现

neuronika采用编译时类型节点+动态反向磁带的混合模式,核心思路是:

  1. 正向计算:用泛型结构体表示每个计算节点(比如Add<Lhs, Rhs>),编译时保留完整类型信息,确保类型安全和高性能。
  2. 反向传播:每个节点在正向计算时,将对应的反向传播逻辑(捕获了输入值、操作参数等的闭包)注册到动态磁带中,这些闭包被类型擦除为Box<dyn FnMut(&mut GradientTape)>。
  3. 反向传播时,只需遍历磁带中的闭包并执行,闭包内部已知道如何处理对应类型的梯度计算,无需手动转型。

具体到你的场景,可以调整如下:

  • 给Expr<'t, T, D>添加一个反向传播闭包字段,该闭包负责根据输出梯度计算输入梯度并更新。
  • Tape不再存储节点的完整类型,而是存储所有反向传播闭包和梯度存储结构。
  • 每个操作(比如Add)在正向计算时,生成对应的反向闭包并注册到磁带中。

这种方式既保留了编译时的类型安全,又能支持异构节点的反向传播,是实现复杂自动微分计算图的更优方案。

内容的提问来源于stack exchange,提问作者nekechs

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.27 16:08:14