support GLM-V vision model dp (#14097)

This commit is contained in:
Yuxuan Zhang
2025-12-05 21:03:54 +08:00
committed by GitHub
parent 5347732219
commit 8fce9e7b2a
4 changed files with 93 additions and 54 deletions
+3 -1
View File
@@ -220,7 +220,9 @@ class Glm4DecoderLayer(nn.Module):
rope_scaling = getattr(config, "rope_scaling", None) rope_scaling = getattr(config, "rope_scaling", None)
max_position_embeddings = getattr(config, "max_position_embeddings", 32768) max_position_embeddings = getattr(config, "max_position_embeddings", 32768)
head_dim = getattr(config, "head_dim", None) head_dim = getattr(config, "head_dim", None)
partial_rotary_factor = getattr(config, "partial_rotary_factor", None) partial_rotary_factor = getattr(
getattr(config, "rope_parameters", None), "partial_rotary_factor", None
) or getattr(config, "partial_rotary_factor", 0.5)
dual_chunk_attention_config = getattr( dual_chunk_attention_config = getattr(
config, "dual_chunk_attention_config", None config, "dual_chunk_attention_config", None
) )
+3 -1
View File
@@ -682,7 +682,9 @@ class Glm4MoeDecoderLayer(nn.Module):
self.config = config self.config = config
rope_theta = getattr(config, "rope_theta", 10000) rope_theta = getattr(config, "rope_theta", 10000)
rope_scaling = getattr(config, "rope_scaling", None) rope_scaling = getattr(config, "rope_scaling", None)
partial_rotary_factor = getattr(config, "partial_rotary_factor", 0.5) partial_rotary_factor = getattr(
getattr(config, "rope_parameters", None), "partial_rotary_factor", None
) or getattr(config, "partial_rotary_factor", 0.5)
max_position_embeddings = getattr(config, "max_position_embeddings", 8192) max_position_embeddings = getattr(config, "max_position_embeddings", 8192)
head_dim = getattr( head_dim = getattr(
config, "head_dim", config.hidden_size // config.num_attention_heads config, "head_dim", config.hidden_size // config.num_attention_heads
+86 -52
View File
@@ -27,18 +27,24 @@ import torch.nn.functional as F
from einops import rearrange from einops import rearrange
from transformers.models.glm4v.configuration_glm4v import Glm4vConfig, Glm4vVisionConfig from transformers.models.glm4v.configuration_glm4v import Glm4vConfig, Glm4vVisionConfig
from sglang.srt.distributed import (
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
)
from sglang.srt.distributed.parallel_state import get_pp_group
from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.activation import SiluAndMul
from sglang.srt.layers.attention import vision_utils from sglang.srt.layers.attention import vision_utils
from sglang.srt.layers.attention.vision import VisionAttention from sglang.srt.layers.attention.vision import VisionAttention
from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.layernorm import LayerNorm, RMSNorm
from sglang.srt.layers.linear import ( from sglang.srt.layers.linear import (
ColumnParallelLinear,
MergedColumnParallelLinear, MergedColumnParallelLinear,
ReplicatedLinear,
RowParallelLinear, RowParallelLinear,
) )
from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.logits_processor import LogitsProcessor
from sglang.srt.layers.pooler import Pooler, PoolingType from sglang.srt.layers.pooler import Pooler, PoolingType
from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.utils import PPMissingLayer
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
from sglang.srt.managers.mm_utils import ( from sglang.srt.managers.mm_utils import (
MultiModalityDataPaddingPatternMultimodalTokens, MultiModalityDataPaddingPatternMultimodalTokens,
@@ -48,6 +54,8 @@ from sglang.srt.managers.schedule_batch import MultimodalDataItem, MultimodalInp
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.glm4 import Glm4Model from sglang.srt.models.glm4 import Glm4Model
from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix from sglang.srt.utils import add_prefix
from sglang.srt.utils.hf_transformers_utils import get_processor from sglang.srt.utils.hf_transformers_utils import get_processor
@@ -73,14 +81,21 @@ class Glm4vVisionMLP(nn.Module):
bias: bool = False, bias: bool = False,
quant_config: Optional[QuantizationConfig] = None, quant_config: Optional[QuantizationConfig] = None,
prefix: str = "", prefix: str = "",
use_data_parallel: bool = False,
): ):
super().__init__() super().__init__()
self.tp_size = (
1 if use_data_parallel else get_tensor_model_parallel_world_size()
)
self.tp_rank = 0 if use_data_parallel else get_tensor_model_parallel_rank()
self.gate_up_proj = MergedColumnParallelLinear( self.gate_up_proj = MergedColumnParallelLinear(
input_size=in_features, input_size=in_features,
output_sizes=[hidden_features] * 2, # [gate_proj, up_proj] output_sizes=[hidden_features] * 2, # [gate_proj, up_proj]
bias=bias, bias=bias,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("gate_up_proj", prefix), prefix=add_prefix("gate_up_proj", prefix),
tp_size=self.tp_size,
tp_rank=self.tp_rank,
) )
self.down_proj = RowParallelLinear( self.down_proj = RowParallelLinear(
hidden_features, hidden_features,
@@ -88,6 +103,8 @@ class Glm4vVisionMLP(nn.Module):
bias=bias, bias=bias,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("down_proj", prefix), prefix=add_prefix("down_proj", prefix),
tp_size=self.tp_size,
tp_rank=self.tp_rank,
) )
self.act_fn = SiluAndMul() self.act_fn = SiluAndMul()
@@ -108,6 +125,7 @@ class Glm4vVisionBlock(nn.Module):
prefix: str = "", prefix: str = "",
num_dummy_heads: int = 0, num_dummy_heads: int = 0,
rms_norm_eps: float = 1e-5, rms_norm_eps: float = 1e-5,
use_data_parallel: bool = False,
) -> None: ) -> None:
super().__init__() super().__init__()
self.norm1 = RMSNorm(dim, eps=rms_norm_eps) self.norm1 = RMSNorm(dim, eps=rms_norm_eps)
@@ -123,12 +141,14 @@ class Glm4vVisionBlock(nn.Module):
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("attn", prefix), prefix=add_prefix("attn", prefix),
num_dummy_heads=num_dummy_heads, num_dummy_heads=num_dummy_heads,
use_data_parallel=use_data_parallel,
) )
self.mlp = Glm4vVisionMLP( self.mlp = Glm4vVisionMLP(
dim, dim,
intermediate_dim, intermediate_dim,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("mlp", prefix), prefix=add_prefix("mlp", prefix),
use_data_parallel=use_data_parallel,
) )
def forward( def forward(
@@ -206,24 +226,28 @@ class Glm4vPatchMerger(nn.Module):
quant_config: Optional[QuantizationConfig] = None, quant_config: Optional[QuantizationConfig] = None,
bias: bool = False, bias: bool = False,
prefix: str = "", prefix: str = "",
use_data_parallel: bool = False,
) -> None: ) -> None:
super().__init__() super().__init__()
self.hidden_size = d_model self.hidden_size = d_model
self.proj = ColumnParallelLinear( tp_size = 1 if use_data_parallel else get_tensor_model_parallel_world_size()
tp_rank = 0 if use_data_parallel else get_tensor_model_parallel_rank()
self.proj = ReplicatedLinear(
self.hidden_size, self.hidden_size,
self.hidden_size, self.hidden_size,
bias=bias, bias=bias,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("proj", prefix), prefix=add_prefix("proj", prefix),
gather_output=True,
) )
self.post_projection_norm = nn.LayerNorm(self.hidden_size) self.post_projection_norm = LayerNorm(self.hidden_size)
self.gate_up_proj = MergedColumnParallelLinear( self.gate_up_proj = MergedColumnParallelLinear(
input_size=self.hidden_size, input_size=self.hidden_size,
output_sizes=[context_dim] * 2, output_sizes=[context_dim] * 2,
bias=bias, bias=bias,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("gate_up_proj", prefix), prefix=add_prefix("gate_up_proj", prefix),
tp_size=tp_size,
tp_rank=tp_rank,
) )
self.down_proj = RowParallelLinear( self.down_proj = RowParallelLinear(
context_dim, context_dim,
@@ -231,6 +255,8 @@ class Glm4vPatchMerger(nn.Module):
bias=bias, bias=bias,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("down_proj", prefix), prefix=add_prefix("down_proj", prefix),
tp_size=tp_size,
tp_rank=tp_rank,
) )
self.extra_activation_func = nn.GELU() self.extra_activation_func = nn.GELU()
@@ -379,6 +405,7 @@ class Glm4vVisionModel(nn.Module):
vision_config: Glm4vVisionConfig, vision_config: Glm4vVisionConfig,
quant_config: Optional[QuantizationConfig] = None, quant_config: Optional[QuantizationConfig] = None,
prefix: str = "", prefix: str = "",
use_data_parallel: bool = False,
) -> None: ) -> None:
super().__init__() super().__init__()
@@ -392,6 +419,7 @@ class Glm4vVisionModel(nn.Module):
self.patch_size = vision_config.patch_size self.patch_size = vision_config.patch_size
self.spatial_merge_size = vision_config.spatial_merge_size self.spatial_merge_size = vision_config.spatial_merge_size
self.out_hidden_size = vision_config.out_hidden_size self.out_hidden_size = vision_config.out_hidden_size
self.use_data_parallel = use_data_parallel
self.patch_embed = Glm4vVisionPatchEmbed( self.patch_embed = Glm4vVisionPatchEmbed(
patch_size=patch_size, patch_size=patch_size,
@@ -412,6 +440,7 @@ class Glm4vVisionModel(nn.Module):
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix(f"blocks.{layer_idx}", prefix), prefix=add_prefix(f"blocks.{layer_idx}", prefix),
rms_norm_eps=vision_config.rms_norm_eps, rms_norm_eps=vision_config.rms_norm_eps,
use_data_parallel=use_data_parallel,
) )
for layer_idx in range(depth) for layer_idx in range(depth)
] ]
@@ -423,6 +452,7 @@ class Glm4vVisionModel(nn.Module):
quant_config=quant_config, quant_config=quant_config,
bias=False, bias=False,
prefix=add_prefix("merger", prefix), prefix=add_prefix("merger", prefix),
use_data_parallel=use_data_parallel,
) )
self.embeddings = Glm4vVisionEmbeddings(vision_config) self.embeddings = Glm4vVisionEmbeddings(vision_config)
@@ -527,11 +557,14 @@ class Glm4vForConditionalGeneration(nn.Module):
) -> None: ) -> None:
super().__init__() super().__init__()
self.pp_group = get_pp_group()
self.config = config self.config = config
self.use_data_parallel = get_global_server_args().mm_enable_dp_encoder
self.visual = Glm4vVisionModel( self.visual = Glm4vVisionModel(
config.vision_config, config.vision_config,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("visual", prefix), prefix=add_prefix("visual", prefix),
use_data_parallel=self.use_data_parallel,
) )
vision_utils.update_vit_attn_dummy_heads_config(self.config) vision_utils.update_vit_attn_dummy_heads_config(self.config)
@@ -542,15 +575,19 @@ class Glm4vForConditionalGeneration(nn.Module):
prefix=add_prefix("model", prefix), prefix=add_prefix("model", prefix),
) )
if config.tie_word_embeddings: if self.pp_group.is_last_rank:
self.lm_head = self.model.embed_tokens if self.pp_group.world_size == 1 and self.config.tie_word_embeddings:
self.lm_head = self.model.embed_tokens
else:
self.lm_head = ParallelLMHead(
self.config.vocab_size,
self.config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
)
else: else:
self.lm_head = ParallelLMHead( # ranks other than the last rank will have a placeholder layer
config.vocab_size, self.lm_head = PPMissingLayer()
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
)
self.is_mrope_enabled = "mrope_section" in self.config.rope_scaling self.is_mrope_enabled = "mrope_section" in self.config.rope_scaling
@@ -565,45 +602,36 @@ class Glm4vForConditionalGeneration(nn.Module):
return pattern.pad_input_tokens(input_ids, mm_inputs) return pattern.pad_input_tokens(input_ids, mm_inputs)
def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor: def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
pixel_values = torch.cat( # in GLM-V, last dim is the same
[item.feature.squeeze(0) for item in items], dim=0 pixel_values = torch.cat([item.feature for item in items], dim=0).type(
).type(self.visual.dtype) self.visual.dtype
)
image_grid_thw = torch.concat([item.image_grid_thw for item in items], dim=0) image_grid_thw = torch.concat([item.image_grid_thw for item in items], dim=0)
# For multi-image, pixel_values is [num_of_images, L, C] shape assert pixel_values.dim() == 2, pixel_values.dim()
# assert pixel_values.dim() == 2, pixel_values.dim()
assert image_grid_thw.dim() == 2, image_grid_thw.dim() assert image_grid_thw.dim() == 2, image_grid_thw.dim()
image_embeds = self.visual(pixel_values, grid_thw=image_grid_thw) if self.use_data_parallel:
split_sizes = ( return run_dp_sharded_mrope_vision_model(
image_grid_thw.prod(-1) // self.visual.spatial_merge_size**2 self.visual, pixel_values, image_grid_thw.tolist(), rope_type="rope_3d"
).tolist() )
image_embeds = torch.split(image_embeds, split_sizes) else:
return torch.cat(image_embeds) image_embeds = self.visual(pixel_values, grid_thw=image_grid_thw)
return image_embeds
def get_video_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor: def get_video_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
pixel_values_videos = torch.cat( # in GLM-V, last dim is the same
[item.feature.squeeze(0) for item in items], dim=0 pixel_values = torch.cat([item.feature for item in items], dim=0).type(
).type(self.visual.dtype) self.visual.dtype
video_grid_thw = torch.concat([item.video_grid_thw for item in items], dim=0)
# For multi-video, pixel_values_videos is [num_of_videos, L, C] shape
# assert pixel_values_videos.dim() == 2, pixel_values_videos.dim()
assert video_grid_thw.dim() == 2, video_grid_thw.dim()
# reshape video_grid_thw -> [b, 3] -> [1, h, w] * frames
temp_frames_hw = []
for t, h, w in video_grid_thw:
repeated_row = (
torch.tensor([1, h.item(), w.item()]).unsqueeze(0).repeat(t, 1)
)
temp_frames_hw.append(repeated_row)
flattened_video_grid_thw = torch.cat(temp_frames_hw, dim=0)
video_embeds = self.visual(
pixel_values_videos, grid_thw=flattened_video_grid_thw
) )
split_sizes = ( video_grid_thw = torch.concat([item.video_grid_thw for item in items], dim=0)
video_grid_thw.prod(-1) // self.visual.spatial_merge_size**2 assert pixel_values.dim() == 2, pixel_values.dim()
).tolist() assert video_grid_thw.dim() == 2, video_grid_thw.dim()
video_embeds = torch.split(video_embeds, split_sizes) if self.use_data_parallel:
return torch.cat(video_embeds) return run_dp_sharded_mrope_vision_model(
self.visual, pixel_values, video_grid_thw.tolist(), rope_type="rope_3d"
)
else:
video_embeds = self.visual(pixel_values, grid_thw=video_grid_thw)
return video_embeds
def get_input_embeddings(self): def get_input_embeddings(self):
return self.model.embed_tokens return self.model.embed_tokens
@@ -653,12 +681,18 @@ class Glm4vForConditionalGeneration(nn.Module):
if self.capture_aux_hidden_states: if self.capture_aux_hidden_states:
hidden_states, aux_hidden_states = hidden_states hidden_states, aux_hidden_states = hidden_states
if not get_embedding: if self.pp_group.is_last_rank:
return self.logits_processor( if not get_embedding:
input_ids, hidden_states, self.lm_head, forward_batch, aux_hidden_states return self.logits_processor(
) input_ids,
hidden_states,
self.lm_head,
forward_batch,
)
else:
return self.pooler(hidden_states, forward_batch)
else: else:
return self.pooler(hidden_states, forward_batch) return hidden_states
def _pad_vit_attn_dummy_heads(self, name: str, loaded_weight: torch.Tensor): def _pad_vit_attn_dummy_heads(self, name: str, loaded_weight: torch.Tensor):
"""pad attn qkv weights for dummy heads""" """pad attn qkv weights for dummy heads"""
+1
View File
@@ -24,6 +24,7 @@ MODELS = [
SimpleNamespace(model="Qwen/Qwen2.5-VL-72B-Instruct", mmmu_accuracy=0.55), SimpleNamespace(model="Qwen/Qwen2.5-VL-72B-Instruct", mmmu_accuracy=0.55),
SimpleNamespace(model="Qwen/Qwen3-VL-32B-Instruct", mmmu_accuracy=0.55), SimpleNamespace(model="Qwen/Qwen3-VL-32B-Instruct", mmmu_accuracy=0.55),
SimpleNamespace(model="OpenGVLab/InternVL2_5-8B", mmmu_accuracy=0.52), SimpleNamespace(model="OpenGVLab/InternVL2_5-8B", mmmu_accuracy=0.52),
SimpleNamespace(model="zai-org/GLM-4.1V-9B-Thinking", mmmu_accuracy=0.68),
] ]