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

Class ComplexBatchNorm

complexnn.py:219–397  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

217# from https://github.com/IMLHF/SE_DCUNet/blob/f28bf1661121c8901ad38149ea827693f1830715/models/layers/complexnn.py#L55
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:
262 self.RMr.zero_()
263 self.RMi.zero_()
264 self.RVrr.fill_(1)
265 self.RVri.zero_()
266 self.RVii.fill_(1)
267 self.num_batches_tracked.zero_()
268
269 def reset_parameters(self):
270 self.reset_running_stats()
271 if self.affine:
272 self.Br.data.zero_()
273 self.Bi.data.zero_()
274 self.Wrr.data.fill_(1)
275 self.Wri.data.uniform_(-.9, +.9) # W will be positive-definite
276 self.Wii.data.fill_(1)

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected