38 lines
1.1 KiB
C++
38 lines
1.1 KiB
C++
// checkpoint.h — сохранение и загрузка модели
|
||
#pragma once
|
||
#include "model.h"
|
||
#include <string>
|
||
#include <unordered_map>
|
||
#include <vector>
|
||
|
||
namespace xt {
|
||
|
||
// Токенизатор: словарь хранится рядом с весами в том же файле
|
||
struct Tokenizer {
|
||
std::vector<std::string> id2tok; // id -> строка
|
||
|
||
int bos = 1, eos = 2, unk = 0;
|
||
int size() const { return (int)id2tok.size(); }
|
||
|
||
void build_from_text(const std::string& text, int max_vocab);
|
||
std::vector<int> encode(const std::string& text, bool add_bos = true) const;
|
||
std::string decode(const std::vector<int>& ids, bool skip_bos = true) const;
|
||
std::string decode_one(int id) const;
|
||
};
|
||
|
||
struct Checkpoint {
|
||
Params params;
|
||
Tokenizer tok;
|
||
int step = 0;
|
||
float loss = 0;
|
||
uint64_t seed = 1337;
|
||
};
|
||
|
||
bool save_checkpoint(const std::string& path, const Checkpoint& ck, std::string& err);
|
||
bool load_checkpoint(const std::string& path, Checkpoint& ck, std::string& err);
|
||
|
||
// Читаемый дамп конфигурации
|
||
void print_model_info(const Checkpoint& ck);
|
||
|
||
} // namespace xt
|