Not a member of Pastebin yet?
Sign Up,
it unlocks many cool features!
- class SimpleBatchNorm2d(nn.Module):
- """
- Поведение как у torch.nn.BatchNorm2d для 4D входов (N,C,H,W):
- - train(): нормализуем по biased var, обновляем running_mean и running_var (unbiased)
- - eval(): используем running_* если track_running_stats=True, иначе — батч-статы
- """
- def __init__(
- self, num_features, eps=1e-5, momentum=0.1, affine=True, track_running_stats=True, device=None, dtype=None
- ):
- super().__init__()
- kw = {"device": device, "dtype": dtype}
- self.num_features = int(num_features)
- self.eps = eps
- self.momentum = float(momentum)
- self.affine = bool(affine)
- self.track_running_stats = bool(track_running_stats)
- if self.affine:
- self.weight = nn.Parameter(torch.ones(self.num_features, **kw))
- self.bias = nn.Parameter(torch.zeros(self.num_features, **kw))
- else:
- self.register_parameter("weight", None)
- self.register_parameter("bias", None)
- if self.track_running_stats:
- self.register_buffer("running_mean", torch.zeros(self.num_features, **kw))
- self.register_buffer("running_var", torch.ones(self.num_features, **kw))
- self.register_buffer("num_batches_tracked", torch.tensor(0, dtype=torch.long, device=device))
- else:
- self.register_buffer("running_mean", None)
- self.register_buffer("running_var", None)
- self.register_buffer("num_batches_tracked", None)
- def forward(self, x: torch.Tensor) -> torch.Tensor:
- if x.dim() != 4:
- raise ValueError(f"BatchNorm2d expects 4D input (N,C,H,W), got {tuple(x.shape)}")
- reduce_dims = (0, 2, 3) # N, H, W
- batch_mean = x.mean(dim=reduce_dims)
- batch_var_biased = x.var(dim=reduce_dims, unbiased=False)
- if self.training:
- if self.track_running_stats:
- self.num_batches_tracked = self.num_batches_tracked + 1
- M = 1
- for d in reduce_dims:
- M *= x.size(d)
- correction = float(M) / float(M - 1) if M > 1 else 1.0
- batch_var_unbiased = batch_var_biased * correction
- m = self.momentum
- self.running_mean = (1 - m) * self.running_mean + m * batch_mean.detach()
- self.running_var = (1 - m) * self.running_var + m * batch_var_unbiased.detach()
- mean_for_norm = batch_mean
- var_for_norm = batch_var_biased
- else:
- if self.track_running_stats and self.running_mean is not None:
- mean_for_norm = self.running_mean
- var_for_norm = self.running_var
- else:
- mean_for_norm = batch_mean
- var_for_norm = batch_var_biased
- x_hat = (x - mean_for_norm.view(1, self.num_features, 1, 1)) / torch.sqrt(
- var_for_norm.view(1, self.num_features, 1, 1) + self.eps
- )
- if self.affine:
- x_hat = x_hat * self.weight.view(1, self.num_features, 1, 1) + self.bias.view(1, self.num_features, 1, 1)
- return x_hat
Advertisement
Add Comment
Please, Sign In to add comment