| 1371 | } |
| 1372 | |
| 1373 | void train_opt_callback(void * vdata, int accum_step, float * sched, bool * cancel) { |
| 1374 | struct train_opt_callback_data * data = (struct train_opt_callback_data *) vdata; |
| 1375 | struct train_params_common * params = data->params; |
| 1376 | struct train_state * train = data->train; |
| 1377 | struct ggml_opt_context * opt = train->opt; |
| 1378 | int n_batch = params->n_batch; |
| 1379 | int n_ctx = params->n_ctx; |
| 1380 | |
| 1381 | if (accum_step == 0) { |
| 1382 | // time measurement |
| 1383 | int64_t now = ggml_time_ms(); |
| 1384 | if (now > data->last_time && opt->iter > data->first_iter) { |
| 1385 | double dt = (double) (now - data->last_time); |
| 1386 | if (data->millis_per_iter == 0.0) { |
| 1387 | data->millis_per_iter = dt; |
| 1388 | } else { |
| 1389 | const double gain = 0.7; |
| 1390 | data->millis_per_iter = data->millis_per_iter*(1.0-gain) + dt*gain; |
| 1391 | } |
| 1392 | } |
| 1393 | |
| 1394 | double remaining_millis = 0.0; |
| 1395 | if (data->millis_per_iter > 0.0) { |
| 1396 | const int n_iter = params->adam_n_iter; |
| 1397 | const int done_iter = opt->iter - data->first_iter; |
| 1398 | const int remaining_iter = n_iter - done_iter; |
| 1399 | remaining_millis = remaining_iter * data->millis_per_iter; |
| 1400 | } |
| 1401 | |
| 1402 | // file saving |
| 1403 | const bool save_now = (params->save_every > 0) && (opt->iter - data->last_save_iter >= params->save_every); |
| 1404 | if (save_now) { |
| 1405 | int new_iters = opt->iter - data->last_save_iter; |
| 1406 | train->train_its += new_iters; |
| 1407 | train->train_tokens += new_iters * opt->params.n_gradient_accumulation * n_batch * n_ctx; |
| 1408 | |
| 1409 | if (data->save_cb) { |
| 1410 | data->save_cb(data->save_data, train); |
| 1411 | } |
| 1412 | |
| 1413 | data->last_save_iter = opt->iter; |
| 1414 | } |
| 1415 | |
| 1416 | // exclude file saving from time measurement, by measuring last_time after saving |
| 1417 | data->last_time = ggml_time_ms(); |
| 1418 | |
| 1419 | *sched = learning_schedule( |
| 1420 | opt->iter, |
| 1421 | params->warmup, |
| 1422 | params->cos_decay_steps, |
| 1423 | params->adam_alpha, |
| 1424 | params->adam_min_alpha, |
| 1425 | params->cos_decay_min, |
| 1426 | params->cos_decay_restart, |
| 1427 | params->enable_restart); |
| 1428 | |
| 1429 | int impr_plot = -(int)(1 + (opt->loss_before - opt->loss_after) * 10.0f + 0.5f); |
| 1430 | if (impr_plot > 0) impr_plot = 0; |
nothing calls this directly
no test coverage detected