135 lines
4.9 KiB
C++
135 lines
4.9 KiB
C++
#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);
|
||
float loss = backward(p, fc, xb.data(), yb.data(), B * block);
|
||
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);
|
||
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;
|
||
}
|
||
|
||
}
|