Octave卷积适配BatchNorm2d出现running_mean维度不匹配RuntimeError
Let's break down what's going wrong here and how to fix it:
The Root Cause
Looking at your forward pass code, there's a critical typo in how you apply batch normalization to the low-frequency (lf) tensor:
lf = self.bnh(lf) if type(lf) == torch.Tensor else lf #THIS IS THE LINE ACCUSING THE ERROR
You're using self.bnh (the batch norm for high-frequency features) instead of self.bnl (the one meant for low-frequency features) here.
Why It Only Triggers When alpha=1
When alpha_out=1, your calculation for hf_ch becomes:
hf_ch = int(num_features * (1 - alpha_out)) = int(64 * (1-1)) = 0
This means self.bnh is initialized as a BatchNorm2d(0) (PyTorch allows this, but it's functionally useless here), while self.bnl is correctly set to BatchNorm2d(64) matching your lf tensor's channel count.
When you try to pass your 64-channel lf tensor into self.bnh (which expects 0 channels), PyTorch throws the running_mean should contain 64 elements not 0 error because the running stats for self.bnh are initialized for 0 channels, but your input has 64.
For other alpha values, hf_ch is non-zero, so self.bnh has a valid channel count. If your lf channel count happens to match hf_ch by coincidence, you might not hit an error immediately—but it's still incorrect behavior, as you're applying the wrong batch norm to the lf tensor.
Fixed Code
Here's the corrected forward method, plus a few extra robustness improvements:
def forward(self, x): if isinstance(x, tuple): hf, lf = x print("IN ON BN: ",lf.shape if lf is not None else None) #DEBUGGING PRINT print(self.bnl) #DEBUGGING PRINT # Apply correct batch norm to each tensor hf = self.bnh(hf) if isinstance(hf, torch.Tensor) else hf lf = self.bnl(lf) if isinstance(lf, torch.Tensor) else lf # Fixed this line! print("ENDED BN") return hf, lf else: return self.bnh(x)
Additional Notes
- Use
isinstance(hf, torch.Tensor)instead oftype(hf) == torch.Tensor—it's more flexible if you ever use subclassed tensors. - Double-check your
alpha_inparameter: right now you're not using it in the__init__method. Ifalpha_inaffects the input tensor split, you'll need to handle that to ensure the input hf/lf channels match your batch norm channel counts. - When
alpha_out=0,lf_chbecomes 0, soself.bnlwill beBatchNorm2d(0). Make sure your forward pass handles cases where lf isNoneto avoid trying to apply batch norm to a non-existent tensor.
内容的提问来源于stack exchange,提问作者Victor Lundgren

