Xenith/xenith/core/checkpoint.cpp
2026-10-02 19:30:49 +07:00

354 lines
12 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.

#include "checkpoint.hpp"
#include <cstdio>
#include <cstring>
#include <fstream>
#include <sstream>
#include <unordered_map>
#include <algorithm>
#include <cctype>
#include <iostream>
namespace xt {
namespace {
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');
}
}
void Tokenizer::build_from_text(const std::string& text, int max_vocab, bool debug) {
id2tok.clear();
id2tok.push_back("<unk>");
id2tok.push_back("<bos>");
id2tok.push_back("<eos>");
id2tok.push_back("\n");
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);
}
if (debug)
{
std::cout << "=== Vocabulary (Size: " << id2tok.size() << ") ===" << std::endl;
for (size_t i = 0; i < id2tok.size(); ++i) {
// Для наглядности заменяем невидимые символы на их представления
std::string display_tok = id2tok[i];
if (display_tok == "\n") display_tok = "\\n";
if (display_tok == " ") display_tok = "[space]";
std::cout << " [" << i << "] \"" << display_tok << "\"" << std::endl;
}
std::cout << "===============================" << std::endl;
}
}
std::vector<int> Tokenizer::encode(const std::string& text, bool add_bos) const {
std::unordered_map<std::string, int> word2id;
std::unordered_map<std::string, int> char2id;
const int base = 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;
}
}
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));
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);
}
}