練實(shí)戰(zhàn):從 opt_level 到統(tǒng)一 API 的完整指南)
人工智能大模型音樂(lè)生成音頻預(yù)訓(xùn)練【免費(fèi)下載鏈接】jukeboxCode for the paper Jukebox: A Generative Model for Music項(xiàng)目地址https://gitcode.com/gh_mirrors/ju/jukebox點(diǎn)擊查看免費(fèi)下載本文以 NVIDIA Apex 倉(cāng)庫(kù)中 amp.rst 文檔為主線系統(tǒng)講解apex.ampAutomatic Mixed Precision這一自動(dòng)混合精度工具如何在只改動(dòng) 3 行代碼的前提下啟用 Tensor Core 加速訓(xùn)練如何理解 O0–O3 四個(gè)優(yōu)化級(jí)別及其背后的六個(gè)核心屬性以及新舊 API 的遷移要點(diǎn)。讀完本文你將能夠?yàn)樽约旱?PyTorch 訓(xùn)練腳本正確接入amp.initialize/amp.scale_loss并針對(duì)不同模型挑選最合適的精度策略。背景Amp 解決什么問(wèn)題混合精度訓(xùn)練的核心矛盾在于FP16 能利用 NVIDIA Tensor Core 成倍加速 GEMM 與卷積等算子但 FP16 的動(dòng)態(tài)范圍約 5 個(gè)數(shù)量級(jí)遠(yuǎn)小于 FP32直接以 FP16 訓(xùn)練容易導(dǎo)致梯度下溢underflow、精度損失甚至發(fā)散。手動(dòng)在腳本中穿插.half()轉(zhuǎn)換不僅繁瑣還極易出錯(cuò)。apex.amp的思路是讓用戶以默認(rèn)的 FP32 方式構(gòu)建模型與數(shù)據(jù)由 Amp 在后臺(tái)自動(dòng)完成何時(shí)轉(zhuǎn) FP16、何時(shí)保留 FP32、如何縮放損失的決策。正如倉(cāng)庫(kù) README 所描述的apex.amp是通過(guò)只修改腳本中的 3 行代碼即可啟用混合精度訓(xùn)練的工具。在 amp.rst 中Amp 給出了一個(gè)完整的最小示例# 以默認(rèn) FP32 精度聲明模型和優(yōu)化器 model torch.nn.Linear(D_in, D_out).cuda() optimizer torch.optim.SGD(model.parameters(), lr1e-3) # 讓 Amp 按 opt_level 的要求執(zhí)行類型轉(zhuǎn)換 model, optimizer amp.initialize(model, optimizer, opt_levelO1) ... # loss.backward() 改為 with amp.scale_loss(loss, optimizer) as scaled_loss: scaled_loss.backward()關(guān)鍵原則是無(wú)論選擇哪個(gè)opt_level用戶都不應(yīng)手動(dòng)對(duì)模型或數(shù)據(jù)調(diào)用.half()。Amp 希望用戶從一段現(xiàn)有的 FP32 腳本出發(fā)加入上述三行代碼即可開始混合精度訓(xùn)練當(dāng)需要關(guān)閉 Amp 時(shí)enabledFalse原腳本的行為與未接入 Amp 時(shí)完全一致。因此遵循 Amp API 沒(méi)有風(fēng)險(xiǎn)卻可能帶來(lái)可觀的性能收益。由于無(wú)需手動(dòng)轉(zhuǎn)換模型與輸入數(shù)據(jù)一段遵循新 API 的腳本可以僅通過(guò)更換opt_level在多種純精度/混合精度模式間切換而不需要改動(dòng)任何其他代碼。opt_level 與屬性理解 Amp 的配置模型Amp 允許用戶方便地實(shí)驗(yàn)不同的純精度與混合精度模式。常用默認(rèn)模式通過(guò)選擇優(yōu)化級(jí)別opt_level來(lái)指定每個(gè)opt_level定義了一組控制 Amp 實(shí)現(xiàn)純精度或混合精度訓(xùn)練的屬性properties。你也可以向amp.initialize直接傳入特定屬性的值實(shí)現(xiàn)比默認(rèn)配置更細(xì)粒度的控制——手動(dòng)指定的值會(huì)覆蓋opt_level建立的默認(rèn)值。從源碼看這六個(gè)屬性在 frontend.py 的Properties類中統(tǒng)一管理并注冊(cè)為一個(gè)options字典默認(rèn)值分別為enabledFalse、opt_levelNone、cast_model_typeNone、patch_torch_functionsFalse、keep_batchnorm_fp32None、master_weightsNone、loss_scale1.0。屬性及其語(yǔ)義如下屬性作用cast_model_type將模型的參數(shù)與緩沖區(qū)buffers轉(zhuǎn)換為目標(biāo)類型。patch_torch_functions對(duì)所有 Torch 函數(shù)與 Tensor 方法進(jìn)行打補(bǔ)丁patch將適合 Tensor Core 的算子如 GEMM、卷積以 FP16 執(zhí)行將受益于 FP32 精度的算子以 FP32 執(zhí)行。keep_batchnorm_fp32即便模型其余部分為 FP16也建議將 BatchNorm 權(quán)重保留在 FP32——這既提升精度又能啟用 cudnn BatchNorm改善性能。master_weights為 FP16 模型權(quán)重維護(hù)一份 FP32 主權(quán)重master weights。優(yōu)化器對(duì) FP32 主權(quán)重執(zhí)行 step以提升精度并捕獲小梯度。loss_scale若為浮點(diǎn)數(shù)則作為靜態(tài)固定損失縮放值若為字符串dynamic則由 Amp 自動(dòng)隨時(shí)間自適應(yīng)調(diào)整縮放值。屬性覆蓋與一致性檢查選中某個(gè)opt_level后你可以通過(guò)屬性關(guān)鍵字參數(shù)做手動(dòng)覆蓋。但如果你試圖覆蓋一個(gè)對(duì)所選opt_level而言沒(méi)有意義的屬性Amp 會(huì)拋出帶解釋的錯(cuò)誤。例如為opt_levelO1指定master_weightsTrue是不合理的O1是在 Torch 函數(shù)周圍插入轉(zhuǎn)換cast而不是轉(zhuǎn)換模型權(quán)重?cái)?shù)據(jù)、激活和權(quán)重在流經(jīng)被 patch 的函數(shù)時(shí)被就地動(dòng)態(tài)重新轉(zhuǎn)換out-of-place recast on the fly。因此 O1 下模型權(quán)重本身可以也應(yīng)該保持 FP32無(wú)需維護(hù)獨(dú)立的 FP32 主權(quán)重。這一檢查在源碼中有明確體現(xiàn)。frontend.py 的Properties.__setattr__對(duì)每個(gè)屬性做了守衛(wèi)例如設(shè)置master_weights時(shí)若當(dāng)前opt_level O1且值為非 None會(huì)調(diào)用warn_or_err報(bào)錯(cuò)設(shè)置cast_model_type時(shí)若為 O1 且值既非False也非torch.float32同樣會(huì)警告。此外keep_batchnorm_fp32被限定為布爾值、字符串True/False或None否則觸發(fā)斷言。四個(gè)優(yōu)化級(jí)別 O0–O3 詳解Amp 識(shí)別的opt_level為O0、O1、O2、O3四種。其中O0與O3并非真正的混合精度但分別對(duì)建立精度基線與速度基線很有價(jià)值O1與O2是混合精度的兩種不同實(shí)現(xiàn)建議兩種都嘗試看哪種對(duì)具體模型帶來(lái)最好的加速與精度。注意O0/O1中的前綴是大寫字母 Othe letter O不是數(shù)字 0。源碼 frontend.py 的initialize中對(duì)非法的opt_level會(huì)拋出RuntimeError并明確提示這一點(diǎn)。在 frontend.py 中四個(gè)級(jí)別分別由O0、O1、O2、O3類實(shí)現(xiàn)它們的__call__方法直接對(duì)Properties對(duì)象寫入默認(rèn)值注冊(cè)在opt_levels字典中供initialize調(diào)用。O0FP32 訓(xùn)練精度基線O0對(duì)你的模型而言通常是 no-op你的模型本來(lái)就是 FP32可用于建立精度基線。O0默認(rèn)屬性值cast_model_typetorch.float32patch_torch_functionsFalsekeep_batchnorm_fp32None不適用一切皆為 FP32master_weightsFalseloss_scale1.0O1混合精度常規(guī)使用推薦O1對(duì)所有 Torch 函數(shù)與 Tensor 方法打補(bǔ)丁按白名單-黑名單whitelist-blacklist模型轉(zhuǎn)換其輸入白名單算子如適合 Tensor Core 的 GEMM、卷積以 FP16 執(zhí)行黑名單中受益于 FP32 精度的算子如 softmax以 FP32 執(zhí)行。O1默認(rèn)使用動(dòng)態(tài)損失縮放除非被覆蓋。O1默認(rèn)屬性值cast_model_typeNone不適用patch_torch_functionsTruekeep_batchnorm_fp32None不適用所有模型權(quán)重保持 FP32master_weightsNone不適用模型權(quán)重保持 FP32loss_scaledynamicO1 的白名單/黑名單機(jī)制源碼級(jí)白名單與黑名單列表定義在 lists 目錄 下的三個(gè)文件中torch_overrides.pyFP16_FUNCS包含conv1d/2d/3d、conv_transpose*、prelu以及 BLAS 類算子addmm、addmv、addr、matmul、mm、mvFP32_FUNCS包含exp、log、mean、sum、std、var、norm、pow等CASTS需要類型提升的多 Tensor 運(yùn)算包含add、div、mul、比較類算子等SEQUENCE_CASTS包含cat、stack。值得注意的是bmm/addbmm/baddbmm會(huì)根據(jù) CUDA 版本動(dòng)態(tài)歸類CUDA ≥ 9.1 才具備快速 FP16 batched matmul 內(nèi)核因此版本足夠新時(shí)加入 FP16 列表否則歸入 FP32 列表。functional_overrides.pyFP16_FUNCS為torch.nn.functional下的conv1d/2d/3d、conv_transpose*、linearFP32_FUNCS包括softmax、log_softmax、softmin、softplus、layer_norm、group_norm、各類 losscross_entropy、mse_loss、nll_loss等以及interpolate。tensor_overrides.pyTensor 方法層面的對(duì)應(yīng)列表。此外amp.py 中的init()會(huì)遍歷這些列表對(duì)白名單函數(shù)用cached_cast包裝maybe_halfFP16 轉(zhuǎn)換帶緩存、對(duì)黑名單函數(shù)用cached_cast包裝maybe_float不緩存并處理CASTS/SEQUENCE_CASTS的類型提升以及對(duì) RNN 及其 cell 的專門白名單處理。對(duì)于BANNED_FUNCS如F.binary_cross_entropyAmp 默認(rèn)在遇到 FP16 輸入時(shí)報(bào)錯(cuò)并給出替換建議改用BCEWithLogitsLoss或注冊(cè)amp.register_float_function(torch, sigmoid)。O2幾乎全 FP16 混合精度O2將模型權(quán)重轉(zhuǎn)換為 FP16patch 模型的forward方法以將輸入數(shù)據(jù)轉(zhuǎn)換為 FP16將 BatchNorm 保留在 FP32維護(hù) FP32 主權(quán)重并更新優(yōu)化器的param_groups使optimizer.step()直接作用于 FP32 權(quán)重必要時(shí)隨后執(zhí)行 FP32 主權(quán)重→FP16 模型權(quán)重的拷貝同時(shí)實(shí)現(xiàn)動(dòng)態(tài)損失縮放除非被覆蓋。與O1不同O2不patch Torch 函數(shù)或 Tensor 方法。O2默認(rèn)屬性值cast_model_typetorch.float16patch_torch_functionsFalsekeep_batchnorm_fp32Truemaster_weightsTrueloss_scaledynamicO2 的 master weights 機(jī)制源碼級(jí)在 _process_optimizer.py 的lazy_init_with_master_weights中每個(gè)torch.cuda.HalfTensor參數(shù)都會(huì)被替換為一份 FP32 主參數(shù)param.detach().clone().float()且requires_gradTrue并用load_state_dict(state_dict())技巧將原有的 per-param 狀態(tài)張量重鑄為 FP32。optimizer.step被包裝為先執(zhí)行原old_step()再調(diào)用_master_params_to_model_params()有 C 擴(kuò)展時(shí)通過(guò)multi_tensor_applier調(diào)用amp_C.multi_tensor_scale一次完成全部參數(shù)拷貝最后清空主梯度。zero_grad也被重寫為只清零模型梯度而置空主梯度。O3FP16 訓(xùn)練速度基線O3可能無(wú)法達(dá)到真正混合精度選項(xiàng)O1/O2的穩(wěn)定性但它可用于建立模型的速度基線供O1/O2的性能對(duì)比參考。若模型使用 BatchNorm為建立光速speed of light基線可嘗試在O3上額外覆蓋keep_batchnorm_fp32True如前所述這會(huì)啟用 cudnn BatchNorm。O3默認(rèn)屬性值cast_model_typetorch.float16patch_torch_functionsFalsekeep_batchnorm_fp32Falsemaster_weightsFalseloss_scale1.0從 frontend.py 中O1、O2、O3類的brief/more文檔字符串還可以看到官方的推薦口徑O1是首次嘗試混合精度訓(xùn)練時(shí)最安全的方式O2的 FP32 主權(quán)重還可以改善收斂與穩(wěn)定性O(shè)3用于建立性能上限。實(shí)戰(zhàn)中通常的做法是先用O0確認(rèn)精度基線用O3觀察速度上限然后在O1/O2中選擇精度與速度平衡最佳者。統(tǒng)一 APIinitialize、scale_loss 與 master_params新版 Amp 將所有功能收斂到apex.amp模塊下核心入口有三個(gè)amp.initialize、amp.scale_loss、amp.master_params。它們由init.py 統(tǒng)一導(dǎo)出initialize來(lái)自 frontend.pyscale_loss/disable_casts來(lái)自 handle.pymaster_params來(lái)自 _amp_state.py。amp.initialize 的完整簽名與參數(shù)amp.initialize(models, optimizersNone, enabledTrue, opt_levelO1, cast_model_typeNone, patch_torch_functionsNone, keep_batchnorm_fp32None, master_weightsNone, loss_scaleNone, cast_model_outputsNone, num_losses1, verbosity1, min_loss_scaleNone, max_loss_scale2.**24)各參數(shù)含義依據(jù) frontend.py 的 docstring 與實(shí)現(xiàn)modelstorch.nn.Module或模塊列表即待修改/轉(zhuǎn)換的模型。optimizers可選優(yōu)化器或優(yōu)化器列表訓(xùn)練時(shí)必需推理時(shí)可選。enabled默認(rèn)True若為False所有 Amp 調(diào)用變?yōu)?no-op腳本如同未接入 Amp 一樣運(yùn)行。opt_level默認(rèn)O1接受O0/O1/O2/O3。cast_model_type/patch_torch_functions/keep_batchnorm_fp32/master_weights/loss_scale均為可選屬性覆蓋任何非None的關(guān)鍵字參數(shù)都會(huì)被解釋為手動(dòng)覆蓋。cast_model_outputs默認(rèn)None確保模型輸出始終轉(zhuǎn)換為指定類型無(wú)論opt_level。num_losses默認(rèn)1預(yù)先告知 Amp 將使用多少個(gè) loss/backward 次數(shù)配合amp.scale_loss的loss_id參數(shù)可為每個(gè) loss/backward 使用不同的損失縮放值提升穩(wěn)定性。verbosity默認(rèn)1設(shè)為 0 可抑制 Amp 相關(guān)輸出。min_loss_scale默認(rèn)None動(dòng)態(tài)損失縮放可選值的地板不用動(dòng)態(tài)縮放時(shí)忽略。max_loss_scale默認(rèn)2.**24動(dòng)態(tài)損失縮放可選值的上限不用動(dòng)態(tài)縮放時(shí)忽略。合法調(diào)用形式模型與優(yōu)化器均可為單個(gè)對(duì)象或列表返回形式與傳入形式保持一致model, optim amp.initialize(model, optim, ...) model, [optim1, optim2] amp.initialize(model, [optim1, optim2], ...) [model1, model2], optim amp.initialize([model1, model2], optim, ...) [model1, model2], [optim1, optim2] amp.initialize([model1, model2], [optim1, optim2], ...)一個(gè)很實(shí)用的特性是loss_scale與keep_batchnorm_fp32都接受數(shù)值或字符串如128.0、dynamic、True/False便于與 argparse 直接互操作——這一點(diǎn)在 Imagenet 示例 中體現(xiàn)得很清楚它通過(guò)--opt-level、--keep-batchnorm-fp32、--loss-scale三個(gè)命令行參數(shù)接收字符串然后原樣傳給amp.initialize。initialize 內(nèi)部做了什么frontend.py 的initialize依次完成檢查torch.backends.cudnn.enabled若被關(guān)閉則拋出RuntimeErrorAmp 依賴 cudnn校驗(yàn)opt_level合法性并從opt_levels字典取回對(duì)應(yīng)默認(rèn)屬性應(yīng)用用戶手動(dòng)覆蓋調(diào)用 _initialize.py 的_initialize執(zhí)行真正的初始化。_initialize內(nèi)部的關(guān)鍵步驟源碼證據(jù)前置檢查check_models拒絕傳入已被torch.nn.parallel.DistributedDataParallel、apex.parallel.DistributedDataParallel或torch.nn.parallel.DataParallel包裝的模型因?yàn)椴⑿邪b必須在amp.initialize返回之后再加如示例代碼注釋所強(qiáng)調(diào)如果先 DDP(model) 再 amp.initialize可能破壞 DDP 的 allreduce hookscheck_params_fp32檢查模型參數(shù)/緩沖區(qū)必須位于 CUDA 且為 FP32若為 Half 或非 CUDA 會(huì)warn_or_err提醒不需要手動(dòng) .half()check_optimizers拒絕傳入已被FP16_Optimizer包裝的優(yōu)化器。模型轉(zhuǎn)換當(dāng)cast_model_type存在時(shí)若keep_batchnorm_fp32True則調(diào)用apex.fp16_utils.convert_network跳過(guò) BatchNorm完成轉(zhuǎn)換否則直接model.to(cast_model_type)隨后用patch_forward包裝每個(gè)模型的forward進(jìn)入的輸入經(jīng)applier遞歸轉(zhuǎn)換為目標(biāo)類型applier支持 Tensor、字符串、ndarray、自定義 batch 類的.to()方法、Mapping 與 Iterable 容器輸出的結(jié)果默認(rèn)轉(zhuǎn)回torch.float32——這正是用戶永遠(yuǎn)不需要調(diào)用 .half()的實(shí)現(xiàn)基礎(chǔ)。優(yōu)化器處理FusedAdam走wrap_fused_adam目前限制為opt_levelO2、keep_batchnorm_fp32False、loss_scale為數(shù)值或dynamic其余優(yōu)化器走 _process_optimizer.py 的_process_optimizer按master_weights分支安裝 lazy-init /_prepare_amp_backward/_post_amp_backward/_master_params_to_model_params等鉤子方法。損失縮放器為每個(gè)num_losses創(chuàng)建LossScaler見 scaler.py。O1 補(bǔ)丁當(dāng)patch_torch_functionsTrue時(shí)調(diào)用amp_init完成 Torch 命名空間的 patch見下文并把optimizer.step包裝為在disable_casts()上下文內(nèi)執(zhí)行——因?yàn)?step 只應(yīng)作用于 FP32 主參數(shù)。amp.scale_loss統(tǒng)一的反向傳播入口amp.scale_loss是一個(gè)上下文管理器context manager用法為with amp.scale_loss(loss, optimizer) as scaled_loss: scaled_loss.backward()其語(yǔ)義依據(jù) handle.py 的實(shí)現(xiàn)與 docstring進(jìn)入上下文時(shí)創(chuàng)建scaled_loss loss.float() * 當(dāng)前損失縮放值yield 給用戶調(diào)用backward()退出上下文時(shí)若delay_unscaleFalse梯度會(huì)被檢查是否含 inf/NaN 并進(jìn)行反縮放unscale隨后即可調(diào)用optimizer.step()。完整簽名amp.scale_loss(loss, optimizers, loss_id0, modelNone, delay_unscaleFalse, delay_overflow_checkFalse)。其中l(wèi)oss_id配合initialize的num_losses實(shí)現(xiàn)每個(gè) loss 獨(dú)立的縮放器delay_unscaleTrue是小眾的忍者級(jí)性能優(yōu)化通常強(qiáng)烈建議保持默認(rèn)False因?yàn)樗鼤?huì)引入多模型/優(yōu)化器/loss 場(chǎng)景下的怪異陷阱。兩個(gè)重要警告來(lái)自 docstring若 Amp 使用顯式 FP32 主參數(shù)O2默認(rèn)如此或手動(dòng)master_weightsTrue任何 FP16 梯度都會(huì)先拷貝為 FP32 主梯度再反縮放optimizer.step()應(yīng)用反縮放后的主梯度到主參數(shù)上。此時(shí)只有 FP32 主梯度會(huì)被反縮放FP16 模型參數(shù)的直接.grad屬性在退出上下文后仍是縮放過(guò)的——這一細(xì)節(jié)影響梯度裁剪裁剪應(yīng)基于主參數(shù)進(jìn)行。梯度溢出時(shí)的處理源碼級(jí)scaler.py 的LossScaler承擔(dān)縮放/反縮放與溢出檢測(cè)。當(dāng)檢測(cè)到溢出inf/NaN且為動(dòng)態(tài)縮放時(shí)update_scale()將損失縮放減半受min_loss_scale地板約束并返回should_skipTrue此時(shí) handle.py 的scale_loss會(huì)把optimizer.step臨時(shí)替換為skip_step打印 Gradient overflow. Skipping step... 并跳過(guò)本次更新同時(shí)清空不會(huì)由zero_grad清零的主梯度。若連續(xù)scale_window默認(rèn) 2000次迭代無(wú)溢出則損失縮放加倍上限max_loss_scale默認(rèn)2**24。此外LossScaler在有 C 擴(kuò)展--cuda_ext --cpp_ext安裝時(shí)使用amp_C.multi_tensor_scale/multi_tensor_axpby融合內(nèi)核批量完成 unscale 與梯度拷貝通過(guò)multi_tensor_applier否則回退到 Python 實(shí)現(xiàn)并打印警告。amp.master_params遍歷主參數(shù)def master_params(optimizer): for group in optimizer.param_groups: for p in group[params]: yield p這是 _amp_state.py 中定義的一個(gè)生成器迭代amp.initialize返回的優(yōu)化器所擁有的參數(shù)。在 O2master weights模式下它產(chǎn)出的是 FP32 主參數(shù)——當(dāng)需要對(duì)這些參數(shù)做梯度裁剪時(shí)見下文高級(jí)用法應(yīng)使用master_params(optimizer)而不是model.parameters()。高級(jí)用法新版統(tǒng)一 Amp API 支持跨迭代的梯度累積、單迭代多次反向傳播、多模型/多優(yōu)化器、自定義用戶定義autograd 函數(shù)以及自定義數(shù)據(jù) batch 類。梯度裁剪與 GAN 也需要特殊處理但這些處理方式不隨opt_level改變。文檔中對(duì)應(yīng)的高級(jí)主題詳見倉(cāng)庫(kù)內(nèi)的 advanced.rst。梯度裁剪基于主參數(shù)由于 O2 模式下 FP16 模型參數(shù)的.grad仍是縮放過(guò)的正確做法是使用amp.master_params(optimizer)拿到 FP32 主參數(shù)再對(duì)主參數(shù)做裁剪# 推薦對(duì)主參數(shù)進(jìn)行梯度裁剪 torch.nn.utils.clip_grad_norm_(amp.master_params(optimizer), max_norm)梯度累積與多次 backwardamp.scale_loss支持跨迭代梯度累積與單次迭代多次反向傳播。若使用delay_unscaleTrue則本次 backward 退出后還不能調(diào)用optimizer.step()必須等后續(xù)某次 backward 以delay_unscaleFalse默認(rèn)結(jié)束后再 step。多模型 / 多優(yōu)化器 / 多 lossamp.initialize接受模型列表與優(yōu)化器列表amp.scale_loss的optimizers參數(shù)可傳單個(gè)優(yōu)化器或優(yōu)化器列表。當(dāng)存在多個(gè) loss 時(shí)通過(guò)num_lossesN與loss_idi組合Amp 為每個(gè) loss 維護(hù)獨(dú)立的損失縮放器model1, model2 ... optim1, optim2 ... [model1, model2], [optim1, optim2] amp.initialize( [model1, model2], [optim1, optim2], opt_levelO1, num_losses2) with amp.scale_loss(loss1, [optim1, optim2], loss_id0) as scaled_loss1: scaled_loss1.backward() with amp.scale_loss(loss2, [optim1, optim2], loss_id1) as scaled_loss2: scaled_loss2.backward() optimizer1.step(); optimizer2.step()即便num_losses保持默認(rèn) 1Amp 也支持多個(gè) loss/backward只是所有 loss 共享同一個(gè)全局損失縮放值。GAN 與自定義 batch 類GAN生成器與判別器可分別作為獨(dú)立模型傳入amp.initialize當(dāng)需要生成器不更新只做判別之類的凍結(jié)場(chǎng)景時(shí)_initialize.py的to_type對(duì)requires_grad的輸入 Tensor 有專門注釋GANs require this說(shuō)明其對(duì)輸入數(shù)據(jù)的處理已考慮 GAN 場(chǎng)景。自定義數(shù)據(jù) batch 類_initialize.py的applier函數(shù)會(huì)檢查對(duì)象是否有.to()方法并調(diào)用它來(lái)轉(zhuǎn)換 batch 內(nèi)的浮點(diǎn) Tensor——因此自定義 batch 類只需實(shí)現(xiàn)to(dtype)方法即可被patch_forward正確轉(zhuǎn)換。從舊 API 遷移Transition GuideAmp 強(qiáng)烈建議遷移到新 API因?yàn)樗ㄓ?、更易用、面向未?lái)原有的FP16_Optimizer類與舊版 Amp API 已被棄用deprecated隨時(shí)可能被移除。舊 Amp API 用戶在新 API 中opt_levelO1執(zhí)行與舊 Amp 相同的 Torch 命名空間 patch。區(qū)別在于新 API 支持靜態(tài)或動(dòng)態(tài)損失縮放而舊 API 只支持動(dòng)態(tài)損失縮放。遷移要點(diǎn)刪除舊的amp_handle amp.init()調(diào)用及返回的amp_handle——amp.initialize()已承擔(dān)并超越了amp.init()的職責(zé)原來(lái)通過(guò)amp_handle暴露的函數(shù)現(xiàn)在是amp模塊下的自由函數(shù)反向傳播上下文管理器需相應(yīng)修改# 舊 API with amp_handle.scale_loss(loss, optimizer) as scaled_loss: scaled_loss.backward() # - # 新 API with amp.scale_loss(loss, optimizer) as scaled_loss: scaled_loss.backward()值得注意的是舊 API 中為用戶函數(shù)標(biāo)注特定精度的調(diào)用如amp.half_function、amp.float_function、amp.promote_function以及注冊(cè)形式的amp.register_half_function/register_float_function/register_promote_function實(shí)現(xiàn)見 amp.py在新 API 中仍然被支持。舊 FP16_Optimizer 用戶opt_levelO2等價(jià)于FP16_Optimizer配合dynamic_loss_scaleTrue。反向傳播同樣需要改為統(tǒng)一版本optimizer.backward(loss) # - with amp.scale_loss(loss, optimizer) as scaled_loss: scaled_loss.backward()舊FP16_Optimizer的一個(gè)痛點(diǎn)在于用戶必須手動(dòng)將模型轉(zhuǎn)為 half調(diào)用.half()或使用apex.fp16_utils的包裝函數(shù)還必須手動(dòng)對(duì)輸入數(shù)據(jù)調(diào)用.half()。新 API 中這兩者都不再需要——無(wú)論選擇哪個(gè)--opt-level你都應(yīng)該用默認(rèn) FP32 格式構(gòu)建模型和傳遞輸入數(shù)據(jù)Amp 會(huì)在model, optimizer amp.initialize(...)期間根據(jù)opt_level及覆蓋標(biāo)志自動(dòng)完成正確的轉(zhuǎn)換。浮點(diǎn)輸入數(shù)據(jù)可以是 FP32 或 FP16但更簡(jiǎn)單的做法是直接保持 FP32——amp.initialize返回的模型其forward已被 patch會(huì)自動(dòng)將輸入數(shù)據(jù)轉(zhuǎn)換到合適類型。一個(gè)完整的實(shí)踐示例結(jié)合倉(cāng)庫(kù)中的 Imagenet 示例一個(gè)接入 Amp 的典型訓(xùn)練腳本骨架如下省略數(shù)據(jù)加載等細(xì)節(jié)import torch from apex import amp # 1. 以 FP32 構(gòu)建模型與優(yōu)化器 model model.cuda() optimizer torch.optim.SGD(model.parameters(), lrargs.lr, momentumargs.momentum, weight_decayargs.weight_decay) # 2. 初始化 Amp先于任何 DDP 包裝 model, optimizer amp.initialize(model, optimizer, opt_levelargs.opt_level, # 如 O1 keep_batchnorm_fp32args.keep_batchnorm_fp32, loss_scaleargs.loss_scale) # 如 dynamic # 3. 分布式場(chǎng)景包裝 DDP 必須在 amp.initialize 之后 if args.distributed: model DDP(model, delay_allreduceTrue) # 4. 訓(xùn)練循環(huán) for epoch in range(args.start_epoch, args.epochs): for i, (input, target) in enumerate(train_loader): output model(input) loss criterion(output, target) with amp.scale_loss(loss, optimizer) as scaled_loss: scaled_loss.backward() optimizer.step() optimizer.zero_grad()命令行調(diào)用示例對(duì)應(yīng)示例倉(cāng)庫(kù)中由 argparse 接收的參數(shù)python main_amp.py --arch resnet50 --opt-level O1 --loss-scale dynamic python main_amp.py --arch resnet50 --opt-level O2 --keep-batchnorm-fp32 TrueAmp 的安裝方式參見倉(cāng)庫(kù) README推薦帶 CUDA/C 擴(kuò)展安裝pip install -v --no-cache-dir --global-option--cpp_ext --global-option--cuda_ext .這樣可啟用multi_tensor_apply融合內(nèi)核unscale、梯度拷貝、master weight 拷貝均受益純 Python 安裝PyTorch 0.4 時(shí)代的需求下amp仍然可用但會(huì)退化為 Python 回退實(shí)現(xiàn)并可能更慢。小結(jié)apex.amp的設(shè)計(jì)哲學(xué)是讓用戶始終以 FP32 視角編寫訓(xùn)練腳本把精度策略交給 opt_level。本文梳理了文檔 amp.rst 的核心內(nèi)容六個(gè)底層屬性cast_model_type、patch_torch_functions、keep_batchnorm_fp32、master_weights、loss_scale構(gòu)成了精度策略的原子單元O0–O3 四個(gè)級(jí)別提供了開箱即用的默認(rèn)組合amp.initialize/amp.scale_loss/amp.master_params三個(gè)入口覆蓋了訓(xùn)練的全部關(guān)鍵環(huán)節(jié)而舊 API 的遷移路徑清晰且無(wú)損舊式函數(shù)標(biāo)注仍然有效。結(jié)合 frontend.py、_initialize.py、_process_optimizer.py、scaler.py 與 handle.py 等源碼可以看到每個(gè) API 承諾背后都有完整的實(shí)現(xiàn)支撐。初次嘗試時(shí)從O1最安全出發(fā)用O0/O3建立精度與速度基線再探索O2的 master weights 收益是一個(gè)穩(wěn)妥的實(shí)踐路徑。贊分享人工智能大模型音樂(lè)生成音頻預(yù)訓(xùn)練【免費(fèi)下載鏈接】jukeboxCode for the paper Jukebox: A Generative Model for Music項(xiàng)目地址https://gitcode.com/gh_mirrors/ju/jukebox點(diǎn)擊查看免費(fèi)下載相關(guān)推薦Ignite混合精度訓(xùn)練實(shí)戰(zhàn)AMP與Apex的完整對(duì)比Ignite混合精度訓(xùn)練實(shí)戰(zhàn)AMP與Apex的完整對(duì)比 PyTorch Ignite是一個(gè)高級(jí)庫(kù)用于幫助在PyTorch中靈活透明地訓(xùn)練和評(píng)估神經(jīng)網(wǎng)絡(luò)。在深深度學(xué)習(xí)模型訓(xùn)練使用 Apex AMP 混合精度訓(xùn)練 DCGANexamples/dcgan/main_amp.py 實(shí)戰(zhàn)指南使用 Apex AMP 混合精度訓(xùn)練 DCGAN examples/dcgan/main_amp.py 實(shí)戰(zhàn)指南 本指南圍繞 Apex 倉(cāng)庫(kù)中 example人工智能深度學(xué)習(xí)分布式訓(xùn)練模型優(yōu)化Colossal-AI 混合精度訓(xùn)練完全指南基于 Booster 的 AMP 配置與實(shí)戰(zhàn)Torch/Apex/Naive AMPColossal AI 混合精度訓(xùn)練完全指南基于 Booster 的 AMP 配置與實(shí)戰(zhàn)Torch/Apex/Naive AMP 導(dǎo)讀 本文以 Colos人工智能大模型分布式訓(xùn)練深度學(xué)習(xí)模型優(yōu)化高性能計(jì)算上一篇PageMenu菜單項(xiàng)居中布局打造優(yōu)雅視覺平衡的終極指南下一篇嵌入式數(shù)據(jù)庫(kù)sled高級(jí)應(yīng)用合并操作符與觸發(fā)器實(shí)現(xiàn)創(chuàng)作聲明:本文部分內(nèi)容由AI輔助生成(AIGC),僅供參考