直接调用model.half()会报错,因其强制将所有参数转为torch.float16,但batchnorm2d等算子内部依赖fp32运算,导致输入类型与算子期望不匹配而崩溃;正确做法是使用autocast+gradscaler的amp机制,保持主权重fp32、自动切换前向/反向精度,并确保三步缩放更新完整执行。

能降,但不是套个autocast就完事——漏掉GradScaler或类型错配,显存不降反崩。
为什么直接 model.half() 会报错 RuntimeError: expected scalar type Float but found Half
PyTorch 的混合精度不是简单把模型转成 FP16。手动调 model.half() 会让所有参数变成 torch.float16,但很多算子(比如 nn.BatchNorm2d、nn.LayerNorm、nn.CrossEntropyLoss)内部依赖 FP32 运算,输入一进就炸。AMP 的设计是:模型权重保持 torch.float32(主权重),前向/反向在 autocast 内自动切 FP16,损失计算和优化器更新仍在 FP32 上做。所以:
• 不要手动调 model.half()
• 输入 data 和 target 保持原始 dtype(通常是 float32)
• loss 计算必须在 autocast 外,或显式退出上下文
必须写全的三步:scaler.scale → scaler.step → scaler.update
GradScaler 不是可选项,它是对抗 FP16 梯度下溢的唯一防线。漏掉任意一步都会导致训练失败或精度崩溃:
• scaler.scale(loss).backward():放大梯度,避免小梯度被截断为 0
• scaler.step(optimizer):用缩放后的梯度更新参数;如果返回 None(说明梯度溢出),参数不更新
• scaler.update():动态调整缩放因子(比如连续成功就增大 scale,失败则衰减)
常见错误:
• 把 optimizer.step() 写在 scaler.step() 前面 → RuntimeError: trying to call step() on optimizer with no gradients
• 忘了 scaler.update() → 下次 scale 不会自适应调整,容易 NaN
autocast 的 dtype 别硬设 torch.float16,Ampere 卡优先试 torch.bfloat16
在 A100 / RTX 3090 / RTX 4090 等 Ampere 及更新架构上,torch.bfloat16 比 torch.float16 更稳:
• 指数位宽和 FP32 一致(8 bit),不容易 overflow/underflow
• Tensor Core 原生支持,性能不输 FP16
• 不需要改代码逻辑,只需把 autocast() 改成 autocast(dtype=torch.bfloat16)
但注意:
• V100(Volta)不支持 bfloat16,强制用会报错
• 如果模型含大量 softmax 或 logsumexp 类算子,bfloat16 数值更鲁棒
• 用 torch.cuda.amp.autocast(enabled=False) 临时退出上下文,只对特定模块禁用 AMP(如自定义 loss)
显存省多少?别只看参数,盯住激活值和优化器状态
FP16 看似减半显存,但实际节省取决于瓶颈在哪:
• 模型参数:7B 模型从 ~28GB(FP32)→ ~14GB(FP16)
• 激活值(activation):高分辨率图像或长序列 Transformer 中,这部分常比参数还占显存,FP16 直接砍半
• 优化器状态:Adam 默认存 3 份 FP32(param + grad + momentum + variance),共 12× 参数量字节;AMP 通常保留 FP32 主权重 + FP16 梯度计算,状态显存不降,但激活值节省仍可观
实测提示:
• ResNet-50 训练显存从 9.8GB → 4.2GB(降幅约 57%)
• 若发现开启 AMP 后显存没明显下降,大概率是 batch size 太大或梯度检查点没开,激活值仍是瓶颈
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











