为什么PyTorch模型转float16后summary显示的参数大小未变化?
问题原因
这个现象是你使用的summary_工具的统计逻辑缺陷导致的,和模型实际的精度切换、显存占用变化无关,具体原因如下:
- 绝大多数PyTorch模型结构统计工具(包括你用到的
summary_、常见的torchsummary/torchinfo等)的参数量大小换算逻辑是硬编码的:默认按单精度浮点数(float32,每个数值占4字节)计算,公式为参数总个数 * 4字节 / 1024 / 1024,不会动态读取每个参数的实际dtype来调整字节换算系数。因此哪怕你的参数已经切换为float16(每个数值占2字节),工具还是会按原来的逻辑计算,输出的大小自然不会变化。 - Apex O3优化是真实将模型参数的dtype切换为float16的,你可以通过PyTorch原生API验证实际显存占用变化:在模型初始化后、优化前调用
print(torch.cuda.memory_allocated())记录初始值,开启O3优化后再调用一次,就会看到实际显存占用已经降低了约一半,符合半精度的预期效果。
注意:
summary_输出的参数大小仅作为参数总规模的参考值,不能用来判断半精度、量化等精度优化的实际效果,真实显存占用请以PyTorch官方显存统计接口的输出为准。
内容的提问来源于stack exchange,提问作者Yeonsik Choi
相关产品推荐
相关产品推荐

