MCPcopy Create free account
hub / github.com/RolnickLab/climart / get_normalization_layer

Function get_normalization_layer

climart/utils/utils.py:51–65  ·  view source on GitHub ↗
(name, dims, num_groups=None, *args, **kwargs)

Source from the content-addressed store, hash-verified

49
50
51def 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
68def identity(X):

Callers 4

__init__Method · 0.90
__init__Method · 0.90
__init__Method · 0.90
__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected