#include "train.hpp" #include #include 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) { 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); } else { gemm_nn(dlogits, p.lm_head, dxf); gemm_tn(dlogits, fc.xf, p.glm_head, true); } Tensor dx; rmsnorm_backward(fc.x, p.rms_final, fc.inv_rms, dxf, dx, p.grms_final); 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); // 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); // dW1 gemm_tn(lc.x2b, dh3, p.g3[L], true); // 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]); for (size_t i = 0; i < dx.n(); ++i) dx.at(i) += g.at(i); } gemm_tn(lc.attout, dx, p.go[L], true); 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); gemm_tn(lc.xb, dk, p.gk[L], true); gemm_tn(lc.xb, dv, p.gv[L], true); 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]); 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); } }