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

342 lines
12 KiB
C++

#include "model.hpp"
#include <cstdio>
#include <cmath>
#include <algorithm>
namespace xt {
bool Config::valid(std::string& err) const {
if (n_head <= 0 || n_embd <= 0 || n_layer <= 0 || vocab_size <= 0 || block_size <= 0) {
err = "размеры должны быть положительными";
return false;
}
if (n_embd % n_head != 0) {
err = "n_embd (" + std::to_string(n_embd) + ") не делится на n_head (" +
std::to_string(n_head) + ")";
return false;
}
if (head_dim() % 2 != 0) {
err = "head_dim должен быть чётным для RoPE";
return false;
}
if (rope_base <= 0) { err = "rope_base должен быть > 0"; return false; }
return true;
}
void Params::init(uint64_t seed) {
Rng rng(seed);
const int d = cfg.n_embd, f = cfg.ffn();
wte.resize(cfg.vocab_size, d);
for (auto& t : wte.d) t = rng.normal() * cfg.init_std;
const bool tied = (cfg.n_vocab_out == 0 || cfg.n_vocab_out == cfg.vocab_size);
if (!tied) {
lm_head.resize(cfg.vocab_out(), d);
for (auto& t : lm_head.d) t = rng.normal() * cfg.init_std;
}
wq.resize(cfg.n_layer, Tensor(d, d));
wk.resize(cfg.n_layer, Tensor(d, d));
wv.resize(cfg.n_layer, Tensor(d, d));
wo.resize(cfg.n_layer, Tensor(d, d));
rms_attn.resize(cfg.n_layer, Tensor(1, d));
rms_ffn.resize(cfg.n_layer, Tensor(1, d));
w1.resize(cfg.n_layer, Tensor(d, f));
w2.resize(cfg.n_layer, Tensor(f, d));
w3.resize(cfg.n_layer, Tensor(d, f));
rms_final = Tensor(1, d);
for (int L = 0; L < cfg.n_layer; ++L) {
for (auto& t : wq[L].d) t = rng.normal() * cfg.init_std;
for (auto& t : wk[L].d) t = rng.normal() * cfg.init_std;
for (auto& t : wv[L].d) t = rng.normal() * cfg.init_std;
for (auto& t : wo[L].d) t = rng.normal() * cfg.init_std;
for (auto& t : w1[L].d) t = rng.normal() * cfg.init_std;
for (auto& t : w2[L].d) t = rng.normal() * cfg.init_std;
for (auto& t : w3[L].d) t = rng.normal() * cfg.init_std;
}
// RMS-веса = 1 (не 0.02 — иначе нормализация убивает сигнал на старте)
for (int L = 0; L < cfg.n_layer; ++L) {
for (auto& t : rms_attn[L].d) t = 1.0f;
for (auto& t : rms_ffn[L].d) t = 1.0f;
}
for (auto& t : rms_final.d) t = 1.0f;
alloc_grads();
m.clear();
v.clear();
}
bool Params::tied() const {
return cfg.n_vocab_out == 0 || cfg.n_vocab_out == cfg.vocab_size;
}
size_t Params::n_params() const {
size_t n = wte.n();
if (!tied()) n += lm_head.n();
for (int L = 0; L < cfg.n_layer; ++L) {
n += wq[L].n() + wk[L].n() + wv[L].n() + wo[L].n();
n += rms_attn[L].n() + rms_ffn[L].n();
n += w1[L].n() + w2[L].n() + w3[L].n();
}
n += rms_final.n();
return n;
}
void Params::alloc_grads() {
const int d = cfg.n_embd, f = cfg.ffn();
gwte.resize(cfg.vocab_size, d);
if (!tied()) glm_head.resize(cfg.vocab_out(), d);
gq.resize(cfg.n_layer, Tensor(d, d));
gk.resize(cfg.n_layer, Tensor(d, d));
gv.resize(cfg.n_layer, Tensor(d, d));
go.resize(cfg.n_layer, Tensor(d, d));
grms_attn.resize(cfg.n_layer, Tensor(1, d));
grms_ffn.resize(cfg.n_layer, Tensor(1, d));
g1.resize(cfg.n_layer, Tensor(d, f));
g2.resize(cfg.n_layer, Tensor(f, d));
g3.resize(cfg.n_layer, Tensor(d, f));
grms_final = Tensor(1, d);
}
std::vector<const Tensor*> Params::all() const {
std::vector<const Tensor*> v;
v.push_back(&wte);
if (!tied()) v.push_back(&lm_head);
for (int L = 0; L < cfg.n_layer; ++L) {
v.push_back(&wq[L]); v.push_back(&wk[L]); v.push_back(&wv[L]); v.push_back(&wo[L]);
v.push_back(&rms_attn[L]); v.push_back(&rms_ffn[L]);
v.push_back(&w1[L]); v.push_back(&w2[L]); v.push_back(&w3[L]);
}
v.push_back(&rms_final);
return v;
}
std::vector<Tensor*> Params::all_w() {
std::vector<Tensor*> v;
v.push_back(&wte);
if (!tied()) v.push_back(&lm_head);
for (int L = 0; L < cfg.n_layer; ++L) {
v.push_back(&wq[L]); v.push_back(&wk[L]); v.push_back(&wv[L]); v.push_back(&wo[L]);
v.push_back(&rms_attn[L]); v.push_back(&rms_ffn[L]);
v.push_back(&w1[L]); v.push_back(&w2[L]); v.push_back(&w3[L]);
}
v.push_back(&rms_final);
return v;
}
std::vector<Tensor*> Params::all_grads() {
std::vector<Tensor*> v;
v.push_back(&gwte);
if (!tied()) v.push_back(&glm_head);
for (int L = 0; L < cfg.n_layer; ++L) {
v.push_back(&gq[L]); v.push_back(&gk[L]); v.push_back(&gv[L]); v.push_back(&go[L]);
v.push_back(&grms_attn[L]); v.push_back(&grms_ffn[L]);
v.push_back(&g1[L]); v.push_back(&g2[L]); v.push_back(&g3[L]);
}
v.push_back(&grms_final);
return v;
}
void Params::zero_grad() {
for (Tensor* t : all_grads()) t->zero();
}
void Params::adam_step(float lr, float b1, float b2, float eps, float wd, int step, float clip) {
auto params = all_w();
auto grads = all_grads();
if (grads.size() != params.size()) alloc_grads();
if (m.size() != params.size()) {
m.clear(); v.clear();
for (Tensor* t : params) { m.push_back(Tensor(t->R(), t->C())); v.push_back(Tensor(t->R(), t->C())); }
}
if (clip > 0.0f) {
double sq = 0;
bool dirty = false;
for (Tensor* g : grads)
for (float& x : g->d) {
if (!std::isfinite(x)) { x = 0.0f; dirty = true; continue; }
sq += (double)x * x;
}
if (dirty)
std::fprintf(stderr,
"предупреждение: в градиенте были nan/inf, они обнулены\n");
float nrm = (float)std::sqrt(sq);
if (std::isfinite(nrm) && nrm > clip) {
const float k = clip / (nrm + 1e-6f);
for (Tensor* g : grads) for (float& x : g->d) x *= k;
}
}
const float bc1 = 1.0f - std::pow(b1, (float)step);
const float bc2 = 1.0f - std::pow(b2, (float)step);
for (size_t i = 0; i < params.size(); ++i) {
Tensor& P = *params[i];
Tensor& G = *grads[i];
Tensor& M = m[i];
Tensor& V = v[i];
for (size_t j = 0; j < P.d.size(); ++j) {
float g = G.d[j];
if (!std::isfinite(g)) continue;
if (wd > 0.0f) g += wd * P.d[j];
M.d[j] = b1 * M.d[j] + (1.0f - b1) * g;
V.d[j] = b2 * V.d[j] + (1.0f - b2) * g * g;
const float upd = (M.d[j] / bc1) / (std::sqrt(V.d[j] / bc2) + eps);
if (std::isfinite(upd)) P.d[j] -= lr * upd;
}
}
}
void rope_tables(int block, int head_dim, int base, std::vector<float>& cs, std::vector<float>& sn) {
const int half = head_dim / 2;
cs.resize((size_t)block * half);
sn.resize((size_t)block * half);
for (int pos = 0; pos < block; ++pos) {
for (int i = 0; i < half; ++i) {
float freq = 1.0f / std::pow((float)base, (2.0f * i) / head_dim);
float ang = pos * freq;
cs[(size_t)pos * half + i] = std::cos(ang);
sn[(size_t)pos * half + i] = std::sin(ang);
}
}
}
void forward(const Params& p, const int* tokens, int n_tok, int seq_len, ForwardCache& fc) {
const Config& c = p.cfg;
const int d = c.n_embd, H = c.n_head, hd = c.head_dim(), f = c.ffn();
int B = n_tok;
int S = seq_len;
if (S <= 0) S = B;
if (S > c.block_size) S = c.block_size;
if (B % S != 0) B = (B / S) * S;
const int n_seq = B / S;
fc.block = S;
fc.n_seq = n_seq;
fc.seq_len = S;
fc.layers.resize(c.n_layer);
const int half = hd / 2;
if ((size_t)fc.rope_cos.size() < (size_t)S * half) {
rope_tables(S, hd, c.rope_base, fc.rope_cos, fc.rope_sin);
}
const float* CS = fc.rope_cos.data();
const float* SN = fc.rope_sin.data();
fc.x.resize(B, d);
for (int t = 0; t < B; ++t) {
float* xr = fc.x.row(t);
const float* er = p.wte.row(tokens[t]);
for (int j = 0; j < d; ++j) xr[j] = er[j];
}
const float scale = 1.0f / std::sqrt((float)hd);
Tensor scores(S, S);
for (int L = 0; L < c.n_layer; ++L) {
LayerCache& lc = fc.layers[L];
lc.x.resize(B, d);
for (size_t i = 0; i < fc.x.n(); ++i) lc.x.at(i) = fc.x.at(i);
rmsnorm_forward(lc.x, p.rms_attn[L], c.rms_eps, lc.xb, lc.inv_rms);
gemm_nn(lc.xb, p.wq[L], lc.q);
gemm_nn(lc.xb, p.wk[L], lc.k);
gemm_nn(lc.xb, p.wv[L], lc.v);
for (int b = 0; b < n_seq; ++b) {
for (int t = 0; t < S; ++t) {
const int row = b * S + t;
const float* cs = CS + (size_t)t * half;
const float* sn = SN + (size_t)t * half;
for (int hh = 0; hh < H; ++hh) {
float* q = lc.q.row(row) + hh * hd;
float* k = lc.k.row(row) + hh * hd;
for (int i = 0; i < half; ++i) {
const float c0 = cs[i], s0 = sn[i];
float q0 = q[2 * i], q1 = q[2 * i + 1];
q[2 * i] = q0 * c0 - q1 * s0;
q[2 * i + 1] = q0 * s0 + q1 * c0;
float k0 = k[2 * i], k1 = k[2 * i + 1];
k[2 * i] = k0 * c0 - k1 * s0;
k[2 * i + 1] = k0 * s0 + k1 * c0;
}
}
}
}
lc.att.resize3d(n_seq * H, S, S);
lc.attout.resize(B, d);
lc.attout.zero();
for (int s = 0; s < n_seq; ++s) {
for (int hh = 0; hh < H; ++hh) {
const int hp = s * H + hh;
for (int t1 = 0; t1 < S; ++t1) {
const float* qr = lc.q.row(s * S + t1) + hh * hd;
float* sr = scores.row(t1);
float mx = -1e30f;
for (int t2 = 0; t2 <= t1; ++t2) {
const float* kr = lc.k.row(s * S + t2) + hh * hd;
float s2 = 0.0f;
for (int j = 0; j < hd; ++j) s2 += qr[j] * kr[j];
s2 *= scale;
sr[t2] = s2;
if (s2 > mx) mx = s2;
}
float sum = 0.0f;
for (int t2 = 0; t2 <= t1; ++t2) { sr[t2] = std::exp(sr[t2] - mx); sum += sr[t2]; }
const float inv = 1.0f / sum;
float* ar = lc.att.row(hp * S + t1);
for (int t2 = 0; t2 < S; ++t2) ar[t2] = (t2 <= t1) ? sr[t2] * inv : 0.0f;
}
for (int t1 = 0; t1 < S; ++t1) {
const float* ar = lc.att.row(hp * S + t1);
float* orow = lc.attout.row(s * S + t1) + hh * hd;
for (int t2 = 0; t2 <= t1; ++t2) {
const float a = ar[t2];
if (a == 0.0f) continue;
const float* vr = lc.v.row(s * S + t2) + hh * hd;
for (int j = 0; j < hd; ++j) orow[j] += a * vr[j];
}
}
}
}
gemm_nn(lc.attout, p.wo[L], lc.proj);
for (size_t i = 0; i < fc.x.n(); ++i) fc.x.at(i) += lc.proj.at(i);
lc.x2.resize(B, d);
for (size_t i = 0; i < fc.x.n(); ++i) lc.x2.at(i) = fc.x.at(i);
rmsnorm_forward(lc.x2, p.rms_ffn[L], c.rms_eps, lc.x2b, lc.inv_rms2);
gemm_nn(lc.x2b, p.w1[L], lc.h1);
gemm_nn(lc.x2b, p.w3[L], lc.h3);
lc.ha.resize(B, f);
for (size_t i = 0; i < lc.ha.n(); ++i)
lc.ha.at(i) = silu(lc.h1.at(i)) * lc.h3.at(i);
gemm_nn(lc.ha, p.w2[L], lc.fo);
for (size_t i = 0; i < fc.x.n(); ++i) fc.x.at(i) += lc.fo.at(i);
}
rmsnorm_forward(fc.x, p.rms_final, c.rms_eps, fc.xf, fc.inv_rms);
fc.logits.resize(B, c.vocab_out());
if (p.tied()) gemm_nt(fc.xf, p.wte, fc.logits);
else gemm_nn(fc.xf, p.lm_head, fc.logits);
}
void forward(const Params& p, const int* tokens, int n_tok, ForwardCache& fc) {
forward(p, tokens, n_tok, n_tok, fc);
}
void forward_last(const Params& p, const int* tokens, int n_tok, Tensor& logits) {
ForwardCache fc;
forward(p, tokens, n_tok, n_tok, fc);
const int V = p.cfg.vocab_out();
logits.resize(1, V);
for (int j = 0; j < V; ++j) logits.at(j) = fc.logits.at((n_tok - 1) * V + j);
}
}