如何获取自定义LSTMModel模型的结构摘要?
获取PyTorch LSTM模型的结构摘要方法
以下是几种获取你提供的LSTMModel结构摘要的实用方法:
1. 直接打印模型实例
这是最基础的方式,创建模型实例后直接打印,会输出各层的类型、参数配置:
import torch import torch.nn as nn # 你的LSTMModel类定义... model = LSTMModel() print(model)
输出会包含LSTM层、全连接层、激活层的详细结构信息。
2. 使用torchsummary库(显示输入输出形状)
这个库能直观展示每层的输入输出形状、参数总量,需要先安装:
pip install torchsummary
然后编写代码:
from torchsummary import summary model = LSTMModel() # 由于模型设置了batch_first=True,input_size指定为(序列长度, 输入维度) summary(model, input_size=(10, 99)) # 序列长度可根据实际需求调整
输出会列出每层的名称、输入输出形状、参数数量,以及总参数量。
3. 使用torchinfo库(更详细的模型统计)
torchinfo比torchsummary更灵活,支持复杂模型结构,还能显示显存占用等信息,安装命令:
pip install torchinfo
使用示例:
from torchinfo import summary model = LSTMModel() # input_size格式为(批量大小, 序列长度, 输入维度) summary(model, input_size=(32, 10, 99))
输出会包含每层的详细统计,包括可训练参数、非可训练参数、输入输出形状等。
额外提示:你的模型forward函数中
h0和c0的定义未考虑批量大小,实际运行时会报错,建议修改为:h0 = torch.zeros(self.n_layers, input.size(0), self.hidden_dim).to(input.device) c0 = torch.zeros(self.n_layers, input.size(0), self.hidden_dim).to(input.device)这样能适配不同批量大小的输入,同时将张量移到模型所在设备(CPU/GPU)。
内容的提问来源于stack exchange,提问作者Bibek Chalise
相关产品推荐
相关产品推荐

