diff --git a/turboquant-poc/bench_tq_dequant.py b/turboquant-poc/bench_tq_dequant.py new file mode 100644 index 0000000..22e1940 --- /dev/null +++ b/turboquant-poc/bench_tq_dequant.py @@ -0,0 +1,169 @@ +"""Correctness + benchmark harness for the fused TurboQuant dequantize kernel. + +STATUS: the CUDA half of this file has NEVER BEEN RUN. It was written on a +CPU-only host (no `nvidia-smi`, no `nvcc`), so this repo contains NO measured +numbers for `csrc/tq_dequant.cu`. Run this on a CUDA box to produce them; do not +quote a speedup that did not come out of this script. + +Three modes: + + --check-math-only CPU only, no CUDA, no nvcc. Verifies the algebraic claim + the kernel rests on -- that folding the per-vector norm + into the codebook gather gives bit-comparable results to + the existing `TurboQuantMSE.dequantize`. This is the part + that CAN be verified without a GPU, and it is. + + --check Builds the extension and checks the kernel's output + against `TurboQuantMSE.dequantize` on real KV-shaped + tensors. Requires CUDA + nvcc. + + --bench Times three implementations at KV-cache decode shapes: + torch-eager : the current dequantize (gather -> GEMM -> scale) + fused-cuda : csrc/tq_dequant.cu gather+scale -> one GEMM + triton : ONLY if a Triton implementation exists; this + repo has none (turboquant_poc.py says so in + its module docstring: "no triton, no custom + CUDA, no bit-packing"), so that row is + reported as ABSENT rather than invented. + Requires CUDA. + +Usage: + python turboquant-poc/bench_tq_dequant.py --check-math-only + python turboquant-poc/bench_tq_dequant.py --check --bench +""" + +from __future__ import annotations + +import argparse +import os +import sys +import time + +import torch + +sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) +from turboquant_poc import TurboQuantMSE # noqa: E402 + +CSRC = os.path.join(os.path.dirname(os.path.abspath(__file__)), "csrc", "tq_dequant.cu") + +# (rows, head_dim) at decode: rows = batch * num_kv_heads * seq_len. +# Qwen2.5-14B: 8 KV heads, head_dim 128, 48 layers. +SHAPES = [ + (8 * 1024, 128), # 1k context + (8 * 8192, 128), # 8k context + (8 * 32768, 128), # 32k context +] + + +def _reference_fused(tq: TurboQuantMSE, idx: torch.Tensor, norms: torch.Tensor): + """What the kernel computes, expressed in torch: fold the norm into the + gather, then a single GEMM. Used to check the algebra without a GPU.""" + y_scaled = tq.centroids[idx.long()] * norms.reshape(-1, 1) + return y_scaled @ tq.Pi + + +def check_math_only(bit_width: int = 4, head_dim: int = 128, rows: int = 4096) -> int: + torch.manual_seed(0) + tq = TurboQuantMSE(bit_width=bit_width, head_dim=head_dim, device="cpu") + x = torch.randn(rows, head_dim) + idx, norms = tq.quantize(x) + + baseline = tq.dequantize(idx, norms) # gather -> GEMM -> scale + fused = _reference_fused(tq, idx, norms) # (gather * scale) -> GEMM + + abs_err = (baseline - fused).abs().max().item() + scale = baseline.abs().max().item() + rel_err = abs_err / scale if scale else 0.0 + print(f"rows={rows} head_dim={head_dim} bit_width={bit_width}") + print(f" max |baseline - fused| = {abs_err:.3e} (relative {rel_err:.3e})") + print(f" float32 eps = {torch.finfo(torch.float32).eps:.3e}") + ok = rel_err < 1e-6 + print(" ALGEBRA VERIFIED" if ok else " ALGEBRA MISMATCH") + print("\nNOTE: this verifies the kernel's MATH only. csrc/tq_dequant.cu itself " + "is UNRUN;\n no GPU was available. Use --check --bench on a CUDA host.") + return 0 if ok else 1 + + +def _load_extension(): + from torch.utils.cpp_extension import load + + return load(name="tq_dequant_ext", sources=[CSRC], verbose=True) + + +def check(ext, bit_width: int, head_dim: int, rows: int) -> None: + tq = TurboQuantMSE(bit_width=bit_width, head_dim=head_dim, device="cuda") + x = torch.randn(rows, head_dim, device="cuda") + idx, norms = tq.quantize(x) + baseline = tq.dequantize(idx, norms) + y = ext.tq_dequant_gather_scale( + idx.reshape(-1, head_dim).contiguous(), tq.centroids, + norms.reshape(-1).float().contiguous(), torch.float32) + fused = y @ tq.Pi + err = (baseline - fused).abs().max().item() + print(f" check rows={rows}: max abs err = {err:.3e}") + assert err < 1e-4, f"kernel disagrees with TurboQuantMSE.dequantize: {err}" + + +def _time(fn, iters: int = 50, warmup: int = 10) -> float: + for _ in range(warmup): + fn() + torch.cuda.synchronize() + t0 = time.perf_counter() + for _ in range(iters): + fn() + torch.cuda.synchronize() + return (time.perf_counter() - t0) / iters * 1e3 # ms + + +def bench(ext, bit_width: int) -> None: + print(f"\n{'rows':>9} {'dim':>5} {'torch-eager ms':>15} {'fused-cuda ms':>14} " + f"{'speedup':>8} {'triton ms':>10}") + for rows, dim in SHAPES: + tq = TurboQuantMSE(bit_width=bit_width, head_dim=dim, device="cuda") + x = torch.randn(rows, dim, device="cuda") + idx, norms = tq.quantize(x) + flat_idx = idx.reshape(-1, dim).contiguous() + flat_norms = norms.reshape(-1).float().contiguous() + + eager = _time(lambda: tq.dequantize(idx, norms)) + fused = _time(lambda: ext.tq_dequant_gather_scale( + flat_idx, tq.centroids, flat_norms, torch.float32) @ tq.Pi) + print(f"{rows:>9} {dim:>5} {eager:>15.3f} {fused:>14.3f} " + f"{eager / fused:>7.2f}x {'ABSENT':>10}") + print("\ntriton column is ABSENT because this repo has no Triton implementation " + "of\nTurboQuant to compare against (turboquant_poc.py: 'no triton, no custom " + "CUDA').\nNothing is estimated in its place.") + + +def main() -> int: + ap = argparse.ArgumentParser() + ap.add_argument("--check-math-only", action="store_true") + ap.add_argument("--check", action="store_true") + ap.add_argument("--bench", action="store_true") + ap.add_argument("--bit-width", type=int, default=4) + args = ap.parse_args() + + if args.check_math_only: + return check_math_only(bit_width=args.bit_width) + + if not (args.check or args.bench): + ap.error("pick --check-math-only, --check, and/or --bench") + + if not torch.cuda.is_available(): + print("BLOCKED: no CUDA device. csrc/tq_dequant.cu cannot be built or timed " + "here.\nPrerequisite: an NVIDIA GPU with a matching CUDA toolkit " + "(nvcc) on PATH.\nRun --check-math-only for the part that does work on " + "CPU.") + return 2 + + ext = _load_extension() + if args.check: + for rows, dim in SHAPES: + check(ext, args.bit_width, dim, rows) + if args.bench: + bench(ext, args.bit_width) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/turboquant-poc/csrc/tq_dequant.cu b/turboquant-poc/csrc/tq_dequant.cu new file mode 100644 index 0000000..92ff3b7 --- /dev/null +++ b/turboquant-poc/csrc/tq_dequant.cu @@ -0,0 +1,137 @@ +// TurboQuant dequantize: fused codebook gather + per-vector norm scaling. +// +// STATUS: UNRUN. This file has never been compiled or executed. The host this +// was written on has no GPU -- `nvidia-smi` and `nvcc` are both absent -- so +// there are NO performance numbers for it anywhere in this repo, and none may +// be quoted until it has actually been built and benchmarked on a CUDA device. +// `bench_tq_dequant.py` is the harness that produces those numbers; it refuses +// to run without CUDA rather than estimating. +// +// WHAT IT REPLACES +// ---------------- +// `TurboQuantMSE.dequantize` (turboquant_poc.py) currently does: +// +// y_hat = self.centroids[flat_idx.long()] # (1) uint8 -> int64 copy, +// # then a gather +// x_hat = y_hat @ self.Pi # (2) GEMM +// x_hat = x_hat * norms.reshape(-1, 1) # (3) row scale +// +// Step (1) materialises an int64 index tensor 8x the size of the uint8 codes +// and then a second fp32 tensor for y_hat; step (3) is another full pass over +// the output. On the decode path this runs for every layer on every token, and +// the repo's own note (turboquant-poc/README.md) is that the pure-PyTorch +// dequantize path is why token rate sits below fp16. +// +// THE FUSION +// ---------- +// Row scaling COMMUTES with the rotation, because the rotation is a row-wise +// linear map: +// +// (y_hat @ Pi) * norms[:, None] == (y_hat * norms[:, None]) @ Pi +// +// So the scale can be folded into the gather, leaving exactly one GEMM and no +// int64 index copy: +// +// y_scaled[r][c] = centroids[idx[r][c]] * norms[r] <- this kernel +// x_hat = y_scaled @ Pi <- cuBLAS, unchanged +// +// The codebook is at most 2**bit_width entries (256 at 8 bits, 16 at 4 bits), +// so it is staged in shared memory once per block and every lookup after that +// is a shared-memory read instead of a global gather. +// +// The equivalence above is verified numerically against the existing +// implementation on CPU by `bench_tq_dequant.py --check-math-only`, which needs +// no GPU. The KERNEL is still unrun. + +#include + +#include +#include + +namespace { + +constexpr int kMaxCodebook = 256; // 8-bit codes are the widest TurboQuant uses + +// One thread per output element. `rows` is the flattened vector count +// (batch * heads * seq_len) and `dim` is head_dim. +template +__global__ void tq_dequant_gather_scale_kernel( + const uint8_t* __restrict__ idx, // (rows, dim) + const float* __restrict__ centroids, // (n_levels,) + const float* __restrict__ norms, // (rows,) + scalar_t* __restrict__ out, // (rows, dim) + const long rows, + const int dim, + const int n_levels) { + extern __shared__ float s_centroids[]; + + for (int i = threadIdx.x; i < n_levels; i += blockDim.x) { + s_centroids[i] = centroids[i]; + } + __syncthreads(); + + const long total = rows * static_cast(dim); + const long stride = static_cast(blockDim.x) * gridDim.x; + for (long t = blockIdx.x * static_cast(blockDim.x) + threadIdx.x; + t < total; t += stride) { + const long row = t / dim; + // uint8 load: no int64 index tensor is ever materialised. + const float centroid = s_centroids[static_cast(idx[t])]; + out[t] = static_cast(centroid * norms[row]); + } +} + +} // namespace + +// Returns y_scaled = centroids[idx] * norms[:, None], ready for a single +// `y_scaled @ Pi` GEMM. Caller keeps the GEMM in cuBLAS. +torch::Tensor tq_dequant_gather_scale( + torch::Tensor idx, // (rows, dim), uint8, contiguous, CUDA + torch::Tensor centroids, // (n_levels,), float32, CUDA + torch::Tensor norms, // (rows,), float32, CUDA + c10::ScalarType out_dtype) { + TORCH_CHECK(idx.is_cuda() && centroids.is_cuda() && norms.is_cuda(), + "tq_dequant_gather_scale: all inputs must be CUDA tensors"); + TORCH_CHECK(idx.scalar_type() == torch::kUInt8, "idx must be uint8"); + TORCH_CHECK(centroids.scalar_type() == torch::kFloat32, "centroids must be float32"); + TORCH_CHECK(norms.scalar_type() == torch::kFloat32, "norms must be float32"); + TORCH_CHECK(idx.dim() == 2, "idx must be 2-D (rows, dim)"); + TORCH_CHECK(norms.numel() == idx.size(0), "norms must have one entry per row"); + TORCH_CHECK(centroids.numel() <= kMaxCodebook, + "codebook larger than ", kMaxCodebook, " entries"); + + idx = idx.contiguous(); + centroids = centroids.contiguous(); + norms = norms.contiguous(); + + const long rows = idx.size(0); + const int dim = static_cast(idx.size(1)); + const int n_levels = static_cast(centroids.numel()); + + auto out = torch::empty({rows, dim}, idx.options().dtype(out_dtype)); + + const int threads = 256; + const long total = rows * static_cast(dim); + const int blocks = static_cast(std::min((total + threads - 1) / threads, 65535L)); + const size_t shmem = static_cast(n_levels) * sizeof(float); + auto stream = at::cuda::getCurrentCUDAStream(); + + AT_DISPATCH_FLOATING_TYPES_AND2( + at::ScalarType::Half, at::ScalarType::BFloat16, out_dtype, + "tq_dequant_gather_scale", [&] { + tq_dequant_gather_scale_kernel + <<>>( + idx.data_ptr(), + centroids.data_ptr(), + norms.data_ptr(), + out.data_ptr(), + rows, dim, n_levels); + }); + C10_CUDA_KERNEL_LAUNCH_CHECK(); + return out; +} + +PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { + m.def("tq_dequant_gather_scale", &tq_dequant_gather_scale, + "TurboQuant fused codebook gather + per-vector norm scale (CUDA)"); +}