| 780 | } |
| 781 | |
| 782 | void ggml_opt_eval(ggml_opt_context_t opt_ctx, ggml_opt_result_t result) { |
| 783 | GGML_ASSERT(opt_ctx->eval_ready); |
| 784 | if (opt_ctx->allocated_graph == opt_ctx->gb_opt) { |
| 785 | const ggml_opt_optimizer_params & opt_pars = opt_ctx->get_opt_pars(opt_ctx->get_opt_pars_ud); |
| 786 | |
| 787 | switch (opt_ctx->optimizer) { |
| 788 | case GGML_OPT_OPTIMIZER_TYPE_ADAMW: { |
| 789 | GGML_ASSERT(opt_pars.adamw.alpha > 0.0f); |
| 790 | GGML_ASSERT(opt_pars.adamw.beta1 >= 0.0f); |
| 791 | GGML_ASSERT(opt_pars.adamw.beta1 <= 1.0f); |
| 792 | GGML_ASSERT(opt_pars.adamw.beta2 >= 0.0f); |
| 793 | GGML_ASSERT(opt_pars.adamw.beta2 <= 1.0f); |
| 794 | GGML_ASSERT(opt_pars.adamw.eps >= 0.0f); |
| 795 | GGML_ASSERT(opt_pars.adamw.wd >= 0.0f); |
| 796 | GGML_ASSERT(opt_pars.adamw.wd <= 1.0f); |
| 797 | |
| 798 | // beta1, beta2 after applying warmup |
| 799 | const float beta1h = 1.0f / (1.0f - powf(opt_pars.adamw.beta1, opt_ctx->iter)); |
| 800 | const float beta2h = 1.0f / (1.0f - powf(opt_pars.adamw.beta2, opt_ctx->iter)); |
| 801 | |
| 802 | float * adamw_par_data = ggml_get_data_f32(opt_ctx->opt_step_params); |
| 803 | adamw_par_data[0] = opt_pars.adamw.alpha; |
| 804 | adamw_par_data[1] = opt_pars.adamw.beta1; |
| 805 | adamw_par_data[2] = opt_pars.adamw.beta2; |
| 806 | adamw_par_data[3] = opt_pars.adamw.eps; |
| 807 | adamw_par_data[4] = opt_pars.adamw.wd; |
| 808 | adamw_par_data[5] = beta1h; |
| 809 | adamw_par_data[6] = beta2h; |
| 810 | } break; |
| 811 | case GGML_OPT_OPTIMIZER_TYPE_SGD: { |
| 812 | GGML_ASSERT(opt_pars.sgd.alpha > 0.0f); |
| 813 | GGML_ASSERT(opt_pars.sgd.wd >= 0.0f); |
| 814 | GGML_ASSERT(opt_pars.sgd.wd <= 1.0f); |
| 815 | float * sgd = ggml_get_data_f32(opt_ctx->opt_step_params); |
| 816 | sgd[0] = opt_pars.sgd.alpha; |
| 817 | sgd[1] = opt_pars.sgd.wd; |
| 818 | } break; |
| 819 | default: |
| 820 | GGML_ABORT("fatal error"); |
| 821 | } |
| 822 | } |
| 823 | |
| 824 | ggml_backend_sched_graph_compute(opt_ctx->backend_sched, opt_ctx->allocated_graph_copy); |
| 825 | opt_ctx->iter += opt_ctx->allocated_graph == opt_ctx->gb_opt; |
| 826 | opt_ctx->opt_i = (opt_ctx->opt_i + 1) % opt_ctx->opt_period; |
| 827 | |
| 828 | if (!opt_ctx->static_graphs) { |
| 829 | opt_ctx->gf = nullptr; |
| 830 | opt_ctx->gb_grad = nullptr; |
| 831 | opt_ctx->gb_opt = nullptr; |
| 832 | opt_ctx->allocated_graph = nullptr; |
| 833 | opt_ctx->allocated_graph_copy = nullptr; |
| 834 | } |
| 835 | |
| 836 | opt_ctx->eval_ready = false; |
| 837 | |
| 838 | if (!result) { |
| 839 | return; |