[CPU] Padding for dim divisibility in TP3/6 cases (#20072)

Co-authored-by: Ma Mingfei <mingfei.ma@intel.com>
This commit is contained in:
Zaili Wang
2026-07-02 13:14:39 +08:00
committed by GitHub
co-authored by Ma Mingfei
parent 80ac11eda3
commit aff44d748d
6 changed files with 119 additions and 18 deletions
+34 -3
View File
@@ -155,15 +155,23 @@ def update_intermediate_size(model_config, attr_name, intermediate_padding_size)
if attr_value % intermediate_padding_size != 0: if attr_value % intermediate_padding_size != 0:
from sglang.srt.layers.vocab_parallel_embedding import pad_vocab_size from sglang.srt.layers.vocab_parallel_embedding import pad_vocab_size
origin_value = attr_value
origin_name = "original_" + attr_name
attr_value = pad_vocab_size(attr_value, intermediate_padding_size) attr_value = pad_vocab_size(attr_value, intermediate_padding_size)
if hasattr(model_config, "hf_config"): if hasattr(model_config, "hf_config"):
update_config(model_config.hf_config, attr_name, attr_value) update_config(model_config.hf_config, attr_name, attr_value)
update_config(model_config.hf_config, origin_name, origin_value)
if hasattr(model_config, "hf_text_config"): if hasattr(model_config, "hf_text_config"):
update_config(model_config.hf_text_config, attr_name, attr_value) update_config(model_config.hf_text_config, attr_name, attr_value)
update_config(model_config.hf_text_config, origin_name, origin_value)
if hasattr(model_config.hf_config, "text_config"): if hasattr(model_config.hf_config, "text_config"):
update_config(model_config.hf_config.text_config, attr_name, attr_value) update_config(model_config.hf_config.text_config, attr_name, attr_value)
update_config(
model_config.hf_config.text_config, origin_name, origin_value
)
else: else:
update_config(model_config, attr_name, attr_value) update_config(model_config, attr_name, attr_value)
update_config(model_config, origin_name, origin_value)
return model_config return model_config
@@ -261,10 +269,18 @@ def adjust_config_with_unaligned_cpu_tp(
"siglip_vision_model", "siglip_vision_model",
"num_attention_heads", "num_attention_heads",
], ],
[model_config.hf_config, "vision_config", "qwen2_5_vl", "num_heads"],
[model_config.hf_config, "vision_config", "qwen3_vl_moe", "num_heads"], [model_config.hf_config, "vision_config", "qwen3_vl_moe", "num_heads"],
[model_config.hf_config, "vision_config", "qwen3_vl", "num_heads"], [model_config.hf_config, "vision_config", "qwen3_vl", "num_heads"],
[model_config.hf_config, "vision_config", "qwen3_5_moe", "num_heads"], [model_config.hf_config, "vision_config", "qwen3_5_moe", "num_heads"],
[model_config.hf_config, "vision_config", "qwen3_5", "num_heads"], [model_config.hf_config, "vision_config", "qwen3_5", "num_heads"],
[model_config.hf_config, "vision_config", "mllama", "attention_heads"],
[
model_config.hf_config,
"vision_config",
"llama4_vision_model",
"num_attention_heads",
],
] ]
if hasattr(model_config.hf_config, "thinker_config"): if hasattr(model_config.hf_config, "thinker_config"):
multimodal_config.append( multimodal_config.append(
@@ -285,11 +301,12 @@ def adjust_config_with_unaligned_cpu_tp(
) )
for m_config, config_name, model_type, num_head_str in multimodal_config: for m_config, config_name, model_type, num_head_str in multimodal_config:
if ( if hasattr(m_config, config_name) and (
hasattr(m_config, config_name) m_config.model_type == model_type
and getattr(m_config, config_name).model_type == model_type or getattr(m_config, config_name).model_type == model_type
): ):
num_heads = getattr(getattr(m_config, config_name), num_head_str) num_heads = getattr(getattr(m_config, config_name), num_head_str)
update_config( update_config(
getattr(m_config, config_name), "original_" + num_head_str, num_heads getattr(m_config, config_name), "original_" + num_head_str, num_heads
) )
@@ -316,4 +333,18 @@ def adjust_config_with_unaligned_cpu_tp(
), ),
) )
# Pad projector_input_dim for Llama4 vision if needed
if model_type == "llama4_vision_model":
proj_inp_dim = getattr(m_config, config_name).projector_input_dim
if proj_inp_dim % tp_size != 0:
from sglang.srt.layers.vocab_parallel_embedding import (
pad_vocab_size,
)
update_config(
getattr(m_config, config_name),
"projector_input_dim",
pad_vocab_size(proj_inp_dim, tp_size),
)
return model_config return model_config
+6 -3
View File
@@ -933,6 +933,9 @@ class GptOssForCausalLM(nn.Module):
moe_ep_size = get_parallel().moe_ep_size moe_ep_size = get_parallel().moe_ep_size
intermediate_size = self.config.intermediate_size intermediate_size = self.config.intermediate_size
original_intermediate_size = getattr(
self.config, "original_intermediate_size", intermediate_size
)
assert ( assert (
intermediate_size % mxfp4_block == 0 intermediate_size % mxfp4_block == 0
), f"{intermediate_size=} must be divisible by {mxfp4_block=}" ), f"{intermediate_size=} must be divisible by {mxfp4_block=}"
@@ -951,7 +954,7 @@ class GptOssForCausalLM(nn.Module):
moe_tp_rank_start = moe_tp_rank * per_rank_intermediate_size moe_tp_rank_start = moe_tp_rank * per_rank_intermediate_size
moe_tp_rank_end = min( moe_tp_rank_end = min(
(moe_tp_rank + 1) * per_rank_intermediate_size, intermediate_size (moe_tp_rank + 1) * per_rank_intermediate_size, original_intermediate_size
) )
moe_ep_rank_start = moe_ep_rank * moe_num_local_experts moe_ep_rank_start = moe_ep_rank * moe_num_local_experts
@@ -968,7 +971,7 @@ class GptOssForCausalLM(nn.Module):
# flat weight from (E, 2 * N, block_size, entry_per_block) # flat weight from (E, 2 * N, block_size, entry_per_block)
# to (E, 2 * N, -1), shouldn't trigger copy for contiguous # to (E, 2 * N, -1), shouldn't trigger copy for contiguous
weight = weight.view( weight = weight.view(
moe_num_global_experts, 2 * intermediate_size, -1 moe_num_global_experts, 2 * original_intermediate_size, -1
).contiguous() ).contiguous()
narrow_weight = weight[ narrow_weight = weight[
@@ -994,7 +997,7 @@ class GptOssForCausalLM(nn.Module):
# same flatten here, but since 2 mx4 value are packed in 1 # same flatten here, but since 2 mx4 value are packed in 1
# uint8, divide by 2 # uint8, divide by 2
weight = weight.view( weight = weight.view(
moe_num_global_experts, -1, intermediate_size // 2 moe_num_global_experts, -1, original_intermediate_size // 2
).contiguous() ).contiguous()
narrow_weight = weight[ narrow_weight = weight[
moe_ep_rank_start:moe_ep_rank_end, moe_ep_rank_start:moe_ep_rank_end,
+24 -2
View File
@@ -198,9 +198,16 @@ class MllamaVisionEncoderLayer(nn.Module):
super().__init__() super().__init__()
self.hidden_size = config.hidden_size self.hidden_size = config.hidden_size
self.num_attention_heads = config.attention_heads self.num_attention_heads = (
config.original_attention_heads
if hasattr(config, "original_attention_heads")
else config.attention_heads
)
self.is_gated = is_gated self.is_gated = is_gated
self.intermediate_size = config.intermediate_size self.intermediate_size = config.intermediate_size
num_dummy_heads = 0
if hasattr(config, "original_attention_heads"):
num_dummy_heads = config.attention_heads - config.original_attention_heads
self.self_attn = VisionAttention( self.self_attn = VisionAttention(
self.hidden_size, self.hidden_size,
@@ -210,6 +217,7 @@ class MllamaVisionEncoderLayer(nn.Module):
quant_config=quant_config, quant_config=quant_config,
flatten_batch=False, flatten_batch=False,
prefix=add_prefix("self_attn", prefix), prefix=add_prefix("self_attn", prefix),
num_dummy_heads=num_dummy_heads,
) )
self.mlp = MllamaVisionMLP( self.mlp = MllamaVisionMLP(
config, quant_config, prefix=add_prefix("mlp", prefix) config, quant_config, prefix=add_prefix("mlp", prefix)
@@ -310,10 +318,17 @@ class MllamaVisionModel(nn.Module):
self.num_patches = (self.image_size // self.patch_size) ** 2 + 1 self.num_patches = (self.image_size // self.patch_size) ** 2 + 1
self.scale = config.hidden_size**-0.5 self.scale = config.hidden_size**-0.5
out_channels = (
config.hidden_size
// config.original_attention_heads
* config.attention_heads
if hasattr(config, "original_attention_heads")
else config.hidden_size
)
self.patch_embedding = ColumnParallelConv2dPatch( self.patch_embedding = ColumnParallelConv2dPatch(
in_channels=config.num_channels, in_channels=config.num_channels,
out_channels=self.hidden_size, out_channels=out_channels,
kernel_size=self.patch_size, kernel_size=self.patch_size,
stride=self.patch_size, stride=self.patch_size,
bias=False, bias=False,
@@ -382,6 +397,10 @@ class MllamaVisionModel(nn.Module):
# tile embeddings # tile embeddings
_, num_patches, dim = hidden_state.shape _, num_patches, dim = hidden_state.shape
# slice off the padded part
if dim > self.hidden_size:
hidden_state = hidden_state[:, :, : self.hidden_size]
dim = self.hidden_size
hidden_state = hidden_state.reshape( hidden_state = hidden_state.reshape(
batch_size * num_concurrent_media, num_tiles, -1, dim batch_size * num_concurrent_media, num_tiles, -1, dim
) )
@@ -501,6 +520,9 @@ class MllamaTextCrossAttention(nn.Module):
self.dropout = config.dropout self.dropout = config.dropout
self.hidden_size = config.hidden_size self.hidden_size = config.hidden_size
self.head_dim = config.hidden_size // self.num_heads self.head_dim = config.hidden_size // self.num_heads
# Use original head_dim since num_heads might be changed for TP num divisibility
if hasattr(config, "head_dim"):
self.head_dim = config.head_dim
self.layer_id = layer_id self.layer_id = layer_id
self.num_key_value_groups = self.num_heads // self.num_key_value_heads self.num_key_value_groups = self.num_heads // self.num_key_value_heads
self.q_local_size = self.num_local_heads * self.head_dim self.q_local_size = self.num_local_heads * self.head_dim
+36 -4
View File
@@ -126,7 +126,9 @@ class Llama4VisionPixelShuffleMLP(nn.Module):
super().__init__() super().__init__()
self.pixel_shuffle_ratio = config.pixel_shuffle_ratio self.pixel_shuffle_ratio = config.pixel_shuffle_ratio
self.mlp = Llama4VisionMLP( self.mlp = Llama4VisionMLP(
input_size=config.intermediate_size, input_size=getattr(
config, "original_intermediate_size", config.intermediate_size
),
intermediate_size=config.projector_input_dim, intermediate_size=config.projector_input_dim,
output_size=config.projector_output_dim, output_size=config.projector_output_dim,
bias=config.multi_modal_projector_bias, bias=config.multi_modal_projector_bias,
@@ -163,8 +165,17 @@ class Llama4VisionEncoderLayer(nn.Module):
): ):
super().__init__() super().__init__()
self.hidden_size = config.hidden_size self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads self.num_attention_heads = (
config.original_num_attention_heads
if hasattr(config, "original_num_attention_heads")
else config.num_attention_heads
)
self.intermediate_size = config.intermediate_size self.intermediate_size = config.intermediate_size
num_dummy_heads = 0
if hasattr(config, "original_num_attention_heads"):
num_dummy_heads = (
config.num_attention_heads - config.original_num_attention_heads
)
self.self_attn = VisionAttention( self.self_attn = VisionAttention(
self.hidden_size, self.hidden_size,
@@ -175,6 +186,7 @@ class Llama4VisionEncoderLayer(nn.Module):
quant_config=None, quant_config=None,
flatten_batch=False, flatten_batch=False,
prefix=add_prefix("self_attn", prefix), prefix=add_prefix("self_attn", prefix),
num_dummy_heads=num_dummy_heads,
qkv_bias=True, qkv_bias=True,
customized_position_embedding_applier=apply_position_embedding, customized_position_embedding_applier=apply_position_embedding,
) )
@@ -273,9 +285,16 @@ class Llama4UnfoldConvolution(nn.Module):
if isinstance(kernel_size, int): if isinstance(kernel_size, int):
kernel_size = (kernel_size, kernel_size) kernel_size = (kernel_size, kernel_size)
self.unfold = torch.nn.Unfold(kernel_size=kernel_size, stride=config.patch_size) self.unfold = torch.nn.Unfold(kernel_size=kernel_size, stride=config.patch_size)
output_size = (
config.hidden_size
// config.original_num_attention_heads
* config.num_attention_heads
if hasattr(config, "original_num_attention_heads")
else config.hidden_size
)
params = { params = {
"input_size": config.num_channels * kernel_size[0] * kernel_size[1], "input_size": config.num_channels * kernel_size[0] * kernel_size[1],
"output_size": config.hidden_size, "output_size": output_size,
"bias": False, "bias": False,
"quant_config": quant_config, "quant_config": quant_config,
"prefix": f"{prefix}.linear", "prefix": f"{prefix}.linear",
@@ -303,7 +322,12 @@ class Llama4VisionRotaryEmbedding(nn.Module):
img_idx[-1, -1] = -2 # ID_CLS_TOKEN img_idx[-1, -1] = -2 # ID_CLS_TOKEN
frequencies_x = img_idx % idx # get the coordinates of the 2d matrix along x frequencies_x = img_idx % idx # get the coordinates of the 2d matrix along x
frequencies_y = img_idx // idx # get the coordinates of the 2d matrix along y frequencies_y = img_idx // idx # get the coordinates of the 2d matrix along y
freq_dim = config.hidden_size // config.num_attention_heads // 2 num_attention_heads = (
config.original_num_attention_heads
if hasattr(config, "original_num_attention_heads")
else config.num_attention_heads
)
freq_dim = config.hidden_size // num_attention_heads // 2
rope_freq = 1.0 / ( rope_freq = 1.0 / (
config.rope_parameters["rope_theta"] config.rope_parameters["rope_theta"]
** (torch.arange(0, freq_dim, 2)[: (freq_dim // 2)].float() / freq_dim) ** (torch.arange(0, freq_dim, 2)[: (freq_dim // 2)].float() / freq_dim)
@@ -378,6 +402,14 @@ class Llama4VisionModel(nn.Module):
) -> torch.Tensor: ) -> torch.Tensor:
# Patch embedding # Patch embedding
hidden_state = self.patch_embedding(pixel_values) hidden_state = self.patch_embedding(pixel_values)
# If padded in patch embedding linear part, only retrieve valid slice
if (
hasattr(self.config, "original_num_attention_heads")
and self.config.num_attention_heads
> self.config.original_num_attention_heads
):
hidden_state = hidden_state[:, :, : self.config.hidden_size]
num_tiles, num_patches, hidden_dim = hidden_state.shape num_tiles, num_patches, hidden_dim = hidden_state.shape
# Add cls token # Add cls token
+3
View File
@@ -210,6 +210,9 @@ class Qwen2DecoderLayer(nn.Module):
self.start_layer = start_layer self.start_layer = start_layer
rope_theta, rope_scaling = get_rope_config(config) rope_theta, rope_scaling = get_rope_config(config)
max_position_embeddings = getattr(config, "max_position_embeddings", 32768) max_position_embeddings = getattr(config, "max_position_embeddings", 32768)
if hasattr(config, "original_num_attention_heads"):
head_dim = config.hidden_size // config.original_num_attention_heads
else:
head_dim = getattr(config, "head_dim", None) head_dim = getattr(config, "head_dim", None)
dual_chunk_attention_config = getattr( dual_chunk_attention_config = getattr(
config, "dual_chunk_attention_config", None config, "dual_chunk_attention_config", None
+14 -4
View File
@@ -74,9 +74,10 @@ from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model
from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel
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 add_prefix, is_cuda, is_npu from sglang.srt.utils import add_prefix, is_cpu, is_cuda, is_npu
_is_cuda = is_cuda() _is_cuda = is_cuda()
_is_cpu = is_cpu()
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -139,6 +140,7 @@ class Qwen2_5_VisionBlock(nn.Module):
dim: int, dim: int,
intermediate_dim: int, intermediate_dim: int,
num_heads: int, num_heads: int,
head_size: int,
hidden_act="silu", hidden_act="silu",
norm_layer: Type[nn.Module] = None, norm_layer: Type[nn.Module] = None,
quant_config: Optional[QuantizationConfig] = None, quant_config: Optional[QuantizationConfig] = None,
@@ -154,7 +156,8 @@ class Qwen2_5_VisionBlock(nn.Module):
self.attn = VisionAttention( self.attn = VisionAttention(
embed_dim=dim, embed_dim=dim,
num_heads=num_heads, num_heads=num_heads,
projection_size=dim, head_size=head_size,
projection_size=num_heads * head_size,
use_qkv_parallel=True, use_qkv_parallel=True,
proj_bias=True, proj_bias=True,
flatten_batch=True, flatten_batch=True,
@@ -212,6 +215,7 @@ class Qwen2_5_VisionPatchMerger(nn.Module):
self, self,
dim: int, dim: int,
context_dim: int, context_dim: int,
padded_context_dim: int,
spatial_merge_size: int = 2, spatial_merge_size: int = 2,
quant_config: Optional[QuantizationConfig] = None, quant_config: Optional[QuantizationConfig] = None,
prefix: str = "", prefix: str = "",
@@ -219,6 +223,7 @@ class Qwen2_5_VisionPatchMerger(nn.Module):
) -> None: ) -> None:
super().__init__() super().__init__()
self.hidden_size = context_dim * (spatial_merge_size**2) self.hidden_size = context_dim * (spatial_merge_size**2)
self.padded_context_dim = padded_context_dim * (spatial_merge_size**2)
self.ln_q = RMSNorm(context_dim, eps=1e-6) self.ln_q = RMSNorm(context_dim, eps=1e-6)
tp_size = 1 if use_data_parallel else get_parallel().tp_size tp_size = 1 if use_data_parallel else get_parallel().tp_size
tp_rank = 0 if use_data_parallel else get_parallel().tp_rank tp_rank = 0 if use_data_parallel else get_parallel().tp_rank
@@ -226,7 +231,7 @@ class Qwen2_5_VisionPatchMerger(nn.Module):
[ [
ColumnParallelLinear( ColumnParallelLinear(
self.hidden_size, self.hidden_size,
self.hidden_size, self.padded_context_dim,
bias=True, bias=True,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("mlp.0", prefix), prefix=add_prefix("mlp.0", prefix),
@@ -235,7 +240,7 @@ class Qwen2_5_VisionPatchMerger(nn.Module):
), ),
nn.GELU(), nn.GELU(),
RowParallelLinear( RowParallelLinear(
self.hidden_size, self.padded_context_dim,
dim, dim,
bias=True, bias=True,
quant_config=quant_config, quant_config=quant_config,
@@ -295,6 +300,9 @@ class Qwen2_5_VisionTransformer(nn.Module, RotaryPosMixin):
) )
norm_layer = partial(nn.LayerNorm, eps=norm_eps) norm_layer = partial(nn.LayerNorm, eps=norm_eps)
if _is_cpu and hasattr(vision_config, "original_num_heads"):
head_dim = hidden_size // vision_config.original_num_heads
else:
head_dim = hidden_size // num_heads head_dim = hidden_size // num_heads
self.rotary_pos_emb = Qwen2_5_VisionRotaryEmbedding(head_dim // 2) self.rotary_pos_emb = Qwen2_5_VisionRotaryEmbedding(head_dim // 2)
self.blocks = nn.ModuleList( self.blocks = nn.ModuleList(
@@ -303,6 +311,7 @@ class Qwen2_5_VisionTransformer(nn.Module, RotaryPosMixin):
dim=hidden_size, dim=hidden_size,
intermediate_dim=mlp_hidden_size, intermediate_dim=mlp_hidden_size,
num_heads=num_heads, num_heads=num_heads,
head_size=head_dim,
hidden_act=vision_config.hidden_act, hidden_act=vision_config.hidden_act,
norm_layer=norm_layer, norm_layer=norm_layer,
quant_config=quant_config, quant_config=quant_config,
@@ -315,6 +324,7 @@ class Qwen2_5_VisionTransformer(nn.Module, RotaryPosMixin):
self.merger = Qwen2_5_VisionPatchMerger( self.merger = Qwen2_5_VisionPatchMerger(
dim=vision_config.out_hidden_size, dim=vision_config.out_hidden_size,
context_dim=hidden_size, context_dim=hidden_size,
padded_context_dim=num_heads * head_dim,
spatial_merge_size=spatial_merge_size, spatial_merge_size=spatial_merge_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("merger", prefix), prefix=add_prefix("merger", prefix),