ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

burn-store 实战指南:Burn 框架的模型存储、序列化与跨框架权重导入

burn-store 实战指南:Burn 框架的模型存储、序列化与跨框架权重导入 burn-store 实战指南Burn 框架的模型存储、序列化与跨框架权重导入【免费下载链接】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/burnburn-store是 Burn 深度学习框架的存储与序列化基础设施 crate负责模型的保存、加载、跨框架权重互操作PyTorch / SafeTensors与内存高效的张量管理。本文基于仓库中 crates/burn-store/README.md 及其配套源码展开覆盖三种存储后端的完整构建器 API、跨框架 Adapter 的转换规则、过滤与张量重命名机制、零拷贝/流式/原子写入等底层实现以及从旧burn-import迁移的完整对照帮助你掌握在 Burn 中导入、保存、转换模型权重的全套方案。一、burn-store 的定位与功能总览根据 READMEburn-store 提供以下核心能力Burnpack 格式.bpkBurn 原生格式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.tomlstd是文件 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)中的模块组织印证了功能划分核心抽象在traitsModuleSnapshot/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模块侧的扩展 traitModuleSnapshot对所有实现burn_core::Module的类型提供了 blanket impltraits.rs因此任何 Burn 模型都能直接调用其方法。关键方法collect(filter, adapter, skip_enum_variants)遍历模块收集张量返回Vecburn_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 进程而不是让被移走的模块被二次 droptraits.rs注释引用了 issue #3754 与 #5477save_into(store)/load_from(store)用户日常打到的便捷入口分别委托给store.collect_from(self)与store.apply_to(self)traits.rs。ModuleStore存储侧的 traitModuleStore为不同格式提供统一接口traits.rscollect_from(module)按 store 配置过滤、重映射、Adapter收集并写出apply_to(mut module)读回并应用返回结构化的ApplyResult包含applied成功应用、missing模块需要但文件中没有、skipped被过滤掉的模块参数、unused文件中多余且无人匹配的张量、errors非致命错误五类信息并实现Displayprintln!({}, result)即可得到带修复建议如使用allow_partial(true)的摘要get_tensor(name)/get_all_tensors()/keys()不依赖模块直接检查存储内容结果首次访问后缓存。ApplyResult是加载诊断的核心is_success()判断是否完整成功出错时LoadResult的Display输出会针对常见问题给出建议见 MIGRATION.md 的 Load Results 一节。三、三种存储后端及其构建器 APIBurnpackStoreBurn 原生格式BurnpackStoreburnpack.rs支持三种数据源对应三种典型部署场景// 文件模式stdfrom_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)自动含formatburnpack、producerburn、versioncrate 版本追加元数据clear_metadata()可全部清空见 burnpack.rsallow_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_adapterIdentityAdapter加载/保存方向的张量转换保存行为在源码中有两点值得强调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.rsmap_indices_contiguous处理 PyTorchnn.Sequential混合层类型导致的索引空洞例如fc.2.weight - fc.1.weight、fc.4.weight - fc.2.weightstore.rsget_bytes()内存模式下取回保存结果文件模式调用会报错。PytorchStore直接读 .pt/.pthPytorchStore由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.rsLinear 权重的布局转置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:CustomLayerenum 模块需用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后nextget_alternative_param_name先问self有备选名则再问next否则回退adapter.rs。五、过滤与重命名PathFilter 与 KeyRemapperPathFilterfilter.rs支持with_regex可多条OR 逻辑、with_full_path精确路径、with_predicate(fn(str, str) - bool)自定义谓词入参为张量路径与容器路径以及match_all关闭过滤。过滤在apply_to时生效而get_all_tensors/keys无视过滤便于全量检查文件内容KeyRemapperkeyremapper.rsstd门控一组 (正则, 替换串) 规则替换串支持$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需显式挂上PyTorchToBurnAdapterBurn 原生导出的则不需要。API 映射表MIGRATION.mdburn-importburn-storeLoadArgs::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 章节。八、适用前提与限制小结文件相关 APIfrom_file、KeyRemapper、map_indices_contiguous等依赖stdfeatureno-std 下可用 Burnpack 与 SafeTensors 的内存/静态字节路径及with_full_path等无正则过滤overwrite(false)默认意味着对已存在文件保存会报错训练框架中反复保存 checkpoint 时需显式.overwrite(true)或用 Burnpack 的原子保存语义理解替换即覆盖skip_enum_variants与PyTorchToBurnAdapter解决的是命名与布局差异而张量内容本身的转换如 dtype 归一建议叠加FloatCastAdapterBurnpack 的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 doesnt compromise on flexibility, efficiency and portability.项目地址: https://gitcode.com/GitHub_Trending/bu/burn创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表