(in_channels, gather=False, **kwargs)
| 442 | |
| 443 | |
| 444 | def Normalize(in_channels, gather=False, **kwargs): # same for 3D and 2D |
| 445 | if gather: |
| 446 | return ContextParallelGroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True) |
| 447 | else: |
| 448 | return torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True) |
| 449 | |
| 450 | |
| 451 | class SpatialNorm3D(nn.Module): |
no test coverage detected