diff --git a/conversion/base.py b/conversion/base.py index bfcef9d48b..18c3ea08ca 100644 --- a/conversion/base.py +++ b/conversion/base.py @@ -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)