(self, minibatches, unlabeled=None)
| 1124 | self.register_buffer('update_count', torch.tensor([0])) |
| 1125 | |
| 1126 | def update(self, minibatches, unlabeled=None): |
| 1127 | |
| 1128 | mean_loss = 0 |
| 1129 | param_gradients = [[] for _ in self.network.parameters()] |
| 1130 | for i, (x, y) in enumerate(minibatches): |
| 1131 | logits = self.network(x) |
| 1132 | |
| 1133 | env_loss = F.cross_entropy(logits, y) |
| 1134 | mean_loss += env_loss.item() / len(minibatches) |
| 1135 | env_grads = autograd.grad(env_loss, self.network.parameters(), retain_graph=True) |
| 1136 | for grads, env_grad in zip(param_gradients, env_grads): |
| 1137 | grads.append(env_grad) |
| 1138 | |
| 1139 | self.optimizer.zero_grad() |
| 1140 | # gradient masking applied here |
| 1141 | self.mask_grads(param_gradients, self.network.parameters()) |
| 1142 | self.optimizer.step() |
| 1143 | self.update_count += 1 |
| 1144 | |
| 1145 | return {'loss': mean_loss} |
| 1146 | |
| 1147 | def mask_grads(self, gradients, params): |
| 1148 | ''' |
nothing calls this directly
no test coverage detected