如何让Pytorch summary正确统计多次调用同一模块时的可训练参数量
问题原因
torchsummary的原生统计逻辑是跟踪前向传播过程中每一次层调用动作,每次触发层计算都会累加对应层的参数量,不会判断该层的参数是否已经被统计过,因此同一个模块在forward中被多次调用时就会出现参数重复计数的问题。
可行解决方案
方案1:替换为torchinfo库(最简便,兼容原有调用逻辑)
torchinfo是torchsummary的社区优化分支,默认实现了参数去重统计,同时保留了torchsummary的所有常用功能(分层输出、输入输出尺寸统计、显存占用计算等),调用方式几乎完全一致,完全符合你沿用原有功能的需求。
如果未安装可以先执行pip install torchinfo完成安装,之后只需要修改导入语句即可:
# 替换原有的 from torchsummary import summary from torchinfo import summary
用你提供的示例代码运行,输出的总参数量会直接显示为12,和真实可训练参数量一致。
方案2:给原生torchsummary加去重补丁(必须用原生库时可选)
如果必须使用原版torchsummary,可以修改其统计逻辑,新增参数去重判断:
- 找到torchsummary安装目录下的
torchsummary.py文件 - 找到参数统计的循环逻辑,新增一个集合存储已统计过的参数id,修改累加规则:
# 原统计逻辑修改后参考: summary_dict = {} counted_params = set() # 新增:存储已经统计过的参数的内存id total_params = 0 trainable_params = 0 # ... 遍历每层的原有代码 ... for param in layer.parameters(): param_id = id(param) if param_id not in counted_params: counted_params.add(param_id) param_size = np.prod(param.size()) total_params += param_size if param.requires_grad: trainable_params += param_size
方案3:原生summary输出后追加正确参数量
不需要修改库代码,在调用完summary之后直接运行你已经掌握的参数量统计代码,追加打印即可,兼顾torchsummary的分层输出和正确的总参数量统计。
内容的提问来源于stack exchange,提问作者Dadeslam
相关产品推荐
相关产品推荐

