在Rust Polars中计算DataFrame的协方差矩阵
计算Polars DataFrame的协方差矩阵
原生Polars实现(无需转换ndarray)
Polars的协方差矩阵计算功能需要启用stats crate特性,启用后可直接调用DataFrame.cov方法得到M×M维度的协方差矩阵(对应输入N×M的DataFrame):
use polars::prelude::*; fn covariance_matrix(df: &DataFrame) -> Result<DataFrame, PolarsError> { // None表示使用无偏估计(除以N-1),传入Some(1.0)则为有偏估计(除以N) df.cov(None) }
- 需在
Cargo.toml中配置Polars特性:polars = { version = "0.35", features = ["stats"] }
优化你的ndarray转换方案
你之前的代码返回的二维数组(M×M)正是N×M DataFrame对应的协方差矩阵维度(协方差矩阵维度为特征数×特征数)。若要实现泛型支持,可调整代码如下:
use polars::prelude::*; use ndarray::Array2; use ndarray_stats::CorrelationExt; fn covariance_matrix<T>(matrix: &DataFrame) -> Result<Array2<f64>, PolarsError> where T: PolarsNumericType, <T as PolarsNumericType>::Native: num_traits::Float + ndarray::ScalarOperand, { let matrix_arr = matrix.to_ndarray::<T>(IndexOrder::C)?; // 用ok_or_else替代unwrap,避免panic let covariance = matrix_arr.cov(1.0) .ok_or_else(|| PolarsError::ComputeError("计算协方差矩阵失败".into()))?; Ok(covariance) }
这里的泛型约束保证了输入列是浮点类型,同时兼容ndarray_stats的计算要求。
内容的提问来源于stack exchange,提问作者transonic912
相关产品推荐
相关产品推荐

