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%
4.0 KB · 97 lines cpp
Raw Blame History
1// Author: Simon-Pierre Boucher — contact@spboucher.ai2//3// Per-kernel attention timing at real training shapes, so optimization work4// targets whatever actually dominates a step instead of what looks slow.5// Causal attention does T(T+1)/2 of the T^2 pairs; each pair costs 2*hd MACs6// in the forward (QK then PV), so the reported TFLOPs use 4*hd flops/pair.7#include <Foundation/Foundation.hpp>8#include <Metal/Metal.hpp>910#include "core/device.h"11#include "core/tensor.h"12#include "ops/metal/metal_ops.h"1314#include <cstdio>15#include <random>1617namespace {1819double now_gpu() { return forge::metal::Stream::get().gpu_seconds(); }2021void fill(forge::Tensor& t, std::mt19937& rng) {22    std::uniform_real_distribution<float> d(-1.0f, 1.0f);23    for (int64_t i = 0; i < t.numel(); ++i) t.data<float>()[i] = d(rng);24}2526} // namespace2728int main() {29    NS::AutoreleasePool* pool = NS::AutoreleasePool::alloc()->init();30    std::printf("device: %s\n\n", forge::Device::get().name().c_str());3132    struct Case { int64_t B, T, H, HKV, HD; const char* tag; };33    const Case cases[] = {34        {64, 512, 6, 6, 64, "gpt-10m   B64 T512 H6"},35        {32, 128, 4, 2, 64, "gpt-smoke B32 T128 H4"},36        {64, 1024, 8, 8, 64, "gpt-25m   B64 T1024 H8"},37    };3839    std::printf("%-24s %9s %9s %9s %9s %9s\n", "case", "scalar", "mma", "bwd ms",40                "scalarTF", "mmaTF");41    for (const Case& c : cases) {42        const int64_t Cq = c.H * c.HD, Ckv = c.HKV * c.HD;43        std::mt19937 rng(1);44        forge::Tensor q = forge::Tensor::empty({c.B, c.T, Cq});45        forge::Tensor k = forge::Tensor::empty({c.B, c.T, Ckv});46        forge::Tensor v = forge::Tensor::empty({c.B, c.T, Ckv});47        forge::Tensor o = forge::Tensor::empty({c.B, c.T, Cq});48        forge::Tensor dO = forge::Tensor::empty({c.B, c.T, Cq});49        forge::Tensor lse = forge::Tensor::empty({c.B, c.H, c.T});50        forge::Tensor dq = forge::Tensor::zeros({c.B, c.T, Cq});51        forge::Tensor dk = forge::Tensor::zeros({c.B, c.T, Ckv});52        forge::Tensor dv = forge::Tensor::zeros({c.B, c.T, Ckv});53        fill(q, rng); fill(k, rng); fill(v, rng); fill(dO, rng);54        const float scale = 1.0f / std::sqrt(float(c.HD));55        const int iters = 3;5657        using FK = forge::metal::FlashKernel;58        double fwd_by_kernel[2];59        const FK kernels[2] = {FK::Scalar, FK::MMA};60        for (int ki = 0; ki < 2; ++ki) {61            forge::metal::flash_attention(q, k, v, c.H, c.HKV, true, scale, o, lse,62                                          kernels[ki]);63            forge::metal::sync(); // warm up + build pipeline64            const double t = now_gpu();65            for (int i = 0; i < iters; ++i)66                forge::metal::flash_attention(q, k, v, c.H, c.HKV, true, scale, o, lse,67                                              kernels[ki]);68            forge::metal::sync();69            fwd_by_kernel[ki] = (now_gpu() - t) / iters;70        }71        const double fwd = fwd_by_kernel[1];72        double t0;7374        // The backward wrapper runs D + dq + dkv; time the whole thing, then75        // the D+dq part alone, and take dkv as the difference.76        forge::metal::flash_attention_backward(q, k, v, o, lse, dO, c.H, c.HKV, true,77                                               scale, dq, dk, dv);78        forge::metal::sync();79        t0 = now_gpu();80        for (int i = 0; i < iters; ++i)81            forge::metal::flash_attention_backward(q, k, v, o, lse, dO, c.H, c.HKV,82                                                   true, scale, dq, dk, dv);83        forge::metal::sync();84        const double bwd = (now_gpu() - t0) / iters;8586        const double pairs = double(c.B) * c.H * double(c.T) * (c.T + 1) / 2.0;87        const double flops = pairs * 4.0 * double(c.HD);88        std::printf("%-24s %9.2f %9.2f %9.2f %8.2fT %8.2fT\n", c.tag,89                    fwd_by_kernel[0] * 1e3, fwd_by_kernel[1] * 1e3, bwd * 1e3,90                    flops / fwd_by_kernel[0] / 1e12, flops / fwd_by_kernel[1] / 1e12);91        std::fflush(stdout);92    }9394    pool->drain();95    return 0;96}97