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实例,需要手动构造模型:
- 给
ModelJson派生序列化/反序列化 trait:
use serde::{Serialize, Deserialize}; #[derive(Serialize, Deserialize)] struct ModelJson { coefficients: Vec<f64>, intercept: f64, }
- 加载参数并构建模型:
// 读取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库,需补充以下配置和代码:
- 更新Cargo.toml,添加serde依赖并开启derive特性:
serde = { version = "1.0", features = ["derive"] } linfa-linear = { version = "0.7.0", features = ["serde"] }
- 正确编写序列化/反序列化代码:
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
相关产品推荐
相关产品推荐

