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

如何查找PyTorch模型中中间层对应的输入层名称

实现方案:通过PyTorch FX符号图获取层输入连接关系

PyTorch为动态图框架,默认执行过程中不会保留完整的层间连接静态信息,我们可以通过官方提供的torch.fx工具将模型转换为符号计算图,通过遍历图节点的方式快速获取Concat层的所有输入层信息,无需修改原有模型结构。

完整可运行代码

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

class Concat(nn.Module):
    def __init__(self, dimension=1):
        super().__init__()
        self.d = dimension

    def forward(self, x):
        return torch.cat(x, self.d)


class SomeModel(nn.Module):
    def __init__(self):
        super(SomeModel, self).__init__()
        self.conv1 = nn.Conv2d(3, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
        self.bn1 = nn.BatchNorm2d(64)
        self.conv2 = nn.Conv2d(3, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)
        self.conc = Concat(1)
        self.linear = nn.Linear(8192, 1)

    def forward(self, x):
        out1 = F.relu(self.bn1(self.conv1(x)))
        out2 = F.relu(self.conv2(x))
        out = self.conc([out1, out2])
        out = F.avg_pool2d(out, 4)
        out = out.view(out.size(0), -1)
        out = self.linear(out)
        return out


if __name__ == '__main__':
    model = SomeModel()
    # 符号追踪得到计算图
    traced = symbolic_trace(model)
    # 先构建所有层名到模块的映射
    name_to_module = dict(traced.named_modules())

    # 遍历计算图所有节点
    for node in traced.graph.nodes:
        # 匹配Concat层对应的节点
        if node.op == 'call_module' and isinstance(name_to_module[node.target], Concat):
            print(f"找到Concat层,层名:{node.target}")
            print("该层的所有输入层/操作如下:")
            # 遍历该节点的所有输入参数
            for idx, arg in enumerate(node.args[0]):
                # 如果输入是模块调用,直接取模块名
                if arg.op == 'call_module':
                    print(f"- 输入{idx+1}对应层名:{arg.target},对应层类型:{name_to_module[arg.target].__class__.__name__}")
                # 如果输入是函数调用(比如示例里的relu),可以打印函数名
                elif arg.op == 'call_function':
                    print(f"- 输入{idx+1}对应操作:{arg.target.__name__},其输入层为:")
                    # 还可以继续递归找函数调用的输入层,这里示例找relu的输入
                    for sub_arg in arg.args:
                        if sub_arg.op == 'call_module':
                            print(f"  - {sub_arg.target},类型:{name_to_module[sub_arg.target].__class__.__name__}")

代码运行输出示例

找到Concat层,层名:conc
该层的所有输入层/操作如下:
- 输入1对应操作:relu,其输入层为:
  - bn1,类型:BatchNorm2d
- 输入2对应操作:relu,其输入层为:
  - conv2,类型:Conv2d

可选方案:通过前向钩子记录输入来源

如果不想使用FX工具,也可以在Concat层的__init__中注册前向钩子,在运行前向的时候记录输入张量的grad_fn来反向推导来源,但这种方式需要实际跑一次前向,且推导逻辑更复杂。如果你的模型包含复杂动态控制流导致FX无法正常追踪,可选用该方案。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.03 00:15:02