mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-11 07:20:33 +02:00
Allow remaining FP8 tensors to be dequantized during conversion
This commit is contained in:
+35
-1
@@ -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":
|
||||
|
||||
Reference in New Issue
Block a user