burn-store 实战指南: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
burn-store是 Burn 深度学习框架的存储与序列化基础设施 crate,负责模型的保存、加载、跨框架权重互操作(PyTorch / SafeTensors)与内存高效的张量管理。本文基于仓库中 crates/burn-store/README.md 及其配套源码展开,覆盖三种存储后端的完整构建器 API、跨框架 Adapter 的转换规则、过滤与张量重命名机制、零拷贝/流式/原子写入等底层实现,以及从旧burn-import迁移的完整对照,帮助你掌握在 Burn 中导入、保存、转换模型权重的全套方案。
一、burn-store 的定位与功能总览
根据 README,burn-store 提供以下核心能力:
- Burnpack 格式(.bpk):Burn 原生格式,CBOR 元数据、内存映射加载、用于有状态训练(如 Adam 优化器状态)的 ParamId 持久化、no-std 支持;
- SafeTensors 格式:行业标准的张量序列化格式,兼顾安全与效率;
- PyTorch 支持:直接读取
.pth/.pt文件并自动完成权重变换; - 零拷贝加载:文件内存映射 + 张量惰性实例化(lazy materialization);
- 流式保存:Burnpack 文件保存时每次只从设备读回一个张量,峰值内存由最大单张量而非整个模型决定;且目标文件只在容器完整写入后才被替换,保存失败时旧文件保持完整;
- 灵活过滤:用正则、精确路径或自定义谓词加载/保存模型子集;
- 张量重映射(Remapping):在加载/保存时重命名张量,解决框架间命名差异;
- 半精度存储:F32/F16 自动转换,文件体积约缩小 50%;
- no-std 支持:Burnpack 与 SafeTensors 格式可用于嵌入式和 WASM 环境。
依赖声明上,crate 分类标注为no-std、embedded、wasm(见 Cargo.toml),默认 feature 为std、pytorch、safetensors、memmap。关键 feature 及含义如下:
| Feature | 默认 | 作用(见 Cargo.toml) |
|---|---|---|
std | 是 | 文件 I/O 等标准库功能;KeyRemapper的正则重映射也依赖它 |
pytorch | 是 | 启用.pt/.pth读取(依赖zip、serde、tar),并导出nested模块用于反序列化外部格式 |
safetensors | 是 | 启用 SafeTensors 读写 |
memmap | 是 | 隐含于std;开启std后读取 SafeTensors 文件一律走内存映射 |
cuda/metal/wgpu/tch | 否 | 透传给burn-core,指定目标后端 |
[lib.rs](https://link.gitcode.com/i/25649bd77e2dbd06ac84aed497ec3c44)中的模块组织印证了功能划分:核心抽象在traits(ModuleSnapshot/ModuleStore),张量收集与回写分别由collector与applier完成,结果汇总在apply_result,跨框架转换在adapter,重命名在keyremapper,过滤在filter;三个具体存储BurnpackStore、SafetensorsStore、PytorchStore各自独立,且burn_packcrate 被直接再导出——它是 burn-store 的张量运输类型(burn_pack::Tensor)的来源。
二、核心抽象:ModuleSnapshot 与 ModuleStore
所有 API 都围绕 traits.rs 中的两个 trait 展开。
ModuleSnapshot:模块侧的扩展 trait
ModuleSnapshot对所有实现burn_core::Module的类型提供了 blanket impl(traits.rs),因此任何 Burn 模型都能直接调用其方法。关键方法:
collect(filter, adapter, skip_enum_variants):遍历模块收集张量,返回Vec<burn_pack::Tensor>。收集是惰性的——返回的每个张量只有在被真正读取数据时才从设备取回(doc 注释见 traits.rs)。第三个参数skip_enum_variants控制路径中是否省略 enum 变体名,例如导出为feature.weight而非feature.BaseConv.weight,这正是与 PyTorch/SafeTensors 命名兼容所需的;apply(tensors, filter, adapter, skip_enum_variants):将张量按名字匹配回写到模块,返回ApplyResult。从源码看,其实现为了**避免克隆整个模块(那会使内存翻倍)**使用了ptr::read/ptr::write的"读出—map改写—写回"技巧,并配了一个AbortOnUnwind守卫:若在map期间适配器或后端代码 panic,则直接 abort 进程而不是让被移走的模块被二次 drop(traits.rs,注释引用了 issue #3754 与 #5477);save_into(store)/load_from(store):用户日常打到的便捷入口,分别委托给store.collect_from(self)与store.apply_to(self)(traits.rs)。
ModuleStore:存储侧的 trait
ModuleStore为不同格式提供统一接口(traits.rs):
collect_from(&module):按 store 配置(过滤、重映射、Adapter)收集并写出;apply_to(&mut module):读回并应用,返回结构化的ApplyResult,包含applied(成功应用)、missing(模块需要但文件中没有)、skipped(被过滤掉的模块参数)、unused(文件中多余且无人匹配的张量)、errors(非致命错误)五类信息,并实现Display,println!("{}", result)即可得到带修复建议(如"使用allow_partial(true)")的摘要;get_tensor(name)/get_all_tensors()/keys():不依赖模块直接检查存储内容,结果首次访问后缓存。
ApplyResult是加载诊断的核心:is_success()判断是否完整成功,出错时LoadResult的Display输出会针对常见问题给出建议(见 MIGRATION.md 的 Load Results 一节)。
三、三种存储后端及其构建器 API
BurnpackStore:Burn 原生格式
BurnpackStore(burnpack.rs)支持三种数据源,对应三种典型部署场景:
// 文件模式(std):from_file 会自动补 .bpk 扩展名 let mut store = BurnpackStore::from_file("model.bpk"); // 内存模式(no-std 可用):传 Some(bytes) 读、None 写 let store = BurnpackStore::from_bytes(None); // 静态字节(嵌入式):数据留在二进制的 .rodata 段,零拷贝切片 static MODEL_DATA: &[u8] = include_bytes!("model.bpk"); let store = BurnpackStore::from_static(MODEL_DATA);from_static的实现细节值得注意:它用Bytes::from_static把静态字节包装成共享句柄,张量数据"slice 而不 copy",因此适合把权重直接编进固件(burnpack.rs)。
完整构建器选项(默认值均来自源码):
| 方法 | 默认 | 说明 |
|---|---|---|
metadata(key, value) | 自动含format=burnpack、producer=burn、version=<crate 版本> | 追加元数据;clear_metadata()可全部清空(见 burnpack.rs) |
allow_partial(bool) | false | 允许缺失张量;关闭时缺任何张量都会以ValidationError硬失败 |
validate(bool) | true | 加载时校验 shape 与 dtype;关闭可提速但数据损坏会延迟暴露 |
overwrite(bool) | false | 目标文件已存在时,保存报错并提示使用.overwrite(true) |
auto_extension(bool) | true | 路径无扩展名时自动补.bpk;已有扩展名则原样使用 |
with_regex/with_full_path/match_all | — | 过滤:正则 / 精确路径 / 匹配全部 |
remap(KeyRemapper)/with_remap_pattern(from, to) | 空 | 加载时张量名正则重映射 |
with_from_adapter/with_to_adapter | IdentityAdapter | 加载/保存方向的张量转换 |
保存行为在源码中有两点值得强调(burnpack.rs):
- 流式:
collect阶段没有任何数据被物化,Writer逐个触达张量时才从设备读回。文件模式下峰值主机内存由最大单张量界定,而非整个模型;字节模式则会在内存中构建整个容器; - 原子:文件模式走
write_to_file_atomic,张量在写入中途物化失败时,不会截断目标路径上已有的内容——注释明确说明"设备回读中途失败不应截断已写入的内容"。
SafetensorsStore:行业标准格式
SafetensorsStore是一个File/Memory双变体枚举(store.rs)。文件保存同样是原子的:先在目标路径旁边创建 scratch 文件(<file_name>.<pid>-<n>.tmp命名),safetensors::serialize_to_file流式写出(张量逐个物化),全部成功后再 rename 到位(store.rs)。源码 doc 注释还提示了几个工程后果:新文件是新 inode,原路径的硬链接不保留(Unix 上权限位保留);符号链接会被替换为普通文件;覆盖保存需要两倍空间;保存中途被 SIGKILL/OOM 杀死会留下可安全删除的.tmp残留。
其构建器比 BurnpackStore 多几个与互操作直接相关的选项:
let mut store = SafetensorsStore::from_file("model.safetensors") // 过滤:多模式为 OR 逻辑 .with_regex(r"^encoder\..*") .with_full_path("decoder.output.bias") .with_predicate(|path, _| path.ends_with(".bias")) // 重命名 .with_key_remapping(r"^encoder\.", "transformer.encoder.") .with_key_remapping(r"\.gamma$", ".weight") // 元数据(默认已含 format/producer/version) .metadata("subset", "encoder_only") // PyTorch 互操作两个关键开关 .skip_enum_variants(true) // 路径中省略 Burn 的 enum 变体名 .map_indices_contiguous(true) // 把 0,2,4 这类跳号层重排为 0,1,2 .allow_partial(true) .validate(true) .overwrite(false) .with_from_adapter(PyTorchToBurnAdapter) .with_to_adapter(HalfPrecisionAdapter::new());skip_enum_variants:加载时让不含枚举名的外部路径匹配到 Burn 模块路径(feature.weight匹配feature.BaseConv.weight);保存时则导出与 PyTorch 约定一致的短路径(store.rs);map_indices_contiguous:处理 PyTorchnn.Sequential混合层类型导致的索引空洞,例如fc.2.weight -> fc.1.weight、fc.4.weight -> fc.2.weight(store.rs);get_bytes():内存模式下取回保存结果,文件模式调用会报错。
PytorchStore:直接读 .pt/.pth
PytorchStore由pytorchfeature 门控导出(lib.rs),用于直接读取 PyTorch pickle 权重。README 与 lib.rs 中的标准用法:
use burn_store::{ModuleSnapshot, PytorchStore}; let mut model = Model::init(&device); let mut store = PytorchStore::from_file("pytorch_model.pth") .with_top_level_key("state_dict") // 权重嵌套在 state_dict 键下时 .allow_partial(true); // 跳过未知张量 model.load_from(&mut store)?;四、跨框架 Adapter:张量级别的转换管线
跨框架兼容的"脏活"集中在 adapter.rs。ModuleAdaptertrait 接收张量及其容器栈上下文ModuleContext(如["Struct:Model", "Vec", "Struct:Linear"]),适配器根据"张量属于哪个用户自定义模块"决定如何变换;module_type()会跳过Vec等集合包装,直接命中最内层的Struct:/Enum:模块(adapter.rs)。
PyTorchToBurnAdapter / BurnToPyTorchAdapter
双向适配器处理两类差异(adapter.rs):
- Linear 权重的布局转置:PyTorch 的
[out, in]<-> Burn 的[in, out]。转置被延迟合成到字节源上(bridge::map_data),只在数据最终被取回时执行,passthrough 零成本;量化类型QFloat因字节布局特殊会被显式放行(adapter.rs); - 归一化层参数改名:在 BatchNorm/LayerNorm/GroupNorm/RmsNorm 中,PyTorch 的
weight/bias<-> Burn 的gamma/beta。除改名外还实现了get_alternative_param_name,在匹配阶段直接尝试备用名,保证norm.weight能命中模块里的gamma参数。
crate 内测试test_pytorch_to_burn_linear_weight、rename_keeps_the_enclosing_path等验证了转置确实移动数据而不仅换 shape、以及改名只替换路径最后一段(adapter.rs)。
HalfPrecisionAdapter:半精度存储
README 的 Quick Start 演示了半精度保存/加载:
use burn_store::{BurnpackStore, HalfPrecisionAdapter, ModuleSnapshot}; // 保存为 F16(约 50% 更小) let adapter = HalfPrecisionAdapter::new(); let mut store = BurnpackStore::from_file("model_f16.bpk") .with_to_adapter(adapter.clone()); model.save_into(&mut store)?; // 加载时同一个适配器自动反向(F16 -> F32) let mut store = BurnpackStore::from_file("model_f16.bpk") .with_from_adapter(adapter); model.load_from(&mut store)?;其智能默认(adapter.rs):按源 dtype 自动判断方向(F32->F16 保存、F16->F32 加载,其他 dtype 原样通过);默认转换 Linear、Embedding、全部 Conv 变体、LayerNorm、GroupNorm、InstanceNorm、RmsNorm、PRelu;默认排除 BatchNorm,因为其running_var在 F16 下会下溢。可用with_module("CustomLayer")(短名自动映射为Struct:CustomLayer,enum 模块需用Enum:X限定形式)与without_module(...)增删清单。
另有FloatCastAdapter::to(dtype)(adapter.rs):目标驱动,把所有浮点张量(F64/F32/Flex32/F16/BF16)统一转为指定 dtype,适合"BF16 检查点载入 F16 后端模型"这类场景——HalfPrecisionAdapter只处理 F32/F16 互转且限定模块清单,会把 BF16 张量原样放行。两者都支持chain组合,例如:
let adapter = PyTorchToBurnAdapter .chain(FloatCastAdapter::to(burn_core::tensor::DType::F16));ChainAdapter的管线语义:adapt为先self后next;get_alternative_param_name先问self,有备选名则再问next,否则回退(adapter.rs)。
五、过滤与重命名:PathFilter 与 KeyRemapper
- PathFilter(filter.rs):支持
with_regex(可多条,OR 逻辑)、with_full_path精确路径、with_predicate(fn(&str, &str) -> bool)自定义谓词(入参为张量路径与容器路径),以及match_all关闭过滤。过滤在apply_to时生效,而get_all_tensors/keys无视过滤,便于全量检查文件内容; - KeyRemapper(keyremapper.rs,
std门控):一组 (正则, 替换串) 规则,替换串支持$1捕获组。例如KeyRemapper::new().add_pattern(r"^pytorch\.(.*)", "burn.$1")。注意重映射作用于加载/保存的名字,get_tensor需要传重映射后的名字;而过滤不参与名字缓存阶段——源码注释强调"重映射在缓存阶段应用,过滤在 apply 时应用"(burnpack.rs)。
组合起来,"从外部检查点只导入某一部分并改名"一条链即可表达:
let mut store = SafetensorsStore::from_file("checkpoint.safetensors") .with_regex(r"^encoder\..*") // 只要 encoder .with_key_remapping(r"^encoder\.", "model.encoder.") .allow_partial(true); model.load_from(&mut store)?;六、内存模型:零拷贝加载与流式保存
README 宣称的"Zero-Copy Loading"由两条路径实现:
- SafeTensors 文件:
std打开时即内存映射(memmap2,见 Cargo.toml 与 lib.rs 的 feature 说明),张量按需从映射中切片; - Burnpack 文件:文件型 Reader 保持张量惰性,"数据在物化时才读取;内存中的共享源则是零拷贝"(burnpack.rs);静态字节(
from_static)则完全不离开.rodata段。
保存方向的对应机制即前述的流式写入与原子替换。仓库提供了四个基准直接度量这些特性(Cargo.toml 的[[bench]]声明):resnet18_loading、unified_loading、unified_saving、zero_copy_loading。运行方式与 README 一致:
# 生成模型文件(一次性) uv run benches/generate_unified_models.py # 加载 / 保存基准 cargo bench --bench unified_loading cargo bench --bench unified_saving # 指定后端 cargo bench --bench unified_loading --features metal基准文件位于 benches/unified_loading.rs、benches/unified_saving.rs、benches/zero_copy_loading.rs,模型生成脚本见 benches/generate_unified_models.py 与 benches/download_resnet18.py。
七、从 burn-import 迁移
README 顶部提示:从burn-import迁移的读者应查阅 MIGRATION.md。核心变化是从"record 中转"变为"直接载入模型":
PyTorch 文件(.pt/.pth):
// burn-import(旧) let record: ModelRecord = PyTorchFileRecorder::<FullPrecisionSettings>::default() .load("model.pt".into(), &device)?; let model = Model::init(&device).load_record(record); // burn-store(新) let mut model = Model::init(&device); let mut store = PytorchStore::from_file("model.pt"); model.load_from(&mut store)?;SafeTensors 文件:PyTorch 导出的.safetensors需显式挂上PyTorchToBurnAdapter,Burn 原生导出的则不需要。
API 映射表(MIGRATION.md):
| burn-import | burn-store |
|---|---|
LoadArgs::new(path) | PytorchStore::from_file(path)/SafetensorsStore::from_file(path) |
.with_key_remap(pattern, replacement) | .with_key_remapping(pattern, replacement) |
.with_top_level_key(key) | .with_top_level_key(key) |
.with_adapter_type(AdapterType::PyTorch) | .with_from_adapter(PyTorchToBurnAdapter) |
Recorder::<FullPrecisionSettings> | 精度由张量 dtype 自动处理 |
.with_debug_print() | 改用 tracing/logging |
新 API 相对旧版新增的能力(迁移指南"New Features"一节):allow_partial(true)部分加载、with_regex过滤、save_into保存(旧 recorders 不支持保存)、以及带applied/skipped/missing/errors的LoadResult。依赖声明相应从burn-import = { features = ["pytorch", "safetensors"] }换成burn-store = { features = ["pytorch", "safetensors"] }。
一个可运行的端到端示例在 examples/import-model-weights:提供pytorch、safetensors、convert三个二进制,分别演示从weights/mnist.pt、weights/mnist.safetensors导入权重做 MNIST 推理,以及把两种格式转换为 Burnpack。系统性的用法(含从 PyTorch 导出权重)见 Burn Book 的 Saving and Loading 章节。
八、适用前提与限制小结
- 文件相关 API(
from_file、KeyRemapper、map_indices_contiguous等)依赖stdfeature;no-std 下可用 Burnpack 与 SafeTensors 的内存/静态字节路径及with_full_path等无正则过滤; overwrite(false)(默认)意味着对已存在文件保存会报错,训练框架中反复保存 checkpoint 时需显式.overwrite(true),或用 Burnpack 的原子保存语义理解"替换即覆盖";skip_enum_variants与PyTorchToBurnAdapter解决的是命名与布局差异,而张量内容本身的转换(如 dtype 归一)建议叠加FloatCastAdapter;- Burnpack 的
validate(true)(默认)在加载时检查 shape/dtype,关闭校验可提速但会把数据损坏问题推迟到运行时。
burn-store 以"两个 trait + 三个 Store + 一个 Adapter 管线"的薄层设计,把 Burn 模块的张量存储收敛到了统一接口之下:保存/加载代码与具体格式解耦,格式间转换靠可组合的 Adapter 完成,性能特性(零拷贝、流式、原子写)则由 burn-pack 底座与std下的内存映射提供支撑。
【免费下载链接】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),仅供参考