// Author: Simon-Pierre Boucher — contact@spboucher.ai // // CPU-vs-Metal parity for every GPU kernel (CLAUDE.md protocol #1: max abs // error <= 1e-4 for f32), plus tensor mechanics and the CPU matmul oracle // check. GPU ops encode into the batched Stream; comparisons happen after // metal::sync() — the readback boundary. #include #include #include "core/tensor.h" #include "nn/transformer.h" #include "ops/cpu/cpu_ops.h" #include "ops/metal/metal_ops.h" #include "ops/ops.h" #include #include #include #include #include namespace { int g_failures = 0; constexpr float kTol = 1e-4f; void expect(bool cond, const char* what) { if (cond) { std::printf(" ok: %s\n", what); } else { std::printf(" FAIL: %s\n", what); ++g_failures; } } void expect_close(const forge::Tensor& got, const forge::Tensor& want, const char* what, float tol = kTol) { const float* pg = got.data(); const float* pw = want.data(); float m = 0.0f; for (int64_t i = 0; i < got.numel(); ++i) m = std::max(m, std::fabs(pg[i] - pw[i])); if (m <= tol) { std::printf(" ok: %s (max abs err %.2e)\n", what, double(m)); } else { std::printf(" FAIL: %s (max abs err %.2e > %.0e)\n", what, double(m), double(tol)); ++g_failures; } } void fill_random(forge::Tensor& t, std::mt19937& rng, float lo = -1.0f, float hi = 1.0f) { std::uniform_real_distribution dist(lo, hi); float* p = t.data(); for (int64_t i = 0; i < t.numel(); ++i) p[i] = dist(rng); } // Deliberately dumb oracle: index-arithmetic triple loop with double acc. void matmul_oracle(const forge::Tensor& a, const forge::Tensor& b, forge::Tensor& c, bool ta, bool tb) { const int64_t M = c.size(0), N = c.size(1); const int64_t K = ta ? a.size(0) : a.size(1); const int64_t lda = a.size(1), ldb = b.size(1); const float* A = a.data(); const float* B = b.data(); float* C = c.data(); for (int64_t i = 0; i < M; ++i) for (int64_t j = 0; j < N; ++j) { double acc = 0.0; for (int64_t k = 0; k < K; ++k) { const float av = ta ? A[k * lda + i] : A[i * lda + k]; const float bv = tb ? B[j * ldb + k] : B[k * ldb + j]; acc += double(av) * double(bv); } C[i * N + j] = float(acc); } } void test_tensor_basics() { std::printf("tensor basics\n"); forge::Tensor t = forge::Tensor::zeros({4, 8}); expect(t.numel() == 32, "numel"); expect(t.is_contiguous(), "contiguous"); expect(t.strides()[0] == 8 && t.strides()[1] == 1, "row-major strides"); forge::Tensor v = t.view({8, 4}); v.set_item(0, 42.0f); expect(t.item_at(0) == 42.0f, "view shares storage"); forge::Tensor s = t.slice0(1, 2); s.set_item(0, 7.0f); expect(t.item_at(8) == 7.0f, "slice0 shares storage at offset"); forge::Tensor h = forge::Tensor::full({3}, 1.5f, forge::DType::F16); expect(std::fabs(h.item_at(2) - 1.5f) < 1e-6f, "f16 roundtrip"); forge::Tensor bf = forge::Tensor::full({3}, 1.5f, forge::DType::BF16); expect(std::fabs(bf.item_at(1) - 1.5f) < 1e-6f, "bf16 roundtrip"); } void test_cpu_matmul() { std::printf("cpu matmul vs oracle (all transpose variants)\n"); std::mt19937 rng(1234); const int64_t M = 17, K = 23, N = 13; for (int ta = 0; ta <= 1; ++ta) for (int tb = 0; tb <= 1; ++tb) { forge::Tensor a = ta ? forge::Tensor::empty({K, M}) : forge::Tensor::empty({M, K}); forge::Tensor b = tb ? forge::Tensor::empty({N, K}) : forge::Tensor::empty({K, N}); forge::Tensor c = forge::Tensor::empty({M, N}); forge::Tensor ref = forge::Tensor::empty({M, N}); fill_random(a, rng); fill_random(b, rng); forge::cpu::matmul(a, b, c, ta, tb); matmul_oracle(a, b, ref, ta, tb); char label[64]; std::snprintf(label, sizeof(label), "cpu matmul ta=%d tb=%d", ta, tb); expect_close(c, ref, label, 1e-5f); } } void test_metal_matmul() { std::printf("metal matmul parity (all kernels, all transpose variants)\n"); std::mt19937 rng(4321); using K_t = forge::metal::MatmulKernel; struct KernelCase { K_t kernel; const char* name; }; const KernelCase kernels[] = {{K_t::Naive, "naive"}, {K_t::Tiled, "tiled"}, {K_t::Simdgroup, "simd"}}; // Ragged (non-multiple of any tile dim) and exactly-tiled shapes, so the // simdgroup kernel's predicated and ALIGNED fast paths both get covered. struct Shape { int64_t M, K, N; const char* tag; }; const Shape shapes[] = {{67, 129, 45, "ragged"}, {128, 64, 192, "aligned"}}; for (const auto& kc : kernels) for (const auto& sh : shapes) for (int ta = 0; ta <= 1; ++ta) for (int tb = 0; tb <= 1; ++tb) { forge::Tensor a = ta ? forge::Tensor::empty({sh.K, sh.M}) : forge::Tensor::empty({sh.M, sh.K}); forge::Tensor b = tb ? forge::Tensor::empty({sh.N, sh.K}) : forge::Tensor::empty({sh.K, sh.N}); forge::Tensor gpu = forge::Tensor::empty({sh.M, sh.N}); forge::Tensor ref = forge::Tensor::empty({sh.M, sh.N}); fill_random(a, rng); fill_random(b, rng); forge::cpu::matmul(a, b, ref, ta, tb); forge::metal::matmul(a, b, gpu, ta, tb, false, kc.kernel); forge::metal::sync(); char label[80]; std::snprintf(label, sizeof(label), "%s matmul %s ta=%d tb=%d", kc.name, sh.tag, ta, tb); expect_close(gpu, ref, label); } // accumulate flag (ragged shape: both epilogue paths get predication) const int64_t M = 67, K = 129, N = 45; forge::Tensor a = forge::Tensor::empty({M, K}); forge::Tensor b = forge::Tensor::empty({K, N}); forge::Tensor acc_gpu = forge::Tensor::empty({M, N}); forge::Tensor acc_ref = forge::Tensor::empty({M, N}); fill_random(a, rng); fill_random(b, rng); fill_random(acc_gpu, rng); std::memcpy(acc_ref.raw(), acc_gpu.raw(), acc_gpu.nbytes()); forge::cpu::matmul(a, b, acc_ref, false, false, true); forge::metal::matmul(a, b, acc_gpu, false, false, true, K_t::Tiled); forge::metal::sync(); expect_close(acc_gpu, acc_ref, "tiled matmul accumulate=true"); // simdgroup accumulate takes the staged (non-fast) epilogue path forge::Tensor sacc_gpu = forge::Tensor::empty({M, N}); forge::Tensor sacc_ref = forge::Tensor::empty({M, N}); fill_random(sacc_gpu, rng); std::memcpy(sacc_ref.raw(), sacc_gpu.raw(), sacc_gpu.nbytes()); forge::cpu::matmul(a, b, sacc_ref, false, false, true); forge::metal::matmul(a, b, sacc_gpu, false, false, true, K_t::Simdgroup); forge::metal::sync(); expect_close(sacc_gpu, sacc_ref, "simd matmul accumulate=true"); } void test_metal_elementwise() { std::printf("metal elementwise parity\n"); std::mt19937 rng(777); const int64_t n = 65537; // non-multiple of the threadgroup size forge::Tensor x = forge::Tensor::empty({n}); forge::Tensor y = forge::Tensor::empty({n}); fill_random(x, rng, -4.0f, 4.0f); fill_random(y, rng, -4.0f, 4.0f); forge::Tensor gpu = forge::Tensor::empty({n}); forge::Tensor ref = forge::Tensor::empty({n}); forge::cpu::add(x, y, ref); forge::metal::add(x, y, gpu); forge::metal::sync(); expect_close(gpu, ref, "add"); forge::cpu::mul(x, y, ref); forge::metal::mul(x, y, gpu); forge::metal::sync(); expect_close(gpu, ref, "mul"); forge::cpu::scale(x, 0.37f, ref); forge::metal::scale(x, 0.37f, gpu); forge::metal::sync(); expect_close(gpu, ref, "scale"); forge::cpu::silu(x, ref); forge::metal::silu(x, gpu); forge::metal::sync(); expect_close(gpu, ref, "silu"); forge::cpu::gelu(x, ref); forge::metal::gelu(x, gpu); forge::metal::sync(); expect_close(gpu, ref, "gelu"); // add_bias: [N, C] + [C] const int64_t rows = 513, C = 127; forge::Tensor xb = forge::Tensor::empty({rows, C}); forge::Tensor bias = forge::Tensor::empty({C}); fill_random(xb, rng); fill_random(bias, rng); forge::Tensor gpub = forge::Tensor::empty({rows, C}); forge::Tensor refb = forge::Tensor::empty({rows, C}); forge::cpu::add_bias(xb, bias, refb); forge::metal::add_bias(xb, bias, gpub); forge::metal::sync(); expect_close(gpub, refb, "add_bias"); } void test_metal_rowops() { std::printf("metal softmax/norm parity\n"); std::mt19937 rng(31337); // Row lengths: tiny (< one simdgroup), odd, large (> threadgroup size, // vocab-like). for (int64_t C : {7LL, 63LL, 384LL, 4099LL}) { const int64_t rows = 129; forge::Tensor x = forge::Tensor::empty({rows, C}); fill_random(x, rng, -8.0f, 8.0f); forge::Tensor gpu = forge::Tensor::empty({rows, C}); forge::Tensor ref = forge::Tensor::empty({rows, C}); char label[64]; forge::cpu::softmax(x, ref); forge::metal::softmax(x, gpu); forge::metal::sync(); std::snprintf(label, sizeof(label), "softmax C=%lld", C); expect_close(gpu, ref, label); forge::Tensor w = forge::Tensor::empty({C}); forge::Tensor b = forge::Tensor::empty({C}); fill_random(w, rng); fill_random(b, rng); forge::cpu::rmsnorm(x, w, 1e-6f, ref); forge::metal::rmsnorm(x, w, 1e-6f, gpu); forge::metal::sync(); std::snprintf(label, sizeof(label), "rmsnorm C=%lld", C); expect_close(gpu, ref, label); forge::cpu::layernorm(x, w, b, 1e-6f, ref); forge::metal::layernorm(x, w, b, 1e-6f, gpu); forge::metal::sync(); std::snprintf(label, sizeof(label), "layernorm C=%lld", C); expect_close(gpu, ref, label); } } void test_metal_moe_quant_ops() { std::printf("metal QAT/MoE op parity\n"); std::mt19937 rng(4242); // fake_quant: both modes, odd and wide rows for (int mode : {0, 1}) { forge::Tensor w = forge::Tensor::empty({33, 257}); fill_random(w, rng, -0.2f, 0.2f); forge::Tensor ref = forge::Tensor::empty(w.shape()); forge::Tensor gpu = forge::Tensor::empty(w.shape()); forge::cpu::fake_quant(w, mode, ref); forge::metal::fake_quant(w, mode, gpu); forge::metal::sync(); expect_close(gpu, ref, mode == 0 ? "fake_quant int8" : "fake_quant ternary"); } // topk_renorm forward + backward (biased selection + both norm modes), // row_scale family const int64_t N = 130, E = 8, C = 37, K = 2; forge::Tensor probs = forge::Tensor::empty({N, E}); fill_random(probs, rng, 0.01f, 1.0f); forge::Tensor bias = forge::Tensor::empty({E}); fill_random(bias, rng, -0.3f, 0.3f); forge::Tensor ref = forge::Tensor::empty({N, E}); forge::Tensor gpu = forge::Tensor::empty({N, E}); forge::cpu::topk_renorm(probs, bias, K, true, ref); forge::metal::topk_renorm(probs, bias, K, true, gpu); forge::metal::sync(); expect_close(gpu, ref, "topk_renorm fwd (biased)"); forge::Tensor dout = forge::Tensor::empty({N, E}); fill_random(dout, rng); forge::Tensor dref = forge::Tensor::zeros({N, E}); forge::Tensor dgpu = forge::Tensor::zeros({N, E}); forge::cpu::topk_renorm_backward(probs, bias, dout, K, true, dref); forge::metal::topk_renorm_backward(probs, bias, dout, K, true, dgpu); forge::metal::sync(); expect_close(dgpu, dref, "topk_renorm bwd (biased)"); forge::Tensor ref_nn = forge::Tensor::empty({N, E}); forge::Tensor gpu_nn = forge::Tensor::empty({N, E}); forge::cpu::topk_renorm(probs, bias, K, false, ref_nn); forge::metal::topk_renorm(probs, bias, K, false, gpu_nn); forge::metal::sync(); expect_close(gpu_nn, ref_nn, "topk fwd (no renorm)"); forge::Tensor sig_ref = forge::Tensor::empty({N, E}); forge::Tensor sig_gpu = forge::Tensor::empty({N, E}); forge::cpu::sigmoid(probs, sig_ref); forge::metal::sigmoid(probs, sig_gpu); forge::metal::sync(); expect_close(sig_gpu, sig_ref, "sigmoid fwd"); forge::Tensor cnt_ref = forge::Tensor::zeros({E}); forge::Tensor cnt_gpu = forge::Tensor::zeros({E}); forge::cpu::expert_counts(ref, cnt_ref); forge::metal::expert_counts(ref, cnt_gpu); forge::metal::sync(); expect_close(cnt_gpu, cnt_ref, "expert_counts"); forge::Tensor x = forge::Tensor::empty({N, C}); fill_random(x, rng); forge::Tensor yref = forge::Tensor::empty({N, C}); forge::Tensor ygpu = forge::Tensor::empty({N, C}); forge::cpu::row_scale(x, ref, 3, yref); forge::metal::row_scale(x, ref, 3, ygpu); forge::metal::sync(); expect_close(ygpu, yref, "row_scale fwd"); forge::Tensor aref = forge::Tensor::zeros({N, C}); forge::Tensor agpu = forge::Tensor::zeros({N, C}); forge::cpu::row_scale_accumulate(x, ref, 3, aref); forge::metal::row_scale_accumulate(x, ref, 3, agpu); forge::metal::sync(); expect_close(agpu, aref, "row_scale accumulate"); forge::Tensor dxo = forge::Tensor::empty({N, C}); fill_random(dxo, rng); forge::Tensor gref = forge::Tensor::zeros({N, E}); forge::Tensor ggpu = forge::Tensor::zeros({N, E}); forge::cpu::row_scale_gate_backward(dxo, x, 3, gref); forge::metal::row_scale_gate_backward(dxo, x, 3, ggpu); forge::metal::sync(); expect_close(ggpu, gref, "row_scale gate bwd"); // generic softmax backward (router-sized rows) forge::Tensor sm = forge::Tensor::empty({N, E}); forge::cpu::softmax(probs, sm); forge::Tensor sref = forge::Tensor::zeros({N, E}); forge::Tensor sgpu = forge::Tensor::zeros({N, E}); forge::cpu::softmax_backward(sm, dout, sref); forge::metal::softmax_backward(sm, dout, sgpu); forge::metal::sync(); expect_close(sgpu, sref, "softmax bwd"); } void test_batched_encoding() { std::printf("batched encoding (many dispatches, one sync)\n"); std::mt19937 rng(55); const int64_t n = 4096; forge::Tensor x = forge::Tensor::empty({n}); fill_random(x, rng); // chain of 20 dependent ops in ONE command buffer: out = ((x+x)*0.5) etc. forge::Tensor cur = forge::Tensor::empty({n}); forge::metal::add(x, x, cur); // cur = 2x for (int i = 0; i < 19; ++i) { forge::Tensor next = forge::Tensor::empty({n}); forge::metal::scale(cur, 0.9f, next); cur = next; } forge::metal::sync(); // expected: 2 * 0.9^19 * x const float k = 2.0f * std::pow(0.9f, 19.0f); forge::Tensor ref = forge::Tensor::empty({n}); forge::cpu::scale(x, k, ref); expect_close(cur, ref, "20-dispatch serial chain", 1e-5f); } // Full-model forward+backward on CPU vs Metal: one test that exercises // every training kernel (embedding fwd/bwd, rope, attention fwd/dq/dkv, // matmul all variants, norms fwd/bwd, activations fwd/bwd, CE fwd+grad, // accumulate/axpy) through the real autograd tape. void test_backend_parity_model(const char* label, const forge::ModelConfig& cfg) { std::printf("backend parity: %s\n", label); forge::nn::Transformer model(cfg, 7); std::mt19937_64 rng(21); forge::Tensor ids = forge::Tensor::empty({2, 6}, forge::DType::I32); forge::Tensor targets = forge::Tensor::empty({2 * 6}, forge::DType::I32); forge::cpu::fill_uniform_int(ids, 0, cfg.vocab_size, rng); forge::cpu::fill_uniform_int(targets, 0, cfg.vocab_size, rng); targets.data()[5] = -1; // exercise ignore_index struct Saved { std::string name; forge::Tensor grad; }; std::vector cpu_grads; float cpu_loss = 0.0f; { forge::ops::set_backend(forge::ops::Backend::CPU); model.zero_grad(); forge::Var loss = model.loss(ids, targets); forge::Tape::get().backward(loss); cpu_loss = loss.value().data()[0]; std::unordered_map seen; for (const auto& [name, p] : model.named_parameters()) { if (!seen.emplace(p.id(), true).second) continue; forge::Tensor copy = forge::Tensor::empty(p.grad().shape()); std::memcpy(copy.raw(), p.grad().raw(), copy.nbytes()); cpu_grads.push_back({name, copy}); } } float metal_loss = 0.0f; { forge::ops::set_backend(forge::ops::Backend::Metal); model.zero_grad(); forge::Var loss = model.loss(ids, targets); forge::Tape::get().backward(loss); forge::metal::sync(); metal_loss = loss.value().data()[0]; forge::ops::set_backend(forge::ops::Backend::CPU); } char lbl[96]; std::snprintf(lbl, sizeof(lbl), "%s loss (cpu %.5f vs gpu %.5f)", label, double(cpu_loss), double(metal_loss)); expect(std::fabs(cpu_loss - metal_loss) <= 1e-4f, lbl); size_t idx = 0; float worst = 0.0f; std::string worst_name; std::unordered_map seen; for (const auto& [name, p] : model.named_parameters()) { if (!seen.emplace(p.id(), true).second) continue; const forge::Tensor& want = cpu_grads[idx++].grad; const float* pg = p.grad().data(); const float* pw = want.data(); for (int64_t i = 0; i < want.numel(); ++i) { const float e = std::fabs(pg[i] - pw[i]); if (e > worst) { worst = e; worst_name = name; } } } std::snprintf(lbl, sizeof(lbl), "%s all grads (worst %.2e @ %s)", label, double(worst), worst_name.c_str()); expect(worst <= kTol, lbl); } // Fused flash attention vs the CPU reference: forward output, the saved // logsumexp, and all three input gradients. Covers MHA and GQA, causal and // non-causal, and a head_dim on the supported list (the model-level parity // tests above use head_dim 8, which deliberately falls back to the unfused // kernel — so without this the fused path would go untested). void test_flash_attention() { std::printf("flash attention parity (fused vs cpu reference)\n"); std::mt19937 rng(2024); struct Case { int64_t B, T, H, HKV, HD; bool causal; const char* tag; }; const Case cases[] = { {2, 37, 4, 4, 64, true, "MHA causal hd64 (ragged T)"}, {2, 64, 4, 2, 64, true, "GQA causal hd64 (2q/1kv)"}, {1, 33, 2, 2, 32, true, "MHA causal hd32"}, {2, 24, 2, 2, 64, false, "MHA non-causal hd64"}, {1, 40, 1, 1, 128, true, "single head hd128"}, }; for (const auto& c : cases) { const int64_t Cq = c.H * c.HD, Ckv = c.HKV * c.HD; forge::Tensor q = forge::Tensor::empty({c.B, c.T, Cq}); forge::Tensor k = forge::Tensor::empty({c.B, c.T, Ckv}); forge::Tensor v = forge::Tensor::empty({c.B, c.T, Ckv}); fill_random(q, rng, -1.5f, 1.5f); fill_random(k, rng, -1.5f, 1.5f); fill_random(v, rng, -1.5f, 1.5f); const float scale = 1.0f / std::sqrt(float(c.HD)); forge::Tensor ref_out = forge::Tensor::empty({c.B, c.T, Cq}); forge::Tensor probs = forge::Tensor::empty({c.B, c.H, c.T, c.T}); forge::cpu::attention(q, k, v, c.H, c.HKV, c.causal, scale, ref_out, &probs); forge::Tensor gpu_out = forge::Tensor::empty({c.B, c.T, Cq}); forge::Tensor lse = forge::Tensor::empty({c.B, c.H, c.T}); forge::metal::flash_attention(q, k, v, c.H, c.HKV, c.causal, scale, gpu_out, lse); forge::metal::sync(); char label[96]; std::snprintf(label, sizeof(label), "%s fwd", c.tag); expect_close(gpu_out, ref_out, label); // backward: same upstream gradient into both paths forge::Tensor dout = forge::Tensor::empty({c.B, c.T, Cq}); fill_random(dout, rng); forge::Tensor rdq = forge::Tensor::zeros({c.B, c.T, Cq}); forge::Tensor rdk = forge::Tensor::zeros({c.B, c.T, Ckv}); forge::Tensor rdv = forge::Tensor::zeros({c.B, c.T, Ckv}); forge::cpu::attention_backward(q, k, v, probs, ref_out, dout, c.H, c.HKV, scale, rdq, rdk, rdv); forge::Tensor gdq = forge::Tensor::zeros({c.B, c.T, Cq}); forge::Tensor gdk = forge::Tensor::zeros({c.B, c.T, Ckv}); forge::Tensor gdv = forge::Tensor::zeros({c.B, c.T, Ckv}); forge::metal::flash_attention_backward(q, k, v, gpu_out, lse, dout, c.H, c.HKV, c.causal, scale, gdq, gdk, gdv); forge::metal::sync(); std::snprintf(label, sizeof(label), "%s dq", c.tag); expect_close(gdq, rdq, label); std::snprintf(label, sizeof(label), "%s dk", c.tag); expect_close(gdk, rdk, label); std::snprintf(label, sizeof(label), "%s dv", c.tag); expect_close(gdv, rdv, label); } } void test_adamw_kernel() { std::printf("adamw kernel parity\n"); std::mt19937 rng(99); const int64_t n = 4097; forge::Tensor w_gpu = forge::Tensor::empty({n}); forge::Tensor g = forge::Tensor::empty({n}); fill_random(w_gpu, rng); fill_random(g, rng); forge::Tensor w_cpu = forge::Tensor::empty({n}); std::memcpy(w_cpu.raw(), w_gpu.raw(), w_gpu.nbytes()); forge::Tensor m_gpu = forge::Tensor::zeros({n}); forge::Tensor v_gpu = forge::Tensor::zeros({n}); forge::Tensor m_cpu = forge::Tensor::zeros({n}); forge::Tensor v_cpu = forge::Tensor::zeros({n}); const float lr = 1e-3f, b1 = 0.9f, b2 = 0.95f, eps = 1e-8f, wd = 0.1f, gs = 0.7f; for (int64_t t = 1; t <= 3; ++t) { forge::metal::adamw_step(w_gpu, g, m_gpu, v_gpu, lr, b1, b2, t, eps, wd, gs); forge::metal::sync(); const float bc1 = 1.0f - std::pow(b1, float(t)); const float bc2 = 1.0f - std::pow(b2, float(t)); float* w = w_cpu.data(); float* m = m_cpu.data(); float* v = v_cpu.data(); const float* gp = g.data(); for (int64_t i = 0; i < n; ++i) { const float grad = gp[i] * gs; m[i] = b1 * m[i] + (1 - b1) * grad; v[i] = b2 * v[i] + (1 - b2) * grad * grad; w[i] -= lr * ((m[i] / bc1) / (std::sqrt(v[i] / bc2) + eps) + wd * w[i]); } } expect_close(w_gpu, w_cpu, "adamw 3 steps", 1e-5f); // sumsq + sum reductions forge::Tensor out = forge::Tensor::empty({1}); forge::metal::sumsq(g, out); forge::metal::sync(); double want = 0.0; for (int64_t i = 0; i < n; ++i) { const double x = double(g.data()[i]); want += x * x; } expect(std::fabs(out.data()[0] - float(want)) <= 1e-2f, "sumsq"); forge::metal::sum(g, out, 0.5f); forge::metal::sync(); double s = 0.0; for (int64_t i = 0; i < n; ++i) s += double(g.data()[i]); expect(std::fabs(out.data()[0] - float(s * 0.5)) <= 1e-2f, "sum"); } } // namespace int main() { NS::AutoreleasePool* pool = NS::AutoreleasePool::alloc()->init(); test_tensor_basics(); test_cpu_matmul(); test_metal_matmul(); test_metal_elementwise(); test_metal_rowops(); test_batched_encoding(); { forge::ModelConfig cfg; cfg.n_layers = 2; cfg.d_model = 16; cfg.n_heads = 2; cfg.n_kv_heads = 1; cfg.d_ff = 24; cfg.vocab_size = 11; cfg.context_length = 8; cfg.tied_embeddings = true; test_backend_parity_model("rmsnorm/swiglu/rope/gqa/tied [unfused hd8]", cfg); // head_dim 64 -> the whole model runs through the fused attention path forge::ModelConfig fcfg; fcfg.n_layers = 2; fcfg.d_model = 128; fcfg.n_heads = 2; fcfg.n_kv_heads = 1; fcfg.d_ff = 96; fcfg.vocab_size = 11; fcfg.context_length = 8; fcfg.tied_embeddings = true; test_backend_parity_model("full model via fused attention [hd64]", fcfg); cfg.norm = "layernorm"; cfg.activation = "gelu"; cfg.use_rope = false; cfg.n_kv_heads = 2; cfg.tied_embeddings = false; test_backend_parity_model("layernorm/gelu/pos-emb/mha", cfg); // QAT: full model with fake-quantized linears (STE backward) forge::ModelConfig qcfg; qcfg.n_layers = 2; qcfg.d_model = 16; qcfg.n_heads = 2; qcfg.n_kv_heads = 1; qcfg.d_ff = 24; qcfg.vocab_size = 11; qcfg.context_length = 8; qcfg.tied_embeddings = true; qcfg.quant = "ternary"; test_backend_parity_model("QAT ternary linears", qcfg); qcfg.quant = "int8"; test_backend_parity_model("QAT int8 linears", qcfg); // MoE: router + top-2 of 4 experts + 1 shared expert + aux loss forge::ModelConfig mcfg; mcfg.n_layers = 2; mcfg.d_model = 16; mcfg.n_heads = 2; mcfg.n_kv_heads = 1; mcfg.d_ff = 24; mcfg.vocab_size = 11; mcfg.context_length = 8; mcfg.tied_embeddings = true; mcfg.n_experts = 4; mcfg.moe_top_k = 2; mcfg.n_shared_experts = 1; test_backend_parity_model("MoE 4+1shared top-2 + aux", mcfg); // V3-style MoE: sigmoid scoring, no renorm, routed scaling, own d_ff, // first layer dense, noaux bias tracking on mcfg.moe_scoring = "sigmoid"; mcfg.moe_norm_topk = false; mcfg.routed_scaling_factor = 2.5f; mcfg.moe_d_ff = 16; mcfg.first_k_dense = 1; mcfg.moe_bias_gamma = 0.001f; mcfg.moe_aux_weight = 0.0f; test_backend_parity_model("MoE V3-style sigmoid/noaux/scaled", mcfg); // Architecture-variant knobs: QK-norm + logit softcap + embed scaling forge::ModelConfig vcfg; vcfg.n_layers = 2; vcfg.d_model = 16; vcfg.n_heads = 2; vcfg.n_kv_heads = 1; vcfg.d_ff = 24; vcfg.vocab_size = 11; vcfg.context_length = 8; vcfg.tied_embeddings = true; vcfg.qk_norm = true; vcfg.final_softcap = 30.0f; vcfg.scale_embeddings = true; test_backend_parity_model("qk-norm + softcap + embed-scale", vcfg); // Wave-1 knobs: ReLU² + QKV bias + decoupled head_dim + NoPE + rope scaling forge::ModelConfig w1; w1.n_layers = 2; w1.d_model = 16; w1.n_heads = 2; w1.n_kv_heads = 1; w1.d_ff = 24; w1.vocab_size = 11; w1.context_length = 8; w1.tied_embeddings = true; w1.activation = "relu2"; w1.attention_bias = true; w1.head_dim_override = 10; w1.nope_every = 2; w1.rope_scale_factor = 8.0f; w1.rope_scale_orig_ctx = 4; test_backend_parity_model("relu2 + qkv-bias + head_dim10 + nope + rope-scale", w1); // Norm placements: OLMo2 post and Gemma sandwich forge::ModelConfig np; np.n_layers = 2; np.d_model = 16; np.n_heads = 2; np.n_kv_heads = 2; np.d_ff = 24; np.vocab_size = 11; np.context_length = 8; np.tied_embeddings = true; np.norm_placement = "post"; test_backend_parity_model("post-norm (OLMo2)", np); np.norm_placement = "sandwich"; test_backend_parity_model("sandwich norm (Gemma)", np); // Wave-2: sliding window (Mistral) through the fused scalar path // (head_dim 64), + local/global pattern, + Gemma3 dual rope theta forge::ModelConfig sw; sw.n_layers = 3; sw.d_model = 128; sw.n_heads = 2; sw.n_kv_heads = 1; sw.d_ff = 96; sw.vocab_size = 11; sw.context_length = 8; sw.tied_embeddings = true; sw.sliding_window = 3; sw.sliding_global_every = 2; sw.rope_theta_global = 1e6f; test_backend_parity_model("sliding-window 3 + global-every-2 [hd64]", sw); // Wave-2: attention softcap (Gemma2) — unfused path on both backends forge::ModelConfig sc; sc.n_layers = 2; sc.d_model = 16; sc.n_heads = 2; sc.n_kv_heads = 1; sc.d_ff = 24; sc.vocab_size = 11; sc.context_length = 8; sc.tied_embeddings = true; sc.attn_softcap = 20.0f; test_backend_parity_model("attn softcap 20 (Gemma2)", sc); // Sliding window through the UNFUSED kernel too (head_dim 8) forge::ModelConfig swu; swu.n_layers = 2; swu.d_model = 16; swu.n_heads = 2; swu.n_kv_heads = 2; swu.d_ff = 24; swu.vocab_size = 11; swu.context_length = 8; swu.tied_embeddings = true; swu.sliding_window = 3; swu.attn_softcap = 20.0f; test_backend_parity_model("sliding-window 3 + softcap [unfused hd8]", swu); } test_metal_moe_quant_ops(); test_flash_attention(); test_adamw_kernel(); pool->drain(); if (g_failures) { std::printf("\n%d test(s) FAILED\n", g_failures); return 1; } std::printf("\nall tests passed\n"); return 0; }