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

PyTorch实现ConvLSTM中in_channels参数设置的疑问

ConvLSTM实现疑问解答

问题描述

我在看一段PyTorch实现的ConvLSTM代码时存在疑问:代码中nn.Conv2d的in_channels设置为in_channels + out_channels,对应输入X_t和隐藏状态H_(t-1)。但根据Shi等人2015年的ConvLSTM原论文《Convolutional LSTM Network: A Machine Learning Approach for Precipitation Nowcasting》,我以为ConvLSTM应包含X_t、H_(t-1)、C_(t-1)三个输入,因此觉得in_channels应该设置为in_channels + 2*out_channels,想确认是否因为C_(t-1)属于窥视孔结构,所以未作为输入传入卷积层?

相关代码如下:

import torch
import torch.nn as nn

class ConvLSTMCell(nn.Module):
    def __init__(self, in_channels, out_channels, kernel_size, bias):
        """
        Initialize ConvLSTM cell.
        Parameters
        ----------
        in_channels: int 输入特征图的通道数
        out_channels: int 输出特征图的通道数
        kernel_size: (int, int) 卷积核的宽和高
        bias: bool 是否使用偏置
        """
        super().__init__()
        self.in_channels = in_channels
        self.out_channels = out_channels

        self.kernel_size = kernel_size # wsg: 此处 kernel_size 需设为列表
        self.padding = kernel_size[0] // 2, kernel_size[1] // 2
        # 需要强制进行padding以保证每次卷积后形状不发生变化
        # 在stride=1时,padding = kernel_size // 2
        # 如:卷积核为3×3则需要padding=1即可
        # 在下面的卷积操作中stride使用的是默认值1
        self.bias = bias
        self.conv = nn.Conv2d(in_channels=self.in_channels + self.out_channels,
                              out_channels=4 * self.out_channels,
                              kernel_size=self.kernel_size,
                              padding=self.padding,
                              bias=self.bias)
        # wsg:初始化输入通道,in+out vs. in+out*2?C_t-1 不算?因为是窥视结构?

    def forward(self, input_tensor, last_state):
        """
        params
            input_tensor: 当前时刻的输入x_t, 形状为
                         [batch_size, in_channels, height, width]
            last_state: 上一时刻的状态c_{t-1}和h_{t-1}, 形状均为 
                        [batch_size, out_channels, height, width]
        return
        """
        
        h_last, c_last = last_state # 前一时刻
        combined_input = torch.cat([input_tensor, h_last], dim=1)
        # [batch_size, in_channels+out_channels, height, width]
        combined_conv = self.conv(combined_input)  
        # [batch_size, 4 * out_channels, height, width]
        # wsg:输入只有 x_t 与 h_(t-1)?没有 C_(t-1) 吗?
        cc_i, cc_f, cc_o, cc_g = torch.split(combined_conv,
                                             self.out_channels, dim=1)
        # 分割得到每个门对应的卷积计算结果,形状均为
        #   [batch_size, out_channels, height, width]
        i = torch.sigmoid(cc_i)
        f = torch.sigmoid(cc_f)
        o = torch.sigmoid(cc_o)
        g = torch.tanh(cc_g) 
        c_next = f * c_last + i * g  # [batch_size, out_channels, height, width]
        h_next = o * torch.tanh(c_next)  # [batch_size, out_channels, height, width]
        return h_next, c_next

    def init_hidden(self, batch_size, image_size):
        height, width = image_size
        return (torch.zeros(batch_size, self.out_channels, height, width, device=self.conv.weight.device),
                torch.zeros(batch_size, self.out_channels, height, width, device=self.conv.weight.device))

解答

  • 这段代码的实现完全符合Shi 2015年原论文的ConvLSTM定义,C_(t-1)本来就不需要作为卷积层的输入
  • 原论文中,输入门i、遗忘门f、候选细胞状态g的计算仅依赖当前输入X_t和上一时刻的隐藏状态H_(t-1),卷积层只需要接收这两个输入拼接后的特征图,因此in_channels设置为in_channels + out_channels是正确的
  • C_(t-1)的作用是直接参与细胞状态的更新计算(对应代码中的c_next = f * c_last + i * g),它并没有通过卷积层处理,这和窥视孔结构无关——窥视孔LSTM是将门的计算和C_(t-1)关联,而标准ConvLSTM并没有这个设计

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 15:23:19