You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

自定义Conv2d_init的FLOPS计算疑问:ptflops显示为0是否合理?

自定义Conv2d_init的FLOPS计算异常问题

问题描述

基于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。需解决以下问题:

  1. 该0值是否真实有效?
  2. 如何正确计算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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.10 01:01:10