[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:
Jacob0226
2026-08-24 00:53:52 -07:00
committed by GitHub
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