如何在PyTorch计算图可视化中为中间操作显示有意义的名称?
Torchviz可视化PyTorch计算图:显示自定义节点名称问题
尝试使用torchviz可视化PyTorch模型计算图时,希望中间节点显示有意义的自定义名称,但当前仅生成AddBackward0、MulBackward0这类通用操作名,难以解读。即使使用自定义nn.Module类标记中间步骤,节点名称仍未体现自定义标识。以下是复现问题的极简代码:
import torch import torch.nn as nn from torchviz import make_dot # 自定义中间计算模块 class CustomAdd(nn.Module): def forward(self, x, y): return x + y class CustomMul(nn.Module): def forward(self, x, y): return x * y class SimpleModel(nn.Module): def __init__(self): super(SimpleModel, self).__init__() self.custom_add = CustomAdd() self.custom_mul = CustomMul() self.linear = nn.Linear(1, 1, bias=False) def forward(self, x): # 中间计算步骤 a = self.custom_add(x, x) # a = x + x b = self.custom_mul(a, x) # b = a * x out = self.linear(b) return out # 实例化模型 model = SimpleModel() # 输入张量 x = torch.tensor([[2.0]], requires_grad=True) # 前向传播 output = model(x) # 计算损失 loss = (output - torch.tensor([[1.0]])) ** 2 # 反向传播 loss.backward() # 可视化计算图 dot = make_dot(loss, params=dict(model.named_parameters())) dot.format = 'png' dot.render('simple_model_graph')
希望对应中间计算a、b的节点能显示CustomAdd、CustomMul这类自定义名称,以下是可行的修改方案:
方案1:给中间张量命名并传入可视化工具
Torchviz可以识别张量的.name属性,或通过扩展params参数显式传入中间变量,从而在图中显示自定义名称:
修改模型的forward方法,为中间张量添加名称,并在可视化时传入这些变量:
class SimpleModel(nn.Module): def __init__(self): super(SimpleModel, self).__init__() self.custom_add = CustomAdd() self.custom_mul = CustomMul() self.linear = nn.Linear(1, 1, bias=False) def forward(self, x): a = self.custom_add(x, x) a = a.clone() # 克隆张量以确保可修改属性 a.name = "CustomAdd(a=x+x)" b = self.custom_mul(a, x) b = b.clone() b.name = "CustomMul(b=a*x)" out = self.linear(b) return out, a, b # 前向传播时获取中间变量 output, a, b = model(x) loss = (output - torch.tensor([[1.0]])) ** 2 # 将模型参数和中间变量合并传入params dot = make_dot(loss, params={**dict(model.named_parameters()), "a": a, "b": b}) dot.format = 'png' dot.render('simple_model_graph_named')
方案2:自定义autograd.Function并指定操作名称
通过自定义autograd.Function并设置其__name__属性,可让反向传播节点显示自定义名称:
import torch import torch.nn as nn from torchviz import make_dot # 自定义Add操作的autograd Function class CustomAddFunc(torch.autograd.Function): @staticmethod def forward(ctx, x, y): ctx.save_for_backward(x, y) return x + y @staticmethod def backward(ctx, grad_output): x, y = ctx.saved_tensors return grad_output, grad_output # 包装为Module类 class CustomAdd(nn.Module): def forward(self, x, y): return CustomAddFunc.apply(x, y) # 设置自定义名称,反向节点会显示为CustomAddBackward CustomAddFunc.__name__ = "CustomAdd" # 自定义Mul操作的autograd Function class CustomMulFunc(torch.autograd.Function): @staticmethod def forward(ctx, x, y): ctx.save_for_backward(x, y) return x * y @staticmethod def backward(ctx, grad_output): x, y = ctx.saved_tensors return grad_output * y, grad_output * x class CustomMul(nn.Module): def forward(self, x, y): return CustomMulFunc.apply(x, y) CustomMulFunc.__name__ = "CustomMul" # 后续模型定义和可视化代码与原示例一致 class SimpleModel(nn.Module): def __init__(self): super(SimpleModel, self).__init__() self.custom_add = CustomAdd() self.custom_mul = CustomMul() self.linear = nn.Linear(1, 1, bias=False) def forward(self, x): a = self.custom_add(x, x) b = self.custom_mul(a, x) out = self.linear(b) return out model = SimpleModel() x = torch.tensor([[2.0]], requires_grad=True) output = model(x) loss = (output - torch.tensor([[1.0]])) ** 2 loss.backward() dot = make_dot(loss, params=dict(model.named_parameters())) dot.format = 'png' dot.render('simple_model_graph_custom_func')
此时生成的计算图中,反向传播节点会显示为CustomAddBackward和CustomMulBackward,清晰对应自定义操作。
方案3:开启Torchviz的属性显示功能
调用make_dot时添加show_attrs=True和show_saved=True参数,可显示模块的属性信息,帮助识别节点对应的自定义模块:
dot = make_dot(loss, params=dict(model.named_parameters()), show_attrs=True, show_saved=True) dot.format = 'png' dot.render('simple_model_graph_with_attrs')
内容的提问来源于stack exchange,提问作者Andreas Schuldei
相关产品推荐
相关产品推荐

