Rust autodiff:`autodiff_forward` 与 `autodiff_reverse` 自动微分宏的技术指南
发布时间:2026/9/10 0:58:52来源:尧图网络
Rust autodiffautodiff_forward与autodiff_reverse自动微分宏的技术指南【免费下载链接】rustEmpowering everyone to build reliable and efficient software.项目地址: https://gitcode.com/GitHub_Trending/ru/rust本文围绕 Rust 标准库中不稳定的autodiff模块文档展开完整讲解#[autodiff_forward]前向模式与#[autodiff_reverse]反向模式两类自动微分宏的参数语法、活动标注activity、在泛型函数、嵌套函数、trait 与 impl 块中的用法以及通过薄包装生成高阶导数如 Hessian的技巧并结合rustc_builtin_macros的宏展开实现与core::intrinsics::autodiff内建函数说明这些宏在编译器内部的工作机制帮助你在 Rust 中以声明式方式直接生成求导函数。功能定位与启用方式自动微分Automatic Differentiation, AD支持目前以独立模块的形式集成在core标准库中模块文档即 library/core/src/autodiff.md通过include_str!内联到模块文档中。在 library/core/src/lib.rs 中可以看到该模块的定义// We dont export this through #[macro_export] for now, to avoid breakage. #[unstable(feature autodiff, issue 124509)] #[doc include_str!(../../core/src/autodiff.md)] pub mod autodiff { #[unstable(feature autodiff, issue 124509)] pub use crate::macros::builtin::{autodiff_forward, autodiff_reverse}; }从源码结构看该模块处于不稳定状态追踪 issue 为124509因此使用前提是nightly 工具链并开启#![feature(autodiff)]然后通过use std::autodiff::*;引入两个宏core中的定义经由标准库重导出。两个宏本体定义在 library/core/src/macros/mod.rs 中标记为#[rustc_builtin_macro]即由编译器内建扩展实现而非普通过程宏。需要注意的适用前提来自宏的文档注释与模块文档文档注释明确注明示例代码标记为rust,ignore (autodiff requires a -Z flag as well as fat-lto for testing)即该功能依赖特定的-Z编译器开关且测试/构建需要fat LTOlto fat才能完成求导不支持在构建中关闭 fat LTO见下文“当前限制”一节。宏语法与活动Activity标注两类宏的调用签名在 library/core/src/macros/mod.rs 中有统一定义#[autodiff_forward(NAME, INPUT_ACTIVITIES, OUTPUT_ACTIVITY)]#[autodiff_reverse(NAME, INPUT_ACTIVITIES, OUTPUT_ACTIVITY)]其中NAME生成的求导函数的函数名INPUT_ACTIVITIES为每个输入参数各指定一个活动标注OUTPUT_ACTIVITY若函数显式返回- ()或隐式无返回值则必须省略否则必须指定。活动标注的取值与含义如下活动适用宏含义Dualforward用于浮点标量参数或引用/裸指针等间接参数。用在输入上时会紧跟原参数生成一个同类型的“影子shadow参数”用在返回值上时生成的函数将返回一个两个浮点标量组成的元组原值 导数Constforward / reverse用于非浮点参数或对浮点参数作为优化——表示不关心该参数方向上的导数Activereverse用于浮点标量。用在输入上时会在生成函数的返回元组中追加一个浮点用在返回上时会向参数列表追加一个浮点种子seedDuplicatedreverse用于引用、裸指针等间接参数紧随原参数生成一个同类型影子参数对 const 引用/指针参数其影子类型为可变的引用或指针前向模式与反向模式的选择依据宏文档给出的经验法则是若被标记为活跃的输出多于输入前向模式通常更高效对应 “Vector-Jacobian product”VJP反之若输入多于输出反向模式更高效对应 “Jacobian-Vector product”JVP。前向模式的调用约定是想追踪某个输入的影响时把该输入影子初始化为1.0、其余输入影子为0.0、输出影子为0.0调用后输入影子被清零输出影子中即为导数值。反向模式则是输出影子seed初始化为1.0调用后输入影子中即包含导数且与不同之处在于调用不会重置输入的影子。一般用法同一函数叠加多个求导宏宏文档autodiff.md指出autodiff 宏几乎可以应用于所有函数定义支持接受结构体、数组、切片、向量、元组等参数。一个函数上可以同时叠加多个 autodiff 宏例如分别独立计算关于x与y的偏导#[autodiff_forward(dsquare1, Dual, Const, Dual)] #[autodiff_forward(dsquare2, Const, Dual, Dual)] #[autodiff_forward(dsquare3, Active, Active, Active)] fn square(x: f64, y: f64) - f64 { x * x 2.0 * y }这里dsquare1计算关于x的导数dsquare2计算关于y的导数dsquare3则同时对两个输入播种。求导后的函数与原函数位于同一作用域并具有与原函数一致的pub/私有可见性。宏文档注释中还给出了一个经典的 Rosenbrock 函数示例见 library/core/src/macros/mod.rs 与 L1594-L1613直观展示了前向/反向生成的函数在运行时如何调用#![feature(autodiff)] use std::autodiff::*; #[autodiff_forward(rb_fwd1, Dual, Const, Dual)] #[autodiff_forward(rb_fwd2, Const, Dual, Dual)] #[autodiff_forward(rb_fwd3, Dual, Dual, Dual)] fn rosenbrock(x: f64, y: f64) - f64 { (1.0 - x).powi(2) 100.0 * (y - x.powi(2)).powi(2) } #[autodiff_reverse(rb_rev, Active, Active, Active)] fn rb_rev_demo(x: f64, y: f64) - f64 { (1.0 - x).powi(2) 100.0 * (y - x.powi(2)).powi(2) } fn main() { let x0 rosenbrock(1.0, 3.0); // 400.0 let (x1, dx1) rb_fwd1(1.0, 1.0, 3.0); // (400.0, -800.0) let (x2, dy1) rb_fwd2(1.0, 3.0, 1.0); // (400.0, 400.0) let (x3, dxy) rb_fwd3(1.0, 1.0, 3.0, 1.0); // (400.0, -400.0) let (out, dxr, dyr) rb_rev(1.0, 3.0, 1.0); // (400.0, -800.0, 400.0) }对于带引用输出的函数如out: mut f64前向模式将其标注为Dual时会追加mut影子输出反向模式则标注为Duplicated并同样追加影子且调用后种子会被重置为0.0。泛型函数与嵌套函数文档进一步确认了两个进阶场景。其一支持带泛型参数的函数#[autodiff_forward(generic_derivative, Duplicated, Active)] fn generic_fT: std::ops::MulOutput T Copy(x: T) - T { x * x }其二支持对函数体内的嵌套函数求导fn outer(x: f64) - f64 { #[autodiff_forward(inner_derivative, Dual, Const)] fn inner(y: f64) - f64 { y * y } inner_derivative(x, 1.0) } fn main() { assert_eq!(outer(3.14), 6.28); }从 compiler/rustc_builtin_macros/src/autodiff.rs 的展开入口可以看到宏对Annotatable::Item、Annotatable::Stmt语句中的项即嵌套函数的情形以及Annotatable::AssocItemtrait 或 impl 中的关联函数三种挂载点都会提取可见性、函数签名、标识符与泛型参数来生成求导版本这与文档描述的三种用法一一对应。与 trait 和 impl 的配合宏文档给出了 autodiff 与 trait 结合的三种方式1. 在 trait 声明上标注提供默认求导实现struct Foo { a: f64 } trait MyTrait { #[autodiff_reverse(df, Const, Active, Active)] fn f(self, x: f64) - f64; } impl MyTrait for Foo { fn f(self, x: f64) - f64 { x.sin() } } fn main() { let foo Foo { a: 3.0 }; assert_eq!(foo.f(2.0), 2.0_f64.sin()); assert_eq!(foo.df(2.0, 1.0).1, 2.0_f64.cos()); }此时df成为 trait 的一部分由 trait 定义方提供默认实现实现MyTrait的用户可以选择沿用默认实现也可以覆盖它来自定义导数——这相当于一种“自定义导数”机制。2. 用生成的函数去实现 trait 方法trait MyTrait { fn f(self, x: f64) - f64; fn df(self, x: f64, seed: f64) - (f64, f64); } impl MyTrait for Foo { #[autodiff_reverse(df, Const, Active, Active)] fn f(self, x: f64) - f64 { self.a * 0.25 * (x * x - 1.0 - 2.0 * x.ln()) } }这里df直接由宏在 impl 中生成从而满足 trait 约束。3. 普通 impl 块结构体字段需要“影子结构体”对不带 trait 的普通impl块求导时若需要追踪结构体字段方向的导数必须用一个与原结构体同构的影子结构体shadow struct承载各字段的导数struct OptProblem { a: f64, b: f64 } impl OptProblem { #[autodiff_reverse(d_objective, Duplicated, Duplicated, Duplicated)] fn objective(self, x: [f64], out: mut f64) { *out self.a x[0].sqrt() * self.b } } fn main() { let p OptProblem { a: 1., b: 2. }; let mut p_shadow OptProblem { a: 0., b: 0. }; let mut dx [0.0]; let mut out 0.0; let mut dout 1.0; p.d_objective(mut p_shadow, x, mut dx, mut out, mut dout); }注意调用约定self对应的影子结构体以mut形式传入且排在原参数之前x与out的影子紧随各自原参数。高阶导数通过薄包装对求导函数再求导要生成二阶导数例如 Hessian可以在一个由 autodiff 宏生成的函数外层再套一个“薄包装”函数并对包装函数应用另一个 autodiff 宏。文档示例为“前向模式套在反向模式之上”#[autodiff_reverse(df, Duplicated, Duplicated)] fn f(x: [f64;2], y: mut f64) { *y x[0] * x[0] x[1] * x[0] } #[autodiff_forward(h, Dual, Dual, Dual, Dual)] fn wrapper(x: [f64;2], dx: mut [f64;2], y: mut f64, dy: mut f64) { df(x, dx, y, dy); } fn main() { let mut y 0.0; let x [2.0, 2.0]; let mut dy 0.0; let mut dx [1.0, 0.0]; let mut bx [0.0, 0.0]; let mut by 1.0; let mut dbx [0.0, 0.0]; let mut dby 0.0; h(x, mut dx, mut bx, mut dbx, mut y, mut dy, mut by, mut dby); assert_eq!(dbx, [2.0, 1.0]); }df先沿反向模式把雅可比信息写入dx、dy影子外层h再沿前向模式追踪这些影子本身的变化最终dbx中即为关于输入方向的二阶导数f x0*x0 x1*x0在dx (1, 0)方向上得到[2.0, 1.0]。源码级机制宏展开与内建函数从仓库源码可以确认这套宏的完整工作链路其核心文件是 compiler/rustc_builtin_macros/src/autodiff.rs。第一步属性解析。宏展开入口expand_forward/expand_reverseL164-L180将属性转发到expand_with_mode后者把用户写入的元组解析为内部结构RustcAutodiffL85-L154第一个元组项是生成函数名随后为各参数的活动标注若函数有非单元返回值最后一个活动标注会被切分出来作为返回值活动否则自动补一个None占位。值得一提的是从from_ast中“batch/vector mode or scalar mode”的注释与width解析逻辑看L93-L113底层解析器还预留了在函数名之后写入整数width的批量向量化求导形式未显式给出时默认width 1的标量模式这一内部形式尚未出现在公开文档中可推断为后续将暴露的能力。第二步改写原函数、生成求导占位函数。expand_with_modeL205-L344的注释说明了展开形态原函数被加上#[rustc_autodiff]属性、并保留可见性与泛型参数原函数体在语义上被替换为不可执行的占位体同时在同作用域生成名为NAME的新函数打上#[rustc_autodiff(Mode, width, 活动列表)]属性。若原函数带有#[inline(never)]展开代码会检测并把它同步到求导函数上L432-L437。第三步内建函数触发 LLVM 层的求导。生成的求导函数体是一次对core::intrinsics::autodiff的调用。该内建函数的文档在 library/core/src/intrinsics/mod.rs 中给出了最直观的展开示例#[autodiff_forward(df1, Dual, Const, Dual)] pub fn f1(x: [f64], y: f64) - f64 { unimplemented!() } // 展开为 #[rustc_autodiff] #[inline(never)] pub fn f1(x: [f64], y: f64) - f64 { ::core::panicking::panic(not implemented) } #[rustc_autodiff(Forward, 1, Dual, Const, Dual)] pub fn df1(x: [f64], bx_0: [f64], y: f64) - (f64, f64) { ::core::intrinsics::autodiff(f1::, df1::, (x, bx_0, y)) }内建函数签名pub const fn autodiffF, G, T: Tuple, R(f: F, df: G, args: T) - R;的文档说明它“使用 Enzyme 生成f的自动微分 LLVM 函数体df作为求导函数、args作为其参数”。也就是说真正的求导变换并不发生在 Rust 词法层面而是在 LLVM 代码生成阶段由 Enzyme AD pass 基于标注完成——这也解释了为何必须启用 fat LTO求导发生在最终链接单元的 IR 层面函数体必须在同一优化上下文中可见。验证入口。仓库内与该功能直接相关的测试包括tests/pretty/autodiff/autodiff_forward.rs 与 tests/pretty/autodiff/autodiff_reverse.rs校验两类宏展开后的 pretty 打印形态对应.pp期望文件tests/codegen-llvm/autodiff/autodiffv2.rs校验 LLVM 层的求导产物tests/ui/autodiff/autodiff_illegal.rs非法用法的报错诊断tests/ui/feature-gates/feature-gate-autodiff.rs 及feature-gate-autodiff-use.rs确认autodiff在未开启 feature gate 时按预期被门禁拦截。当前限制与适用前提模块文档明确列出了当前截至该仓库快照的限制使用时必须对照不支持对接受dyn Trait参数的函数求导构建必须启用lto fat未启用 fat LTO 的构建尚不支持与“求导在 LLVM 链接单元层面进行”的机制一致;debug 模式构建更容易出现编译失败即当前实现更依赖优化构建路径该功能整体处于 unstable 状态featureautodiffissue 124509仅在 nightly 上可用且如宏文档注释所述还需要额外的-Z开关配合活动标注目前仅支持文档所列的Dual/Const前向与Active/Duplicated/Const反向源码与文档均注明“更多选项将在后续暴露”。在上述前提下这套宏的价值在于无需手工推导求导公式也无需引入第三方 AD 框架即可在纯 Rust 中以属性注解的方式为标量函数、带引用输出的函数、trait 方法乃至嵌套/泛型函数批量生成前向、反向乃至二阶求导函数并且生成物与普通 Rust 函数一样可组合、可被 trait 化。【免费下载链接】rustEmpowering everyone to build reliable and efficient software.项目地址: https://gitcode.com/GitHub_Trending/ru/rust创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
网站建设高端定制企业官网