SPB Git

spb/forge Public MIT

Forge — LLM training from scratch in pure C++20 + Metal on Apple Silicon.

C++ 61.2% C 23% Python 7.6% TeX 7.2% CMake 1.1%
8.2 KB · 238 lines cpp
Raw Blame History
1// Author: Simon-Pierre Boucher — contact@spboucher.ai2//3// Numerical gradient checking (central differences) for every op with4// parameters, plus a full tiny-model loss. Forward math is f32, so the5// difference quotient carries ~1e-4-relative noise; eps is chosen per the6// usual sqrt(machine-eps) rule for O(1) inputs and tolerances follow7// CLAUDE.md (rel error <= 1e-3 with an absolute floor).8#include "core/autograd.h"9#include "nn/transformer.h"10#include "ops/ops.h"1112#include <cmath>13#include <cstdio>14#include <functional>15#include <random>16#include <vector>1718using namespace forge;1920namespace {2122int g_failures = 0;23constexpr float kEps = 1e-2f;24constexpr float kRelTol = 1e-3f;25constexpr float kAbsTol = 1e-4f;2627// loss_fn must run forward and return the scalar loss Var. Gradcheck28// perturbs every element of every checked Var and compares the analytic29// gradient against (f(x+e) - f(x-e)) / 2e.30void gradcheck(const char* name, const std::function<Var()>& loss_fn,31               const std::vector<Var>& checked, float eps = kEps, float tol = kRelTol) {32    // analytic33    for (const Var& p : checked) p.zero_grad();34    Tape::get().clear();35    Var loss = loss_fn();36    Tape::get().backward(loss);3738    float worst = 0.0f;39    for (const Var& p : checked) {40        const int64_t n = p.value().numel();41        float* x = p.value().data<float>();42        const float* g = p.grad().data<float>();43        for (int64_t i = 0; i < n; ++i) {44            const float saved = x[i];45            float fp, fm;46            {47                NoGrad ng;48                x[i] = saved + eps;49                fp = loss_fn().value().data<float>()[0];50                x[i] = saved - eps;51                fm = loss_fn().value().data<float>()[0];52            }53            x[i] = saved;54            const float num = (fp - fm) / (2.0f * eps);55            const float ana = g[i];56            const float err = std::fabs(num - ana);57            const float rel = err / std::max(kAbsTol / kRelTol, std::fabs(num) + std::fabs(ana));58            worst = std::max(worst, rel);59        }60    }61    if (worst <= tol) {62        std::printf("  ok: %s (max rel err %.2e)\n", name, double(worst));63    } else {64        std::printf("  FAIL: %s (max rel err %.2e > %.0e)\n", name, double(worst),65                    double(tol));66        ++g_failures;67    }68}6970Var make_var(std::vector<int64_t> shape, std::mt19937_64& rng, float std = 1.0f) {71    Tensor t = Tensor::empty(std::move(shape));72    cpu::fill_normal(t, 0.0f, std, rng);73    return Var(std::move(t), /*requires_grad=*/true);74}7576// Reduce an op output to a scalar via a fixed random projection so every77// output element influences the loss.78Var project(const Var& y, const Tensor& r) {79    Var rv(r, /*requires_grad=*/false);80    const int64_t n = y.value().numel();81    Var y2 = y.reshaped({1, n});82    Var r2 = rv.reshaped({1, n});83    return ops::matmul(y2, r2, false, true); // [1,1]84}8586Tensor random_proj(int64_t n, std::mt19937_64& rng) {87    Tensor r = Tensor::empty({n});88    cpu::fill_normal(r, 0.0f, 1.0f, rng);89    return r;90}9192} // namespace9394int main() {95    std::mt19937_64 rng(42);9697    // matmul, all transpose variants98    for (int ta = 0; ta <= 1; ++ta) {99        for (int tb = 0; tb <= 1; ++tb) {100            const int64_t M = 3, K = 4, N = 5;101            Var a = ta ? make_var({K, M}, rng) : make_var({M, K}, rng);102            Var b = tb ? make_var({N, K}, rng) : make_var({K, N}, rng);103            Tensor r = random_proj(M * N, rng);104            char name[64];105            std::snprintf(name, sizeof(name), "matmul ta=%d tb=%d", ta, tb);106            gradcheck(name, [&]() { return project(ops::matmul(a, b, ta, tb), r); },107                      {a, b});108        }109    }110111    // add_bias112    {113        Var x = make_var({4, 6}, rng);114        Var b = make_var({6}, rng);115        Tensor r = random_proj(24, rng);116        gradcheck("add_bias", [&]() { return project(ops::add_bias(x, b), r); }, {x, b});117    }118119    // mul120    {121        Var a = make_var({3, 5}, rng);122        Var b = make_var({3, 5}, rng);123        Tensor r = random_proj(15, rng);124        gradcheck("mul", [&]() { return project(ops::mul(a, b), r); }, {a, b});125    }126127    // silu / gelu128    {129        Var x = make_var({4, 5}, rng);130        Tensor r = random_proj(20, rng);131        gradcheck("silu", [&]() { return project(ops::silu(x), r); }, {x});132        gradcheck("gelu", [&]() { return project(ops::gelu(x), r); }, {x});133    }134135    // rmsnorm / layernorm136    {137        Var x = make_var({3, 8}, rng);138        Var w = make_var({8}, rng);139        Var b = make_var({8}, rng);140        Tensor r = random_proj(24, rng);141        gradcheck("rmsnorm", [&]() { return project(ops::rmsnorm(x, w, 1e-6f), r); }, {x, w});142        gradcheck("layernorm",143                  [&]() { return project(ops::layernorm(x, w, b, 1e-6f), r); }, {x, w, b});144    }145146    // embedding147    {148        Var w = make_var({7, 4}, rng);149        Tensor ids = Tensor::empty({2, 3}, DType::I32);150        std::mt19937_64 idrng(7);151        cpu::fill_uniform_int(ids, 0, 7, idrng);152        Tensor r = random_proj(2 * 3 * 4, rng);153        gradcheck("embedding", [&]() { return project(ops::embedding(w, ids), r); }, {w});154    }155156    // rope157    {158        Var x = make_var({2, 5, 2 * 6}, rng); // B=2, T=5, H=2, hd=6159        Tensor r = random_proj(2 * 5 * 12, rng);160        gradcheck("rope", [&]() { return project(ops::rope(x, 2, 10000.0f, 0), r); }, {x});161    }162163    // attention: MHA causal, then GQA164    {165        const int64_t B = 2, T = 4, H = 2, hd = 4;166        Var q = make_var({B, T, H * hd}, rng, 0.5f);167        Var k = make_var({B, T, H * hd}, rng, 0.5f);168        Var v = make_var({B, T, H * hd}, rng, 0.5f);169        Tensor r = random_proj(B * T * H * hd, rng);170        const float scale = 1.0f / std::sqrt(float(hd));171        gradcheck("attention causal MHA",172                  [&]() { return project(ops::attention(q, k, v, H, H, true, scale), r); },173                  {q, k, v});174175        Var kg = make_var({B, T, 1 * hd}, rng, 0.5f);176        Var vg = make_var({B, T, 1 * hd}, rng, 0.5f);177        gradcheck("attention causal GQA (2q/1kv)",178                  [&]() { return project(ops::attention(q, kg, vg, H, 1, true, scale), r); },179                  {q, kg, vg});180    }181182    // cross entropy (with an ignored index)183    {184        Var logits = make_var({6, 9}, rng);185        Tensor targets = Tensor::empty({6}, DType::I32);186        std::mt19937_64 idrng(11);187        cpu::fill_uniform_int(targets, 0, 9, idrng);188        targets.data<int32_t>()[3] = -1; // ignore_index189        gradcheck("cross_entropy", [&]() { return ops::cross_entropy(logits, targets); },190                  {logits});191    }192193    // full tiny model: every parameter of a 2-layer transformer194    {195        ModelConfig cfg;196        cfg.n_layers = 2;197        cfg.d_model = 16;198        cfg.n_heads = 2;199        cfg.n_kv_heads = 1; // exercise GQA200        cfg.d_ff = 24;201        cfg.vocab_size = 11;202        cfg.context_length = 8;203        cfg.tied_embeddings = true;204        nn::Transformer model(cfg, 123);205206        Tensor ids = Tensor::empty({2, 5}, DType::I32);207        Tensor targets = Tensor::empty({2 * 5}, DType::I32);208        std::mt19937_64 idrng(13);209        cpu::fill_uniform_int(ids, 0, cfg.vocab_size, idrng);210        cpu::fill_uniform_int(targets, 0, cfg.vocab_size, idrng);211212        std::vector<Var> params;213        std::unordered_map<const void*, bool> seen;214        for (const auto& [name, p] : model.named_parameters()) {215            if (!seen.emplace(p.id(), true).second) continue;216            params.push_back(p);217        }218        // Composite check: the analytic gradient is exact per the per-op219        // checks above; the f32 forward pass through ~100 ops limits the220        // central-difference quotient itself (error scales as eps^2 —221        // verified 2.6e-2 @ eps=1e-2 -> 2.5e-3 @ eps=3e-3), so this runs222        // with a documented noise-bound tolerance, not the per-op 1e-3.223        gradcheck("tiny transformer (all params)",224                  [&]() {225                      Var logits = model.forward(ids);226                      return ops::cross_entropy(logits, targets);227                  },228                  params, /*eps=*/3e-3f, /*tol=*/5e-3f);229    }230231    if (g_failures) {232        std::printf("\n%d gradcheck(s) FAILED\n", g_failures);233        return 1;234    }235    std::printf("\nall gradchecks passed\n");236    return 0;237}238