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

Rust-linfa中线性回归模型保存与加载的实现问题

Linfa线性回归模型保存与加载问题

我正在用Rust的linfa库做机器学习开发,重点使用Linear Regression(线性回归)模型。现在想实现训练好的模型的保存和加载,但试了两种方法都没成功。

方案一:手动提取参数序列化

我尝试提取linfa线性回归的核心训练参数,存入自定义结构体后用serde_json保存为JSON,但加载后无法还原成可用的模型进行预测。

实现细节

  • 存储参数的自定义结构体:
struct ModelJson {
    coefficients: Vec<f64>,
    intercept: f64,
}
  • 存储流程代码:
let model = lin_reg.fit(&dataset)?;
let model_json = ModelJson {
    coefficients: model.params().to_vec(),
    intercept: model.intercept(),
};
// 序列化保存
let json_str = serde_json::to_string(&model_json).unwrap();
std::fs::write("model.json", json_str).unwrap();
  • 保存后的JSON示例:
{"coefficients":[-0.00017907873576254802,-0.00100659702068151,-0.0008275037845519519,0.0004613216043979551,0.0010300634934599436],"intercept":50.525680622870084}

修复方法

问题在于加载JSON后,无法直接将参数映射回FittedLinearRegression实例,需要手动构造模型:

  1. 给ModelJson派生序列化/反序列化 trait:
use serde::{Serialize, Deserialize};

#[derive(Serialize, Deserialize)]
struct ModelJson {
    coefficients: Vec<f64>,
    intercept: f64,
}
  1. 加载参数并构建模型:
// 读取JSON文件
let json_str = std::fs::read_to_string("model.json").unwrap();
let model_json: ModelJson = serde_json::from_str(&json_str).unwrap();

// 将Vec<f64>转为linfa依赖的Array1<f64>
use ndarray::Array1;
let params = Array1::from_vec(model_json.coefficients);

// 构造FittedLinearRegression实例
use linfa_linear::FittedLinearRegression;
let loaded_model = FittedLinearRegression {
    intercept: model_json.intercept,
    params,
};

// 用加载后的模型预测
let predictions = loaded_model.predict(&test_dataset);

注意:需确保ndarray依赖已添加到Cargo.toml中。


方案二:使用linfa-linear的serde特性序列化

我知道linfa支持模型序列化,尝试开启linfa-linear的serde特性,但编译报错。

实现细节

  • Cargo.toml依赖配置:
linfa-linear = {version="0.7.0", features=["serde"]}
  • 序列化代码:
let model = lin_reg.fit(&dataset)?;
let serialized = serde_json::to_string(&model).unwrap();
  • 编译错误:
the trait bound `FittedLinearRegression<f64>: serde::ser::Serialize` is not satisfied
the following other types implement trait `serde::ser::Serialize`:
  bool
  char
  isize
  i8
  i16
  i32
  i64
  i128
and 133 othersrustcClick for full compiler diagnostic
main.rs(82, 22): required by a bound introduced by this call

修复方法

报错原因是缺少serde的derive支持,且linfa-linear的serde特性依赖外部serde库,需补充以下配置和代码:

  1. 更新Cargo.toml,添加serde依赖并开启derive特性:
serde = { version = "1.0", features = ["derive"] }
linfa-linear = { version = "0.7.0", features = ["serde"] }
  1. 正确编写序列化/反序列化代码:
use serde::{Serialize, Deserialize};
use linfa_linear::FittedLinearRegression;

// 序列化模型
let model: FittedLinearRegression<f64> = lin_reg.fit(&dataset)?;
let serialized = serde_json::to_string(&model).expect("序列化模型失败");
std::fs::write("model.json", serialized).expect("写入模型文件失败");

// 反序列化模型
let serialized_str = std::fs::read_to_string("model.json").expect("读取模型文件失败");
let loaded_model: FittedLinearRegression<f64> = serde_json::from_str(&serialized_str).expect("反序列化模型失败");

// 使用模型预测
let predictions = loaded_model.predict(&test_dataset);

如果追求更高的序列化效率,也可以用bincode做二进制序列化:

  • 添加bincode依赖:
bincode = "1.3"
  • 序列化/反序列化代码:
// 序列化
let serialized = bincode::serialize(&model).unwrap();
std::fs::write("model.bincode", serialized).unwrap();

// 反序列化
let loaded_model: FittedLinearRegression<f64> = bincode::deserialize(&std::fs::read("model.bincode").unwrap()).unwrap();

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 03:57:47