Performs a single optimization step. Args: param(Tensor): param values to be update grad(Tensor): param gradients
(self, param, grad)
| 736 | self.backward_and_update(loss) |
| 737 | |
| 738 | def update(self, param, grad): |
| 739 | """Performs a single optimization step. |
| 740 | |
| 741 | Args: |
| 742 | param(Tensor): param values to be update |
| 743 | grad(Tensor): param gradients |
| 744 | """ |
| 745 | grad /= self.world_size |
| 746 | self.opt.update(param, grad) |
| 747 | |
| 748 | def all_reduce(self, tensor): |
| 749 | """Performs all reduce of a tensor for distributed training. |
no outgoing calls