Xenith/xenith/core/checkpoint.cpp
2026-09-29 19:55:06 +07:00

355 lines
13 KiB
C++
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

// checkpoint.cpp — бинарный формат XNH1, токенизатор уровня байт/слова
#include "checkpoint.h"
#include <cstdio>
#include <cstring>
#include <fstream>
#include <sstream>
#include <unordered_map>
#include <algorithm>
#include <cctype>
namespace xt {
// ------------------------------------------------------- токенизатор
// Уровень 1: частое слова целиком. Уровень 2: отдельные байты.
// Так словарь остаётся маленьким, а любой текст кодируется без потерь.
namespace {
// Разбивает UTF-8 строку на "символы" (кодпоинты)
std::vector<std::string> utf8_chars(const std::string& s) {
std::vector<std::string> out;
for (size_t i = 0; i < s.size();) {
unsigned char c = (unsigned char)s[i];
size_t len = 1;
if ((c & 0xF8) == 0xF0) len = 4;
else if ((c & 0xF0) == 0xE0) len = 3;
else if ((c & 0xE0) == 0xC0) len = 2;
if (i + len > s.size()) len = 1;
out.push_back(s.substr(i, len));
i += len;
}
return out;
}
bool is_space(const std::string& t) {
return t.size() == 1 && (t[0] == ' ' || t[0] == '\t' || t[0] == '\n' || t[0] == '\r');
}
} // namespace
void Tokenizer::build_from_text(const std::string& text, int max_vocab) {
id2tok.clear();
id2tok.push_back("<unk>"); // 0
id2tok.push_back("<bos>"); // 1
id2tok.push_back("<eos>"); // 2
id2tok.push_back("\n"); // 3
// Считаем частоты слов (с сохранением регистра)
std::unordered_map<std::string, int> freq;
std::string word;
auto flush = [&]() {
if (!word.empty()) { freq[word]++; word.clear(); }
};
for (const std::string& ch : utf8_chars(text)) {
if (is_space(ch)) flush();
else word += ch;
}
flush();
std::vector<std::pair<std::string, int>> items(freq.begin(), freq.end());
std::sort(items.begin(), items.end(),
[](const std::pair<std::string, int>& a, const std::pair<std::string, int>& b) {
if (a.second != b.second) return a.second > b.second;
return a.first < b.first; // детерминированный порядок при равенстве
});
const int room = max_vocab - (int)id2tok.size();
for (int i = 0; i < (int)items.size() && i < room; ++i) id2tok.push_back(items[i].first);
// Добираем оставшиеся слоты одиночными байтами, чтобы покрыть весь алфавит
std::vector<std::string> chars = utf8_chars(text);
std::sort(chars.begin(), chars.end());
chars.erase(std::unique(chars.begin(), chars.end()), chars.end());
for (const std::string& c : chars) {
if ((int)id2tok.size() >= max_vocab) break;
if (std::find(id2tok.begin(), id2tok.end(), c) == id2tok.end())
id2tok.push_back(c);
}
}
std::vector<int> Tokenizer::encode(const std::string& text, bool add_bos) const {
// Карта: слово -> id. Строится лениво, один раз на вызов.
std::unordered_map<std::string, int> word2id;
std::unordered_map<std::string, int> char2id;
const int base = 4; // первые 4 — служебные
for (int i = base; i < size(); ++i) {
const std::string& t = id2tok[i];
if (t.size() == 1) char2id[t] = i;
else word2id[t] = i;
}
std::vector<int> out;
if (add_bos) out.push_back(bos);
std::string word;
auto emit_word = [&]() {
if (word.empty()) return;
auto it = word2id.find(word);
if (it != word2id.end()) {
out.push_back(it->second);
} else {
// разбираем неизвестное слово на символы
for (const std::string& c : utf8_chars(word)) {
auto ic = char2id.find(c);
if (ic != char2id.end()) out.push_back(ic->second);
else {
auto ib = char2id.find(" ");
out.push_back(ib != char2id.end() ? ib->second : unk);
}
}
}
word.clear();
};
for (const std::string& ch : utf8_chars(text)) {
if (is_space(ch)) {
emit_word();
if (ch == "\n") {
auto it = char2id.find("\n");
if (it != char2id.end()) out.push_back(it->second);
else out.push_back(3);
}
// обычные пробелы не кодируем — слова склеиваются по границам словаря
} else {
word += ch;
}
}
emit_word();
out.push_back(eos);
return out;
}
std::string Tokenizer::decode_one(int id) const {
if (id < 0 || id >= size()) return "";
const std::string& t = id2tok[id];
if (id == 0) return "";
if (id == 1) return "";
if (id == 2) return "\n";
if (id == 3) return "\n";
return t;
}
std::string Tokenizer::decode(const std::vector<int>& ids, bool skip_bos) const {
std::string out;
for (int id : ids) {
if (skip_bos && (id == bos || id == unk)) continue;
if (id == eos) break;
out += decode_one(id);
if (id >= 4 && id2tok[id].size() > 1) out += " ";
}
return out;
}
// ------------------------------------------------------- ввод-вывод
namespace {
template <typename T>
bool wr(std::ofstream& f, const T& v) {
f.write(reinterpret_cast<const char*>(&v), sizeof(T));
return (bool)f;
}
template <typename T>
bool rd(std::ifstream& f, T& v) {
f.read(reinterpret_cast<char*>(&v), sizeof(T));
return (bool)f;
}
} // namespace
bool save_checkpoint(const std::string& path, const Checkpoint& ck, std::string& err) {
std::ofstream f(path, std::ios::binary);
if (!f) { err = "не удалось открыть для записи: " + path; return false; }
if (!wr(f, XNH_MAGIC)) { err = "ошибка записи магии"; return false; }
if (!wr(f, (uint32_t)1)) { err = "ошибка записи версии"; return false; }
const Config& c = ck.params.cfg;
wr(f, (int32_t)c.vocab_size);
wr(f, (int32_t)c.n_layer);
wr(f, (int32_t)c.n_head);
wr(f, (int32_t)c.n_embd);
wr(f, (int32_t)c.block_size);
wr(f, (int32_t)c.ffn_dim);
wr(f, (int32_t)c.rope_base);
wr(f, c.rms_eps);
wr(f, c.init_std);
wr(f, (int32_t)c.n_vocab_out);
wr(f, (int32_t)ck.step);
wr(f, ck.loss);
wr(f, ck.seed);
wr(f, (int32_t)ck.tok.bos);
wr(f, (int32_t)ck.tok.eos);
wr(f, (int32_t)ck.tok.unk);
// словарь
const int nv = ck.tok.size();
wr(f, (int32_t)nv);
for (int i = 0; i < nv; ++i) {
const uint32_t len = (uint32_t)ck.tok.id2tok[i].size();
wr(f, len);
f.write(ck.tok.id2tok[i].data(), len);
}
// веса
for (const Tensor* t : ck.params.all()) {
wr(f, (int32_t)t->R());
wr(f, (int32_t)t->C());
f.write(reinterpret_cast<const char*>(t->ptr()), (std::streamsize)t->n() * sizeof(float));
}
// оптимизатор: только если просили
if (!ck.params.m.empty()) {
wr(f, (uint32_t)1);
for (size_t i = 0; i < ck.params.m.size(); ++i) {
wr(f, (int32_t)ck.params.m[i].R());
wr(f, (int32_t)ck.params.m[i].C());
f.write(reinterpret_cast<const char*>(ck.params.m[i].ptr()),
(std::streamsize)ck.params.m[i].n() * sizeof(float));
wr(f, (int32_t)ck.params.v[i].R());
wr(f, (int32_t)ck.params.v[i].C());
f.write(reinterpret_cast<const char*>(ck.params.v[i].ptr()),
(std::streamsize)ck.params.v[i].n() * sizeof(float));
}
} else {
wr(f, (uint32_t)0);
}
f.flush();
if (!f) { err = "ошибка при сбросе на диск"; return false; }
return true;
}
bool load_checkpoint(const std::string& path, Checkpoint& ck, std::string& err) {
std::ifstream f(path, std::ios::binary);
if (!f) { err = "не удалось открыть: " + path; return false; }
uint32_t magic = 0, ver = 0;
if (!rd(f, magic) || !rd(f, ver)) { err = "битый заголовок"; return false; }
if (magic != XNH_MAGIC) {
char buf[64];
std::snprintf(buf, sizeof buf,
"плохая магия 0x%08X (ожидалась 0x%08X) — это не Xenith-модель",
magic, XNH_MAGIC);
err = buf;
return false;
}
if (ver != 1) { err = "неподдерживаемая версия: " + std::to_string(ver); return false; }
Config& c = ck.params.cfg;
int32_t v;
rd(f, v); c.vocab_size = v;
rd(f, v); c.n_layer = v;
rd(f, v); c.n_head = v;
rd(f, v); c.n_embd = v;
rd(f, v); c.block_size = v;
rd(f, v); c.ffn_dim = v;
rd(f, v); c.rope_base = v;
rd(f, c.rms_eps);
rd(f, c.init_std);
rd(f, v); c.n_vocab_out = v;
rd(f, v); ck.step = v;
rd(f, ck.loss);
rd(f, ck.seed);
rd(f, v); ck.tok.bos = v;
rd(f, v); ck.tok.eos = v;
rd(f, v); ck.tok.unk = v;
if (!c.valid(err)) return false;
int32_t nv = 0;
rd(f, nv);
ck.tok.id2tok.clear();
ck.tok.id2tok.reserve(nv);
for (int i = 0; i < nv; ++i) {
uint32_t len = 0;
rd(f, len);
std::string s(len, '\0');
if (len) f.read(&s[0], len);
ck.tok.id2tok.push_back(s);
}
if (ck.tok.size() != c.vocab_size) {
err = "размер словаря (" + std::to_string(ck.tok.size()) +
") не совпадает с vocab_size (" + std::to_string(c.vocab_size) + ")";
return false;
}
// аллоцируем и читаем веса
const int d = c.n_embd, ffn = c.ffn();
ck.params.wte.resize(c.vocab_size, d);
if (!ck.params.tied()) ck.params.lm_head.resize(c.vocab_out(), d);
ck.params.wq.resize(c.n_layer, Tensor(d, d)); // q/k/v — три отдельных тензора [d,d]
ck.params.wk.resize(c.n_layer, Tensor(d, d));
ck.params.wv.resize(c.n_layer, Tensor(d, d));
ck.params.wo.resize(c.n_layer, Tensor(d, d));
ck.params.rms_attn.resize(c.n_layer, Tensor(1, d));
ck.params.rms_ffn.resize(c.n_layer, Tensor(1, d));
ck.params.w1.resize(c.n_layer, Tensor(d, ffn));
ck.params.w2.resize(c.n_layer, Tensor(ffn, d));
ck.params.w3.resize(c.n_layer, Tensor(d, ffn));
ck.params.rms_final = Tensor(1, d);
for (Tensor* t : ck.params.all_w()) {
int32_t r, cc;
if (!rd(f, r) || !rd(f, cc)) { err = "файл обрывается в весах"; return false; }
if (r != t->R() || cc != t->C()) {
err = "размер веса не совпал: в файле " + std::to_string(r) + "x" + std::to_string(cc) +
", ожидалось " + std::to_string(t->R()) + "x" + std::to_string(t->C());
return false;
}
f.read(reinterpret_cast<char*>(t->ptr()), (std::streamsize)t->n() * sizeof(float));
if (!f) { err = "файл обрывается при чтении весов"; return false; }
}
// моменты оптимизатора (необязательно)
uint32_t has_opt = 0;
if (rd(f, has_opt) && has_opt) {
auto ps = ck.params.all();
ck.params.m.clear();
ck.params.v.clear();
for (size_t i = 0; i < ps.size(); ++i) {
int32_t r, cc;
if (!rd(f, r) || !rd(f, cc)) { err = "файл обрывается в оптимизаторе"; return false; }
ck.params.m.emplace_back(r, cc);
f.read(reinterpret_cast<char*>(ck.params.m[i].ptr()),
(std::streamsize)(r * cc) * sizeof(float));
rd(f, r); rd(f, cc);
ck.params.v.emplace_back(r, cc);
f.read(reinterpret_cast<char*>(ck.params.v[i].ptr()),
(std::streamsize)(r * cc) * sizeof(float));
}
}
ck.params.alloc_grads();
return true;
}
void print_model_info(const Checkpoint& ck) {
const Config& c = ck.params.cfg;
size_t tot = ck.params.n_params();
printf(" конфигурация:\n");
printf(" слоёв : %d\n", c.n_layer);
printf(" голов : %d (head_dim=%d)\n", c.n_head, c.head_dim());
printf(" размер эмбеддинга: %d\n", c.n_embd);
printf(" FFN : %d (SwiGLU)\n", c.ffn());
printf(" словарь : %d\n", c.vocab_size);
printf(" контекст : %d\n", c.block_size);
printf(" RoPE base : %d\n", c.rope_base);
printf(" выход: %s\n", ck.params.tied() ? "tied (wte)" : "отдельная матрица");
printf(" параметров : %zu (%.2f М)\n", tot, tot / 1048576.0);
printf(" шаг : %d, loss %.4f\n", ck.step, ck.loss);
}
} // namespace xt