354 lines
12 KiB
C++
354 lines
12 KiB
C++
#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);
|
||
}
|
||
|
||
}
|