Rust实现高斯约旦消元法结果异常排查及代码优化咨询
高斯约旦消元法Rust版本结果不符问题排查与优化
问题描述
将自行编写的Python高斯约旦消元法代码转译为Rust版本后,代码运行无报错与警告,但输出结果与正确值不符。怀疑是精度舍入问题但无法定位原因,同时需要该Rust代码的优化方案。
Python原代码
# Python equivalent the Rust version was transcoded from by me import numpy as np M = np.array([[8, -8, -9, -1], [-10, 15, -9, -25], [-9, -1, 7, 3]], float) if M[0][0] == 0: M[0] += 1 for j in range(len(M)): # For each column if M[j][j] != 1: # If the jth element of that column is not 1 M[j] = M[j] / M[j][j] # Then divide the row by that element for i in range(len(M)): # For each row if i != j: # If we are not at the row we want to have the pivot 1 M[i] -= M[j] * M[i][j] np.set_printoptions(precision=20) display(M[:, -1])
Rust转译代码
fn main () { let mut m = [[8.0, -8.0, -9.0, -1.0], [-10.0, 15.0, -9.0, -25.0], [-9.0, -1.0, 7.0, 3.0]]; if m[0][0] == 0.0 { for i in 0..m[0].len() { m[0][i] += 1.0; } } for c in 0..m.len() { if m[c][c] != 1.0 { for e in 0..m[c].len() { m[c][e] = m[c][e] / m[c][c]; } } for r in 0..m.len() { if r != c { for e in 0..m[r].len() { m[r][e] -= m[c][e] * m[r][c] } } } } for element in m { for element_ in element { print!("{}", element_); } print!("\n"); } }
正确结果
0.6303724928366758, -0.5501432664756454, 1.160458452722063
问题定位
并非精度问题,而是逻辑错误:
- Python中
M[i] -= M[j] * M[i][j]是向量级操作,会先基于消元前M[i][j]的原始值计算完整的M[j] * M[i][j]向量,再整体从M[i]中减去。 - Rust代码中,遍历行元素时,当
e == c会将m[r][c]修改为0(因m[c][c]已归一化为1),后续e > c的元素计算时,使用的是已被修改为0的m[r][c]而非原始消元系数,导致消元不彻底,结果错误。
修复后的Rust代码
核心修改是提前保存消元系数factor,避免实时修改导致的错误:
fn main() { let mut m = vec![ vec![8.0, -8.0, -9.0, -1.0], vec![-10.0, 15.0, -9.0, -25.0], vec![-9.0, -1.0, 7.0, 3.0], ]; // 正确处理主元为0的情况:寻找非零主元行交换 for c in 0..m.len() { // 查找当前列下方的非零主元行 let pivot_row = (c..m.len()) .find(|&r| m[r][c].abs() > 1e-9) .expect("Matrix is singular (no unique solution)"); if pivot_row != c { m.swap(c, pivot_row); } // 归一化主元行 let pivot_val = m[c][c]; for e in 0..m[c].len() { m[c][e] /= pivot_val; } // 消去其他行的当前列元素 for r in 0..m.len() { if r != c { let factor = m[r][c]; // 提前保存原始消元系数 for e in 0..m[r].len() { m[r][e] -= m[c][e] * factor; } } } } // 输出结果(保留16位精度) println!("计算结果:"); for row in m { println!("{:.16}", row.last().unwrap()); } }
代码优化方案
- 使用动态数组
Vec<Vec<f64>>:替代固定大小数组,支持任意维度矩阵,灵活性更高。 - 正确处理主元为0的场景:原代码直接修改整行会改变方程,正确做法是寻找下方非零主元行交换,避免矩阵奇异时崩溃。
- 浮点精度判断:用
abs() > 1e-9替代== 0.0或==1.0,避免浮点精度误差导致的逻辑错误。 - 迭代器简化代码:利用Rust迭代器特性替代手动索引循环,代码更简洁易读(如寻找主元行的逻辑)。
- 输出精度控制:通过
{:.16}格式化输出,避免科学计数法或精度丢失。
内容的提问来源于stack exchange,提问作者Arcturus
相关产品推荐
相关产品推荐

