Xenith/xenith/core/trainer.cpp
2026-10-04 14:53:45 +07:00

135 lines
4.9 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 "train.hpp"
#include "checkpoint.hpp"
#include <cstdio>
#include <cmath>
#include <fstream>
#include <sstream>
#include <chrono>
#include <algorithm>
namespace xt {
namespace {
std::string read_file(const std::string& path, bool& ok) {
std::ifstream f(path, std::ios::binary);
if (!f) { ok = false; return {}; }
std::ostringstream ss;
ss << f.rdbuf();
ok = true;
return ss.str();
}
float lr_at(const TrainConfig& tc, int step) {
if (step < tc.warmup) return tc.lr * (step + 1) / (float)std::max(1, tc.warmup);
const int t = step - tc.warmup;
const int total = std::max(1, tc.steps - tc.warmup);
const float f = (float)t / total;
const float cos = 0.5f * (1.0f + std::cos(3.14159265f * std::min(1.0f, f)));
return tc.lr * (tc.lr_min_frac + (1.0f - tc.lr_min_frac) * cos);
}
}
TrainStats train_model(Params& p, const Tokenizer& tok, const std::string& corpus_path,
const TrainConfig& tc, const std::string& out_path,
int resume_step) {
TrainStats st;
bool ok = false;
std::string text = read_file(corpus_path, ok);
if (!ok) { fprintf(stderr, "ошибка: не читается корпус %s\n", corpus_path.c_str()); return st; }
std::vector<int> ids = tok.encode(text, true);
if (ids.size() < 4) {
fprintf(stderr, "ошибка: корпус слишком короткий (%zu токенов)\n", ids.size());
return st;
}
fprintf(stderr, "корпус: %zu токенов из %zu байт\n", ids.size(), text.size());
Rng rng(tc.seed);
const int block = std::min(tc.block, p.cfg.block_size);
const int B = tc.batch_size;
if (B * (block + 1) > (int)ids.size()) {
fprintf(stderr, "ошибка: батч %d x %d не помещается в корпус\n", B, block);
return st;
}
std::vector<int> val_ids;
bool has_val = false;
if (tc.val_every > 0 && ids.size() > (size_t)tc.val_tokens + block + 1) {
val_ids.assign(ids.end() - tc.val_tokens, ids.end());
ids.resize(ids.size() - tc.val_tokens);
has_val = true;
fprintf(stderr, "валидация: %zu токенов\n\033[?25l", val_ids.size());
}
std::vector<int> xb(B * block), yb(B * block);
ForwardCache fc;
auto t0 = std::chrono::steady_clock::now();
const int start = resume_step;
for (int step = start; step < tc.steps; ++step) {
for (int b = 0; b < B; ++b) {
size_t off = rng.below(ids.size() - block - 1);
for (int t = 0; t < block; ++t) {
xb[b * block + t] = ids[off + t];
yb[b * block + t] = ids[off + t + 1];
}
}
forward(p, xb.data(), B * block, block, fc, tc.fast);
float loss = backward(p, fc, xb.data(), yb.data(), B * block, tc.fast);
p.adam_step(lr_at(tc, step), tc.beta1, tc.beta2, tc.eps,
tc.weight_decay, step + 1, tc.clip);
st.loss = loss;
st.step = step;
st.tokens_seen += (double)B * block;
auto now = std::chrono::steady_clock::now();
double dt = std::chrono::duration<double>(now - t0).count();
double tps = dt > 0 ? st.tokens_seen / dt : 0;
fprintf(stderr, "шаг %5d/%d loss %.4f ppl %8.2f lr %.2e %6.0f ток/с\r",
step, tc.steps, loss, std::exp(loss), lr_at(tc, step), tps);
if (tc.val_every > 0 && has_val && (step + 1) % tc.val_every == 0) {
float vl = 0;
const int nwin = 8;
for (int w = 0; w < nwin; ++w) {
size_t off = (val_ids.size() - block - 1) * w / nwin;
std::vector<int> vx(block);
for (int t = 0; t < block; ++t) vx[t] = val_ids[off + t];
forward(p, vx.data(), block, block, fc, tc.fast);
Tensor lg;
softmax_eval(p, fc, vx.data(), val_ids.data() + off + 1, block, lg);
for (int t = 0; t < block; ++t) vl += lg.at(t);
}
st.val_loss = vl / (nwin * block);
fprintf(stderr, " val loss %.4f ppl %.2f\n", st.val_loss, std::exp(st.val_loss));
}
if (tc.ckpt_every > 0 && (step + 1) % tc.ckpt_every == 0 && step + 1 < tc.steps) {
Checkpoint ck{p, tok, step + 1, loss, tc.seed};
std::string path = out_path + ".step" + std::to_string(step + 1);
std::string err;
if (save_checkpoint(path, ck, err)) fprintf(stderr, " чекпоинт -> %s\n", path.c_str());
else fprintf(stderr, " чекпоинт не сохранён: %s\n", err.c_str());
}
}
Checkpoint ck{p, tok, tc.steps, st.loss, tc.seed};
std::string err;
if (!out_path.empty()) {
if (save_checkpoint(out_path, ck, err))
fprintf(stderr, "модель сохранена: %s\n", out_path.c_str());
else
fprintf(stderr, "ошибка сохранения: %s\n", err.c_str());
}
return st;
}
}