Burn 框架 Record 机制与 burnpack 序列化格式深度解析
【免费下载链接】burnBurn is a next generation tensor library and Deep Learning Framework that doesn't compromise on flexibility, efficiency and portability.项目地址: https://gitcode.com/GitHub_Trending/bu/burn
导读
本文围绕 Burn 深度学习框架的Record(记录)机制展开,讲解训练状态如何以与后端解耦的纯张量数据形式被保存与加载,并深入剖析其底层容器格式burnpack(.bpk)的三段式文件结构。读完本文,你将掌握ModuleRecord、OptimizerRecord、LrSchedulerRecord三种记录类型的使用方法(含save/load与内存字节缓冲两种 I/O 路径)、加载期行为配置(部分加载、校验开关、dtype 策略),以及如何借助Learner检查点机制自动完成训练中断与恢复——同时理解这些 API 背后burn-core、burn-pack、burn-optim各 crate 的源码实现。
一、Record 是什么:训练状态的"可移植快照"
在 Burn 中,Record 是训练状态(模型参数、优化器状态、学习率调度器状态)的序列化载体。它的核心设计有两点:
- 与后端解耦:Record 持有的是纯张量数据(
TensorData),而不是绑定在某个后端上的张量对象。因此,用burn-cuda训练保存的权重,可以直接加载到burn-ndarray或burn-wgpu上运行,无需任何转换。 - 参数初始化保持惰性:加载 Record 时并不真正触发张量分配或内核执行,只是把参数值"登记"到模块上,实际计算要等到模块真正被使用才发生(详见下文"从记录的权重初始化"一节)。
所有 Record 统一序列化为burnpack格式(扩展名.bpk),该格式由独立的burn-packcrate 实现。burn-pack刻意保持"极简且与张量库无关":它只依赖burn-std(提供DType/Bytes)、serde和一个 CBOR 编解码器,本身并不理解 Burn 的模块或张量概念,而是由上层(如burn-core)在Tensor条目与自身记录类型之间做桥接(见 crates/burn-pack/src/lib.rs 的 crate 文档)。
从源码结构看(crates/burn-core/src/store/mod.rs 的模块文档),burn-core中的 Record 系统被刻意设计得"小而直白":通过ModuleVisitor/ModuleMapper按参数路径遍历模块,不做过滤、适配器或惰性快照——更丰富的快照与导入工具(过滤、键重映射、PyTorch/SafeTensors 适配器、跨框架存储)全部集中在burn-storecrate 中。
二、burnpack 文件格式:三段式二进制容器
一个 burnpack 文件由三个部分组成,所有多字节整数均为小端序:
| 组成部分 | 内容 | 说明 |
|---|---|---|
| 固定大小头部(Header,10 字节) | 魔数"BURN"(0x4255524E)、格式版本(u16)、元数据长度(u32) | 头部各字段字节范围由magic_range()、version_range()、metadata_size_range()定义 |
| 元数据块(CBOR) | 每个张量的描述(名称、dtype、shape、数据偏移、可选参数 id)、任意命名类型化标量、用户自定义 key/value 对 | 使用 CBOR 序列化,长度记录在头部 |
| 张量数据区 | 每个张量的原始字节,起始位置对齐到256 字节边界 | 支持零拷贝 / 内存映射(mmap)读取 |
对应实现位于 crates/burn-pack/src/base.rs:
MAGIC_NUMBER: u32 = 0x4255524E,即 ASCII 的"BURN"(写成小端字节时文件里呈现为NRUB);FORMAT_VERSION: u16 = 0x0001;HEADER_SIZE = 4 + 2 + 4 = 10字节;TENSOR_ALIGNMENT: u64 = 256,数据区起点通过aligned_data_section_start(metadata_size)计算,确保所有张量偏移(相对数据区)换算成绝对文件位置后依然满足 256 字节对齐。
为什么是 256 字节对齐
对齐到 256 字节边界带来多重收益(见base.rs中TENSOR_ALIGNMENT的注释):
- 满足所有元素类型的指针对齐要求(如
f64需要 8 字节对齐); - 缓存行友好(主流 CPU 缓存行为 64 字节);
- GPU 合并访存友好(CUDA 偏好 256 字节对齐);
- 为更宽的 SIMD 预留空间(AVX-512 为 64 字节,未来的 AVX-1024 为 128 字节);
- 与业界主流格式保持一致:GGUF、MLX、ncnn、MNN、TNN、vLLM-AWQ、Marlin 等 15+ 格式均采用 256 字节对齐,而 SafeTensors 采用 64 字节(AVX-512 最低要求),Core ML 采用 4096 字节。
256 字节对齐对典型张量尺寸而言开销可忽略不计,同时最大程度兼容当前与未来的硬件。
元数据中的类型化标量(Scalar)
burnpack 的元数据区不仅描述张量,还支持存储命名类型化标量——包括有符号整数、无符号整数、浮点数、布尔值(Scalar枚举的Int(i64)、UInt(u64)、Float(f64)、Bool(bool)四个变体)。标量存放在 CBOR 元数据区而非张量数据区,因此不产生对齐开销。
正是这一能力让优化器和学习率调度器能够把非张量状态(步数计数器、当前学习率、动量超参数等)用同一种格式持久化。对于纯标量状态,Scalar的转换是类型安全的:例如i32::try_from(Scalar::from(-5i32))可以成功,但u8::try_from(Scalar::from(300u32))会因超出范围而失败,i64::try_from(Scalar::Float(1.5))会因变体不匹配而失败(相关测试见base.rs的scalar_tests模块)。
内置安全限制
为防止恶意或损坏的输入导致资源耗尽,burnpack 读取端在分配内存前会拒绝超出以下任一上限的文件(见 crates/burn-pack/src/lib.rs 的"Safety limits"一节):
| 常量 | 限制值 | 目的 |
|---|---|---|
MAX_METADATA_SIZE | 100 MB | 防止过大的元数据声明耗尽内存 |
MAX_TENSOR_COUNT | 100,000 | 防止过多张量导致资源耗尽 |
MAX_TENSOR_SIZE | 32 位平台 2 GB / 64 位平台 10 GB | 防止单个张量声明过大 |
MAX_CBOR_RECURSION_DEPTH | 128 层 | 防止深层 CBOR 嵌套导致栈溢出 |
MAX_FILE_SIZE | 100 GB(仅 std) | 文件加载器的文件大小上限 |
同时,读取器还会校验文件大小是否足以容纳其声明的每个张量,否则返回Error::ValidationError。
延迟张量字节写入
值得一提的实现细节:张量的字节并不需要在写入时已经存在。Reader只在访问时才从数据源读取张量字节;而调用方如果已知张量长度却尚未持有其数据(例如模块快照、ONNX 初始化器),可以用Tensor::deferred构造条目,把大于宿主内存的模型流式写入文件(相关契约见 crates/burn-pack/src/lib.rs 的 crate 文档)。
三、三种 Record 类型
Burn 将训练状态划分为三类 Record,分别对应训练中三个不同角色的持久化需求:
| Record | 持有内容 | 产生方式 |
|---|---|---|
ModuleRecord | 模块的参数 | module.into_record() |
OptimizerRecord | 优化器状态 | optimizer.to_record() |
LrSchedulerRecord | 学习率调度器状态 | scheduler.to_record() |
每种 Record 都支持两条 I/O 路径:
- 文件路径:
save(path)/load(path)——当路径没有扩展名时自动追加.bpk; - 内存字节缓冲:
into_bytes()/from_bytes(bytes)——对no-std部署特别有用,字节可以被include_bytes!直接嵌入编译产物(例如 examples/mnist-inference-web 中嵌入model.bpk的做法)。
四、ModuleRecord:模块参数的保存与加载
ModuleRecord(位于burn::store)以模块内参数路径为键持有模块的参数。在源码中,每个被记录的张量是一个RecordTensor { path, id, data }三元组(crates/burn-core/src/store/mod.rs):path是模块内的点分路径,id是参数 id(ParamId),data是TensorData。
它通过Moduletrait 自身产生和应用:
use burn::store::ModuleRecord; // 取出记录并保存(写出 model.bpk) model.into_record().save("model")?; // writes model.bpk // 加载回来并应用到已初始化的模块 let record = ModuleRecord::load("model")?; let model = ModelConfig::new().init(&device).load_record(record);收集与回放的实现原理
- 收集方向:
ModuleRecord::from_module通过Collector实现ModuleVisitor,在enter_module/exit_module时维护路径栈,遇到Float/Int/Bool参数时调用record()把(路径, 参数id, 张量数据)压入列表。注意它记录的是参数的保存形态(transform_for_save),这正是加载端校验与回放所依据的形态——对于像Col布局Linear权重这种通过 mapper 改变形状的参数,这一选择保证了形状映射参数能正确往返(有专门的round_trip_a_shape_mapped_param测试佐证)。 - 回放方向:
ModuleRecord::apply通过ModuleRecordMapper实现ModuleMapper,按模块路径查找记录中的张量,命中则把张量装载回参数,同时恢复持久化的ParamId。这一点很关键:优化器状态是按ParamId索引的,恢复参数 id 才能让优化器状态在保存/加载循环后依然有效(对应load_record_preserves_param_id测试)。
加载期行为配置(保存时忽略)
ModuleRecord提供一组 builder 方法,用于配置加载时的行为(这些设置在保存时被忽略):
.allow_partial(true)—— 即使记录中缺少某些模块参数也允许加载(例如加载用into_record_group取得的局部记录,或把旧版本 checkpoint 加载到新增了层的新模型上);.allow_unused(true)—— 允许记录中含有匹配不到任何模块参数的张量。它与allow_partial互为镜像,默认被拒绝:一条落不到任何位置的记录条目意味着目标模块并非该记录来源,看似成功的加载其实静默少做了事情。仅在明确场景下放开——把 checkpoint 加载到它来源模块的一部分上;.validate(false)—— 跳过形状不匹配 / 张量缺失的校验;.cast_to_module_dtype()/.with_dtype_policy(..)—— 加载时把记录数据转换为模块参数的 dtype(默认策略是参数采用记录的 dtype)。
对应的DTypePolicy枚举定义了两个变体(crates/burn-core/src/store/mod.rs):
| 策略 | 行为 |
|---|---|
FromRecord(默认) | 模块参数采用记录的 dtype,数据原样加载 |
CastToModule | 记录数据在加载时转换为模块参数当前的 dtype(会物化目标参数以读取其 dtype) |
// 允许部分加载 + 转换为模块 dtype let model = ModelConfig::new() .init(&device) .load_record(record.allow_partial(true).cast_to_module_dtype());保存侧 dtype 不可配置
保存侧的 dtype 是不可配置的:记录保存的是模块当前持有的 dtype。若要控制加载时的 dtype,有两个入口:
- 保存前调用
model.cast(dtype),让记录直接保存目标 dtype; - 加载时使用
.cast_to_module_dtype()/.with_dtype_policy(..)做转换。
store/mod.rs的测试给出了两种策略的行为差异:记录保存 f32 数据,目标模块参数是 f64——默认策略(FromRecord)加载后参数保持 f32(数据原样),而.cast_to_module_dtype()会把 f32 数据转换为 f64 并保持数值不变。
失败处理:try_load_record
load_record在校验失败时会 panic;需要可失败语义时使用try_load_record,它返回Result<Self, RecordError>。RecordError有两个变体:
Io(String)—— 读写记录时的 I/O 或格式错误(由burn_pack::Error转换而来);Validation(String)—— 应用记录时校验失败(形状不匹配、不允许部分加载时张量缺失、不允许未使用条目时记录张量无参数匹配)。
match model.clone().try_load_record(record) { Ok(model) => { /* 加载成功 */ } Err(e) => eprintln!("加载失败: {e}"), }从源码看,apply的校验逻辑会汇总三类问题:errors(形状不匹配)、missing(缺失张量)、unused(记录中无参数匹配的条目,排序后命名输出),分别由validate、allow_partial、allow_unused三个开关控制是否放行。
用into_record_group记录模块的局部
当只需要记录模块的某一部分参数时,可以用into_record_group(ParamGroup)。从Collector的实现看,ParamGroup会在读取参数数据之前就过滤掉不匹配的参数,因此"一个组的记录永远不会物化模块的其余部分"(crates/burn-core/src/module/base.rs 中into_record_group的文档)。对应的测试验证:只记录weight组的记录应用回完整模块时,weight落地、bias保持原初始化值,且需配合.allow_partial(true)使用。
从记录的权重初始化
一个实用技巧:由于参数初始化是惰性的,init(device)后紧跟load_record(record)并不会产生实际的张量分配与 GPU/CPU 内核执行开销。因此完全可以用Model::init(&device).load_record(record)这种"先初始化再覆盖"的方式加载权重,而不必担心性能代价。更完整的保存/加载流程见 Burn 中的模型保存与加载。
五、OptimizerRecord与LrSchedulerRecord:检查点与恢复训练
优化器和学习率调度器暴露了形状相同的 API,用于训练检查点与恢复:
// 优化器状态(加载时无需设备;状态会在下一步迁移到每个参数的设备上) optimizer.save("optim")?; let optimizer = optimizer.load("optim")?; // 学习率调度器状态(仅标量) scheduler.to_record().save("scheduler")?; let scheduler = scheduler.load_record(LrSchedulerRecord::load("scheduler")?);OptimizerRecord:按参数键控的状态分解
与按模块路径键控的ModuleRecord不同,OptimizerRecord按参数(ParamId)键控(crates/burn-optim/src/optim/module/record/mod.rs):
- 每个参数的优化器状态被分解为名为
"{param_id}.{field}"的张量(携带来源param_id),外加若干存放在 burnpack 标量图中的类型化标量条目; - 结构包含
tensors(状态张量)、scalars(类型化标量)、paths(元数据字符串映射)。
从ModuleOptimizer::to_record取到记录后,同样支持save/load(文件)与into_bytes/from_bytes(内存)。加载时不需要设备:状态张量会在优化器下一步step时迁移到每个参数所在的设备上。
LrSchedulerRecord:纯标量状态
学习率调度器的状态只是少量标量(步数计数器、当前学习率等),因此LrSchedulerRecord的 burnpack 记录只含命名类型化标量、不含张量(crates/burn-optim/src/lr_scheduler/base.rs):
with_scalar(key, value)/scalar(key)读写单个标量;- 组合式调度器(如
ComposedLrScheduler)通过with_record(prefix, record)/record(prefix)以索引前缀嵌套子调度器的记录; from_state/into_state复用与优化器状态相同的RecordState分解机制(state_flatten/state_unflatten),并在 debug 构建中断言不会产生张量叶子节点——"调度器状态应为纯标量"这一不变量由代码强制保证。
lr_scheduler/base.rs的测试工具check_save_load验证了保存/加载往返的正确语义:先推进若干步,保存并重新加载记录后,调度器必须从离开的位置继续产生与未保存副本完全一致的学习率序列。
Learner自动检查点
当使用Learner训练时,上述三类记录由检查点机制(checkpointer)自动保存与恢复,无需手工管理。相关用法详见 Learner 指南;LearnerConfig的num_epochs、checkpoint配置与save/load流程可参考 crates/burn-train/src/learner/base.rs 及 crates/burn-train/src/checkpoint 目录下的实现。
六、跨框架格式:burn-store
ModuleRecordAPI 适用于基本的保存/加载,但以下高级需求需要使用burn-storecrate(基于同一 burnpack 格式):
- 从其他生态导入权重:PyTorch(
.pt/.pth,只读)、SafeTensors(.safetensors); - 更高级的 store 功能:键重映射(
with_key_remapping/KeyRemapper)、过滤(with_regex/with_full_path)、半精度存储(HalfPrecisionAdapter)、零拷贝内存映射加载、部分加载与ApplyResult结构化结果检查、模型手术(快照收集与重放)。
例如从 PyTorch 加载并重映射键名:
use burn_store::{ModuleSnapshot, PytorchStore}; let mut model = MyModel::init(&device); let mut store = PytorchStore::from_file("pytorch_model.pt") .with_top_level_key("state_dict") .with_key_remapping(r"^model\.", ""); model.load_from(&mut store)?;完整的示例(含 PyTorch 导出注意事项、SafeTensors 适配器、元数据写入、大模型流式保存、非连续层索引映射等)见 Burn 中的模型保存与加载,对应实现位于 crates/burn-store(其pytorch-tests与safetensors-tests目录包含与 PyTorch 导出脚本对应的端到端测试)。仓库中的可运行示例还包括 examples/import-model-weights(PyTorch/SafeTensors 权重导入)与 examples/mnist-inference-web(.bpk模型嵌入 WebAssembly 推理)。
七、典型工作流总结
- 保存模型:
model.into_record().save("model")—— 得到model.bpk; - 加载模型:
Model::init(&device).load_record(ModuleRecord::load("model")?); - 迁移精度:保存前
model.cast(dtype),或加载时record.cast_to_module_dtype(); - 部分加载:
record.allow_partial(true)配合try_load_record; - 断点续训:用
Learner检查点自动持久化ModuleRecord+OptimizerRecord+LrSchedulerRecord三类状态; - 嵌入式 / no-std:
into_bytes()/from_bytes()在内存中完成序列化,字节可嵌入编译产物; - 跨框架互操作:通过
burn-store的PytorchStore/SafetensorsStore导入导出权重。
无论选择哪条路径,底层都是同一种紧凑、可零拷贝读取、具备安全上限校验的 burnpack 容器——这正是 Burn 训练状态"可移植、可恢复、可互操作"的根基。
【免费下载链接】burnBurn is a next generation tensor library and Deep Learning Framework that doesn't compromise on flexibility, efficiency and portability.项目地址: https://gitcode.com/GitHub_Trending/bu/burn
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考