Rust IndexSet与Python Tuple性能对比:自动微分加法优化疑问
自动微分Dual类加法:Rust与Python性能差异分析与优化请求
我分别用Python(以Tuple存储变量)和Rust(以IndexSet存储变量)实现了自动微分的Dual类加法逻辑。基准测试显示:
- 当两个Dual实例变量完全一致时,Rust版本的加法性能远低于Python(100变量时慢4倍,1000变量时慢7倍)
- 当变量不同时,Rust性能反而更优
我怀疑性能差异源于两者判断变量是否一致的逻辑:
- Python中直接使用
self.vars == argument.vars - Rust中则通过长度判断+逐元素比对的方式
现需要明确该逻辑为何导致显著性能差距,并寻求Rust侧的优化方案。
基准测试数据
Rust
- 浮点加法:265 ps
- 100不同变量Dual加法:97 us
- 1000不同变量Dual加法:953 us
- 100相同变量Dual加法:13 us(比Python慢4倍)
- 1000相同变量Dual加法:129 us(比Python慢7倍)
Python
- 浮点加法:77 ns
- 100不同变量Dual加法:353 us
- 1000不同变量Dual加法:26.4 ms
- 100相同变量Dual加法:3.2 us
- 1000相同变量Dual加法:18.6 us
代码实现
Python实现
import numpy as np class Dual: def __init__( self, real: float = 0.0, vars: tuple[str, ...] = (), dual: np.ndarray = np.ones(0), ): self.real, self.vars, self.dual = real, vars, dual def __add__(self, argument): if self.vars == argument.vars: return Dual(self.real + argument.real, self.vars, self.dual + argument.dual) else: x, y = self._to_combined_vars(argument) return x + y def _to_combined_vars(self, other): combined_vars = sorted(list(set(self.vars).union(set(other.vars)))) x = self if combined_vars == self.vars else self._to_new_vars(combined_vars) y = other if combined_vars == other.vars else other._to_new_vars(combined_vars) return x, y def _to_new_vars(self, new_vars): dual = np.zeros(len(new_vars)) ix_ = list(map(lambda x: new_vars.index(x), self.vars)) dual[ix_] = self.dual return Dual(self.real, new_vars, dual)
(注:原代码中_to_new_vars__为笔误,已修正为_to_new_vars)
Rust实现
use indexmap::set::IndexSet; use ndarray::{Array1, Array}; use auto_ops::{impl_op, impl_op_ex}; #[derive(Clone, Debug)] pub struct Dual { pub real: f64, pub vars: IndexSet<String>, pub dual: Array1<f64>, } impl Dual { fn to_combined_vars(&self, other: &Dual) -> (Dual, Dual) { let comb_vars = IndexSet::from_iter(self.vars.union(&other.vars).map(|x| x.clone())); (self.to_new_vars(&comb_vars), other.to_new_vars(&comb_vars)) } fn to_new_vars(&self, new_vars: &IndexSet<String>) -> Dual { let mut dual = Array::zeros(new_vars.len()); for (i, index) in new_vars.iter().map(|x| self.vars.get_index_of(x)).enumerate() { match index { Some(value) => { dual[[i]] = self.dual[[value]] } None => {} } } Dual { vars: new_vars.clone(), real: self.real, dual } } } impl_op_ex!(+ |a: &Dual, b: &Dual| -> Dual { if a.vars.len() == b.vars.len() && a.vars.iter().zip(b.vars.iter()).all(|(a,b)| a==b) { Dual { real: a.real + b.real, dual: a.dual.clone() + b.dual.clone(), vars: a.vars.clone() } } else { let (x, y) = a.to_combined_vars(b); Dual { real: x.real + y.real, dual: x.dual + y.dual, vars: x.vars } } });
(注:原代码存在语法缺失,已补全impl Dual的闭合花括号,并导入Array)
核心疑问与需求
- 为什么Rust中
长度判断+逐元素比对的变量相等逻辑,会在变量完全一致的场景下比Python的Tuple相等判断慢这么多? - 针对Rust的这一性能瓶颈,有哪些具体的优化方案?
内容的提问来源于stack exchange,提问作者Attack68
相关产品推荐
相关产品推荐

