如何查找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
相关产品推荐
相关产品推荐

