35 lines
926 B
C++
35 lines
926 B
C++
#pragma once
|
|
#include "model.hpp"
|
|
#include <string>
|
|
#include <unordered_map>
|
|
#include <vector>
|
|
|
|
namespace xt {
|
|
|
|
struct Tokenizer {
|
|
std::vector<std::string> id2tok;
|
|
|
|
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, bool debug = false);
|
|
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
|