#include "model.hpp" #include #include #include 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 Params::all() const { std::vector 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 Params::all_w() { std::vector 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 Params::all_grads() { std::vector 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& cs, std::vector& 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, bool fast) { 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(); if (fast) fc.x.resize_fast(B, d); else 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]; if (fast) { lc.x.resize_fast(B, d); lc.xb.resize_fast(B, d); lc.q.resize_fast(B, d); lc.k.resize_fast(B, d); lc.v.resize_fast(B, d); lc.attout.resize_fast(B, d); lc.proj.resize_fast(B, d); lc.x2.resize_fast(B, d); lc.x2b.resize_fast(B, d); lc.h1.resize_fast(B, f); lc.h3.resize_fast(B, f); lc.ha.resize_fast(B, f); lc.fo.resize_fast(B, d); } else { lc.x.resize(B, d); lc.xb.resize(B, d); lc.q.resize(B, d); lc.k.resize(B, d); lc.v.resize(B, d); lc.attout.resize(B, d); lc.proj.resize(B, d); lc.x2.resize(B, d); lc.x2b.resize(B, d); lc.h1.resize(B, f); lc.h3.resize(B, f); lc.ha.resize(B, f); lc.fo.resize(B, d); } lc.att.resize3d(n_seq * H, S, S); if (fast) fc.xf.resize_fast(B, d); else fc.xf.resize(B, d); rmsnorm_forward(lc.x, p.rms_attn[L], c.rms_eps, lc.xb, lc.inv_rms, fast); 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); if (fast) fc.x.resize_fast(B, d); else fc.x.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); if (fast) fc.x.resize_fast(B, d); else fc.x.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, fast); gemm_nn(lc.x2b, p.w1[L], lc.h1); gemm_nn(lc.x2b, p.w3[L], lc.h3); if (fast) fc.x.resize_fast(B, f); else fc.x.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, fast); if (fast) fc.x.resize_fast(B, c.vocab_out()); else fc.x.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, bool fast) { forward(p, tokens, n_tok, n_tok, fc, fast); } void forward_last(const Params& p, const int* tokens, int n_tok, Tensor& logits, bool fast) { ForwardCache fc; forward(p, tokens, n_tok, n_tok, fc, fast); 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); } }