使用Tensorboard可视化网络结构(基于pytorch)

前言

我们在搭建网络模型的时候,通常希望可以对自己搭建好的网络模型有一个比较好的直观感受,从而更好地了解网络模型的结构,Tensorboard工具的使用就给我们提供了方便的途径

Tensorboard概况

Tensorboard是由Google公司开源的一款可视化工具,是TensorFlow的一个附属组件,但在pytorch项目中也可以使用。

它有以下主要功能:

  • 可视化网络模型:您可以直观地了解模型的结构,包括层的堆栈方式,激活函数等。
  • 记录和绘制训练过程中各项指标的变化,例如loss曲线、准确率曲线等。
  • 可视化特征空间和高维度数据。
  • 可视化梯度、权重乃至激活函数输出分布等等。

这篇博客主要介绍Tensorboard可视化网络模型的功能

代码实现

我们搭建一个简单的神经网络,依赖的库环境

python 复制代码
import torch
import torch.nn as nn
import torch.nn.functional as F

from torch.utils.tensorboard import SummaryWriter

搭建网络模型

python 复制代码
class Net(nn.Module):
    def __init__(self,input_dim,layer1_dim,layer2_dim,output_dim):  
        super(Net,self).__init__()
        self.flatten = nn.Flatten() 
        self.layer1 = nn.Sequential(nn.Linear(input_dim,layer1_dim),nn.ReLU())
        self.layer2 = nn.Sequential(nn.Linear(layer1_dim,layer2_dim),nn.ReLU())
        self.out = nn.Sequential(nn.Linear(layer2_dim,output_dim),nn.Softmax(dim=-1))

    def forward(self,x):
        x = self.flatten(x)
        x = self.layer1(x)
        x = self.layer2(x)
        x = self.out(x)
        return x

# 初始化网络中的值
input_dim,layer1_dim,layer2_dim,output_dim=32*32,512,128,10
model = Net(input_dim,layer1_dim,layer2_dim,output_dim)

网络由三层全连接层组成,输入的数据形状为

python 复制代码
# 定义输入模型的数据,1表示批次,这里可以忽略
input_data = torch.rand(1,32,32)

# 模型输出的数据
output_data = model(input_data)

接下俩就是使用SummaryWriter创建日志保存搭建好的网络模型

python 复制代码
with SummaryWriter(log_dir=r"D:\CSDN_point\12_22\logs", comment="Net") as w:
    w.add_graph(model, input_data)

log_dir参数就是日志文件的本地保存路径,comment就是日志的备注,add_graph()传入网络模型和输入数据,运行后就会在指定路径上生成对应的文件,打开log文件所在的文件位置,在顶部路径上输入cmd,打开命令行窗口

在命令行窗口输入

python 复制代码
tensorboard --logdir logs

复制返回的网址在浏览器打开,就可以得到对应的网络可视化结果了

欢迎大家讨论交流~


相关推荐
AKAMAI2 分钟前
云成本困境:开支激增正阻碍欧洲AI创新
人工智能·云原生·云计算
大模型真好玩13 分钟前
LangGraph实战项目:从零手搓DeepResearch(一)——DeepResearch应用体系详细介绍
人工智能·python·mcp
IT古董20 分钟前
【第五章:计算机视觉-项目实战之生成式算法实战:扩散模型】3.生成式算法实战:扩散模型-(4)在新数据集上微调现有扩散模型
人工智能
嵌入式-老费26 分钟前
Easyx图形库使用(潜力无限的图像处理)
图像处理·人工智能
Goona_28 分钟前
PyQt批量年龄计算工具:从身份证到指定日期的周岁处理
python·小程序·交互·pyqt
JXY_AI36 分钟前
AI问答与搜索引擎:信息获取的现状
人工智能·搜索引擎
B站_计算机毕业设计之家1 小时前
Python+Flask+Prophet 汽车之家二手车系统 逻辑回归 二手车推荐系统 机器学习(逻辑回归+Echarts 源码+文档)✅
大数据·人工智能·python·机器学习·数据分析·汽车·大屏端
MoRanzhi12031 小时前
SciPy傅里叶变换与信号处理教程:数学原理与Python实现
python·机器学习·数学建模·数据分析·信号处理·傅里叶分析·scipy
XXX-X-XXJ1 小时前
三、从 MinIO 存储到 OCR 提取,再到向量索引生成
人工智能·后端·python·ocr
AI人工智能+1 小时前
行驶证识别技术通过OCR和AI实现信息自动化采集与处理,涵盖图像预处理、文字识别及结构化校验,提升效率与准确性
人工智能·深度学习·ocr·行驶证识别