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

如何获取自定义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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.21 07:22:55