PyTorch实现Batch Normalization逆计算多Batch运行异常问题求助
问题原因分析
- 核心bug:正向传播计算得到的原始输入的批次均值、方差没有存储,逆计算时错误地对正向输出的结果重新计算了新的均值方差,两套统计量完全不一致,导致逆计算无法还原原始输入。
- 单batch测试能通过的原因:代码里针对batch size=1的场景,硬编码了均值为全0、方差为eps,不管正向还是逆计算阶段用的都是相同的固定统计量,所以恰好能还原。
- 多batch场景下的错误流程:
- 正向阶段输入原始x(shape 3,10,100,100),计算得到的m、v是原始x的均值方差
- 调用reverse时输入的是正向输出x1,此时代码重新对x1计算均值方差,这个新的m、v和正向用的完全不同
- 用错误的统计量做逆计算,自然结果和原始x偏差极大。
修复方案
你需要在正向传播计算完均值方差后,把当前的m和v暂存到类的实例变量里,逆计算的时候直接读取暂存的统计量,不要重新计算。修改点如下:
- 在正向传播计算得到m、v之后,添加代码把它们存到
self.current_m和self.current_v里 - reverse函数里删除重新计算m、v的逻辑,直接读取
self.current_m和self.current_v使用 - 验证模式下也要保证正向和逆计算用的是同一套全局统计量。
核心修改后的代码参考:
class BatchNorm(nn.Module): def __init__(self, dim, eps=1e-5): super().__init__() self.eps = eps # 可直接初始化为1,不需要后续做exp运算,更符合BN常规实现 self.gamma = nn.Parameter(torch.ones(1, dim), requires_grad=True) self.beta = nn.Parameter(torch.zeros(1, dim), requires_grad=True) self.batch_mean = None self.batch_var = None # 新增:存储当前步的统计量给逆计算用 self.current_m = None self.current_v = None def forward(self, x, reverse=False): if reverse == True: return self.reverse(x) B, C, W, H = x.shape if self.training: if B>1: m = x.mean(dim=0) v = x.var(dim=0) + self.eps else: # 补充device参数,避免x在GPU上时出现设备不匹配错误 m = torch.zeros(C, W, H, device=x.device) v = torch.zeros(C, W, H, device=x.device) + self.eps self.batch_mean = None else: if self.batch_mean is None: self.set_batch_stats_func(x) m = self.batch_mean.clone() v = self.batch_var.clone() # 暂存当前统计量 self.current_m = m self.current_v = v # 不需要repeat_interleave,直接用广播机制即可适配W、H维度,更省内存 gamma = self.gamma[..., None, None] beta = self.beta[..., None, None] x_hat = (x - m) / torch.sqrt(v) x_hat = x_hat * gamma + beta log_det = torch.sum(torch.log(gamma) - 0.5 * torch.log(v)) return x_hat, log_det def reverse(self, x): B, C, W, H = x.shape # 直接读取正向暂存的统计量,不要重新计算 m = self.current_m v = self.current_v gamma = self.gamma[..., None, None] beta = self.beta[..., None, None] x_hat = (x - beta) / gamma * torch.sqrt(v) + m log_det = torch.sum(-torch.log(gamma) + 0.5 * torch.log(v)) return x_hat, log_det def set_batch_stats_func(self, x): print("setting batch stats for validation") self.batch_mean = x.mean(dim=0) self.batch_var = x.var(dim=0) + self.eps
内容的提问来源于stack exchange,提问作者Nikoo_Ebrahimi
相关产品推荐
相关产品推荐

