| 49 | |
| 50 | |
| 51 | def get_normalization_layer(name, dims, num_groups=None, *args, **kwargs): |
| 52 | if not isinstance(name, str) or name.lower() == 'none': |
| 53 | return None |
| 54 | elif 'batch' in name: |
| 55 | return nn.BatchNorm1d(num_features=dims, *args, **kwargs) |
| 56 | elif 'layer' in name: |
| 57 | return nn.LayerNorm(dims, *args, **kwargs) |
| 58 | elif 'inst' in name: |
| 59 | return nn.InstanceNorm1d(num_features=dims, *args, **kwargs) |
| 60 | elif 'group' in name: |
| 61 | if num_groups is None: |
| 62 | num_groups = int(dims / 10) |
| 63 | return nn.GroupNorm(num_groups=num_groups, num_channels=dims) |
| 64 | else: |
| 65 | raise ValueError("Unknown normalization name", name) |
| 66 | |
| 67 | |
| 68 | def identity(X): |