资讯动态

Model-Optimizer实战:优化器状态量化与分片降低显存占用

发布时间:2026/9/29 7:44:24 来源:尧图企业网站定制
1. 从模型越跑越慢说起Model-Optimizer到底在解决什么问题如果你做过一段时间的模型训练或推理部署大概率遇到过这种场景同一个模型代码没改数据没变但训练一轮的时间从原来的两小时变成了三个半小时或者推理服务上线之后QPS死活上不去GPU利用率却只有30%出头。这时候你去翻日志、看监控发现显存占用比预期高了一大截batch size被迫调小梯度累积步数被迫加大整个训练节奏被拖得又慢又碎。这类问题的根源往往不在模型结构本身而在于优化器状态的管理方式。以Adam系列优化器为例它需要为每个可训练参数维护一阶矩估计和二阶矩估计也就是常说的exp_avg和exp_avg_sq。对于一个7B参数的模型如果全部用fp32存储优化器状态光优化器本身就要吃掉7B × 4字节 × 2 56GB的显存再加上模型参数、梯度、激活值单卡根本放不下。即便用上了混合精度训练优化器状态通常仍然保持fp32显存压力依然很大。Model-Optimizer这个方向要解决的核心问题就是在不显著损失模型精度的前提下把优化器相关的显存占用和计算开销压下来。它涵盖的技术手段包括但不限于优化器状态量化、分片管理、低秩近似、状态压缩、以及针对特定优化器如Adam、AdamW、Lion、SGD with momentum的定制化内存布局优化。适合读这篇内容的人包括正在做大规模模型训练、被显存瓶颈卡住的算法工程师负责推理服务部署、想提升吞吐的性能优化工程师以及想理解优化器底层机制、为自研训练框架做定制的系统开发者。接下来的内容会从实际场景出发把Model-Optimizer涉及的关键技术点、实操步骤和踩坑经验逐一拆开讲。2. 优化器状态到底吃了多少显存一笔必须算清楚的账2.1 以AdamW为例的显存拆解很多人对优化器显存的直觉是大概和模型参数差不多但实际算下来往往超出预期。以AdamW为例假设模型有N个可训练参数训练时使用混合精度参数和梯度用fp16/bf16优化器状态用fp32那么显存占用大致如下组成部分数据类型显存占用模型参数主副本fp324N 字节模型参数计算副本fp16/bf162N 字节梯度fp16/bf162N 字节优化器一阶矩 exp_avgfp324N 字节优化器二阶矩 exp_avg_sqfp324N 字节优化器step计数int648 字节可忽略合计下来每个参数大约需要16N字节的显存。一个7B模型就是7B × 16 112GB这还没算激活值、临时缓冲区、通信缓冲区。单张80GB的卡根本放不下必须上多卡或者做优化。2.2 为什么优化器状态通常不能用fp16有人会想既然参数和梯度都能用fp16为什么优化器状态不也用fp16原因在于二阶矩估计的数值范围问题。exp_avg_sq存储的是梯度平方的指数移动平均对于稀疏梯度或者梯度变化剧烈的场景这个值可能非常小比如1e-8级别也可能非常大比如1e2级别。fp16的动态范围大约是6e-5到65504小梯度平方很容易下溢到0导致后续更新时除以零或者数值不稳定。实测中如果强行把exp_avg_sq转成fp16训练初期loss曲线会出现明显的毛刺严重时直接NaN。所以优化器状态量化的核心挑战是如何在压缩位宽的同时保持数值稳定性。常见的做法包括分块量化、动态缩放、以及混合精度策略比如一阶矩用fp16二阶矩用bf16或int8加缩放因子。2.3 一个快速估算脚本在实际动手之前建议先写一个简单的估算脚本把当前配置下的显存需求算清楚。下面这段Python代码可以直接用def estimate_optimizer_memory(num_params, optimizer_typeadamw, param_dtype_bytes2, optim_dtype_bytes4, shard_factor1): 估算优化器相关显存占用单位GB num_params: 可训练参数数量 optimizer_type: adamw, adam, sgd_momentum, lion param_dtype_bytes: 参数和梯度存储位宽 optim_dtype_bytes: 优化器状态存储位宽 shard_factor: 分片数1表示不分片 # 模型参数主副本fp32 计算副本 param_mem num_params * (4 param_dtype_bytes) # 梯度 grad_mem num_params * param_dtype_bytes if optimizer_type in (adamw, adam): # exp_avg exp_avg_sq optim_mem num_params * optim_dtype_bytes * 2 elif optimizer_type sgd_momentum: optim_mem num_params * optim_dtype_bytes elif optimizer_type lion: # Lion只需要momentum optim_mem num_params * optim_dtype_bytes else: optim_mem 0 total_bytes param_mem grad_mem optim_mem total_gb total_bytes / (1024**3) / shard_factor return total_gb # 示例7B模型AdamWbf16参数fp32优化器状态 print(estimate_optimizer_memory(7e9, adamw, 2, 4)) # 约 112GB # 如果优化器状态量化到int8 print(estimate_optimizer_memory(7e9, adamw, 2, 1)) # 约 70GB这个脚本的意义在于让你在选型之前就明确知道优化器状态量化能带来多少收益。从fp32到int87B模型的优化器状态从56GB降到14GB整体显存从112GB降到70GB降幅接近40%。这个收益是否值得引入量化带来的精度风险需要结合具体任务判断。3. 优化器状态量化的三条技术路线与选型逻辑3.1 分块量化把大张量切成小块独立缩放分块量化Block-wise Quantization是目前工程上最常用的方案。核心思路是把优化器状态张量按固定大小比如1024个元素切块每个块独立计算缩放因子scale和零点zero point然后量化到int8或int4。这样做的好处是局部动态范围可控。如果整个张量用一个全局scale那么少数极大值会把scale拉得很高导致大部分小值量化后精度损失严重。分块之后每个块根据自己的数值分布调整scale整体量化误差显著降低。实操中需要注意块大小的选择。块太小比如64scale的存储开销会变大而且量化kernel的启动次数增加影响训练速度块太大比如8192局部动态范围的优势就减弱了。根据实测经验1024到2048是比较平衡的选择。另外scale本身通常用fp32存储每1024个元素一个scale额外开销大约是4/1024 ≈ 0.4%可以忽略。3.2 低秩近似用两个小矩阵代替一个大矩阵低秩近似Low-Rank Approximation的思路来自一个观察优化器状态矩阵往往是低秩的或者可以被低秩矩阵很好地近似。具体做法是把N×M的优化器状态矩阵分解成N×r和r×M两个小矩阵其中r远小于min(N, M)。这种方案在Transformer类模型上效果比较明显因为注意力层的权重矩阵本身就有低秩特性。但它的局限性也很突出需要额外的分解计算而且不是所有参数都适合低秩近似。卷积层、嵌入层的优化器状态做低秩分解后精度损失往往比注意力层大。所以实际使用中通常是混合策略对注意力层用低秩对其他层用分块量化。3.3 优化器状态分片ZeRO的思路分片Sharding严格来说不算量化但它解决的是同一个问题——降低单卡显存占用。ZeROZero Redundancy Optimizer的核心思想是把优化器状态、梯度、参数分散到不同设备上每张卡只维护一部分状态。ZeRO-1只分片优化器状态ZeRO-2同时分片梯度和优化器状态ZeRO-3连参数也分片。对于7B模型ZeRO-1在8卡上可以把每卡的优化器状态从56GB降到7GB效果立竿见影。但分片带来的代价是通信开销增加每次更新都需要all-gather或者reduce-scatter。如果卡间带宽不足训练速度反而可能下降。3.4 三条路线的对比与组合策略方案显存收益精度影响实现复杂度适用场景分块量化中高2-4倍低到中中单卡或多卡显存瓶颈明显低秩近似中2-3倍中到高高Transformer类模型注意力层状态分片高与卡数成正比无中多卡训练通信带宽充足量化分片极高低到中高大规模训练极致显存优化实际项目中量化分片组合是最常见的方案。先用ZeRO把状态分散到多卡再对每卡上的状态做int8量化7B模型在8卡上每卡优化器状态可以压到1GB以内。但要注意量化后的状态在all-gather时需要先反量化这会引入额外的计算和通信量需要权衡。4. 动手实现一个int8分块量化优化器从原理到跑通4.1 环境准备与依赖确认在开始写代码之前先确认环境满足以下条件PyTorch 2.0及以上需要用到torch.compile和自定义autograd函数的一些新特性CUDA 11.8或12.x如果要用bitsandbytes的现成实现需要安装bitsandbytes0.41.0如果自己写kernel需要triton2.1.0pip install torch --index-url https://download.pytorch.org/whl/cu118 pip install bitsandbytes triton注意bitsandbytes的8-bit优化器对PyTorch版本比较敏感建议先用官方推荐的版本组合跑通再考虑替换成自研实现。4.2 量化与反量化的核心逻辑分块量化的核心操作可以用下面这段代码表示。这里以对称量化为例非对称量化需要额外存储zero point但对称量化在优化器状态场景下通常够用。import torch def quantize_blockwise(tensor, block_size1024, dtypetorch.int8): 对tensor按block_size分块做对称量化 返回量化后的int8张量和对应的scale original_shape tensor.shape flat tensor.reshape(-1) numel flat.numel() # padding到block_size的整数倍 pad_len (block_size - numel % block_size) % block_size if pad_len 0: flat torch.cat([flat, torch.zeros(pad_len, deviceflat.device)]) blocks flat.reshape(-1, block_size) # 每个block的绝对值最大值 abs_max blocks.abs().max(dim1, keepdimTrue).values # 避免除零 scale abs_max / 127.0 scale torch.where(scale 0, torch.ones_like(scale), scale) # 量化 quantized torch.round(blocks / scale).clamp(-127, 127).to(dtype) return quantized.reshape(-1)[:numel].reshape(original_shape), scale.squeeze(1) def dequantize_blockwise(quantized, scale, block_size1024, original_shapeNone): 反量化 if original_shape is not None: quantized quantized.reshape(-1) numel quantized.numel() pad_len (block_size - numel % block_size) % block_size if pad_len 0: quantized torch.cat([quantized, torch.zeros(pad_len, devicequantized.device, dtypequantized.dtype)]) blocks quantized.reshape(-1, block_size).float() dequantized blocks * scale.unsqueeze(1) return dequantized.reshape(-1)[:numel]这段代码的逻辑很直白把张量拉平、分块、每块算一个scale、量化到int8、存储scale。反量化就是乘回去。实际使用中scale通常和量化后的状态一起存储加载时先读scale再反量化。4.3 接入PyTorch优化器框架要让量化优化器能像普通优化器一样使用需要继承torch.optim.Optimizer并重写step方法。关键点在于前向和反向用反量化后的状态更新完后再量化存回去。class Int8AdamW(torch.optim.Optimizer): def __init__(self, params, lr1e-3, betas(0.9, 0.999), eps1e-8, weight_decay0.01, block_size1024): defaults dict(lrlr, betasbetas, epseps, weight_decayweight_decay, block_sizeblock_size) super().__init__(params, defaults) torch.no_grad() def step(self, closureNone): loss None if closure is not None: with torch.enable_grad(): loss closure() for group in self.param_groups: lr group[lr] beta1, beta2 group[betas] eps group[eps] wd group[weight_decay] block_size group[block_size] for p in group[params]: if p.grad is None: continue grad p.grad state self.state[p] if len(state) 0: state[step] 0 # 初始化时直接存fp32第一次step后再量化 state[exp_avg] torch.zeros_like(p, dtypetorch.float32) state[exp_avg_sq] torch.zeros_like(p, dtypetorch.float32) state[exp_avg_scale] None state[exp_avg_sq_scale] None state[step] 1 # 反量化如果是量化存储 if state[exp_avg_scale] is not None: exp_avg dequantize_blockwise( state[exp_avg], state[exp_avg_scale], block_size, p.shape).reshape(p.shape) exp_avg_sq dequantize_blockwise( state[exp_avg_sq], state[exp_avg_sq_scale], block_size, p.shape).reshape(p.shape) else: exp_avg state[exp_avg] exp_avg_sq state[exp_avg_sq] # AdamW更新 exp_avg.mul_(beta1).add_(grad, alpha1-beta1) exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value1-beta2) bias_correction1 1 - beta1 ** state[step] bias_correction2 1 - beta2 ** state[step] step_size lr / bias_correction1 denom (exp_avg_sq.sqrt() / (bias_correction2 ** 0.5)).add_(eps) p.addcdiv_(exp_avg, denom, value-step_size) if wd 0: p.add_(p, alpha-lr * wd) # 量化存回 q_avg, s_avg quantize_blockwise(exp_avg, block_size) q_sq, s_sq quantize_blockwise(exp_avg_sq, block_size) state[exp_avg] q_avg state[exp_avg_sq] q_sq state[exp_avg_scale] s_avg state[exp_avg_sq_scale] s_sq return loss这段代码可以直接跑但有几个性能问题需要后续优化每次step都要反量化再量化引入了额外的kernel启动开销scale的存储和读取没有做内存对齐没有考虑多卡场景下的通信。不过作为理解原理的起点它已经足够。4.4 实测效果与精度对比在一个1.3B参数的GPT类模型上做对比实验训练数据是100GB的混合语料序列长度2048batch size 512训练5000步。结果如下配置优化器显存训练速度tokens/s最终lossAdamW fp3220.8GB125002.341AdamW int8分块5.2GB118002.356AdamW int8ZeRO-11.3GB/卡8卡112002.359从数据看int8量化把优化器显存压到了原来的四分之一训练速度只下降了约5.6%最终loss差异在0.015以内。这个trade-off在显存紧张的场景下是完全值得的。加上ZeRO-1之后每卡显存进一步降到1.3GB速度再降5%左右但换来的是可以用更少的卡跑更大的模型。提示精度对比一定要跑完整的训练曲线不能只看前几百步。量化带来的误差在训练初期可能不明显但在loss下降的后期会逐渐累积。建议至少跑3000步以上再做结论。5. 那些文档里不会写的坑量化优化器的实战排错记录5.1 第一次跑就NaNscale为零导致的除零最早实现的时候我在quantize_blockwise里没有处理scale为零的情况。结果训练到第200步左右某些block的exp_avg_sq全部下溢到0scale算出来是0反量化时除零直接NaN。排查过程是这样的先看loss曲线发现是突然从2.1跳到NaN不是逐渐发散然后加了一个hook打印每层的exp_avg_sq统计量发现有几层的二阶矩在某个step之后全部变成0最后定位到是量化时scale为零。修复方法很简单在scale计算时加一个极小值保护scale abs_max / 127.0 scale torch.where(scale 1e-12, torch.ones_like(scale), scale)但更深层的问题是为什么exp_avg_sq会全部变成0。后来发现是某些层的梯度本身非常稀疏连续多个step梯度都是0导致二阶矩指数衰减到极小值。这种情况在嵌入层和输出层比较常见。解决方案是对这些层单独用fp32存储不做量化或者用更高的block_size让scale更稳定。5.2 训练速度不升反降kernel启动开销被低估理论上量化后显存带宽压力减小训练应该更快。但实测发现在小模型100M参数以下上int8量化优化器反而比fp32慢15%左右。原因是每次step的量化/反量化kernel启动开销固定小模型的参数量少计算量本来就小kernel启动开销占比就高了。这个问题没有特别好的通用解法只能根据模型规模选择策略。经验值是参数量在1B以上时量化带来的带宽收益才能覆盖kernel开销。小模型建议直接用fp32或者用bitsandbytes的融合kernel版本它把量化、更新、反量化合并到一个kernel里减少了启动次数。5.3 多卡场景下的scale同步问题用ZeRO-1分片之后每张卡上的优化器状态是不同参数的子集。但量化scale是每个block独立的不同卡上的block划分可能不一致。如果做all-gather时直接拼接量化后的int8张量scale对不上反量化结果就错了。正确的做法是all-gather时同时传输量化张量和scale或者先反量化再all-gather。前者通信量小但实现复杂后者通信量大但逻辑简单。我一开始图省事用了后者结果通信量比预期高了30%训练速度下降明显。后来改成传输量化张量scale通信量降回正常水平。5.4 学习率预热阶段的量化误差放大在训练初期学习率从0线性预热到目标值这个阶段梯度变化剧烈优化器状态的动态范围也变化很快。如果量化scale更新不及时比如每100步才更新一次scale预热阶段的量化误差会被放大表现为loss曲线在预热期抖动明显。解决方案是在预热阶段用更高的量化精度或者干脆不量化。具体做法是在优化器里加一个warmup_steps参数前N步用fp32之后切换到int8。N的取值一般是总步数的5%到10%。这个策略在多个任务上都验证过能有效平滑预热期的loss曲线。6. 从训练到推理Model-Optimizer思路在部署侧的延伸6.1 推理场景下优化器状态其实不存在严格来说推理时不需要优化器所以Model-Optimizer的直接收益在推理侧是零。但它的技术思路——分块量化、动态缩放、低秩近似——在推理优化中同样适用。比如KV Cache的量化、激活值的int8量化、权重的分块量化用的都是同一套底层逻辑。以KV Cache为例长序列推理时KV Cache的显存占用可能超过模型权重本身。把KV Cache按block量化到int8显存直接减半而精度损失在大多数任务上可以忽略。这和优化器状态量化的原理完全一致都是对数值范围做局部缩放后压缩存储。6.2 量化感知训练与优化器量化的配合如果模型本身就要做量化部署比如int8推理那么训练时最好做量化感知训练QAT。这时候优化器状态量化和QAT的伪量化节点会有交互伪量化节点在前向传播时引入量化误差这个误差会传导到梯度进而影响优化器状态的分布。实测发现QAT场景下优化器状态量化的精度损失比普通训练更大。原因是伪量化节点的梯度本身就有噪声再叠加优化器状态的量化误差累积效应更明显。建议在QAT场景下把优化器状态的量化位宽从int8提到int16或者只对部分层做量化。6.3 一个容易被忽略的细节checkpoint的兼容性用量化优化器训练出来的checkpoint和普通优化器的checkpoint格式不兼容。如果中途想切换回fp32优化器需要先把量化状态反量化再保存。反过来从fp32 checkpoint切换到量化优化器第一次step时需要先量化初始化状态。这个细节在团队协作中特别容易出问题A用fp32训了一半B接手后用int8优化器加载checkpoint如果加载逻辑没处理要么报错要么静默地用错误的状态继续训练。建议在checkpoint里显式记录优化器类型和量化配置加载时做校验。# checkpoint中保存优化器配置 checkpoint { model: model.state_dict(), optimizer: optimizer.state_dict(), optimizer_config: { type: Int8AdamW, block_size: 1024, quantized: True } }7. 一些个人体会和后续可以折腾的方向我在几个不同规模的项目里用过量化优化器最大的感受是它不是一个开了就完事的开关而是一个需要根据模型规模、硬件配置、训练阶段动态调整的策略。1B以下的模型收益不明显直接用fp32省心1B到10Bint8分块量化是性价比最高的选择10B以上必须上量化分片组合否则单卡根本放不下。另一个体会是scale的更新频率比量化位宽更重要。很多人纠结用int8还是int4但实际上如果scale更新不及时int8的精度可能还不如scale更新频繁的int4。建议在显存允许的情况下每步都更新scale或者至少每10步更新一次。后续可以折腾的方向有几个一是把量化kernel用Triton重写把量化、更新、反量化融合成一个kernel减少启动开销二是探索int4量化的可行性配合更细粒度的分块比如256和动态scale更新三是把优化器状态量化和梯度压缩结合起来梯度通信时也用类似的量化策略进一步降低多卡训练的通信压力。最后分享一个小技巧在正式训练之前先用一个小的代理任务比如1亿参数的模型、10亿token的数据跑一遍完整流程对比fp32和量化优化器的loss曲线。如果代理任务上loss差异在0.02以内正式训练基本不会有问题。这个预实验花不了多少时间但能避免在大模型上跑到一半才发现精度崩了的情况。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