// Author: Simon-Pierre Boucher โ€” contact@spboucher.ai #include "ops/metal/metal_ops.h" #include "core/device.h" #include #include #include #include #include namespace forge::metal { namespace { void check(bool cond, const char* msg) { if (!cond) { std::fprintf(stderr, "forge/metal: %s\n", msg); std::abort(); } } struct MatmulParams { uint32_t M, N, K; uint32_t lda, ldb; uint32_t accumulate; }; struct NormParams { uint32_t C; float eps; }; // Reduction-style kernels: one threadgroup per row, strided row walk. MTL::Size row_threadgroup(int64_t C, MTL::ComputePipelineState* pso) { NS::UInteger tg = 256; tg = std::min(tg, pso->maxTotalThreadsPerThreadgroup()); tg = std::min(tg, NS::UInteger((C + 31) / 32) * 32); return MTL::Size(std::max(tg, 32), 1, 1); } void encode_flat(const char* kernel, std::initializer_list tensors, const void* params, size_t params_len, int64_t n) { Device& dev = Device::get(); MTL::ComputePipelineState* pso = dev.pipeline(kernel); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); int idx = 0; for (const Tensor* t : tensors) enc->setBuffer(t->buffer(), t->buffer_offset(), idx++); if (params) enc->setBytes(params, params_len, idx); const NS::UInteger tg = std::min(256, pso->maxTotalThreadsPerThreadgroup()); enc->dispatchThreads(MTL::Size(NS::UInteger(n), 1, 1), MTL::Size(tg, 1, 1)); } } // namespace // ---- Stream ----------------------------------------------------------------- Stream& Stream::get() { static Stream stream; return stream; } MTL::ComputeCommandEncoder* Stream::encoder() { return encoder_of(concurrent_ ? MTL::DispatchTypeConcurrent : MTL::DispatchTypeSerial); } void Stream::set_concurrent(bool on) { concurrent_ = on; } MTL::ComputeCommandEncoder* Stream::concurrent_encoder() { return encoder_of(MTL::DispatchTypeConcurrent); } MTL::ComputeCommandEncoder* Stream::encoder_of(int dispatch_type) { if (enc_ && dispatch_type_ == dispatch_type) return enc_; if (enc_) { // Close the current encoder; the boundary orders tracked resources. enc_->endEncoding(); enc_->release(); enc_ = nullptr; } if (!cmd_) { Device::get().allocator().set_defer(true); // park releases until sync cmd_ = Device::get().queue()->commandBuffer(); check(cmd_ != nullptr, "commandBuffer creation failed"); cmd_->retain(); // survives autorelease-pool drains between ops } enc_ = cmd_->computeCommandEncoder(MTL::DispatchType(dispatch_type)); check(enc_ != nullptr, "encoder creation failed"); enc_->retain(); dispatch_type_ = dispatch_type; return enc_; } void Stream::sync() { if (!cmd_) return; if (enc_) enc_->endEncoding(); cmd_->commit(); cmd_->waitUntilCompleted(); if (cmd_->status() == MTL::CommandBufferStatusError) { NS::Error* err = cmd_->error(); std::fprintf(stderr, "forge/metal: command buffer failed: %s\n", err ? err->localizedDescription()->utf8String() : "?"); std::abort(); } gpu_seconds_ += cmd_->GPUEndTime() - cmd_->GPUStartTime(); if (enc_) enc_->release(); cmd_->release(); enc_ = nullptr; cmd_ = nullptr; dispatch_type_ = -1; Allocator& alloc = Device::get().allocator(); alloc.flush_retired(); // GPU is idle: parked buffers may recycle now alloc.set_defer(false); } // ---- matmul ------------------------------------------------------------------- void matmul(const Tensor& a, const Tensor& b, Tensor& c, bool ta, bool tb, bool accumulate, MatmulKernel kernel) { check(a.dtype() == DType::F32 && b.dtype() == DType::F32 && c.dtype() == DType::F32, "matmul: f32 only"); const int64_t M = ta ? a.size(1) : a.size(0); const int64_t K = ta ? a.size(0) : a.size(1); const int64_t N = tb ? b.size(0) : b.size(1); check((tb ? b.size(1) : b.size(0)) == K, "matmul: inner dims mismatch"); check(c.size(0) == M && c.size(1) == N, "matmul: output shape mismatch"); // Auto: simdgroup kernel unless the problem is too small to fill even one // 64x64 block, where the 16x16-tiled kernel's finer granularity wins. if (kernel == MatmulKernel::Auto) kernel = (M >= 64 && N >= 64 && K >= 16) ? MatmulKernel::Simdgroup : MatmulKernel::Tiled; Device& dev = Device::get(); const bool simd = kernel == MatmulKernel::Simdgroup; const bool aligned = simd && (M % 64 == 0) && (N % 64 == 0) && (K % 16 == 0); // Function constants: 0/1 = transposes, 3 = alignment fast path. MTL::FunctionConstantValues* constants = MTL::FunctionConstantValues::alloc()->init(); constants->setConstantValue(&ta, MTL::DataTypeBool, NS::UInteger(0)); constants->setConstantValue(&tb, MTL::DataTypeBool, NS::UInteger(1)); std::string key = std::string(ta ? "t" : "n") + (tb ? "t" : "n"); if (simd) { constants->setConstantValue(&aligned, MTL::DataTypeBool, NS::UInteger(3)); key += aligned ? "/a" : "/u"; } const char* name = simd ? "matmul_simd_f32" : (kernel == MatmulKernel::Naive ? "matmul_naive_f32" : "matmul_tiled_f32"); MTL::ComputePipelineState* pso = dev.pipeline(name, constants, key); constants->release(); MatmulParams p{uint32_t(M), uint32_t(N), uint32_t(K), uint32_t(a.size(1)), uint32_t(b.size(1)), accumulate ? 1u : 0u}; MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(a.buffer(), a.buffer_offset(), 0); enc->setBuffer(b.buffer(), b.buffer_offset(), 1); enc->setBuffer(c.buffer(), c.buffer_offset(), 2); enc->setBytes(&p, sizeof(p), 3); switch (kernel) { case MatmulKernel::Naive: enc->dispatchThreads(MTL::Size(NS::UInteger(N), NS::UInteger(M), 1), MTL::Size(16, 16, 1)); break; case MatmulKernel::Simdgroup: { constexpr NS::UInteger BM = 64, BN = 64, THREADS = 128; enc->dispatchThreadgroups( MTL::Size((NS::UInteger(N) + BN - 1) / BN, (NS::UInteger(M) + BM - 1) / BM, 1), MTL::Size(THREADS, 1, 1)); break; } default: { constexpr NS::UInteger TILE = 16; enc->dispatchThreadgroups( MTL::Size((NS::UInteger(N) + TILE - 1) / TILE, (NS::UInteger(M) + TILE - 1) / TILE, 1), MTL::Size(TILE, TILE, 1)); break; } } } // ---- elementwise --------------------------------------------------------------- void add(const Tensor& a, const Tensor& b, Tensor& out) { check(a.numel() == b.numel() && a.numel() == out.numel(), "add: numel mismatch"); encode_flat("add_f32", {&a, &b, &out}, nullptr, 0, a.numel()); } void mul(const Tensor& a, const Tensor& b, Tensor& out) { check(a.numel() == b.numel() && a.numel() == out.numel(), "mul: numel mismatch"); encode_flat("mul_f32", {&a, &b, &out}, nullptr, 0, a.numel()); } void scale(const Tensor& a, float s, Tensor& out) { encode_flat("scale_f32", {&a, &out}, &s, sizeof(s), a.numel()); } void add_bias(const Tensor& x, const Tensor& bias, Tensor& out) { const uint32_t C = uint32_t(bias.numel()); encode_flat("add_bias_f32", {&x, &bias, &out}, &C, sizeof(C), x.numel()); } void silu(const Tensor& x, Tensor& out) { encode_flat("silu_f32", {&x, &out}, nullptr, 0, x.numel()); } void gelu(const Tensor& x, Tensor& out) { encode_flat("gelu_f32", {&x, &out}, nullptr, 0, x.numel()); } void relu2(const Tensor& x, Tensor& out) { encode_flat("relu2_f32", {&x, &out}, nullptr, 0, x.numel()); } void relu2_backward(const Tensor& x, const Tensor& dout, Tensor& dx) { encode_flat("relu2_bwd_f32", {&x, &dout, &dx}, nullptr, 0, x.numel()); } void softcap(const Tensor& x, float cap, Tensor& out) { encode_flat("softcap_f32", {&x, &out}, &cap, sizeof(cap), x.numel()); } void softcap_backward(const Tensor& y, const Tensor& dout, float cap, Tensor& dx) { encode_flat("softcap_bwd_f32", {&y, &dout, &dx}, &cap, sizeof(cap), y.numel()); } // ---- row reductions -------------------------------------------------------------- void softmax(const Tensor& x, Tensor& out) { const int64_t C = x.shape().back(); const int64_t rows = x.numel() / C; Device& dev = Device::get(); MTL::ComputePipelineState* pso = dev.pipeline("softmax_f32"); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(x.buffer(), x.buffer_offset(), 0); enc->setBuffer(out.buffer(), out.buffer_offset(), 1); const uint32_t c32 = uint32_t(C); enc->setBytes(&c32, sizeof(c32), 2); enc->dispatchThreadgroups(MTL::Size(NS::UInteger(rows), 1, 1), row_threadgroup(C, pso)); } void rmsnorm(const Tensor& x, const Tensor& w, float eps, Tensor& out) { const int64_t C = w.numel(); const int64_t rows = x.numel() / C; Device& dev = Device::get(); MTL::ComputePipelineState* pso = dev.pipeline("rmsnorm_f32"); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(x.buffer(), x.buffer_offset(), 0); enc->setBuffer(w.buffer(), w.buffer_offset(), 1); enc->setBuffer(out.buffer(), out.buffer_offset(), 2); NormParams p{uint32_t(C), eps}; enc->setBytes(&p, sizeof(p), 3); enc->dispatchThreadgroups(MTL::Size(NS::UInteger(rows), 1, 1), row_threadgroup(C, pso)); } void layernorm(const Tensor& x, const Tensor& w, const Tensor& b, float eps, Tensor& out) { const int64_t C = w.numel(); const int64_t rows = x.numel() / C; Device& dev = Device::get(); MTL::ComputePipelineState* pso = dev.pipeline("layernorm_f32"); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(x.buffer(), x.buffer_offset(), 0); enc->setBuffer(w.buffer(), w.buffer_offset(), 1); enc->setBuffer(b.buffer(), b.buffer_offset(), 2); enc->setBuffer(out.buffer(), out.buffer_offset(), 3); NormParams p{uint32_t(C), eps}; enc->setBytes(&p, sizeof(p), 4); enc->dispatchThreadgroups(MTL::Size(NS::UInteger(rows), 1, 1), row_threadgroup(C, pso)); } // ---- QAT / MoE ------------------------------------------------------------------- void fake_quant(const Tensor& w, int mode, Tensor& out) { check(w.ndim() == 2 && w.numel() == out.numel(), "fake_quant: bad shapes"); const uint32_t p[2] = {uint32_t(w.size(1)), uint32_t(mode)}; encode_flat("fake_quant_f32", {&w, &out}, p, sizeof(p), w.size(0)); } void softmax_backward(const Tensor& p, const Tensor& dout, Tensor& dx) { const int64_t C = p.shape().back(); const uint32_t c32 = uint32_t(C); encode_flat("softmax_bwd_f32", {&p, &dout, &dx}, &c32, sizeof(c32), p.numel() / C); } void sigmoid(const Tensor& x, Tensor& out) { encode_flat("sigmoid_f32", {&x, &out}, nullptr, 0, x.numel()); } void sigmoid_backward(const Tensor& y, const Tensor& dout, Tensor& dx) { encode_flat("sigmoid_bwd_f32", {&y, &dout, &dx}, nullptr, 0, y.numel()); } void topk_renorm(const Tensor& p, const Tensor& bias, int64_t k, bool norm, Tensor& out) { const int64_t E = p.shape().back(); check(E <= 64, "topk_renorm: E must be <= 64 (kernel's kept[] bound)"); const uint32_t pr[3] = {uint32_t(E), uint32_t(k), norm ? 1u : 0u}; encode_flat("topk_renorm_f32", {&p, &bias, &out}, pr, sizeof(pr), p.numel() / E); } void topk_renorm_backward(const Tensor& p, const Tensor& bias, const Tensor& dout, int64_t k, bool norm, Tensor& dp) { const int64_t E = p.shape().back(); check(E <= 64, "topk_renorm_backward: E must be <= 64"); const uint32_t pr[3] = {uint32_t(E), uint32_t(k), norm ? 1u : 0u}; encode_flat("topk_renorm_bwd_f32", {&p, &bias, &dout, &dp}, pr, sizeof(pr), p.numel() / E); } void expert_counts(const Tensor& gates, Tensor& counts) { const int64_t E = gates.shape().back(); const uint32_t pr[2] = {uint32_t(E), uint32_t(gates.numel() / E)}; encode_flat("expert_counts_f32", {&gates, &counts}, pr, sizeof(pr), E); } namespace { struct RowScaleParams { uint32_t C, E, e; }; } // namespace void row_scale(const Tensor& x, const Tensor& gates, int64_t e, Tensor& out) { const RowScaleParams p{uint32_t(x.shape().back()), uint32_t(gates.shape().back()), uint32_t(e)}; encode_flat("row_scale_f32", {&x, &gates, &out}, &p, sizeof(p), x.numel()); } void row_scale_accumulate(const Tensor& x, const Tensor& gates, int64_t e, Tensor& dst) { const RowScaleParams p{uint32_t(x.shape().back()), uint32_t(gates.shape().back()), uint32_t(e)}; encode_flat("row_scale_acc_f32", {&x, &gates, &dst}, &p, sizeof(p), x.numel()); } void row_scale_gate_backward(const Tensor& dout, const Tensor& x, int64_t e, Tensor& dgates) { const RowScaleParams p{uint32_t(x.shape().back()), uint32_t(dgates.shape().back()), uint32_t(e)}; encode_flat("row_scale_gate_bwd_f32", {&dout, &x, &dgates}, &p, sizeof(p), x.numel() / x.shape().back()); } // ---- backward / training ops --------------------------------------------------- void accumulate(Tensor& dst, const Tensor& src) { encode_flat("accum_f32", {&dst, &src}, nullptr, 0, dst.numel()); } void axpy(Tensor& dst, const Tensor& src, const Tensor& s) { encode_flat("axpy_f32", {&dst, &src, &s}, nullptr, 0, dst.numel()); } void silu_backward(const Tensor& x, const Tensor& dout, Tensor& dx) { encode_flat("silu_bwd_f32", {&x, &dout, &dx}, nullptr, 0, x.numel()); } void gelu_backward(const Tensor& x, const Tensor& dout, Tensor& dx) { encode_flat("gelu_bwd_f32", {&x, &dout, &dx}, nullptr, 0, x.numel()); } void add_bias_backward(const Tensor& dout, Tensor& dbias) { const uint32_t C = uint32_t(dbias.numel()); const uint32_t N = uint32_t(dout.numel() / C); const uint32_t nc[2] = {N, C}; encode_flat("add_bias_bwd_f32", {&dout, &dbias}, nc, sizeof(nc), C); } void rmsnorm_backward(const Tensor& x, const Tensor& w, float eps, const Tensor& dout, Tensor& dx, Tensor& dw) { const int64_t C = w.numel(); const int64_t rows = x.numel() / C; Tensor inv_rms = Tensor::empty({rows}); Device& dev = Device::get(); { MTL::ComputePipelineState* pso = dev.pipeline("rmsnorm_bwd_dx_f32"); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(x.buffer(), x.buffer_offset(), 0); enc->setBuffer(w.buffer(), w.buffer_offset(), 1); enc->setBuffer(dout.buffer(), dout.buffer_offset(), 2); enc->setBuffer(dx.buffer(), dx.buffer_offset(), 3); enc->setBuffer(inv_rms.buffer(), inv_rms.buffer_offset(), 4); NormParams p{uint32_t(C), eps}; enc->setBytes(&p, sizeof(p), 5); enc->dispatchThreadgroups(MTL::Size(NS::UInteger(rows), 1, 1), row_threadgroup(C, pso)); } { MTL::ComputePipelineState* pso = dev.pipeline("rmsnorm_bwd_dw_f32"); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(x.buffer(), x.buffer_offset(), 0); enc->setBuffer(dout.buffer(), dout.buffer_offset(), 1); enc->setBuffer(inv_rms.buffer(), inv_rms.buffer_offset(), 2); enc->setBuffer(dw.buffer(), dw.buffer_offset(), 3); NormParams p{uint32_t(C), eps}; enc->setBytes(&p, sizeof(p), 4); const uint32_t r32 = uint32_t(rows); enc->setBytes(&r32, sizeof(r32), 5); enc->dispatchThreads(MTL::Size(NS::UInteger(C), 1, 1), MTL::Size(64, 1, 1)); } } void layernorm_backward(const Tensor& x, const Tensor& w, float eps, const Tensor& dout, Tensor& dx, Tensor& dw, Tensor& db) { const int64_t C = w.numel(); const int64_t rows = x.numel() / C; Tensor mean = Tensor::empty({rows}); Tensor istd = Tensor::empty({rows}); Device& dev = Device::get(); { MTL::ComputePipelineState* pso = dev.pipeline("layernorm_bwd_dx_f32"); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(x.buffer(), x.buffer_offset(), 0); enc->setBuffer(w.buffer(), w.buffer_offset(), 1); enc->setBuffer(dout.buffer(), dout.buffer_offset(), 2); enc->setBuffer(dx.buffer(), dx.buffer_offset(), 3); enc->setBuffer(mean.buffer(), mean.buffer_offset(), 4); enc->setBuffer(istd.buffer(), istd.buffer_offset(), 5); NormParams p{uint32_t(C), eps}; enc->setBytes(&p, sizeof(p), 6); enc->dispatchThreadgroups(MTL::Size(NS::UInteger(rows), 1, 1), row_threadgroup(C, pso)); } { MTL::ComputePipelineState* pso = dev.pipeline("layernorm_bwd_dwdb_f32"); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(x.buffer(), x.buffer_offset(), 0); enc->setBuffer(dout.buffer(), dout.buffer_offset(), 1); enc->setBuffer(mean.buffer(), mean.buffer_offset(), 2); enc->setBuffer(istd.buffer(), istd.buffer_offset(), 3); enc->setBuffer(dw.buffer(), dw.buffer_offset(), 4); enc->setBuffer(db.buffer(), db.buffer_offset(), 5); NormParams p{uint32_t(C), eps}; enc->setBytes(&p, sizeof(p), 6); const uint32_t r32 = uint32_t(rows); enc->setBytes(&r32, sizeof(r32), 7); enc->dispatchThreads(MTL::Size(NS::UInteger(C), 1, 1), MTL::Size(64, 1, 1)); } } // ---- rope ----------------------------------------------------------------------- namespace { struct RopeParams { uint32_t T, H, HD; uint32_t pos_offset; }; void rope_encode(const Tensor& in, Tensor& out, int64_t n_heads, const Tensor& freqs, int64_t pos_offset, bool inverse) { const int64_t B = in.size(0), T = in.size(1), C = in.size(2); const int64_t hd = C / n_heads; check(freqs.numel() == hd / 2, "rope: freqs must have head_dim/2 entries"); Device& dev = Device::get(); MTL::FunctionConstantValues* constants = MTL::FunctionConstantValues::alloc()->init(); // slots 0/1 belong to matmul TA/TB; RoPE uses slot 2 constants->setConstantValue(&inverse, MTL::DataTypeBool, NS::UInteger(2)); MTL::ComputePipelineState* pso = dev.pipeline("rope_f32", constants, inverse ? "inv" : "fwd"); constants->release(); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(in.buffer(), in.buffer_offset(), 0); enc->setBuffer(out.buffer(), out.buffer_offset(), 1); RopeParams p{uint32_t(T), uint32_t(n_heads), uint32_t(hd), uint32_t(pos_offset)}; enc->setBytes(&p, sizeof(p), 2); enc->setBuffer(freqs.buffer(), freqs.buffer_offset(), 3); const int64_t pairs = B * T * n_heads * (hd / 2); enc->dispatchThreads(MTL::Size(NS::UInteger(pairs), 1, 1), MTL::Size(256, 1, 1)); } } // namespace void rope(const Tensor& x, int64_t n_heads, const Tensor& freqs, int64_t pos_offset, Tensor& out) { rope_encode(x, out, n_heads, freqs, pos_offset, /*inverse=*/false); } void rope_backward(const Tensor& dout, int64_t n_heads, const Tensor& freqs, int64_t pos_offset, Tensor& dx) { rope_encode(dout, dx, n_heads, freqs, pos_offset, /*inverse=*/true); } // ---- embedding ------------------------------------------------------------------- void embedding(const Tensor& weight, const Tensor& ids, Tensor& out) { check(ids.dtype() == DType::I32, "embedding: ids must be i32 on GPU"); const int64_t C = weight.size(1); const int64_t N = ids.numel(); Device& dev = Device::get(); MTL::ComputePipelineState* pso = dev.pipeline("embedding_fwd_f32"); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(weight.buffer(), weight.buffer_offset(), 0); enc->setBuffer(ids.buffer(), ids.buffer_offset(), 1); enc->setBuffer(out.buffer(), out.buffer_offset(), 2); const uint32_t c32 = uint32_t(C); enc->setBytes(&c32, sizeof(c32), 3); enc->dispatchThreads(MTL::Size(NS::UInteger(C), NS::UInteger(N), 1), MTL::Size(std::min(NS::UInteger(C), 64), 4, 1)); } void embedding_backward(const Tensor& ids, const Tensor& dout, Tensor& dweight) { check(ids.dtype() == DType::I32, "embedding_backward: ids must be i32 on GPU"); const int64_t C = dweight.size(1); const int64_t N = ids.numel(); Device& dev = Device::get(); MTL::ComputePipelineState* pso = dev.pipeline("embedding_bwd_f32"); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(ids.buffer(), ids.buffer_offset(), 0); enc->setBuffer(dout.buffer(), dout.buffer_offset(), 1); enc->setBuffer(dweight.buffer(), dweight.buffer_offset(), 2); const uint32_t c32 = uint32_t(C); enc->setBytes(&c32, sizeof(c32), 3); enc->dispatchThreads(MTL::Size(NS::UInteger(C), NS::UInteger(N), 1), MTL::Size(std::min(NS::UInteger(C), 64), 4, 1)); } // ---- attention ------------------------------------------------------------------- namespace { struct AttnParams { uint32_t B, T, H, HKV, HD; float scale; uint32_t causal; uint32_t window; float softcap; }; } // namespace void attention(const Tensor& q, const Tensor& k, const Tensor& v, int64_t n_heads, int64_t n_kv_heads, bool causal, float scale, Tensor& out, Tensor* probs_out, int64_t window, float attn_softcap) { check(probs_out != nullptr, "attention: probs_out required (unfused path)"); const int64_t B = q.size(0), T = q.size(1); const int64_t hd = q.size(2) / n_heads; Device& dev = Device::get(); MTL::ComputePipelineState* pso = dev.pipeline("attention_fwd_f32"); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(q.buffer(), q.buffer_offset(), 0); enc->setBuffer(k.buffer(), k.buffer_offset(), 1); enc->setBuffer(v.buffer(), v.buffer_offset(), 2); enc->setBuffer(out.buffer(), out.buffer_offset(), 3); enc->setBuffer(probs_out->buffer(), probs_out->buffer_offset(), 4); AttnParams p{uint32_t(B), uint32_t(T), uint32_t(n_heads), uint32_t(n_kv_heads), uint32_t(hd), scale, causal ? 1u : 0u, uint32_t(window), attn_softcap}; enc->setBytes(&p, sizeof(p), 5); enc->dispatchThreads(MTL::Size(NS::UInteger(B * n_heads * T), 1, 1), MTL::Size(64, 1, 1)); } void attention_backward(const Tensor& q, const Tensor& k, const Tensor& v, const Tensor& probs, const Tensor& out, const Tensor& dout, int64_t n_heads, int64_t n_kv_heads, float scale, Tensor& dq, Tensor& dk, Tensor& dv, float attn_softcap) { const int64_t B = q.size(0), T = q.size(1); const int64_t hd = q.size(2) / n_heads; Device& dev = Device::get(); // window is irrelevant here: masked positions carry prob 0 in `probs`. AttnParams p{uint32_t(B), uint32_t(T), uint32_t(n_heads), uint32_t(n_kv_heads), uint32_t(hd), scale, 1u, 0u, attn_softcap}; // D[b,h,i] = dO_i ยท O_i == rowsum(dP โˆ˜ P): computed once per query row so // neither backward kernel needs the O(T) inner recompute. Tensor d_term = Tensor::empty({B, n_heads, T}); { MTL::ComputePipelineState* pso = dev.pipeline("attention_bwd_d_f32"); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(out.buffer(), out.buffer_offset(), 0); enc->setBuffer(dout.buffer(), dout.buffer_offset(), 1); enc->setBuffer(d_term.buffer(), d_term.buffer_offset(), 2); enc->setBytes(&p, sizeof(p), 3); enc->dispatchThreads(MTL::Size(NS::UInteger(B * n_heads * T), 1, 1), MTL::Size(64, 1, 1)); } { MTL::ComputePipelineState* pso = dev.pipeline("attention_bwd_dq_f32"); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(q.buffer(), q.buffer_offset(), 0); enc->setBuffer(k.buffer(), k.buffer_offset(), 1); enc->setBuffer(v.buffer(), v.buffer_offset(), 2); enc->setBuffer(probs.buffer(), probs.buffer_offset(), 3); enc->setBuffer(dout.buffer(), dout.buffer_offset(), 4); enc->setBuffer(d_term.buffer(), d_term.buffer_offset(), 5); enc->setBuffer(dq.buffer(), dq.buffer_offset(), 6); enc->setBytes(&p, sizeof(p), 7); enc->dispatchThreads(MTL::Size(NS::UInteger(B * n_heads * T), 1, 1), MTL::Size(64, 1, 1)); } { MTL::ComputePipelineState* pso = dev.pipeline("attention_bwd_dkv_f32"); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(q.buffer(), q.buffer_offset(), 0); enc->setBuffer(k.buffer(), k.buffer_offset(), 1); enc->setBuffer(v.buffer(), v.buffer_offset(), 2); enc->setBuffer(probs.buffer(), probs.buffer_offset(), 3); enc->setBuffer(dout.buffer(), dout.buffer_offset(), 4); enc->setBuffer(d_term.buffer(), d_term.buffer_offset(), 5); enc->setBuffer(dk.buffer(), dk.buffer_offset(), 6); enc->setBuffer(dv.buffer(), dv.buffer_offset(), 7); enc->setBytes(&p, sizeof(p), 8); enc->dispatchThreads(MTL::Size(NS::UInteger(B * n_kv_heads * T), 1, 1), MTL::Size(64, 1, 1)); } } // ---- cross entropy ----------------------------------------------------------------- void cross_entropy(const Tensor& logits, const Tensor& targets, int64_t n_valid, Tensor& losses, Tensor& loss_out, Tensor* dlogits) { check(targets.dtype() == DType::I32, "cross_entropy: targets must be i32 on GPU"); const int64_t V = logits.size(1); const int64_t N = logits.size(0); Device& dev = Device::get(); struct CEParams { uint32_t V; float inv_n; uint32_t want_grad; } p{uint32_t(V), n_valid > 0 ? 1.0f / float(n_valid) : 0.0f, dlogits ? 1u : 0u}; MTL::ComputePipelineState* pso = dev.pipeline("cross_entropy_f32"); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(logits.buffer(), logits.buffer_offset(), 0); enc->setBuffer(targets.buffer(), targets.buffer_offset(), 1); enc->setBuffer(losses.buffer(), losses.buffer_offset(), 2); // dlogits slot must be bound even when unused const Tensor& dl = dlogits ? *dlogits : losses; enc->setBuffer(dl.buffer(), dl.buffer_offset(), 3); enc->setBytes(&p, sizeof(p), 4); enc->dispatchThreadgroups(MTL::Size(NS::UInteger(N), 1, 1), row_threadgroup(V, pso)); sum(losses, loss_out, n_valid > 0 ? 1.0f / float(n_valid) : 0.0f); } // ---- optimizer / reductions ---------------------------------------------------------- void adamw_step(Tensor& w, const Tensor& g, Tensor& m, Tensor& v, float lr, float beta1, float beta2, int64_t t, float eps, float wd, float grad_scale) { struct AdamWParams { float lr, beta1, beta2, bc1, bc2, eps, wd, grad_scale; } p{lr, beta1, beta2, 1.0f - std::pow(beta1, float(t)), 1.0f - std::pow(beta2, float(t)), eps, wd, grad_scale}; encode_flat("adamw_f32", {&w, &g, &m, &v}, &p, sizeof(p), w.numel()); } void sumsq(const Tensor& x, Tensor& out) { Device& dev = Device::get(); MTL::ComputePipelineState* pso = dev.pipeline("sumsq_f32"); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(x.buffer(), x.buffer_offset(), 0); enc->setBuffer(out.buffer(), out.buffer_offset(), 1); const uint32_t n = uint32_t(x.numel()); enc->setBytes(&n, sizeof(n), 2); // single threadgroup: partials[0] is the total enc->dispatchThreadgroups(MTL::Size(1, 1, 1), MTL::Size(1024, 1, 1)); } void sum(const Tensor& x, Tensor& out, float mul) { Device& dev = Device::get(); MTL::ComputePipelineState* pso = dev.pipeline("sum_f32"); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(x.buffer(), x.buffer_offset(), 0); enc->setBuffer(out.buffer(), out.buffer_offset(), 1); const uint32_t n = uint32_t(x.numel()); enc->setBytes(&n, sizeof(n), 2); enc->setBytes(&mul, sizeof(mul), 3); enc->dispatchThreadgroups(MTL::Size(1, 1, 1), MTL::Size(1024, 1, 1)); } // ---- fused (flash) attention ------------------------------------------------- namespace { struct FlashParams { uint32_t B, T, H, HKV; float scale; uint32_t causal; uint32_t window; // sliding window; 0 = full. Scalar kernels only โ€” the // caller routes window > 0 away from MMA. }; // Must match the INSTANTIATE_FLASH list in flash_attention.metal. constexpr int64_t kFlashHeadDims[] = {16, 32, 48, 64, 80, 96, 128}; std::string flash_kernel(const char* stem, int64_t head_dim) { return std::string(stem) + "_hd" + std::to_string(head_dim); } } // namespace bool flash_supported(int64_t head_dim) { for (int64_t hd : kFlashHeadDims) if (hd == head_dim) return true; return false; } void flash_attention(const Tensor& q, const Tensor& k, const Tensor& v, int64_t n_heads, int64_t n_kv_heads, bool causal, float scale, Tensor& out, Tensor& lse, FlashKernel kernel, int64_t window) { const int64_t B = q.size(0), T = q.size(1); const int64_t hd = q.size(2) / n_heads; check(flash_supported(hd), "flash_attention: unsupported head_dim"); if (kernel == FlashKernel::Auto) kernel = window > 0 ? FlashKernel::Scalar : FlashKernel::MMA; const bool mma = kernel == FlashKernel::MMA; check(!(mma && window > 0), "flash_attention: MMA kernel has no window support"); Device& dev = Device::get(); MTL::ComputePipelineState* pso = dev.pipeline( flash_kernel(mma ? "flash_attn_fwd_mma_f32" : "flash_attn_fwd_f32", hd)); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(q.buffer(), q.buffer_offset(), 0); enc->setBuffer(k.buffer(), k.buffer_offset(), 1); enc->setBuffer(v.buffer(), v.buffer_offset(), 2); enc->setBuffer(out.buffer(), out.buffer_offset(), 3); enc->setBuffer(lse.buffer(), lse.buffer_offset(), 4); FlashParams p{uint32_t(B), uint32_t(T), uint32_t(n_heads), uint32_t(n_kv_heads), scale, causal ? 1u : 0u, uint32_t(window)}; enc->setBytes(&p, sizeof(p), 5); // One threadgroup per (query block, head, batch). Block size and thread // count must match the kernel's enums: the scalar kernel is one thread // per query row (TGQ=64), the MMA kernel is BQ=32 rows across 4 // simdgroups (128 threads). const NS::UInteger block_q = mma ? 32 : 64; const NS::UInteger threads = mma ? 128 : 64; check(pso->maxTotalThreadsPerThreadgroup() >= threads, "flash_attention: pipeline cannot host the required threadgroup size"); const NS::UInteger q_blocks = (NS::UInteger(T) + block_q - 1) / block_q; enc->dispatchThreadgroups( MTL::Size(q_blocks * NS::UInteger(n_heads) * NS::UInteger(B), 1, 1), MTL::Size(threads, 1, 1)); } void flash_attention_backward(const Tensor& q, const Tensor& k, const Tensor& v, const Tensor& out, const Tensor& lse, const Tensor& dout, int64_t n_heads, int64_t n_kv_heads, bool causal, float scale, Tensor& dq, Tensor& dk, Tensor& dv, FlashKernel kernel, int64_t window) { const int64_t B = q.size(0), T = q.size(1); const int64_t hd = q.size(2) / n_heads; check(flash_supported(hd), "flash_attention_backward: unsupported head_dim"); if (kernel == FlashKernel::Auto) kernel = window > 0 ? FlashKernel::Scalar : FlashKernel::MMA; const bool mma = kernel == FlashKernel::MMA; check(!(mma && window > 0), "flash_attention_backward: MMA kernel has no window support"); Device& dev = Device::get(); FlashParams p{uint32_t(B), uint32_t(T), uint32_t(n_heads), uint32_t(n_kv_heads), scale, causal ? 1u : 0u, uint32_t(window)}; // D[b,h,i] = dO_i . O_i (shared with the unfused path). Tensor d_term = Tensor::empty({B, n_heads, T}); { AttnParams ap{uint32_t(B), uint32_t(T), uint32_t(n_heads), uint32_t(n_kv_heads), uint32_t(hd), scale, causal ? 1u : 0u, 0u, 0.0f}; MTL::ComputePipelineState* pso = dev.pipeline("attention_bwd_d_f32"); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(out.buffer(), out.buffer_offset(), 0); enc->setBuffer(dout.buffer(), dout.buffer_offset(), 1); enc->setBuffer(d_term.buffer(), d_term.buffer_offset(), 2); enc->setBytes(&ap, sizeof(ap), 3); enc->dispatchThreads(MTL::Size(NS::UInteger(B * n_heads * T), 1, 1), MTL::Size(64, 1, 1)); } { MTL::ComputePipelineState* pso = dev.pipeline(flash_kernel( mma ? "flash_attn_bwd_dq_mma_f32" : "flash_attn_bwd_dq_f32", hd)); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(q.buffer(), q.buffer_offset(), 0); enc->setBuffer(k.buffer(), k.buffer_offset(), 1); enc->setBuffer(v.buffer(), v.buffer_offset(), 2); enc->setBuffer(dout.buffer(), dout.buffer_offset(), 3); enc->setBuffer(lse.buffer(), lse.buffer_offset(), 4); enc->setBuffer(d_term.buffer(), d_term.buffer_offset(), 5); enc->setBuffer(dq.buffer(), dq.buffer_offset(), 6); enc->setBytes(&p, sizeof(p), 7); if (mma) { const NS::UInteger q_blocks = (NS::UInteger(T) + 32 - 1) / 32; enc->dispatchThreadgroups( MTL::Size(q_blocks * NS::UInteger(n_heads) * NS::UInteger(B), 1, 1), MTL::Size(128, 1, 1)); } else { const NS::UInteger tg = std::min(64, pso->maxTotalThreadsPerThreadgroup()); enc->dispatchThreads(MTL::Size(NS::UInteger(B * n_heads * T), 1, 1), MTL::Size(tg, 1, 1)); } } if (mma) { // dV then dK as separate kernels: fused, the thread held K, V, dK and // dV as fragments and spilled 4352 bytes (measured with gpudebug). { MTL::ComputePipelineState* pso = dev.pipeline(flash_kernel("flash_attn_bwd_dv_mma_f32", hd)); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(q.buffer(), q.buffer_offset(), 0); enc->setBuffer(k.buffer(), k.buffer_offset(), 1); enc->setBuffer(dout.buffer(), dout.buffer_offset(), 2); enc->setBuffer(lse.buffer(), lse.buffer_offset(), 3); enc->setBuffer(dv.buffer(), dv.buffer_offset(), 4); enc->setBytes(&p, sizeof(p), 5); const NS::UInteger kv_blocks = (NS::UInteger(T) + 32 - 1) / 32; enc->dispatchThreadgroups( MTL::Size(kv_blocks * NS::UInteger(n_kv_heads) * NS::UInteger(B), 1, 1), MTL::Size(128, 1, 1)); } { MTL::ComputePipelineState* pso = dev.pipeline(flash_kernel("flash_attn_bwd_dk_mma_f32", hd)); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(q.buffer(), q.buffer_offset(), 0); enc->setBuffer(k.buffer(), k.buffer_offset(), 1); enc->setBuffer(v.buffer(), v.buffer_offset(), 2); enc->setBuffer(dout.buffer(), dout.buffer_offset(), 3); enc->setBuffer(lse.buffer(), lse.buffer_offset(), 4); enc->setBuffer(d_term.buffer(), d_term.buffer_offset(), 5); enc->setBuffer(dk.buffer(), dk.buffer_offset(), 6); enc->setBytes(&p, sizeof(p), 7); const NS::UInteger kv_blocks = (NS::UInteger(T) + 32 - 1) / 32; enc->dispatchThreadgroups( MTL::Size(kv_blocks * NS::UInteger(n_kv_heads) * NS::UInteger(B), 1, 1), MTL::Size(128, 1, 1)); } } else { { MTL::ComputePipelineState* pso = dev.pipeline(flash_kernel("flash_attn_bwd_dv_f32", hd)); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(q.buffer(), q.buffer_offset(), 0); enc->setBuffer(k.buffer(), k.buffer_offset(), 1); enc->setBuffer(dout.buffer(), dout.buffer_offset(), 2); enc->setBuffer(lse.buffer(), lse.buffer_offset(), 3); enc->setBuffer(dv.buffer(), dv.buffer_offset(), 4); enc->setBytes(&p, sizeof(p), 5); const NS::UInteger tg = std::min(64, pso->maxTotalThreadsPerThreadgroup()); enc->dispatchThreads(MTL::Size(NS::UInteger(B * n_kv_heads * T), 1, 1), MTL::Size(tg, 1, 1)); } { MTL::ComputePipelineState* pso = dev.pipeline(flash_kernel("flash_attn_bwd_dk_f32", hd)); MTL::ComputeCommandEncoder* enc = Stream::get().encoder(); enc->setComputePipelineState(pso); enc->setBuffer(q.buffer(), q.buffer_offset(), 0); enc->setBuffer(k.buffer(), k.buffer_offset(), 1); enc->setBuffer(v.buffer(), v.buffer_offset(), 2); enc->setBuffer(dout.buffer(), dout.buffer_offset(), 3); enc->setBuffer(lse.buffer(), lse.buffer_offset(), 4); enc->setBuffer(d_term.buffer(), d_term.buffer_offset(), 5); enc->setBuffer(dk.buffer(), dk.buffer_offset(), 6); enc->setBytes(&p, sizeof(p), 7); const NS::UInteger tg = std::min(64, pso->maxTotalThreadsPerThreadgroup()); enc->dispatchThreads(MTL::Size(NS::UInteger(B * n_kv_heads * T), 1, 1), MTL::Size(tg, 1, 1)); } } } } // namespace forge::metal