3.9. 权重保存与加载RustyML 只用 2 个方法持久化训练好的Sequential模型save_to_path和load_from_path。这两个方法的职责刻意收得很窄。Keras 的model.save()会写出一个自描述的打包文件能重建计算图、编译配置和优化器状态。RustyML 不这么做它只保存各层的权重文件里没有足够的信息能重建出一个模型。加载分两步第一步用代码搭出一模一样的层栈第二步把保存下来的数组加载进这个层栈。本节讲清楚文件里存了什么、这条边界为什么划在这里、有哪些出错方式以及“只存权重”模型的工作流范式。3.9.1. 磁盘上究竟写了什么save_to_path遍历所有层。对每一层它记录一小段元数据标签外加这一层的权重。然后它用 postcard 把整个向量序列化成紧凑的二进制块再由缓冲写入器把这块数据写盘。load_from_path把文件读回来、反序列化核对你搭出的模型是否与文件吻合然后逐层把数组应用回去。函数签名如下pub fn save_to_path(self, path: impl AsRefstd::path::Path) - RustymlResult(); pub fn load_from_path(mut self, path: impl AsRefstd::path::Path) - RustymlResult();两个方法都接受任何实现了AsRefPath的类型str、String、Path或PathBuf。下面示例里的model.bin字面量只是最常见的写法。扩展名对格式没有任何影响因为不管你给文件起什么名字postcard 写进去的都是裸字节。save_to_path用的是File::create它会截断并覆盖文件所以往同一路径再次保存会替换掉上一个 checkpoint。这个行为正好适合“只保留最优”这种循环。逐层元数据只有一个作用校验而不是重建。每一层贡献一个类型名字符串比如Dense或Conv2D再加一个输出形状字符串。这些字符串是校验用的标签不是搭建层的配方。文件不记录任何一层的激活函数也不记录epsilon、momentum、步长、卷积核大小、膨胀系数或分组数。你没法把一个空的Sequential丢给load_from_path就换回一个能用的模型。RustyML 把这种格式称为“只存权重”正是因为这一点架构活在你的源码里文件只承载填充这套架构所需的那些数字。写进文件的不写进文件的逐层的类型名标签加载时校验激活函数、超参数epsilon、momentum、步长……逐层的输出形状标签仅供参考优化器及其累积状态Adam 的矩、SGD 的动量每一层的权重数组见 3.9.2损失函数与编译状态BatchNormalization 的运行均值与运行方差训练时的打乱种子、训练模式标志3.9.2.LayerWeight枚举逐层载荷磁盘上的权重格式是一个 Rust 枚举LayerWeighta。它为每种受支持的层类型设一个变体另有一个Empty变体留给无参数的层。每个变体包裹一个小结构体其中的数组以Cow存储。Cow让同一个类型能兼顾两个方向保存时Sequential::get_weights借用活着的数组Cow::Borrowed不做克隆加载时load_from_path反序列化成拥有所有权的数组Cow::Owned。枚举采用 serde 默认的外部标签externally tagged表示因为 postcard 不是自描述格式需要把判别标记显式写出来。每个变体的载荷决定了一次往返是否精确变体载荷Denseweight (in, out)、bias (1, out)SimpleRNNkernel、recurrent_kernel、biasLSTM/GRU融合的kernel、recurrent_kernel、bias门块[i|f|g|o]/[z|r|h]Conv1D/Conv2D/Conv3D卷积weight核与biasSeparableConv2Ddepthwiseweight核、pointwiseweight核以及biasDepthwiseConv2Ddepthwiseweight核与bias没有 pointwise 核BatchNormalizationgamma、beta、running_mean、running_varLayerNormalization/InstanceNormalization/GroupNormalization仅gamma、betaEmpty无Dropout、池化、flatten、纯激活层归一化层之间的这种差异是有意为之也合乎设计。BatchNormalization 在训练中累积运行统计量并在推理时用到它们。这两个数组是训练出来的状态的一部分所以必须挺过序列化。如果重新加载时丢掉它们模型就会拿默认统计量去归一化推理模式的输入得到错误的结果。层归一化、实例归一化和组归一化每次前向传播都从当前输入现算统计量不保留任何运行状态能存的只有gamma和beta。3.9.4 会详细展示BatchNormalization这种情形。3.9.3. 一次完整的往返完整的流程有 5 步搭建、简单训练、保存、重建同样的层栈、加载。最后一步要确认恢复出的模型预测出同样的值。留意make_arch函数它把架构定义一次让活着的模型和重新加载的目标都调用它。这个习惯是“只存权重”持久化里最有用的实践因为它保证两个层栈绝不会各走各的。usendarray::Array;userustyml::error::Error;userustyml::neural_network::Tensor;userustyml::neural_network::layers::activation::linear::Linear;userustyml::neural_network::layers::dense::Dense;userustyml::neural_network::losses::MeanSquaredError;userustyml::neural_network::optimizers::SGD;userustyml::neural_network::sequential::Sequential;// 架构的唯一真实来源实时模型和重新加载的模型共用同一份定义。fnmake_arch()-Sequential{letmutmSequential::new();m.add(Dense::new(4,3,Linear::new()).unwrap()).add(Dense::new(3,2,Linear::new()).unwrap());m}fnmain()-Result(),Error{letx:TensorArray::from_shape_vec((2,4),vec![0.1f32,0.2,0.3,0.4,-0.1,-0.2,-0.3,-0.4]).unwrap().into_dyn();lety:TensorArray::from_shape_vec((2,2),vec![1.0f32,0.0,0.0,1.0]).unwrap().into_dyn();letmutmodelmake_arch();model.compile(SGD::new(0.01,0.0,false,0.0).unwrap(),MeanSquaredError::new());model.fit(x,y,5)?;letbeforemodel.predict(x)?;// 保存权重然后重建完全相同的层栈并加载进去。letpathroundtrip_demo.bin;model.save_to_path(path)?;letmutrestoredmake_arch();restored.load_from_path(path)?;letafterrestored.predict(x)?;// 两个预测张量必须逐元素相等。letmax_diff(after-before).mapv(f32::abs).iter().cloned().fold(0.0f32,f32::max);println!(max abs difference after round-trip: {max_diff:e});assert!(max_diff1e-6);std::fs::remove_file(path).unwrap();Ok(())}这次往返是精确的不是近似的postcard 无损地保存每一个f32load_from_path又把数组直接写进各层所以重新加载的模型在逐位层面执行的是完全相同的计算。上面的1e-6容差是防御性的余量不是在给可能的漂移留后路。3.9.4. 归一化的运行统计量能完整往返BatchNormalization的推理路径会读取running_mean和running_var。这个测试把这些统计量训练到偏离默认值保存模型再把权重加载进一个全新、未训练过的模型。然后检查推理模式下的预测是否依然吻合。答案是吻合的。这些运行数组随BatchNormalization变体一起走。usendarray::Array;userustyml::neural_network::Tensor;userustyml::neural_network::layers::regularization::normalization::batch_normalization::BatchNormalization;userustyml::neural_network::losses::MeanSquaredError;userustyml::neural_network::optimizers::SGD;userustyml::neural_network::sequential::Sequential;fnmake_arch()-Sequential{letmutmSequential::new();m.add(BatchNormalization::new(vec![4,3],0.9,1e-5).unwrap());m}fnmain(){letx:TensorArray::from_shape_vec((4,3),vec![0.5f32,-1.0,2.0,1.5,0.2,-0.7,-1.2,0.8,1.1,0.3,-0.4,0.9],).unwrap().into_dyn();// 训练一轮让 running_mean / running_var 偏离初始值。letmutmodelmake_arch();model.compile(SGD::new(0.001,0.0,false,0.0).unwrap(),MeanSquaredError::new());model.fit(x,x,8).unwrap();letbeforemodel.predict(x).unwrap();// 推理模式使用运行统计量letpathbatchnorm_demo.bin;model.save_to_path(path).unwrap();letmutrestoredmake_arch();// 全新模型运行统计量仍是默认值restored.load_from_path(path).unwrap();letafterrestored.predict(x).unwrap();letmax_diff(after-before).mapv(f32::abs).iter().cloned().fold(0.0f32,f32::max);println!(running-stat round-trip max abs difference: {max_diff:e});assert!(max_diff1e-6);std::fs::remove_file(path).unwrap();}要是运行统计量没被持久化after会和before大幅偏离因为全新模型的默认值对不上那 8 个 epoch 累积下来的批统计量。3.9.5. 哪些东西没被保存又为何会咬你一口优化器、它的累积状态、损失函数以及整套编译配置都会被留在原地。这是与 Keras 差别最大的地方Keras 的默认保存格式会把优化器一并打包好让fit无缝续训。在 RustyML 里加载给你的是一个空白模型上的权重optimizer和loss字段都是None由此引出两个后果。其一在模型能fit或开始训练之前你必须重新调用一次compile。预测则可以立刻用因为predict不需要优化器。其二加载之后续训会从零重启优化器这一点很容易被忽略。Adam 的一阶矩和二阶矩估计归零SGD 的动量缓冲也会清空Adam 用于偏差修正的时间步同样重新从 1 开始。对几步微调来说这次重置无伤大雅。但如果一段长训练被一次保存和加载切成两半加载后的头几步会比不间断训练走出更大、阻尼更弱的更新。损失曲线上会冒出一个暂时的凸起这是持久化留下的假象不是数据的问题。想在不出现这种断裂的前提下暂停再恢复长时间训练就让进程一直运行不要绕道磁盘保存再加载。RustyML 没有提供序列化优化器矩的 API。训练时的打乱种子同样不会被保存。如果你想让恢复后的打乱可复现加载后用set_seed重新设置一次。3.9.6. 加载时的错误变体每一种失败都以Error::Io(...)的形式出现见 1.6. 错误处理。恰好有 4 个变体每一个都对应一类不同的失败原因错误原因IoError::Std文件无法打开/读取路径不存在、权限不足或无法写入IoError::UnsupportedModelFormat文件没有带上本次构建的魔数和格式版本它要么不是 RustyML 模型要么出自一个磁盘权重布局不同的版本——见 3.9.8IoError::Serialization在头部合法的前提下字节流对预期的 schema 而言不是合法的 postcard损坏或截断IoError::ModelStructureMismatch你搭出的模型和文件对不上层数不对、某个位置的层类型不对或某个权重的形状和目标层不一致加载流程会先校验头部所以这 4 种错误有先后顺序根本不是模型的文件走不到 postcard 解码器出自不兼容版本的文件走不到结构检查。ModelStructureMismatch是开发模型时最常撞上的错误。RustyML 在 3 处会抛出它层数检查、逐位置的类型名检查以及权重应用过程中的形状检查。第三种情形不太显眼某一层set_weights报出的形状不一致会被包进同一个变体里。比如说一个Dense::new(2, 2, ...)的文件加载进Dense::new(3, 3, ...)的目标会通过层数和类型检查却栽在形状上报的仍是ModelStructureMismatch。消息字符串会告诉你到底是哪种检查失败了。匹配这个错误很直接userustyml::error::{Error,IoError};userustyml::neural_network::layers::activation::linear::Linear;userustyml::neural_network::layers::dense::Dense;userustyml::neural_network::sequential::Sequential;fnmain(){// 保存一个单层模型。letmutsavedSequential::new();saved.add(Dense::new(2,2,Linear::new()).unwrap());letpathmismatch_demo.bin;saved.save_to_path(path).unwrap();// 用错误的层数重建然后尝试加载。letmutwrongSequential::new();wrong.add(Dense::new(2,2,Linear::new()).unwrap()).add(Dense::new(2,2,Linear::new()).unwrap());matchwrong.load_from_path(path){Err(Error::Io(IoError::ModelStructureMismatch(msg))){println!(rejected as expected: {msg});}Err(other)panic!(unexpected error: {other:?}),Ok(())panic!(load must not succeed on a structure mismatch),}std::fs::remove_file(path).unwrap();}类型名检查是按字符串比对的。它能抓住常见错误本该是 Dense 的地方冒出一个 Conv 层或者少一层、多一层。但它抓不住那些不改变类型名和权重形状的超参数变化。两个形状相同、激活函数不同的Dense层会毫无怨言地加载成功因为激活函数根本不在磁盘上。这就是“只存元数据标签、不存完整架构”的实打实的代价也是“坚持用make_arch函数”这条纪律要紧的原因。3.9.7. postcard 格式体积、速度、可移植性postcard 是一种紧凑、非自描述的二进制格式。非自描述意味着流里不含字段名只按声明顺序写入各个值。这个设计让文件体积很小同时也把每个文件绑定在写出它的那个 crate 版本的确切结构体和枚举布局上。文件体积很好预估每个参数是一个f32占 4 字节再加上一点固定开销用于记录数组维度varint 编码、枚举判别标记这些小枚举各占 1 字节以及那几段简短的元数据字符串。因此Sequential::summary打印的Total params: N (N*4 B)那一行就是文件体积一个相当接近的上界估计一个 10 万参数的模型大约落在 400 KB。序列化是单趟线性扫描加一次缓冲写入所以瓶颈在 I/O 时间而非 CPU 时间。这个格式可以跨机器移植。postcard 规定了自己的字节序而不是直接倾泻本机字节序的内存。所以在一种架构上写出的 checkpoint能在另一种架构上正确反序列化与主机的字节序无关。这种可移植性不需要任何转换步骤也没有分平台的变体。它不能跨越的是 crate 版本这一点是下一节的主题。3.9.8. 工作流范式为最优模型打 checkpoint。分成若干短轮次训练每轮结束后用evaluate给模型打分只要分数改善就覆盖同一个文件。save_to_path会截断文件所以它始终装着目前为止最好的权重。内存里最终的那个模型可能已经越过最优点、开始过拟合了所以要丢弃它改用重新加载的版本。打分要用evaluate而不是fit返回的History的最后一项。history 里的每一项都是在该 epoch进行当中、取自那次权重更新之前的前向传播测出的损失。所以最后一项描述的是模型已经不再持有的权重照着它挑 checkpoint 会整整晚存一轮。evaluate在整份数据上跑一次推理模式的前向传播用编译进去的损失函数打分。它什么都不更新不算梯度、不动参数也不碰BatchNormalization的滑动统计量。所以建立在evaluate之上的筛选规则不会改变它正在度量的这次训练。usendarray::Array;userustyml::error::Error;userustyml::neural_network::Tensor;userustyml::neural_network::layers::activation::linear::Linear;userustyml::neural_network::layers::dense::Dense;userustyml::neural_network::losses::MeanSquaredError;userustyml::neural_network::optimizers::SGD;userustyml::neural_network::sequential::Sequential;fnmake_arch()-Sequential{letmutmSequential::new();m.add(Dense::new(4,3,Linear::new()).unwrap()).add(Dense::new(3,2,Linear::new()).unwrap());m}fnmain()-Result(),Error{letx:TensorArray::from_shape_vec((2,4),vec![0.1f32,0.2,0.3,0.4,-0.1,-0.2,-0.3,-0.4]).unwrap().into_dyn();lety:TensorArray::from_shape_vec((2,2),vec![1.0f32,0.0,0.0,1.0]).unwrap().into_dyn();letmutmodelmake_arch();model.compile(SGD::new(0.05,0.0,false,0.0).unwrap(),MeanSquaredError::new());letpathbest.bin;letmutbestf32::INFINITY;forroundin0..10{model.fit(x,y,2)?;letvalmodel.evaluate(x,y)?;ifvalbest{bestval;model.save_to_path(path)?;// 覆盖上一个 checkpointprintln!(round {round}: new best {val:.6}, checkpoint written);}}// best.bin 保存的是损失最低的权重未必是最后一轮的权重。letmutdeployedmake_arch();deployed.load_from_path(path)?;// evaluate 要的是编译进来的损失函数顺手传给它的优化器它一次都不会碰。deployed.compile(SGD::new(0.05,0.0,false,0.0).unwrap(),MeanSquaredError::new());assert_eq!(deployed.evaluate(x,y)?,best);std::fs::remove_file(path).unwrap();Ok(())}最后这一步用的是精确相等而不是容差。重新加载的权重与保存下去的逐位相同而evaluate在不含 dropout 的模型上是确定性的。所以恢复出的 checkpoint 打出的分数必须等于当初让它被写下的那个分数。在程序之间迁移权重。一个训练用的二进制搭建架构、训练模型、调用save_to_path另一个独立的服务用二进制搭建完全相同的架构、调用load_from_path。把make_arch函数放进公共模块或公共 crate 里给两边共用这样两个程序在架构上就不可能产生分歧。类型名和层数检查确实能抓到分歧但要等到一次失败的加载之后。共用同一个构造函数则能在编译期就抓到同样的分歧。服务用的二进制永远不需要调用compile因为它只调用predict。跨 crate 升级的版本管理。在文件头出现之前一个过时的 checkpoint 可能会悄无声息地加载失败。现在文件头挡住了这一点。每个保存下来的模型都以一个魔数和一个格式版本开头load_from_path会在解码任何其他内容之前先校验这两者。所以一个出自不同权重布局版本的 checkpoint会立刻以IoError::UnsupportedModelFormat报错。要是没有这道检查加载过程可能会把文件解析成看起来正确、实际数值却错误的数组。头部之后运行的结构检查比的是层数、类型名和权重的尺寸。一个过时文件可能凑巧同时满足这 3 项检查一个方形卷积核在轴被置换后尺寸不变而一个Dense权重的形状无论上游是什么张量布局都长得一样。这道防线能不能起作用取决于背后的开发纪律。只要某个权重容器的张量布局、秩或字段顺序发生变化RustyML 就要把版本号往上跳一格。跟着跳了的漂移会被拦下来忘了跳的则拦不住因为 postcard 依然是非自描述的文件里再没有别的东西带标签。对于任何需要在几周后、几个版本后重新加载的 checkpoint都要在Cargo.toml里锁定 RustyML 的版本。经过一次有意的升级之后要么重新跑一遍训练要么在旧版本下加载文件、再用新版本重新保存不要拿一个旧文件去赌新代码。文件头带来的效果是让版本相关的失败变得又响又早而不是悄无声息。想更深入了解这套持久化机制见 7.2. 深入模型持久化那一节讲了get_weights的查看路径、SerializableSequential包装器以及 RustyML 如何通过向下转型downcasting把权重应用回去。