深度嵌套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为例,可修改其源码中的模块遍历逻辑:
- 找到库中负责遍历模块的函数(如
_add_hooks) - 将原有的
for module in self.model.children()替换为for _, module in self.model.named_modules() - 调整钩子注册的逻辑,确保每个子模块都能被正确统计
方法三:使用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
相关产品推荐
相关产品推荐

