| 754 | } |
| 755 | |
| 756 | void ggml_opt_eval(ggml_opt_context_t opt_ctx, ggml_opt_result_t result) { |
| 757 | GGML_ASSERT(opt_ctx->eval_ready); |
| 758 | if (opt_ctx->allocated_graph == opt_ctx->gb_opt) { |
| 759 | struct ggml_opt_optimizer_params opt_pars = opt_ctx->get_opt_pars(opt_ctx->get_opt_pars_ud); |
| 760 | |
| 761 | GGML_ASSERT(opt_pars.adamw.alpha > 0.0f); |
| 762 | GGML_ASSERT(opt_pars.adamw.beta1 >= 0.0f); |
| 763 | GGML_ASSERT(opt_pars.adamw.beta1 <= 1.0f); |
| 764 | GGML_ASSERT(opt_pars.adamw.beta2 >= 0.0f); |
| 765 | GGML_ASSERT(opt_pars.adamw.beta2 <= 1.0f); |
| 766 | GGML_ASSERT(opt_pars.adamw.eps >= 0.0f); |
| 767 | GGML_ASSERT(opt_pars.adamw.wd >= 0.0f); |
| 768 | GGML_ASSERT(opt_pars.adamw.wd <= 1.0f); |
| 769 | |
| 770 | // beta1, beta2 after applying warmup |
| 771 | const float beta1h = 1.0f/(1.0f - powf(opt_pars.adamw.beta1, opt_ctx->iter)); |
| 772 | const float beta2h = 1.0f/(1.0f - powf(opt_pars.adamw.beta2, opt_ctx->iter)); |
| 773 | |
| 774 | float * adamw_par_data = ggml_get_data_f32(opt_ctx->adamw_params); |
| 775 | adamw_par_data[0] = opt_pars.adamw.alpha; |
| 776 | adamw_par_data[1] = opt_pars.adamw.beta1; |
| 777 | adamw_par_data[2] = opt_pars.adamw.beta2; |
| 778 | adamw_par_data[3] = opt_pars.adamw.eps; |
| 779 | adamw_par_data[4] = opt_pars.adamw.wd; |
| 780 | adamw_par_data[5] = beta1h; |
| 781 | adamw_par_data[6] = beta2h; |
| 782 | } |
| 783 | |
| 784 | ggml_backend_sched_graph_compute(opt_ctx->backend_sched, opt_ctx->allocated_graph_copy); |
| 785 | opt_ctx->iter += opt_ctx->allocated_graph == opt_ctx->gb_opt; |
| 786 | opt_ctx->opt_i = (opt_ctx->opt_i + 1) % opt_ctx->opt_period; |
| 787 | |
| 788 | if (!opt_ctx->static_graphs) { |
| 789 | opt_ctx->gf = nullptr; |
| 790 | opt_ctx->gb_grad = nullptr; |
| 791 | opt_ctx->gb_opt = nullptr; |
| 792 | opt_ctx->allocated_graph = nullptr; |
| 793 | opt_ctx->allocated_graph_copy = nullptr; |
| 794 | } |
| 795 | |
| 796 | opt_ctx->eval_ready = false; |
| 797 | |
| 798 | if (!result) { |
| 799 | return; |
| 800 | } |
| 801 | |
| 802 | if (result->ndata == 0) { |
| 803 | result->loss_per_datapoint = opt_ctx->loss_per_datapoint; |
| 804 | result->opt_period = opt_ctx->opt_period; |
| 805 | } else { |
| 806 | GGML_ASSERT(result->loss_per_datapoint == opt_ctx->loss_per_datapoint); |
| 807 | GGML_ASSERT(result->opt_period == opt_ctx->opt_period); |
| 808 | } |
| 809 | |
| 810 | const int64_t ndata = opt_ctx->outputs->ne[1]; |
| 811 | GGML_ASSERT(result->ndata == ndata*int64_t(result->loss.size()) && "varying batch size not supported"); |
| 812 | result->ndata += ndata; |
| 813 | |