恒美微站
首页
关于我们
建站服务
主题模板
案例展示
资讯中心
联系我们
初创-词表角度如何减少显存消耗
首页
资讯中心
/
初创-词表角度如何减少显存消耗
初创-词表角度如何减少显存消耗
发布时间:2026/9/28 4:55:42
ChunkedLinearCrossEntropy:分块词表交叉熵(Qwen3.5训练显存削峰方案)概述问题根源:LMHead词表乘法带来巨大显存尖峰大模型CausalLM训练标准流程:主干Transformer输出hidden_states[B, L, H]过lm_head = nn.Linear(H, vocab_size):logits=hidden@WlmTlogits = hidden @ W_{lm}^Tlogits=hidden@WlmTlogits转fp32,计算F.cross_entropyQwen3.5-0.8B:vocab_size=248320。对于 packing 得到的长序列L=16384:bf16 logits张量:16384 × 248320 × 2 bytes ≈ 8GBcross_entropy计算需要fp32,直接翻倍到16GB4080只有16GB显存,仅仅这一步logits就直接OOM/触发GPU重置。更糟的:prompt部分label=-100,本来不需要参与loss,但原生实现依然全部送入lm_head,完整算出整段logits,白白占用显存。核心痛点总结:原生实现:一次性生成完整[L, V]logits,全部保存在计算图,显存尖峰由序列长度 × 词表大小决定;提示token也参与lm_head矩阵乘。✅分块方案两个核心优化:过滤:只取label≠-100的token,prompt直接剔除,不送入lm_head;分块流式计算:有效token拆成小块(本例_CHUNK_TOKENS=2048),每块单独算logits,计算完立刻del logits释放显存;计算图不保存完整logits;前向no_grad + 反向重算logits:前向只累加loss,不保留logits激活;反向阶段,利用保存的hidden、lm_head权重、labels,重新计算当前chunk的logits来求梯度。👉 显存上限:仅单个chunk的logits常驻显存,不再是全局全长logits。性能特点:词表矩阵乘不在模型关键路径上;分块几乎不增加step耗时,只压低显存尖峰。实验数据回顾(A100,stage2,无梯度检查点)Case步内峰值显存step耗时普通,完整logits,非pack短序列21.2 GB0.349 s普通 + 2048分块,非pack短序列18.5 GB0.349 s16k packing长序列,完整logits35.5 GB0.784 s16k packing长序列 +2048分块29.1 GB0.762 s短样本场景:省2.7GB;长pack序列收益巨大,省6.4GB;耗时几乎不变;4080上chunk=256/2048,step时间稳定≈2.6s。补充:这和前面线性注意力、FA varlen空block是独立优化。这个模块只管lm_head词表矩阵乘+交叉熵loss显存,不影响Transformer主干、FlashAttention、线性注意力。原理拆解(autograd custom Function)PyTorch自定义autograd Function,把前向、反向逻辑拆开。Forward(前向)输入:hidden(有效token对应的hidden,已经过滤掉label=-100),lm_head.weight,labels保存hidden, weight, labels到ctx,不保存logits;torch.no_grad()上下文内,循环遍历chunk:每块做F.linear算出当前chunk的logits,转fp32;单块计算cross_entropy,sum累加loss;计算完成立刻del logits释放显存;最终返回平均loss(总loss / 有效token数量)。关键点:no_grad,logits不进入autograd计算图,不会占用计算图的激活显存。Backward(反向)ctx里只有hidden, weight, labels,没有logits。必须重算logits求梯度:遍历同样chunk;每块重新计算logits;手动推导cross-entropy梯度:dLdlogit=softmax(logit)−one_hot(label)\frac{d\mathcal{L}}{dlogit}=softmax(logit)-one\_hot(label)dlogitdL=softmax(logit)−one_hot(label)代码里用scatter_add实现onehot减法,避免实例化超大onehot矩阵(节省显存);算出dlogits,链式求导:dhidden=dlogits@Wlmdhidden = dlogits @ W_{lm}dhidden=dlogits@W