OrbitQuant Packed Matmul
Packed low-bit matrix multiplication kernel for OrbitQuant inference.
This kernel consumes OrbitQuant packed weight indices, per-row norms, and Lloyd-Max centroids directly, avoiding a full BF16/FP16 dequantized weight cache before the linear projection.
API
import torch
from orbitquant_packed_matmul import matmul_packed_weight
out = matmul_packed_weight(
x,
packed_weight_indices,
row_norms,
centroids,
bits=4,
out_features=3072,
in_features=3072,
bias=bias,
)
Inputs:
x: contiguous or reshapeable tensor with shape[..., in_features].packed_weight_indices:uint8low-bit packed row-major codebook indices.row_norms: row norms with shape[out_features]; CUDA consumes the artifact'sbfloat16values directly, while Metal usesfloat32internally.centroids:float32Lloyd-Max centroids with shape[2**bits].bias: optional projection bias.
x may be float32, float16, or bfloat16. The output has shape
[..., out_features] and the same dtype as x.
The CUDA package also exports the operations used by OrbitQuant's W4A4 runtime:
quantize_activations_int8: token norm, RPBH/FWHT, nearest-codebook assignment, and INT8-surrogate output in one native launch.quantize_activations_packed_w4: the same activation path with packed 4-bit output for the direct packed matmul fallback.matmul_packed_w4a4_int8: direct packed A4/W4 CUDA MMA with fused token norm, row norm, surrogate scales, and bias epilogue.
On CUDA compute capability 8.0 or newer, OrbitQuant normally combines
quantize_activations_int8 with chunked packed-weight decode and Torch's
CUTLASS-backed INT8 matmul. The direct packed MMA operation remains available
when that path is unsupported. Neither path materializes a complete BF16/FP16
weight matrix.
Build And Test
nix --option sandbox relaxed --option max-jobs 1 --option cores 8 \
run .#build-and-copy -L
nix --option sandbox relaxed --option max-jobs 1 --option cores 8 \
run .#ci-test -L
The build produces ABI3 Hugging Face Kernels artifacts under build/ for the
supported backend variants on the current platform. On macOS, sandbox relaxed
or enabled Nix sandboxing is required by kernel-builder. The commands build
local files only; they do not upload to Kernel Hub.
For a faster CUDA-only development build on a machine with a matching Torch and CUDA toolchain:
cargo install --git https://github.com/huggingface/kernels hf-kernel-builder
kernel-builder check-config .
kernel-builder create-pyproject -f .
TORCH_CUDA_ARCH_LIST="8.9" CUDACXX=/usr/local/cuda/bin/nvcc \
python setup.py build_kernel
For a local Metal build compatible with macOS 15 and newer:
cargo install --git https://github.com/huggingface/kernels hf-kernel-builder
kernel-builder check-config .
kernel-builder create-pyproject -f .
MACOSX_DEPLOYMENT_TARGET=15.0 \
CMAKE_ARGS="-DCMAKE_OSX_DEPLOYMENT_TARGET=15.0" \
python setup.py build_kernel
kernel-builder check-abi --macos 15.0 --python-abi 3.9 .
This generated project is for local testing and must not be committed or
distributed without a successful kernel-builder check-abi. Use the Nix build
for redistributable variants.
For direct local imports, add the matching build/torch*-<backend>-<platform>
directory to PYTHONPATH; the torch* variant must match the runtime PyTorch
version:
export PYTHONPATH="/path/to/build/torch212-metal-aarch64-darwin:$PYTHONPATH"
python -c "import orbitquant_packed_matmul; print(orbitquant_packed_matmul)"
For PyTorch 2.9 CUDA inference, set the allocator before starting Python when minimum reserved memory is important:
PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True python generate.py
OrbitQuant detects that importable package before trying any Hub loader. For
Hugging Face kernels local loading instead, set LOCAL_KERNELS to the same
built variant directory containing metadata.json, not to the source package
root:
export LOCAL_KERNELS="WaveCut/orbitquant-packed-matmul=/path/to/build/torch212-metal-aarch64-darwin"
Benchmark
The benchmark reports two PyTorch references:
predequantized_f_linear_seconds_per_iter:torch.nn.functional.linearover a full dequantized weight matrix that was materialized before timing.dequantize_then_f_linear_seconds_per_iter: materialize the full dequantized weight matrix inside each timed iteration, then calltorch.nn.functional.linear.
PYTHONPATH=/path/to/build/torch212-metal-aarch64-darwin \
python benchmarks/benchmark.py \
--device mps \
--bits 4 \
--rows 512 \
--in-features 3072 \
--out-features 3072 \
--iters 20
--rows accepts a comma-separated sweep (the default covers decode-bound
small batches and GEMM-bound large batches), and --dtype selects the
activation dtype explicitly. Headline timings are hot-loop medians.
The script prints JSON with packed_seconds_per_iter,
predequantized_f_linear_seconds_per_iter,
dequantize_then_f_linear_seconds_per_iter (all hot-loop medians), the
per-path *_hot_mean_seconds, *_hot_median_seconds, and
*_hot_p95_seconds distributions,
packed_vs_predequantized_f_linear_speedup,
packed_vs_dequantize_then_f_linear_speedup, compatibility aliases
reference_seconds_per_iter and packed_vs_reference_speedup, and
max_abs_error.
It also reports storage accounting for the packed weight path:
packed_weight_indices_bytes, row_norms_bytes, centroid_bytes,
packed_weight_path_bytes, materialized_weight_bytes, and
packed_weight_path_vs_materialized_weight_ratio. These values describe only
the weight-side storage used by this operator; they are not end-to-end model
VRAM measurements.
Metal reference results
Measured on an Apple M2 Max with Torch 2.12.1, FP16 activations, W4 packed weights, transformer-scale shapes (hot-loop medians over 30 iterations; each iteration synchronizes, so sub-millisecond rows include the MPS submit floor):
| Shape (rows x in x out) | Packed Metal | Resident FP16 F.linear |
Materialize + F.linear |
Packed vs resident | Packed vs materialize |
|---|---|---|---|---|---|
| 1 x 3072 x 3072 | 0.330 ms | 0.200 ms | 1.791 ms | 0.60x | 5.42x |
| 4 x 3072 x 3072 | 0.377 ms | 0.215 ms | 1.768 ms | 0.57x | 4.69x |
| 32 x 3072 x 3072 | 0.367 ms | 0.312 ms | 1.700 ms | 0.85x | 4.63x |
| 512 x 3072 x 3072 | 1.355 ms | 1.077 ms | 2.331 ms | 0.79x | 1.72x |
| 512 x 3072 x 12288 | 5.138 ms | 3.907 ms | 9.312 ms | 0.76x | 1.81x |
| 1 x 3072 x 12288 | 0.417 ms | 0.383 ms | 6.202 ms | 0.92x | 14.88x |
Batches of at most four rows dispatch a skinny-batch GEMV (one simdgroup per output column, decoded weight segments reused across the batch); larger batches dispatch the simdgroup-matrix tiles. The packed weight payload, row norms, and centroids occupy about 25% of the materialized FP16 weight size at W4. The resident reference excludes weight materialization time and retains the complete FP16 matrix in memory — it is the throughput ceiling for a kernel that decodes weights on the fly, not a like-for-like memory configuration.
End-to-end FLUX.2 Klein 9B measurements and the SDNQ comparison are recorded in
docs/flux2-klein-9b-sdnq-vs-orbitquant.md.