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

如何在PyTorch中正确使用LSTM解决3D输入要求及模型效果退化问题

问题根因分析

你当前代码的核心错误是使用了nn.EmbeddingBag作为嵌入层,它的输出已经丢失了LSTM必需的序列维度:

  • nn.EmbeddingBag的作用是将每个样本对应的所有token embedding直接做求和/平均/最值聚合,输出形状固定为[batch_size, embed_dim],你当前得到的[10, 32]就是10个样本各自的全局句子级向量,完全没有时序维度信息。
  • 你强行用view(1, x.shape[0], -1)把它转成[1, 10, 32]喂给LSTM,相当于把10个样本当成了长度为1的序列的10个batch,完全打乱了数据的语义逻辑,因此模型效果极差。
正确实现方案

要使用LSTM做文本分类,需要调整如下核心逻辑:

1. 替换嵌入层

将nn.EmbeddingBag替换为普通nn.Embedding,保留每个token的独立嵌入向量,维持序列维度。

2. 调整输入预处理逻辑

原有的text+offsets是适配EmbeddingBag的输入格式,现在需要将同一batch内的不同长度句子做填充(padding),统一为[batch_size, 最大序列长度]的二维张量作为输入,同时记录每个句子的真实长度用于后续变长序列优化。

3. 调整LSTM参数

添加batch_first=True参数,让LSTM适配[batch_size, seq_len, embed_dim]的输入格式,更符合常规数据排布逻辑。

4. 处理LSTM输出

文本分类场景下,一般取LSTM输出的最后一个有效时间步的隐藏状态,或者对所有时间步输出做全局池化,再传入后续全连接层。

修改后代码示例
import torch
import torch.nn as nn
from torch.nn.utils.rnn import pack_padded_sequence, pad_packed_sequence

torch.manual_seed(random_state)

class Net(nn.Module):
    def __init__(self, vocab_size= len(vocab), embed_dim= 32, num_class= 3, lstm_layers=5):
        super().__init__()
        # 替换为普通Embedding,padding_idx指定填充位的id
        self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=0)

        # 添加batch_first=True,适配[batch, seq_len, embed_dim]格式输入
        self.lstm = nn.LSTM(input_size= embed_dim, hidden_size= embed_dim, num_layers=lstm_layers, batch_first=True)

        # 全连接层添加激活函数,避免线性堆叠失效
        self.lin = nn.Sequential(
            nn.Linear(in_features= embed_dim, out_features= embed_dim),
            nn.ReLU(),
            nn.Linear(in_features= embed_dim, out_features= embed_dim),
            nn.ReLU(),
            nn.Linear(in_features= embed_dim, out_features= 16),
            nn.ReLU(),
            nn.Linear(in_features= 16, out_features= 16),
            nn.ReLU(),
            nn.Linear(in_features= 16, out_features= 8),
            nn.ReLU(),
            nn.Linear(in_features= 8, out_features= 8),
            nn.ReLU(),
        )

        self.out = nn.Linear(in_features= 8, out_features= num_class)

        self.init_weights()

    def init_weights(self):
        initrange = 0.5
        self.embedding.weight.data.uniform_(-initrange, initrange)
        for layer in self.lin:
            if isinstance(layer, nn.Linear):
                layer.weight.data.uniform_(-initrange, initrange)
                layer.bias.data.zero_()

    def forward(self, text, seq_len):
        # text形状为[batch_size, max_seq_len],seq_len为每个句子的真实长度
        x = self.embedding(text) # 输出形状[batch_size, max_seq_len, embed_dim]
        
        # 可选:使用pack_padded_sequence优化变长序列计算,忽略pad部分的无效计算
        packed_x = pack_padded_sequence(x, seq_len.cpu(), batch_first=True, enforce_sorted=False)
        packed_out, _ = self.lstm(packed_x)
        x, _ = pad_packed_sequence(packed_out, batch_first=True) # 输出形状[batch_size, max_seq_len, embed_dim]
        
        # 取最后一个有效时间步的输出作为句子特征
        last_idx = seq_len - 1
        last_output = x[torch.arange(x.shape[0]), last_idx, :] # 形状[batch_size, embed_dim]

        x = self.lin(last_output)
        return self.out(x)

model = Net()
model.to(device)
输入适配说明

训练时的输入需要调整为:

  • text:填充后的token id张量,形状[batch_size, max_seq_len]
  • seq_len:一维张量,记录batch内每个句子的原始长度,长度等于batch size

内容的提问来源于stack exchange,提问作者user9102437

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 19:15:06