[EPD][VLM] Support Kimi VL EPD (#22490)

Signed-off-by: LHXuuu <xulianhao.xlh@antgroup.com>
This commit is contained in:
LHXuuu
2026-04-16 12:40:02 +08:00
committed by GitHub
parent 58c6b871b2
commit e7ad7c587a
6 changed files with 268 additions and 102 deletions
@@ -13,6 +13,7 @@ from http import HTTPStatus
from typing import TYPE_CHECKING, Dict, List, Optional from typing import TYPE_CHECKING, Dict, List, Optional
import aiohttp import aiohttp
import numpy as np
import torch import torch
import zmq import zmq
import zmq.asyncio import zmq.asyncio
@@ -196,12 +197,27 @@ _VIDEO_META_ATTRS = ("video_timestamps", "second_per_grid_ts")
def _cat_grid(dims, flatten_items=False): def _cat_grid(dims, flatten_items=False):
"""Concatenate non-None tensors from a list; optionally flatten each before cat.""" """Concatenate non-None grid entries; supports tensor/ndarray/list inputs."""
valid = (
[g.flatten() for g in dims if g is not None] def _to_tensor(g):
if flatten_items if isinstance(g, torch.Tensor):
else [g for g in dims if g is not None] return g.cpu() if g.is_cuda else g
) if isinstance(g, np.ndarray):
return torch.from_numpy(g)
return torch.as_tensor(g)
valid = []
for g in dims:
if g is None:
continue
t = _to_tensor(g)
if flatten_items:
t = t.flatten()
elif t.ndim == 0:
# Keep cat semantics stable for scalar-like metadata.
t = t.unsqueeze(0)
valid.append(t)
return torch.cat(valid, dim=0) if valid else None return torch.cat(valid, dim=0) if valid else None
@@ -1021,6 +1037,26 @@ class MMReceiverBase(ABC):
return num_items_assigned return num_items_assigned
def _extract_url_data(self, request_obj) -> List[Dict]: def _extract_url_data(self, request_obj) -> List[Dict]:
def flatten_mm_items(items):
if not isinstance(items, list):
return [items]
flat = []
for item in items:
if isinstance(item, (list, tuple)):
flat.extend(flatten_mm_items(list(item)))
else:
flat.append(item)
return flat
def to_raw_url(mm_item):
if isinstance(mm_item, ImageData):
return mm_item.url
if isinstance(mm_item, dict):
# tolerate {"url": ...} shaped payloads
return mm_item.get("url", mm_item)
return mm_item
mm_data = [] mm_data = []
for attr, modality in [ for attr, modality in [
("image_data", Modality.IMAGE), ("image_data", Modality.IMAGE),
@@ -1029,16 +1065,11 @@ class MMReceiverBase(ABC):
]: ]:
mm_items = getattr(request_obj, attr, None) mm_items = getattr(request_obj, attr, None)
if mm_items: if mm_items:
if not isinstance(mm_items, list): mm_items = flatten_mm_items(mm_items)
mm_items = [mm_items]
for mm_item in mm_items: for mm_item in mm_items:
mm_data.append( mm_data.append(
{ {
"url": ( "url": to_raw_url(mm_item),
mm_item.url
if isinstance(mm_item, ImageData)
else mm_item
),
"modality": modality, "modality": modality,
} }
) )
@@ -141,7 +141,10 @@ def _get_mm_grid_dim(mm_inputs, modality, model_type: Optional[str] = None):
# Kimi K2.5 vision processor only emits `grid_thws`; prefer it over generic keys # Kimi K2.5 vision processor only emits `grid_thws`; prefer it over generic keys
# so we never pick a mis-typed or stale `image_grid_hws` field from kwargs. # so we never pick a mis-typed or stale `image_grid_hws` field from kwargs.
attrs = _mm_grid_attrs[modality] attrs = _mm_grid_attrs[modality]
if (model_type or "").lower() == "kimi_k25" and modality == Modality.IMAGE: if (model_type or "").lower() in [
"kimi_k25",
"kimi_vl",
] and modality == Modality.IMAGE:
attrs = ("grid_thws", "image_grid_thw", "image_grid_hws") attrs = ("grid_thws", "image_grid_thw", "image_grid_hws")
for attr in attrs: for attr in attrs:
if attr in mm_inputs and mm_inputs[attr] is not None: if attr in mm_inputs and mm_inputs[attr] is not None:
@@ -583,9 +586,7 @@ class MMEncoder:
else: else:
return int(grid[0] * grid[1] * grid[2]) return int(grid[0] * grid[1] * grid[2])
def _kimi_k25_tokens_from_patch_grid( def _kimi_tokens_from_patch_grid(self, grid: Union[torch.Tensor, List[int]]) -> int:
self, grid: Union[torch.Tensor, List[int]]
) -> int:
"""MoonViT + tpool: output len is (h//mh)*(w//mw); temporal dim is pooled (not t*h*w/merge^2).""" """MoonViT + tpool: output len is (h//mh)*(w//mw); temporal dim is pooled (not t*h*w/merge^2)."""
if isinstance(grid, torch.Tensor): if isinstance(grid, torch.Tensor):
flat = grid.flatten() flat = grid.flatten()
@@ -603,8 +604,11 @@ class MMEncoder:
input_length = self.get_num_patches(grid, modality) input_length = self.get_num_patches(grid, modality)
return self._get_feat_extract_output_lengths(input_length) return self._get_feat_extract_output_lengths(input_length)
else: else:
if self.model_type == "kimi_k25" and modality == Modality.IMAGE: if (
return self._kimi_k25_tokens_from_patch_grid(grid) self.model_type in ["kimi_k25", "kimi_vl"]
and modality == Modality.IMAGE
):
return self._kimi_tokens_from_patch_grid(grid)
merge_size = getattr(self.image_processor, "merge_size", 2) merge_size = getattr(self.image_processor, "merge_size", 2)
return self.get_num_patches(grid, modality) // (merge_size**2) return self.get_num_patches(grid, modality) // (merge_size**2)
@@ -869,30 +873,79 @@ class MMEncoder:
] ]
return timestamps return timestamps
def _normalize_kimi_k25_encoder_images(self, images): @staticmethod
"""KimiK25VisionProcessor.preprocess expects MediaInput dicts, not raw PIL.""" def _flatten_nested_items(items):
if not isinstance(items, (list, tuple)):
return [items]
flat = []
for item in items:
if isinstance(item, (list, tuple)):
flat.extend(MMEncoder._flatten_nested_items(item))
else:
flat.append(item)
return flat
def _normalize_kimi_encoder_images(self, images):
"""Normalize Kimi image inputs for the image processor call."""
from PIL import Image as PILImage from PIL import Image as PILImage
def wrap_one(img): def wrap_one(img):
if isinstance(img, dict) and img.get("type") in ("image", "video_chunk"): if isinstance(img, dict) and img.get("type") in ("image", "video_chunk"):
return img return [img]
if isinstance(img, PILImage.Image): if isinstance(img, PILImage.Image):
return {"type": "image", "image": img} return [{"type": "image", "image": img}]
return img return [img]
if not images: if not images:
return images return images
# Disagg may supply a nested list; Kimi preprocess expects a flat list of media.
if isinstance(images[0], (list, tuple)): # Disagg may supply nested lists from grouped routing.
images = [x for group in images for x in group] images = self._flatten_nested_items(images)
return [wrap_one(img) for img in images]
# Kimi-VL image processor expects a flat list of concrete images.
if self.model_type == "kimi_vl":
normalized = []
for img in images:
if (
isinstance(img, dict)
and img.get("type") == "image"
and "image" in img
):
inner = img["image"]
if isinstance(inner, (list, tuple)):
normalized.extend(self._flatten_nested_items(inner))
else:
normalized.append(inner)
else:
normalized.append(img)
return normalized
# Kimi-K2.5 vision processor expects media dicts.
normalized = []
for img in images:
wrapped = wrap_one(img)
for media in wrapped:
# Some pipelines may produce {"type": "image", "image": [PIL]}.
# Split it into one media item per concrete image object.
if (
isinstance(media, dict)
and media.get("type") == "image"
and isinstance(media.get("image"), (list, tuple))
):
for inner in self._flatten_nested_items(media["image"]):
normalized.append({**media, "image": inner})
else:
normalized.append(media)
return normalized
async def _process_mm_items(self, mm_items, modality): async def _process_mm_items(self, mm_items, modality):
if modality == Modality.IMAGE and self.image_processor: if modality == Modality.IMAGE and self.image_processor:
images = await self._flatten_and_load_images(mm_items) images = await self._flatten_and_load_images(mm_items)
image_config = self.vision_config.get("image", {}) image_config = self.vision_config.get("image", {})
if self.model_type == "kimi_k25": if self.model_type in ["kimi_k25", "kimi_vl"]:
images = self._normalize_kimi_k25_encoder_images(images) images = self._normalize_kimi_encoder_images(images)
processor_input = self.image_processor(images=images, **image_config) processor_input = self.image_processor(images=images, **image_config)
if hasattr(self.model, "thinker"): # for omni models if hasattr(self.model, "thinker"): # for omni models
get_feature_method = self.model.thinker.get_image_feature get_feature_method = self.model.thinker.get_image_feature
+23 -8
View File
@@ -128,13 +128,16 @@ class KimiVLForConditionalGeneration(nn.Module):
self.multi_modal_projector = KimiVLMultiModalProjector(config=config) self.multi_modal_projector = KimiVLMultiModalProjector(config=config)
self.quant_config = quant_config self.quant_config = quant_config
text_config = copy.deepcopy(config.text_config)
text_config.architectures = ["DeepseekV2ForCausalLM"] self.language_model = None
self.language_model = DeepseekV2ForCausalLM( if not config.encoder_only:
config=text_config, text_config = copy.deepcopy(config.text_config)
quant_config=quant_config, text_config.architectures = ["DeepseekV2ForCausalLM"]
prefix=add_prefix("language_model", prefix), self.language_model = DeepseekV2ForCausalLM(
) config=text_config,
quant_config=quant_config,
prefix=add_prefix("language_model", prefix),
)
def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor: def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
pixel_values = ( pixel_values = (
@@ -215,6 +218,13 @@ class KimiVLForConditionalGeneration(nn.Module):
for args in weights: for args in weights:
name, loaded_weight = args[:2] name, loaded_weight = args[:2]
kwargs = args[2] if len(args) > 2 else {} kwargs = args[2] if len(args) > 2 else {}
is_vision_weight = ("vision" in name) or ("multi_modal_projector" in name)
if self.config.encoder_only and not is_vision_weight:
continue
if self.config.language_only and is_vision_weight:
continue
if "rotary_emb.inv_freq" in name: if "rotary_emb.inv_freq" in name:
continue continue
@@ -251,6 +261,8 @@ class KimiVLForConditionalGeneration(nn.Module):
# Skip loading extra bias for GPTQ models. # Skip loading extra bias for GPTQ models.
if name.endswith(".bias") and name not in params_dict: if name.endswith(".bias") and name not in params_dict:
continue continue
if name not in params_dict:
continue
param = params_dict[name] param = params_dict[name]
weight_loader = param.weight_loader weight_loader = param.weight_loader
@@ -266,6 +278,8 @@ class KimiVLForConditionalGeneration(nn.Module):
if weight_name not in name: if weight_name not in name:
continue continue
name = name.replace(weight_name, param_name) name = name.replace(weight_name, param_name)
if name not in params_dict:
continue
param = params_dict[name] param = params_dict[name]
weight_loader = param.weight_loader weight_loader = param.weight_loader
@@ -295,7 +309,8 @@ class KimiVLForConditionalGeneration(nn.Module):
param = params_dict[name] param = params_dict[name]
weight_loader = getattr(param, "weight_loader", default_weight_loader) weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight, **kwargs) weight_loader(param, loaded_weight, **kwargs)
self.language_model.post_load_weights() if self.language_model is not None:
self.language_model.post_load_weights()
def get_spec_layer_idx_from_weight_name( def get_spec_layer_idx_from_weight_name(
@@ -0,0 +1,113 @@
"""Kimi-specific grid-based multimodal data helpers.
Shared by KimiVLImageProcessor and KimiK2_5VLImageProcessor.
"""
from typing import Union
import numpy as np
import torch
from sglang.srt.managers.schedule_batch import (
Modality,
MultimodalDataItem,
MultimodalProcessorOutput,
)
class KimiGridMMDataMixin:
"""Mixin providing Kimi-specific grid-based multimodal data helpers.
Expects the concrete class to supply:
- self.hf_config (with vision_config.merge_kernel_size)
- self._tokenizer (with .encode())
"""
def _num_image_tokens_from_grid(
self, grid_thw: Union[torch.Tensor, np.ndarray, list, tuple]
) -> int:
"""Compute Kimi-style image token count from 2D/3D grid metadata."""
merge_h, merge_w = self.hf_config.vision_config.merge_kernel_size
if isinstance(grid_thw, torch.Tensor):
vals = grid_thw.flatten().tolist()
elif isinstance(grid_thw, np.ndarray):
vals = grid_thw.reshape(-1).tolist()
elif isinstance(grid_thw, (list, tuple)):
vals = list(np.array(grid_thw).reshape(-1).tolist())
else:
raise TypeError(
f"Unsupported grid type for kimi image tokens: {type(grid_thw)}"
)
if len(vals) >= 3:
_t, h, w = vals[-3], vals[-2], vals[-1]
elif len(vals) == 2:
_t, h, w = 1, vals[0], vals[1]
else:
raise ValueError(
f"Invalid grid metadata for kimi image tokens: {vals} "
"(expected [t,h,w] or [h,w])"
)
h, w = int(h), int(w)
return (h * w) // (merge_h * merge_w)
def _build_kimi_mm_data_from_grids(
self, prompt, embeddings, **kwargs
) -> MultimodalProcessorOutput:
image_token_id = kwargs.get("image_token_id", 0)
img_grid_thw = kwargs.get("img_grid_thw", None)
if not isinstance(prompt, list):
prompt = self._tokenizer.encode(prompt)
image_token_counts = [
self._num_image_tokens_from_grid(grid) for grid in img_grid_thw
]
input_ids = []
offsets = []
img_idx = 0
for token in prompt:
if token != image_token_id:
input_ids.append(token)
continue
if img_idx >= len(image_token_counts):
raise ValueError(
"The number of image placeholders exceeds img_grid_thw entries."
)
num_tokens = image_token_counts[img_idx]
start = len(input_ids)
input_ids.extend([image_token_id] * num_tokens)
offsets.append((start, len(input_ids) - 1))
img_idx += 1
if img_idx != len(image_token_counts):
raise ValueError(
"The number of image placeholders does not match img_grid_thw entries."
)
image_embeddings = embeddings[Modality.IMAGE]
mm_items = []
consumed = 0
for start, end in offsets:
num_tokens = end - start + 1
embedding_slice = image_embeddings[consumed : consumed + num_tokens]
consumed += num_tokens
mm_items.append(
MultimodalDataItem(
modality=Modality.IMAGE,
offsets=[(start, end)],
precomputed_embeddings=embedding_slice,
)
)
return MultimodalProcessorOutput(
input_ids=input_ids,
mm_items=mm_items,
im_token_id=image_token_id,
)
@@ -9,8 +9,6 @@ import torch.nn.functional as F
from PIL import Image from PIL import Image
from sglang.srt.managers.schedule_batch import ( from sglang.srt.managers.schedule_batch import (
Modality,
MultimodalDataItem,
MultimodalProcessorOutput, MultimodalProcessorOutput,
) )
from sglang.srt.models.kimi_k25 import KimiK25ForConditionalGeneration from sglang.srt.models.kimi_k25 import KimiK25ForConditionalGeneration
@@ -20,6 +18,7 @@ from sglang.srt.multimodal.processors.base_processor import (
from sglang.srt.multimodal.processors.base_processor import ( from sglang.srt.multimodal.processors.base_processor import (
MultimodalSpecialTokens, MultimodalSpecialTokens,
) )
from sglang.srt.multimodal.processors.kimi_common import KimiGridMMDataMixin
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# GPU image preprocessing utilities (resize, pad, normalize, patchify on CUDA) # GPU image preprocessing utilities (resize, pad, normalize, patchify on CUDA)
@@ -333,7 +332,7 @@ class KimiGPUProcessorWrapper:
# Compatible with KimiVLForConditionalGeneration # Compatible with KimiVLForConditionalGeneration
class KimiK2_5VLImageProcessor(SGLangBaseProcessor): class KimiK2_5VLImageProcessor(KimiGridMMDataMixin, SGLangBaseProcessor):
models = [KimiK25ForConditionalGeneration] models = [KimiK25ForConditionalGeneration]
gpu_image_decode = True # nvJPEG for JPEG, PIL fallback for others gpu_image_decode = True # nvJPEG for JPEG, PIL fallback for others
@@ -386,66 +385,11 @@ class KimiK2_5VLImageProcessor(SGLangBaseProcessor):
im_token_id=self.mm_tokens.image_token_id, im_token_id=self.mm_tokens.image_token_id,
) )
def _num_image_tokens_from_grid(self, grid_thw: torch.Tensor) -> int:
# Kimi-K2.5 applies temporal pooling and spatial 2D merge in vision tower.
# The output sequence length per image is h*w/(merge_h*merge_w).
merge_h, merge_w = self.hf_config.vision_config.merge_kernel_size
_t, h, w = grid_thw.tolist()
return (h * w) // (merge_h * merge_w)
def get_mm_data(self, prompt, embeddings, **kwargs): def get_mm_data(self, prompt, embeddings, **kwargs):
img_grid_thw = kwargs.get("img_grid_thw", None) img_grid_thw = kwargs.get("img_grid_thw", None)
return self._build_kimi_mm_data_from_grids(
if not isinstance(prompt, list): prompt=prompt,
prompt = self._tokenizer.encode(prompt) embeddings=embeddings,
image_token_id=self.mm_tokens.image_token_id,
image_token_id = self.mm_tokens.image_token_id img_grid_thw=img_grid_thw,
image_token_counts = [
self._num_image_tokens_from_grid(grid) for grid in img_grid_thw
]
input_ids = []
offsets = []
img_idx = 0
for token in prompt:
if token != image_token_id:
input_ids.append(token)
continue
if img_idx >= len(image_token_counts):
raise ValueError(
"The number of image placeholders exceeds img_grid_thw entries."
)
num_tokens = image_token_counts[img_idx]
start = len(input_ids)
input_ids.extend([image_token_id] * num_tokens)
offsets.append((start, len(input_ids) - 1))
img_idx += 1
if img_idx != len(image_token_counts):
raise ValueError(
"The number of image placeholders does not match img_grid_thw entries."
)
image_embeddings = embeddings[Modality.IMAGE]
mm_items = []
consumed = 0
for start, end in offsets:
num_tokens = end - start + 1
embedding_slice = image_embeddings[consumed : consumed + num_tokens]
consumed += num_tokens
mm_items.append(
MultimodalDataItem(
modality=Modality.IMAGE,
offsets=[(start, end)],
precomputed_embeddings=embedding_slice,
)
)
return MultimodalProcessorOutput(
input_ids=input_ids,
mm_items=mm_items,
im_token_id=image_token_id,
) )
@@ -9,10 +9,11 @@ from sglang.srt.multimodal.processors.base_processor import (
from sglang.srt.multimodal.processors.base_processor import ( from sglang.srt.multimodal.processors.base_processor import (
MultimodalSpecialTokens, MultimodalSpecialTokens,
) )
from sglang.srt.multimodal.processors.kimi_common import KimiGridMMDataMixin
# Compatible with KimiVLForConditionalGeneration # Compatible with KimiVLForConditionalGeneration
class KimiVLImageProcessor(SGLangBaseProcessor): class KimiVLImageProcessor(KimiGridMMDataMixin, SGLangBaseProcessor):
models = [KimiVLForConditionalGeneration] models = [KimiVLForConditionalGeneration]
gpu_image_decode = False # KimiVL HF processor does not support tensor inputs gpu_image_decode = False # KimiVL HF processor does not support tensor inputs
@@ -48,3 +49,12 @@ class KimiVLImageProcessor(SGLangBaseProcessor):
mm_items=mm_items, mm_items=mm_items,
im_token_id=self.mm_tokens.image_token_id, im_token_id=self.mm_tokens.image_token_id,
) )
def get_mm_data(self, prompt, embeddings, **kwargs):
img_grid_thw = kwargs.get("img_grid_thw", None)
return self._build_kimi_mm_data_from_grids(
prompt=prompt,
embeddings=embeddings,
image_token_id=self.mm_tokens.image_token_id,
img_grid_thw=img_grid_thw,
)