[AMD] Quark shared-experts gate: recognise a trailing MTP layer (#36124)
Co-authored-by: Cursor <cursoragent@cursor.com> Co-authored-by: Thomas Wang <thomawan@amd.com>
This commit is contained in:
co-authored by
Cursor
Thomas Wang
parent
f464e77d17
commit
7bbd0ddeb5
@@ -322,6 +322,14 @@ class QuarkConfig(QuantizationConfig):
|
||||
self.kv_cache_config = kv_cache_config
|
||||
self.pack_method = pack_method
|
||||
self.exclude_layers = cast(list[str], self.quant_config.get("exclude", []))
|
||||
# Both are consumed by _is_draft_layer(), which has to tell an appended
|
||||
# MTP/NextN draft layer from a target-model one. "No draft stack" is
|
||||
# spelled None as often as it is spelled absent -- ModelConfig defaults
|
||||
# the same field to None -- so coerce rather than let range() raise.
|
||||
self.num_hidden_layers = getattr(hf_config, "num_hidden_layers", None)
|
||||
self.num_nextn_predict_layers = int(
|
||||
getattr(hf_config, "num_nextn_predict_layers", 0) or 0
|
||||
)
|
||||
self.is_prequantized = is_prequantized
|
||||
self.dequantization_config = dequantization_config
|
||||
# Load-as-is FP8 config for excluded layers of a mixed-precision source
|
||||
@@ -555,6 +563,10 @@ class QuarkConfig(QuantizationConfig):
|
||||
|
||||
return cls(
|
||||
quant_config=config,
|
||||
# The requantization branch above takes hf_config off the same dict;
|
||||
# the prequantized path needs it too, so _is_draft_layer() has the
|
||||
# layer counts it compares against.
|
||||
hf_config=config.get("hf_config"),
|
||||
kv_cache_group=kv_cache_group,
|
||||
kv_cache_config=kv_cache_config,
|
||||
pack_method=pack_method,
|
||||
@@ -895,12 +907,34 @@ class QuarkConfig(QuantizationConfig):
|
||||
def get_scaled_act_names(self) -> List[str]:
|
||||
return []
|
||||
|
||||
def _is_draft_layer(self, layer: str) -> bool:
|
||||
"""Whether an excluded layer belongs to the MTP/NextN draft stack.
|
||||
|
||||
A draft layer is excluded from quantization in most checkpoints, and
|
||||
says nothing about how the target model stores its own shared experts,
|
||||
so it must not veto shared-expert fusion for the target model's layers.
|
||||
|
||||
Checkpoints spell it either as "mtp.*" or as extra entries appended to
|
||||
the main decoder, model.layers.[num_hidden_layers ..
|
||||
+ num_nextn_predict_layers) -- the same range
|
||||
get_spec_layer_idx_from_weight_name() walks in the MTP models.
|
||||
"""
|
||||
if layer.startswith("mtp."):
|
||||
return True
|
||||
if self.num_hidden_layers is None:
|
||||
return False
|
||||
base = self.num_hidden_layers
|
||||
return any(
|
||||
layer.startswith(f"model.layers.{base + i}.")
|
||||
for i in range(self.num_nextn_predict_layers)
|
||||
)
|
||||
|
||||
def can_fuse_shared_expert(self) -> bool:
|
||||
# Shared-expert body excluded from quant; the gate must not veto fusion.
|
||||
if any(
|
||||
"shared_expert" in layer
|
||||
and "shared_expert_gate" not in layer
|
||||
and not layer.startswith("mtp.")
|
||||
and not self._is_draft_layer(layer)
|
||||
for layer in self.exclude_layers
|
||||
):
|
||||
return False
|
||||
|
||||
Reference in New Issue
Block a user