179 lines
5.2 KiB
C++
179 lines
5.2 KiB
C++
#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;
|
||
if (fread(&map_size, sizeof(map_size), 1, mdl.fd))
|
||
{
|
||
|
||
} else
|
||
{
|
||
std::cerr << "Ошибка чтения файла" << std::endl;
|
||
return mdl;
|
||
}
|
||
|
||
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;
|
||
} |