MCPcopy Create free account
hub / github.com/LBANN/lbann / BatchNormModule

Class BatchNormModule

applications/selfsupervised/modules.py:5–49  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

3import resnet
4
5class BatchNormModule(lbann.modules.Module):
6
7 global_count = 0 # Static counter, used for default names
8
9 def __init__(self,
10 statistics_group_size=1,
11 name=None,
12 data_layout='data_parallel'):
13 super().__init__()
14 BatchNormModule.global_count += 1
15 self.instance = 0
16 self.statistics_group_size = statistics_group_size
17 self.name = (name
18 if name
19 else 'bnmodule{0}'.format(BatchNormModule.global_count))
20 self.data_layout = data_layout
21
22 # Initialize weights
23 self.scale = lbann.Weights(
24 initializer=lbann.ConstantInitializer(value=1.0),
25 name=self.name + '_scale')
26 self.bias = lbann.Weights(
27 initializer=lbann.ConstantInitializer(value=0.0),
28 name=self.name + '_bias')
29 self.running_mean = lbann.Weights(
30 initializer=lbann.ConstantInitializer(value=0.0),
31 name=self.name + '_running_mean')
32 self.running_variance = lbann.Weights(
33 initializer=lbann.ConstantInitializer(value=1.0),
34 name=self.name + '_running_variance')
35
36 def forward(self, x):
37 self.instance += 1
38 name = '{0}_instance{1}'.format(self.name, self.instance)
39 return lbann.BatchNormalization(
40 x,
41 weights=[self.scale, self.bias,
42 self.running_mean, self.running_variance],
43 decay=0.9,
44 scale_init=1.0,
45 bias_init=0.0,
46 epsilon=1e-5,
47 statistics_group_size=self.statistics_group_size,
48 name=name,
49 data_layout=self.data_layout)
50
51class ConvBnRelu(lbann.modules.Module):
52

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected