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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 20:40:54