| 122 | return lr |
| 123 | |
| 124 | def apply_gradients(self, dummy_overflow_buf, lr_scheduled, per_param_decay, grad_list, param_list, momentum, velocity, update): |
| 125 | # Compute global gradient norm |
| 126 | global_grad_norm = multi_tensor_applier( |
| 127 | multi_tensor_l2norm, |
| 128 | dummy_overflow_buf, |
| 129 | [grad_list], |
| 130 | False)[0].item() |
| 131 | |
| 132 | # Compute per parameter norm |
| 133 | param_norms = multi_tensor_applier( |
| 134 | multi_tensor_l2norm, |
| 135 | dummy_overflow_buf, |
| 136 | [param_list], |
| 137 | True)[1] |
| 138 | |
| 139 | # Compute LAMB update |
| 140 | multi_tensor_applier( |
| 141 | lamb_compute_update, |
| 142 | dummy_overflow_buf, |
| 143 | [grad_list, param_list, momentum, velocity, update], |
| 144 | torch.cuda.FloatTensor(per_param_decay), |
| 145 | self.step_count, |
| 146 | self.b1, |
| 147 | self.b2, |
| 148 | self.epsilon, |
| 149 | global_grad_norm, |
| 150 | self.max_global_grad_norm, |
| 151 | ) |
| 152 | |
| 153 | # Computer per parameter update norm |
| 154 | update_norms = multi_tensor_applier( |
| 155 | multi_tensor_l2norm, |
| 156 | dummy_overflow_buf, |
| 157 | [update], |
| 158 | True)[1] |
| 159 | |
| 160 | # Apply LAMB update on parameters |
| 161 | multi_tensor_applier( |
| 162 | lamb_apply_update, |
| 163 | dummy_overflow_buf, |
| 164 | [param_list, update], |
| 165 | param_norms, |
| 166 | update_norms, |
| 167 | lr_scheduled, |
| 168 | ) |
| 169 | |
| 170 | def step(self, closure=None): |
| 171 | """Performs a single optimization step. |