Avoid fusing FP8 for --fuse-gate-up-exps as we may have different ws

This commit is contained in:
Oliver Simons
2026-10-01 17:59:18 +02:00
parent b2860f80b1
commit e703584f32
+21 -1
View File
@@ -154,6 +154,7 @@ class ModelBase:
self.fuse_gate_up_exps = fuse_gate_up_exps
self._gate_exp_buffer: dict[int, Tensor] = {}
self._up_exp_buffer: dict[int, Tensor] = {}
self._unfusable_gate_up_layers: set[int] = set()
self.fuse_qkv = fuse_qkv
self._q_buffer: dict[int, Tensor] = {}
self._k_buffer: dict[int, Tensor] = {}
@@ -759,6 +760,23 @@ class ModelBase:
raise ValueError(f"Can not map tensor {name!r}")
return new_name
def prepare_gate_up_fusion(self) -> None:
self._unfusable_gate_up_layers.clear()
if not self.fuse_gate_up_exps or self._fp8_as_q8:
return
for name, gen in self.model_tensors.items():
merged_name = re.sub(r"(\.experts)\.\d+\.", r"\1.", name)
mapped = self.tensor_map.get_type_and_name(merged_name, try_suffixes=(".weight",))
if mapped is None:
continue
tensor_type, new_name = mapped
if tensor_type not in (gguf.MODEL_TENSOR.FFN_GATE_EXP, gguf.MODEL_TENSOR.FFN_UP_EXP) or not new_name.endswith(".weight"):
continue
if gen().dtype == torch.float8_e4m3fn:
bid = next(int(part) for part in new_name.split(".") if part.isdecimal())
self._unfusable_gate_up_layers.add(bid)
def prepare_qkv_fusion(self) -> None:
self._fusable_qkv_weight_layers.clear()
self._fusable_qkv_bias_layers.clear()
@@ -822,7 +840,8 @@ class ModelBase:
new_name = self.map_tensor_name(name)
# Handle gate/up expert tensor fusion if enabled
if self.fuse_gate_up_exps and bid is not None:
# Preserved FP8 gate/up scales must stay with their separate projections.
if self.fuse_gate_up_exps and bid is not None and bid not in self._unfusable_gate_up_layers:
if self.match_model_tensor_name(new_name, gguf.MODEL_TENSOR.FFN_GATE_EXP, bid):
self._gate_exp_buffer[bid] = data_torch
elif self.match_model_tensor_name(new_name, gguf.MODEL_TENSOR.FFN_UP_EXP, bid):
@@ -1193,6 +1212,7 @@ class ModelBase:
self._prepare_fp8_e4m3_tensors()
self.dequant_model()
self.prepare_gate_up_fusion()
self.prepare_qkv_fusion()
# Handle empty tensor_map for models with block_count=0 (like MobileNetV5)