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

PyTorch双向ConvLSTM实现代码报错求助:initialize_weights参数缺失TypeError问题排查

Fixing the Bidirectional ConvLSTM Implementation in PyTorch

Let's break down and fix the issues in your bidirectional ConvLSTM implementation step by step:

1. The initialize_weights Function Parameter Mismatch

The error TypeError: initialize_weights() missing 1 required positional argument: 'layer' happens because you defined initialize_weights with a self parameter (like a class method), but it's actually a standalone utility function. When using nn.Module.apply(), PyTorch passes each submodule as the first argument to the function, so the extra self causes a parameter mismatch.

Fix: Remove the self parameter from the function definition:

def initialize_weights(layer):
    """Initialize a layer's weights and biases.
    Args:
        layer: A PyTorch Module's layer."""
    if isinstance(layer, (nn.BatchNorm2d, nn.BatchNorm1d)):
        pass
    else:
        try:
            nn.init.xavier_normal_(layer.weight)
        except AttributeError:
            pass
        try:
            nn.init.uniform_(layer.bias)
        except (ValueError, AttributeError):
            pass

2. Missing forward Method in ConvLSTMCell

Your ConvLSTMCell class doesn't have a forward method, which is required for PyTorch modules to execute the forward pass. This would cause an error even after fixing the weight initialization.

Fix: Implement the forward method following the ConvLSTM formula from the original paper:

def forward(self, x, states=None):
    # Initialize hidden and cell states if not provided
    batch_size = x.size(0)
    if states is None:
        h = torch.zeros(batch_size, self.kernels, self.input_dim, self.input_dim, device=x.device)
        c = torch.zeros(batch_size, self.kernels, self.input_dim, self.input_dim, device=x.device)
    else:
        h, c = states

    # Apply batch normalization if enabled
    if self.batch_norm_layer is not None:
        x = self.batch_norm_layer(x)

    # Calculate gates
    i = torch.sigmoid(self.W_xi(x) + self.W_hi(h) + self.W_ci(c))  # Input gate
    f = torch.sigmoid(self.W_xf(x) + self.W_hf(h) + self.W_cf(c))  # Forget gate
    c = f * c + i * torch.tanh(self.W_xc(x) + self.W_hc(h))        # Cell state update
    o = torch.sigmoid(self.W_xo(x) + self.W_ho(h) + self.W_co(c))  # Output gate
    h = o * torch.tanh(c)                                          # Hidden state update

    # Apply dropout
    h = self.H_drop(h)
    c = self.C_drop(c)

    return h, (h, c)

3. Unused Bias in HadamardProduct

Your HadamardProduct class initializes a bias parameter but doesn't use it in the forward pass. This is a minor oversight that wastes parameters.

Fix: Update the forward method to include the bias:

def forward(self, x):
    return x * self.weights + self.bias

4. Input Dimension Mismatch & Sequence Indexing Bug

  • Your test code defines ConvLSTM with input_dim=128, but the input tensor x has spatial dimensions (224,224), causing a shape mismatch.
  • The backward sequence indexing in the ConvLSTM.forward method was incorrect (x[:,-seq_idx,::] would skip the first element when reversing the sequence).

Fix: Align input dimensions and correct the backward sequence index:

# In ConvLSTM.forward, replace the backward sequence line with:
layer_in_out_bwd = x[:,-seq_idx-1,::]

Full Corrected Code

Here's the complete working implementation:

import torch
from torch import nn

def initialize_weights(layer):
    """Initialize a layer's weights and biases.
    Args:
        layer: A PyTorch Module's layer."""
    if isinstance(layer, (nn.BatchNorm2d, nn.BatchNorm1d)):
        pass
    else:
        try:
            nn.init.xavier_normal_(layer.weight)
        except AttributeError:
            pass
        try:
            nn.init.uniform_(layer.bias)
        except (ValueError, AttributeError):
            pass

class HadamardProduct(nn.Module):
    """A Hadamard product layer.
    Args:
        shape: The shape of the layer."""
    def __init__(self, shape):
        super().__init__()
        self.weights = nn.Parameter(torch.empty(*shape))
        self.bias = nn.Parameter(torch.empty(*shape))
        # Initialize weights and bias for Hadamard layer
        nn.init.xavier_normal_(self.weights)
        nn.init.uniform_(self.bias)
        
    def forward(self, x):
        return x * self.weights + self.bias

