Automatic MIXED PRECISION(AMP)参数疑问及fp16 checkpoint保存方法咨询
关于自动混合精度(AMP)的常见问题解答
1. AMP会将模型参数减半吗?
不会。AMP的核心逻辑是保持模型参数以FP32格式存储,仅在计算过程中自动将部分算子的输入、输出临时转为FP16以提升计算速度,同时通过梯度缩放技术避免梯度下溢问题。模型参数本身不会转为FP16,自然不存在参数内存占用减半的情况。
2. 手动调用half()训练效果差,但AMP效果好,如何保存FP16格式的checkpoint?
手动将整个模型和数据转为FP16时,所有算子(包括对数值精度敏感的BatchNorm、Softmax等)都被迫用FP16计算,极易出现精度丢失、梯度下溢,导致训练效果崩溃。而AMP会自动识别适合FP16的算子,保留敏感算子的FP32计算,同时配合梯度缩放,平衡了速度与精度。
保存FP16格式的checkpoint完全可行,实现步骤如下:
- 加载原FP32格式的checkpoint
- 遍历参数,将每个参数转换为FP16格式后保存
示例代码:
import torch # 加载FP32 checkpoint fp32_ckpt = torch.load("your_fp32_checkpoint.pth") # 转换为FP16 fp16_ckpt = {key: val.half() for key, val in fp32_ckpt.items()} # 保存FP16 checkpoint torch.save(fp16_ckpt, "your_fp16_checkpoint.pth")
注意:加载FP16 checkpoint时,需先将模型转为FP16(调用model.half()),或在加载时完成类型转换,避免参数与模型的 dtype 不匹配。
3. autocast模式下执行conv2d时,参数是否会被转换为半精度?
在autocast上下文环境中,conv2d的输入数据会被自动转为FP16,但模型的conv2d参数始终以FP32格式存储。不过计算过程中,PyTorch会临时将FP32的参数转换为FP16参与运算,运算结束后参数仍恢复为FP32存储。这种设计既利用了FP16的计算效率,又避免了参数长期存为FP16带来的精度损耗。
内容的提问来源于stack exchange,提问作者lee Lin
相关产品推荐
相关产品推荐

