MCPcopy Create free account
hub / github.com/dangf15/THLNet / __init__

Method __init__

complexnn.py:220–258  ·  view source on GitHub ↗
(self, num_features, eps=1e-5, momentum=0.1, affine=True,
            track_running_stats=True, complex_axis=1)

Source from the content-addressed store, hash-verified

218
219class ComplexBatchNorm(torch.nn.Module):
220 def __init__(self, num_features, eps=1e-5, momentum=0.1, affine=True,
221 track_running_stats=True, complex_axis=1):
222 super(ComplexBatchNorm, self).__init__()
223 self.num_features = num_features//2
224 self.eps = eps
225 self.momentum = momentum
226 self.affine = affine
227 self.track_running_stats = track_running_stats
228
229 self.complex_axis = complex_axis
230
231 if self.affine:
232 self.Wrr = torch.nn.Parameter(torch.Tensor(self.num_features))
233 self.Wri = torch.nn.Parameter(torch.Tensor(self.num_features))
234 self.Wii = torch.nn.Parameter(torch.Tensor(self.num_features))
235 self.Br = torch.nn.Parameter(torch.Tensor(self.num_features))
236 self.Bi = torch.nn.Parameter(torch.Tensor(self.num_features))
237 else:
238 self.register_parameter('Wrr', None)
239 self.register_parameter('Wri', None)
240 self.register_parameter('Wii', None)
241 self.register_parameter('Br', None)
242 self.register_parameter('Bi', None)
243
244 if self.track_running_stats:
245 self.register_buffer('RMr', torch.zeros(self.num_features))
246 self.register_buffer('RMi', torch.zeros(self.num_features))
247 self.register_buffer('RVrr', torch.ones (self.num_features))
248 self.register_buffer('RVri', torch.zeros(self.num_features))
249 self.register_buffer('RVii', torch.ones (self.num_features))
250 self.register_buffer('num_batches_tracked', torch.tensor(0, dtype=torch.long))
251 else:
252 self.register_parameter('RMr', None)
253 self.register_parameter('RMi', None)
254 self.register_parameter('RVrr', None)
255 self.register_parameter('RVri', None)
256 self.register_parameter('RVii', None)
257 self.register_parameter('num_batches_tracked', None)
258 self.reset_parameters()
259
260 def reset_running_stats(self):
261 if self.track_running_stats:

Callers

nothing calls this directly

Calls 2

reset_parametersMethod · 0.95
__init__Method · 0.45

Tested by

no test coverage detected