[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.kv_cache_config = kv_cache_config
|
||||||
self.pack_method = pack_method
|
self.pack_method = pack_method
|
||||||
self.exclude_layers = cast(list[str], self.quant_config.get("exclude", []))
|
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.is_prequantized = is_prequantized
|
||||||
self.dequantization_config = dequantization_config
|
self.dequantization_config = dequantization_config
|
||||||
# Load-as-is FP8 config for excluded layers of a mixed-precision source
|
# Load-as-is FP8 config for excluded layers of a mixed-precision source
|
||||||
@@ -555,6 +563,10 @@ class QuarkConfig(QuantizationConfig):
|
|||||||
|
|
||||||
return cls(
|
return cls(
|
||||||
quant_config=config,
|
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_group=kv_cache_group,
|
||||||
kv_cache_config=kv_cache_config,
|
kv_cache_config=kv_cache_config,
|
||||||
pack_method=pack_method,
|
pack_method=pack_method,
|
||||||
@@ -895,12 +907,34 @@ class QuarkConfig(QuantizationConfig):
|
|||||||
def get_scaled_act_names(self) -> List[str]:
|
def get_scaled_act_names(self) -> List[str]:
|
||||||
return []
|
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:
|
def can_fuse_shared_expert(self) -> bool:
|
||||||
# Shared-expert body excluded from quant; the gate must not veto fusion.
|
# Shared-expert body excluded from quant; the gate must not veto fusion.
|
||||||
if any(
|
if any(
|
||||||
"shared_expert" in layer
|
"shared_expert" in layer
|
||||||
and "shared_expert_gate" not 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
|
for layer in self.exclude_layers
|
||||||
):
|
):
|
||||||
return False
|
return False
|
||||||
|
|||||||
Reference in New Issue
Block a user