(in_channels, gather=False, **kwargs)
| 341 | |
| 342 | |
| 343 | def Normalize(in_channels, gather=False, **kwargs): # same for 3D and 2D |
| 344 | if gather: |
| 345 | return ContextParallelGroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True) |
| 346 | else: |
| 347 | return torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True) |
| 348 | |
| 349 | |
| 350 | class SpatialNorm3D(nn.Module): |
no test coverage detected