class ConvLSTMCell(nn.Module):
    """A convolutional LSTM cell.
    Implementation details follow closely the ConvLSTM paper by Shi et al. (2015)."""
    def __init__(self, input_bands, input_dim, kernels, dropout, batch_norm):
        super().__init__()
        self.input_bands = input_bands
        self.input_dim = input_dim
        self.kernels = kernels
        self.dropout = dropout
        self.batch_norm = batch_norm
        self.kernel_size = 3
        self.padding = 1 # Preserve spatial dimensions
        self.input_conv_params = {
            'in_channels': self.input_bands,
            'out_channels': self.kernels,
            'kernel_size': self.kernel_size,
            'padding': self.padding,
            'bias': True
        }
        self.hidden_conv_params = {
            'in_channels': self.kernels,
            'out_channels': self.kernels,
            'kernel_size': self.kernel_size,
            'padding': self.padding,
            'bias': True
        }
        self.state_shape = (1, self.kernels, self.input_dim, self.input_dim)
        self.batch_norm_layer = nn.BatchNorm2d(num_features=self.input_bands) if self.batch_norm else None
        
        # Input Gates
        self.W_xi = nn.Conv2d(**self.input_conv_params)
        self.W_hi = nn.Conv2d(**self.hidden_conv_params)
        self.W_ci = HadamardProduct(self.state_shape)
        # Forget Gates
        self.W_xf = nn.Conv2d(**self.input_conv_params)
        self.W_hf = nn.Conv2d(**self.hidden_conv_params)
        self.W_cf = HadamardProduct(self.state_shape)
        # Memory Gates
        self.W_xc = nn.Conv2d(**self.input_conv_params)
        self.W_hc = nn.Conv2d(**self.hidden_conv_params)
        # Output Gates
        self.W_xo = nn.Conv2d(**self.input_conv_params)
        self.W_ho = nn.Conv2d(**self.hidden_conv_params)
        self.W_co = HadamardProduct(self.state_shape)
        # Dropouts
        self.H_drop = nn.Dropout2d(p=self.dropout)
        self.C_drop = nn.Dropout2d(p=self.dropout)
        
        self.apply(initialize_weights)
        
    def forward(self, x, states=None):
        # Initialize hidden and cell states if not provided
        batch_size = x.size(0)
        if states is None:
            h = torch.zeros(batch_size, self.kernels, self.input_dim, self.input_dim, device=x.device)
            c = torch.zeros(batch_size, self.kernels, self.input_dim, self.input_dim, device=x.device)
        else:
            h, c = states

        # Apply batch normalization if enabled
        if self.batch_norm_layer is not None:
            x = self.batch_norm_layer(x)

        # Calculate gates
        i = torch.sigmoid(self.W_xi(x) + self.W_hi(h) + self.W_ci(c))  # Input gate
        f = torch.sigmoid(self.W_xf(x) + self.W_hf(h) + self.W_cf(c))  # Forget gate
        c = f * c + i * torch.tanh(self.W_xc(x) + self.W_hc(h))        # Cell state update
        o = torch.sigmoid(self.W_xo(x) + self.W_ho(h) + self.W_co(c))  # Output gate
        h = o * torch.tanh(c)                                          # Hidden state update

        # Apply dropout
        h = self.H_drop(h)
        c = self.C_drop(c)

        return h, (h, c)

class ConvLSTM(nn.Module):
    def __init__(self, input_bands, input_dim, kernels, num_layers, bidirectional, dropout):
        super().__init__()
        self.input_bands = input_bands
        self.input_dim = input_dim
        self.kernels = kernels
        self.num_layers = num_layers
        self.bidirectional = bidirectional
        self.dropout = dropout
        self.layers_fwd = self.initialize_layers()
        self.layers_bwd = self.initialize_layers() if self.bidirectional else None
        
        self.fc_output = nn.Sequential(
            nn.Flatten(),
            nn.Linear(
                in_features=self.kernels*self.input_dim**2*(1 if not self.bidirectional else 2),
                out_features=1024
            ),
            nn.Linear(
                in_features=1024,
                out_features=1
            )
        )
        self.apply(initialize_weights)
        
    def initialize_layers(self):
        """Initialize a single direction of the model's layers."""
        layers = nn.ModuleList()
        for i in range(self.num_layers):
            layers.append(
                ConvLSTMCell(
                    input_bands=self.input_bands if i == 0 else self.kernels,
                    input_dim=self.input_dim,
                    dropout=self.dropout if i+1 < self.num_layers else 0,
                    kernels=self.kernels,
                    batch_norm=False
                )
            )
        return layers
        
    def forward(self, x):
        """Perform forward pass with the model.
        Input shape: [Batch, Seq, Band, Dim, Dim]
        Returns: Batch of predictions"""
        seq_len = x.shape[1]
        final_out = None
        
        for seq_idx in range(seq_len):
            # Forward direction processing
            layer_in_out = x[:,seq_idx,::]
            states = None
            for layer in self.layers_fwd:
                layer_in_out, states = layer(layer_in_out, states)
            
            if not self.bidirectional:
                final_out = layer_in_out
                continue
            
            # Backward direction processing (reverse sequence)
            layer_in_out_bwd = x[:,-seq_idx-1,::]
            states = None
            for layer in self.layers_bwd:
                layer_in_out_bwd, states = layer(layer_in_out_bwd, states)
            
            # Concatenate forward and backward outputs
            layer_in_out = torch.cat((layer_in_out, layer_in_out_bwd), dim=1)
            final_out = layer_in_out
        
        return self.fc_output(final_out)

Test the Fixed Code

Run this test code to verify:

import torch

# Initialize model with input_dim matching input spatial size
ConvLSTM2D = ConvLSTM(input_bands=128, input_dim=224, kernels=3, num_layers=1, bidirectional=True, dropout=0.0)
# Input shape: [Batch, Seq, Band, Dim, Dim]
x = torch.randn([5, 1, 128, 224, 224])
# Forward pass
t1 = ConvLSTM2D(x)
print(t1.shape)  # Should output torch.Size([5, 1])

This should now run without errors and produce the expected output shape.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.29 21:57:44