(self, num_groups, num_channels)
| 20 | |
| 21 | class BaseGroupNorm(nn.GroupNorm): |
| 22 | def __init__(self, num_groups, num_channels): |
| 23 | super().__init__(num_groups=num_groups, num_channels=num_channels) |
| 24 | |
| 25 | def forward(self, x, zero_pad=False, **kwargs): |
| 26 | if zero_pad: |