File size: 5,175 Bytes
e9c8366 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 | """fp6 — FP6 (E3M2 / E2M3) 6-bit floating-point format primitives.
FP6 E3M2: 1 sign + 3 exponent + 2 mantissa. Bias=3. Wider range.
FP6 E2M3: 1 sign + 2 exponent + 3 mantissa. Bias=1. More precision.
Used in MXFP6 (OCP MX: FP6 + E8M0 block scale, block_size=32) and
NVFP6 (NVIDIA: FP6 + FP8 E4M3 per-block scale).
Storage: 6-bit values (padded to byte boundary: 4 values per 3 bytes, or
simple 8-bit storage with 2 bits unused).
"""
from __future__ import annotations
import torch
def _make_fp6_lut(exp_bits: int, mant_bits: int, bias: int) -> torch.Tensor:
"""Build the 64-entry lookup table for an FP6 format (6-bit codes 0..63).
Code layout: bit 5 = sign, bits 4..(5-exp) = exponent, low bits = mantissa.
"""
n = 1 << 6 # 64 entries
lut = torch.zeros(n, dtype=torch.float32)
exp_mask = (1 << exp_bits) - 1
mant_mask = (1 << mant_bits) - 1
max_exp = exp_mask
emax = (max_exp - 1) - bias
max_finite = (2.0 ** emax) * (1.0 + mant_mask / (2.0 ** mant_bits))
for code in range(n):
sign = -1.0 if (code & 0x20) else 1.0
exp = (code >> mant_bits) & exp_mask
mant = code & mant_mask
if exp == 0:
if mant == 0:
val = 0.0
else:
val = sign * (mant / (2.0 ** mant_bits)) * (2.0 ** (1 - bias))
elif exp == max_exp:
val = sign * max_finite
else:
val = sign * (1.0 + mant / (2.0 ** mant_bits)) * (2.0 ** (exp - bias))
lut[code] = val
return lut
FP6_E3M2_LUT: torch.Tensor = _make_fp6_lut(exp_bits=3, mant_bits=2, bias=3)
FP6_E2M3_LUT: torch.Tensor = _make_fp6_lut(exp_bits=2, mant_bits=3, bias=1)
def quantize_fp6(w: torch.Tensor, lut: torch.Tensor, group_size: int = 32) -> tuple[torch.Tensor, torch.Tensor]:
"""Quantize weights to FP6 codes + per-group absmax scale.
Args:
w: float [out, in] (or flat). in is padded to be divisible by group_size.
lut: 64-entry FP6 LUT (FP6_E3M2_LUT or FP6_E2M3_LUT).
group_size: elements per scale group.
Returns:
(codes, scale): codes int64 [out, in_padded] values [0,63],
scale fp32 [out, num_groups].
"""
w = w.detach().float()
out_features = w.shape[0]
in_features = w.shape[1] if w.dim() > 1 else w.numel()
if w.dim() > 1:
flat = w
else:
flat = w.reshape(1, -1)
out_features = 1
pad = (group_size - (flat.shape[1] % group_size)) % group_size
if pad > 0:
flat = torch.nn.functional.pad(flat, (0, pad))
in_padded = flat.shape[1]
num_groups = in_padded // group_size
grouped = flat.reshape(out_features, num_groups, group_size)
scale = grouped.abs().amax(dim=2).clamp(min=1e-8)
scale_exp = scale.unsqueeze(2).expand_as(grouped)
w_norm = grouped / scale_exp
lut_dev = lut.to(w_norm.device)
lut_max = lut_dev.abs().amax().item()
w_norm = w_norm.clamp(-lut_max, lut_max)
diff = w_norm.unsqueeze(-1) - lut_dev
codes = diff.abs().argmin(dim=-1).to(torch.int64)
codes = codes.reshape(out_features, in_padded)
return codes, scale
def dequantize_fp6(
codes: torch.Tensor, scale: torch.Tensor, lut: torch.Tensor,
group_size: int = 32, in_features: int | None = None,
) -> torch.Tensor:
"""Reconstruct float weight from FP6 codes + per-group scale."""
out_features = codes.shape[0]
in_padded = codes.shape[1]
num_groups = in_padded // group_size
lut_dev = lut.to(codes.device)
w_norm = lut_dev[codes]
w_grouped = w_norm.reshape(out_features, num_groups, group_size)
scale_exp = scale.to(torch.float32).unsqueeze(2).expand_as(w_grouped)
w_deq = (w_grouped * scale_exp).reshape(out_features, in_padded)
if in_features is not None and in_features < in_padded:
w_deq = w_deq[:, :in_features]
return w_deq
def pack_fp6(codes: torch.Tensor) -> torch.Tensor:
"""Pack 6-bit codes: 4 values per 3 bytes (24 bits).
Returns uint8 tensor of packed bytes.
"""
c = codes.to(torch.int32)
assert c.shape[-1] % 4 == 0, "in_features must be divisible by 4 for FP6 packing"
n = c.shape[-1] // 4
# 4 codes -> 3 bytes: [c0(6) c1(6) c2(6) c3(6)] -> [b0: c0|c1h, b1: c1l|c2, b2: c3]
flat = c.reshape(-1, n, 4)
b0 = flat[..., 0] | ((flat[..., 1] & 0x3F) << 6)
b1 = (flat[..., 1] >> 2) | ((flat[..., 2] & 0x0F) << 4)
b2 = (flat[..., 2] >> 4) | (flat[..., 3] << 2)
packed = torch.stack([b0, b1, b2], dim=-1).reshape(*c.shape[:-1], n * 3)
return packed.to(torch.uint8)
def unpack_fp6(packed: torch.Tensor) -> torch.Tensor:
"""Unpack 3 bytes -> 4 FP6 codes."""
p = packed.to(torch.int32)
n = p.shape[-1] // 3
flat = p.reshape(*p.shape[:-1], n, 3)
c0 = flat[..., 0] & 0x3F
c1 = (flat[..., 0] >> 6) | ((flat[..., 1] & 0x0F) << 2)
c2 = (flat[..., 1] >> 4) | ((flat[..., 2] & 0x03) << 4)
c3 = flat[..., 2] >> 2
codes = torch.stack([c0, c1, c2, c3], dim=-1).reshape(*p.shape[:-1], n * 4)
return codes.to(torch.int64)
__all__ = [
"FP6_E3M2_LUT",
"FP6_E2M3_LUT",
"quantize_fp6",
"dequantize_fp6",
"pack_fp6",
"unpack_fp6",
] |