| 11 | |
| 12 | |
| 13 | class 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 | |
| 39 | def Activation(activation=None, size=None, dim=-1, inplace=False): |
nothing calls this directly
no outgoing calls
no test coverage detected