PyTorch 神经网络模型可视化(Netron)

简介: PyTorch 神经网络模型可视化(Netron)

PyTorch 神经网络模型可视化(Netron

Netron 是一个用于可视化深度学习模型的工具,可以帮助我们更好地理解模型的结构和参数。

支持以下格式的模型存储文件:

格式 模板(文件) 免下载打开
ONNX squeezenet open
TensorFlow Lite yamnet open
TensorFlow chessbot open
Keras mobilenet open
TorchScript traced_online_pred_layer open
Core ML exermote open
Darknet yolo open

GitHub 链接:https://github.com/lutzroeder/netron

官网:https://netron.app


ONNX

(1)在 PyTorch 中,可以使用 torch.onnx.export 函数将模型导出为 ONNX 格式:

import torch
import netron
# 定义 PyTorch 模型
class MyModel(torch.nn.Module):
    def __init__(self):
        super(MyModel, self).__init__()
        self.conv = torch.nn.Conv2d(3, 64, kernel_size=3, stride=1, padding=1)
        self.bn = torch.nn.BatchNorm2d(64)
        self.relu = torch.nn.ReLU(inplace=True)
        self.pool = torch.nn.MaxPool2d(kernel_size=2, stride=2)
        self.fc = torch.nn.Linear(64 * 8 * 8, 10)
    def forward(self, x):
        x = self.conv(x)
        x = self.bn(x)
        x = self.relu(x)
        x = self.pool(x)
        x = x.view(-1, 64 * 8 * 8)
        x = self.fc(x)
        return x
# 创建模型实例并加载预训练权重
model = MyModel()
# 设置示例输入
input = torch.randn(1, 3, 32, 32)
# 将模型导出为 ONNX 格式
torch.onnx.export(model, input, './model/Test/onnx_model.onnx')  # 导出后 netron.start(path) 打开

(2)再使用 Netron 的 netron.start 指令打开导出的 ONNX 模型文件:

import netron
# 打开导出的 ONNX 模型文件
netron.start('./model/Test/onnx_model.onnx')
Serving './model/Test/onnx_model.onnx' at http://localhost:8080

将在浏览器中自动启动 Netron 工具,并对该模型文件进行可视化。

注意:

当模型被导出为 ONNX 格式,会在指定目录生成以 .onnx 为后缀的文件,只需将其上传至 Netron 官网 也可实现可视化:

在 Netron 中,可以查看模型的结构、参数和输入输出等信息。可以通过缩放、旋转和平移等操作来调整模型的可视化效果,以更好地理解模型的结构和参数。

torch.save

当使用 torch.save 对保存的模型进行可视化时:

# 保存模型
torch.save(model.state_dict(), './model/Test/saved_model.pt')
# 可视化
netron.start('./model/Test/saved_model.pt')

如下图,这种方式并不能显示该模型的详细信息:

所以: Netron 不支持 PyTorch 通过 torch.save 方式导出的模型文件。

torch.jit.script

可参考:torch.jit.script 与 torch.jit.trace

使用 torch.jit.script 先将模型转换为脚本,再使用 torch.jit.save 保存模型,最后进行可视化:

# TorchScript:script
scripted_model = torch.jit.script(model)
# 保存模型
torch.jit.save(scripted_model, './model/Test/scripted_model.pth')
# 可视化
netron.start('./model/Test/scripted_model.pth')

torch.jit.trace

可参考:torch.jit.script 与 torch.jit.trace

使用 torch.jit.trace 先将模型转换为跟踪模型执行的工具,再使用 torch.jit.save 保存模型,最后进行可视化:

# TorchScript:trace
traced_model = torch.jit.trace(model, torch.randn(1, 3, 32, 32))
# 保存模型
torch.jit.save(traced_model, './model/Test/traced_model.pth')
# 可视化
netron.start('./model/Test/traced_model.pth')

目录
相关文章
|
4天前
|
存储 分布式计算 监控
应用层---网络模型
应用层---网络模型
15 3
|
3天前
|
机器学习/深度学习 存储 算法
使用Python实现深度学习模型:强化学习与深度Q网络(DQN)
使用Python实现深度学习模型:强化学习与深度Q网络(DQN)
18 2
|
11天前
|
机器学习/深度学习 搜索推荐 算法
基于深度学习神经网络协同过滤模型(NCF)的图书推荐系统
登录注册 热门图书 图书分类 图书推荐 借阅图书 购物图书 个人中心 可视化大屏 后台管理
12825 1
基于深度学习神经网络协同过滤模型(NCF)的图书推荐系统
|
4天前
|
机器学习/深度学习 数据采集 TensorFlow
使用Python实现深度学习模型:图神经网络(GNN)
使用Python实现深度学习模型:图神经网络(GNN)
11 1
|
4天前
|
网络协议 网络性能优化 数据安全/隐私保护
计算机网络基础知识和术语(二)---分层结构模型
计算机网络基础知识和术语(二)---分层结构模型
8 1
|
2天前
|
NoSQL Java Redis
Redis系列学习文章分享---第十八篇(Redis原理篇--网络模型,通讯协议,内存回收)
Redis系列学习文章分享---第十八篇(Redis原理篇--网络模型,通讯协议,内存回收)
10 0
|
2天前
|
存储 消息中间件 缓存
Redis系列学习文章分享---第十七篇(Redis原理篇--数据结构,网络模型)
Redis系列学习文章分享---第十七篇(Redis原理篇--数据结构,网络模型)
7 0
|
3天前
|
并行计算 PyTorch 程序员
老程序员分享:Pytorch入门之Siamese网络
老程序员分享:Pytorch入门之Siamese网络
|
17天前
|
机器学习/深度学习 自然语言处理 算法
【从零开始学习深度学习】49.Pytorch_NLP项目实战:文本情感分类---使用循环神经网络RNN
【从零开始学习深度学习】49.Pytorch_NLP项目实战:文本情感分类---使用循环神经网络RNN
|
17天前
|
机器学习/深度学习 PyTorch 算法框架/工具
【从零开始学习深度学习】30. 神经网络中批量归一化层(batch normalization)的作用及其Pytorch实现
【从零开始学习深度学习】30. 神经网络中批量归一化层(batch normalization)的作用及其Pytorch实现

热门文章

最新文章