| 18943 | } |
| 18944 | |
| 18945 | GGML_API void ggml_opt_init( |
| 18946 | struct ggml_context * ctx, |
| 18947 | struct ggml_opt_context * opt, |
| 18948 | struct ggml_opt_params params, |
| 18949 | int64_t nx) { |
| 18950 | opt->ctx = ctx; |
| 18951 | opt->params = params; |
| 18952 | opt->iter = 0; |
| 18953 | opt->nx = nx; |
| 18954 | opt->just_initialized = true; |
| 18955 | if (opt->ctx == NULL) { |
| 18956 | struct ggml_init_params ctx_opt_params; |
| 18957 | if (opt->params.type == GGML_OPT_ADAM) { |
| 18958 | ctx_opt_params.mem_size = GGML_MEM_ALIGN*3 + ggml_tensor_overhead()*3 + ggml_type_size(GGML_TYPE_F32)*nx*3; |
| 18959 | if (opt->params.past > 0) { |
| 18960 | ctx_opt_params.mem_size += GGML_MEM_ALIGN + ggml_tensor_overhead() + ggml_type_size(GGML_TYPE_F32)*opt->params.past; |
| 18961 | } |
| 18962 | } else if (opt->params.type == GGML_OPT_LBFGS) { |
| 18963 | ctx_opt_params.mem_size = GGML_MEM_ALIGN*9 + ggml_tensor_overhead()*9 + ggml_type_size(GGML_TYPE_F32)*(nx*5 + opt->params.lbfgs.m*2 + nx*opt->params.lbfgs.m*2); |
| 18964 | if (opt->params.past > 0) { |
| 18965 | ctx_opt_params.mem_size += GGML_MEM_ALIGN + ggml_tensor_overhead() + ggml_type_size(GGML_TYPE_F32)*opt->params.past; |
| 18966 | } |
| 18967 | } |
| 18968 | ctx_opt_params.mem_buffer = NULL; |
| 18969 | ctx_opt_params.no_alloc = false; |
| 18970 | |
| 18971 | opt->ctx = ggml_init(ctx_opt_params); |
| 18972 | } |
| 18973 | switch (opt->params.type) { |
| 18974 | case GGML_OPT_ADAM: |
| 18975 | { |
| 18976 | opt->adam.g = ggml_new_tensor_1d(opt->ctx, GGML_TYPE_F32, nx); |
| 18977 | opt->adam.m = ggml_new_tensor_1d(opt->ctx, GGML_TYPE_F32, nx); |
| 18978 | opt->adam.v = ggml_new_tensor_1d(opt->ctx, GGML_TYPE_F32, nx); |
| 18979 | opt->adam.pf = params.past > 0 |
| 18980 | ? ggml_new_tensor_1d(opt->ctx, GGML_TYPE_F32, params.past) |
| 18981 | : NULL; |
| 18982 | ggml_set_zero(opt->adam.m); |
| 18983 | ggml_set_zero(opt->adam.v); |
| 18984 | if (opt->adam.pf) { |
| 18985 | ggml_set_zero(opt->adam.pf); |
| 18986 | } |
| 18987 | } break; |
| 18988 | case GGML_OPT_LBFGS: |
| 18989 | { |
| 18990 | opt->lbfgs.x = ggml_new_tensor_1d(opt->ctx, GGML_TYPE_F32, nx); |
| 18991 | opt->lbfgs.xp = ggml_new_tensor_1d(opt->ctx, GGML_TYPE_F32, nx); |
| 18992 | opt->lbfgs.g = ggml_new_tensor_1d(opt->ctx, GGML_TYPE_F32, nx); |
| 18993 | opt->lbfgs.gp = ggml_new_tensor_1d(opt->ctx, GGML_TYPE_F32, nx); |
| 18994 | opt->lbfgs.d = ggml_new_tensor_1d(opt->ctx, GGML_TYPE_F32, nx); |
| 18995 | opt->lbfgs.pf = params.past > 0 |
| 18996 | ? ggml_new_tensor_1d(opt->ctx, GGML_TYPE_F32, params.past) |
| 18997 | : NULL; |
| 18998 | opt->lbfgs.lmal = ggml_new_tensor_1d(opt->ctx, GGML_TYPE_F32, params.lbfgs.m); |
| 18999 | opt->lbfgs.lmys = ggml_new_tensor_1d(opt->ctx, GGML_TYPE_F32, params.lbfgs.m); |
| 19000 | opt->lbfgs.lms = ggml_new_tensor_2d(opt->ctx, GGML_TYPE_F32, nx, params.lbfgs.m); |
| 19001 | opt->lbfgs.lmy = ggml_new_tensor_2d(opt->ctx, GGML_TYPE_F32, nx, params.lbfgs.m); |
| 19002 | ggml_set_zero(opt->lbfgs.x); |
no test coverage detected