convert: NVFP4 (modelopt e2m1) dequant path for convert_fp8_to_int4.py (#83)
convert_fp8_to_int4.py could ingest FP8-blockscale, bf16 and f32
checkpoints, but not NVFP4 — the format NVIDIA modelopt emits for
REAP-pruned GLM-5.2 (quant_algo=NVFP4). Those expert weights are U8
(two e2m1 nibbles/byte) with a per-16-block FP8 scale sidecar
(.weight_scale) plus a per-tensor F32 global (.weight_scale_2).
Adds dequant_nvfp4(): e2m1 LUT (verified 1:1 against
ml_dtypes.float4_e2m1fn), low-nibble=even / high-nibble=odd unpacking,
per-block scale repeat-interleaved over group_size=16, times the global.
classify() now consumes the .weight_scale/.weight_scale_2/.input_scale
sidecars, and dequant() routes U8 tensors with a scale sidecar to the
NVFP4 path (keys is now required, so a stray U8 can't fall through to it).
Existing FP8/bf16/f32 paths are untouched.
Guards against silent corruption rather than trusting the input:
* group_size is fixed at 16 and the block-scale column count is
asserted == ceil(I/16); the old code inferred it from the data
(I // ncols), which misaligns silently on a padded/swizzled scale
layout and hard-crashes on a non-multiple-of-16 I — after a multi-GB
shard download. A partial trailing block is handled by slicing to I.
* .weight_scale_2 is asserted < 1: modelopt stores the small global and
MULTIPLIES; llm-compressor/compressed-tensors stores the reciprocal
(>= 1) and DIVIDES. The two are dtype-identical, so without this a
compressed-tensors checkpoint would corrupt every tensor by ~gscale^2.
--selftest-nvfp4 (no network) asserts the 16-code LUT, an encode->dequant
round-trip to <1e-9, and a dequant->colibri-int4 requant bound (<0.30;
the inherent ~0.17 is informational, not a precision claim).
Claude-Session: https://claude.ai/code/session_01DS7oc65c5RdA9V99otRCwt
Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
@@ -64,6 +64,16 @@ def quant_int2(w, bits): # -> (qbytes U8 [O*ceil(I/4)], s
|
||||
out[:, :vk.shape[1]] |= ((vk + 2).astype(np.uint8) << (k * 2))
|
||||
return out.reshape(-1), s[:, 0].astype(np.float32)
|
||||
|
||||
# ---------- NVFP4 (modelopt) : LUT e2m1 ----------
|
||||
# FP4 e2m1 = 1 sign + 2 exp + 1 mantissa. 16 codici, magnitudini {0,.5,1,1.5,2,3,4,6}.
|
||||
# Bit 3 = segno. Ordine impacchettato (compressed_tensors/vLLM): nibble BASSO = elemento
|
||||
# pari, nibble ALTO = elemento dispari. LUT verificata 1:1 con ml_dtypes.float4_e2m1fn.
|
||||
# EN: FP4 e2m1 = 1 sign + 2 exp + 1 mantissa. 16 codes, magnitudes {0,.5,1,1.5,2,3,4,6}.
|
||||
# EN: bit 3 = sign. Packed order (compressed_tensors/vLLM): LOW nibble = even element,
|
||||
# EN: HIGH nibble = odd element. LUT verified 1:1 against ml_dtypes.float4_e2m1fn.
|
||||
_E2M1 = [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0,
|
||||
-0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0]
|
||||
|
||||
# ---------- classificazione dei tensori ----------
|
||||
def layer_idx(name):
|
||||
p = name.split(".")
|
||||
@@ -73,7 +83,10 @@ def layer_idx(name):
|
||||
return -1
|
||||
|
||||
def classify(name, n_layers, keep_mtp=False, keep_idx=False):
|
||||
if name.endswith("_scale_inv"): return "consumed" # gestito col suo peso
|
||||
if name.endswith("_scale_inv"): return "consumed" # FP8 base: gestito col suo peso
|
||||
# NVFP4 (modelopt): i sidecar delle scale sono consumati insieme al loro .weight U8.
|
||||
# EN: NVFP4 (modelopt): scale sidecars are consumed together with their U8 .weight.
|
||||
if name.endswith((".weight_scale", ".weight_scale_2", ".input_scale")): return "consumed"
|
||||
li = layer_idx(name)
|
||||
if keep_idx:
|
||||
# modalita' --indexer: SOLO i pesi del DSA lightning indexer dei layer principali
|
||||
@@ -95,10 +108,59 @@ def classify(name, n_layers, keep_mtp=False, keep_idx=False):
|
||||
if name.endswith(".weight"): return "q" # attn/dense-mlp/shared (residente)
|
||||
return "f32"
|
||||
|
||||
# ---------- dequant di un tensore (fp8+scale a blocchi / bf16 / f32) ----------
|
||||
def dequant(f, name):
|
||||
# ---------- dequant NVFP4 (modelopt) di UN tensore expert -> f32 [O,I] ----------
|
||||
def dequant_nvfp4(f, name):
|
||||
"""NVFP4 di NVIDIA modelopt (quant_algo=NVFP4, quant_method=modelopt).
|
||||
- `name` U8 [O, I/2] : due nibble e2m1 per byte lungo la dim di
|
||||
contrazione (input); pari=nibble basso, dispari=alto.
|
||||
- `name.weight_scale` F8_E4M3 [O, I/16] : scala per-BLOCCO di 16 elementi (group_size=16),
|
||||
lungo la dim di input. Decodifica f8e4m3 -> f32.
|
||||
- `name.weight_scale_2` F32 [] : scala GLOBALE per-tensore, ~amax/(6*448) (piccola).
|
||||
Dequant (convenzione modelopt = MOLTIPLICA, NON dividere):
|
||||
W[o,i] = e2m1_lut[nibble] * f8_block_scale[o, i//16] * weight_scale_2
|
||||
FOOTGUN: llm-compressor/compressed-tensors memorizza il RECIPROCO (global grande) e DIVIDE;
|
||||
modelopt memorizza il valore piccolo e MOLTIPLICA. Questo checkpoint e' modelopt -> moltiplica.
|
||||
EN: NVIDIA modelopt NVFP4. LOW nibble=even elem, HIGH=odd. weight_scale = per-16-block FP8
|
||||
EN: (group_size 16) along the input dim; weight_scale_2 = per-tensor global FP32 (~amax/2688,
|
||||
EN: small). Dequant MULTIPLIES both scales. FOOTGUN: llm-compressor stores the reciprocal
|
||||
EN: (large global) and DIVIDES; modelopt stores the small value and MULTIPLIES."""
|
||||
import torch
|
||||
GS = 16 # NVFP4: block scale ogni 16 elementi
|
||||
packed = f.get_tensor(name) # uint8 [O, I/2]
|
||||
bscale = f.get_tensor(name + "_scale").to(torch.float32) # [O, ceil(I/16)] da f8e4m3
|
||||
gscale = f.get_tensor(name + "_scale_2").to(torch.float32) # scalare per-tensore
|
||||
O, Ih = packed.shape; I = Ih * 2
|
||||
# Convenzione: modelopt memorizza il global PICCOLO e MOLTIPLICA. Se e' >=1 e'
|
||||
# quasi certamente il reciproco di compressed-tensors (che DIVIDE) -> ci fermiamo
|
||||
# invece di corrompere silenziosamente ogni tensore. EN: guard modelopt-vs-CT.
|
||||
assert float(gscale) < 1.0, (
|
||||
f"{name}: weight_scale_2={float(gscale):.4g} >= 1 sembra il reciproco "
|
||||
"(compressed-tensors, che DIVIDE); questo path assume modelopt (MOLTIPLICA)")
|
||||
# Il layout deve essere lo scale per-blocco piatto di modelopt: una colonna ogni
|
||||
# 16 elementi di input (niente swizzle cutlass/TensorRT). Verifichiamo, non deduciamo:
|
||||
# dedurre gs = I // ncol misallinea in silenzio su layout paddati/swizzati.
|
||||
nb = (I + GS - 1) // GS
|
||||
assert bscale.shape[1] == nb, (
|
||||
f"{name}: weight_scale ha {bscale.shape[1]} colonne, attese {nb} = ceil({I}/{GS}); "
|
||||
"layout scale inatteso (swizzled/paddato?), rifiuto per non corrompere")
|
||||
lut = torch.tensor(_E2M1, dtype=torch.float32)
|
||||
nib = torch.empty((O, I), dtype=torch.long)
|
||||
nib[:, 0::2] = (packed & 0x0F).to(torch.long) # elemento pari = nibble basso
|
||||
nib[:, 1::2] = ((packed >> 4) & 0x0F).to(torch.long) # elemento dispari = nibble alto
|
||||
w4 = lut[nib] # [O, I] valori e2m1
|
||||
sc = bscale.repeat_interleave(GS, dim=1)[:, :I] # blocco parziale di coda: slice a I
|
||||
return (w4 * sc * gscale).numpy()
|
||||
|
||||
# ---------- dequant di un tensore (nvfp4 / fp8+scale a blocchi / bf16 / f32) ----------
|
||||
def dequant(f, name, keys):
|
||||
import torch
|
||||
sl = f.get_slice(name); dt = sl.get_dtype()
|
||||
# NVFP4 (modelopt): pesi expert U8 con sidecar `.weight_scale`. In questo checkpoint gli
|
||||
# UNICI tensori U8 sono gli expert NVFP4, ma richiediamo comunque il sidecar (keys e'
|
||||
# obbligatorio: senza, un qualunque U8 verrebbe decodificato come NVFP4).
|
||||
# EN: NVFP4 expert weights are U8 with a `.weight_scale` sidecar; require the sidecar.
|
||||
if dt in ("U8", "uint8") and (name + "_scale") in keys:
|
||||
return dequant_nvfp4(f, name)
|
||||
if dt in ("F8_E4M3", "float8_e4m3fn"):
|
||||
w = f.get_tensor(name).to(torch.float32)
|
||||
sc = f.get_tensor(name + "_scale_inv").to(torch.float32) # [ceil(O/128),ceil(I/128)]
|
||||
@@ -110,10 +172,11 @@ def dequant(f, name):
|
||||
def convert_shard(path, out_dict, n_layers, ebits, io_bits, xbits, keep_mtp=False, keep_idx=False):
|
||||
from safetensors import safe_open
|
||||
with safe_open(path, framework="pt") as f:
|
||||
keys = set(f.keys())
|
||||
for name in f.keys():
|
||||
kind = classify(name, n_layers, keep_mtp, keep_idx)
|
||||
if kind in ("skip", "consumed"): continue
|
||||
w = dequant(f, name)
|
||||
w = dequant(f, name, keys)
|
||||
if kind == "f32":
|
||||
out_dict[name] = w.astype(np.float32)
|
||||
else:
|
||||
@@ -138,6 +201,8 @@ def main():
|
||||
ap.add_argument("--n-layers", type=int, default=78)
|
||||
ap.add_argument("--min-free-gb", type=float, default=20.0)
|
||||
ap.add_argument("--selftest", action="store_true")
|
||||
ap.add_argument("--selftest-nvfp4", action="store_true",
|
||||
help="unit-test del dequant NVFP4 (LUT e2m1 + round-trip), nessun download / no network")
|
||||
ap.add_argument("--mtp", action="store_true",
|
||||
help="download and convert ONLY the MTP head (model.layers.<n_layers>.*) -> out-mtp-*.safetensors")
|
||||
ap.add_argument("--indexer", action="store_true",
|
||||
@@ -152,6 +217,66 @@ def main():
|
||||
a.ebits = 8 if (a.mtp or a.indexer) else 4
|
||||
if a.xbits is None: a.xbits = a.ebits
|
||||
|
||||
if a.selftest_nvfp4:
|
||||
import torch
|
||||
# 1) LUT e2m1: i 16 codici devono decodificare esattamente ai valori attesi.
|
||||
lut = torch.tensor(_E2M1, dtype=torch.float32)
|
||||
expect = [0.0,0.5,1.0,1.5,2.0,3.0,4.0,6.0,-0.0,-0.5,-1.0,-1.5,-2.0,-3.0,-4.0,-6.0]
|
||||
assert lut.tolist() == expect, "LUT e2m1 errata"
|
||||
print("[nvfp4] LUT e2m1: 16/16 codici OK")
|
||||
# 2) round-trip: costruisco un tensore ai SOLI valori rappresentabili (scala nota per
|
||||
# blocco+globale), impacchetto come modelopt, poi dequant deve tornare ESATTO.
|
||||
import numpy as np, io
|
||||
from safetensors.torch import save as st_save
|
||||
from safetensors import safe_open
|
||||
rng = np.random.default_rng(0); O, I, GS = 8, 64, 16
|
||||
codes = rng.integers(0, 16, size=(O, I)).astype(np.uint8) # nibble e2m1 casuali
|
||||
w4 = np.array(_E2M1, np.float32)[codes] # [O,I]
|
||||
# scale per-blocco (rappresentabili in f8e4m3) + globale piccola (stile modelopt)
|
||||
blk = rng.choice([0.5,1.0,2.0,4.0,8.0], size=(O, I//GS)).astype(np.float32)
|
||||
gscale = np.float32(3.9e-5)
|
||||
W = w4 * np.repeat(blk, GS, axis=1) * gscale # riferimento esatto
|
||||
# impacchetto: pari->nibble basso, dispari->alto
|
||||
packed = (codes[:, 0::2] | (codes[:, 1::2] << 4)).astype(np.uint8)
|
||||
import ml_dtypes # solo per il test: encode f8e4m3 delle scale di blocco
|
||||
tens = {name: torch.from_numpy(arr) for name, arr in {
|
||||
"w.weight": packed,
|
||||
"w.weight_scale": blk.astype(ml_dtypes.float8_e4m3fn).view(np.uint8), # placeholder
|
||||
}.items()}
|
||||
# torch non ha un costruttore da bytes f8: passo via file safetensors scritto a mano.
|
||||
# piu' semplice: uso direttamente dequant_nvfp4 su un finto 'f' in-memory.
|
||||
class _F:
|
||||
def __init__(s, d): s.d = d
|
||||
def get_tensor(s, n): return s.d[n]
|
||||
def get_slice(s, n): return None
|
||||
blk_f8 = blk.astype(ml_dtypes.float8_e4m3fn) # quantizza le scale a f8
|
||||
f = _F({"w.weight": torch.from_numpy(packed),
|
||||
"w.weight_scale": torch.from_numpy(blk_f8.view(np.uint8)).view(torch.float8_e4m3fn),
|
||||
"w.weight_scale_2": torch.tensor(gscale)})
|
||||
got = dequant_nvfp4(f, "w.weight")
|
||||
# riferimento con scale gia' quantizzate a f8 (per confronto esatto)
|
||||
Wq = w4 * np.repeat(blk_f8.astype(np.float32), GS, axis=1) * gscale
|
||||
maxerr = float(np.abs(got - Wq).max())
|
||||
print(f"[nvfp4] round-trip encode->dequant: max abs err = {maxerr:.3e} "
|
||||
f"({'OK' if maxerr < 1e-9 else 'FAIL'})")
|
||||
assert maxerr < 1e-9
|
||||
# 3) requant colibri int4 su valori dequantati -> errore piccolo atteso
|
||||
q, s = quant_int4(got.astype(np.float32), 4)
|
||||
rb = (I + 1)//2; qb = q.reshape(O, rb)
|
||||
lo = (qb & 0x0F).astype(np.int32) - 8; hi = ((qb >> 4) & 0x0F).astype(np.int32) - 8
|
||||
deq = np.empty((O, I), np.float32); deq[:, 0::2] = lo; deq[:, 1::2] = hi[:, :I-I//2]
|
||||
deq = deq * s[:, None]
|
||||
rel = np.abs(deq - got).mean() / (np.abs(got).mean() + 1e-12)
|
||||
# Informativo, NON un test di uguaglianza: requantizzare int4 per-riga dati che
|
||||
# spaziano 16x per il block-scale costa ~0.17 di errore relativo di suo. La soglia
|
||||
# larga becca solo una corruzione grossolana, non e' un bound di precisione.
|
||||
# EN: informational — per-row int4 requant of 16x-block-range data inherently ~0.17.
|
||||
print(f"[nvfp4] dequant->colibri int4->dequant: errore rel medio = {rel:.4f} "
|
||||
f"(atteso ~0.17; {'OK' if rel < 0.30 else 'ANOMALO'})")
|
||||
assert rel < 0.30, f"requant rel err {rel:.3f} troppo alto: dequant probabilmente corrotto"
|
||||
print("[nvfp4] SELFTEST OK")
|
||||
return
|
||||
|
||||
if a.selftest:
|
||||
import torch
|
||||
w = (torch.randn(256, 256) * 0.3)
|
||||
|
||||
Reference in New Issue
Block a user