(self, x, y, dist_option, spars)
| 50 | return y |
| 51 | |
| 52 | def train_one_batch(self, x, y, dist_option, spars): |
| 53 | out = self.forward(x) |
| 54 | loss = self.softmax_cross_entropy(out, y) |
| 55 | |
| 56 | if dist_option == 'plain': |
| 57 | self.optimizer(loss) |
| 58 | elif dist_option == 'half': |
| 59 | self.optimizer.backward_and_update_half(loss) |
| 60 | elif dist_option == 'partialUpdate': |
| 61 | self.optimizer.backward_and_partial_update(loss) |
| 62 | elif dist_option == 'sparseTopK': |
| 63 | self.optimizer.backward_and_sparse_update(loss, |
| 64 | topK=True, |
| 65 | spars=spars) |
| 66 | elif dist_option == 'sparseThreshold': |
| 67 | self.optimizer.backward_and_sparse_update(loss, |
| 68 | topK=False, |
| 69 | spars=spars) |
| 70 | return out, loss |
| 71 | |
| 72 | def set_optimizer(self, optimizer): |
| 73 | self.optimizer = optimizer |
no test coverage detected