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

深度嵌套PyTorch模型的每层FLOPs及乘加运算量计算方法

深度嵌套PyTorch模型的每层FLOPs与乘加次数计算方案

常用库(如flopth、ptflops、pytorch-OpCounter)失效的核心原因是深度嵌套的子模块(比如多层Sequential嵌套、自定义Module嵌套)无法被库的遍历逻辑完全覆盖,以下是三种可行的解决方法:

方法一:手动递归遍历所有子模块并统计

利用PyTorch自带的named_modules()方法可以递归遍历模型的所有层级子模块,无需担心嵌套深度。针对不同类型的层,手动套用对应公式计算FLOPs和乘加次数(MACs)。

示例代码

import torch
import torch.nn as nn

def calculate_layer_stats(model, input_size):
    # 生成输入张量
    dummy_input = torch.randn(*input_size)
    flops_record = {}
    macs_record = {}

    # 注册钩子获取模块输入形状
    def register_shape_hook(module):
        input_shape = None
        def hook(_, inputs, __):
            nonlocal input_shape
            input_shape = inputs[0].shape
        handle = module.register_forward_hook(hook)
        # 前向传播一次获取输入形状
        with torch.no_grad():
            model(dummy_input)
        handle.remove()
        return input_shape

    # 遍历所有子模块
    for module_name, module in model.named_modules():
        if module_name == "":  # 跳过模型本身
            continue
        
        input_shape = register_shape_hook(module)
        if not input_shape:
            flops_record[module_name] = 0
            macs_record[module_name] = 0
            continue

        # 卷积层计算逻辑
        if isinstance(module, nn.Conv2d):
            batch, in_ch, in_h, in_w = input_shape
            out_ch, _, k_h, k_w = module.weight.shape
            groups = module.groups
            out_h = (in_h + 2*module.padding[0] - k_h) // module.stride[0] + 1
            out_w = (in_w + 2*module.padding[1] - k_w) // module.stride[1] + 1

            macs = out_ch * (in_ch // groups) * k_h * k_w * out_h * out_w
            flops = 2 * macs  # 一次乘加算2个FLOP
            if module.bias is not None:
                flops += out_ch * out_h * out_w  # 偏置加法操作
        # 全连接层计算逻辑
        elif isinstance(module, nn.Linear):
            batch, in_feat = input_shape
            out_feat = module.out_features

            macs = in_feat * out_feat
            flops = 2 * macs
            if module.bias is not None:
                flops += out_feat
        # BN层计算逻辑(减均值、除标准差、乘gamma、加beta)
        elif isinstance(module, nn.BatchNorm2d):
            batch, ch, h, w = input_shape
            flops = 4 * ch * h * w
            macs = 2 * ch * h * w  # 乘gamma+加beta算一次乘加
        # 其他无乘加操作的层(如ReLU、MaxPool)直接记0
        else:
            flops = 0
            macs = 0

        flops_record[module_name] = flops
        macs_record[module_name] = macs

    return flops_record, macs_record

# 测试用深度嵌套模型
class DeepNestedModel(nn.Module):
    def __init__(self):
        super().__init__()
        self.block1 = nn.Sequential(
            nn.Conv2d(3, 16, 3),
            nn.BatchNorm2d(16),
            nn.Sequential(
                nn.Conv2d(16, 32, 3),
                nn.ReLU()
            )
        )
        self.fc = nn.Linear(32*26*26, 10)

if __name__ == "__main__":
    model = DeepNestedModel()
    layer_flops, layer_macs = calculate_layer_stats(model, (1, 3, 32, 32))
    
    print("每层FLOPs统计:")
    for name, val in layer_flops.items():
        print(f"{name}: {val:,}")
    
    print("\n每层乘加次数统计:")
    for name, val in layer_macs.items():
        print(f"{name}: {val:,}")

方法二:修改现有库的遍历逻辑

部分库无法处理深度嵌套是因为只遍历直接子模块(children()),而非递归遍历所有子模块。以pytorch-OpCounter为例,可修改其源码中的模块遍历逻辑:

  1. 找到库中负责遍历模块的函数(如_add_hooks)
  2. 将原有的for module in self.model.children()替换为for _, module in self.model.named_modules()
  3. 调整钩子注册的逻辑,确保每个子模块都能被正确统计

方法三:使用PyTorch FX进行静态图追踪

PyTorch FX可以将动态模型转换为静态计算图,自动展开所有嵌套模块,之后遍历图节点即可统计每层的运算量。

示例代码

import torch
import torch.nn as nn
from torch.fx import symbolic_trace

def count_stats_with_fx(model, input_size):
    traced_model = symbolic_trace(model)
    dummy_input = torch.randn(*input_size)
    flops_record = {}
    macs_record = {}

    def calc_node_stats(node):
        if node.op != "call_module":
            return 0, 0
        
        module = traced_model.get_submodule(node.target)
        input_shape = node.args[0].shape

        if isinstance(module, nn.Conv2d):
            in_ch = input_shape[1]
            out_ch = module.out_channels
            k_h, k_w = module.kernel_size
            out_h = (input_shape[2] + 2*module.padding[0] - k_h) // module.stride[0] + 1
            out_w = (input_shape[3] + 2*module.padding[1] - k_w) // module.stride[1] + 1
            groups = module.groups

            macs = out_ch * (in_ch // groups) * k_h * k_w * out_h * out_w
            flops = 2 * macs
            flops += out_ch * out_h * out_w if module.bias is not None else 0
            return flops, macs
        elif isinstance(module, nn.Linear):
            in_feat = input_shape[1]
            out_feat = module.out_features
            macs = in_feat * out_feat
            flops = 2 * macs + (out_feat if module.bias is not None else 0)
            return flops, macs
        return 0, 0

    for node in traced_model.graph.nodes:
        flops, macs = calc_node_stats(node)
        if flops > 0:
            flops_record[node.target] = flops
            macs_record[node.target] = macs

    return flops_record, macs_record

# 测试使用
if __name__ == "__main__":
    model = DeepNestedModel()
    layer_flops, layer_macs = count_stats_with_fx(model, (1, 3, 32, 32))
    
    print("FX追踪的每层FLOPs:")
    for name, val in layer_flops.items():
        print(f"{name}: {val:,}")

注意事项

  • 若模型包含动态分支(如if/else、循环),需确保FX能正确追踪,或改用TorchScript进行静态转换。
  • FLOPs与MACs的定义需统一:通常一次乘加(a*b + c)算2个FLOP,MACs则统计乘加操作的次数。
  • 自定义模块需手动补充对应的运算量计算逻辑,避免遗漏。

内容的提问来源于stack exchange,提问作者afsara_ben

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 23:10:43