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

166 lines
5.5 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.

// gradcheck.cpp — проверяем, что обратный проход совпадает с численным
// и заодно меряем скорость.
#include "gradcheck.hpp"
#include "generate.hpp"
#include "train.hpp"
#include <cstdio>
#include <cmath>
#include <chrono>
#include <vector>
#include <random>
namespace xt {
namespace {
// loss на фиксированном батче, без обратного прохода
float eval_loss(const Params& p, const std::vector<int>& x, const std::vector<int>& y, int n, bool fast) {
ForwardCache fc;
forward(p, x.data(), n, fc, fast);
Tensor loss;
return softmax_eval(p, fc, x.data(), y.data(), n, loss);
}
struct Probe {
Tensor* w;
int idx;
float analytic;
float numeric;
};
}
int run_gradcheck(bool verbose, bool fast) {
Config c;
c.vocab_size = 24;
c.n_layer = 2;
c.n_head = 2;
c.n_embd = 8;
c.block_size = 12;
c.ffn_dim = 16;
c.rope_base = 10000;
c.init_std = 0.3f;
Params p;
p.cfg = c;
p.init(777);
const int n = 9;
std::vector<int> x(n), y(n);
for (int i = 0; i < n; ++i) {
x[i] = 1 + (i * 7) % (c.vocab_size - 1);
y[i] = 1 + (i * 5 + 3) % (c.vocab_size - 1);
}
ForwardCache fc;
forward(p, x.data(), n, n, fc, fast);
backward(p, fc, x.data(), y.data(), n, fast);
auto tensors = p.all_w();
auto grads = p.all_grads();
const float eps = 1e-2f;
double worst_rel = 0.0;
std::string worst_name;
int checked = 0, failed = 0;
printf("gradcheck: модель %dL x %dH x %dE, ffn=%d, %d пар\n",
c.n_layer, c.n_head, c.n_embd, c.ffn(), n);
printf("проверяю %zu тензоров, eps=%g\n\n", tensors.size(), eps);
for (size_t ti = 0; ti < tensors.size(); ++ti) {
Tensor& W = *tensors[ti];
const Tensor& G = *grads[ti];
const int sz = (int)W.n();
Rng rng(1000 + (uint64_t)ti);
int picks[3] = {(int)rng.below(sz), (int)rng.below(sz), (int)rng.below(sz)};
for (int k = 0; k < 3; ++k) {
const int i = picks[k];
const float orig = W.at(i);
W.at(i) = orig + eps;
const float lp = eval_loss(p, x, y, n, fast);
W.at(i) = orig - eps;
const float lm = eval_loss(p, x, y, n, fast);
W.at(i) = orig;
const float num = (lp - lm) / (2.0f * eps);
const float ana = G.at(i);
const float denom = std::max(1e-4f, G.L2() / std::sqrt((float)std::max(1, (int)G.n())));
const float rel = std::fabs(num - ana) / denom;
checked++;
if (rel > 0.05f) failed++;
if (rel > worst_rel) { worst_rel = rel; char b[64]; std::snprintf(b, sizeof b, "тензор %zu[%d]", ti, i); worst_name = b; }
if (verbose || rel > 0.05f) {
printf(" тензор %2zu (%3dx%-3d) [%5d] аналит %+.6f числ %+.6f rel %.2e%s\n",
ti, W.R(), W.C(), i, ana, num, rel, rel > 0.05f ? " <-- РАСХОЖДЕНИЕ" : "");
}
}
}
printf("\nпроверено %d значений, расхождений %d\n", checked, failed);
printf("худшая относительная ошибка: %.3e (%s)\n", worst_rel, worst_name.c_str());
if (failed == 0 && worst_rel < 0.05) {
printf("ГОДНО: обратный проход сходится с численным.\n");
return 0;
}
printf("НЕ ГОДНО: обратный проход содержит ошибку.\n");
return 1;
}
int run_bench(Checkpoint& ck, int iters, int block, bool fast) {
const Config& c = ck.params.cfg;
if (block > c.block_size) block = c.block_size;
const int n = block;
std::vector<int> x(n), y(n);
Rng rng(42);
for (int i = 0; i < n; ++i) {
x[i] = (int)rng.below(ck.tok.size());
y[i] = (int)rng.below(ck.tok.size());
}
printf("бенчмарк: %dL x %dH x %dE, ffn=%d, вокно %d, словарь %d, потоков %d\n",
c.n_layer, c.n_head, c.n_embd, c.ffn(), n, c.vocab_size, hw_threads());
// прогрев
ForwardCache fc;
forward(ck.params, x.data(), n, n, fc, fast);
backward(ck.params, fc, x.data(), y.data(), n, fast);
auto t0 = std::chrono::steady_clock::now();
for (int i = 0; i < iters; ++i) forward(ck.params, x.data(), n, n, fc, fast);
auto t1 = std::chrono::steady_clock::now();
double d_fwd = std::chrono::duration<double>(t1 - t0).count();
float fwd_tps = (double)n * iters / d_fwd;
t0 = std::chrono::steady_clock::now();
for (int i = 0; i < iters; ++i) {
forward(ck.params, x.data(), n, n, fc, fast);
backward(ck.params, fc, x.data(), y.data(), n, fast);
}
t1 = std::chrono::steady_clock::now();
double d_all = std::chrono::duration<double>(t1 - t0).count();
float all_tps = (double)n * iters / d_all;
GenConfig gc;
gc.max_new_tokens = 32;
gc.temperature = 0.0f;
gc.stream = false;
t0 = std::chrono::steady_clock::now();
generate(ck.params, ck.tok, "тест", gc, nullptr);
t1 = std::chrono::steady_clock::now();
double d_gen = std::chrono::duration<double>(t1 - t0).count();
printf(" форвард : %8.0f ток/с (%.2f мс на %d токов)\n", fwd_tps, d_fwd / iters * 1000, n);
printf(" форвард+назад : %8.0f ток/с (%.2f мс)\n", all_tps, d_all / iters * 1000);
printf(" генерация : %8.1f ток/с\n", 32.0 / d_gen);
return 0;
}
}