自定义Conv2d_init的FLOPS计算疑问:ptflops显示为0是否合理?
问题描述
基于PyTorch的nn.Conv2d自定义了Conv2d_init类,仅重写reset_parameters方法实现自定义权重初始化:
import math import torch import torch.nn as nn import torch.nn.init as init class Conv2d_init(nn.Conv2d): def __init__(self, in_channels, out_channels, kernel_size, stride=1, padding=0, dilation=1, groups=1, bias=True, padding_mode="zeros"): super(Conv2d_init, self).__init__(in_channels, out_channels, kernel_size, stride, padding, dilation, groups, bias, padding_mode) def reset_parameters(self): init.xavier_normal_(self.weight) if self.bias is not None: fan_in, _ = init._calculate_fan_in_and_fan_out(self.weight) bound = 1 / math.sqrt(fan_in) init.uniform_(self.bias, -bound, bound)
随后构建包含该类的SegDecNet模型:
class SegDecNet(nn.Module): def __init__(self, device, input_width, input_height, input_channels): super(SegDecNet, self).__init__() self.vo=Conv2d_init(in_channels=3, out_channels=256,kernel_size=5, padding=2, bias=False) self.vo2=Conv2d_init(in_channels=256, out_channels=512,kernel_size=5, padding=2, bias=False) # self.vo=nn.Conv2d(3,256,kernel_size=5,stride=1,padding=2) # self.vo2=nn.Conv2d(256,512,kernel_size=5,stride=1,padding=2) def forward(self, input): x=self.vo(input) x=self.vo2(x) return x
使用ptflops工具计算模型复杂度时,发现使用Conv2d_init时FLOPS显示为0,替换为原生nn.Conv2d时FLOPS为500 GMac。需解决以下问题:
- 该0值是否真实有效?
- 如何正确计算
Conv2d_init的FLOPS?
解答
1. 0值无效的原因
0值是错误结果,并非真实计算量。因为Conv2d_init完全继承了nn.Conv2d的前向传播逻辑,仅修改了权重初始化方式,其计算量(FLOPS)和原生nn.Conv2d完全一致。ptflops返回0是因为工具未识别到自定义的Conv2d_init模块,没有触发对应的FLOPS计算规则。
ptflops的核心逻辑是通过模块类型匹配来调用对应的计算函数,它默认只支持PyTorch原生的nn.Conv2d等模块,自定义类Conv2d_init不在默认支持列表中,因此被判定为无计算量的模块。
2. 正确计算FLOPS的方法
以下是三种可行的解决方式:
方法一:给自定义类添加模块类型标记
在Conv2d_init类中添加__module__属性,让ptflops将其识别为torch.nn下的原生Conv2d类型,从而复用原生模块的FLOPS计算逻辑:
class Conv2d_init(nn.Conv2d): __module__ = 'torch.nn' # 添加该行,标记为原生nn模块 def __init__(self, in_channels, out_channels, kernel_size, stride=1, padding=0, dilation=1, groups=1, bias=True, padding_mode="zeros"): super(Conv2d_init, self).__init__(in_channels, out_channels, kernel_size, stride, padding, dilation, groups, bias, padding_mode) def reset_parameters(self): init.xavier_normal_(self.weight) if self.bias is not None: fan_in, _ = init._calculate_fan_in_and_fan_out(self.weight) bound = 1 / math.sqrt(fan_in) init.uniform_(self.bias, -bound, bound)
方法二:手动注册FLOPS计算函数
直接将原生nn.Conv2d的FLOPS计算函数注册到自定义类上,让ptflops明确知道如何计算该模块的复杂度:
from ptflops import get_model_complexity_info from ptflops.flops_counter import _DEFAULT_SUPPORTED_OPS # 复制原生Conv2d的计算函数 conv_flops_func = _DEFAULT_SUPPORTED_OPS[nn.Conv2d] # 将自定义类注册到支持列表中 _DEFAULT_SUPPORTED_OPS[Conv2d_init] = conv_flops_func # 初始化模型并计算FLOPS model = SegDecNet(device='cpu', input_width=256, input_height=256, input_channels=3) flops, params = get_model_complexity_info(model, (3, 256, 256), as_strings=True, print_per_layer_stat=True) print(f"FLOPS: {flops}")
方法三:临时替换自定义模块为原生模块
由于Conv2d_init和nn.Conv2d的前向逻辑完全一致,仅初始化不同,可在计算FLOPS前,将模型中的Conv2d_init替换为nn.Conv2d(不影响计算量结果):
def replace_conv2d_init(model): for name, module in model.named_children(): if isinstance(module, Conv2d_init): # 创建对应的原生Conv2d模块 new_conv = nn.Conv2d( in_channels=module.in_channels, out_channels=module.out_channels, kernel_size=module.kernel_size, stride=module.stride, padding=module.padding, dilation=module.dilation, groups=module.groups, bias=module.bias is not None, padding_mode=module.padding_mode ) setattr(model, name, new_conv) else: # 递归处理子模块 replace_conv2d_init(module) # 替换后计算FLOPS model = SegDecNet(device='cpu', input_width=256, input_height=256, input_channels=3) replace_conv2d_init(model) flops, params = get_model_complexity_info(model, (3, 256, 256), as_strings=True, print_per_layer_stat=True) print(f"FLOPS: {flops}")
内容的提问来源于stack exchange,提问作者y ttt

