MCPcopy Create free account
hub / github.com/0xShug0/audio.cpp / ggml_opt_eval

Function ggml_opt_eval

external/ggml/src/ggml-opt.cpp:782–877  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

780}
781
782void 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;

Callers 5

ggml_opt_epochFunction · 0.85
test_gradFunction · 0.85
test_forward_backwardFunction · 0.85
test_idata_splitFunction · 0.85

Calls 9

ggml_get_data_f32Function · 0.85
ggml_is_scalarFunction · 0.85
ggml_backend_tensor_getFunction · 0.85
ggml_nbytesFunction · 0.85
sizeMethod · 0.45
dataMethod · 0.45
endMethod · 0.45
beginMethod · 0.45

Tested by 4

test_gradFunction · 0.68
test_forward_backwardFunction · 0.68
test_idata_splitFunction · 0.68