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

172 lines
5.0 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 "converter.hpp"
#include <cstring>
#include <iostream>
bool save_to_bfr(const std::string& filename, const Model& model)
{
std::ofstream out(filename, std::ios::binary);
if (!out.is_open()) {
std::cerr << "Критическая ошибка! Не удалось создать файл: " << filename << std::endl;
std::cerr << "Причина от ОС: " << std::strerror(errno) << " (код ошибки: " << errno << ")" << std::endl;
return false;
}
out.write(reinterpret_cast<const char*>(&model.config), sizeof(ModelConfig));
uint32_t map_size = static_cast<uint32_t>(model.vocab.size());
out.write(reinterpret_cast<const char*>(&map_size), sizeof(map_size));
for (const auto& [token, id] : model.vocab)
{
out.write(reinterpret_cast<const char*>(&id), sizeof(id));
uint32_t token_len = static_cast<uint32_t>(token.length());
out.write(reinterpret_cast<const char*>(&token_len), sizeof(token_len));
out.write(token.data(), token_len);
}
size_t w_emb_bytes = model.w_emb.size() * sizeof(float);
out.write(reinterpret_cast<const char*>(model.w_emb.data()), w_emb_bytes);
size_t w_qkv_bytes = model.w_qkv.size() * sizeof(float);
out.write(reinterpret_cast<const char*>(model.w_qkv.data()), w_qkv_bytes);
size_t w_o_bytes = model.w_o.size() * sizeof(float);
out.write(reinterpret_cast<const char*>(model.w_emb.data()), w_o_bytes);
return true;
}
Model load_from_bfr(const std::string& filename)
{
Model mdl;
mdl.fd = fopen(filename.c_str(), "rb");
if (!mdl.fd) {
perror("fopen failed");
return mdl;
}
size_t items_read = fread(&mdl.config, sizeof(ModelConfig), 1, mdl.fd);
if (items_read != 1) {
std::cerr << "Ошибка: прочитано " << items_read
<< " структур вместо 1\n";
fclose(mdl.fd);
mdl.fd = nullptr;
return mdl;
}
uint32_t map_size;
fread(&map_size, sizeof(map_size), 1, mdl.fd);
if (map_size != mdl.config.vocab_size) {
std::cerr << "Warning: vocab size mismatch! "
<< "Header says " << mdl.config.vocab_size
<< ", file has " << map_size << "\n";
}
mdl.vocab.reserve(map_size);
for (uint32_t i = 0; i < map_size; ++i) {
uint32_t id, token_len;
fread(&id, sizeof(id), 1, mdl.fd);
fread(&token_len, sizeof(token_len), 1, mdl.fd);
std::string token(token_len, '\0');
fread(&token[0], 1, token_len, mdl.fd);
mdl.vocab.emplace(std::move(token), id);
}
mdl.offset_emb = ftell(mdl.fd);
mdl.offset_qkv = mdl.offset_emb + (mdl.config.vocab_size * mdl.config.embedding_dim * sizeof(float));
mdl.offset_o = mdl.offset_qkv + (mdl.config.num_layers * mdl.config.embedding_dim * 3 * mdl.config.num_heads * sizeof(float));
return mdl;
}
void Model::close_fd()
{
if (fd != nullptr)
{
fclose(fd);
fd = nullptr;
}
}
void Model::load(size_t num, std::vector<float>& target, uint64_t offset)
{
target.resize(num);
if (fseek(fd, offset, SEEK_SET) != 0) {
perror("fseek failed");
target.clear();
}
size_t bt_read = fread(target.data(), 1, num * sizeof(float), fd);
if (bt_read != num * sizeof(float)) {
std::cerr << "Read error: expected " << num * sizeof(float)
<< " bytes, got " << bt_read << "\n";
target.clear();
}
}
void Model::load_edm()
{
load(
config.vocab_size * config.embedding_dim,
w_emb,
offset_emb
);
}
void Model::load_qkv()
{
load(
config.num_layers * config.embedding_dim * 3 * config.num_heads,
w_qkv,
offset_qkv
);
}
void Model::load_o()
{
load(
config.num_layers * (config.num_heads * config.head_dim) * config.embedding_dim,
w_o,
offset_o
);
}
void Model::free_edm() { w_emb.resize(0); w_emb.shrink_to_fit(); }
void Model::free_qkv() { w_qkv.resize(0); w_qkv.shrink_to_fit(); }
void Model::free_o() { w_o.resize(0); w_o.shrink_to_fit(); }
void Model::allocate_kv_cache() {
size_t total_elements =
(size_t)config.max_ctx *
config.num_layers *
2 *
config.num_heads *
config.head_dim;
kv_cache = (float*)malloc(total_elements * sizeof(float));
if (!kv_cache) {
perror("Allocation failed");
return;
}
memset(kv_cache, 0, total_elements);
ctx_len = 0;
}
void Model::free_kv_cache()
{
free(kv_cache);
ctx_len = 0;
}
float* Model::get_kv(uint64_t token, uint32_t layer, cache_type kv, uint32_t head, uint32_t dim)
{
size_t head_stride = config.head_dim;
size_t kv_stride = config.num_heads * config.head_dim;
size_t layer_stride = 2 * kv_stride;
size_t token_stride = config.num_layers * layer_stride;
size_t offset = (token * token_stride) +
(layer * layer_stride) +
(static_cast<int>(kv) * kv_stride) +
(head * head_stride) +
dim;
return kv_cache + offset;
}