(self, loss: Tensor)
| 77 | |
| 78 | |
| 79 | def loss_wrapper(self, loss: Tensor) -> Tensor: |
| 80 | # parameter activation: it is a l2 loss with 0 weight |
| 81 | for param in self.parameters(): |
| 82 | loss += 0 * torch.sum(param ** 2) |
| 83 | return loss |
| 84 | |
| 85 | |
| 86 | def train_one_step(self, data_dict: dict) -> dict: |