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%
3.2 KB · 86 lines cpp
Raw Blame History
1// Author: Simon-Pierre Boucher — contact@spboucher.ai2//3// GEMM throughput benchmark. Times each kernel over the shapes a transformer4// step actually issues (fwd X·Wᵀ, dX = dY·W, dW = dYᵀ·X) plus square shapes5// for a clean TFLOPs number. GPU time comes from command-buffer6// GPUStartTime/GPUEndTime, so it excludes CPU encode overhead.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>16#include <vector>1718using forge::metal::MatmulKernel;1920namespace {2122struct Case {23    int64_t M, K, N;24    bool ta, tb;25    const char* label;26};2728double bench(const Case& c, MatmulKernel kernel, int iters) {29    forge::Tensor a = c.ta ? forge::Tensor::empty({c.K, c.M})30                           : forge::Tensor::empty({c.M, c.K});31    forge::Tensor b = c.tb ? forge::Tensor::empty({c.N, c.K})32                           : forge::Tensor::empty({c.K, c.N});33    forge::Tensor out = forge::Tensor::empty({c.M, c.N});34    std::mt19937 rng(7);35    std::uniform_real_distribution<float> dist(-1.0f, 1.0f);36    for (int64_t i = 0; i < a.numel(); ++i) a.data<float>()[i] = dist(rng);37    for (int64_t i = 0; i < b.numel(); ++i) b.data<float>()[i] = dist(rng);3839    // warmup (also builds the pipeline)40    forge::metal::matmul(a, b, out, c.ta, c.tb, false, kernel);41    forge::metal::sync();4243    const double t0 = forge::metal::Stream::get().gpu_seconds();44    for (int i = 0; i < iters; ++i)45        forge::metal::matmul(a, b, out, c.ta, c.tb, false, kernel);46    forge::metal::sync();47    const double elapsed = forge::metal::Stream::get().gpu_seconds() - t0;4849    const double flops = 2.0 * double(c.M) * double(c.N) * double(c.K) * iters;50    return flops / elapsed / 1e12; // TFLOP/s51}5253} // namespace5455int main() {56    NS::AutoreleasePool* pool = NS::AutoreleasePool::alloc()->init();57    std::printf("device: %s\n\n", forge::Device::get().name().c_str());5859    // gpt-25m-ish shapes: batch*seq = 64*1024 rows, d_model 512, d_ff 140860    const std::vector<Case> cases = {61        {4096, 4096, 4096, false, false, "square 4096"},62        {2048, 2048, 2048, false, false, "square 2048"},63        {1024, 1024, 1024, false, false, "square 1024"},64        {65536, 512, 1408, false, true, "fwd  mlp  X·W1ᵀ"},65        {65536, 1408, 512, false, true, "fwd  mlp  H·W2ᵀ"},66        {65536, 512, 512, false, true, "fwd  attn X·Wqᵀ"},67        {65536, 512, 1408, false, false, "bwd  dX = dY·W"},68        {512, 65536, 1408, true, false, "bwd  dW = dYᵀ·X"},69        {65536, 512, 4096, false, true, "lm head (vocab 4096)"},70    };7172    std::printf("%-24s %10s %10s %10s\n", "case", "naive", "tiled", "simd");73    for (const Case& c : cases) {74        const double gflop = 2.0 * double(c.M) * double(c.N) * double(c.K) / 1e9;75        const int iters = gflop > 50.0 ? 3 : 10;76        const double n = bench(c, MatmulKernel::Naive, iters);77        const double t = bench(c, MatmulKernel::Tiled, iters);78        const double s = bench(c, MatmulKernel::Simdgroup, iters);79        std::printf("%-24s %9.2fT %9.2fT %9.2fT\n", c.label, n, t, s);80        std::fflush(stdout);81    }8283    pool->drain();84    return 0;85}86