MCPcopy Create free account
hub / github.com/HazyResearch/spacetime / OurModule

Class OurModule

model/components.py:13–36  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

11
12
13class OurModule(nn.Module):
14 def __init__(self):
15 super().__init__()
16
17 def register(self, name, tensor, trainable=False, lr=None, wd=None):
18 """Utility method: register a tensor as a buffer or trainable parameter"""
19 if trainable:
20 try:
21 self.register_parameter(name, nn.Parameter(tensor))
22 except KeyError:
23 delattr(self, name)
24 self.register_parameter(name, nn.Parameter(tensor))
25 else:
26
27 try:
28 self.register_buffer(name, tensor)
29 except KeyError:
30 delattr(self, name)
31 self.register_buffer(name, tensor)
32
33 optim = {}
34 if trainable and lr is not None: optim["lr"] = lr
35 if trainable and wd is not None: optim["weight_decay"] = wd
36 if len(optim) > 0: setattr(getattr(self, name), "_optim", optim)
37
38
39def Activation(activation=None, size=None, dim=-1, inplace=False):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected