Xenith/xenith/core/backward.cpp
2026-09-29 19:55:06 +07:00

243 lines
10 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.

// backward.cpp — ручной обратный проход трансформера
#include "train.h"
#include <cmath>
#include <algorithm>
namespace xt {
// softmax + cross-entropy на одном ряду; сразу отдаёт градиент по логитам
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; // теперь dlogits = вероятности
float loss = -std::log(std::max(dlogits[target], 1e-30f));
dlogits[target] -= 1.0f; // dL/dlogits = p - onehot
// Второго умножения на inv здесь быть НЕ должно: вероятности уже нормированы.
return loss;
}
float backward(Params& p, const ForwardCache& fc,
const int* x, const int* y, int n_pairs) {
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();
// ------------------------------------------- 1. логиты -> dxf, градиенты выхода
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); // dxf = dlogits @ wte
gemm_tn(dlogits, fc.xf, p.gwte, true); // dwte += dlogits^T @ xf
} else {
gemm_nn(dlogits, p.lm_head, dxf);
gemm_tn(dlogits, fc.xf, p.glm_head, true);
}
// ------------------------------------------- 2. финальный RMSNorm
Tensor dx;
rmsnorm_backward(fc.x, p.rms_final, fc.inv_rms, dxf, dx, p.grms_final);
// ------------------------------------------- 3. слои в обратном порядке
//
// Структура слоя (та же, что в forward):
// x_mid = x_in + proj proj = attout @ wo
// x2 = x_mid (x2 копия residual'а)
// x_out = x2 + fo fo = FFN(x2b), x2b = RMS(x2)
//
// dfo = dL/dx_out = dres
// dproj = dL/dx_mid = dres + вклад FFN (proj добавляется ДО fo)
// dL/dx_in = dproj + вклад attention
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];
// dres — градиент по ВЫХОДУ слоя residual'а.
const Tensor dres = dx;
// ================= FFN =================
// fo = ha @ w2
gemm_tn(lc.ha, dres, p.g2[L], true); // dW2
gemm_nt(dres, p.w2[L], dha); // dha = dres @ w2^T
// ha = silu(h1) * h3
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);
}
}
// dW = X^T @ dY => первым аргументом gemm_tn идёт X
gemm_tn(lc.x2b, dh1, p.g1[L], true); // dW1
gemm_tn(lc.x2b, dh3, p.g3[L], true); // dW3
// dx2b = dh1 @ w1^T + dh3 @ w3^T
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);
// RMSNorm перед FFN: вклад идёт в dx (градиент по x2 == x_mid)
{
Tensor g;
rmsnorm_backward(lc.x2, p.rms_ffn[L], lc.inv_rms2, dx2b, g, p.grms_ffn[L]);
for (size_t i = 0; i < dx.n(); ++i) dx.at(i) += g.at(i);
}
// ================= attention =================
// proj = attout @ wo
// Здесь читаем dx (== dL/dx_mid, уже с вкладом FFN), а НЕ dres.
gemm_tn(lc.attout, dx, p.go[L], true); // dWo
dattout.resize(B, d);
gemm_nt(dx, p.wo[L], dattout); // dattout = dx @ wo^T
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;
// dv[t2] += att[t1][t2] * dattout[t1]
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];
}
// datt[t1][t2] = dot(dattout[t1], v[t2])
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;
}
// якобиан softmax: dscores = a * (dp - sum(a*dp)) * scale
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;
// dq[t1] += dscores * k[t2]; dk[t2] += dscores * q[t1]
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];
}
}
}
}
}
// ================= RoPE backward =================
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;
}
}
}
}
// ================= QKV backward =================
gemm_tn(lc.xb, dq, p.gq[L], true); // dWq
gemm_tn(lc.xb, dk, p.gk[L], true); // dWk
gemm_tn(lc.xb, dv, p.gv[L], true); // dWv
// dxb = dq @ wq^T + dk @ wk^T + dv @ wv^T
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);
// ================= RMSNorm перед attention =================
{
Tensor g;
rmsnorm_backward(lc.x, p.rms_attn[L], lc.inv_rms, dxb, g, p.grms_attn[L]);
for (size_t i = 0; i < dx.n(); ++i) dx.at(i) += g.at(i);
}
}
// ------------------------------------------- 4. градиент входных эмбеддингов
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);
}
} // namespace xt