【问题标题】:Pytorch - Incorrect dimensions when using LSTM networkPytorch - 使用 LSTM 网络时尺寸不正确
【发布时间】:2018-11-06 11:55:07
【问题描述】:

我开始使用 pytorch,并使用其中一个教程作为参考,使用了一些转换来构建以下模型:

model = torch.nn.Sequential( 
     torch.nn.Linear(D_in, H),
     torch.nn.ReLU(),
     torch.nn.Linear(H, D_out),
)

我想使用 LSTM 网络,所以我尝试了以下操作:

model = torch.nn.Sequential(
      torch.nn.LSTM(D_in, H),
      torch.nn.Linear(H, D_out) 
)

这给了我这个错误:

RuntimeError:输入必须有 3 个维度,得到 2 个维度

为什么我会看到这个错误?我预计我对如何在 pytorch 中链接转换(网络?)的理解存在根本性的错误......

编辑

在遵循@esBee 的建议后,我发现以下运行正常。这是因为 LSTM 期望输入具有以下维度:

input of shape (seq_len, batch, input_size):包含输入序列特征的张量。输入也可以是打包的变长序列

local_x = local_x.unsqueeze(0)
y_pred, (hn, cn) = layerA(local_x)
y_pred = y_pred.squeeze(0)
y_pred = layerB(y_pred)

但是,我的原始训练/测试数据集只有序列长度 1 的事实让我觉得我做错了什么。这个参数在神经网络中的作用是什么?

【问题讨论】:

    标签: deep-learning lstm pytorch


    【解决方案1】:

    这里需要注意的是,与 torch.nn.Linear 等线性层相反,torch.nn.LSTM 等重复层有多个输出。

    虽然torch.nn.Linear 只返回y = Ax + b 中的y,但torch.nn.LSTMs 返回output, (h_n, c_n)(更详细地解释in the docs)让您选择要处理的输出。因此,在您的示例中发生的情况是,您将所有这几种类型的输出输入到 LSTM 层之后的层中(导致您看到的错误)。相反,您应该选择 LSTM 输出的特定部分,并将其仅提供给下一层。

    遗憾的是,我不知道如何在 Sequential 中选择 LSTM 的输出(欢迎提出建议),但您可以重写

    model = torch.nn.Sequential(
        torch.nn.LSTM(D_in, H),
        torch.nn.Linear(H, D_out) 
    )
    
    model(x)
    

    作为

    layerA = torch.nn.LSTM(D_in, H)
    layerB = torch.nn.Linear(H, D_out)
    
    x = layerA(x)
    x = layerB(x)
    

    然后通过编写选择 LSTM 最后一层的输出特征 (h_n) 来纠正它

    layerA = torch.nn.LSTM(D_in, H)
    layerB = torch.nn.Linear(H, D_out)
    
    x = layerA(x)[0]
    x = layerB(x)
    

    【讨论】:

    • 请看我的编辑。我很想听听你的意见。
    【解决方案2】:

    错误消息告诉您输入需要三个维度。

    查看pytorchdocumentation,他们提供的例子是这样的:

    lstm = nn.LSTM(3, 3)  # Input dim is 3, output dim is 3
    

    D_inH 都没有三个维度。

    【讨论】:

    • 所有 LSTM 网络都需要三维输入吗?
    • PyTorch RNNs 一般采用 3-dim 输入,但这当然不是 LSTMs 的一般要求,你可以构造不同输入形状的 LSTM。举个例子,如果您不使用批处理(将 batch-dim 设置为 1),您实际上只需使用二维。你可以看这里:stackoverflow.com/questions/50399055/…
    猜你喜欢
    • 2020-02-19
    • 2022-01-19
    • 2019-06-05
    • 1970-01-01
    • 2019-03-28
    • 1970-01-01
    • 2011-08-17
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多