Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
169 changes: 169 additions & 0 deletions turboquant-poc/bench_tq_dequant.py
Original file line number Diff line number Diff line change
@@ -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())
137 changes: 137 additions & 0 deletions turboquant-poc/csrc/tq_dequant.cu
Original file line number Diff line number Diff line change
@@ -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 <torch/extension.h>

#include <cuda.h>
#include <cuda_runtime.h>

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 <typename scalar_t>
__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<long>(dim);
const long stride = static_cast<long>(blockDim.x) * gridDim.x;
for (long t = blockIdx.x * static_cast<long>(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<int>(idx[t])];
out[t] = static_cast<scalar_t>(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<int>(idx.size(1));
const int n_levels = static_cast<int>(centroids.numel());

auto out = torch::empty({rows, dim}, idx.options().dtype(out_dtype));

const int threads = 256;
const long total = rows * static_cast<long>(dim);
const int blocks = static_cast<int>(std::min<long>((total + threads - 1) / threads, 65535L));
const size_t shmem = static_cast<size_t>(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<scalar_t>
<<<blocks, threads, shmem, stream>>>(
idx.data_ptr<uint8_t>(),
centroids.data_ptr<float>(),
norms.data_ptr<float>(),
out.data_ptr<scalar_t>(),
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)");
}
Loading