PyTorch双向ConvLSTM实现代码报错求助:initialize_weights参数缺失TypeError问题排查
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
ConvLSTMwithinput_dim=128, but the input tensorxhas spatial dimensions(224,224), causing a shape mismatch. - The backward sequence indexing in the
ConvLSTM.forwardmethod 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

