Add DP ViT support for Kimi K2.5 (#18689)

This commit is contained in:
Yuhao Yang
2026-02-18 23:03:07 +08:00
committed by GitHub
parent eb6ff4a940
commit 5a7ae059e3

View File

@@ -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,