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

如何在PyTorch中将MLP结构修改为输入输出一致的RNN结构

调整RNN类结构以匹配现有MLP的输入输出

现有MLP类代码

import torch.nn as nn

class sample(nn.Module):
    def __init__(self):
        super(sample, self).__init__()  # 修正原代码的拼写错误:init() → __init__()
        self.linear = nn.Linear(1, 20)
    def forward(self, t, is_train = False, y = None):
        a = self.linear(t)
        return a

原始RNN代码的问题

  • PyTorch中的循环神经网络类是nn.RNN(首字母大写),而非nn.rnn
  • 初始化时将RNN实例赋值给了self.linear,但forward方法中调用的是self.rnn,属性名不匹配
  • nn.RNN没有batch_size参数,正确的参数顺序为input_size, hidden_size, num_layers
  • RNN要求输入张量形状为(seq_len, batch_size, input_size)或(batch_size, seq_len, input_size)(需设置batch_first=True),但MLP的输入是普通二维张量(batch_size, input_size),需调整输入形状适配
  • RNN的返回值是元组(output, hidden),直接返回无法和MLP的输出维度匹配,需处理输出格式

修改后的RNN类代码

以下代码保证输入输出格式与MLP完全一致(输入为(batch_size, 1),输出为(batch_size, 20)):

import torch.nn as nn

class sample(nn.Module):
    def __init__(self, input_size=1, hidden_dim=20, num_layer=1):
        super(sample, self).__init__()
        self.input_size = input_size
        self.hidden_dim = hidden_dim
        self.num_layer = num_layer
        
        # 初始化RNN,设置batch_first=True适配batch维度在前的输入格式
        self.rnn = nn.RNN(input_size=input_size, 
                          hidden_size=hidden_dim, 
                          num_layers=num_layer,
                          batch_first=True)
        
        # 线性层用于将RNN隐藏状态映射到与MLP一致的输出维度(20维)
        self.fc = nn.Linear(hidden_dim, 20)

    def forward(self, t, is_train=False, y=None):
        # 将MLP格式的输入(batch_size, input_size)转换为RNN要求的(batch_size, seq_len, input_size)
        # 这里把单个样本视为序列长度为1的序列
        t = t.unsqueeze(1)
        
        # RNN前向传播,获取输出和最后一层隐藏状态
        output, hidden = self.rnn(t)
        
        # 取序列最后一步的输出,再通过线性层得到目标维度的结果
        last_output = output[:, -1, :]
        a = self.fc(last_output)
        
        return a

补充说明

  • 若输入本身是序列数据,可根据实际seq_len调整输入形状的处理逻辑
  • 若将hidden_dim直接设置为20,也可以直接使用hidden[-1, :, :](最后一层的隐藏状态)作为输出,无需额外线性层

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 08:40:36