数据管线与模型存储
从样本到张量批次
ruda-dataset 提供 Dataset<I> 抽象,包括 get(index)、len() 与迭代。ruda-model 的 dataset feature 提供组批和数据加载器。Batcher<B, I, O> 接收样本及目标后端设备,构造输出批次。
创建二进制项目并配置:
[dependencies]
ruda-tensor = { version = "0.21", default-features = false, features = ["api-std"] }
ruda-tensor-host = { version = "0.21", default-features = false, features = ["std"] }
ruda-model = { version = "0.21", default-features = false, features = ["std", "dataset"] }
将下面程序放入 src/main.rs,执行 cargo run:
use ruda_model::data::{
dataloader::{batcher::Batcher, DataLoaderBuilder},
dataset::InMemDataset,
};
use ruda_tensor::{api::Tensor, DType, TensorData};
use ruda_tensor_host::{Host, HostDevice};
struct PairBatcher;
impl Batcher<Host, [f32; 2], Tensor<Host, 2>> for PairBatcher {
fn batch(&self, items: Vec<[f32; 2]>, device: &HostDevice) -> Tensor<Host, 2> {
let rows = items.len();
let values: Vec<f32> = items.into_iter().flatten().collect();
Tensor::from_data(TensorData::new(values, [rows, 2]), (device, DType::F32))
}
}
fn main() {
let dataset = InMemDataset::new(vec![
[1.0f32, 2.0], [3.0, 4.0], [5.0, 6.0], [7.0, 8.0], [9.0, 10.0],
]);
let loader = DataLoaderBuilder::<Host, [f32; 2], Tensor<Host, 2>>::new(PairBatcher)
.batch_size(2)
.shuffle(42)
.num_workers(0)
.set_device(HostDevice)
.build(dataset);
let mut rows = 0;
for batch in loader.iter() {
assert_eq!(batch.dims()[1], 2);
rows += batch.dims()[0];
}
assert_eq!(rows, 5);
}
最后一批只有一行,因此形状应根据 items.len() 生成,而不是始终使用配置的 batch size。num_workers(0) 在调用线程组批,正数启用后台 worker。set_device 决定传给 batcher 的设备,不代表组批结束后还有一次自动设备迁移。
shuffle(seed) 控制加载器的洗牌随机数生成器。再次迭代同一个加载器会推进生成器状态;仅保存 seed 不能完整恢复到训练中的某个批次。
数据集选择
| 输入 | 入口 | 行为 |
|---|---|---|
| 已有内存样本 | InMemDataset::new(Vec<I>) |
持有样本,读取时克隆条目 |
| 其他数据集 | InMemDataset::from_dataset |
把全部样本载入内存 |
| JSON Lines | InMemDataset::from_json_rows |
每行反序列化一个 JSON 样本 |
| CSV | InMemDataset::from_csv |
根据传入的 CSV reader 配置反序列化各行 |
| SQLite | SqliteDataset |
启用 sqlite 或 sqlite-bundled |
| 自定义来源 | 实现 Dataset<I> |
定义按索引访问与数据集长度 |
ruda-dataset::transform 提供数据集组合适配器。样本解码放在 dataset/transform 中,张量组装放在 batcher 中。下载资源可使用启用了 network feature 的 ruda-io 所提供的 network::downloader::download_file_as_bytes;它与模型、数据存储是不同职责。
保存和加载模型
在上述项目中追加依赖:
ruda-nn = { version = "0.21", default-features = false, features = ["std"] }
ruda-store = { version = "0.21", default-features = false, features = ["std", "rudapack"] }
将 src/main.rs 替换为下面的保存/恢复程序。它在工作目录写入 linear.bpk,重复运行会覆盖该文件。
use ruda_nn::LinearConfig;
use ruda_store::{ModuleSnapshot, RudapackStore};
use ruda_tensor::{api::Tensor, DType};
use ruda_tensor_host::{Host, HostDevice};
fn main() -> Result<(), Box<dyn std::error::Error>> {
let device = HostDevice;
let model = LinearConfig::new(2, 1).init::<Host>(&device);
let input = Tensor::<Host, 2>::from_data([[1.0f32, 2.0]], (&device, DType::F32));
let expected = model.forward(input.clone()).into_data();
let mut output = RudapackStore::from_file("linear.bpk").overwrite(true);
model.save_into(&mut output)?;
let mut restored = LinearConfig::new(2, 1).init::<Host>(&device);
let mut source = RudapackStore::from_file("linear.bpk");
let result = restored.load_from(&mut source)?;
assert!(result.errors.is_empty());
assert!(result.missing.is_empty());
let actual = restored.forward(input).into_data();
assert_eq!(actual.as_slice::<f32>()?, expected.as_slice::<f32>()?);
Ok(())
}
加载前需要构造相同架构的模型。ModuleSnapshot 遍历模型参数;权重文件不会自动重建任意 Rust 模型代码。
存储格式选择
| Store | Feature | 用途 |
|---|---|---|
RudapackStore |
rudapack |
保存参数 ID 的 RUDA 张量快照,默认扩展名 .bpk |
SafetensorsStore |
safetensors |
读写 safetensors 权重和元数据 |
PytorchStore |
pytorch |
导入 PyTorch 检查点与 state dictionary |
ruda-store 默认启用三种格式。PyTorch 权重嵌套在键下时,使用 with_top_level_key("state_dict")。布局和命名惯例通过 PyTorch 适配路径处理,不能假定只复制原始字节即可完成转换。
with_regex 等过滤器选择参数路径。Safetensors 和 PyTorch Store 用 with_key_remapping 映射参数名;Rudapack 使用 with_remap_pattern 或 remap(KeyRemapper)。allow_partial(true) 显式允许部分加载,默认值是 false。检查 ApplyResult.applied、skipped、missing、unused 和 errors,区分已载入、未匹配及不兼容的条目。形状不匹配和 dtype 不匹配是不同错误。
HalfPrecisionAdapter 对浮点快照进行存储精度转换:保存时使用 with_to_adapter,读取时使用 with_from_adapter。without_module("LayerNorm") 可使归一化参数保留全精度,见完整半精度示例。
权重与训练状态的区别
加载权重恢复的是模型参数,不是完整训练位置。续训还需要训练循环使用的优化器状态、调度器/步数,以及数据顺序和随机状态。Rudapack 保留参数 ID,但不会自行捕获应用的全部状态。Record 与优化器恢复流程见训练与状态保存。