SPB Git

spb/forge Public MIT

Forge — LLM training from scratch in pure C++20 + Metal on Apple Silicon.

C++ 61.2% C 23% Python 7.6% TeX 7.2% CMake 1.1%
38.7 KB · 905 lines cpp
Raw Blame History
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