From 5a7ae059e37f2c481462a2ad965467264718c461 Mon Sep 17 00:00:00 2001 From: Yuhao Yang <47235274+yhyang201@users.noreply.github.com> Date: Wed, 18 Feb 2026 23:03:07 +0800 Subject: [PATCH] Add DP ViT support for Kimi K2.5 (#18689) --- python/sglang/srt/models/kimi_k25.py | 24 ++++++++++++++++++++---- 1 file changed, 20 insertions(+), 4 deletions(-) diff --git a/python/sglang/srt/models/kimi_k25.py b/python/sglang/srt/models/kimi_k25.py index 1b7f20aca..d8399a691 100644 --- a/python/sglang/srt/models/kimi_k25.py +++ b/python/sglang/srt/models/kimi_k25.py @@ -35,6 +35,8 @@ from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.deepseek_v2 import DeepseekV3ForCausalLM from sglang.srt.models.kimi_vl_moonvit import MLP2 from sglang.srt.models.utils import WeightsMapper +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 KIMIV_VT_INFER_MAX_PATCH_NUM = 16328 @@ -475,9 +477,10 @@ class MoonViT3dPretrainedModel(nn.Module): _supports_flash_attn_2 = True _supports_sdpa = True - def __init__(self, config, *inputs, **kwargs): + def __init__(self, config, *inputs, use_data_parallel: bool = False, **kwargs): super().__init__() config = deepcopy(config) + self.config = config self.merge_kernel_size = config.merge_kernel_size self.patch_size = config.patch_size self.merge_type = config.merge_type @@ -500,6 +503,7 @@ class MoonViT3dPretrainedModel(nn.Module): "mlp_dim": config.intermediate_size, "activation": PytorchGELUTanh(), "attn_bias": True, + "use_data_parallel": use_data_parallel, }, video_attn_type=config.video_attn_type, ) @@ -541,11 +545,9 @@ class K2VLMultiModalProjector(nn.Module): def __init__( self, config: KimiK25VisionConfig, - use_data_parallel: bool = False, prefix: str = "", ): super().__init__() - self.use_data_parallel = use_data_parallel # Hidden size after patch merging merge_h, merge_w = config.merge_kernel_size @@ -663,8 +665,11 @@ class KimiK25ForConditionalGeneration(nn.Module): super().__init__() self.config = config self.quant_config = quant_config + self.use_data_parallel = get_global_server_args().mm_enable_dp_encoder # Create vision tower - self.vision_tower = MoonViT3dPretrainedModel(config.vision_config) + self.vision_tower = MoonViT3dPretrainedModel( + config.vision_config, use_data_parallel=self.use_data_parallel + ) # Create mm projector self.mm_projector = K2VLMultiModalProjector(config.vision_config) @@ -687,6 +692,17 @@ class KimiK25ForConditionalGeneration(nn.Module): target_dtype = self.vision_tower.patch_embed.proj.weight.dtype pixel_values = pixel_values.to(target_dtype) + + if self.use_data_parallel: + image_embeds = run_dp_sharded_mrope_vision_model( + self.vision_tower, + pixel_values, + grid_thws.tolist(), + rope_type="rope_2d", + ) + image_features = self.mm_projector(image_embeds) + return image_features + image_features = vision_tower_forward_auto( self.vision_tower, pixel_values,