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%
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