Skip to content

vllm.v1.attention.ops.ultraquant.format

FP4 codepoint + UE8M0 scale constants for the UltraQuant KV cache.

K/V codes are FP4 E2M1. Per-group scales are UE8M0 (one byte, a power of two): s = 2^round(log2(c · absmax)) with c = 0.156. Decode consumes Q as FP8 E4M3 so scaled F8F6F4 MFMA can run natively on CDNA4.

AoS slot layout (group_size=32), per (token, head): bytes [0 .. D/2) : K codes (FP4 nibbles, 2/byte) bytes [D/2 .. D/2 + Gk) : K scales (UE8M0, 1 byte × Gk groups) bytes [D/2 + Gk .. D + Gk) : V codes bytes [D + Gk .. D + 2*Gk) : V scales D=256 → Gk=8 → 272 B/slot. D=128 → Gk=4 → 136 B/slot.

No per-token norm-fold. No V rotation. K is Hadamard-rotated at store.

Functions:

  • get_constant_c –

    Return the fixed UltraQuant scale constant.

  • get_group_size –

    Return the fixed group size required by scaled MFMA.

  • k_codes_bytes –

    Bytes for one head's packed K codes.

  • k_scales_bytes –

    Bytes for one head's K scales (one E8M0 byte per group).

  • slot_size –

    Bytes per (token, head) slot.

  • ue8m0_decode –

    Decode a UE8M0 byte back to fp32. 0 → 0.0 (zero sentinel).

  • ue8m0_encode –

    Snap a positive fp32 scale s to the nearest power of 2 and

Attributes:

  • DEFAULT_CONSTANT_C (float) –

    MSE-optimal scale constant; s_raw = c * absmax, then UE8M0-snapped.

  • FP4_MAX (float) –

    Largest representable FP4 magnitude.

  • GROUP_SIZE (int) –

    Elements per group; one E8M0 scale per group. Fixed at 32 because the

DEFAULT_CONSTANT_C = 0.156 module-attribute

MSE-optimal scale constant; s_raw = c * absmax, then UE8M0-snapped.

FP4_MAX = 6.0 module-attribute

Largest representable FP4 magnitude.

GROUP_SIZE = 32 module-attribute

Elements per group; one E8M0 scale per group. Fixed at 32 because the AMD scaled F8F6F4 MFMA instruction also consumes one E8M0 scale per 32 elements along K. Other group sizes would force a fallback to plain MFMA + accumulator-side scale multiply (no hardware fast-path).

get_constant_c()

Return the fixed UltraQuant scale constant.

Source code in vllm/v1/attention/ops/ultraquant/format.py
def get_constant_c() -> float:
    """Return the fixed UltraQuant scale constant."""
    return DEFAULT_CONSTANT_C

get_group_size()

Return the fixed group size required by scaled MFMA.

Source code in vllm/v1/attention/ops/ultraquant/format.py
def get_group_size() -> int:
    """Return the fixed group size required by scaled MFMA."""
    return GROUP_SIZE

k_codes_bytes(head_dim)

Bytes for one head's packed K codes.

Source code in vllm/v1/attention/ops/ultraquant/format.py
def k_codes_bytes(head_dim: int) -> int:
    """Bytes for one head's packed K codes."""
    if head_dim % 2 != 0:
        raise ValueError(f"head_dim must be even, got {head_dim}")
    return head_dim // 2

k_scales_bytes(head_dim, group_size=None)

Bytes for one head's K scales (one E8M0 byte per group).

Source code in vllm/v1/attention/ops/ultraquant/format.py
def k_scales_bytes(head_dim: int, group_size: int | None = None) -> int:
    """Bytes for one head's K scales (one E8M0 byte per group)."""
    return n_groups(head_dim, group_size)

slot_size(head_dim, group_size=None)

Bytes per (token, head) slot.

Source code in vllm/v1/attention/ops/ultraquant/format.py
def slot_size(head_dim: int, group_size: int | None = None) -> int:
    """Bytes per (token, head) slot."""
    return (
        k_codes_bytes(head_dim)
        + k_scales_bytes(head_dim, group_size)
        + v_codes_bytes(head_dim)
        + v_scales_bytes(head_dim, group_size)
    )

ue8m0_decode(byte)

Decode a UE8M0 byte back to fp32. 0 → 0.0 (zero sentinel).

Source code in vllm/v1/attention/ops/ultraquant/format.py
def ue8m0_decode(byte: int) -> float:
    """Decode a UE8M0 byte back to fp32. 0 → 0.0 (zero sentinel)."""
    if byte == 0:
        return 0.0
    return 2.0 ** (byte - UE8M0_BIAS)

ue8m0_encode(s)

Snap a positive fp32 scale s to the nearest power of 2 and encode as a UE8M0 byte. s <= 0 encodes as 0 (zero sentinel).

Source code in vllm/v1/attention/ops/ultraquant/format.py
def ue8m0_encode(s: float) -> int:
    """Snap a positive fp32 scale `s` to the nearest power of 2 and
    encode as a UE8M0 byte. `s <= 0` encodes as 0 (zero sentinel)."""
    if s <= 0.0 or not math.isfinite(s):
        return 0
    exp = int(round(math.log2(s)))
    if exp < _UE8M0_MIN_EXP:
        return 0
    if exp > _UE8M0_MAX_EXP:
        exp = _UE8M0_MAX_EXP
    return exp + UE8M0_BIAS