(self, x, y, dist_option, spars)
| 45 | return y |
| 46 | |
| 47 | def train_one_batch(self, x, y, dist_option, spars): |
| 48 | out = self.forward(x) |
| 49 | loss = self.softmax_cross_entropy(out, y) |
| 50 | |
| 51 | if dist_option == "plain": |
| 52 | self.optimizer(loss) |
| 53 | elif dist_option == "half": |
| 54 | self.optimizer.backward_and_update_half(loss) |
| 55 | elif dist_option == "partialUpdate": |
| 56 | self.optimizer.backward_and_partial_update(loss) |
| 57 | elif dist_option == "sparseTopK": |
| 58 | self.optimizer.backward_and_sparse_update(loss, topK=True, spars=spars) |
| 59 | elif dist_option == "sparseThreshold": |
| 60 | self.optimizer.backward_and_sparse_update(loss, topK=False, spars=spars) |
| 61 | return out, loss |
| 62 | |
| 63 | def set_optimizer(self, optimizer): |
| 64 | self.optimizer = optimizer |
nothing calls this directly
no test coverage detected