[Fix] Fix Qwen3.5 MoE model loading and Mamba cache sharding in PP mode (#21448)
Co-authored-by: zhangxiaolei123456 <zhangxiaolei.666@bytedance.com>
This commit is contained in:
co-authored by
zhangxiaolei123456
parent
c06ca1526c
commit
9b4dd27478
@@ -175,11 +175,13 @@ class HybridMambaDecodeReqToTokenPool(HybridReqToTokenPool):
|
|||||||
device: str,
|
device: str,
|
||||||
enable_memory_saver: bool,
|
enable_memory_saver: bool,
|
||||||
cache_params: "Mamba2CacheParams",
|
cache_params: "Mamba2CacheParams",
|
||||||
|
mamba_layer_ids: List[int],
|
||||||
speculative_num_draft_tokens: int,
|
speculative_num_draft_tokens: int,
|
||||||
enable_mamba_extra_buffer: bool,
|
enable_mamba_extra_buffer: bool,
|
||||||
pre_alloc_size: int,
|
pre_alloc_size: int,
|
||||||
enable_overlap_schedule: bool,
|
enable_overlap_schedule: bool,
|
||||||
mamba_size: int = None,
|
mamba_size: int = None,
|
||||||
|
start_layer: int = None,
|
||||||
):
|
):
|
||||||
DecodeReqToTokenPool.__init__(
|
DecodeReqToTokenPool.__init__(
|
||||||
self,
|
self,
|
||||||
@@ -196,13 +198,13 @@ class HybridMambaDecodeReqToTokenPool(HybridReqToTokenPool):
|
|||||||
effective_mamba_size = (
|
effective_mamba_size = (
|
||||||
mamba_size if mamba_size is not None else size
|
mamba_size if mamba_size is not None else size
|
||||||
) + pre_alloc_size
|
) + pre_alloc_size
|
||||||
# TODO: Support PP
|
self.start_layer = start_layer if start_layer is not None else 0
|
||||||
self.start_layer = 0
|
|
||||||
self.layer_transfer_counter = None
|
self.layer_transfer_counter = None
|
||||||
self._init_mamba_pool(
|
self._init_mamba_pool(
|
||||||
size=effective_mamba_size,
|
size=effective_mamba_size,
|
||||||
mamba_spec_state_size=size + pre_alloc_size,
|
mamba_spec_state_size=size + pre_alloc_size,
|
||||||
cache_params=cache_params,
|
cache_params=cache_params,
|
||||||
|
mamba_layer_ids=mamba_layer_ids,
|
||||||
device=device,
|
device=device,
|
||||||
enable_mamba_extra_buffer=self.enable_mamba_extra_buffer,
|
enable_mamba_extra_buffer=self.enable_mamba_extra_buffer,
|
||||||
speculative_num_draft_tokens=speculative_num_draft_tokens,
|
speculative_num_draft_tokens=speculative_num_draft_tokens,
|
||||||
|
|||||||
@@ -220,6 +220,7 @@ class MambaPool:
|
|||||||
size: int,
|
size: int,
|
||||||
spec_state_size: int,
|
spec_state_size: int,
|
||||||
cache_params: BaseLinearStateParams,
|
cache_params: BaseLinearStateParams,
|
||||||
|
mamba_layer_ids: List[int],
|
||||||
device: str,
|
device: str,
|
||||||
enable_memory_saver: bool = False,
|
enable_memory_saver: bool = False,
|
||||||
speculative_num_draft_tokens: Optional[int] = None,
|
speculative_num_draft_tokens: Optional[int] = None,
|
||||||
@@ -231,7 +232,7 @@ class MambaPool:
|
|||||||
self.memory_saver_adapter = TorchMemorySaverAdapter.create(
|
self.memory_saver_adapter = TorchMemorySaverAdapter.create(
|
||||||
enable=enable_memory_saver
|
enable=enable_memory_saver
|
||||||
)
|
)
|
||||||
num_mamba_layers = len(cache_params.layers)
|
num_mamba_layers = len(mamba_layer_ids)
|
||||||
|
|
||||||
self.size = size
|
self.size = size
|
||||||
self.device = device
|
self.device = device
|
||||||
@@ -454,9 +455,11 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
device: str,
|
device: str,
|
||||||
enable_memory_saver: bool,
|
enable_memory_saver: bool,
|
||||||
cache_params: BaseLinearStateParams,
|
cache_params: BaseLinearStateParams,
|
||||||
|
mamba_layer_ids: List[int],
|
||||||
enable_mamba_extra_buffer: bool,
|
enable_mamba_extra_buffer: bool,
|
||||||
speculative_num_draft_tokens: int = None,
|
speculative_num_draft_tokens: int = None,
|
||||||
enable_overlap_schedule: bool = True,
|
enable_overlap_schedule: bool = True,
|
||||||
|
start_layer: Optional[int] = None,
|
||||||
):
|
):
|
||||||
super().__init__(
|
super().__init__(
|
||||||
size=size,
|
size=size,
|
||||||
@@ -468,13 +471,13 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
self.mamba_ping_pong_track_buffer_size = 2 if enable_overlap_schedule else 1
|
self.mamba_ping_pong_track_buffer_size = 2 if enable_overlap_schedule else 1
|
||||||
self.enable_mamba_extra_buffer = enable_mamba_extra_buffer
|
self.enable_mamba_extra_buffer = enable_mamba_extra_buffer
|
||||||
self.enable_memory_saver = enable_memory_saver
|
self.enable_memory_saver = enable_memory_saver
|
||||||
# TODO: Support PP
|
self.start_layer = start_layer if start_layer is not None else 0
|
||||||
self.start_layer = 0
|
|
||||||
self.layer_transfer_counter = None
|
self.layer_transfer_counter = None
|
||||||
self._init_mamba_pool(
|
self._init_mamba_pool(
|
||||||
size=mamba_size,
|
size=mamba_size,
|
||||||
mamba_spec_state_size=mamba_spec_state_size,
|
mamba_spec_state_size=mamba_spec_state_size,
|
||||||
cache_params=cache_params,
|
cache_params=cache_params,
|
||||||
|
mamba_layer_ids=mamba_layer_ids,
|
||||||
device=device,
|
device=device,
|
||||||
enable_mamba_extra_buffer=enable_mamba_extra_buffer,
|
enable_mamba_extra_buffer=enable_mamba_extra_buffer,
|
||||||
speculative_num_draft_tokens=speculative_num_draft_tokens,
|
speculative_num_draft_tokens=speculative_num_draft_tokens,
|
||||||
@@ -485,6 +488,7 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
size: int,
|
size: int,
|
||||||
mamba_spec_state_size: int,
|
mamba_spec_state_size: int,
|
||||||
cache_params: BaseLinearStateParams,
|
cache_params: BaseLinearStateParams,
|
||||||
|
mamba_layer_ids: List[int],
|
||||||
device: str,
|
device: str,
|
||||||
enable_mamba_extra_buffer: bool,
|
enable_mamba_extra_buffer: bool,
|
||||||
speculative_num_draft_tokens: int = None,
|
speculative_num_draft_tokens: int = None,
|
||||||
@@ -493,11 +497,12 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
size=size,
|
size=size,
|
||||||
spec_state_size=mamba_spec_state_size,
|
spec_state_size=mamba_spec_state_size,
|
||||||
cache_params=cache_params,
|
cache_params=cache_params,
|
||||||
|
mamba_layer_ids=mamba_layer_ids,
|
||||||
device=device,
|
device=device,
|
||||||
enable_memory_saver=self.enable_memory_saver,
|
enable_memory_saver=self.enable_memory_saver,
|
||||||
speculative_num_draft_tokens=speculative_num_draft_tokens,
|
speculative_num_draft_tokens=speculative_num_draft_tokens,
|
||||||
)
|
)
|
||||||
self.mamba_map = {layer_id: i for i, layer_id in enumerate(cache_params.layers)}
|
self.mamba_map = {layer_id: i for i, layer_id in enumerate(mamba_layer_ids)}
|
||||||
|
|
||||||
self.device = device
|
self.device = device
|
||||||
self.req_index_to_mamba_index_mapping: torch.Tensor = torch.zeros(
|
self.req_index_to_mamba_index_mapping: torch.Tensor = torch.zeros(
|
||||||
@@ -1235,13 +1240,14 @@ class HybridLinearKVPool(KVCache):
|
|||||||
use_mla: bool = False,
|
use_mla: bool = False,
|
||||||
kv_lora_rank: int = None,
|
kv_lora_rank: int = None,
|
||||||
qk_rope_head_dim: int = None,
|
qk_rope_head_dim: int = None,
|
||||||
|
start_layer: Optional[int] = None,
|
||||||
):
|
):
|
||||||
self.size = size
|
self.size = size
|
||||||
self.dtype = dtype
|
self.dtype = dtype
|
||||||
self.device = device
|
self.device = device
|
||||||
self.full_layer_nums = len(full_attention_layer_ids)
|
self.full_layer_nums = len(full_attention_layer_ids)
|
||||||
self.page_size = page_size
|
self.page_size = page_size
|
||||||
self.start_layer = 0 # TODO: Support PP
|
self.start_layer = start_layer if start_layer is not None else 0
|
||||||
self.layer_transfer_counter = None
|
self.layer_transfer_counter = None
|
||||||
self.head_num = head_num
|
self.head_num = head_num
|
||||||
self.head_dim = head_dim
|
self.head_dim = head_dim
|
||||||
|
|||||||
@@ -401,11 +401,19 @@ class ModelRunnerKVCacheMixin:
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||||
cache_params=config.mamba2_cache_params,
|
cache_params=config.mamba2_cache_params,
|
||||||
|
mamba_layer_ids=(
|
||||||
|
[
|
||||||
|
i
|
||||||
|
for i in config.mamba2_cache_params.layers
|
||||||
|
if self.start_layer <= i < self.end_layer
|
||||||
|
]
|
||||||
|
),
|
||||||
speculative_num_draft_tokens=self.server_args.speculative_num_draft_tokens,
|
speculative_num_draft_tokens=self.server_args.speculative_num_draft_tokens,
|
||||||
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
|
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
|
||||||
pre_alloc_size=pre_alloc_size,
|
pre_alloc_size=pre_alloc_size,
|
||||||
enable_overlap_schedule=not self.server_args.disable_overlap_schedule,
|
enable_overlap_schedule=not self.server_args.disable_overlap_schedule,
|
||||||
mamba_size=self.server_args.max_mamba_cache_size,
|
mamba_size=self.server_args.max_mamba_cache_size,
|
||||||
|
start_layer=self.start_layer,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.req_to_token_pool = DecodeReqToTokenPool(
|
self.req_to_token_pool = DecodeReqToTokenPool(
|
||||||
@@ -426,9 +434,17 @@ class ModelRunnerKVCacheMixin:
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||||
cache_params=config.mamba2_cache_params,
|
cache_params=config.mamba2_cache_params,
|
||||||
|
mamba_layer_ids=(
|
||||||
|
[
|
||||||
|
i
|
||||||
|
for i in config.mamba2_cache_params.layers
|
||||||
|
if self.start_layer <= i < self.end_layer
|
||||||
|
]
|
||||||
|
),
|
||||||
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
|
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
|
||||||
speculative_num_draft_tokens=self.server_args.speculative_num_draft_tokens,
|
speculative_num_draft_tokens=self.server_args.speculative_num_draft_tokens,
|
||||||
enable_overlap_schedule=not self.server_args.disable_overlap_schedule,
|
enable_overlap_schedule=not self.server_args.disable_overlap_schedule,
|
||||||
|
start_layer=self.start_layer,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.req_to_token_pool = ReqToTokenPool(
|
self.req_to_token_pool = ReqToTokenPool(
|
||||||
@@ -643,6 +659,7 @@ class ModelRunnerKVCacheMixin:
|
|||||||
mamba_pool=self.req_to_token_pool.mamba_pool,
|
mamba_pool=self.req_to_token_pool.mamba_pool,
|
||||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||||
use_mla=self.use_mla_backend,
|
use_mla=self.use_mla_backend,
|
||||||
|
start_layer=self.start_layer,
|
||||||
**extra_args,
|
**extra_args,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -67,7 +67,7 @@ from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
|||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
from sglang.srt.layers.radix_linear_attention import RadixLinearAttention
|
from sglang.srt.layers.radix_linear_attention import RadixLinearAttention
|
||||||
from sglang.srt.layers.rotary_embedding import get_rope
|
from sglang.srt.layers.rotary_embedding import get_rope
|
||||||
from sglang.srt.layers.utils import PPMissingLayer
|
from sglang.srt.layers.utils import PPMissingLayer, get_layer_id
|
||||||
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||||
from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
|
from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||||
@@ -1038,6 +1038,13 @@ class Qwen3_5ForCausalLM(nn.Module):
|
|||||||
name = name.replace(r"model.language_model.", r"model.")
|
name = name.replace(r"model.language_model.", r"model.")
|
||||||
if ".self_attn." in name:
|
if ".self_attn." in name:
|
||||||
name = name.replace(".self_attn", "")
|
name = name.replace(".self_attn", "")
|
||||||
|
layer_id = get_layer_id(name)
|
||||||
|
if (
|
||||||
|
layer_id is not None
|
||||||
|
and hasattr(self, "start_layer")
|
||||||
|
and (layer_id < self.start_layer or layer_id >= self.end_layer)
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
|
||||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||||
if weight_name not in name:
|
if weight_name not in name:
|
||||||
@@ -1175,6 +1182,14 @@ class Qwen3_5MoeForCausalLM(Qwen3_5ForCausalLM):
|
|||||||
if ".self_attn." in name:
|
if ".self_attn." in name:
|
||||||
name = name.replace(".self_attn", "")
|
name = name.replace(".self_attn", "")
|
||||||
|
|
||||||
|
layer_id = get_layer_id(name)
|
||||||
|
if (
|
||||||
|
layer_id is not None
|
||||||
|
and hasattr(self, "start_layer")
|
||||||
|
and (layer_id < self.start_layer or layer_id >= self.end_layer)
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
|
||||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||||
if "experts.gate_up_proj" in name or "experts.down_proj" in name:
|
if "experts.gate_up_proj" in name or "experts.down_proj" in name:
|
||||||
is_fused_expert = True
|
is_fused_expert = True
|
||||||
@@ -1355,6 +1370,13 @@ class Qwen3_5ForConditionalGeneration(Qwen3VLForConditionalGeneration):
|
|||||||
name = name.replace(r"model.language_model.", r"model.")
|
name = name.replace(r"model.language_model.", r"model.")
|
||||||
if ".self_attn." in name:
|
if ".self_attn." in name:
|
||||||
name = name.replace(".self_attn", "")
|
name = name.replace(".self_attn", "")
|
||||||
|
layer_id = get_layer_id(name)
|
||||||
|
if (
|
||||||
|
layer_id is not None
|
||||||
|
and hasattr(self, "start_layer")
|
||||||
|
and (layer_id < self.start_layer or layer_id >= self.end_layer)
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
|
||||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||||
if weight_name not in name:
|
if weight_name not in name:
|
||||||
@@ -1510,6 +1532,14 @@ class Qwen3_5MoeForConditionalGeneration(Qwen3VLForConditionalGeneration):
|
|||||||
if ".self_attn." in name:
|
if ".self_attn." in name:
|
||||||
name = name.replace(".self_attn", "")
|
name = name.replace(".self_attn", "")
|
||||||
|
|
||||||
|
layer_id = get_layer_id(name)
|
||||||
|
if (
|
||||||
|
layer_id is not None
|
||||||
|
and hasattr(self, "start_layer")
|
||||||
|
and (layer_id < self.start_layer or layer_id >= self.end_layer)
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
|
||||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||||
if name.endswith("experts.gate_up_proj") or name.endswith(
|
if name.endswith("experts.gate_up_proj") or name.endswith(
|
||||||
"experts.down_proj"
|
"experts.down_proj"
|
||||||
|
|||||||
@@ -1141,6 +1141,19 @@ class Qwen3VLForConditionalGeneration(nn.Module):
|
|||||||
input_deepstack_embeds = embedding[:, separate_index:]
|
input_deepstack_embeds = embedding[:, separate_index:]
|
||||||
return input_embeds, input_deepstack_embeds
|
return input_embeds, input_deepstack_embeds
|
||||||
|
|
||||||
|
@property
|
||||||
|
def start_layer(self) -> int:
|
||||||
|
return getattr(getattr(self, "model", None), "start_layer", 0)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def end_layer(self) -> int:
|
||||||
|
model = getattr(self, "model", None)
|
||||||
|
end_layer = getattr(model, "end_layer", None)
|
||||||
|
if end_layer is not None:
|
||||||
|
return end_layer
|
||||||
|
cfg = getattr(model, "config", None)
|
||||||
|
return int(getattr(cfg, "num_hidden_layers", 0))
|
||||||
|
|
||||||
def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs):
|
def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs):
|
||||||
pattern = MultiModalityDataPaddingPatternMultimodalTokens()
|
pattern = MultiModalityDataPaddingPatternMultimodalTokens()
|
||||||
return pattern.pad_input_tokens(input_ids, mm_inputs)
|
return pattern.pad_input_tokens(input_ids, mm_inputs)
|
||||||
|
|||||||
@@ -99,6 +99,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
device=device,
|
device=device,
|
||||||
enable_memory_saver=False,
|
enable_memory_saver=False,
|
||||||
cache_params=mamba2_cache_params,
|
cache_params=mamba2_cache_params,
|
||||||
|
mamba_layer_ids=mamba_layers,
|
||||||
enable_mamba_extra_buffer=False,
|
enable_mamba_extra_buffer=False,
|
||||||
speculative_num_draft_tokens=3,
|
speculative_num_draft_tokens=3,
|
||||||
)
|
)
|
||||||
@@ -340,6 +341,7 @@ class TestMamba(unittest.TestCase):
|
|||||||
device=device,
|
device=device,
|
||||||
enable_memory_saver=False,
|
enable_memory_saver=False,
|
||||||
cache_params=mamba2_cache_params,
|
cache_params=mamba2_cache_params,
|
||||||
|
mamba_layer_ids=mamba_layers,
|
||||||
enable_mamba_extra_buffer=False,
|
enable_mamba_extra_buffer=False,
|
||||||
speculative_num_draft_tokens=3,
|
speculative_num_draft_tokens=3,
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user