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#include "ops/metal/metal_ops.h"34#include "core/device.h"56#include <Metal/Metal.hpp>78#include <algorithm>9#include <cstdio>10#include <cstdlib>11#include <string>1213namespace forge::metal {1415namespace {1617void check(bool cond, const char* msg) {18 if (!cond) {19 std::fprintf(stderr, "forge/metal: %s\n", msg);20 std::abort();21 }22}2324struct MatmulParams {25 uint32_t M, N, K;26 uint32_t lda, ldb;27 uint32_t accumulate;28};2930struct NormParams {31 uint32_t C;32 float eps;33};3435// Reduction-style kernels: one threadgroup per row, strided row walk.36MTL::Size row_threadgroup(int64_t C, MTL::ComputePipelineState* pso) {37 NS::UInteger tg = 256;38 tg = std::min<NS::UInteger>(tg, pso->maxTotalThreadsPerThreadgroup());39 tg = std::min<NS::UInteger>(tg, NS::UInteger((C + 31) / 32) * 32);40 return MTL::Size(std::max<NS::UInteger>(tg, 32), 1, 1);41}4243void encode_flat(const char* kernel, std::initializer_list<const Tensor*> tensors,44 const void* params, size_t params_len, int64_t n) {45 Device& dev = Device::get();46 MTL::ComputePipelineState* pso = dev.pipeline(kernel);47 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();48 enc->setComputePipelineState(pso);49 int idx = 0;50 for (const Tensor* t : tensors)51 enc->setBuffer(t->buffer(), t->buffer_offset(), idx++);52 if (params) enc->setBytes(params, params_len, idx);53 const NS::UInteger tg =54 std::min<NS::UInteger>(256, pso->maxTotalThreadsPerThreadgroup());55 enc->dispatchThreads(MTL::Size(NS::UInteger(n), 1, 1), MTL::Size(tg, 1, 1));56}5758} // namespace5960// ---- Stream -----------------------------------------------------------------6162Stream& Stream::get() {63 static Stream stream;64 return stream;65}6667MTL::ComputeCommandEncoder* Stream::encoder() {68 return encoder_of(concurrent_ ? MTL::DispatchTypeConcurrent69 : MTL::DispatchTypeSerial);70}7172void Stream::set_concurrent(bool on) { concurrent_ = on; }7374MTL::ComputeCommandEncoder* Stream::concurrent_encoder() {75 return encoder_of(MTL::DispatchTypeConcurrent);76}7778MTL::ComputeCommandEncoder* Stream::encoder_of(int dispatch_type) {79 if (enc_ && dispatch_type_ == dispatch_type) return enc_;80 if (enc_) {81 // Close the current encoder; the boundary orders tracked resources.82 enc_->endEncoding();83 enc_->release();84 enc_ = nullptr;85 }86 if (!cmd_) {87 Device::get().allocator().set_defer(true); // park releases until sync88 cmd_ = Device::get().queue()->commandBuffer();89 check(cmd_ != nullptr, "commandBuffer creation failed");90 cmd_->retain(); // survives autorelease-pool drains between ops91 }92 enc_ = cmd_->computeCommandEncoder(MTL::DispatchType(dispatch_type));93 check(enc_ != nullptr, "encoder creation failed");94 enc_->retain();95 dispatch_type_ = dispatch_type;96 return enc_;97}9899void Stream::sync() {100 if (!cmd_) return;101 if (enc_) enc_->endEncoding();102 cmd_->commit();103 cmd_->waitUntilCompleted();104 if (cmd_->status() == MTL::CommandBufferStatusError) {105 NS::Error* err = cmd_->error();106 std::fprintf(stderr, "forge/metal: command buffer failed: %s\n",107 err ? err->localizedDescription()->utf8String() : "?");108 std::abort();109 }110 gpu_seconds_ += cmd_->GPUEndTime() - cmd_->GPUStartTime();111 if (enc_) enc_->release();112 cmd_->release();113 enc_ = nullptr;114 cmd_ = nullptr;115 dispatch_type_ = -1;116 Allocator& alloc = Device::get().allocator();117 alloc.flush_retired(); // GPU is idle: parked buffers may recycle now118 alloc.set_defer(false);119}120121// ---- matmul -------------------------------------------------------------------122123void matmul(const Tensor& a, const Tensor& b, Tensor& c,124 bool ta, bool tb, bool accumulate, MatmulKernel kernel) {125 check(a.dtype() == DType::F32 && b.dtype() == DType::F32 && c.dtype() == DType::F32,126 "matmul: f32 only");127 const int64_t M = ta ? a.size(1) : a.size(0);128 const int64_t K = ta ? a.size(0) : a.size(1);129 const int64_t N = tb ? b.size(0) : b.size(1);130 check((tb ? b.size(1) : b.size(0)) == K, "matmul: inner dims mismatch");131 check(c.size(0) == M && c.size(1) == N, "matmul: output shape mismatch");132133 // Auto: simdgroup kernel unless the problem is too small to fill even one134 // 64x64 block, where the 16x16-tiled kernel's finer granularity wins.135 if (kernel == MatmulKernel::Auto)136 kernel = (M >= 64 && N >= 64 && K >= 16) ? MatmulKernel::Simdgroup137 : MatmulKernel::Tiled;138139 Device& dev = Device::get();140 const bool simd = kernel == MatmulKernel::Simdgroup;141 const bool aligned = simd && (M % 64 == 0) && (N % 64 == 0) && (K % 16 == 0);142143 // Function constants: 0/1 = transposes, 3 = alignment fast path.144 MTL::FunctionConstantValues* constants = MTL::FunctionConstantValues::alloc()->init();145 constants->setConstantValue(&ta, MTL::DataTypeBool, NS::UInteger(0));146 constants->setConstantValue(&tb, MTL::DataTypeBool, NS::UInteger(1));147 std::string key = std::string(ta ? "t" : "n") + (tb ? "t" : "n");148 if (simd) {149 constants->setConstantValue(&aligned, MTL::DataTypeBool, NS::UInteger(3));150 key += aligned ? "/a" : "/u";151 }152 const char* name = simd ? "matmul_simd_f32"153 : (kernel == MatmulKernel::Naive ? "matmul_naive_f32"154 : "matmul_tiled_f32");155 MTL::ComputePipelineState* pso = dev.pipeline(name, constants, key);156 constants->release();157158 MatmulParams p{uint32_t(M), uint32_t(N), uint32_t(K),159 uint32_t(a.size(1)), uint32_t(b.size(1)), accumulate ? 1u : 0u};160161 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();162 enc->setComputePipelineState(pso);163 enc->setBuffer(a.buffer(), a.buffer_offset(), 0);164 enc->setBuffer(b.buffer(), b.buffer_offset(), 1);165 enc->setBuffer(c.buffer(), c.buffer_offset(), 2);166 enc->setBytes(&p, sizeof(p), 3);167168 switch (kernel) {169 case MatmulKernel::Naive:170 enc->dispatchThreads(MTL::Size(NS::UInteger(N), NS::UInteger(M), 1),171 MTL::Size(16, 16, 1));172 break;173 case MatmulKernel::Simdgroup: {174 constexpr NS::UInteger BM = 64, BN = 64, THREADS = 128;175 enc->dispatchThreadgroups(176 MTL::Size((NS::UInteger(N) + BN - 1) / BN,177 (NS::UInteger(M) + BM - 1) / BM, 1),178 MTL::Size(THREADS, 1, 1));179 break;180 }181 default: {182 constexpr NS::UInteger TILE = 16;183 enc->dispatchThreadgroups(184 MTL::Size((NS::UInteger(N) + TILE - 1) / TILE,185 (NS::UInteger(M) + TILE - 1) / TILE, 1),186 MTL::Size(TILE, TILE, 1));187 break;188 }189 }190}191192// ---- elementwise ---------------------------------------------------------------193194void add(const Tensor& a, const Tensor& b, Tensor& out) {195 check(a.numel() == b.numel() && a.numel() == out.numel(), "add: numel mismatch");196 encode_flat("add_f32", {&a, &b, &out}, nullptr, 0, a.numel());197}198199void mul(const Tensor& a, const Tensor& b, Tensor& out) {200 check(a.numel() == b.numel() && a.numel() == out.numel(), "mul: numel mismatch");201 encode_flat("mul_f32", {&a, &b, &out}, nullptr, 0, a.numel());202}203204void scale(const Tensor& a, float s, Tensor& out) {205 encode_flat("scale_f32", {&a, &out}, &s, sizeof(s), a.numel());206}207208void add_bias(const Tensor& x, const Tensor& bias, Tensor& out) {209 const uint32_t C = uint32_t(bias.numel());210 encode_flat("add_bias_f32", {&x, &bias, &out}, &C, sizeof(C), x.numel());211}212213void silu(const Tensor& x, Tensor& out) {214 encode_flat("silu_f32", {&x, &out}, nullptr, 0, x.numel());215}216217void gelu(const Tensor& x, Tensor& out) {218 encode_flat("gelu_f32", {&x, &out}, nullptr, 0, x.numel());219}220221void relu2(const Tensor& x, Tensor& out) {222 encode_flat("relu2_f32", {&x, &out}, nullptr, 0, x.numel());223}224225void relu2_backward(const Tensor& x, const Tensor& dout, Tensor& dx) {226 encode_flat("relu2_bwd_f32", {&x, &dout, &dx}, nullptr, 0, x.numel());227}228229void softcap(const Tensor& x, float cap, Tensor& out) {230 encode_flat("softcap_f32", {&x, &out}, &cap, sizeof(cap), x.numel());231}232233void softcap_backward(const Tensor& y, const Tensor& dout, float cap, Tensor& dx) {234 encode_flat("softcap_bwd_f32", {&y, &dout, &dx}, &cap, sizeof(cap), y.numel());235}236237// ---- row reductions --------------------------------------------------------------238239void softmax(const Tensor& x, Tensor& out) {240 const int64_t C = x.shape().back();241 const int64_t rows = x.numel() / C;242 Device& dev = Device::get();243 MTL::ComputePipelineState* pso = dev.pipeline("softmax_f32");244 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();245 enc->setComputePipelineState(pso);246 enc->setBuffer(x.buffer(), x.buffer_offset(), 0);247 enc->setBuffer(out.buffer(), out.buffer_offset(), 1);248 const uint32_t c32 = uint32_t(C);249 enc->setBytes(&c32, sizeof(c32), 2);250 enc->dispatchThreadgroups(MTL::Size(NS::UInteger(rows), 1, 1),251 row_threadgroup(C, pso));252}253254void rmsnorm(const Tensor& x, const Tensor& w, float eps, Tensor& out) {255 const int64_t C = w.numel();256 const int64_t rows = x.numel() / C;257 Device& dev = Device::get();258 MTL::ComputePipelineState* pso = dev.pipeline("rmsnorm_f32");259 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();260 enc->setComputePipelineState(pso);261 enc->setBuffer(x.buffer(), x.buffer_offset(), 0);262 enc->setBuffer(w.buffer(), w.buffer_offset(), 1);263 enc->setBuffer(out.buffer(), out.buffer_offset(), 2);264 NormParams p{uint32_t(C), eps};265 enc->setBytes(&p, sizeof(p), 3);266 enc->dispatchThreadgroups(MTL::Size(NS::UInteger(rows), 1, 1),267 row_threadgroup(C, pso));268}269270void layernorm(const Tensor& x, const Tensor& w, const Tensor& b, float eps, Tensor& out) {271 const int64_t C = w.numel();272 const int64_t rows = x.numel() / C;273 Device& dev = Device::get();274 MTL::ComputePipelineState* pso = dev.pipeline("layernorm_f32");275 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();276 enc->setComputePipelineState(pso);277 enc->setBuffer(x.buffer(), x.buffer_offset(), 0);278 enc->setBuffer(w.buffer(), w.buffer_offset(), 1);279 enc->setBuffer(b.buffer(), b.buffer_offset(), 2);280 enc->setBuffer(out.buffer(), out.buffer_offset(), 3);281 NormParams p{uint32_t(C), eps};282 enc->setBytes(&p, sizeof(p), 4);283 enc->dispatchThreadgroups(MTL::Size(NS::UInteger(rows), 1, 1),284 row_threadgroup(C, pso));285}286287288// ---- QAT / MoE -------------------------------------------------------------------289290void fake_quant(const Tensor& w, int mode, Tensor& out) {291 check(w.ndim() == 2 && w.numel() == out.numel(), "fake_quant: bad shapes");292 const uint32_t p[2] = {uint32_t(w.size(1)), uint32_t(mode)};293 encode_flat("fake_quant_f32", {&w, &out}, p, sizeof(p), w.size(0));294}295296void softmax_backward(const Tensor& p, const Tensor& dout, Tensor& dx) {297 const int64_t C = p.shape().back();298 const uint32_t c32 = uint32_t(C);299 encode_flat("softmax_bwd_f32", {&p, &dout, &dx}, &c32, sizeof(c32), p.numel() / C);300}301302void sigmoid(const Tensor& x, Tensor& out) {303 encode_flat("sigmoid_f32", {&x, &out}, nullptr, 0, x.numel());304}305306void sigmoid_backward(const Tensor& y, const Tensor& dout, Tensor& dx) {307 encode_flat("sigmoid_bwd_f32", {&y, &dout, &dx}, nullptr, 0, y.numel());308}309310void topk_renorm(const Tensor& p, const Tensor& bias, int64_t k, bool norm,311 Tensor& out) {312 const int64_t E = p.shape().back();313 check(E <= 64, "topk_renorm: E must be <= 64 (kernel's kept[] bound)");314 const uint32_t pr[3] = {uint32_t(E), uint32_t(k), norm ? 1u : 0u};315 encode_flat("topk_renorm_f32", {&p, &bias, &out}, pr, sizeof(pr), p.numel() / E);316}317318void topk_renorm_backward(const Tensor& p, const Tensor& bias, const Tensor& dout,319 int64_t k, bool norm, Tensor& dp) {320 const int64_t E = p.shape().back();321 check(E <= 64, "topk_renorm_backward: E must be <= 64");322 const uint32_t pr[3] = {uint32_t(E), uint32_t(k), norm ? 1u : 0u};323 encode_flat("topk_renorm_bwd_f32", {&p, &bias, &dout, &dp}, pr, sizeof(pr),324 p.numel() / E);325}326327void expert_counts(const Tensor& gates, Tensor& counts) {328 const int64_t E = gates.shape().back();329 const uint32_t pr[2] = {uint32_t(E), uint32_t(gates.numel() / E)};330 encode_flat("expert_counts_f32", {&gates, &counts}, pr, sizeof(pr), E);331}332333namespace {334struct RowScaleParams { uint32_t C, E, e; };335} // namespace336337void row_scale(const Tensor& x, const Tensor& gates, int64_t e, Tensor& out) {338 const RowScaleParams p{uint32_t(x.shape().back()), uint32_t(gates.shape().back()),339 uint32_t(e)};340 encode_flat("row_scale_f32", {&x, &gates, &out}, &p, sizeof(p), x.numel());341}342343void row_scale_accumulate(const Tensor& x, const Tensor& gates, int64_t e, Tensor& dst) {344 const RowScaleParams p{uint32_t(x.shape().back()), uint32_t(gates.shape().back()),345 uint32_t(e)};346 encode_flat("row_scale_acc_f32", {&x, &gates, &dst}, &p, sizeof(p), x.numel());347}348349void row_scale_gate_backward(const Tensor& dout, const Tensor& x, int64_t e,350 Tensor& dgates) {351 const RowScaleParams p{uint32_t(x.shape().back()), uint32_t(dgates.shape().back()),352 uint32_t(e)};353 encode_flat("row_scale_gate_bwd_f32", {&dout, &x, &dgates}, &p, sizeof(p),354 x.numel() / x.shape().back());355}356357// ---- backward / training ops ---------------------------------------------------358359void accumulate(Tensor& dst, const Tensor& src) {360 encode_flat("accum_f32", {&dst, &src}, nullptr, 0, dst.numel());361}362363void axpy(Tensor& dst, const Tensor& src, const Tensor& s) {364 encode_flat("axpy_f32", {&dst, &src, &s}, nullptr, 0, dst.numel());365}366367void silu_backward(const Tensor& x, const Tensor& dout, Tensor& dx) {368 encode_flat("silu_bwd_f32", {&x, &dout, &dx}, nullptr, 0, x.numel());369}370371void gelu_backward(const Tensor& x, const Tensor& dout, Tensor& dx) {372 encode_flat("gelu_bwd_f32", {&x, &dout, &dx}, nullptr, 0, x.numel());373}374375void add_bias_backward(const Tensor& dout, Tensor& dbias) {376 const uint32_t C = uint32_t(dbias.numel());377 const uint32_t N = uint32_t(dout.numel() / C);378 const uint32_t nc[2] = {N, C};379 encode_flat("add_bias_bwd_f32", {&dout, &dbias}, nc, sizeof(nc), C);380}381382void rmsnorm_backward(const Tensor& x, const Tensor& w, float eps,383 const Tensor& dout, Tensor& dx, Tensor& dw) {384 const int64_t C = w.numel();385 const int64_t rows = x.numel() / C;386 Tensor inv_rms = Tensor::empty({rows});387 Device& dev = Device::get();388389 {390 MTL::ComputePipelineState* pso = dev.pipeline("rmsnorm_bwd_dx_f32");391 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();392 enc->setComputePipelineState(pso);393 enc->setBuffer(x.buffer(), x.buffer_offset(), 0);394 enc->setBuffer(w.buffer(), w.buffer_offset(), 1);395 enc->setBuffer(dout.buffer(), dout.buffer_offset(), 2);396 enc->setBuffer(dx.buffer(), dx.buffer_offset(), 3);397 enc->setBuffer(inv_rms.buffer(), inv_rms.buffer_offset(), 4);398 NormParams p{uint32_t(C), eps};399 enc->setBytes(&p, sizeof(p), 5);400 enc->dispatchThreadgroups(MTL::Size(NS::UInteger(rows), 1, 1),401 row_threadgroup(C, pso));402 }403 {404 MTL::ComputePipelineState* pso = dev.pipeline("rmsnorm_bwd_dw_f32");405 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();406 enc->setComputePipelineState(pso);407 enc->setBuffer(x.buffer(), x.buffer_offset(), 0);408 enc->setBuffer(dout.buffer(), dout.buffer_offset(), 1);409 enc->setBuffer(inv_rms.buffer(), inv_rms.buffer_offset(), 2);410 enc->setBuffer(dw.buffer(), dw.buffer_offset(), 3);411 NormParams p{uint32_t(C), eps};412 enc->setBytes(&p, sizeof(p), 4);413 const uint32_t r32 = uint32_t(rows);414 enc->setBytes(&r32, sizeof(r32), 5);415 enc->dispatchThreads(MTL::Size(NS::UInteger(C), 1, 1), MTL::Size(64, 1, 1));416 }417}418419void layernorm_backward(const Tensor& x, const Tensor& w, float eps,420 const Tensor& dout, Tensor& dx, Tensor& dw, Tensor& db) {421 const int64_t C = w.numel();422 const int64_t rows = x.numel() / C;423 Tensor mean = Tensor::empty({rows});424 Tensor istd = Tensor::empty({rows});425 Device& dev = Device::get();426427 {428 MTL::ComputePipelineState* pso = dev.pipeline("layernorm_bwd_dx_f32");429 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();430 enc->setComputePipelineState(pso);431 enc->setBuffer(x.buffer(), x.buffer_offset(), 0);432 enc->setBuffer(w.buffer(), w.buffer_offset(), 1);433 enc->setBuffer(dout.buffer(), dout.buffer_offset(), 2);434 enc->setBuffer(dx.buffer(), dx.buffer_offset(), 3);435 enc->setBuffer(mean.buffer(), mean.buffer_offset(), 4);436 enc->setBuffer(istd.buffer(), istd.buffer_offset(), 5);437 NormParams p{uint32_t(C), eps};438 enc->setBytes(&p, sizeof(p), 6);439 enc->dispatchThreadgroups(MTL::Size(NS::UInteger(rows), 1, 1),440 row_threadgroup(C, pso));441 }442 {443 MTL::ComputePipelineState* pso = dev.pipeline("layernorm_bwd_dwdb_f32");444 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();445 enc->setComputePipelineState(pso);446 enc->setBuffer(x.buffer(), x.buffer_offset(), 0);447 enc->setBuffer(dout.buffer(), dout.buffer_offset(), 1);448 enc->setBuffer(mean.buffer(), mean.buffer_offset(), 2);449 enc->setBuffer(istd.buffer(), istd.buffer_offset(), 3);450 enc->setBuffer(dw.buffer(), dw.buffer_offset(), 4);451 enc->setBuffer(db.buffer(), db.buffer_offset(), 5);452 NormParams p{uint32_t(C), eps};453 enc->setBytes(&p, sizeof(p), 6);454 const uint32_t r32 = uint32_t(rows);455 enc->setBytes(&r32, sizeof(r32), 7);456 enc->dispatchThreads(MTL::Size(NS::UInteger(C), 1, 1), MTL::Size(64, 1, 1));457 }458}459460// ---- rope -----------------------------------------------------------------------461462namespace {463struct RopeParams {464 uint32_t T, H, HD;465 uint32_t pos_offset;466};467468void rope_encode(const Tensor& in, Tensor& out, int64_t n_heads, const Tensor& freqs,469 int64_t pos_offset, bool inverse) {470 const int64_t B = in.size(0), T = in.size(1), C = in.size(2);471 const int64_t hd = C / n_heads;472 check(freqs.numel() == hd / 2, "rope: freqs must have head_dim/2 entries");473 Device& dev = Device::get();474 MTL::FunctionConstantValues* constants = MTL::FunctionConstantValues::alloc()->init();475 // slots 0/1 belong to matmul TA/TB; RoPE uses slot 2476 constants->setConstantValue(&inverse, MTL::DataTypeBool, NS::UInteger(2));477 MTL::ComputePipelineState* pso =478 dev.pipeline("rope_f32", constants, inverse ? "inv" : "fwd");479 constants->release();480481 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();482 enc->setComputePipelineState(pso);483 enc->setBuffer(in.buffer(), in.buffer_offset(), 0);484 enc->setBuffer(out.buffer(), out.buffer_offset(), 1);485 RopeParams p{uint32_t(T), uint32_t(n_heads), uint32_t(hd), uint32_t(pos_offset)};486 enc->setBytes(&p, sizeof(p), 2);487 enc->setBuffer(freqs.buffer(), freqs.buffer_offset(), 3);488 const int64_t pairs = B * T * n_heads * (hd / 2);489 enc->dispatchThreads(MTL::Size(NS::UInteger(pairs), 1, 1), MTL::Size(256, 1, 1));490}491} // namespace492493void rope(const Tensor& x, int64_t n_heads, const Tensor& freqs, int64_t pos_offset,494 Tensor& out) {495 rope_encode(x, out, n_heads, freqs, pos_offset, /*inverse=*/false);496}497498void rope_backward(const Tensor& dout, int64_t n_heads, const Tensor& freqs,499 int64_t pos_offset, Tensor& dx) {500 rope_encode(dout, dx, n_heads, freqs, pos_offset, /*inverse=*/true);501}502503// ---- embedding -------------------------------------------------------------------504505void embedding(const Tensor& weight, const Tensor& ids, Tensor& out) {506 check(ids.dtype() == DType::I32, "embedding: ids must be i32 on GPU");507 const int64_t C = weight.size(1);508 const int64_t N = ids.numel();509 Device& dev = Device::get();510 MTL::ComputePipelineState* pso = dev.pipeline("embedding_fwd_f32");511 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();512 enc->setComputePipelineState(pso);513 enc->setBuffer(weight.buffer(), weight.buffer_offset(), 0);514 enc->setBuffer(ids.buffer(), ids.buffer_offset(), 1);515 enc->setBuffer(out.buffer(), out.buffer_offset(), 2);516 const uint32_t c32 = uint32_t(C);517 enc->setBytes(&c32, sizeof(c32), 3);518 enc->dispatchThreads(MTL::Size(NS::UInteger(C), NS::UInteger(N), 1),519 MTL::Size(std::min<NS::UInteger>(NS::UInteger(C), 64), 4, 1));520}521522void embedding_backward(const Tensor& ids, const Tensor& dout, Tensor& dweight) {523 check(ids.dtype() == DType::I32, "embedding_backward: ids must be i32 on GPU");524 const int64_t C = dweight.size(1);525 const int64_t N = ids.numel();526 Device& dev = Device::get();527 MTL::ComputePipelineState* pso = dev.pipeline("embedding_bwd_f32");528 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();529 enc->setComputePipelineState(pso);530 enc->setBuffer(ids.buffer(), ids.buffer_offset(), 0);531 enc->setBuffer(dout.buffer(), dout.buffer_offset(), 1);532 enc->setBuffer(dweight.buffer(), dweight.buffer_offset(), 2);533 const uint32_t c32 = uint32_t(C);534 enc->setBytes(&c32, sizeof(c32), 3);535 enc->dispatchThreads(MTL::Size(NS::UInteger(C), NS::UInteger(N), 1),536 MTL::Size(std::min<NS::UInteger>(NS::UInteger(C), 64), 4, 1));537}538539// ---- attention -------------------------------------------------------------------540541namespace {542struct AttnParams {543 uint32_t B, T, H, HKV, HD;544 float scale;545 uint32_t causal;546 uint32_t window;547 float softcap;548};549} // namespace550551void attention(const Tensor& q, const Tensor& k, const Tensor& v,552 int64_t n_heads, int64_t n_kv_heads, bool causal, float scale,553 Tensor& out, Tensor* probs_out, int64_t window, float attn_softcap) {554 check(probs_out != nullptr, "attention: probs_out required (unfused path)");555 const int64_t B = q.size(0), T = q.size(1);556 const int64_t hd = q.size(2) / n_heads;557 Device& dev = Device::get();558 MTL::ComputePipelineState* pso = dev.pipeline("attention_fwd_f32");559 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();560 enc->setComputePipelineState(pso);561 enc->setBuffer(q.buffer(), q.buffer_offset(), 0);562 enc->setBuffer(k.buffer(), k.buffer_offset(), 1);563 enc->setBuffer(v.buffer(), v.buffer_offset(), 2);564 enc->setBuffer(out.buffer(), out.buffer_offset(), 3);565 enc->setBuffer(probs_out->buffer(), probs_out->buffer_offset(), 4);566 AttnParams p{uint32_t(B), uint32_t(T), uint32_t(n_heads), uint32_t(n_kv_heads),567 uint32_t(hd), scale, causal ? 1u : 0u, uint32_t(window),568 attn_softcap};569 enc->setBytes(&p, sizeof(p), 5);570 enc->dispatchThreads(MTL::Size(NS::UInteger(B * n_heads * T), 1, 1),571 MTL::Size(64, 1, 1));572}573574void attention_backward(const Tensor& q, const Tensor& k, const Tensor& v,575 const Tensor& probs, const Tensor& out, const Tensor& dout,576 int64_t n_heads, int64_t n_kv_heads, float scale,577 Tensor& dq, Tensor& dk, Tensor& dv, float attn_softcap) {578 const int64_t B = q.size(0), T = q.size(1);579 const int64_t hd = q.size(2) / n_heads;580 Device& dev = Device::get();581 // window is irrelevant here: masked positions carry prob 0 in `probs`.582 AttnParams p{uint32_t(B), uint32_t(T), uint32_t(n_heads), uint32_t(n_kv_heads),583 uint32_t(hd), scale, 1u, 0u, attn_softcap};584585 // D[b,h,i] = dO_i · O_i == rowsum(dP ∘ P): computed once per query row so586 // neither backward kernel needs the O(T) inner recompute.587 Tensor d_term = Tensor::empty({B, n_heads, T});588 {589 MTL::ComputePipelineState* pso = dev.pipeline("attention_bwd_d_f32");590 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();591 enc->setComputePipelineState(pso);592 enc->setBuffer(out.buffer(), out.buffer_offset(), 0);593 enc->setBuffer(dout.buffer(), dout.buffer_offset(), 1);594 enc->setBuffer(d_term.buffer(), d_term.buffer_offset(), 2);595 enc->setBytes(&p, sizeof(p), 3);596 enc->dispatchThreads(MTL::Size(NS::UInteger(B * n_heads * T), 1, 1),597 MTL::Size(64, 1, 1));598 }599 {600 MTL::ComputePipelineState* pso = dev.pipeline("attention_bwd_dq_f32");601 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();602 enc->setComputePipelineState(pso);603 enc->setBuffer(q.buffer(), q.buffer_offset(), 0);604 enc->setBuffer(k.buffer(), k.buffer_offset(), 1);605 enc->setBuffer(v.buffer(), v.buffer_offset(), 2);606 enc->setBuffer(probs.buffer(), probs.buffer_offset(), 3);607 enc->setBuffer(dout.buffer(), dout.buffer_offset(), 4);608 enc->setBuffer(d_term.buffer(), d_term.buffer_offset(), 5);609 enc->setBuffer(dq.buffer(), dq.buffer_offset(), 6);610 enc->setBytes(&p, sizeof(p), 7);611 enc->dispatchThreads(MTL::Size(NS::UInteger(B * n_heads * T), 1, 1),612 MTL::Size(64, 1, 1));613 }614 {615 MTL::ComputePipelineState* pso = dev.pipeline("attention_bwd_dkv_f32");616 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();617 enc->setComputePipelineState(pso);618 enc->setBuffer(q.buffer(), q.buffer_offset(), 0);619 enc->setBuffer(k.buffer(), k.buffer_offset(), 1);620 enc->setBuffer(v.buffer(), v.buffer_offset(), 2);621 enc->setBuffer(probs.buffer(), probs.buffer_offset(), 3);622 enc->setBuffer(dout.buffer(), dout.buffer_offset(), 4);623 enc->setBuffer(d_term.buffer(), d_term.buffer_offset(), 5);624 enc->setBuffer(dk.buffer(), dk.buffer_offset(), 6);625 enc->setBuffer(dv.buffer(), dv.buffer_offset(), 7);626 enc->setBytes(&p, sizeof(p), 8);627 enc->dispatchThreads(MTL::Size(NS::UInteger(B * n_kv_heads * T), 1, 1),628 MTL::Size(64, 1, 1));629 }630}631632// ---- cross entropy -----------------------------------------------------------------633634void cross_entropy(const Tensor& logits, const Tensor& targets, int64_t n_valid,635 Tensor& losses, Tensor& loss_out, Tensor* dlogits) {636 check(targets.dtype() == DType::I32, "cross_entropy: targets must be i32 on GPU");637 const int64_t V = logits.size(1);638 const int64_t N = logits.size(0);639 Device& dev = Device::get();640641 struct CEParams {642 uint32_t V;643 float inv_n;644 uint32_t want_grad;645 } p{uint32_t(V), n_valid > 0 ? 1.0f / float(n_valid) : 0.0f,646 dlogits ? 1u : 0u};647648 MTL::ComputePipelineState* pso = dev.pipeline("cross_entropy_f32");649 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();650 enc->setComputePipelineState(pso);651 enc->setBuffer(logits.buffer(), logits.buffer_offset(), 0);652 enc->setBuffer(targets.buffer(), targets.buffer_offset(), 1);653 enc->setBuffer(losses.buffer(), losses.buffer_offset(), 2);654 // dlogits slot must be bound even when unused655 const Tensor& dl = dlogits ? *dlogits : losses;656 enc->setBuffer(dl.buffer(), dl.buffer_offset(), 3);657 enc->setBytes(&p, sizeof(p), 4);658 enc->dispatchThreadgroups(MTL::Size(NS::UInteger(N), 1, 1), row_threadgroup(V, pso));659660 sum(losses, loss_out, n_valid > 0 ? 1.0f / float(n_valid) : 0.0f);661}662663// ---- optimizer / reductions ----------------------------------------------------------664665void adamw_step(Tensor& w, const Tensor& g, Tensor& m, Tensor& v,666 float lr, float beta1, float beta2, int64_t t, float eps, float wd,667 float grad_scale) {668 struct AdamWParams {669 float lr, beta1, beta2, bc1, bc2, eps, wd, grad_scale;670 } p{lr, beta1, beta2, 1.0f - std::pow(beta1, float(t)),671 1.0f - std::pow(beta2, float(t)), eps, wd, grad_scale};672 encode_flat("adamw_f32", {&w, &g, &m, &v}, &p, sizeof(p), w.numel());673}674675void sumsq(const Tensor& x, Tensor& out) {676 Device& dev = Device::get();677 MTL::ComputePipelineState* pso = dev.pipeline("sumsq_f32");678 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();679 enc->setComputePipelineState(pso);680 enc->setBuffer(x.buffer(), x.buffer_offset(), 0);681 enc->setBuffer(out.buffer(), out.buffer_offset(), 1);682 const uint32_t n = uint32_t(x.numel());683 enc->setBytes(&n, sizeof(n), 2);684 // single threadgroup: partials[0] is the total685 enc->dispatchThreadgroups(MTL::Size(1, 1, 1), MTL::Size(1024, 1, 1));686}687688void sum(const Tensor& x, Tensor& out, float mul) {689 Device& dev = Device::get();690 MTL::ComputePipelineState* pso = dev.pipeline("sum_f32");691 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();692 enc->setComputePipelineState(pso);693 enc->setBuffer(x.buffer(), x.buffer_offset(), 0);694 enc->setBuffer(out.buffer(), out.buffer_offset(), 1);695 const uint32_t n = uint32_t(x.numel());696 enc->setBytes(&n, sizeof(n), 2);697 enc->setBytes(&mul, sizeof(mul), 3);698 enc->dispatchThreadgroups(MTL::Size(1, 1, 1), MTL::Size(1024, 1, 1));699}700701702703// ---- fused (flash) attention -------------------------------------------------704705namespace {706707struct FlashParams {708 uint32_t B, T, H, HKV;709 float scale;710 uint32_t causal;711 uint32_t window; // sliding window; 0 = full. Scalar kernels only — the712 // caller routes window > 0 away from MMA.713};714715// Must match the INSTANTIATE_FLASH list in flash_attention.metal.716constexpr int64_t kFlashHeadDims[] = {16, 32, 48, 64, 80, 96, 128};717718std::string flash_kernel(const char* stem, int64_t head_dim) {719 return std::string(stem) + "_hd" + std::to_string(head_dim);720}721722} // namespace723724bool flash_supported(int64_t head_dim) {725 for (int64_t hd : kFlashHeadDims)726 if (hd == head_dim) return true;727 return false;728}729730void flash_attention(const Tensor& q, const Tensor& k, const Tensor& v,731 int64_t n_heads, int64_t n_kv_heads, bool causal, float scale,732 Tensor& out, Tensor& lse, FlashKernel kernel, int64_t window) {733 const int64_t B = q.size(0), T = q.size(1);734 const int64_t hd = q.size(2) / n_heads;735 check(flash_supported(hd), "flash_attention: unsupported head_dim");736737 if (kernel == FlashKernel::Auto)738 kernel = window > 0 ? FlashKernel::Scalar : FlashKernel::MMA;739 const bool mma = kernel == FlashKernel::MMA;740 check(!(mma && window > 0), "flash_attention: MMA kernel has no window support");741742 Device& dev = Device::get();743 MTL::ComputePipelineState* pso = dev.pipeline(744 flash_kernel(mma ? "flash_attn_fwd_mma_f32" : "flash_attn_fwd_f32", hd));745 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();746 enc->setComputePipelineState(pso);747 enc->setBuffer(q.buffer(), q.buffer_offset(), 0);748 enc->setBuffer(k.buffer(), k.buffer_offset(), 1);749 enc->setBuffer(v.buffer(), v.buffer_offset(), 2);750 enc->setBuffer(out.buffer(), out.buffer_offset(), 3);751 enc->setBuffer(lse.buffer(), lse.buffer_offset(), 4);752 FlashParams p{uint32_t(B), uint32_t(T), uint32_t(n_heads), uint32_t(n_kv_heads),753 scale, causal ? 1u : 0u, uint32_t(window)};754 enc->setBytes(&p, sizeof(p), 5);755 // One threadgroup per (query block, head, batch). Block size and thread756 // count must match the kernel's enums: the scalar kernel is one thread757 // per query row (TGQ=64), the MMA kernel is BQ=32 rows across 4758 // simdgroups (128 threads).759 const NS::UInteger block_q = mma ? 32 : 64;760 const NS::UInteger threads = mma ? 128 : 64;761 check(pso->maxTotalThreadsPerThreadgroup() >= threads,762 "flash_attention: pipeline cannot host the required threadgroup size");763 const NS::UInteger q_blocks = (NS::UInteger(T) + block_q - 1) / block_q;764 enc->dispatchThreadgroups(765 MTL::Size(q_blocks * NS::UInteger(n_heads) * NS::UInteger(B), 1, 1),766 MTL::Size(threads, 1, 1));767}768769void flash_attention_backward(const Tensor& q, const Tensor& k, const Tensor& v,770 const Tensor& out, const Tensor& lse, const Tensor& dout,771 int64_t n_heads, int64_t n_kv_heads, bool causal,772 float scale, Tensor& dq, Tensor& dk, Tensor& dv,773 FlashKernel kernel, int64_t window) {774 const int64_t B = q.size(0), T = q.size(1);775 const int64_t hd = q.size(2) / n_heads;776 check(flash_supported(hd), "flash_attention_backward: unsupported head_dim");777 if (kernel == FlashKernel::Auto)778 kernel = window > 0 ? FlashKernel::Scalar : FlashKernel::MMA;779 const bool mma = kernel == FlashKernel::MMA;780 check(!(mma && window > 0),781 "flash_attention_backward: MMA kernel has no window support");782783 Device& dev = Device::get();784 FlashParams p{uint32_t(B), uint32_t(T), uint32_t(n_heads), uint32_t(n_kv_heads),785 scale, causal ? 1u : 0u, uint32_t(window)};786787 // D[b,h,i] = dO_i . O_i (shared with the unfused path).788 Tensor d_term = Tensor::empty({B, n_heads, T});789 {790 AttnParams ap{uint32_t(B), uint32_t(T), uint32_t(n_heads),791 uint32_t(n_kv_heads), uint32_t(hd), scale, causal ? 1u : 0u,792 0u, 0.0f};793 MTL::ComputePipelineState* pso = dev.pipeline("attention_bwd_d_f32");794 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();795 enc->setComputePipelineState(pso);796 enc->setBuffer(out.buffer(), out.buffer_offset(), 0);797 enc->setBuffer(dout.buffer(), dout.buffer_offset(), 1);798 enc->setBuffer(d_term.buffer(), d_term.buffer_offset(), 2);799 enc->setBytes(&ap, sizeof(ap), 3);800 enc->dispatchThreads(MTL::Size(NS::UInteger(B * n_heads * T), 1, 1),801 MTL::Size(64, 1, 1));802 }803 {804 MTL::ComputePipelineState* pso = dev.pipeline(flash_kernel(805 mma ? "flash_attn_bwd_dq_mma_f32" : "flash_attn_bwd_dq_f32", hd));806 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();807 enc->setComputePipelineState(pso);808 enc->setBuffer(q.buffer(), q.buffer_offset(), 0);809 enc->setBuffer(k.buffer(), k.buffer_offset(), 1);810 enc->setBuffer(v.buffer(), v.buffer_offset(), 2);811 enc->setBuffer(dout.buffer(), dout.buffer_offset(), 3);812 enc->setBuffer(lse.buffer(), lse.buffer_offset(), 4);813 enc->setBuffer(d_term.buffer(), d_term.buffer_offset(), 5);814 enc->setBuffer(dq.buffer(), dq.buffer_offset(), 6);815 enc->setBytes(&p, sizeof(p), 7);816 if (mma) {817 const NS::UInteger q_blocks = (NS::UInteger(T) + 32 - 1) / 32;818 enc->dispatchThreadgroups(819 MTL::Size(q_blocks * NS::UInteger(n_heads) * NS::UInteger(B), 1, 1),820 MTL::Size(128, 1, 1));821 } else {822 const NS::UInteger tg =823 std::min<NS::UInteger>(64, pso->maxTotalThreadsPerThreadgroup());824 enc->dispatchThreads(MTL::Size(NS::UInteger(B * n_heads * T), 1, 1),825 MTL::Size(tg, 1, 1));826 }827 }828 if (mma) {829 // dV then dK as separate kernels: fused, the thread held K, V, dK and830 // dV as fragments and spilled 4352 bytes (measured with gpudebug).831 {832 MTL::ComputePipelineState* pso =833 dev.pipeline(flash_kernel("flash_attn_bwd_dv_mma_f32", hd));834 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();835 enc->setComputePipelineState(pso);836 enc->setBuffer(q.buffer(), q.buffer_offset(), 0);837 enc->setBuffer(k.buffer(), k.buffer_offset(), 1);838 enc->setBuffer(dout.buffer(), dout.buffer_offset(), 2);839 enc->setBuffer(lse.buffer(), lse.buffer_offset(), 3);840 enc->setBuffer(dv.buffer(), dv.buffer_offset(), 4);841 enc->setBytes(&p, sizeof(p), 5);842 const NS::UInteger kv_blocks = (NS::UInteger(T) + 32 - 1) / 32;843 enc->dispatchThreadgroups(844 MTL::Size(kv_blocks * NS::UInteger(n_kv_heads) * NS::UInteger(B), 1, 1),845 MTL::Size(128, 1, 1));846 }847 {848 MTL::ComputePipelineState* pso =849 dev.pipeline(flash_kernel("flash_attn_bwd_dk_mma_f32", hd));850 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();851 enc->setComputePipelineState(pso);852 enc->setBuffer(q.buffer(), q.buffer_offset(), 0);853 enc->setBuffer(k.buffer(), k.buffer_offset(), 1);854 enc->setBuffer(v.buffer(), v.buffer_offset(), 2);855 enc->setBuffer(dout.buffer(), dout.buffer_offset(), 3);856 enc->setBuffer(lse.buffer(), lse.buffer_offset(), 4);857 enc->setBuffer(d_term.buffer(), d_term.buffer_offset(), 5);858 enc->setBuffer(dk.buffer(), dk.buffer_offset(), 6);859 enc->setBytes(&p, sizeof(p), 7);860 const NS::UInteger kv_blocks = (NS::UInteger(T) + 32 - 1) / 32;861 enc->dispatchThreadgroups(862 MTL::Size(kv_blocks * NS::UInteger(n_kv_heads) * NS::UInteger(B), 1, 1),863 MTL::Size(128, 1, 1));864 }865 } else {866 {867 MTL::ComputePipelineState* pso =868 dev.pipeline(flash_kernel("flash_attn_bwd_dv_f32", hd));869 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();870 enc->setComputePipelineState(pso);871 enc->setBuffer(q.buffer(), q.buffer_offset(), 0);872 enc->setBuffer(k.buffer(), k.buffer_offset(), 1);873 enc->setBuffer(dout.buffer(), dout.buffer_offset(), 2);874 enc->setBuffer(lse.buffer(), lse.buffer_offset(), 3);875 enc->setBuffer(dv.buffer(), dv.buffer_offset(), 4);876 enc->setBytes(&p, sizeof(p), 5);877 const NS::UInteger tg =878 std::min<NS::UInteger>(64, pso->maxTotalThreadsPerThreadgroup());879 enc->dispatchThreads(MTL::Size(NS::UInteger(B * n_kv_heads * T), 1, 1),880 MTL::Size(tg, 1, 1));881 }882 {883 MTL::ComputePipelineState* pso =884 dev.pipeline(flash_kernel("flash_attn_bwd_dk_f32", hd));885 MTL::ComputeCommandEncoder* enc = Stream::get().encoder();886 enc->setComputePipelineState(pso);887 enc->setBuffer(q.buffer(), q.buffer_offset(), 0);888 enc->setBuffer(k.buffer(), k.buffer_offset(), 1);889 enc->setBuffer(v.buffer(), v.buffer_offset(), 2);890 enc->setBuffer(dout.buffer(), dout.buffer_offset(), 3);891 enc->setBuffer(lse.buffer(), lse.buffer_offset(), 4);892 enc->setBuffer(d_term.buffer(), d_term.buffer_offset(), 5);893 enc->setBuffer(dk.buffer(), dk.buffer_offset(), 6);894 enc->setBytes(&p, sizeof(p), 7);895 const NS::UInteger tg =896 std::min<NS::UInteger>(64, pso->maxTotalThreadsPerThreadgroup());897 enc->dispatchThreads(MTL::Size(NS::UInteger(B * n_kv_heads * T), 1, 1),898 MTL::Size(tg, 1, 1));899 }900 }901}902903904} // namespace forge::metal905