如何加速Rust递归函数?求对应缓存优化方案(附Python实现)
在Rust中优化递归函数的重复计算问题
问题背景
我在学习Rust时写了一个递归函数f(n),调用f(50)时速度极慢,运行两小时都没出结果。排查后发现是大量重复计算导致的——类似场景在Python里用functools.lru_cache装饰器就能轻松实现缓存优化,想知道Rust里怎么实现同样的效果。
原Rust代码
fn f(n: i64) -> i64 { if n == 0 { return 1; } if n > 0 { return 2 * f(1 - n) + 3 * f(n - 1) + 2; } -1 * f(-1 * n) } fn main() { println!("{}", f(50)) }
参考的Python缓存实现
from functools import lru_cache @lru_cache def f(n): if n == 0: return 1 return 2 * f(1 - n) + 3 * f(n - 1) + 2 if n > 0 else -1 * f(-n) res = 0 for s in str(f(50)): res += int(s) print(res)
解决方法
1. 手动用HashMap实现缓存
Rust标准库没有内置缓存装饰器,最直接的方式是手动用HashMap存储已计算的结果,避免重复递归:
use std::collections::HashMap; // 把缓存作为可变引用传入函数 fn f(n: i64, cache: &mut HashMap<i64, i64>) -> i64 { // 先查缓存,命中直接返回 if let Some(&val) = cache.get(&n) { return val; } let result = match n { 0 => 1, positive if positive > 0 => 2 * f(1 - positive, cache) + 3 * f(positive - 1, cache) + 2, negative => -1 * f(-negative, cache), }; // 计算完成后存入缓存 cache.insert(n, result); result } fn main() { let mut cache = HashMap::new(); let result = f(50, &mut cache); // 可选:和Python代码一致,计算结果的各位数字之和 let digit_sum: u64 = result.to_string() .chars() .map(|c| c.to_digit(10).unwrap() as u64) .sum(); println!("f(50) = {}", result); println!("各位数字之和 = {}", digit_sum); }
2. 用第三方库简化缓存逻辑
如果不想手动处理缓存,可以用社区成熟的cached库,它提供了类似Python装饰器的语法:
首先在Cargo.toml添加依赖:
[dependencies] cached = "0.47"
然后修改代码:
use cached::proc_macro::cached; // 用#[cached]装饰函数,自动实现缓存 #[cached] fn f(n: i64) -> i64 { if n == 0 { return 1; } if n > 0 { return 2 * f(1 - n) + 3 * f(n - 1) + 2; } -1 * f(-n) } fn main() { let result = f(50); let digit_sum: u64 = result.to_string() .chars() .map(|c| c.to_digit(10).unwrap() as u64) .sum(); println!("f(50) = {}", result); println!("各位数字之和 = {}", digit_sum); }
3. 改写为迭代实现(更高效)
递归加缓存虽然能解决问题,但递归本身存在栈开销。进一步分析函数的递推关系,可以将其改写为迭代形式,效率更高:
观察函数规则:
- 当
n < 0时,f(n) = -f(-n),只需计算正数结果再取反 - 当
n = 1时,f(1) = 2*f(0) + 3*f(0) + 2 = 7 - 当
n >= 2时,1-n是负数,代入规则得f(1-n) = -f(n-1),因此原式可化简为:f(n) = 2*(-f(n-1)) + 3*f(n-1) + 2 = f(n-1) + 2
基于这个简化的递推关系,迭代实现如下:
fn f_iter(n: i64) -> i64 { if n == 0 { return 1; } let target = n.abs(); let mut current = 1; // f(0) = 1 for k in 1..=target { current = if k == 1 { 2 * current + 3 * current + 2 // 计算f(1) } else { current + 2 // f(k) = f(k-1) + 2 }; } if n < 0 { -current } else { current } } fn main() { let result = f_iter(50); let digit_sum: u64 = result.to_string() .chars() .map(|c| c.to_digit(10).unwrap() as u64) .sum(); println!("f(50) = {}", result); println!("各位数字之和 = {}", digit_sum); }
这个迭代版本的时间复杂度是O(n),完全没有递归栈的开销,计算f(50)几乎是瞬间完成的。
内容的提问来源于stack exchange,提问作者Wilper
相关产品推荐
相关产品推荐

