(self, minibatches, unlabeled=None)
| 941 | self.tau = hparams["tau"] |
| 942 | |
| 943 | def update(self, minibatches, unlabeled=None): |
| 944 | mean_loss = 0 |
| 945 | param_gradients = [[] for _ in self.network.parameters()] |
| 946 | for i, (x, y) in enumerate(minibatches): |
| 947 | logits = self.network(x) |
| 948 | |
| 949 | env_loss = F.cross_entropy(logits, y) |
| 950 | mean_loss += env_loss.item() / len(minibatches) |
| 951 | |
| 952 | env_grads = autograd.grad(env_loss, self.network.parameters()) |
| 953 | for grads, env_grad in zip(param_gradients, env_grads): |
| 954 | grads.append(env_grad) |
| 955 | |
| 956 | self.optimizer.zero_grad() |
| 957 | self.mask_grads(self.tau, param_gradients, self.network.parameters()) |
| 958 | self.optimizer.step() |
| 959 | |
| 960 | return {'loss': mean_loss} |
| 961 | |
| 962 | def mask_grads(self, tau, gradients, params): |
| 963 |
nothing calls this directly
no test coverage detected