[AMD] Enable MTP for GLM-5-mxfp4 model (#23219)
This commit is contained in:
@@ -38,6 +38,7 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
is_dp_attention_enabled,
|
is_dp_attention_enabled,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.layernorm import RMSNorm
|
from sglang.srt.layers.layernorm import RMSNorm
|
||||||
|
from sglang.srt.layers.linear import ReplicatedLinear
|
||||||
from sglang.srt.layers.logits_processor import LogitsProcessor
|
from sglang.srt.layers.logits_processor import LogitsProcessor
|
||||||
from sglang.srt.layers.quantization import Fp8Config
|
from sglang.srt.layers.quantization import Fp8Config
|
||||||
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||||
@@ -101,7 +102,18 @@ class DeepseekModelNextN(nn.Module):
|
|||||||
self.enorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
self.enorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||||
self.hnorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
self.hnorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||||
|
|
||||||
self.eh_proj = nn.Linear(2 * config.hidden_size, config.hidden_size, bias=False)
|
if quant_config is not None and quant_config.get_name() == "quark":
|
||||||
|
self.eh_proj = ReplicatedLinear(
|
||||||
|
2 * config.hidden_size,
|
||||||
|
config.hidden_size,
|
||||||
|
bias=False,
|
||||||
|
quant_config=quant_config,
|
||||||
|
prefix=add_prefix("eh_proj", prefix),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self.eh_proj = nn.Linear(
|
||||||
|
2 * config.hidden_size, config.hidden_size, bias=False
|
||||||
|
)
|
||||||
|
|
||||||
self.rot_weight = None
|
self.rot_weight = None
|
||||||
if _is_npu:
|
if _is_npu:
|
||||||
@@ -166,21 +178,23 @@ class DeepseekModelNextN(nn.Module):
|
|||||||
hidden_states = input_embeds
|
hidden_states = input_embeds
|
||||||
|
|
||||||
if hidden_states.shape[0] > 0:
|
if hidden_states.shape[0] > 0:
|
||||||
hidden_states = self.eh_proj(
|
eh_input = torch.cat(
|
||||||
torch.cat(
|
(
|
||||||
(
|
self.enorm(hidden_states),
|
||||||
self.enorm(hidden_states),
|
self.hnorm(
|
||||||
self.hnorm(
|
forward_batch.spec_info.hidden_states
|
||||||
forward_batch.spec_info.hidden_states
|
if self.rot_weight is None
|
||||||
if self.rot_weight is None
|
else torch.matmul(
|
||||||
else torch.matmul(
|
forward_batch.spec_info.hidden_states, self.rot_weight
|
||||||
forward_batch.spec_info.hidden_states, self.rot_weight
|
)
|
||||||
)
|
|
||||||
),
|
|
||||||
),
|
),
|
||||||
dim=-1,
|
),
|
||||||
)
|
dim=-1,
|
||||||
)
|
)
|
||||||
|
if isinstance(self.eh_proj, ReplicatedLinear):
|
||||||
|
hidden_states, _ = self.eh_proj(eh_input)
|
||||||
|
else:
|
||||||
|
hidden_states = self.eh_proj(eh_input)
|
||||||
|
|
||||||
if nsa_use_prefill_cp(forward_batch, self.nsa_enable_prefill_cp):
|
if nsa_use_prefill_cp(forward_batch, self.nsa_enable_prefill_cp):
|
||||||
hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states)
|
hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states)
|
||||||
@@ -249,8 +263,20 @@ class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM):
|
|||||||
self.cp_rank = None
|
self.cp_rank = None
|
||||||
self.cp_size = None
|
self.cp_size = None
|
||||||
|
|
||||||
|
nextn_quant_config = quant_config
|
||||||
|
# For quark, if the MTP layer is listed in exclude_layers, set quant_config to None.
|
||||||
|
if nextn_quant_config is not None and nextn_quant_config.get_name() == "quark":
|
||||||
|
from sglang.srt.layers.quantization.quark.utils import (
|
||||||
|
should_ignore_layer,
|
||||||
|
)
|
||||||
|
|
||||||
|
ckpt_prefix = f"model.layers.{config.num_hidden_layers}"
|
||||||
|
mapped_prefix = self.hf_to_sglang_mapper._map_name(ckpt_prefix)
|
||||||
|
if should_ignore_layer(mapped_prefix, nextn_quant_config.exclude_layers):
|
||||||
|
nextn_quant_config = None
|
||||||
|
|
||||||
self.model = DeepseekModelNextN(
|
self.model = DeepseekModelNextN(
|
||||||
config, quant_config, prefix=add_prefix("model", prefix)
|
config, nextn_quant_config, prefix=add_prefix("model", prefix)
|
||||||
)
|
)
|
||||||
self.lm_head = ParallelLMHead(
|
self.lm_head = ParallelLMHead(
|
||||||
config.vocab_size,
|
config.vocab_size,
|
||||||
|
|||||||
Reference in New Issue
Block a user