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

Octave卷积适配BatchNorm2d出现running_mean维度不匹配RuntimeError

Fixing the RuntimeError in Custom Octave Convolution BatchNorm2d

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 of type(hf) == torch.Tensor—it's more flexible if you ever use subclassed tensors.
  • Double-check your alpha_in parameter: right now you're not using it in the __init__ method. If alpha_in affects 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_ch becomes 0, so self.bnl will be BatchNorm2d(0). Make sure your forward pass handles cases where lf is None to avoid trying to apply batch norm to a non-existent tensor.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.07 21:47:50