You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Rust中nalgebra DMatrix调用smartcore train_test_split报错求助

问题:使用smartcore的train_test_split处理nalgebra矩阵时的Trait绑定错误

我正在开发一个小型机器学习应用,从CSV文件读取数据并转换为nalgebra库的DMatrix。想借助smartcore库的train_test_split函数拆分数据集为训练集和测试集,但调用该函数时遇到编译错误,求错误原因及解决方法。

我的代码

use std::error::Error;
use std::io::BufReader;
use std::io::BufRead;
use std::fs::File;
use nalgebra::DMatrix;
use std::str::FromStr;

use smartcore::model_selection::train_test_split;

fn read_csv(input: &mut dyn BufRead) -> Result<DMatrix<f64>, Box<dyn Error>> {

    let mut samples = Vec::new();

    let mut rows = 0;

    for line in input.lines().skip(1){
        rows += 1;

        for data in line?.split_terminator(",") {

            let a = f64::from_str(data.trim());

            match a {
                Ok(value) => samples.push(value),
                Err(e) => println!("Error parsing data in row: {}", rows),
            }
        }
    }

    let cols = samples.len() / rows;

    Ok(DMatrix::from_row_slice(rows, cols, &samples[..]))

}

fn main() -> Result<(), Box<dyn Error>> {

    //Load CSV
    let file = File::open("dataset/heart.csv").unwrap();
    let data: DMatrix<f64> = read_csv(&mut BufReader::new(file)).unwrap();

    let x = data.columns(0, 13).into_owned();
    let y = data.column(13).into_owned();

    // ERROR
    let (x_train, x_test, y_train, y_test) = train_test_split(&x, &y.transpose(), 0.2, true);

    println!("{:?}", x_train);
    
    Ok(())

}

错误信息

error[E0277]: the trait bound `nalgebra::Matrix<f64, Dyn, Dyn, VecStorage<f64, Dyn, Dyn>>: smartcore::linalg::Matrix<_>` is not satisfied
   --> src/main.rs:53:63
    |
53  |     let (x_train, x_test, y_train, y_test) = train_test_split(&x, &y.transpose(), 0.2, true);
    |                                              ---------------- ^^ the trait `smartcore::linalg::Matrix<_>` is not implemented for `nalgebra::Matrix<f64, Dyn, Dyn, VecStorage<f64, Dyn, Dyn>>`
    |                                              |
    |                                              required by a bound introduced by this call
    |
    = help: the following other types implement trait `smartcore::linalg::Matrix<T>`:
              DenseMatrix<T>
              nalgebra::base::matrix::Matrix<T, nalgebra::base::dimension::Dynamic, nalgebra::base::dimension::Dynamic, nalgebra::base::vec_storage::VecStorage<T, nalgebra::base::dimension::Dynamic, nalgebra::base::dimension::Dynamic>>
note: required by a bound in `train_test_split`

    |
133 | pub fn train_test_split<T: RealNumber, M: Matrix<T>>(
    |                                           ^^^^^^^^^ required by this bound in `train_test_split`

error[E0277]: the trait bound `nalgebra::Matrix<f64, Dyn, Dyn, VecStorage<f64, Dyn, Dyn>>: BaseMatrix<_>` is not satisfied
  --> src/main.rs:53:46
   |
53  |     let (x_train, x_test, y_train, y_test) = train_test_split(&x, &y.transpose(), 0.2, true);
   |                                              ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ the trait `BaseMatrix<_>` is not implemented for `nalgebra::Matrix<f64, Dyn, Dyn, VecStorage<f64, Dyn, Dyn>>`
   |
   = help: the following other types implement trait `BaseMatrix<T>`:
             DenseMatrix<T>
             nalgebra::base::matrix::Matrix<T, nalgebra::base::dimension::Dynamic, nalgebra::base::dimension::Dynamic, nalgebra::base::vec_storage::VecStorage<T, nalgebra::base::dimension::Dynamic, nalgebra::base::dimension::Dynamic>>

For more information about this error, try `rustc --explain E0277`.
error: could not compile `logistic-regression` due to 2 previous errors

错误原因及解决方法

错误原因

smartcore的train_test_split函数要求输入矩阵实现smartcore::linalg::Matrix trait,虽然错误提示显示nalgebra的动态矩阵理论上支持该trait,但实际编译失败的核心原因是未启用smartcore的nalgebra特性,或依赖版本不兼容。

解决步骤

  1. 启用smartcore的nalgebra特性
    修改Cargo.toml中的smartcore依赖配置,明确启用nalgebra支持:

    [dependencies]
    nalgebra = "0.32.3"
    smartcore = { version = "0.3.2", features = ["nalgebra"] }
    std = { version = "1.0", features = ["full"] }
    

    注意:请保持nalgebra和smartcore的版本兼容,建议使用最新稳定版本。

  2. 修正y参数的维度传递
    原代码将列向量y转置为行矩阵传递,但train_test_split支持直接接收nalgebra的列向量(DVector),无需转置。修改main函数中的调用代码:

    let (x_train, x_test, y_train, y_test) = train_test_split(&x, &y, 0.2, true);
    
  3. 优化CSV读取的错误处理(可选)
    原代码中解析数据错误仅打印提示,未终止流程,可能导致后续构造矩阵时出现维度不匹配问题。建议将解析错误向上返回:

    // 替换原match逻辑
    let a = f64::from_str(data.trim())?;
    samples.push(a);
    

内容的提问来源于stack exchange,提问作者Tutu ツ

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.09 21:54:56