如何在Tensorboard中查看格式规整的PyTorch模型结构摘要
TensorBoard中展示规整PyTorch模型层级结构的实现方法
直接调用writer.add_text("model", str(model))出现格式错乱,核心原因是add_text()默认将传入内容按Markdown规则解析,纯文本的换行、缩进不会被自动保留,只要做简单格式包裹即可实现和控制台print(net)完全一致的展示效果。
方案1:代码块包裹模型字符串写入(推荐用于多版本结构对比)
该方法不需要构造模型输入、不触发前向计算,写入的内容和控制台print输出完全一致,适合多组对比实验下不同版本模型的结构留存。
实现代码如下:
from torch.utils.tensorboard import SummaryWriter # 初始化日志写入器、实例化模型 writer = SummaryWriter(log_dir="./runs/exp_version1") net = Net(input_shape=(1,64,1292), num_classes=4) # 将模型字符串包裹为Markdown代码块格式 formatted_model_str = "```python\n" + str(net) + "\n```" # 写入TensorBoard,可通过global_step参数区分不同实验版本 writer.add_text("model/architecture", formatted_model_str, global_step=1) writer.close()
写入后打开TensorBoard的Text面板,就能看到缩进、换行完全和控制台一致,还带基础语法高亮的模型结构。
方案2:写入交互式计算图
如果需要同时留存可交互的模型计算图,只需要构造一个和模型输入形状匹配的哑张量传入add_graph即可,不需要使用真实数据集:
writer = SummaryWriter(log_dir="./runs/exp_version1") net = Net(input_shape=(1,64,1292), num_classes=4) # 构造带batch维度的哑输入,形状为(batch_size, channels, height, width) dummy_input = torch.randn(1, 1, 64, 1292) writer.add_graph(net, dummy_input) writer.close()
附:你贴出的Net代码存在两处会导致运行报错的问题:
- 计算卷积层输出特征数时使用了
np.prod,但未导入numpy库- forward函数中调用了
self.feature_extractor(x),但你定义的卷积模块属性名为self.conv,属性名不匹配会抛出AttributeError
内容的提问来源于stack exchange,提问作者bikram
相关产品推荐
相关产品推荐

