Burn 自定义训练循环完全指南:手写 Epoch、梯度累积与 Parameter Groups

Burn 自定义训练循环完全指南:手写 Epoch、梯度累积与 Parameter Groups Burn 自定义训练循环完全指南手写 Epoch、梯度累积与 Parameter Groups【免费下载链接】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本篇技术指南聚焦 BurnRust 深度学习框架中如何绕过内置的Learner训练接口完全手写一套属于自己的训练循环。文章以仓库中完整的 custom-training-loop 示例 为骨架逐步拆解前向传播、损失计算、backward()梯度求解、GradientsParams参数梯度映射、optim.step()参数更新以及验证阶段关闭梯度追踪的model.valid()用法并深入讲解GradientsAccumulator梯度累积与ParamGroup分组优化等进阶能力。读完本文你将掌握在不使用任何高层训练抽象的前提下用 Burn 的底层 API 构建可定制、可复现的训练与验证流水线。为什么要手写训练循环Burn 内置了专门用于简化训练流程的 Learner它封装了数据加载、批次迭代、指标记录、Checkpoint、进度渲染等大量样板逻辑。但内置方案并不总能覆盖所有需求某些特殊训练策略如对抗训练、知识蒸馏、自定义的交替优化难以用配置项描述你可能希望在每个迭代粒度上插入自定义逻辑如动态修改损失权重、逐迭代打印详细信息或者你只是更习惯把训练流程的每一行代码都掌握在自己手里。此时直接手写训练循环往往更快、更灵活。Burn 的底层 API 设计上就为这种用法留足了空间——你只需要Module、GradientsParams、Optimizer这几个核心抽象即可拼出完整训练流程。从 basic workflow 示例出发、去掉Learner就是最顺滑的起点。起点配置与数据管道保持不变自定义训练循环的第一步与常规流程完全一致定义训练配置、初始化模型与优化器、构建数据加载器。仓库中的 custom-training-loop 示例 完整展示了这部分代码#[derive(Config, Debug)] pub struct MnistTrainingConfig { #[config(default 10)] pub num_epochs: usize, #[config(default 64)] pub batch_size: usize, #[config(default 4)] pub num_workers: usize, #[config(default 42)] pub seed: u64, #[config(default 1e-4)] pub lr: f64, pub model: ModelConfig, pub optimizer: AdamConfig, } pub fn run(device: Device) { // Create the configuration. let config_model ModelConfig::new(10, 1024); let config_optimizer AdamConfig::new(); let config MnistTrainingConfig::new(config_model, config_optimizer); let device device.autodiff(); device.seed(config.seed); // Create the model and optimizer. let mut model config.model.init(device); let mut optim config.optimizer.init(); // Create the batcher. let batcher MnistBatcher::default(); // Create the dataloaders. let dataloader_train DataLoaderBuilder::new(batcher.clone()) .batch_size(config.batch_size) .shuffle(config.seed) .num_workers(config.num_workers) .build(MnistDataset::train()); let dataloader_test DataLoaderBuilder::new(batcher) .batch_size(config.batch_size) .shuffle(config.seed) .num_workers(config.num_workers) .build(MnistDataset::test()); }这里几个要点值得展开说明#[derive(Config)]与#[config(default ...)]配置结构体由Config派生宏生成构造函数与序列化能力。所有标了default的字段在调用MnistTrainingConfig::new(model, optimizer)时都可以省略只传必须显式指定的model与optimizer。底层Configtrait见 crates/burn-core/src/config.rs要求类型满足Debug Serialize DeserializeOwned因此训练超参数天然可以被保存为 JSON 配置文件、随实验一起归档。设备与随机种子device.autodiff()将普通设备包装成具备自动微分能力的设备模型参数、输入张量都将创建在该设备上device.seed(config.seed)为设备上的随机数生成器播种保证实验可复现。模型与优化器的创建是分离的config.model.init(device)依据ModelConfig初始化网络参数config.optimizer.init()创建AdamConfig对应的优化器实例此时尚未为任何参数分配动量状态状态是惰性创建的。本示例使用的 CNN 模型定义在 examples/guide/src/model.rs包含两层卷积、自适应池化、Dropout 与两个全连接层。Batcher 负责把原始样本变成张量批次MnistBatcher见 examples/guide/src/data.rs实现BatcherMnistItem, MnistBatch将28x28的图像归一化到均值为 0.1307、标准差为 0.3081 的分布并拼装出images: Tensor3与targets: Tensor1, Int的批次结构。手写训练循环前向、损失、梯度与参数更新配置和加载器就绪后就进入核心环节——用for循环手写训练与验证逻辑// Iterate over our training and validation loop for X epochs. for epoch in 1..config.num_epochs 1 { // Implement our training loop. for (iteration, batch) in dataloader_train.iter().map(Result::unwrap).enumerate() { let output model.forward(batch.images); let loss CrossEntropyLoss::new(None, output.device()) .forward(output.clone(), batch.targets.clone()); let accuracy accuracy(output, batch.targets); println!( [Train - Epoch {} - Iteration {}] Loss {:.3} | Accuracy {:.3} %, epoch, iteration, loss.clone().into_scalar::f32(), accuracy, ); // Gradients for the current backward pass let grads loss.backward(); // Gradients linked to each parameter of the model. let grads GradientsParams::from_grads(grads, model); // Update the model using the optimizer. model optim.step(config.lr, model, grads); } }逐行拆解这段训练逻辑外层 Epoch 循环从1到num_epochs含每次迭代都完整走一遍训练集与验证集。前向传播model.forward(batch.images)得到每个样本的类别得分output。由于模型是在device.autodiff()上创建的前向过程会自动记录计算图为反向传播做好准备。损失计算使用 Burn 内置的 CrossEntropyLoss 计算分类损失。output.clone()与batch.targets.clone()是因为后续的accuracy函数还要复用这两个张量。准确率示例自定义了accuracy函数见 examples/custom-training-loop/src/lib.rs内部用output.argmax(1)取预测类别、equal(targets)统计正确数再除以样本总数换算成百分比。反向传播loss.backward()返回一个Gradients对象其中包含计算图中每个变量的梯度。梯度映射GradientsParams::from_grads(grads, model)是关键一步——它把Gradients中按变量节点组织的梯度转换为按模型参数ParamId组织的GradientsParams。参数更新optim.step(config.lr, model, grads)消费掉梯度并返回更新后的模型。注意step返回的是新模型所以需要重新赋值给model。为什么需要GradientsParams映射从源码看GradientsParams本质是一个以ParamId为键的张量容器见 crates/burn-optim/src/optim/grads.rspub struct GradientsParams { container: TensorContainerParamId, }from_grads内部通过GradientsParamsConverter访问者遍历模块的所有参数把Gradients中对应的梯度逐一提出来。这一步是必须的因为在一个训练流程中你可能运行多个不同的自动微分图例如共享特征提取器的多任务场景同一批ParamId的梯度可能分散在多张图里。GradientsParams把这些梯度统一汇总、按参数 ID 索引后续优化器才能精确地把每个梯度应用到对应参数上。它额外提供了get、remove、register、to_device等方法方便你按需读取或搬运梯度。与 PyTorch 的对比无需zero_grad、无需手动注册梯度值得强调的是这套流程与 PyTorch 的习惯有本质区别不需要zero_grad()optim.step()会消费传入的GradientsParams而每轮loss.backward()都会基于当前前向重新构建计算图、重新计算梯度不存在旧梯度残留的问题不需要把梯度注册给优化器梯度以参数为键随GradientsParams显式传入优化器的内部状态如 Adam 的一阶/二阶动量则由优化器自身按ParamId维护对用户完全透明。梯度累积用GradientsAccumulator实现大批次等效微调批次过小或显存受限时常需要累积多个小批次的梯度后再更新一次参数。Burn 提供了GradientsAccumulator让这件事变得非常直接见 crates/burn-optim/src/optim/grad_accum.rslet mut accumulator GradientsAccumulator::new(); let grads model.backward(); let grads GradientsParams::from_grads(grads, model); accumulator.accumulate(model, grads); // 反复调用累积多个小批次 // ... let grads accumulator.grads(); // 弹出累积后的梯度执行一次 step其工作原理是accumulate会遍历模块的每个参数把新梯度张量与容器内已有的梯度做逐元素相加见ModuleGradsAccumulator::visit_floatgrad_accum.rsgrads()则通过mem::swap取出全部累积梯度并重置容器状态。仓库自带的单元测试如test_accumulate_gradients_two_steps验证了连续累积两个小批次后容器中仍按参数 ID 维护着两个参数对应的梯度条目。累积完成后再像普通训练一样调用一次optim.step(lr, model, accumulated_grads)即可学习率上通常还需相应放大或使用 warmup 策略——这部分属于训练策略的权衡Burn 不做任何强制。验证循环用model.valid()关闭梯度追踪每个 epoch 结束后示例对未见过的测试集执行一轮验证// Get the model without autodiff. let model_valid model.valid(); // Implement our validation loop. for (iteration, batch) in dataloader_test.iter().map(Result::unwrap).enumerate() { let output model_valid.forward(batch.images); let loss CrossEntropyLoss::new(None, output.device()) .forward(output.clone(), batch.targets.clone()); let accuracy accuracy(output, batch.targets); println!( [Valid - Epoch {} - Iteration {}] Loss {} | Accuracy {}, epoch, iteration, loss.clone().into_scalar::f32(), accuracy, ); }核心是model.valid()它返回一个参数位于不具备自动微分能力设备上的模型副本从而在验证阶段彻底关闭梯度追踪与计算图记录节省内存与算力。从 crates/burn-core/src/module/base.rs 的AutodiffModule::valid文档可以看到更精细的语义返回的模型中所有张量的梯度需求require_grad被禁用模块自带的训练标志如 Dropout 的开启状态也被关闭这些被禁用的状态会保留在内部之后可用Module::train恢复——也就是说valid()不会破坏模型原本的训练状态在普通非 autodiff设备上调用valid()依然会关闭训练标志因此该操作是幂等的。ParamFlag的valid实现见 crates/burn-core/src/module/param/flag.rs印证了这一点它保留is_active原值但把内部标志值置为关闭状态验证循环中使用的是关闭状态下的网络行为例如 Dropout 不再随机丢弃。验证完成后继续下一个 epoch 时直接使用原始的model训练态即可无需额外操作。进阶一Parameter Groups——对模型不同部分使用不同优化器与学习率当模型的不同部分需要不同优化器或不同学习率时Learner场景下可以通过ModuleOptimizer与模块学习率调度器上的ParamGroup来配置。ParamGroup按参数 ID 或参数在模块中的路径进行选择常用的构造方式如下表详见 burn-book/src/building-blocks/optimizer.md构造函数匹配范围ParamGroup::all()所有参数ParamGroup::from_ids(ids)显式指定的参数 ID 列表ParamGroup::from_path(encoder.weight)精确的模块路径ParamGroup::from_predicate(encoder)路径中包含该关键字的参数ParamGroup::from_regex(pattern)?路径匹配正则表达式的参数组之间可以组合、也可以排除另一组例如选择 encoder 中除 bias 外的所有参数use burn::module::ParamGroup; let encoder ParamGroup::from_predicate(encoder) .exclude(ParamGroup::from_predicate(bias));通过ModuleOptimizer::with_group为特定组挂载专属优化器可附带该组专属的梯度裁剪配置let optimizer default_optimizer.with_group( ParamGroup::from_predicate(encoder), encoder_optimizer, None, // Optional gradient clipping for this group. );with_group的实现见 crates/burn-optim/src/optim/module/module_optimizer.rs说明了几个关键语义第一个初始优化器作为全局兜底必须匹配所有参数若一个参数同时命中多个组最后添加的组优先源码中通过next_back()取最后一个匹配组在优化已经开始后新增分组会清空被该组命中的参数已有优化器状态因为这些状态可能属于不同类型的优化器见self.param_context.retain(...)的清理逻辑。分组配置完成后optim.step(lr, model, grads)会在同一次调用中按组路由每个参数及其状态训练循环无需手动拆分梯度容器或分批对子模块执行 step。学习率调度器也支持相同的分组模型因此优化器选择与学习率策略可以相互独立地分配。在自定义训练循环中只需把optim换成配置好分组的ModuleOptimizer其余训练代码无需任何改动。进阶二自定义类型——用泛型把模型与优化器打包如果你要在一个类型里同时持有模型与优化器可以定义一个对两者都泛化的结构体且无需引入后端类型参数struct LearnerM, O { model: M, optim: O, }trait 约束如M: AutodiffModule、O: Optimizer可以只加在实际使用M和O的方法上而不是堆在结构体定义处这保持了抽象与后端实现的解耦。结合 Burn 的运行时设备分发机制这类泛型封装可以在不同后端CPU、CUDA、WGPU 等间自由迁移而无需改动逻辑。若你的自定义循环需要跨设备移动整个模型可进一步利用GradientsParams::to_device与Module::to_device配合完成梯度与参数的搬运。运行完整示例本文涉及的全部代码都已内置在仓库的 examples/custom-training-loop 目录中examples/custom-training-loop/src/lib.rs训练配置、训练/验证循环与accuracy函数的完整实现examples/custom-training-loop/examples/custom-training-loop.rs入口调用custom_training_loop::run(Device::default())启动配套的模型与 Batcher 复用了 examples/guide 中的ModelConfig与MnistBatcher。编译运行方式与仓库其他示例一致cargo run --example custom-training-loop --release需在对应示例目录下执行或通过 workspace 指定。运行时会在终端逐迭代输出[Train - Epoch x - Iteration y] Loss ... | Accuracy ... %与[Valid - ...]形式的日志你可以据此直观对照文中每一步的执行效果。小结手写训练循环并不意味着重造轮子Burn 把Learner之上的所有高层能力都沉淀成了粒度适中的底层原语——forward/backward()负责梯度求解GradientsParams负责梯度与参数的键值映射optim.step()负责消费梯度并推进优化器状态GradientsAccumulator负责跨批次梯度累积valid()负责验证阶段的梯度关闭。理解并组合这些原语后你既能在需要时完全掌控训练的每一个细节也能在需要时随时回到Learner的便捷轨道上实现可进可退的训练流程设计。【免费下载链接】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),仅供参考