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

PyTorch实现U-Net架构可视化失败,求可行绘图方案

PyTorch版U-Net架构可视化解决方案

问题根源

你的U-Net是用PyTorch实现的,但你尝试的plot_model是Keras/TensorFlow专属工具,框架不兼容,所以无法生成架构图。下面提供三种适配PyTorch的可视化方案:


方案1:用torchviz生成计算图

torchviz可以直接基于PyTorch的计算图生成可视化结构,步骤如下:

  1. 补全代码中缺失的导入(原代码漏了部分PyTorch模块):
from torch.nn import Module, Conv2d, ReLU, MaxPool2d, ConvTranspose2d, ModuleList
# 临时定义config(若你的config模块已存在可忽略)
class Config:
    INPUT_IMAGE_HEIGHT = 256
    INPUT_IMAGE_WIDTH = 256
config = Config()
  1. 安装torchviz:
pip install torchviz
  1. 运行可视化代码:
import torch
from torchviz import make_dot

# 实例化U-Net模型
model = UNet()
# 创建匹配输入尺寸的dummy张量(示例为3通道256x256)
dummy_input = torch.randn(1, 3, 256, 256)
# 前向传播得到输出
output = model(dummy_input)
# 生成可视化图并保存
graph = make_dot(output, params=dict(model.named_parameters()))
graph.render("unet_architecture")  # 保存为PDF文件
graph.view()  # 直接打开查看

方案2:用torchinfo输出结构化摘要

如果不需要图形化,只想快速查看模型层级、参数数量、输入输出尺寸,torchinfo更高效:

  1. 安装torchinfo:
pip install torchinfo
  1. 运行代码:
from torchinfo import summary

model = UNet()
# 输入格式:(batch_size, channels, height, width)
summary(model, input_size=(1, 3, 256, 256))

输出会清晰展示U-Net编码器、解码器的每一层名称、输入输出形状、参数数量。


方案3:用Netron交互式可视化

Netron是通用模型可视化工具,支持PyTorch、ONNX等格式,步骤如下:

  1. 安装Netron:
pip install netron
  1. 将PyTorch模型导出为ONNX格式:
import torch

model = UNet()
dummy_input = torch.randn(1, 3, 256, 256)
# 导出模型
torch.onnx.export(model, dummy_input, "unet_model.onnx")
  1. 启动Netron查看:
import netron
netron.start("unet_model.onnx")

会自动打开浏览器,展示交互式架构图,可展开查看每个子模块的细节。


注意事项

  • 原代码中实例化模型的错误:model = UNet 需改为 model = UNet()(必须调用构造函数创建实例)。
  • 若你的config模块已存在,无需临时定义Config类,直接导入即可。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 02:00:38