vadimk772336

Untitled

Aug 13th, 2025
245
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
Python 3.21 KB | None | 0 0
  1.  
  2. class SimpleBatchNorm2d(nn.Module):
  3.     """
  4.    Поведение как у torch.nn.BatchNorm2d для 4D входов (N,C,H,W):
  5.    - train(): нормализуем по biased var, обновляем running_mean и running_var (unbiased)
  6.    - eval(): используем running_* если track_running_stats=True, иначе — батч-статы
  7.    """
  8.  
  9.     def __init__(
  10.         self, num_features, eps=1e-5, momentum=0.1, affine=True, track_running_stats=True, device=None, dtype=None
  11.     ):
  12.         super().__init__()
  13.         kw = {"device": device, "dtype": dtype}
  14.         self.num_features = int(num_features)
  15.         self.eps = eps
  16.         self.momentum = float(momentum)
  17.         self.affine = bool(affine)
  18.         self.track_running_stats = bool(track_running_stats)
  19.  
  20.         if self.affine:
  21.             self.weight = nn.Parameter(torch.ones(self.num_features, **kw))
  22.             self.bias = nn.Parameter(torch.zeros(self.num_features, **kw))
  23.         else:
  24.             self.register_parameter("weight", None)
  25.             self.register_parameter("bias", None)
  26.  
  27.         if self.track_running_stats:
  28.             self.register_buffer("running_mean", torch.zeros(self.num_features, **kw))
  29.             self.register_buffer("running_var", torch.ones(self.num_features, **kw))
  30.             self.register_buffer("num_batches_tracked", torch.tensor(0, dtype=torch.long, device=device))
  31.         else:
  32.             self.register_buffer("running_mean", None)
  33.             self.register_buffer("running_var", None)
  34.             self.register_buffer("num_batches_tracked", None)
  35.  
  36.     def forward(self, x: torch.Tensor) -> torch.Tensor:
  37.         if x.dim() != 4:
  38.             raise ValueError(f"BatchNorm2d expects 4D input (N,C,H,W), got {tuple(x.shape)}")
  39.  
  40.         reduce_dims = (0, 2, 3)  # N, H, W
  41.         batch_mean = x.mean(dim=reduce_dims)
  42.         batch_var_biased = x.var(dim=reduce_dims, unbiased=False)
  43.  
  44.         if self.training:
  45.             if self.track_running_stats:
  46.                 self.num_batches_tracked = self.num_batches_tracked + 1
  47.                 M = 1
  48.                 for d in reduce_dims:
  49.                     M *= x.size(d)
  50.                 correction = float(M) / float(M - 1) if M > 1 else 1.0
  51.                 batch_var_unbiased = batch_var_biased * correction
  52.                 m = self.momentum
  53.                 self.running_mean = (1 - m) * self.running_mean + m * batch_mean.detach()
  54.                 self.running_var = (1 - m) * self.running_var + m * batch_var_unbiased.detach()
  55.             mean_for_norm = batch_mean
  56.             var_for_norm = batch_var_biased
  57.         else:
  58.             if self.track_running_stats and self.running_mean is not None:
  59.                 mean_for_norm = self.running_mean
  60.                 var_for_norm = self.running_var
  61.             else:
  62.                 mean_for_norm = batch_mean
  63.                 var_for_norm = batch_var_biased
  64.  
  65.         x_hat = (x - mean_for_norm.view(1, self.num_features, 1, 1)) / torch.sqrt(
  66.             var_for_norm.view(1, self.num_features, 1, 1) + self.eps
  67.         )
  68.         if self.affine:
  69.             x_hat = x_hat * self.weight.view(1, self.num_features, 1, 1) + self.bias.view(1, self.num_features, 1, 1)
  70.         return x_hat
  71.  
Advertisement
Add Comment
Please, Sign In to add comment