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%
28.3 KB · 695 lines cpp
Raw Blame History
1// Author: Simon-Pierre Boucher — contact@spboucher.ai2//3// CPU-vs-Metal parity for every GPU kernel (CLAUDE.md protocol #1: max abs4// error <= 1e-4 for f32), plus tensor mechanics and the CPU matmul oracle5// check. GPU ops encode into the batched Stream; comparisons happen after6// metal::sync() — the readback boundary.7#include <Foundation/Foundation.hpp>8#include <Metal/Metal.hpp>910#include "core/tensor.h"11#include "nn/transformer.h"12#include "ops/cpu/cpu_ops.h"13#include "ops/metal/metal_ops.h"14#include "ops/ops.h"1516#include <cmath>17#include <cstdio>18#include <cstdlib>19#include <random>20#include <vector>2122namespace {2324int g_failures = 0;25constexpr float kTol = 1e-4f;2627void expect(bool cond, const char* what) {28    if (cond) {29        std::printf("  ok: %s\n", what);30    } else {31        std::printf("  FAIL: %s\n", what);32        ++g_failures;33    }34}3536void expect_close(const forge::Tensor& got, const forge::Tensor& want, const char* what,37                  float tol = kTol) {38    const float* pg = got.data<float>();39    const float* pw = want.data<float>();40    float m = 0.0f;41    for (int64_t i = 0; i < got.numel(); ++i) m = std::max(m, std::fabs(pg[i] - pw[i]));42    if (m <= tol) {43        std::printf("  ok: %s (max abs err %.2e)\n", what, double(m));44    } else {45        std::printf("  FAIL: %s (max abs err %.2e > %.0e)\n", what, double(m), double(tol));46        ++g_failures;47    }48}4950void fill_random(forge::Tensor& t, std::mt19937& rng, float lo = -1.0f, float hi = 1.0f) {51    std::uniform_real_distribution<float> dist(lo, hi);52    float* p = t.data<float>();53    for (int64_t i = 0; i < t.numel(); ++i) p[i] = dist(rng);54}5556// Deliberately dumb oracle: index-arithmetic triple loop with double acc.57void matmul_oracle(const forge::Tensor& a, const forge::Tensor& b, forge::Tensor& c,58                   bool ta, bool tb) {59    const int64_t M = c.size(0), N = c.size(1);60    const int64_t K = ta ? a.size(0) : a.size(1);61    const int64_t lda = a.size(1), ldb = b.size(1);62    const float* A = a.data<float>();63    const float* B = b.data<float>();64    float* C = c.data<float>();65    for (int64_t i = 0; i < M; ++i)66        for (int64_t j = 0; j < N; ++j) {67            double acc = 0.0;68            for (int64_t k = 0; k < K; ++k) {69                const float av = ta ? A[k * lda + i] : A[i * lda + k];70                const float bv = tb ? B[j * ldb + k] : B[k * ldb + j];71                acc += double(av) * double(bv);72            }73            C[i * N + j] = float(acc);74        }75}7677void test_tensor_basics() {78    std::printf("tensor basics\n");79    forge::Tensor t = forge::Tensor::zeros({4, 8});80    expect(t.numel() == 32, "numel");81    expect(t.is_contiguous(), "contiguous");82    expect(t.strides()[0] == 8 && t.strides()[1] == 1, "row-major strides");8384    forge::Tensor v = t.view({8, 4});85    v.set_item(0, 42.0f);86    expect(t.item_at(0) == 42.0f, "view shares storage");8788    forge::Tensor s = t.slice0(1, 2);89    s.set_item(0, 7.0f);90    expect(t.item_at(8) == 7.0f, "slice0 shares storage at offset");9192    forge::Tensor h = forge::Tensor::full({3}, 1.5f, forge::DType::F16);93    expect(std::fabs(h.item_at(2) - 1.5f) < 1e-6f, "f16 roundtrip");94    forge::Tensor bf = forge::Tensor::full({3}, 1.5f, forge::DType::BF16);95    expect(std::fabs(bf.item_at(1) - 1.5f) < 1e-6f, "bf16 roundtrip");96}9798void test_cpu_matmul() {99    std::printf("cpu matmul vs oracle (all transpose variants)\n");100    std::mt19937 rng(1234);101    const int64_t M = 17, K = 23, N = 13;102    for (int ta = 0; ta <= 1; ++ta)103        for (int tb = 0; tb <= 1; ++tb) {104            forge::Tensor a = ta ? forge::Tensor::empty({K, M}) : forge::Tensor::empty({M, K});105            forge::Tensor b = tb ? forge::Tensor::empty({N, K}) : forge::Tensor::empty({K, N});106            forge::Tensor c = forge::Tensor::empty({M, N});107            forge::Tensor ref = forge::Tensor::empty({M, N});108            fill_random(a, rng);109            fill_random(b, rng);110            forge::cpu::matmul(a, b, c, ta, tb);111            matmul_oracle(a, b, ref, ta, tb);112            char label[64];113            std::snprintf(label, sizeof(label), "cpu matmul ta=%d tb=%d", ta, tb);114            expect_close(c, ref, label, 1e-5f);115        }116}117118void test_metal_matmul() {119    std::printf("metal matmul parity (all kernels, all transpose variants)\n");120    std::mt19937 rng(4321);121    using K_t = forge::metal::MatmulKernel;122    struct KernelCase { K_t kernel; const char* name; };123    const KernelCase kernels[] = {{K_t::Naive, "naive"},124                                  {K_t::Tiled, "tiled"},125                                  {K_t::Simdgroup, "simd"}};126    // Ragged (non-multiple of any tile dim) and exactly-tiled shapes, so the127    // simdgroup kernel's predicated and ALIGNED fast paths both get covered.128    struct Shape { int64_t M, K, N; const char* tag; };129    const Shape shapes[] = {{67, 129, 45, "ragged"}, {128, 64, 192, "aligned"}};130131    for (const auto& kc : kernels)132        for (const auto& sh : shapes)133            for (int ta = 0; ta <= 1; ++ta)134                for (int tb = 0; tb <= 1; ++tb) {135                    forge::Tensor a = ta ? forge::Tensor::empty({sh.K, sh.M})136                                         : forge::Tensor::empty({sh.M, sh.K});137                    forge::Tensor b = tb ? forge::Tensor::empty({sh.N, sh.K})138                                         : forge::Tensor::empty({sh.K, sh.N});139                    forge::Tensor gpu = forge::Tensor::empty({sh.M, sh.N});140                    forge::Tensor ref = forge::Tensor::empty({sh.M, sh.N});141                    fill_random(a, rng);142                    fill_random(b, rng);143                    forge::cpu::matmul(a, b, ref, ta, tb);144                    forge::metal::matmul(a, b, gpu, ta, tb, false, kc.kernel);145                    forge::metal::sync();146                    char label[80];147                    std::snprintf(label, sizeof(label), "%s matmul %s ta=%d tb=%d",148                                  kc.name, sh.tag, ta, tb);149                    expect_close(gpu, ref, label);150                }151152    // accumulate flag (ragged shape: both epilogue paths get predication)153    const int64_t M = 67, K = 129, N = 45;154    forge::Tensor a = forge::Tensor::empty({M, K});155    forge::Tensor b = forge::Tensor::empty({K, N});156    forge::Tensor acc_gpu = forge::Tensor::empty({M, N});157    forge::Tensor acc_ref = forge::Tensor::empty({M, N});158    fill_random(a, rng);159    fill_random(b, rng);160    fill_random(acc_gpu, rng);161    std::memcpy(acc_ref.raw(), acc_gpu.raw(), acc_gpu.nbytes());162    forge::cpu::matmul(a, b, acc_ref, false, false, true);163    forge::metal::matmul(a, b, acc_gpu, false, false, true, K_t::Tiled);164    forge::metal::sync();165    expect_close(acc_gpu, acc_ref, "tiled matmul accumulate=true");166167    // simdgroup accumulate takes the staged (non-fast) epilogue path168    forge::Tensor sacc_gpu = forge::Tensor::empty({M, N});169    forge::Tensor sacc_ref = forge::Tensor::empty({M, N});170    fill_random(sacc_gpu, rng);171    std::memcpy(sacc_ref.raw(), sacc_gpu.raw(), sacc_gpu.nbytes());172    forge::cpu::matmul(a, b, sacc_ref, false, false, true);173    forge::metal::matmul(a, b, sacc_gpu, false, false, true, K_t::Simdgroup);174    forge::metal::sync();175    expect_close(sacc_gpu, sacc_ref, "simd matmul accumulate=true");176}177178void test_metal_elementwise() {179    std::printf("metal elementwise parity\n");180    std::mt19937 rng(777);181    const int64_t n = 65537; // non-multiple of the threadgroup size182    forge::Tensor x = forge::Tensor::empty({n});183    forge::Tensor y = forge::Tensor::empty({n});184    fill_random(x, rng, -4.0f, 4.0f);185    fill_random(y, rng, -4.0f, 4.0f);186187    forge::Tensor gpu = forge::Tensor::empty({n});188    forge::Tensor ref = forge::Tensor::empty({n});189190    forge::cpu::add(x, y, ref);191    forge::metal::add(x, y, gpu);192    forge::metal::sync();193    expect_close(gpu, ref, "add");194195    forge::cpu::mul(x, y, ref);196    forge::metal::mul(x, y, gpu);197    forge::metal::sync();198    expect_close(gpu, ref, "mul");199200    forge::cpu::scale(x, 0.37f, ref);201    forge::metal::scale(x, 0.37f, gpu);202    forge::metal::sync();203    expect_close(gpu, ref, "scale");204205    forge::cpu::silu(x, ref);206    forge::metal::silu(x, gpu);207    forge::metal::sync();208    expect_close(gpu, ref, "silu");209210    forge::cpu::gelu(x, ref);211    forge::metal::gelu(x, gpu);212    forge::metal::sync();213    expect_close(gpu, ref, "gelu");214215    // add_bias: [N, C] + [C]216    const int64_t rows = 513, C = 127;217    forge::Tensor xb = forge::Tensor::empty({rows, C});218    forge::Tensor bias = forge::Tensor::empty({C});219    fill_random(xb, rng);220    fill_random(bias, rng);221    forge::Tensor gpub = forge::Tensor::empty({rows, C});222    forge::Tensor refb = forge::Tensor::empty({rows, C});223    forge::cpu::add_bias(xb, bias, refb);224    forge::metal::add_bias(xb, bias, gpub);225    forge::metal::sync();226    expect_close(gpub, refb, "add_bias");227}228229void test_metal_rowops() {230    std::printf("metal softmax/norm parity\n");231    std::mt19937 rng(31337);232    // Row lengths: tiny (< one simdgroup), odd, large (> threadgroup size,233    // vocab-like).234    for (int64_t C : {7LL, 63LL, 384LL, 4099LL}) {235        const int64_t rows = 129;236        forge::Tensor x = forge::Tensor::empty({rows, C});237        fill_random(x, rng, -8.0f, 8.0f);238        forge::Tensor gpu = forge::Tensor::empty({rows, C});239        forge::Tensor ref = forge::Tensor::empty({rows, C});240241        char label[64];242        forge::cpu::softmax(x, ref);243        forge::metal::softmax(x, gpu);244        forge::metal::sync();245        std::snprintf(label, sizeof(label), "softmax C=%lld", C);246        expect_close(gpu, ref, label);247248        forge::Tensor w = forge::Tensor::empty({C});249        forge::Tensor b = forge::Tensor::empty({C});250        fill_random(w, rng);251        fill_random(b, rng);252253        forge::cpu::rmsnorm(x, w, 1e-6f, ref);254        forge::metal::rmsnorm(x, w, 1e-6f, gpu);255        forge::metal::sync();256        std::snprintf(label, sizeof(label), "rmsnorm C=%lld", C);257        expect_close(gpu, ref, label);258259        forge::cpu::layernorm(x, w, b, 1e-6f, ref);260        forge::metal::layernorm(x, w, b, 1e-6f, gpu);261        forge::metal::sync();262        std::snprintf(label, sizeof(label), "layernorm C=%lld", C);263        expect_close(gpu, ref, label);264    }265}266267void test_metal_moe_quant_ops() {268    std::printf("metal QAT/MoE op parity\n");269    std::mt19937 rng(4242);270271    // fake_quant: both modes, odd and wide rows272    for (int mode : {0, 1}) {273        forge::Tensor w = forge::Tensor::empty({33, 257});274        fill_random(w, rng, -0.2f, 0.2f);275        forge::Tensor ref = forge::Tensor::empty(w.shape());276        forge::Tensor gpu = forge::Tensor::empty(w.shape());277        forge::cpu::fake_quant(w, mode, ref);278        forge::metal::fake_quant(w, mode, gpu);279        forge::metal::sync();280        expect_close(gpu, ref, mode == 0 ? "fake_quant int8" : "fake_quant ternary");281    }282283    // topk_renorm forward + backward (biased selection + both norm modes),284    // row_scale family285    const int64_t N = 130, E = 8, C = 37, K = 2;286    forge::Tensor probs = forge::Tensor::empty({N, E});287    fill_random(probs, rng, 0.01f, 1.0f);288    forge::Tensor bias = forge::Tensor::empty({E});289    fill_random(bias, rng, -0.3f, 0.3f);290    forge::Tensor ref = forge::Tensor::empty({N, E});291    forge::Tensor gpu = forge::Tensor::empty({N, E});292    forge::cpu::topk_renorm(probs, bias, K, true, ref);293    forge::metal::topk_renorm(probs, bias, K, true, gpu);294    forge::metal::sync();295    expect_close(gpu, ref, "topk_renorm fwd (biased)");296297    forge::Tensor dout = forge::Tensor::empty({N, E});298    fill_random(dout, rng);299    forge::Tensor dref = forge::Tensor::zeros({N, E});300    forge::Tensor dgpu = forge::Tensor::zeros({N, E});301    forge::cpu::topk_renorm_backward(probs, bias, dout, K, true, dref);302    forge::metal::topk_renorm_backward(probs, bias, dout, K, true, dgpu);303    forge::metal::sync();304    expect_close(dgpu, dref, "topk_renorm bwd (biased)");305306    forge::Tensor ref_nn = forge::Tensor::empty({N, E});307    forge::Tensor gpu_nn = forge::Tensor::empty({N, E});308    forge::cpu::topk_renorm(probs, bias, K, false, ref_nn);309    forge::metal::topk_renorm(probs, bias, K, false, gpu_nn);310    forge::metal::sync();311    expect_close(gpu_nn, ref_nn, "topk fwd (no renorm)");312313    forge::Tensor sig_ref = forge::Tensor::empty({N, E});314    forge::Tensor sig_gpu = forge::Tensor::empty({N, E});315    forge::cpu::sigmoid(probs, sig_ref);316    forge::metal::sigmoid(probs, sig_gpu);317    forge::metal::sync();318    expect_close(sig_gpu, sig_ref, "sigmoid fwd");319320    forge::Tensor cnt_ref = forge::Tensor::zeros({E});321    forge::Tensor cnt_gpu = forge::Tensor::zeros({E});322    forge::cpu::expert_counts(ref, cnt_ref);323    forge::metal::expert_counts(ref, cnt_gpu);324    forge::metal::sync();325    expect_close(cnt_gpu, cnt_ref, "expert_counts");326327    forge::Tensor x = forge::Tensor::empty({N, C});328    fill_random(x, rng);329    forge::Tensor yref = forge::Tensor::empty({N, C});330    forge::Tensor ygpu = forge::Tensor::empty({N, C});331    forge::cpu::row_scale(x, ref, 3, yref);332    forge::metal::row_scale(x, ref, 3, ygpu);333    forge::metal::sync();334    expect_close(ygpu, yref, "row_scale fwd");335336    forge::Tensor aref = forge::Tensor::zeros({N, C});337    forge::Tensor agpu = forge::Tensor::zeros({N, C});338    forge::cpu::row_scale_accumulate(x, ref, 3, aref);339    forge::metal::row_scale_accumulate(x, ref, 3, agpu);340    forge::metal::sync();341    expect_close(agpu, aref, "row_scale accumulate");342343    forge::Tensor dxo = forge::Tensor::empty({N, C});344    fill_random(dxo, rng);345    forge::Tensor gref = forge::Tensor::zeros({N, E});346    forge::Tensor ggpu = forge::Tensor::zeros({N, E});347    forge::cpu::row_scale_gate_backward(dxo, x, 3, gref);348    forge::metal::row_scale_gate_backward(dxo, x, 3, ggpu);349    forge::metal::sync();350    expect_close(ggpu, gref, "row_scale gate bwd");351352    // generic softmax backward (router-sized rows)353    forge::Tensor sm = forge::Tensor::empty({N, E});354    forge::cpu::softmax(probs, sm);355    forge::Tensor sref = forge::Tensor::zeros({N, E});356    forge::Tensor sgpu = forge::Tensor::zeros({N, E});357    forge::cpu::softmax_backward(sm, dout, sref);358    forge::metal::softmax_backward(sm, dout, sgpu);359    forge::metal::sync();360    expect_close(sgpu, sref, "softmax bwd");361}362363void test_batched_encoding() {364    std::printf("batched encoding (many dispatches, one sync)\n");365    std::mt19937 rng(55);366    const int64_t n = 4096;367    forge::Tensor x = forge::Tensor::empty({n});368    fill_random(x, rng);369    // chain of 20 dependent ops in ONE command buffer: out = ((x+x)*0.5) etc.370    forge::Tensor cur = forge::Tensor::empty({n});371    forge::metal::add(x, x, cur); // cur = 2x372    for (int i = 0; i < 19; ++i) {373        forge::Tensor next = forge::Tensor::empty({n});374        forge::metal::scale(cur, 0.9f, next);375        cur = next;376    }377    forge::metal::sync();378    // expected: 2 * 0.9^19 * x379    const float k = 2.0f * std::pow(0.9f, 19.0f);380    forge::Tensor ref = forge::Tensor::empty({n});381    forge::cpu::scale(x, k, ref);382    expect_close(cur, ref, "20-dispatch serial chain", 1e-5f);383}384385// Full-model forward+backward on CPU vs Metal: one test that exercises386// every training kernel (embedding fwd/bwd, rope, attention fwd/dq/dkv,387// matmul all variants, norms fwd/bwd, activations fwd/bwd, CE fwd+grad,388// accumulate/axpy) through the real autograd tape.389void test_backend_parity_model(const char* label, const forge::ModelConfig& cfg) {390    std::printf("backend parity: %s\n", label);391    forge::nn::Transformer model(cfg, 7);392393    std::mt19937_64 rng(21);394    forge::Tensor ids = forge::Tensor::empty({2, 6}, forge::DType::I32);395    forge::Tensor targets = forge::Tensor::empty({2 * 6}, forge::DType::I32);396    forge::cpu::fill_uniform_int(ids, 0, cfg.vocab_size, rng);397    forge::cpu::fill_uniform_int(targets, 0, cfg.vocab_size, rng);398    targets.data<int32_t>()[5] = -1; // exercise ignore_index399400    struct Saved { std::string name; forge::Tensor grad; };401    std::vector<Saved> cpu_grads;402    float cpu_loss = 0.0f;403404    {405        forge::ops::set_backend(forge::ops::Backend::CPU);406        model.zero_grad();407        forge::Var loss = model.loss(ids, targets);408        forge::Tape::get().backward(loss);409        cpu_loss = loss.value().data<float>()[0];410        std::unordered_map<const void*, bool> seen;411        for (const auto& [name, p] : model.named_parameters()) {412            if (!seen.emplace(p.id(), true).second) continue;413            forge::Tensor copy = forge::Tensor::empty(p.grad().shape());414            std::memcpy(copy.raw(), p.grad().raw(), copy.nbytes());415            cpu_grads.push_back({name, copy});416        }417    }418419    float metal_loss = 0.0f;420    {421        forge::ops::set_backend(forge::ops::Backend::Metal);422        model.zero_grad();423        forge::Var loss = model.loss(ids, targets);424        forge::Tape::get().backward(loss);425        forge::metal::sync();426        metal_loss = loss.value().data<float>()[0];427        forge::ops::set_backend(forge::ops::Backend::CPU);428    }429430    char lbl[96];431    std::snprintf(lbl, sizeof(lbl), "%s loss (cpu %.5f vs gpu %.5f)", label,432                  double(cpu_loss), double(metal_loss));433    expect(std::fabs(cpu_loss - metal_loss) <= 1e-4f, lbl);434435    size_t idx = 0;436    float worst = 0.0f;437    std::string worst_name;438    std::unordered_map<const void*, bool> seen;439    for (const auto& [name, p] : model.named_parameters()) {440        if (!seen.emplace(p.id(), true).second) continue;441        const forge::Tensor& want = cpu_grads[idx++].grad;442        const float* pg = p.grad().data<float>();443        const float* pw = want.data<float>();444        for (int64_t i = 0; i < want.numel(); ++i) {445            const float e = std::fabs(pg[i] - pw[i]);446            if (e > worst) { worst = e; worst_name = name; }447        }448    }449    std::snprintf(lbl, sizeof(lbl), "%s all grads (worst %.2e @ %s)", label,450                  double(worst), worst_name.c_str());451    expect(worst <= kTol, lbl);452}453454// Fused flash attention vs the CPU reference: forward output, the saved455// logsumexp, and all three input gradients. Covers MHA and GQA, causal and456// non-causal, and a head_dim on the supported list (the model-level parity457// tests above use head_dim 8, which deliberately falls back to the unfused458// kernel — so without this the fused path would go untested).459void test_flash_attention() {460    std::printf("flash attention parity (fused vs cpu reference)\n");461    std::mt19937 rng(2024);462463    struct Case { int64_t B, T, H, HKV, HD; bool causal; const char* tag; };464    const Case cases[] = {465        {2, 37, 4, 4, 64, true,  "MHA causal hd64 (ragged T)"},466        {2, 64, 4, 2, 64, true,  "GQA causal hd64 (2q/1kv)"},467        {1, 33, 2, 2, 32, true,  "MHA causal hd32"},468        {2, 24, 2, 2, 64, false, "MHA non-causal hd64"},469        {1, 40, 1, 1, 128, true, "single head hd128"},470    };471472    for (const auto& c : cases) {473        const int64_t Cq = c.H * c.HD, Ckv = c.HKV * c.HD;474        forge::Tensor q = forge::Tensor::empty({c.B, c.T, Cq});475        forge::Tensor k = forge::Tensor::empty({c.B, c.T, Ckv});476        forge::Tensor v = forge::Tensor::empty({c.B, c.T, Ckv});477        fill_random(q, rng, -1.5f, 1.5f);478        fill_random(k, rng, -1.5f, 1.5f);479        fill_random(v, rng, -1.5f, 1.5f);480        const float scale = 1.0f / std::sqrt(float(c.HD));481482        forge::Tensor ref_out = forge::Tensor::empty({c.B, c.T, Cq});483        forge::Tensor probs = forge::Tensor::empty({c.B, c.H, c.T, c.T});484        forge::cpu::attention(q, k, v, c.H, c.HKV, c.causal, scale, ref_out, &probs);485486        forge::Tensor gpu_out = forge::Tensor::empty({c.B, c.T, Cq});487        forge::Tensor lse = forge::Tensor::empty({c.B, c.H, c.T});488        forge::metal::flash_attention(q, k, v, c.H, c.HKV, c.causal, scale, gpu_out, lse);489        forge::metal::sync();490491        char label[96];492        std::snprintf(label, sizeof(label), "%s fwd", c.tag);493        expect_close(gpu_out, ref_out, label);494495        // backward: same upstream gradient into both paths496        forge::Tensor dout = forge::Tensor::empty({c.B, c.T, Cq});497        fill_random(dout, rng);498499        forge::Tensor rdq = forge::Tensor::zeros({c.B, c.T, Cq});500        forge::Tensor rdk = forge::Tensor::zeros({c.B, c.T, Ckv});501        forge::Tensor rdv = forge::Tensor::zeros({c.B, c.T, Ckv});502        forge::cpu::attention_backward(q, k, v, probs, ref_out, dout, c.H, c.HKV, scale,503                                       rdq, rdk, rdv);504505        forge::Tensor gdq = forge::Tensor::zeros({c.B, c.T, Cq});506        forge::Tensor gdk = forge::Tensor::zeros({c.B, c.T, Ckv});507        forge::Tensor gdv = forge::Tensor::zeros({c.B, c.T, Ckv});508        forge::metal::flash_attention_backward(q, k, v, gpu_out, lse, dout, c.H, c.HKV,509                                               c.causal, scale, gdq, gdk, gdv);510        forge::metal::sync();511512        std::snprintf(label, sizeof(label), "%s dq", c.tag);513        expect_close(gdq, rdq, label);514        std::snprintf(label, sizeof(label), "%s dk", c.tag);515        expect_close(gdk, rdk, label);516        std::snprintf(label, sizeof(label), "%s dv", c.tag);517        expect_close(gdv, rdv, label);518    }519}520521void test_adamw_kernel() {522    std::printf("adamw kernel parity\n");523    std::mt19937 rng(99);524    const int64_t n = 4097;525    forge::Tensor w_gpu = forge::Tensor::empty({n});526    forge::Tensor g = forge::Tensor::empty({n});527    fill_random(w_gpu, rng);528    fill_random(g, rng);529    forge::Tensor w_cpu = forge::Tensor::empty({n});530    std::memcpy(w_cpu.raw(), w_gpu.raw(), w_gpu.nbytes());531    forge::Tensor m_gpu = forge::Tensor::zeros({n});532    forge::Tensor v_gpu = forge::Tensor::zeros({n});533    forge::Tensor m_cpu = forge::Tensor::zeros({n});534    forge::Tensor v_cpu = forge::Tensor::zeros({n});535536    const float lr = 1e-3f, b1 = 0.9f, b2 = 0.95f, eps = 1e-8f, wd = 0.1f, gs = 0.7f;537    for (int64_t t = 1; t <= 3; ++t) {538        forge::metal::adamw_step(w_gpu, g, m_gpu, v_gpu, lr, b1, b2, t, eps, wd, gs);539        forge::metal::sync();540        const float bc1 = 1.0f - std::pow(b1, float(t));541        const float bc2 = 1.0f - std::pow(b2, float(t));542        float* w = w_cpu.data<float>();543        float* m = m_cpu.data<float>();544        float* v = v_cpu.data<float>();545        const float* gp = g.data<float>();546        for (int64_t i = 0; i < n; ++i) {547            const float grad = gp[i] * gs;548            m[i] = b1 * m[i] + (1 - b1) * grad;549            v[i] = b2 * v[i] + (1 - b2) * grad * grad;550            w[i] -= lr * ((m[i] / bc1) / (std::sqrt(v[i] / bc2) + eps) + wd * w[i]);551        }552    }553    expect_close(w_gpu, w_cpu, "adamw 3 steps", 1e-5f);554555    // sumsq + sum reductions556    forge::Tensor out = forge::Tensor::empty({1});557    forge::metal::sumsq(g, out);558    forge::metal::sync();559    double want = 0.0;560    for (int64_t i = 0; i < n; ++i) {561        const double x = double(g.data<float>()[i]);562        want += x * x;563    }564    expect(std::fabs(out.data<float>()[0] - float(want)) <= 1e-2f, "sumsq");565    forge::metal::sum(g, out, 0.5f);566    forge::metal::sync();567    double s = 0.0;568    for (int64_t i = 0; i < n; ++i) s += double(g.data<float>()[i]);569    expect(std::fabs(out.data<float>()[0] - float(s * 0.5)) <= 1e-2f, "sum");570}571572} // namespace573574int main() {575    NS::AutoreleasePool* pool = NS::AutoreleasePool::alloc()->init();576577    test_tensor_basics();578    test_cpu_matmul();579    test_metal_matmul();580    test_metal_elementwise();581    test_metal_rowops();582    test_batched_encoding();583584    {585        forge::ModelConfig cfg;586        cfg.n_layers = 2; cfg.d_model = 16; cfg.n_heads = 2; cfg.n_kv_heads = 1;587        cfg.d_ff = 24; cfg.vocab_size = 11; cfg.context_length = 8;588        cfg.tied_embeddings = true;589        test_backend_parity_model("rmsnorm/swiglu/rope/gqa/tied [unfused hd8]", cfg);590591        // head_dim 64 -> the whole model runs through the fused attention path592        forge::ModelConfig fcfg;593        fcfg.n_layers = 2; fcfg.d_model = 128; fcfg.n_heads = 2; fcfg.n_kv_heads = 1;594        fcfg.d_ff = 96; fcfg.vocab_size = 11; fcfg.context_length = 8;595        fcfg.tied_embeddings = true;596        test_backend_parity_model("full model via fused attention [hd64]", fcfg);597598        cfg.norm = "layernorm"; cfg.activation = "gelu"; cfg.use_rope = false;599        cfg.n_kv_heads = 2; cfg.tied_embeddings = false;600        test_backend_parity_model("layernorm/gelu/pos-emb/mha", cfg);601602        // QAT: full model with fake-quantized linears (STE backward)603        forge::ModelConfig qcfg;604        qcfg.n_layers = 2; qcfg.d_model = 16; qcfg.n_heads = 2; qcfg.n_kv_heads = 1;605        qcfg.d_ff = 24; qcfg.vocab_size = 11; qcfg.context_length = 8;606        qcfg.tied_embeddings = true;607        qcfg.quant = "ternary";608        test_backend_parity_model("QAT ternary linears", qcfg);609        qcfg.quant = "int8";610        test_backend_parity_model("QAT int8 linears", qcfg);611612        // MoE: router + top-2 of 4 experts + 1 shared expert + aux loss613        forge::ModelConfig mcfg;614        mcfg.n_layers = 2; mcfg.d_model = 16; mcfg.n_heads = 2; mcfg.n_kv_heads = 1;615        mcfg.d_ff = 24; mcfg.vocab_size = 11; mcfg.context_length = 8;616        mcfg.tied_embeddings = true;617        mcfg.n_experts = 4; mcfg.moe_top_k = 2; mcfg.n_shared_experts = 1;618        test_backend_parity_model("MoE 4+1shared top-2 + aux", mcfg);619620        // V3-style MoE: sigmoid scoring, no renorm, routed scaling, own d_ff,621        // first layer dense, noaux bias tracking on622        mcfg.moe_scoring = "sigmoid"; mcfg.moe_norm_topk = false;623        mcfg.routed_scaling_factor = 2.5f; mcfg.moe_d_ff = 16;624        mcfg.first_k_dense = 1; mcfg.moe_bias_gamma = 0.001f;625        mcfg.moe_aux_weight = 0.0f;626        test_backend_parity_model("MoE V3-style sigmoid/noaux/scaled", mcfg);627628        // Architecture-variant knobs: QK-norm + logit softcap + embed scaling629        forge::ModelConfig vcfg;630        vcfg.n_layers = 2; vcfg.d_model = 16; vcfg.n_heads = 2; vcfg.n_kv_heads = 1;631        vcfg.d_ff = 24; vcfg.vocab_size = 11; vcfg.context_length = 8;632        vcfg.tied_embeddings = true;633        vcfg.qk_norm = true; vcfg.final_softcap = 30.0f; vcfg.scale_embeddings = true;634        test_backend_parity_model("qk-norm + softcap + embed-scale", vcfg);635636        // Wave-1 knobs: ReLU² + QKV bias + decoupled head_dim + NoPE + rope scaling637        forge::ModelConfig w1;638        w1.n_layers = 2; w1.d_model = 16; w1.n_heads = 2; w1.n_kv_heads = 1;639        w1.d_ff = 24; w1.vocab_size = 11; w1.context_length = 8;640        w1.tied_embeddings = true;641        w1.activation = "relu2"; w1.attention_bias = true;642        w1.head_dim_override = 10; w1.nope_every = 2;643        w1.rope_scale_factor = 8.0f; w1.rope_scale_orig_ctx = 4;644        test_backend_parity_model("relu2 + qkv-bias + head_dim10 + nope + rope-scale", w1);645646        // Norm placements: OLMo2 post and Gemma sandwich647        forge::ModelConfig np;648        np.n_layers = 2; np.d_model = 16; np.n_heads = 2; np.n_kv_heads = 2;649        np.d_ff = 24; np.vocab_size = 11; np.context_length = 8;650        np.tied_embeddings = true;651        np.norm_placement = "post";652        test_backend_parity_model("post-norm (OLMo2)", np);653        np.norm_placement = "sandwich";654        test_backend_parity_model("sandwich norm (Gemma)", np);655656        // Wave-2: sliding window (Mistral) through the fused scalar path657        // (head_dim 64), + local/global pattern, + Gemma3 dual rope theta658        forge::ModelConfig sw;659        sw.n_layers = 3; sw.d_model = 128; sw.n_heads = 2; sw.n_kv_heads = 1;660        sw.d_ff = 96; sw.vocab_size = 11; sw.context_length = 8;661        sw.tied_embeddings = true;662        sw.sliding_window = 3; sw.sliding_global_every = 2;663        sw.rope_theta_global = 1e6f;664        test_backend_parity_model("sliding-window 3 + global-every-2 [hd64]", sw);665666        // Wave-2: attention softcap (Gemma2) — unfused path on both backends667        forge::ModelConfig sc;668        sc.n_layers = 2; sc.d_model = 16; sc.n_heads = 2; sc.n_kv_heads = 1;669        sc.d_ff = 24; sc.vocab_size = 11; sc.context_length = 8;670        sc.tied_embeddings = true;671        sc.attn_softcap = 20.0f;672        test_backend_parity_model("attn softcap 20 (Gemma2)", sc);673674        // Sliding window through the UNFUSED kernel too (head_dim 8)675        forge::ModelConfig swu;676        swu.n_layers = 2; swu.d_model = 16; swu.n_heads = 2; swu.n_kv_heads = 2;677        swu.d_ff = 24; swu.vocab_size = 11; swu.context_length = 8;678        swu.tied_embeddings = true;679        swu.sliding_window = 3; swu.attn_softcap = 20.0f;680        test_backend_parity_model("sliding-window 3 + softcap [unfused hd8]", swu);681    }682    test_metal_moe_quant_ops();683    test_flash_attention();684    test_adamw_kernel();685686    pool->drain();687688    if (g_failures) {689        std::printf("\n%d test(s) FAILED\n", g_failures);690        return 1;691    }692    std::printf("\nall tests passed\n");693    return 0;694}695