vllm.v1.attention.ops.ultraquant.reference
¶
PyTorch reference for the UltraQuant KV cache format.
Scale is UE8M0 (power of two): after s = c * absmax snap
s = 2^round(log2(s)) and store one byte. Q is Hadamard-rotated then
cast to FP8 E4M3 before the QK matmul, matching the kernel launcher.
Same FP4 E2M1 grid. c = 0.156. No per-token L2 norm. No V rotation.
K is rotated at encode; Q arrives pre-rotated at decode.
This module is the bit-comparison ground truth for the kernels.
Classes:
-
UltraQuantEncoded–Encoded ultraquant representation of one tensor along its last axis.
Functions:
-
hadamard_matrix–Sylvester Hadamard, normalised so H @ H.T = I.
-
reference_ultraquant_attention–End-to-end reference attention with ultraquant K/V encode→decode in
-
ultraquant_dequant–Decode
encodedto fp32 (no inverse rotation — result in the -
ultraquant_encode–Encode
x(shape [..., D]) to FP4 codes + E8M0 per-group-of-32 scales. -
ultraquant_encode_decode–Full encode→decode round-trip in the rotated basis. Returns fp32.
UltraQuantEncoded
dataclass
¶
Encoded ultraquant representation of one tensor along its last axis.
Shapes assume input shape [..., D]:
- codes_packed: [..., D // 2] uint8 — 2 FP4 nibbles per byte
- scale_bytes: [..., D // GROUP_SIZE] uint8 — E8M0 byte per group
Dequant is code · 2^(byte - 127) (c is folded into the byte).
Attributes:
-
scales_fp32(Tensor) –Decode the E8M0 bytes back to fp32 (for dequant / inspection).
Source code in vllm/v1/attention/ops/ultraquant/reference.py
scales_fp32
property
¶
Decode the E8M0 bytes back to fp32 (for dequant / inspection).
_snap_to_sorted_idx(x_norm, *, dtype=torch.float32)
¶
Snap normalised values to nearest FP4 level, returning sorted index
in [0, 14]. Same semantic as torch.bucketize against midpoints
(reproduced by 14 tl.where in the kernel).
Source code in vllm/v1/attention/ops/ultraquant/reference.py
_ue8m0_snap_tensor(s_raw)
¶
Snap a positive-scale tensor s_raw to a power of two and return
(s_snapped_fp32, byte_uint8).
Matches the kernel: zero / non-finite inputs map to byte=0 / value=0.0.
Round-half-up via floor(log2(s) + 0.5), matching Triton's
tl.floor(... + 0.5) idiom.
Source code in vllm/v1/attention/ops/ultraquant/reference.py
hadamard_matrix(dim, device, dtype=torch.float32)
¶
Sylvester Hadamard, normalised so H @ H.T = I.
Source code in vllm/v1/attention/ops/ultraquant/reference.py
reference_ultraquant_attention(query, key, value, *, scale, sinks=None, constant_c=None)
¶
End-to-end reference attention with ultraquant K/V encode→decode in the rotated basis. Q is Hadamard-rotated AND cast through FP8 E4M3 (scale = 1) before the QK matmul — mirrors what the launcher does.
Output shape: [B, Hq, D] in query.dtype.
Source code in vllm/v1/attention/ops/ultraquant/reference.py
ultraquant_dequant(encoded, head_dim)
¶
Decode encoded to fp32 (no inverse rotation — result in the
rotated basis if rotation was applied at encode).
Dequant is code · 2^(byte - 127) (c is already folded into the byte).
Source code in vllm/v1/attention/ops/ultraquant/reference.py
ultraquant_encode(x, *, rotate=True, constant_c=None)
¶
Encode x (shape [..., D]) to FP4 codes + E8M0 per-group-of-32 scales.
Steps:
1. Optional Hadamard rotation: x_rot = x @ H.T
2. Reshape to [..., G, GROUP_SIZE].
3. absmax = max(|x_rot|) per group.
4. s_raw = c · absmax; s_snapped = 2^round(log2(s_raw)).
Zero-amax groups encode byte=0.
5. sorted_idx = bucketize(x_g / s_snapped, midpoints) ∈ [0, 14].
6. Remap sorted_idx → FP4 E2M1 bit pattern via SORTED_TO_BITS.
7. Pack pairs of 4-bit codes into bytes.
Source code in vllm/v1/attention/ops/ultraquant/reference.py
ultraquant_encode_decode(x, *, rotate=True, constant_c=None)
¶
Full encode→decode round-trip in the rotated basis. Returns fp32.