news 2026/8/8 10:33:45

【RustyML入门】1.4. 第一个端到端模型

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
【RustyML入门】1.4. 第一个端到端模型

1.4. 第一个端到端模型

1.4.1. 完整程序

usendarray::{Array1,Array2};userustyml::machine_learning::LogisticRegression;userustyml::metrics::{ConfusionMatrix,accuracy};userustyml::set_global_seed;userustyml::utils::StandardScaler;userustyml::utils::train_test_split::train_test_split;fnmain(){// 设置全局种子,保证可复现性set_global_seed(42);// 造一个合成的二分类数据集letn_per_class=30usize;letmutrows:Vec<f64>=Vec::with_capacity(2*n_per_class*2);letmutlabels:Vec<f64>=Vec::with_capacity(2*n_per_class);foriin0..n_per_class{leta=(i%6)asf64;letb=(i/6)asf64;// 类别0中心在-0.6附近// 类别1在+0.6附近rows.push(-0.6+0.3*a);rows.push((-0.6+0.3*b)*100.0);labels.push(0.0);rows.push(0.6+0.3*a);rows.push((0.6+0.3*b)*100.0);labels.push(1.0);}letn=labels.len();letx:Array2<f64>=Array2::from_shape_vec((n,2),rows).unwrap();lety:Array1<f64>=Array1::from(labels);// 划分训练和测试集let(x_train,x_test,y_train,y_test)=train_test_split(x,y,Some(0.25),Some(42)).unwrap();// 标准化:统计量只在训练集上拟合,然后应用到两个划分letmutscaler=StandardScaler::new();letx_train_std=scaler.fit_transform(&x_train).unwrap();letx_test_std=scaler.transform(&x_test).unwrap();// 训练letmutmodel=LogisticRegression::new(true,0.5,1000,1e-6).unwrap();model.fit(&x_train_std,&y_train).unwrap();// 在测试集上预测letpreds:Array1<f64>=model.predict(&x_test_std).unwrap();// 评估letacc=accuracy(&y_test,&preds);letcm=ConfusionMatrix::new(&y_test,&preds);let(tp,fp,tn,fn_)=cm.get_counts();println!("test accuracy = {:.3}",acc);println!("TP={tp} FP={fp} TN={tn} FN={fn_}");println!("{}",cm.summary());println!("iterations run: {:?}",model.get_actual_iterations());// 保存模型与加载模型letpath="logreg_model.bin";model.save_to_path(path).unwrap();letloaded=LogisticRegression::load_from_path(path).unwrap();letpreds_loaded=loaded.predict(&x_test_std).unwrap();println!("round-trip predictions identical: {}",preds==preds_loaded);}

这些println!会打印一份小报告,格式类似于:

test accuracy = 0.XXX TP=.. FP=.. TN=.. FN=.. Confusion Matrix: +-----------------+--------------------+--------------------+ | | Predicted Positive | Predicted Negative | +-----------------+--------------------+--------------------+ | Actual Positive | TP: .. | FN: .. | | Actual Negative | FP: .. | TN: .. | +-----------------+--------------------+--------------------+ ... derived metrics ... iterations run: Some(..) round-trip predictions identical: true

1.4.2. 设置种子以确保可复现性

rustyml::set_global_seed接收一个u64。它设置一个线程局部的种子,同一线程上此后构造的每个没设种子(random_state == None)的组件,都会以它为源派生自己的RNG,模仿Keras风格的全局种子模型,详见7.1. 可复现性与随机种子。

在这条流水线里,它对数字没有任何影响。LogisticRegression::fit就是普通的全批量梯度下降,毫无随机性(对同一份数据拟合两次,得到的权重逐位相同),standardize同样是确定的。唯一带随机性的步骤是train_test_split内部的洗牌,而我给它显式传了Some(42)。按局部种子(如果设置了)大于全局的原则,这部分改这个局部种子。建议养成在代码开头设置set_global_seed的习惯。

1.4.3. 数据集与标签规范

特征是形状为(n_samples, n_features)Array2<f64>,标签是Array1<f64>。标签的元素类型不是随手定的,LogisticRegression::fit要求y是一个f64数组,取值必须恰好0.01.0(你可以在它的文档注释中看到这条规定),其他任何值都会返回Error::InvalidInputfit里没有内置的LabelEncoder步骤,如果你的标签是字符串或任意整数,得自己转换,详见4.3. 标签编码。

1.4.4. 划分数据

let (x_train, x_test, y_train, y_test) = train_test_split(x, y, Some(0.25), Some(42)).unwrap();

train_test_split会消耗xy,如果你在这之后还需要原始数据就需要提前克隆。test_size参数传None表示0.3(70/30 划分),random_stateNone表示从系统熵取种子(若设了全局种子,则从全局种子取)。

空数据集是会返回Error::EmptyInputx/y长度不一致会返回Error::DimensionMismatchtest_size值在(0, 1)之外是Error::InvalidParameter,数据集小到无法划分会返回Error::InvalidInput。函数会在两侧各保留至少一个样本,比如对十行数据取test_size = 0.99,训练侧仍会剩1行而不是0行。

类别不平衡时,普通的随机划分可能把某个类别全部划分到测试集或是全部划分到训练集。这时候最好使用train_test_split_stratified,它对每个类别独立划分。

划分的具体机制和分层变体详见4.1. 训练集与测试集划分。

1.4.5. 标准化时不泄露测试信息

逻辑回归靠梯度下降拟合,而梯度下降在意特征的尺度。上面的代码构造出的特征2比特征1大一百倍,这会导致梯度被特征2主导,需要把每个特征标准化到零均值、单位方差。

有两种方法做标准化:

  • 调用函数standardize(data, axis):它是无状态、一次性的变换,它从你传进去的那个数组算出均值和标准差,返回一个新的标准化数组,什么都不保留。
  • 使用StandardScaler结构体以及其方法:它能够存储状态,调用fit会计算统计量并把它们存下来,transform把存储好的统计量施加到之后交给它的任何数组上。

用函数时,常见的错误的做法是对x_train调一次standardize,再对x_test调一次,这样会导致训练集和测试集落到两套不同的尺度上。在划分之前就标准化整个矩阵会把测试集的均值和方差揉进模型训练所用的数字里,导致准确率虚高。

正确的方法是使用StandardScaler结构体以及其方法,在训练集上拟合一次,然后变换所有东西:

use rustyml::utils::StandardScaler; let mut scaler = StandardScaler::new(); let x_train_std = scaler.fit_transform(&x_train).unwrap(); // 在这里从训练集计算到均值/标准差 let x_test_std = scaler.transform(&x_test).unwrap(); // 对测试集套用相同的统计量

只在训练集上使用fit/fit_transform,其余用transform。统计量是可以获取的(scaler.get_mean()get_scale())。 除数是总体标准差(ddof = 0,与scikit-learn的StandardScaler一致)。

当你不需要保持统计量一致,或者你需要结构体不覆盖的Row/Global轴时就可以用函数:

usendarray::array;userustyml::utils::standardize::{standardize,StandardizationAxis};fnmain(){letx=array![[1.0,200.0],[3.0,400.0],[5.0,600.0]];letz=standardize(&x,StandardizationAxis::Column).unwrap();println!("standardized shape: {:?}",z.dim());println!("{:?}",z);}

输出:

standardized shape: (3, 2) [[-1.224744871391589, -1.224744871391589], [0.0, 0.0], [1.224744871391589, 1.224744871391589]], shape=[3, 2], strides=[2, 1], layout=Cc (0x5), const ndim=2

更完整的教程见4.2. 标准化与归一化。

1.4.6. 训练模型

let mut model = LogisticRegression::new(true, 0.5, 1000, 1e-6).unwrap(); model.fit(&x_train_std, &y_train).unwrap();

LogisticRegression::new返回Result<Self, Error>,因为它会先校验超参数,如果不符合要求就会返回错误Error::InvalidParameter。你也可以使用LogisticRegression::default()来使用设定了默认值的模型(fit_intercept = truelearning_rate = 0.01max_iter = 100tol = 1e-4)。相比于默认值new里传的学习率和迭代预算都调大了,因为标准化后的特征容得下更大的步长。正则化默认关闭,如果需要可以用.with_regularization(RegularizationType::L2(alpha))打开,这个方法的返回值同样被Result包裹(负的或非有限的alpha会返回错误)。RustyML的惩罚强度计算与scikit-learn的SGDClassifier/SGDRegressor保持一致。LogisticRegression(C=c)则对应alpha = 1 / (c * n)n是训练样本数。完整的换算表就写在RegularizationType自己的文档注释里。

fit接受特征矩阵和目标向量的引用。它的返回值被Result包裹,用于防范潜在的错误:

  • 0.0/1.0标签检查
  • 空输入(Error::EmptyInput
  • x/y长度不一致(Error::DimensionMismatch
  • 非有限的特征值(Error::NonFinite)。如果权重或损失溢出,它也会在循环中途触发Error::NonFinite

当相邻迭代间的损失变化在tol以下,或者迭代次数等于max_iter之后就会停止迭代。你可以通过model.get_actual_iterations()知道是哪种情况停止的迭代:如果它等于max_iter,说明模型到了迭代次数上限而非收敛,你该调高max_iter或重新设置学习率。更多关于逻辑回归的内容详见2.2. 逻辑回归。

1.4.7. 预测

let preds: Array1<f64> = model.predict(&x_test_std).unwrap();

predict返回Result<Array1<f64>, Error>,每个元素是硬类别标签0.01.0(sigmoid概率在0.5处阈值化)。如果你要的是底层概率而非标签就调用predict_proba

predict的返回值同样被Result包裹,用于防范潜在的错误:

  • fit之前调用它是Error::NotFitted
  • 传入的特征数与模型训练时的不同是Error::DimensionMismatch

1.4.8. 准确率与混淆矩阵

let acc = accuracy(&y_test, &preds); let cm = ConfusionMatrix::new(&y_test, &preds); let (tp, fp, tn, fn_) = cm.get_counts(); println!("{}", cm.summary());

accuracy(y_true, y_pred) -> f64计算准确率,也就是标签相符的比例。通过ConfusionMatrix::new(y_true, y_pred)构建混淆矩阵,get_counts()返回(tp, fp, tn, fn)。混淆矩阵只收标签,两个数组的每个元素都必须恰好是0.01.0,其余一律panic。若标签用的是另一对取值——比如间隔分类器给出的-1.0/+1.0——就用ConfusionMatrix::new_with_labels(y_true, y_pred, negative_label, positive_label),对应scikit-learn的confusion_matrix(..., labels=[neg, pos])。它的方法有accuracy()precision()recall()specificity()f1_score()balanced_accuracy()mcc(),也可以打印summary()得到前面那张格式化表格。数据不平衡时,别用原始准确率,应该使用balanced_accuracy()mcc()

更多的分类指标详见5.2. 分类指标。

1.4.9. 保存与加载模型

model.save_to_path("logreg_model.bin").unwrap(); let loaded = LogisticRegression::load_from_path("logreg_model.bin").unwrap();

save_to_path(&self, path: &str) -> Result<(), Error>用postcard把整个模型(权重、截距开关、学习率、迭代次数、正则化设置)序列化成二进制。通过load_from_path(path: &str) -> Result<Self, Error>可以读取。详细内容见3.9. 权重保存与加载与7.2. 深入模型持久化。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/8/8 10:30:23

高性能中国车牌生成系统架构设计与技术实现方案

高性能中国车牌生成系统架构设计与技术实现方案 【免费下载链接】chinese_license_plate_generator 中国车牌生成器 项目地址: https://gitcode.com/gh_mirrors/ch/chinese_license_plate_generator 中国车牌生成器是一个基于Python的高性能开源工具&#xff0c;专门为计…

作者头像 李华
网站建设 2026/8/8 10:28:29

AI模型推理加速:从量化、图优化到LLM KV Cache的工程实践

这次我们来看一个非常硬核的话题&#xff1a;AI讲AI第76期&#xff1a;推理加速的工程智慧。这期内容不是介绍某个具体的开源模型或工具&#xff0c;而是聚焦于一个更底层、更关键的问题&#xff1a;如何让已经训练好的AI模型&#xff0c;在实际部署时跑得更快、更省资源。对于…

作者头像 李华
网站建设 2026/8/8 10:28:25

AI编程助手实战:从概念到应用,手把手教你提升开发效率

1. 背景与核心概念&#xff1a;当AI遇上少儿编程 最近&#xff0c;一个“8岁孩子用AI Studio四小时自制平台游戏”的案例在技术社区和家长圈里引起了不小的讨论。这听起来有些不可思议&#xff0c;但背后反映的是一个正在发生的技术趋势&#xff1a;AI驱动的低代码/无代码开发…

作者头像 李华
网站建设 2026/8/8 10:27:40

大语言模型指令遵循与自主判断的平衡策略与工程实践

在人工智能技术快速发展的今天&#xff0c;大语言模型&#xff08;LLM&#xff09;的指令遵循能力是其核心价值之一。然而&#xff0c;开发者和研究人员在实际应用中发现&#xff0c;模型在严格遵循用户指令与进行必要的自主判断之间&#xff0c;常常存在一种微妙的张力。这种矛…

作者头像 李华
网站建设 2026/8/8 10:26:54

Claude Code禁用确认提示的3种方法与安全实践

1. 理解Claude Code的交互机制Claude Code作为一款AI编程助手工具&#xff0c;其默认的交互设计确实存在一些不够便捷的地方。最典型的就是每次执行代码时都需要手动输入"yes"确认&#xff0c;这个设计本意是为了防止误操作&#xff0c;但在高频使用场景下反而成了效…

作者头像 李华