Burn 深度学习框架模型保存与加载完整指南:从 ModuleRecord 到 burn-store 跨框架权重迁移

发布时间:2026/9/14 2:18:26
Burn 深度学习框架模型保存与加载完整指南:从 ModuleRecord 到 burn-store 跨框架权重迁移 Burn 深度学习框架模型保存与加载完整指南从 ModuleRecord 到 burn-store 跨框架权重迁移【免费下载链接】burnBurn is a next generation tensor library and Deep Learning Framework that doesnt compromise on flexibility, efficiency and portability.项目地址: https://gitcode.com/GitHub_Trending/bu/burn本文是 Burn 深度学习框架位于仓库burn-book/src/saving-and-loading.md官方文档《Saving and Loading Models》的深度技术解析。文章围绕两条主线展开一是基于ModuleRecord与 burnpack.bpk格式的基础保存/加载流程二是基于burn-storecrate 的高级权重管理能力——包括 SafeTensors / PyTorch 跨框架互操作、键重映射、部分加载、过滤、零拷贝内存映射与半精度存储。读完本文你将掌握在 Burn 项目中保存、恢复、迁移与手术式移植模型权重的完整实战方案。一、基础知识Record 与 burnpack 格式在 Burn 中模型的参数保存在ModuleRecord中并通过 burnpack.bpk格式序列化。ModuleRecord持有的是与后端解耦的纯张量数据因此权重在后端之间可移植用 GPU 后端训练保存的权重可以直接在 CPU 后端上加载使用。burnpack 文件由三个部分组成详见 Record 文档头部header固定大小包含BURN魔数、格式版本号和元数据长度元数据CBOR 编码描述每个张量的名称、dtype、形状、数据偏移与可选的参数 id以及具名类型化标量整数、浮点、布尔和用户自定义键值对张量数据区每个张量的字节从 256 字节边界开始对齐因此数据可以被零拷贝 / 内存映射方式读回。类型化标量的存在使优化器与学习率调度器的非张量状态如 Adam 的步数、动量也能以同一格式持久化这正是断点续训的基础。Burn 共有三种 Record 类型Record持有内容产生方式ModuleRecord模块参数module.into_record()OptimizerRecord优化器状态optimizer.to_record()LrSchedulerRecord学习率调度器状态scheduler.to_record()每种 Record 既可写入文件save/load无扩展名时自动追加.bpk也可写入内存字节缓冲into_bytes/from_bytes适用于no-std部署场景权重随编译代码一起嵌入。二、最基础的保存与加载ModuleRecord API2.1 保存模型调用into_record()获取模型的参数记录再调用save写入磁盘use burn::store::ModuleRecord; // Take a record of the models parameters and save it to disk. model .into_record() .save(model_path) .expect(Should be able to save the model);注意当路径没有扩展名时.bpk扩展名会被自动追加因此只需提供文件路径与基本文件名。例如传入model会生成model.bpk。2.2 加载模型// Load the record from the burnpack file. let record ModuleRecord::load(model_path) .expect(Should be able to load the model weights from the provided file); // Apply the loaded weights to a model. model model.load_record(record);Record 是后端无关的用某个后端保存的权重可以在另一个后端上加载。如果需要转换精度可以在取 Record 之前调用model.cast(dtype)或在加载时使用record.cast_to_module_dtype()。2.3 从已保存权重初始化模型最直接的做法是使用load_recordModuletrait 的方法。由于参数初始化是惰性的——在模块真正被使用之前不会发生实际的张量分配也不会执行 GPU/CPU kernel——因此先init(device)再load_record(record)不会带来有意义的性能开销。use burn::store::ModuleRecord; // Create a dummy initialized model to save. let device Default::default(); let model Model::init(device); // Save its parameters to a burnpack file. model .into_record() .save(model_path) .expect(Should be able to save the model);之后即可从磁盘上的记录恢复模型// Load the model record from the burnpack file. let record ModuleRecord::load(model_path) .expect(Could not load model weights); // Initialize a new model with the loaded record/weights. let model Model::init(device).load_record(record);2.4 部分加载与容错加载如果记录中只包含部分参数例如只保存了部分层可以使用record.allow_partial(true)允许记录中缺少部分模块参数时仍然加载成功model.try_load_record(record)可失败fallible的加载变体。在加载行为上Record 还支持以下 builder 配置保存时被忽略.validate(false)跳过形状不匹配 / 张量缺失校验.cast_to_module_dtype()/.with_dtype_policy(..)加载时把记录数据转换为模块参数 dtype默认情况下参数会采用记录的 dtype。注意保存侧的 dtype 不可配置Record 保存的是模块当前持有的 dtype。要控制加载侧的 dtype请使用上述cast_to_module_dtype()系列方法。关于OptimizerRecord与LrSchedulerRecord它们暴露了相同的 API 形态用于训练断点续训// Optimizer state (no device needed on load; state migrates to each parameters device on the next step). optimizer.save(optim)?; let optimizer optimizer.load(optim)?; // Learning rate scheduler state (scalars only). scheduler.to_record().save(scheduler)?; let scheduler scheduler.load_record(LrSchedulerRecord::load(scheduler)?);使用Learner训练时这些记录由检查点checkpointer自动保存与恢复详见 Learner 文档。三、Model Weight Storeburn-store crate 进阶方案ModuleRecordAPI 适合基础的保存/加载而burn-storecrate 在相同的 burnpack 格式之上增加了内存效率与灵活性。从 lib.rs 的文档可以看出其核心能力Burnpack 格式CBOR 元数据、面向有状态训练的 ParamId 持久化、no-std支持SafeTensors 格式行业标准的张量序列化PyTorch 兼容直接加载 PyTorch 模型权重并自动做权重变换零拷贝加载内存映射文件 惰性张量物化灵活过滤使用正则、精确路径或自定义谓词加载/保存模型子集张量重映射加载/保存时为框架兼容性重命名张量。burn-store提供两个核心 trait见 traits.rsModuleSnapshot扩展 trait为任意Module提供collect()与apply()。collect是惰性的——返回的每个burn_pack::Tensor只在其数据被真正访问时才从设备读回apply则按名称把张量应用到模块参数上返回包含 applied / skipped / missing / unused / errors 信息的ApplyResult。ModuleStore统一存储接口定义collect_from、apply_to、get_tensor、get_all_tensors、keys等方法由BurnpackStore、SafetensorsStore、PytorchStore分别实现。save_into与load_from正是这两个 trait 的便捷封装见 traits.rssave_into调用store.collect_from(self)load_from调用store.apply_to(self)。3.1 支持的格式格式扩展名说明Burnpack.bpkBurn 原生格式加载快、支持零拷贝可持久化训练状态含 ParamIdSafeTensors.safetensorsHugging Face 提出的行业标准格式用于安全张量序列化PyTorch.pt、.pth直接加载 PyTorch 模型权重只读这些格式分别由BurnpackStore、SafetensorsStore、PytorchStore支持并受 Cargo feature flags 控制std默认启用文件 I/O、safetensors默认、pytorch默认、memmap由std隐含只要开启stdSafeTensors 文件加载就会内存映射。3.2 保存模型use burn_store::{ModuleSnapshot, BurnpackStore}; // Save to Burnpack (recommended) let mut store BurnpackStore::from_file(model.bpk); model.save_into(mut store)?; // Or save to SafeTensors use burn_store::SafetensorsStore; let mut store SafetensorsStore::from_file(model.safetensors); model.save_into(mut store)?;从源码burnpack.rs可以看到BurnpackStore在保存时会自动写入默认元数据formatburnpack、producerburn、versionburn-store crate 版本。另外自动扩展名from_file(model)会自动创建model.bpk已有扩展名则原样保留.auto_extension(false)可关闭见 burnpack.rs覆盖保护默认overwrite false目标文件已存在时保存会报错需显式.overwrite(true)见 burnpack.rs。3.3 加载模型use burn_store::{ModuleSnapshot, BurnpackStore}; let device Default::default(); let mut model MyModel::init(device); // Load from Burnpack let mut store BurnpackStore::from_file(model.bpk); model.load_from(mut store)?;3.4 从 PyTorch 加载权重可以直接从 PyTorch 的.pt文件加载权重use burn_store::{ModuleSnapshot, PytorchStore}; let mut model MyModel::init(device); let mut store PytorchStore::from_file(pytorch_model.pt); model.load_from(mut store)?;PyTorch 导出规范只保存权重state_dict不要保存整个模型对象import torch import torch.nn as nn class Net(nn.Module): def __init__(self): super(Net, self).__init__() self.conv1 nn.Conv2d(2, 2, (2, 2)) self.conv2 nn.Conv2d(2, 2, (2, 2), biasFalse) def forward(self, x): return self.conv2(self.conv1(x)) model Net() torch.save(model.state_dict(), model.pt) # Correct: save state_dict # torch.save(model, model.pt) # Wrong: saves entire model访问嵌套的 state_dict某些 PyTorch 检查点会把state_dict嵌套在一个键下例如checkpoint[state_dict]此时用with_top_level_keylet mut store PytorchStore::from_file(checkpoint.pt) .with_top_level_key(state_dict); model.load_from(mut store)?;自动权重变换PytorchStore在加载时会自动应用PyTorchToBurnAdapter逻辑见 adapter.rsLinear 层权重转置PyTorch 的[out, in]布局转换为 Burn 的[in, out]归一化层参数重命名weight → gamma、bias → beta对 BatchNorm / LayerNorm / GroupNorm / RmsNorm 生效。这套适配逻辑通过ModuleContext携带的容器栈识别张量所属的模块类型如Struct:Linear且能穿透Vec等集合包装器——即便VecLinear中的张量也会被正确识别为 Linear 权重。3.5 从 SafeTensors 加载从 PyTorch 导出的 SafeTensors 文件需要使用适配器进行正确的权重变换use burn_store::{ModuleSnapshot, PyTorchToBurnAdapter, SafetensorsStore}; let mut model MyModel::init(device); let mut store SafetensorsStore::from_file(model.safetensors) .with_from_adapter(PyTorchToBurnAdapter); model.load_from(mut store)?;如果 SafeTensors 文件是 Burn 自己创建的则无需适配器let mut store SafetensorsStore::from_file(model.safetensors); model.load_from(mut store)?;PyTorch 导出到 SafeTensorsfrom safetensors.torch import save_file model Net() save_file(model.state_dict(), model.safetensors)3.6 保存为 PyTorch 兼容格式使用BurnToPyTorchAdapter可以让 Burn 保存的权重被 PyTorch 直接消费Linear 权重转置回[out, in]gamma/beta重命名为weight/biasuse burn_store::{BurnToPyTorchAdapter, SafetensorsStore}; let mut store SafetensorsStore::from_file(for_pytorch.safetensors) .with_to_adapter(BurnToPyTorchAdapter) .skip_enum_variants(true); model.save_into(mut store)?;.skip_enum_variants(true)会在路径中跳过枚举变体名称例如 Burn 的feature.BaseConv.weight变为feature.weight因为 PyTorch / SafeTensors 不使用枚举变体。3.7 处理加载结果load_from返回关于加载过程的详细信息ApplyResult见 apply_result.rs包含applied成功应用的张量路径列表skipped因过滤而跳过的张量路径missing模块需要但源中没有的张量路径带容器栈信息unused源中存在但模块没有匹配的张量路径errors应用过程中遇到的非致命错误如形状不匹配、dtype 不匹配、适配器错误、加载错误。注意查看result.missing、result.errors等字段要求 store 配置了.allow_partial(true)。否则缺失张量会在你拿到ApplyResult之前就导致硬性Err。// Use .allow_partial(true) to get an ApplyResult with structured info let mut store PytorchStore::from_file(pretrained.pt) .allow_partial(true); let result model.load_from(mut store)?; // Print a formatted summary with suggestions println!({}, result); // Or inspect individual fields println!(Applied: {} tensors, result.applied.len()); println!(Missing: {:?}, result.missing); println!(Errors: {:?}, result.errors); if result.is_success() { println!(All tensors loaded successfully); }ApplyResult实现了Display会输出一份带 Did you mean? 建议的格式化摘要当存在缺失张量时它会基于 Jaro 相似度相似度阈值 0.7从未使用张量中推荐最可能的候选名并且能识别枚举变体导致的命名差异如 Burn 的field.BaseConv.weight对应 PyTorch 的field.weight给出skip_enum_variants或键重映射两种解决方案见 apply_result.rs。如果不使用.allow_partial(true)则走常规错误处理match model.load_from(mut store) { Ok(result) println!(Loaded successfully: {} tensors, result.applied.len()), Err(e) eprintln!(Failed to load: {e}), }3.8 添加元数据Burnpack 与 SafeTensors 都支持自定义元数据let mut store BurnpackStore::from_file(model.bpk) .metadata(version, 1.0) .metadata(description, My trained model) .metadata(epochs, 100); model.save_into(mut store)?;BurnpackStore还提供clear_metadata()清空所有元数据包括默认的format/producer/version字段见 burnpack.rs。四、高级功能4.1 键重映射Key Remapping当模型结构不匹配时可以用正则表达式重映射参数名let mut store PytorchStore::from_file(model.pt) // Remove prefix: model.conv1.weight - conv1.weight .with_key_remapping(r^model\., ) // Rename: layer1 - encoder.layer1 .with_key_remapping(r^layer, encoder.layer); model.load_from(mut store)?;对于复杂的重映射使用KeyRemapper见 keyremapper.rs它支持带捕获组的正则替换如$1use burn_store::KeyRemapper; let remapper KeyRemapper::new() .add_pattern(r^transformer\.h\.(\d)\., transformer.layer$1.)? .add_pattern(r\.attn\., .attention.)?; let mut store SafetensorsStore::from_file(model.safetensors) .remap(remapper);KeyRemapper的多个模式会按序依次应用于张量名返回值包含变换前后的路径对便于调试见 keyremapper.rs。4.2 部分加载Partial Loading当某些张量缺失时依然可以加载权重缺失的参数保持随机初始化let mut store PytorchStore::from_file(pretrained.pt) .allow_partial(true); let result model.load_from(mut store)?; println!(Missing (initialized randomly): {:?}, result.missing);4.3 过滤张量Filtering只加载或保存特定层// Load only encoder layers let mut store SafetensorsStore::from_file(model.safetensors) .with_regex(r^encoder\..*) .allow_partial(true); // Save only encoder layers let mut store SafetensorsStore::from_file(encoder.safetensors) .with_regex(r^encoder\..*); model.save_into(mut store)?; // Multiple patterns (OR logic) let mut store SafetensorsStore::from_file(model.safetensors) .with_regex(r^encoder\..*) // encoder tensors .with_regex(r.*\.bias$) // OR any bias tensors .with_full_path(decoder.scale); // OR specific tensorPathFilter见 filter.rs采用OR 逻辑——路径命中任意一条规则即被包含。除了正则与精确路径还支持with_predicate(fn)自定义谓词可同时接收路径与容器路径两个参数以及PathFilter::all()/PathFilter::none()快捷方式与or()组合。注意with_regex传入非法正则时会直接 panic见 filter.rs。4.4 非连续层索引Non-Contiguous Layer IndicesPyTorch 的nn.Sequential混用带参层与无参层如 ReLU时会产生非连续的索引。PytorchStore会自动将其重映射为连续索引PyTorch: fc.0.weight, fc.2.weight, fc.4.weight (gaps from ReLU layers) Burn: fc.0.weight, fc.1.weight, fc.2.weight (contiguous)map_indices_contiguous函数见 keyremapper.rs会检测路径中所有数字索引并按位置上下文独立重编号支持嵌套 Sequential 结构例如feature.layers.0.conv_block.2.weight中的每层索引各自重排。该行为默认开启需要时可关闭let mut store PytorchStore::from_file(model.pt) .map_indices_contiguous(false);4.5 零拷贝加载Zero-Copy Loading对嵌入式模型或超大文件零拷贝加载可避免内存复制// Embedded model (compile-time) static MODEL_DATA: [u8] include_bytes!(model.bpk); let mut store BurnpackStore::from_static(MODEL_DATA); model.load_from(mut store)?; // Large file (memory-mapped) let mut store BurnpackStore::from_file(large_model.bpk) .zero_copy(true); model.load_from(mut store)?;从源码看burnpack.rsfrom_static会把静态字节包装为共享字节源张量数据直接切片引用.rodata段而不复制。文件模式下burnpack 的 256 字节对齐设计见 Record 文档保证了内存映射读取的可行性。4.6 保存大模型保存到 Burnpack文件时采用流式写每个参数仅在 writer 到达它时才从设备读回并在读取下一个参数之前释放。因此峰值主机内存由最大的单个张量决定而非整个模型大小——即使模型权重驻留在显存中且超过主机 RAM也能成功保存。// Streams to disk; only one tensor is held in host memory at a time. let mut store BurnpackStore::from_file(large_model.bpk).overwrite(true); model.save_into(mut store)?;这一流式特性只适用于文件保存。BurnpackStore::from_bytes按定义需要在内存中构建整个容器所以大模型请优先使用文件路径。另外通过BurnpackStore的文件保存是**全有或全无all-or-nothing**的由于参数在写入中途才读回容器会先写在目标位置旁边完成后才重命名到位。因此保存失败、panic 或被 kill 都不会用截断文件替换已有文件。抵抗断电是更强的保证且仅限 Unix详见Writer::write_to_file_atomic。4.7 半精度存储Half-Precision Storage用 F16 保存模型可将文件体积缩小约 50%加载时再恢复为全精度use burn_store::{ModuleSnapshot, BurnpackStore, HalfPrecisionAdapter}; let adapter HalfPrecisionAdapter::new(); // Save: F32 - F16 (same adapter for both directions) let mut store BurnpackStore::from_file(model_f16.bpk) .with_to_adapter(adapter.clone()); model.save_into(mut store)?; // Load: F16 - F32 let mut store BurnpackStore::from_file(model_f16.bpk) .with_from_adapter(adapter); model.load_from(mut store)?;从源码看adapter.rsHalfPrecisionAdapter会根据张量原始 dtype自动判断转换方向F32 → F16保存F16 → F32加载其他 dtype 原样通过。默认转换的模块包括Linear、Embedding、Conv*含 ConvTranspose、DeformConv2d、LayerNorm、GroupNorm、InstanceNorm、RmsNorm、PRelu。BatchNorm 默认被排除因为其running_var在 F16 下会下溢见 adapter.rs。可用with_module()/without_module()自定义转换集合// Keep LayerNorm at full precision let adapter HalfPrecisionAdapter::new() .without_module(LayerNorm); // Add a custom module to the conversion set let adapter HalfPrecisionAdapter::new() .with_module(CustomLayer);with_module/without_module同时接受短名LayerNorm自动映射为Struct:LayerNorm与限定名如Enum:MyModulewithout_module在模块不在集合中时会 panic见 adapter.rs。另外还有一个目标驱动的FloatCastAdapter它会将所有浮点张量F64、F32、Flex32、F16、BF16统一转换到指定 dtype不限定模块类型——适合把 BF16 检查点加载到 F16 模型这类场景例如PyTorchToBurnAdapter.chain(FloatCastAdapter::to(DType::F16))。4.8 直接张量访问不加载到模型直接检查文件中的张量use burn_store::ModuleStore; let mut store PytorchStore::from_file(model.pt); // List all tensor names let names store.keys()?; // Get specific tensor if let Some(tensor) store.get_tensor(encoder.layer0.weight)? { println!(Shape: {:?}, DType: {:?}, tensor.shape, tensor.dtype); // Reading the data is a free function, not a method on the tensor let data burn_store::bridge::to_data(tensor)?; }从 traits.rs 可以看到这些方法的语义get_tensor返回惰性加载的引用一次只能持有一个查看多个请用get_all_tensors返回的BTreeMap键重映射会生效但过滤规则不会作用于直接访问过滤只在apply_to时应用结果首次调用后会被缓存。4.9 模型手术Model Surgery在模型之间迁移权重use burn_store::{ModuleSnapshot, PathFilter}; // Transfer all weights let snapshots model1.collect(None, None, false); model2.apply(snapshots, None, None, false); // Transfer only encoder weights let filter PathFilter::new().with_regex(r^encoder\..*); let snapshots model1.collect(Some(filter.clone()), None, false); model2.apply(snapshots, Some(filter), None, false);collect的第三个参数skip_enum_variants用于处理枚举变体collect返回的PackTensor数据是惰性的apply通过Module::map重写模块参数见 traits.rs实现上采用了 unsafe 的移出-映射-写回模式以避免克隆整个模块见注释引用的 issue 3754并在 unwind 窗口内用AbortOnUnwind守卫防止模块被二次 drop见 traits.rs。五、API 参考5.1 Builder 方法类别方法说明过滤with_regex(pattern)按正则模式过滤with_full_path(path)包含指定张量with_predicate(fn)自定义过滤逻辑重映射with_key_remapping(from, to)基于正则的重命名remap(KeyRemapper)复杂重映射规则适配器with_from_adapter(adapter)加载时变换with_to_adapter(adapter)保存时变换HalfPrecisionAdapter::new()F32/F16 混合精度配置allow_partial(bool)缺失张量时继续with_top_level_key(key)访问嵌套字典PyTorchskip_enum_variants(bool)跳过路径中的枚举变体map_indices_contiguous(bool)重映射非连续索引metadata(key, value)添加自定义元数据zero_copy(bool)启用零拷贝加载5.2 直接访问方法方法说明keys()获取有序张量名列表get_all_tensors()以 BTreeMap 获取全部张量get_tensor(name)按名称获取指定张量六、故障排查Troubleshooting6.1 常见问题1. Missing source values 错误你保存了完整的 PyTorch 模型而不是state_dict。请重新导出torch.save(model.state_dict(), model.pt)。2. 形状不匹配Burn 模型与源架构不一致。请核对各层配置通道数、卷积核大小、bias 设置。3. 键未找到参数名不匹配。使用with_key_remapping()或先检查键列表let store PytorchStore::from_file(model.pt); println!(Available keys: {:?}, store.keys()?);6.2 检查文件内容使用Netron可视化.pt与.safetensors文件Netron 是流行的神经网络可视化工具作者 lutzroeder。对 Burnpack 文件仓库提供了现成的检查示例源码位于 crates/burn-store/examples/burnpack_inspect.rscargo run --example burnpack_inspect model.bpk仓库中还有更丰富的参考burn-store的 PyTorch 读取测试覆盖了各种 dtype 与结构见 crates/burn-store/src/pytorch/tests/reader/SafeTensors 的 round-trip、过滤、元数据、错误处理等测试见 crates/burn-store/src/safetensors/tests/。此外examples/import-model-weights 提供了从 PyTorch 导出 MNIST 权重并导入 Burn 的端到端示例含mnist_train_export.py与训练好的.pt/.safetensors权重文件是理解完整导入流程的最佳实战起点。【免费下载链接】burnBurn is a next generation tensor library and Deep Learning Framework that doesnt compromise on flexibility, efficiency and portability.项目地址: https://gitcode.com/GitHub_Trending/bu/burn创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

关于本文作者

来自尧图内容编辑团队

尧图内容编辑团队 内容团队

尧图内容编辑团队

本文由尧图网络内容编辑团队执笔。团队由资深项目经理、前端工程师与设计师组成,所有内容均来自亲手交付的真实项目,先讲清问题、再给出可落地的解法。尧图深耕北京网站建设十年,服务过京华建材集团、智造科技等各行业客户,把一线经验沉淀为可复用的行业观察。

  • 十年建站经验,覆盖建材、制造、服务、文创等
  • 项目经理把关选题与事实准确性
  • 工程师与设计师联合撰写专业细节
  • 统一编辑规范,保证文风与排版一致
  • 每月复盘转化数据,迭代选题方向

延伸阅读

相关资讯与近期热门内容

深度阅读推荐

建站决策前值得细读的三篇

网站改版的5个关键决策
2024-08-12

网站改版的5个关键决策

什么时候该改版、改到什么程度、如何避免流量掉光,京华建材集团改版复盘给出答案。

获取专属建站方案

看完文章,把您的行业与预算告诉我们,免费获取一份量身定制的官网建设方案与报价。

立即免费咨询