AI systems
FlashGraph
A C++/CUDA inference engine that fuses a transformer block's RMSNorm, MatMul and GELU into one INT8 GPU kernel.
- Role
- Personal project
- Type
- AI systems
- Links
- Source on GitHub
- GPU launch for three ops
- 1
- weights, FP16 compute
- INT8
- mallocs after init
- 0
How it fits together
- Python APIfused_inference(), quantize_weights(), cpu_inference() and benchmark()
- zero-copy
- C++ bindingValidates inputs and hands PyTorch tensors to the kernel as raw pointers, zero-copy
- launch
- CUDA fused kernelRMSNorm, then MatMul (INT8 to FP16), then GELU in one launch with shared-memory tiling
- allocate
- Memory arena64-byte aligned bump allocator: no malloc after initialisation
The problem
Run naively, a transformer block launches a separate GPU kernel for each operation and sends every intermediate result back through global memory. FlashGraph is an exercise in how far fusion, quantisation and careful memory management can go on a deliberately small block.
How I approached it
- 1
Fuse
RMSNorm, MatMul and GELU run in a single GPU launch with intermediates kept in on-chip memory, using shared-memory tiling (64x64x32) and FP32 accumulators.
- 2
Quantise
Symmetric per-tensor INT8 weights are stored in global memory and dequantised to FP16 on the fly. FP16 maths uses vectorised half2 operations, accumulating in FP32 so values do not overflow.
- 3
Manage memory
A 64-byte aligned bump-pointer arena (posix_memalign) means zero malloc after initialisation, which also gives a standalone C++ deployment path. Pinned host memory and non-blocking CUDA streams let transfers overlap compute.
- 4
Validate and profile
A CPU baseline checks the GPU output. Nsight Compute roofline and bandwidth commands are included for profiling the fused kernel.
What I built
- A PyTorch extension (flashgraph) with a small Python API, plus a CPU-only build with Make or CMake so the C++ side can be built and tested without a GPU.
- A Colab notebook for building and running on a cloud GPU.
The result
A working C++/CUDA extension callable from PyTorch, with a CPU reference for correctness, four C++ arena unit tests and a Python validation suite.
Built with
- C++
- CUDA
- PyTorch
- Python
- CMake
- Nsight Compute
- pytest