243 lines
10 KiB
C++
243 lines
10 KiB
C++
// 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
|