From 1e8cabbd804b7ce8e60c12391f2f24e38d008e03 Mon Sep 17 00:00:00 2001 From: lgoyal6 Date: Sun, 6 Sep 2026 16:13:18 -0700 Subject: [PATCH] Add an unrun native CUDA dequant kernel for the TurboQuant PoC TurboQuantMSE.dequantize in turboquant_poc.py does three passes: it promotes the uint8 codes to int64 to index the codebook, materialises an fp32 gather result, runs the GEMM, and then scales the rows in a fourth pass. On the decode path that runs for every layer on every token, and turboquant-poc/README.md already names the PyTorch dequantize path as the reason token rate sits below fp16. 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 folds into the gather and leaves exactly one GEMM and no int64 index copy. tq_dequant_gather_scale does the fused half and deliberately leaves the GEMM to cuBLAS rather than reimplementing it. The codebook is at most 256 entries, so it is staged in shared memory once per block. STATUS: UNRUN, and the file says so in its first paragraph. The host this was written on has no GPU - nvidia-smi and nvcc are both absent - so there are no performance numbers for this kernel anywhere in this repo, and none may be quoted until it has been built and benchmarked on a CUDA device. bench_tq_dequant.py is that harness, and it refuses to run without CUDA rather than estimating. What IS verified is the algebra the fusion rests on: --check-math-only compares the fused result against the existing implementation on CPU. Known issue, reported rather than fixed because it cannot be reproduced on a host with no nvcc: the kernel uses three symbols whose headers it does not include - at::cuda::getCurrentCUDAStream() needs , C10_CUDA_KERNEL_LAUNCH_CHECK() needs , and std::min needs . Whoever takes this to a CUDA box should expect to add them. --- turboquant-poc/bench_tq_dequant.py | 169 +++++++++++++++++++++++++++++ turboquant-poc/csrc/tq_dequant.cu | 137 +++++++++++++++++++++++ 2 files changed, 306 insertions(+) create mode 100644 turboquant-poc/bench_tq_dequant.py create mode 100644 turboquant-poc/csrc/tq_dequant.cu 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)"); +}