如何用PyTorch统计模型前向传播的参数量和MACs(含逐元素运算)
问题原因
普通的register_forward_hook只能作用于继承自nn.Module的层实例,你在forward里直接写的+、*等逐元素运算属于原生PyTorch算子,不属于独立的Module实例,所以钩子不会触发,自然统计不到对应的运算量。
可行解决方案
方案1:封装逐元素运算为自定义Module(适配原有钩子逻辑)
如果不想改动你现有的钩子统计逻辑,可以把所有需要统计的逐元素运算封装成独立的自定义Module,之后在forward里调用这些Module实例代替直接写运算符,这样你的现有钩子就能捕获到这些运算的输入输出,进而计算MACs。
示例代码:
import torch.nn as nn import torch # 封装逐元素加法 class ElemAdd(nn.Module): def forward(self, a, b): return a + b # 封装逐元素乘法 class ElemMul(nn.Module): def forward(self, a, b): return a * b # 改造后的模型 class MyModel(nn.Module): def __init__(self): super().__init__() self.l1 = nn.Linear(10,10) self.l2 = nn.Linear(10,10) # 实例化逐元素运算模块 self.add1 = ElemAdd() self.mul1 = ElemMul() self.add2 = ElemAdd() self.add3 = ElemAdd() def forward(self, x): y1 = self.l1(x) y2 = self.l2(x) y3 = self.add1(y1, y2) tmp = self.mul1(y3, y1) tmp2 = self.add2(y2, y2) return self.add3(tmp2, tmp)
之后你可以和之前给Linear层挂钩子一样,给ElemAdd、ElemMul实例也挂钩子,按照你定义的MACs规则统计即可:比如逐元素乘法对应输入张量元素总数个MACs,逐元素加法不单独计数。
方案2:用Torch FX追踪全计算图(无需改模型代码,通用度更高)
如果不想改动原有模型代码,可以用PyTorch自带的FX工具追踪整个模型的前向计算图,遍历所有算子节点(包括逐元素运算的算子节点),直接统计所有运算的MACs,不需要单独挂钩子。
示例代码片段:
from torch.fx import symbolic_trace model = MyModel() input_sample = torch.randn(1, 10) # 符号追踪得到计算图 gm = symbolic_trace(model) # 遍历所有计算节点 for node in gm.graph.nodes: # 过滤出call_function类型的节点,对应原生算子调用 if node.op == 'call_function': # 获取算子类型 op = node.target # 获取输入输出张量的元信息 inputs = [i.meta['tensor_meta'] for i in node.args if 'tensor_meta' in i.meta] output = node.meta['tensor_meta'] # 按你的规则统计MACs if op == torch.add: # 加法按你的规则不计数,可按需调整 pass elif op == torch.mul: # 逐元素乘法的MACs等于输出张量的元素总数 macs = output.numel() print(f"逐元素乘法节点{node.name} MACs: {macs}")
方案3:用PyTorch Profiler统计(无需自己实现统计逻辑)
你也可以直接用PyTorch自带的性能分析工具,直接统计所有算子的运算量,支持直接导出FLOPs数据,按规则转换为MACs即可:
from torch.profiler import profile, ProfilerActivity model = MyModel() input = torch.randn(1,10) with profile(activities=[ProfilerActivity.CPU], with_flops=True) as prof: model(input) # 打印所有算子的FLOPs,按规则1 MAC=2 FLOPs做除法即可得到你需要的MACs print(prof.key_averages().table(sort_by="flops", row_limit=10))
统计规则对齐注意
按照你给出的MACs定义:y1*(y2 + y3)对应1次MACs,统计时注意不要重复计数:
- 逐元素加法本身不单独统计MACs
- 乘加组合运算或者单独的逐元素乘法,按乘法的元素数统计MACs即可
内容的提问来源于stack exchange,提问作者Maxwell Albert
相关产品推荐
相关产品推荐

