| 661 | } |
| 662 | |
| 663 | bool load_train_state_gguf(struct gguf_context * fctx, struct ggml_context * f_ggml_ctx, struct train_state * train) { |
| 664 | if (gguf_find_key(fctx, LLM_KV_TRAINING_FILE_VERSION) < 0) { |
| 665 | return false; |
| 666 | } |
| 667 | |
| 668 | uint32_t file_version; |
| 669 | GGUF_GET_KEY(fctx, file_version, gguf_get_val_u32, GGUF_TYPE_UINT32, true, LLM_KV_TRAINING_FILE_VERSION); |
| 670 | GGML_ASSERT(file_version <= 1); |
| 671 | |
| 672 | if (file_version == 0) { |
| 673 | |
| 674 | GGUF_GET_KEY(fctx, train->train_its, gguf_get_val_u32, GGUF_TYPE_UINT32, true, LLM_KV_TRAINING_ITERATION_COUNT); |
| 675 | GGUF_GET_KEY(fctx, train->train_samples, gguf_get_val_u32, GGUF_TYPE_UINT32, true, LLM_KV_TRAINING_SAMPLE_COUNT); |
| 676 | GGUF_GET_KEY(fctx, train->train_tokens, gguf_get_val_u32, GGUF_TYPE_UINT32, true, LLM_KV_TRAINING_TOKEN_COUNT); |
| 677 | |
| 678 | } else if (file_version == 1) { |
| 679 | |
| 680 | GGUF_GET_KEY(fctx, train->train_its, gguf_get_val_u64, GGUF_TYPE_UINT64, true, LLM_KV_TRAINING_ITERATION_COUNT); |
| 681 | GGUF_GET_KEY(fctx, train->train_samples, gguf_get_val_u64, GGUF_TYPE_UINT64, true, LLM_KV_TRAINING_SAMPLE_COUNT); |
| 682 | GGUF_GET_KEY(fctx, train->train_tokens, gguf_get_val_u64, GGUF_TYPE_UINT64, true, LLM_KV_TRAINING_TOKEN_COUNT); |
| 683 | GGUF_GET_KEY(fctx, train->train_epochs, gguf_get_val_u64, GGUF_TYPE_UINT64, true, LLM_KV_TRAINING_EPOCH_COUNT); |
| 684 | |
| 685 | GGUF_GET_KEY(fctx, train->shuffle_samples_hash, gguf_get_val_u64, GGUF_TYPE_UINT64, false, LLM_KV_TRAINING_SHUFFLE_SAMPLES_HASH); |
| 686 | GGUF_GET_KEY(fctx, train->shuffle_rng_state_current, gguf_get_val_str, GGUF_TYPE_STRING, false, LLM_KV_TRAINING_SHUFFLE_RNG_STATE); |
| 687 | GGUF_GET_KEY(fctx, train->shuffle_sample_count, gguf_get_val_u64, GGUF_TYPE_UINT64, false, LLM_KV_TRAINING_SHUFFLE_SAMPLE_COUNT); |
| 688 | GGUF_GET_KEY(fctx, train->shuffle_next_sample, gguf_get_val_u64, GGUF_TYPE_UINT64, false, LLM_KV_TRAINING_SHUFFLE_NEXT_SAMPLE); |
| 689 | } |
| 690 | |
| 691 | load_opt_context_gguf(fctx, f_ggml_ctx, train->opt); |
| 692 | return true; |
| 693 | } |
| 694 | |
| 695 | void save_train_state_gguf(struct gguf_context * fctx, struct train_state * train) { |
| 696 | gguf_set_val_u32(fctx, LLM_KV_TRAINING_FILE_VERSION, 1); |
no test coverage detected