[AMD] Fix weight load shape mismatch for amd dsr1 0528 mxfp4 (#19425)
This commit is contained in:
@@ -52,6 +52,7 @@ class QuarkConfig(QuantizationConfig):
|
|||||||
self.kv_cache_group = kv_cache_group
|
self.kv_cache_group = kv_cache_group
|
||||||
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.packed_modules_mapping = self.quant_config["packed_modules_mapping"]
|
self.packed_modules_mapping = self.quant_config["packed_modules_mapping"]
|
||||||
|
|
||||||
@@ -69,13 +70,17 @@ class QuarkConfig(QuantizationConfig):
|
|||||||
def get_name(self) -> str:
|
def get_name(self) -> str:
|
||||||
return "quark"
|
return "quark"
|
||||||
|
|
||||||
|
def apply_weight_name_mapper(self, hf_to_sglang_mapper):
|
||||||
|
self.exclude_layers = hf_to_sglang_mapper.apply_list(self.exclude_layers)
|
||||||
|
|
||||||
def get_quant_method(
|
def get_quant_method(
|
||||||
self, layer: torch.nn.Module, prefix: str
|
self, layer: torch.nn.Module, prefix: str
|
||||||
) -> Optional["QuantizeMethodBase"]:
|
) -> Optional["QuantizeMethodBase"]:
|
||||||
# Check if the layer is skipped for quantization.
|
# Check if the layer is skipped for quantization.
|
||||||
exclude_layers = cast(list[str], self.quant_config.get("exclude"))
|
|
||||||
if should_ignore_layer(
|
if should_ignore_layer(
|
||||||
prefix, ignore=exclude_layers, fused_mapping=self.packed_modules_mapping
|
prefix,
|
||||||
|
ignore=self.exclude_layers,
|
||||||
|
fused_mapping=self.packed_modules_mapping,
|
||||||
):
|
):
|
||||||
if isinstance(layer, LinearBase):
|
if isinstance(layer, LinearBase):
|
||||||
return UnquantizedLinearMethod()
|
return UnquantizedLinearMethod()
|
||||||
|
|||||||
@@ -50,6 +50,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
|||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.models.deepseek_common.utils import enable_nextn_moe_bf16_cast_to_fp8
|
from sglang.srt.models.deepseek_common.utils import enable_nextn_moe_bf16_cast_to_fp8
|
||||||
from sglang.srt.models.deepseek_v2 import DeepseekV2DecoderLayer, DeepseekV3ForCausalLM
|
from sglang.srt.models.deepseek_v2 import DeepseekV2DecoderLayer, DeepseekV3ForCausalLM
|
||||||
|
from sglang.srt.models.utils import WeightsMapper
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu
|
from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu
|
||||||
|
|
||||||
@@ -61,6 +62,7 @@ _is_npu = is_npu()
|
|||||||
|
|
||||||
|
|
||||||
class DeepseekModelNextN(nn.Module):
|
class DeepseekModelNextN(nn.Module):
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
config: PretrainedConfig,
|
config: PretrainedConfig,
|
||||||
@@ -192,6 +194,15 @@ class DeepseekModelNextN(nn.Module):
|
|||||||
|
|
||||||
class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM):
|
class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM):
|
||||||
|
|
||||||
|
# Support amd/DeepSeek-R1-0528-MXFP4 renaming: model.layers.61*.
|
||||||
|
# Ref: HF config.json for amd/DeepSeek-R1-0528-MXFP4
|
||||||
|
# https://huggingface.co/amd/DeepSeek-R1-0528-MXFP4/blob/main/config.json
|
||||||
|
hf_to_sglang_mapper = WeightsMapper(
|
||||||
|
orig_to_new_substr={
|
||||||
|
"model.layers.61": "model.decoder",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
config: PretrainedConfig,
|
config: PretrainedConfig,
|
||||||
|
|||||||
Reference in New Issue
Block a user