Rust中如何实现无unsafe、性能达标的地道Floyd-Warshall算法
方案1:添加前置断言消除边界检查
你只需要在进入循环前添加断言确认所有邻接矩阵行的长度都等于顶点数n,编译器在开启release优化时就会自动消除后续所有下标访问的边界检查,不需要使用unsafe,性能和你写的unsafe版本完全一致:
fn floyd_warshall_safe(dist: &mut [Vec<i32>]) { let n = dist.len(); // 给编译器明确提示:所有行的长度都等于n assert!(dist.iter().all(|row| row.len() == n)); for i in 0..n { // 提前提取第i行的不可变引用,避免重复索引外层数组 let row_i = &dist[i]; for j in 0..n { // 提前提取第j行的可变引用,减少外层数组索引次数 let row_j = &mut dist[j]; for k in 0..n { // 编译器已经明确k、i都小于n,不会越界,直接跳过边界检查 row_j[k] = row_j[k].min(row_j[i] + row_i[k]); } } } }
这种写法是地道的Rust安全写法,没有任何unsafe代码,逻辑也和原生实现完全一致,可读性没有损失。
方案2:改用一维连续存储的邻接矩阵,进一步提升性能
嵌套Vec<Vec<i32>>的每一行存储地址不一定连续,缓存命中率低。你可以把邻接矩阵改成连续存储的一维数组,长度为n*n,不仅更方便编译器优化,缓存友好性也更高,实际运行速度通常比嵌套Vec的unsafe版本还要快:
fn floyd_warshall_1d(dist: &mut [i32], n: usize) { // 确认数组长度符合要求,后续访问不会越界 assert!(dist.len() == n * n); for i in 0..n { for j in 0..n { for k in 0..n { let idx_jk = j * n + k; let idx_ji = j * n + i; let idx_ik = i * n + k; dist[idx_jk] = dist[idx_jk].min(dist[idx_ji] + dist[idx_ik]); } } } }
使用这个版本时,你只需要把原来的二维坐标[j][k]转换成j * n + k的一维坐标即可。
额外优化提示
在release模式编译时,你可以在Cargo.toml中添加以下配置,开启更激进的优化,性能还能有10%~20%的提升:
[profile.release] lto = "fat" codegen-units = 1 opt-level = 3
内容的提问来源于stack exchange,提问作者Borys
相关产品推荐
相关产品推荐

