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

207 lines
7.4 KiB
C++

#include "train.hpp"
#include <cmath>
#include <algorithm>
namespace xt {
static float softmax_ce(const float* logits, int V, int target, float* dlogits) {
float mx = -1e30f;
for (int j = 0; j < V; ++j) mx = std::max(mx, logits[j]);
float sum = 0.0f;
for (int j = 0; j < V; ++j) { float e = std::exp(logits[j] - mx); dlogits[j] = e; sum += e; }
const float inv = 1.0f / sum;
for (int j = 0; j < V; ++j) dlogits[j] *= inv;
float loss = -std::log(std::max(dlogits[target], 1e-30f));
dlogits[target] -= 1.0f;
return loss;
}
float backward(Params& p, const ForwardCache& fc,
const int* x, const int* y, int n_pairs, bool fast) {
const Config& c = p.cfg;
const int d = c.n_embd, H = c.n_head, hd = c.head_dim(), f = c.ffn();
const int V = c.vocab_out();
const int B = n_pairs;
const int S = (fc.seq_len > 0) ? fc.seq_len : B;
const int n_seq = (S > 0) ? (B / S) : 1;
const int half = hd / 2;
const float scale = 1.0f / std::sqrt((float)hd);
p.zero_grad();
Tensor dlogits(B, V);
double loss_sum = 0;
for (int t = 0; t < B; ++t)
loss_sum += softmax_ce(fc.logits.row(t), V, y[t], dlogits.row(t));
const float dl = 1.0f / (float)B;
for (auto& t : dlogits.d) t *= dl;
Tensor dxf(B, d);
if (p.tied()) {
gemm_nn(dlogits, p.wte, dxf);
gemm_tn(dlogits, fc.xf, p.gwte, true, fast);
} else {
gemm_nn(dlogits, p.lm_head, dxf);
gemm_tn(dlogits, fc.xf, p.glm_head, true, fast);
}
Tensor dx;
rmsnorm_backward(fc.x, p.rms_final, fc.inv_rms, dxf, dx, p.grms_final, fast);
Tensor tmp, dxb, dx2b, dha, dh1, dh3, dq, dk, dv, datt, dattout;
for (int L = c.n_layer - 1; L >= 0; --L) {
const LayerCache& lc = fc.layers[L];
const Tensor dres = dx;
gemm_tn(lc.ha, dres, p.g2[L], true, fast); // dW2
gemm_nt(dres, p.w2[L], dha); // dha = dres @ w2^T
dh1.resize(B, f);
dh3.resize(B, f);
for (int t = 0; t < B; ++t) {
const float* g1r = lc.h1.row(t);
const float* g3r = lc.h3.row(t);
const float* dar = dha.row(t);
float* d1 = dh1.row(t);
float* d3 = dh3.row(t);
for (int j = 0; j < f; ++j) {
const float g = g1r[j], u = g3r[j], ga = dar[j];
d1[j] = ga * u * dsilu(g);
d3[j] = ga * silu(g);
}
}
gemm_tn(lc.x2b, dh1, p.g1[L], true, fast); // dW1
gemm_tn(lc.x2b, dh3, p.g3[L], true, fast); // dW3
dx2b.resize(B, d);
gemm_nt(dh1, p.w1[L], dx2b);
gemm_nt(dh3, p.w3[L], tmp);
for (size_t i = 0; i < dx2b.n(); ++i) dx2b.at(i) += tmp.at(i);
{
Tensor g;
rmsnorm_backward(lc.x2, p.rms_ffn[L], lc.inv_rms2, dx2b, g, p.grms_ffn[L], fast);
for (size_t i = 0; i < dx.n(); ++i) dx.at(i) += g.at(i);
}
gemm_tn(lc.attout, dx, p.go[L], true, fast);
dattout.resize(B, d);
gemm_nt(dx, p.wo[L], dattout);
dq.resize(B, d); dq.zero();
dk.resize(B, d); dk.zero();
dv.resize(B, d); dv.zero();
datt.resize3d(n_seq * H, S, S);
for (int sq = 0; sq < n_seq; ++sq) {
for (int hh = 0; hh < H; ++hh) {
const int hp = sq * H + hh;
for (int t1 = 0; t1 < S; ++t1) {
const int r1 = sq * S + t1;
const float* dar = dattout.row(r1) + hh * hd;
const float* ar = lc.att.row(hp * S + t1);
float* dq_row = dq.row(r1) + hh * hd;
for (int t2 = 0; t2 <= t1; ++t2) {
const float a = ar[t2];
if (a == 0.0f) continue;
float* dvr = dv.row(sq * S + t2) + hh * hd;
for (int j = 0; j < hd; ++j) dvr[j] += a * dar[j];
}
float* darow = datt.row(hp * S + t1);
for (int t2 = 0; t2 <= t1; ++t2) {
const float* vr = lc.v.row(sq * S + t2) + hh * hd;
float s = 0.0f;
for (int j = 0; j < hd; ++j) s += dar[j] * vr[j];
darow[t2] = s;
}
float dpdot = 0.0f;
for (int t2 = 0; t2 <= t1; ++t2) dpdot += ar[t2] * darow[t2];
for (int t2 = 0; t2 <= t1; ++t2)
darow[t2] = ar[t2] * (darow[t2] - dpdot) * scale;
for (int t2 = 0; t2 <= t1; ++t2) {
const float ds = darow[t2];
if (ds == 0.0f) continue;
const float* kr = lc.k.row(sq * S + t2) + hh * hd;
const float* qr = lc.q.row(r1) + hh * hd;
float* dkr = dk.row(sq * S + t2) + hh * hd;
for (int j = 0; j < hd; ++j) {
dq_row[j] += ds * kr[j];
dkr[j] += ds * qr[j];
}
}
}
}
}
for (int sq = 0; sq < n_seq; ++sq) {
for (int t = 0; t < S; ++t) {
const int row = sq * S + t;
const float* cs = fc.rope_cos.data() + (size_t)t * half;
const float* sn = fc.rope_sin.data() + (size_t)t * half;
for (int hh = 0; hh < H; ++hh) {
float* dq_ = dq.row(row) + hh * hd;
float* dk_ = dk.row(row) + hh * hd;
for (int i = 0; i < half; ++i) {
const float c0 = cs[i], s0 = sn[i];
float a0 = dq_[2 * i], a1 = dq_[2 * i + 1];
dq_[2 * i] = a0 * c0 + a1 * s0;
dq_[2 * i + 1] = -a0 * s0 + a1 * c0;
float b0 = dk_[2 * i], b1 = dk_[2 * i + 1];
dk_[2 * i] = b0 * c0 + b1 * s0;
dk_[2 * i + 1] = -b0 * s0 + b1 * c0;
}
}
}
}
gemm_tn(lc.xb, dq, p.gq[L], true, fast);
gemm_tn(lc.xb, dk, p.gk[L], true, fast);
gemm_tn(lc.xb, dv, p.gv[L], true, fast);
dxb.resize(B, d);
gemm_nt(dq, p.wq[L], dxb);
gemm_nt(dk, p.wk[L], tmp);
for (size_t i = 0; i < dxb.n(); ++i) dxb.at(i) += tmp.at(i);
gemm_nt(dv, p.wv[L], tmp);
for (size_t i = 0; i < dxb.n(); ++i) dxb.at(i) += tmp.at(i);
{
Tensor g;
rmsnorm_backward(lc.x, p.rms_attn[L], lc.inv_rms, dxb, g, p.grms_attn[L], fast);
for (size_t i = 0; i < dx.n(); ++i) dx.at(i) += g.at(i);
}
}
for (int t = 0; t < B; ++t) {
const float* gxr = dx.row(t);
float* gw = p.gwte.row(x[t]);
for (int j = 0; j < d; ++j) gw[j] += gxr[j];
}
return (float)(loss_sum / (double)B);
}
float softmax_eval(const Params& p, const ForwardCache& fc,
const int* x, const int* y, int n_pairs, Tensor& per_token_loss) {
(void)x;
const int V = p.cfg.vocab_out();
Tensor dl(n_pairs, V);
per_token_loss.resize(n_pairs, 1);
double sum = 0;
for (int t = 0; t < n_pairs; ++t) {
const float l = softmax_ce(fc.logits.row(t), V, y[t], dl.row(t));
per_token_loss.at(t) = l;
sum += l;
}
return (float)(sum / (double)n_pairs);
}
}