Apex Amp 高级用法实战:梯度裁剪、多模型/多优化器/多损失训练与自定义数据批处理
发布时间:2026/9/26 8:04:57来源:尧图网络
人工智能大模型音乐生成音频预训练【免费下载链接】jukeboxCode for the paper Jukebox: A Generative Model for Music项目地址https://gitcode.com/gh_mirrors/ju/jukebox点击查看免费下载本文以 advanced.rst 为主线系统讲解 Apex 混合精度训练模块 Amp 在基础用法之上的高级主题master params 概念与正确的梯度裁剪方式、用户自定义 autograd 函数的注册、多模型/多优化器/多损失场景下的初始化与反向传播、跨迭代梯度累积以及如何让 Amp 识别并自动转换自定义数据批类型。读完本文你将掌握 Amp 在这些复杂训练场景下的正确调用姿势并能结合仓库源码apex/apex/amp/目录理解其底层实现原理避免在真实训练脚本中踩坑。引言什么是 Amp 的高级用法AmpAutomatic Mixed Precision的入门用法——调用amp.initialize并选择opt_level——解决的是如何让模型跑起来的问题而本文对应的高级用法解决的是如何在各种真实训练架构下不出错的问题。这些高级主题包括多生成器/判别器与多损失的 GAN 训练、需要梯度裁剪的稳定性训练、跨迭代累积梯度的长序列训练以及数据管线中出现的自定义批数据类型。需要先明确一个贯穿全文的核心概念master params主参数。Amp 把优化器param_groups直接拥有的参数称为 master params。不同opt_level下master params 与model.parameters()的关系可能完全不同在opt_levelO2下amp.initialize会把模型大部分参数强制转换为 FP16同时在模型外部为每个新转换的 FP16 参数创建一个 FP32 master param并让优化器的param_groups指向这些 FP32 参数在opt_level为O0、O1、O3时master params 通常与模型参数完全重合O0全部保持 FP32、O3全部转 FP16、O1模型权重本身保持 FP32 而只在算子层做 cast。这一概念直接决定了梯度裁剪的正确对象是下一节的核心。梯度裁剪Gradient clipping的正确姿势为什么要裁剪 master params 而不是model.parameters()在混合精度训练中梯度是在被 loss scale 放大的状态下产生的且可能分别存在于 FP16 模型参数与 FP32 master 参数两套对象上。正确的实践是始终裁剪优化器param_groups所直接拥有的那些参数即 master params的梯度而不是通过model.parameters()取回的模型参数。此外如果 Amp 使用了 loss scaling梯度必须在完成 unscale取消缩放之后再裁剪。unscale 动作发生在退出amp.scale_loss上下文管理器时。这一点在opt_levelO2场景下尤为关键。从源码看_process_optimizer.py中的lazy_init_with_master_weights会为每个 FP16 参数执行类似master_param param.detach().clone().float()的操作创建独立的 FP32 master 参数并替换param_group[params]中的引用而handle.py中scale_loss的 docstring 明确警告使用显式 FP32 master params 时只有 FP32 master 梯度会被 unscaleFP16 模型参数.grad在退出上下文后仍然是缩放后的状态。因此若对model.parameters()做裁剪裁剪到的将是未取消缩放的梯度数值完全错误。适用于任意 opt_level 的统一模式以下模式对任何opt_levelO0–O3都正确with amp.scale_loss(loss, optimizer) as scaled_loss: scaled_loss.backward() # 梯度在退出上下文管理器时完成 unscale # 现在可以安全地裁剪了。请把 # torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm) # 替换为 torch.nn.utils.clip_grad_norm_(amp.master_params(optimizer), max_norm) # 或者 torch.nn.utils.clip_grad_value_(amp.master_params(optimizer), max_)这里使用了工具函数amp.master_params(optimizer)。查看_amp_state.py可知它是一个生成器表达式遍历优化器param_groups中的所有参数def master_params(optimizer): for group in optimizer.param_groups: for p in group[params]: yield p还需特别强调clip_grad_norm_(amp.master_params(optimizer), max_norm)是替代instead of而非叠加in addition toclip_grad_norm_(model.parameters(), max_norm)。两者同时调用会导致梯度被重复裁剪破坏训练稳定性。自定义/用户自定义 autograd 函数对于用户自定义的 autograd 函数自定义Function或普通函数Amp 的旧 API——注册用户函数——仍然被视为正确且有效的做法。仓库中amp.py提供了注册式与装饰器式两套入口# 注册式Registry form把模块中的指定函数注册为强制 FP16 / FP32 / 提升 amp.register_half_function(module, name) amp.register_float_function(module, name) amp.register_promote_function(module, name) # 装饰器式Decorator form直接装饰函数 amp.half_function(fn) amp.float_function(fn) amp.promote_function(fn)从实现看注册式入口会把(module, name, cast_fn)元组放入全局_USER_CAST_REGISTRY/_USER_PROMOTE_REGISTRY集合见amp.py随后在amp.init()的第 0 步被依次取出通过wrap.cached_cast/wrap.promote包裹为目标函数见amp.py与wrap.py。关键约束这些函数必须在调用amp.initialize之前完成注册因为注册表在初始化阶段才会被消费并清空。若在amp.initialize之后注册包裹动作不会生效。强制特定层/函数到目标类型文档明确指出Amp 目前仍在寻求一种通用化的暴露方式使得该能力在不同opt_level之间不需要用户侧代码分叉因此这一节在文档中属于进行中的工作原文为 Im still working on a generalizable exposure for this...。从当前仓库实现看可以借助以下既有机制实现让某层/某函数以指定类型运行的等效效果但需要注意它们与opt_level的耦合通过frontend.py中amp.initialize的属性覆盖参数如cast_model_type、keep_batchnorm_fp32、master_weights、loss_scale整体调节模型与算子的精度分布例如opt_levelO2默认keep_batchnorm_fp32True使 BatchNorm 保持在 FP32 以提升稳定性对单个自定义函数使用上一节的自定义函数注册机制强制其输入 cast 到 FP16/FP32对自定义数据批类型使用自定义数据批类型一节的to(dtype)机制控制进入模型的数据精度。在使用覆盖参数时需留意frontend.py中Properties.__setattr__的一致性检查例如O1下cast_model_type与master_weights不应被设置为非None值否则会触发warn_or_err。这正是文档所说不同 opt-level 间代码会分叉的体现也是官方尚未将其通用化的原因。多模型 / 多优化器 / 多损失使用多个模型/优化器完成初始化amp.initialize的optimizer参数可以是单个优化器或优化器列表model参数可以是单个模型或模型列表只要返回值的类型与你接受的类型一致即可。以下调用全部合法model, optim amp.initialize(model, optim, ...) model, [optim0, optim1] amp.initialize(model, [optim0, optim1], ...) [model0, model1], optim amp.initialize([model0, model1], optim, ...) [model0, model1], [optim0, optim1] amp.initialize([model0, model1], [optim0, optim1], ...)从源码看_initialize.py会通过isinstance(optimizers, torch.optim.Optimizer)/isinstance(models, torch.nn.Module)判断输入是单个对象还是列表并在返回时保持列表进、列表出单个进、单个出的类型对称性。多优化器的反向传播scale_loss 必须覆盖全部相关优化器无论何时执行一次反向传播amp.scale_loss上下文管理器都必须接收所有拥有本次反向传播所产生梯度的参数的优化器——即使某个优化器只拥有其中一部分参数也是如此。如果某次反向传播只有一个优化器的参数会产生梯度可以直接把该优化器传给amp.scale_loss否则必须传入将要产生梯度的优化器列表# loss0 只把梯度累积到 optim0 拥有的参数上 with amp.scale_loss(loss0, optim0) as scaled_loss: scaled_loss.backward() # loss1 只把梯度累积到 optim1 拥有的参数上 with amp.scale_loss(loss1, optim1) as scaled_loss: scaled_loss.backward() # loss2 同时把梯度累积到 optim0 与 optim1 拥有的参数上 with amp.scale_loss(loss2, [optim0, optim1]) as scaled_loss: scaled_loss.backward()查看handle.py的scale_loss实现可以印证入口处会判断optimizers是否为单个优化器若是则包装成列表统一处理退出上下文时遍历所有传入优化器调用其_post_amp_backward(loss_scaler)完成梯度 unscale。这正是必须传入全部相关优化器的底层原因——漏掉任何一个该优化器拥有的参数梯度就不会被正确取消缩放。可选让 Amp 为每个损失维护独立的 loss scaler默认情况下Amp 维护一个全局 loss scaler所有反向传播所有with amp.scale_loss(...)调用共用它。使用全局 loss scaler 不需要给amp.initialize或amp.scale_loss增加任何参数上面的多优化器示例底层用的就是这一个全局 scaler可以直接工作。不过你也可以让 Amp 为每个损失单独维护一个 loss scaler以获得更大的数值灵活性。做法分两步向amp.initialize传入num_losses参数告诉 Amp 你计划执行多少次反向传播从而创建多少个 loss scaler在每次反向传播中向amp.scale_loss传入loss_id参数指明本次反向传播使用哪一个 scalermodel, [optim0, optim1] amp.initialize(model, [optim0, optim1], ..., num_losses3) with amp.scale_loss(loss0, optim0, loss_id0) as scaled_loss: scaled_loss.backward() with amp.scale_loss(loss1, optim1, loss_id1) as scaled_loss: scaled_loss.backward() with amp.scale_loss(loss2, [optim0, optim1], loss_id2) as scaled_loss: scaled_loss.backward()num_losses与loss_id应纯粹基于损失/反向传播的集合来确定使用多少个优化器、每个反向传播关联单个还是多个优化器与这套编号无关。从源码看这一机制如何落地在_initialize.py中_initialize会清空_amp_state.loss_scalers列表并按num_losses循环创建对应数量的LossScaler实例而handle.py的scale_loss会通过loss_scaler _amp_state.loss_scalers[loss_id]取出与本次反向传播对应的 scaler 来计算loss_scale。每个 scaler 独立跟踪溢出状态与动态缩放轨迹互不干扰。仓库中的测试 test_multiple_models_optimizers_losses.py 对上述场景做了系统验证它覆盖了 2 模型 2 损失 1 优化器含多个 param group、2 模型 2 损失 2 优化器、以及一个损失同时累积到两个优化器参数上等组合并在每个opt_levelO0–O3下分别以num_losses1全局 scaler与num_losses2每损失独立 scaler两种模式对比amp.master_params与参考梯度同时注入 inf 梯度检验动态 loss scaling 的跳过逻辑。这说明多模型/多优化器/多损失组合是 Amp 官方测试覆盖的受支持场景。跨迭代梯度累积Gradient accumulation当显存不足以一次性容纳大批次时常见做法是跨多个迭代累积梯度。Amp 支持这一模式开箱即用并且能正确适配多个模型/优化器/损失以及前述的梯度裁剪规则if iter % iters_to_accumulate 0: # 每 iters_to_accumulate 个迭代取消缩放并执行 step with amp.scale_loss(loss, optimizer) as scaled_loss: scaled_loss.backward() # 如需梯度裁剪 # torch.nn.utils.clip_grad_norm_(amp.master_params(optimizer), max_norm) optimizer.step() optimizer.zero_grad() else: # 其余迭代只累积梯度不取消缩放、不 step with amp.scale_loss(loss, optimizer) as scaled_loss: scaled_loss.backward()作为一项小的性能优化你可以向amp.scale_loss传入delay_unscaleTrue将 unscale 推迟到真正准备step()时再进行。但只有在你完全清楚自己在做什么时才应尝试delay_unscaleTrue因为它与梯度裁剪、多模型/多优化器/多损失的交互会变得很微妙if iter % iters_to_accumulate 0: # 每 iters_to_accumulate 个迭代取消缩放并执行 step with amp.scale_loss(loss, optimizer) as scaled_loss: scaled_loss.backward() optimizer.step() optimizer.zero_grad() else: # 其余迭代只累积梯度不取消缩放、不 step with amp.scale_loss(loss, optimizer, delay_unscaleTrue) as scaled_loss: scaled_loss.backward()源码层面的依据同样在handle.py当delay_unscaleTrue时scale_loss退出时不会调用_post_amp_backward而是给优化器设置params_have_scaled_gradients True标记直到某次delay_unscaleFalse的调用退出时才会真正执行 unscale 并调用loss_scaler.update_scale()更新动态缩放系数。update_scale()的实现位于scaler.py若检测到溢出则把 loss scale 减半若设置了min_loss_scale则取两者较大者并跳过本次 stepshould_skipTrue若连续scale_window默认 2000次未溢出则把 loss scale 加倍但上限为max_loss_scaleamp.initialize默认2**24。自定义数据批类型Custom data batch types原理Amp 通过 patch forward 来转换输入数据Amp 的设计意图是无论opt_level是什么你都不需要手动 cast 输入数据。Amp 通过 patch 模型的forward方法将传入数据按opt_level自动转换为合适类型。查看_initialize.py的patch_forward可以确认在properties.cast_model_type存在时新的forward会对所有输入args 与 kwargs先做applier(args, input_caster)转换再调用原始forward并把输出 cast 回torch.float32除非显式指定cast_model_outputs。要完成输入转换Amp 需要知道怎么转换。被打包的forward能识别并转换浮点型 Tensor非浮点 Tensor如 IntTensor不会被触碰浮点 Tensor 的 Python 容器list、tuple、dict 等。给自定义批类型加一个to方法但如果你把 Tensor 包装在自定义类里转换逻辑无法穿透这层自定义外壳去访问并转换里面的 Tensor。你需要告诉 Amp 如何转换自定义批类型为它实现一个接受torch.dtype如torch.float16或torch.float32的to方法并返回转换到该dtype的自定义批实例。被打包的forward会检查你的to方法是否存在并以该opt_level对应的正确类型调用它。示例class CustomData(object): def __init__(self): self.tensor torch.cuda.FloatTensor([1, 2, 3]) def to(self, dtype): self.tensor self.tensor.to(dtype) return self源码印证在_initialize.py的applier函数中判断顺序为——torch.Tensor走fn(value)转换string_classes字符串与np.ndarray原样返回hasattr(value, to)的自定义类走fn(value)Mapping与Iterable容器递归处理。注意to_type同一文件第 18–32 行对 Tensor 会先检查is_floating_point()只转换浮点 Tensor非浮点 Tensor 直接返回。两个重要注意事项numpy ndarray 会被直接透传而不做转换。Amp 会原样转发 numpy ndarray。如果你把输入数据以裸的、未包装的 ndarray 形式传入然后在model.forward内部用它在 CUDA 上创建 Tensor那么这个 Tensor 的类型将不依赖于opt_level其精度可能正确也可能不正确。文档建议尽可能传可转换的数据输入——即 Tensor、Tensor 集合、或带有to方法的自定义类。Amp 不会替你调用.cuda()。Amp 假设你的原始脚本已经处理好把 Tensor 从主机搬到设备的流程。输入 Tensor 需要已经在正确的设备上Amp 只负责精度dtype的转换不负责设备的迁移。综合案例GAN 训练GAN生成对抗网络是上述多个高级主题的典型综合场景生成器与判别器是两个独立模型甚至各自带优化器生成对抗损失与重建损失等多个损失并存梯度裁剪也常被用于稳定训练。文档注明一个综合性示例尚在建设中仓库中对应的实验目录为 apex/examples/dcgan当前以 README 说明为主。结合本文内容GAN 场景的 Amp 接入要点可以归纳为使用amp.initialize([gen, disc], [opt_gen, opt_disc], opt_level..., num_lossesN)一次性初始化多个模型与优化器每个损失的反向传播通过with amp.scale_loss(loss, opt) as scaled_loss: scaled_loss.backward()完成涉及多个优化器的损失则传入优化器列表需要梯度裁剪时对amp.master_params(optimizer)裁剪并确保裁剪发生在scale_loss上下文退出unscale 完成之后若 GAN 存在多个来源不同的损失可考虑num_lossesloss_id让 Amp 为每个损失维护独立 loss scaler避免某个频繁溢出的损失拖累全局缩放系数。进阶阅读与源码索引Amp 入门与opt_levelO0–O3完整说明amp.rst本文对应的原始高级用法文档advanced.rstamp.initialize/amp.scale_loss/amp.master_params入口与全部参数说明frontend.py、handle.py、_amp_state.py模型/优化器初始化与 forward patch 实现_initialize.pymaster weightsO2创建与拷贝回模型权重的实现_process_optimizer.pyLossScaler动态/静态 loss scaling、溢出检测与缩放调整scaler.py多模型/多优化器/多损失与动态 loss scaling 的系统性测试test_multiple_models_optimizers_losses.py自定义函数注册相关的白名单/黑名单算子表functional_overrides.py、torch_overrides.py、tensor_overrides.py赞分享人工智能大模型音乐生成音频预训练【免费下载链接】jukeboxCode for the paper Jukebox: A Generative Model for Music项目地址https://gitcode.com/gh_mirrors/ju/jukebox点击查看免费下载相关推荐AGiXT安全最佳实践API密钥管理、OAuth认证与数据保护完全指南AGiXT安全最佳实践API密钥管理、OAuth认证与数据保护完全指南 AGiXT是一个动态的AI Agent自动化平台提供强大的AI代理编排和任务执行能力DETR高级技巧自定义数据集训练与模型优化实战DETR高级技巧自定义数据集训练与模型优化实战 引言告别COCO依赖掌握DETR定制化训练 你是否还在为DETR只能处理COCO数据集而烦恼是否想将DE人工智能计算机视觉深度学习告别梯度爆炸Pytorch-UNet模型训练的梯度裁剪实战指南告别梯度爆炸Pytorch UNet模型训练的梯度裁剪实战指南 你是否在训练U Net模型时遇到过损失值突然飙升、模型无法收敛的问题这很可能是梯度爆炸Gr人工智能深度学习计算机视觉图像处理上一篇Joy-Con Toolkit完整指南免费开源的手柄终极定制工具下一篇如何快速清理Windows驱动冗余DriverStore Explorer新手完全指南创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
网站建设高端定制企业官网