| 233 | } |
| 234 | |
| 235 | int ds_adam_step(int optimizer_id, |
| 236 | size_t step, |
| 237 | float lr, |
| 238 | float beta1, |
| 239 | float beta2, |
| 240 | float epsilon, |
| 241 | float weight_decay, |
| 242 | bool bias_correction, |
| 243 | torch::Tensor& params, |
| 244 | torch::Tensor& grads, |
| 245 | torch::Tensor& exp_avg, |
| 246 | torch::Tensor& exp_avg_sq) |
| 247 | { |
| 248 | auto params_c = params.contiguous(); |
| 249 | auto grads_c = grads.contiguous(); |
| 250 | auto exp_avg_c = exp_avg.contiguous(); |
| 251 | auto exp_avg_sq_c = exp_avg_sq.contiguous(); |
| 252 | |
| 253 | std::shared_ptr<Adam_Optimizer> opt = |
| 254 | std::static_pointer_cast<Adam_Optimizer>(s_optimizers[optimizer_id]); |
| 255 | opt->IncrementStep(step, beta1, beta2); |
| 256 | opt->update_state(lr, epsilon, weight_decay, bias_correction); |
| 257 | |
| 258 | invoke(opt, params_c, grads_c, exp_avg_c, exp_avg_sq_c, params_c.numel()); |
| 259 | |
| 260 | return 0; |
| 261 | } |
| 262 | |
| 263 | void adamw_rollback_inplace(float* params, |
| 264 | const float* grads, |
no test coverage detected