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

如何用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 03:48:02