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

如何在PyTorch中创建支持动态序列长度的LSTM

问题描述

我在PyTorch中实现了一个LSTM模型,需要让它支持可变序列长度。以下是我的代码:

class Seq2SeqSingle(nn.Module):
    def __init__(self, input_size, hidden_size, num_layers, in_features, out_features):
        super(Seq2SeqSingle, self).__init__()
        self.out_features = out_features
        self.num_layers = num_layers
        self.input_size = input_size
        self.hidden_size = hidden_size

        self.fc_i = nn.Linear(input_size, out_features)
        self.fc_o = nn.Linear(out_features, input_size)
        self.lstm = nn.LSTM(input_size=input_size, hidden_size=hidden_size, num_layers=num_layers, batch_first=True)
        self.fc_0 = nn.Linear(128*11, out_features)         ## <----------- LOOK HERE
        self.fc_1 = nn.Linear(out_features, out_features)

    def forward(self, x):
        #print(x.shape)
        output = self.fc_i(torch.relu(x))
        output = self.fc_o(torch.relu(output))
        
        h_0 = Variable(torch.zeros(self.num_layers, x.size(0), self.hidden_size)).to(device)
        c_0 = Variable(torch.zeros(self.num_layers, x.size(0), self.hidden_size)).to(device)
        output, (h_out, c_out) = self.lstm(output, (h_0, c_0))
        output = output.reshape(x.size(0), -1)
        output = self.fc_0(torch.relu(output))
        output = self.fc_1(torch.relu(output))
        output = nn.functional.softmax(output, dim = 1)
        return output

当前为匹配LSTM层的输出尺寸,我用隐藏层大小128乘以序列长度11定义fc_0的输入维度,但更换序列长度时程序会崩溃,请问如何避免这种固定尺寸设定?


解决方案

要让模型支持可变序列长度,核心是不要把序列长度硬编码到全连接层的输入维度中,可以通过以下两种常用方法解决:

方法1:使用LSTM的最后时刻隐藏状态

LSTM返回的(h_out, c_out)中,h_out是各层最后时刻的隐藏状态,形状为[num_layers, batch_size, hidden_size]。直接取最后一层的状态作为全连接层输入,无需依赖序列长度。

修改后的代码:

class Seq2SeqSingle(nn.Module):
    def __init__(self, input_size, hidden_size, num_layers, in_features, out_features):
        super(Seq2SeqSingle, self).__init__()
        self.out_features = out_features
        self.num_layers = num_layers
        self.input_size = input_size
        self.hidden_size = hidden_size

        self.fc_i = nn.Linear(input_size, out_features)
        self.fc_o = nn.Linear(out_features, input_size)
        self.lstm = nn.LSTM(input_size=input_size, hidden_size=hidden_size, num_layers=num_layers, batch_first=True)
        # 输入维度改为hidden_size,无需依赖序列长度
        self.fc_0 = nn.Linear(hidden_size, out_features)         
        self.fc_1 = nn.Linear(out_features, out_features)

    def forward(self, x):
        output = self.fc_i(torch.relu(x))
        output = self.fc_o(torch.relu(output))
        
        # 自动获取输入设备,避免手动指定device
        h_0 = torch.zeros(self.num_layers, x.size(0), self.hidden_size).to(x.device)
        c_0 = torch.zeros(self.num_layers, x.size(0), self.hidden_size).to(x.device)
        output, (h_out, c_out) = self.lstm(output, (h_0, c_0))
        
        # 取最后一层的最后时刻隐藏状态,形状变为[batch_size, hidden_size]
        output = h_out[-1, :, :]
        
        output = self.fc_0(torch.relu(output))
        output = self.fc_1(torch.relu(output))
        output = nn.functional.softmax(output, dim = 1)
        return output

说明:该方法适合只关注序列最终状态的任务(如分类),PyTorch 0.4.0后Variable已废弃,直接用tensor即可。

方法2:使用全局池化(平均/最大池化)

对LSTM输出的整个序列维度做池化操作,将[batch_size, seq_len, hidden_size]压缩为[batch_size, hidden_size],适配任意序列长度。

修改后的代码:

class Seq2SeqSingle(nn.Module):
    def __init__(self, input_size, hidden_size, num_layers, in_features, out_features):
        super(Seq2SeqSingle, self).__init__()
        self.out_features = out_features
        self.num_layers = num_layers
        self.input_size = input_size
        self.hidden_size = hidden_size

        self.fc_i = nn.Linear(input_size, out_features)
        self.fc_o = nn.Linear(out_features, input_size)
        self.lstm = nn.LSTM(input_size=input_size, hidden_size=hidden_size, num_layers=num_layers, batch_first=True)
        # 输入维度改为hidden_size,池化后序列维度被压缩
        self.fc_0 = nn.Linear(hidden_size, out_features)         
        self.fc_1 = nn.Linear(out_features, out_features)

    def forward(self, x):
        output = self.fc_i(torch.relu(x))
        output = self.fc_o(torch.relu(output))
        
        h_0 = torch.zeros(self.num_layers, x.size(0), self.hidden_size).to(x.device)
        c_0 = torch.zeros(self.num_layers, x.size(0), self.hidden_size).to(x.device)
        output, (h_out, c_out) = self.lstm(output, (h_0, c_0))
        
        # 全局平均池化:对序列维度(dim=1)做平均
        output = torch.mean(output, dim=1)
        # 或者用全局最大池化:output = torch.max(output, dim=1)[0]
        
        output = self.fc_0(torch.relu(output))
        output = self.fc_1(torch.relu(output))
        output = nn.functional.softmax(output, dim = 1)
        return output

说明:该方法能保留整个序列的信息,适合需要综合序列所有时刻特征的任务。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.10 02:15:43