torch.manual_seed(21) num_channels = 3 # Example tensor so that we can use randn_like() below. y = torch.randn(20, num_channels, 1, 1) model = nn.BatchNorm2d(num_channels) # nb is a dict containing the buffers (non-trainable parameters) # of the BatchNorm2d layer. Since these are non-trainable # parameters, we don't need to run a backward pass to update # these values. They will be updated during the forward pass itself. nb = dict(model.named_buffers()) print(f"Buffers in BatchNorm2d: {nb.keys()}n") stacked = torch.tensor([]).reshape(0, num_channels, 1, 1) for i in range(2000): x = torch.randn_like(y) y_hat = model(x) # Save all the input tensor into 'stacked' so that # we can compute the mean and variance later. stacked = torch.cat([stacked, x], dim=0) # end for print(f"Shape of stackend tensor: {stacked.shape}n") smean = stacked.mean(dim=(0, 2, 3)) svar = stacked.var(dim=(0, 2, 3)) print(f"Manually Computed:") print(f"------------------") print(f"Mean: {smean}nVariance: {svar}n") print(f"Computed by BatchNorm2d:") print(f"------------------------") rm, rv = nb['running_mean'], nb['running_var'] print(f"Mean: {rm}nVariance: {rv}n") print(f"Mean Absolute Differences:") print(f"--------------------------") print(f"Mean: {(smean-rm).abs().mean():.4f}, Variance: {(svar-rv).abs().mean():.4f}")