Allow remaining FP8 tensors to be dequantized during conversion

This commit is contained in:
Oliver Simons
2026-10-09 21:06:17 +02:00
parent ce9ffde2fc
commit 67d467dbfc
+35 -1
View File
@@ -521,7 +521,41 @@ class ModelBase:
nvfp4_compressed_tensors = self._is_nvfp4_compressed_tensors(quant_method, quant_format, groups)
if nvfp4_compressed_tensors:
dequant_fp8()
fp8_configs = [
g["weights"] for g in groups.values()
if isinstance(g, dict) and g.get("format") == "float-quantized"
and g.get("weights", {}).get("type") == "float"
and g.get("weights", {}).get("num_bits") == 8
]
for name in self.model_tensors.keys():
# Preserved FP8 scales and NVFP4 weights were already consumed.
if name.endswith(".weight_scale"):
weight_name = name.removesuffix("_scale")
if weight_name not in self.model_tensors:
tensors_to_remove.append(name)
continue
w = self.model_tensors[weight_name]
s = self.model_tensors[name]
weight = w()
if weight.dtype not in (torch.float8_e4m3fn, torch.float8_e5m2):
raise ValueError(f"Expected remaining weight {weight_name!r} to be FP8, got {weight.dtype}")
scale = s()
block_size = None
is_channel_scale = scale.ndim >= 1 and scale.ndim <= weight.ndim and scale.shape[0] == weight.shape[0] and all(size == 1 for size in scale.shape[1:])
if scale.numel() != 1 and not is_channel_scale:
block_configs = [c for c in fp8_configs if c.get("strategy") == "block"]
if not block_configs or any(not c.get("block_structure") or c.get("group_size") is not None for c in block_configs):
raise ValueError(f"FP8 weight {weight_name!r} with scale shape {list(scale.shape)} requires a block layout")
block_sizes = {tuple(c["block_structure"]) for c in block_configs}
if len(block_sizes) != 1:
raise ValueError(f"FP8 weight {weight_name!r} has ambiguous block layouts")
block_size = next(iter(block_sizes))
self.model_tensors[weight_name] = lambda w=w, s=s, bs=block_size: dequant_simple(w(), s(), bs)
tensors_to_remove.append(name)
if self._fp8_as_q8:
self._fp8_dequantized.add(weight_name)
if name.endswith((".input_scale", ".activation_scale", "_activation_scale", ".k_scale", ".v_scale")):
tensors_to_remove.append(name)
elif len(groups) > 1:
raise NotImplementedError("Can't handle multiple config groups for compressed-tensors yet")
elif quant_format == "float-quantized" or quant_format == "int-quantized" or quant_format == "naive-quantized":