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

条件GAN时间序列预测:LSTM生成器噪声与输入拼接问题求助

多变量时间序列预测GAN生成器的LSTM噪声拼接问题解决

报错原因分析

你遇到的RuntimeError是因为噪声张量的维度扩展逻辑错误:

  • noise.unsqueeze(1)将形状[16,32]的噪声转为[16,1,32]
  • 后续expand(-1, self.time_steps, self.output_dim)试图把第三维度从32改成1,这在PyTorch中不允许(非单元素维度不能直接强制修改尺寸)
  • 另外你选择在dim=1(时间步维度)拼接X和噪声,这会导致时间步长度从10变成11,不符合LSTM对输入序列结构的要求,正确的拼接维度应该是特征维度(dim=2)

正确的噪声拼接方式与代码修正

合理的拼接逻辑是:将噪声扩展为与输入时间序列相同的时间步长度,然后在特征维度拼接,让每个时间步都包含原始特征+噪声特征,这样既保留历史时间序列信息,又注入GAN所需的随机性。

修改后的完整代码如下:

import torch
import torch.nn as nn

class Generator_LSTM_LEVY(nn.Module):
    def __init__(self, hidden_dim, feature_no, time_steps, noise_size, layer_dim=1):
        super().__init__()
        self.hidden_dim = hidden_dim
        self.time_steps = time_steps
        self.feature_no = feature_no
        self.noise_size = noise_size
        self.layer_dim = layer_dim
        
        # LSTM输入维度 = 原始特征数 + 噪声维度
        self.lstm = nn.LSTM(
            input_size=feature_no + noise_size,
            hidden_size=hidden_dim,
            num_layers=layer_dim,
            batch_first=True,
            bidirectional=True
        )
        # 最终输出层:将LSTM输出映射为次日黄金价格预测(形状[batch_size,1])
        self.fc = nn.Linear(hidden_dim * 2, 1)  # 双向LSTM输出维度是2*hidden_dim

    def forward(self, x, noise): 
        # x shape: [batch_size, time_steps, feature_no] = [16,10,2]
        # noise shape: [batch_size, noise_size] = [16,32]
        
        # 扩展噪声维度:[16,32] -> [16,1,32] -> [16,10,32]
        noise_exp = noise.unsqueeze(1).expand(-1, self.time_steps, self.noise_size)
        # 在特征维度拼接x和噪声,得到[16,10,2+32=34]
        x_n = torch.cat((x, noise_exp), dim=2)  

        # LSTM前向传播,无需手动初始化h0/c0,PyTorch会自动初始化全零隐藏状态
        out, _ = self.lstm(x_n)

        # 取最后一个时间步的输出:[16,10,2*hidden_dim] -> [16,2*hidden_dim]
        out = out[:, -1, :]
        # 映射为最终预测值:[16,2*hidden_dim] -> [16,1]
        pred = self.fc(out)
        return pred

关键修正点说明

  1. 维度扩展修正:用expand(-1, self.time_steps, self.noise_size)代替错误的尺寸修改,保持噪声维度不变,仅扩展时间步长度,确保张量形状匹配
  2. 拼接维度修正:选择dim=2(特征维度)拼接,输入LSTM的张量形状保持[batch_size, time_steps, total_features],符合LSTM的输入要求
  3. 初始化参数补全:在__init__中明确传入noise_size和layer_dim,避免未定义变量的问题;双向LSTM的输出维度是2倍hidden_dim,全连接层需对应调整
  4. 隐藏状态处理:无需手动创建h0和c0,PyTorch会自动为每个batch初始化全零的隐藏状态,简化代码

拼接方式的合理性说明

将噪声与每个时间步的特征拼接是时间序列GAN中常用的策略:

  • 每个时间步都注入噪声,让生成器能基于相同的历史序列生成多种合理的预测结果,符合GAN的多样性要求
  • 保留原始时间序列的时序结构,LSTM仍能学习历史数据的时间依赖关系,不会破坏时间序列的特性

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 13:02:10