// Author: Simon-Pierre Boucher — contact@spboucher.ai // // Per-kernel attention timing at real training shapes, so optimization work // targets whatever actually dominates a step instead of what looks slow. // Causal attention does T(T+1)/2 of the T^2 pairs; each pair costs 2*hd MACs // in the forward (QK then PV), so the reported TFLOPs use 4*hd flops/pair. #include #include #include "core/device.h" #include "core/tensor.h" #include "ops/metal/metal_ops.h" #include #include namespace { double now_gpu() { return forge::metal::Stream::get().gpu_seconds(); } void fill(forge::Tensor& t, std::mt19937& rng) { std::uniform_real_distribution d(-1.0f, 1.0f); for (int64_t i = 0; i < t.numel(); ++i) t.data()[i] = d(rng); } } // namespace int main() { NS::AutoreleasePool* pool = NS::AutoreleasePool::alloc()->init(); std::printf("device: %s\n\n", forge::Device::get().name().c_str()); struct Case { int64_t B, T, H, HKV, HD; const char* tag; }; const Case cases[] = { {64, 512, 6, 6, 64, "gpt-10m B64 T512 H6"}, {32, 128, 4, 2, 64, "gpt-smoke B32 T128 H4"}, {64, 1024, 8, 8, 64, "gpt-25m B64 T1024 H8"}, }; std::printf("%-24s %9s %9s %9s %9s %9s\n", "case", "scalar", "mma", "bwd ms", "scalarTF", "mmaTF"); for (const Case& c : cases) { const int64_t Cq = c.H * c.HD, Ckv = c.HKV * c.HD; std::mt19937 rng(1); forge::Tensor q = forge::Tensor::empty({c.B, c.T, Cq}); forge::Tensor k = forge::Tensor::empty({c.B, c.T, Ckv}); forge::Tensor v = forge::Tensor::empty({c.B, c.T, Ckv}); forge::Tensor o = forge::Tensor::empty({c.B, c.T, Cq}); forge::Tensor dO = forge::Tensor::empty({c.B, c.T, Cq}); forge::Tensor lse = forge::Tensor::empty({c.B, c.H, c.T}); forge::Tensor dq = forge::Tensor::zeros({c.B, c.T, Cq}); forge::Tensor dk = forge::Tensor::zeros({c.B, c.T, Ckv}); forge::Tensor dv = forge::Tensor::zeros({c.B, c.T, Ckv}); fill(q, rng); fill(k, rng); fill(v, rng); fill(dO, rng); const float scale = 1.0f / std::sqrt(float(c.HD)); const int iters = 3; using FK = forge::metal::FlashKernel; double fwd_by_kernel[2]; const FK kernels[2] = {FK::Scalar, FK::MMA}; for (int ki = 0; ki < 2; ++ki) { forge::metal::flash_attention(q, k, v, c.H, c.HKV, true, scale, o, lse, kernels[ki]); forge::metal::sync(); // warm up + build pipeline const double t = now_gpu(); for (int i = 0; i < iters; ++i) forge::metal::flash_attention(q, k, v, c.H, c.HKV, true, scale, o, lse, kernels[ki]); forge::metal::sync(); fwd_by_kernel[ki] = (now_gpu() - t) / iters; } const double fwd = fwd_by_kernel[1]; double t0; // The backward wrapper runs D + dq + dkv; time the whole thing, then // the D+dq part alone, and take dkv as the difference. forge::metal::flash_attention_backward(q, k, v, o, lse, dO, c.H, c.HKV, true, scale, dq, dk, dv); forge::metal::sync(); t0 = now_gpu(); for (int i = 0; i < iters; ++i) forge::metal::flash_attention_backward(q, k, v, o, lse, dO, c.H, c.HKV, true, scale, dq, dk, dv); forge::metal::sync(); const double bwd = (now_gpu() - t0) / iters; const double pairs = double(c.B) * c.H * double(c.T) * (c.T + 1) / 2.0; const double flops = pairs * 4.0 * double(c.HD); std::printf("%-24s %9.2f %9.2f %9.2f %8.2fT %8.2fT\n", c.tag, fwd_by_kernel[0] * 1e3, fwd_by_kernel[1] * 1e3, bwd * 1e3, flops / fwd_by_kernel[0] / 1e12, flops / fwd_by_kernel[1] / 1e12); std::fflush(stdout); } pool->drain(); return 0; }