[Perf] Speed up the Kimi-K2.5 vision path and match PIL bicubic in the GPU resize (#33349)

Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
Liangsheng Yin
2026-08-03 21:57:49 -07:00
committed by GitHub
co-authored by Mick
parent 7ba393dd15
commit afc868517b
11 changed files with 761 additions and 181 deletions
+3 -1
View File
@@ -148,6 +148,7 @@ def prepare_vision_attention_metadata(
cu_seqlens: torch.Tensor, cu_seqlens: torch.Tensor,
device: torch.device, device: torch.device,
*, *,
max_seqlen: Optional[int] = None,
packed_indptrs: Optional[torch.Tensor] = None, packed_indptrs: Optional[torch.Tensor] = None,
sequence_lengths: Optional[torch.Tensor] = None, sequence_lengths: Optional[torch.Tensor] = None,
flashinfer_max_seqlen: Optional[int] = None, flashinfer_max_seqlen: Optional[int] = None,
@@ -156,7 +157,8 @@ def prepare_vision_attention_metadata(
cu_seqlens = cu_seqlens.to(device=device, dtype=torch.int32, non_blocking=True) cu_seqlens = cu_seqlens.to(device=device, dtype=torch.int32, non_blocking=True)
seq_lens = cu_seqlens[1:] - cu_seqlens[:-1] seq_lens = cu_seqlens[1:] - cu_seqlens[:-1]
max_seqlen = int(seq_lens.max().item()) if max_seqlen is None:
max_seqlen = int(seq_lens.max().item())
return VisionAttentionMetadata( return VisionAttentionMetadata(
cu_seqlens=cu_seqlens, cu_seqlens=cu_seqlens,
seq_lens=seq_lens, seq_lens=seq_lens,
+55 -56
View File
@@ -8,6 +8,11 @@ import torch.nn.functional as F
from torch import nn from torch import nn
from transformers.activations import PytorchGELUTanh from transformers.activations import PytorchGELUTanh
from sglang.kernels.ops.attention.vision_rope import (
PreparedInplaceComplexRoPE,
apply_fused_qk_complex_rope_inplace,
prepare_fused_qk_complex_rope_inplace,
)
from sglang.srt.configs.kimi_k25 import KimiK25Config, KimiK25VisionConfig from sglang.srt.configs.kimi_k25 import KimiK25Config, KimiK25VisionConfig
from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation
from sglang.srt.layers.attention.vision import ( from sglang.srt.layers.attention.vision import (
@@ -33,32 +38,46 @@ from sglang.srt.managers.schedule_batch import (
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
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.deepseek_v2 import DeepseekV3ForCausalLM from sglang.srt.models.deepseek_v2 import DeepseekV3ForCausalLM
from sglang.srt.models.kimi_vl_moonvit import MLP2 from sglang.srt.models.kimi_vl_moonvit import MLP2, tpool_patch_merger
from sglang.srt.models.utils import WeightsMapper from sglang.srt.models.utils import WeightsMapper
from sglang.srt.multimodal.mm_utils import ( from sglang.srt.multimodal.mm_utils import (
concat_or_single,
materialize_multimodal_features, materialize_multimodal_features,
run_dp_sharded_mrope_vision_model, run_dp_sharded_mrope_vision_model,
) )
from sglang.srt.runtime_context import get_mm, get_parallel, get_server_args from sglang.srt.runtime_context import (
from sglang.srt.utils import add_prefix, is_npu get_exec,
get_mm,
get_parallel,
get_server_args,
)
from sglang.srt.utils import add_prefix, is_cuda, is_npu
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
_is_npu = is_npu() _is_npu = is_npu()
_is_cuda = is_cuda()
def apply_rope( def apply_rope(
xq: torch.Tensor, xk: torch.Tensor, freqs_cis: torch.Tensor, x_shape=None xq: torch.Tensor,
xk: torch.Tensor,
freqs_cis: torch.Tensor | PreparedInplaceComplexRoPE,
x_shape=None,
) -> tuple[torch.Tensor, torch.Tensor]: ) -> tuple[torch.Tensor, torch.Tensor]:
""" """
Args: (The leading dimensions of all inputs should be the same) Args: (The leading dimensions of all inputs should be the same)
xq: query, tensor of shape (..., num_heads, head_dim) xq: query, tensor of shape (..., num_heads, head_dim)
xk: key, tensor of shape (..., num_heads, head_dim) xk: key, tensor of shape (..., num_heads, head_dim)
freqs_cis: tensor of shape (..., head_dim/2), dtype=torch.complex64. It contains the precomputed cis(freqs) for each position in the 2D grid. freqs_cis: Complex frequencies for the portable path, or inputs
prepared once for the contiguous in-place CUDA kernel.
Returns: Returns:
xq_out, xk_out: tensors of shape (..., num_heads, head_dim) xq_out, xk_out: tensors of shape (..., num_heads, head_dim)
""" """
if isinstance(freqs_cis, tuple):
return apply_fused_qk_complex_rope_inplace(xq, xk, freqs_cis)
freqs_cis = freqs_cis.unsqueeze(-2) # ..., 1, head_dim/2 freqs_cis = freqs_cis.unsqueeze(-2) # ..., 1, head_dim/2
# ..., num_heads, head_dim/2 # ..., num_heads, head_dim/2
xq_ = torch.view_as_complex(xq.float().view(*xq.shape[:-1], -1, 2)) xq_ = torch.view_as_complex(xq.float().view(*xq.shape[:-1], -1, 2))
@@ -68,36 +87,6 @@ def apply_rope(
return xq_out.type_as(xq), xk_out.type_as(xk) return xq_out.type_as(xq), xk_out.type_as(xk)
def tpool_patch_merger(
x: torch.Tensor,
grid_thws: torch.Tensor,
merge_kernel_size: tuple[int, int] = (2, 2),
) -> list[torch.Tensor]:
d_model = x.size(-1)
outputs = []
pre_sum = 0
for t, h, w in grid_thws.tolist():
# Get the current sequence
seq = x[pre_sum : pre_sum + t * h * w]
# Reshape along self.merge_kernel_size and concat to the last dimension
kernel_height, kernel_width = merge_kernel_size
new_height, new_width = h // kernel_height, w // kernel_width
reshaped_seq = seq.view(
t, new_height, kernel_height, new_width, kernel_width, d_model
)
reshaped_seq = (
reshaped_seq.permute(0, 1, 3, 2, 4, 5).contiguous().mean(dim=0)
) # temporal pooling
padded_seq = reshaped_seq.view(
new_height * new_width, kernel_height * kernel_width, -1
)
outputs.append(padded_seq)
pre_sum += t * h * w
return outputs
class MoonViTEncoderLayer(nn.Module): class MoonViTEncoderLayer(nn.Module):
def __init__( def __init__(
@@ -433,6 +422,9 @@ class MoonVision3dPatchEmbed(nn.Module):
class MoonViT3dEncoder(nn.Module): class MoonViT3dEncoder(nn.Module):
# Class-level default so forward() stays usable on instances built with
# __new__ (unit tests skip __init__).
use_fused_rope = False
def __init__( def __init__(
self, self,
@@ -452,6 +444,9 @@ class MoonViT3dEncoder(nn.Module):
self.rope_2d = Rope2DPosEmbRepeated( self.rope_2d = Rope2DPosEmbRepeated(
block_cfg["hidden_dim"] // block_cfg["num_heads"], 512, 512 block_cfg["hidden_dim"] // block_cfg["num_heads"], 512, 512
) )
self.use_fused_rope = (
_is_cuda and get_exec().deterministic.rl_on_policy_target is None
)
self.blocks = nn.ModuleList( self.blocks = nn.ModuleList(
[ [
MoonViTEncoderLayer( MoonViTEncoderLayer(
@@ -472,8 +467,18 @@ class MoonViT3dEncoder(nn.Module):
rope_freqs_cis = self.rope_2d.get_freqs_cis( rope_freqs_cis = self.rope_2d.get_freqs_cis(
grid_thws=grid_thws, device=hidden_states.device grid_thws=grid_thws, device=hidden_states.device
) )
# The in-place kernel is a JIT template on the q/k dtype, and only
# fp16/bf16 are exercised by test_vision_rope_inplace. Leave other
# dtypes on the portable path rather than ship an untested one.
if self.use_fused_rope and hidden_states.dtype in (
torch.float16,
torch.bfloat16,
):
rope_freqs_cis = prepare_fused_qk_complex_rope_inplace(rope_freqs_cis)
sequence_lengths = (grid_thws[:, 0] * grid_thws[:, 1] * grid_thws[:, 2]).to( sequence_lengths = grid_thws[:, 0] * grid_thws[:, 1] * grid_thws[:, 2]
max_seqlen = int(sequence_lengths.max().item())
sequence_lengths = sequence_lengths.to(
device=hidden_states.device, dtype=torch.int32 device=hidden_states.device, dtype=torch.int32
) )
lengths = torch.cat( lengths = torch.cat(
@@ -483,14 +488,12 @@ class MoonViT3dEncoder(nn.Module):
) )
) )
# FlashAttention needs a host integer. Compute it once per MoonViT
# forward and pass it to every encoder block instead of synchronizing
# once per block inside the attention backend.
max_seqlen = int(lengths.max().item())
cu_seqlens = lengths.to(hidden_states.device).cumsum(dim=0, dtype=torch.int32) cu_seqlens = lengths.to(hidden_states.device).cumsum(dim=0, dtype=torch.int32)
forward_metadata = prepare_vision_attention_metadata( forward_metadata = prepare_vision_attention_metadata(
cu_seqlens, device=hidden_states.device cu_seqlens,
device=hidden_states.device,
max_seqlen=max_seqlen,
) )
for block in self.blocks: for block in self.blocks:
@@ -625,18 +628,12 @@ class K2VLMultiModalProjector(nn.Module):
@torch.inference_mode() @torch.inference_mode()
def mm_projection_auto( def mm_projection_auto(
mm_projector: torch.nn.Module | None, vt_output: list[torch.Tensor] mm_projector: torch.nn.Module,
): vt_output: Sequence[torch.Tensor],
"""Apply MM projector to vision tower outputs.""" ) -> torch.Tensor:
if mm_projector is None: """Project MoonViT's per-image outputs into one flattened (tokens, dim) feature."""
return vt_output projected = mm_projector(concat_or_single(vt_output, dim=0))
return projected.reshape(-1, projected.shape[-1])
num_embedding_list = [x.shape[0] for x in vt_output]
batched = torch.cat(vt_output, dim=0)
proj_out = mm_projector(batched) if mm_projector else batched
proj_out = proj_out.reshape(-1, proj_out.shape[-1])
proj_out = torch.split(proj_out, num_embedding_list)
return proj_out
class KimiK25ForConditionalGeneration(nn.Module): class KimiK25ForConditionalGeneration(nn.Module):
@@ -769,9 +766,11 @@ class KimiK25ForConditionalGeneration(nn.Module):
return image_features return image_features
pixel_values = materialize_item_features(list(range(len(items)))) pixel_values = materialize_item_features(list(range(len(items))))
image_embeds = self.vision_tower(pixel_values, grid_thws.to(device)) # grid_thws stays on the host: MoonViT3d only reads it as shape metadata
proj_out = mm_projection_auto(self.mm_projector, image_embeds) # (.tolist() in the pos-emb, RoPE and merger), so a device copy would
return torch.cat(proj_out, dim=0) # buy one sync per read. Same contract the encoder-DP path relies on.
image_embeds = self.vision_tower(pixel_values, grid_thws)
return mm_projection_auto(self.mm_projector, image_embeds)
def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs): def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs):
pattern = MultiModalityDataPaddingPatternMultimodalTokens() pattern = MultiModalityDataPaddingPatternMultimodalTokens()
+30 -1
View File
@@ -1,7 +1,7 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
# ruff: noqa: E501 # ruff: noqa: E501
# Adapted from https://huggingface.co/moonshotai/Kimi-VL-A3B-Instruct/blob/main/modeling_kimi_vl.py # Adapted from https://huggingface.co/moonshotai/Kimi-VL-A3B-Instruct/blob/main/modeling_kimi_vl.py
# This file is meant to be used in kimi_vl.py only # Shared MoonViT building blocks for kimi_vl.py and kimi_k25.py
# Copyright 2025 The Moonshot AI Team, DeepSeek-AI, and HuggingFace Inc. team. All rights reserved. # Copyright 2025 The Moonshot AI Team, DeepSeek-AI, and HuggingFace Inc. team. All rights reserved.
# #
# The code is based on llava (llava/modeling_llava.py) and DeepSeek-V3 (DeepSeek-V3/modeling_deepseek.py), but modified for KimiVL. # The code is based on llava (llava/modeling_llava.py) and DeepSeek-V3 (DeepSeek-V3/modeling_deepseek.py), but modified for KimiVL.
@@ -566,6 +566,35 @@ def patch_merger(
return outputs return outputs
def tpool_patch_merger(
x: torch.Tensor,
grid_thws: torch.Tensor,
merge_kernel_size: tuple[int, int] = (2, 2),
) -> List[torch.Tensor]:
"""Group spatial patches and average only across real video frames."""
d_model = x.size(-1)
outputs = []
pre_sum = 0
for t, h, w in grid_thws.tolist():
seq = x[pre_sum : pre_sum + t * h * w]
kernel_height, kernel_width = merge_kernel_size
new_height, new_width = h // kernel_height, w // kernel_width
reshaped_seq = seq.view(
t, new_height, kernel_height, new_width, kernel_width, d_model
)
reshaped_seq = reshaped_seq.permute(0, 1, 3, 2, 4, 5).contiguous()
reshaped_seq = reshaped_seq.squeeze(0) if t == 1 else reshaped_seq.mean(dim=0)
outputs.append(
reshaped_seq.view(
new_height * new_width, kernel_height * kernel_width, d_model
)
)
pre_sum += t * h * w
return outputs
class MoonVitVLProjector(nn.Module): class MoonVitVLProjector(nn.Module):
def __init__( def __init__(
+10 -4
View File
@@ -544,6 +544,12 @@ def run_dp_sharded_vision_model(
return vision_embeddings return vision_embeddings
def concat_or_single(tensors: Sequence[torch.Tensor], dim: int = 0) -> torch.Tensor:
"""Concatenate multiple tensors without copying a singleton input."""
return tensors[0] if len(tensors) == 1 else torch.cat(tensors, dim=dim)
# Adapted from https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/models/vision.py # Adapted from https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/models/vision.py
def run_dp_sharded_mrope_vision_model( def run_dp_sharded_mrope_vision_model(
vision_model: torch.nn.Module, vision_model: torch.nn.Module,
@@ -617,12 +623,12 @@ def run_dp_sharded_mrope_vision_model(
# already concatenates these tensors before returning, so keep the # already concatenates these tensors before returning, so keep the
# TP=1 DP-encoder path on the same projector-facing contract. # TP=1 DP-encoder path on the same projector-facing contract.
if isinstance(image_embeds, list): if isinstance(image_embeds, list):
return torch.cat(image_embeds, dim=0) return concat_or_single(image_embeds, dim=0)
return image_embeds return image_embeds
if rope_type == "rope_2d_packed": if rope_type == "rope_2d_packed":
image_embeds = vision_model(pixel_values, grid_thw) image_embeds = vision_model(pixel_values, grid_thw)
if isinstance(image_embeds, list): if isinstance(image_embeds, list):
return torch.cat(image_embeds, dim=0) return concat_or_single(image_embeds, dim=0)
return image_embeds return image_embeds
return vision_model(pixel_values, grid_thw=grid_thw) return vision_model(pixel_values, grid_thw=grid_thw)
@@ -719,7 +725,7 @@ def run_dp_sharded_mrope_vision_model(
else: else:
image_embeds_local = vision_model(pixel_values_local, local_grid_thw) image_embeds_local = vision_model(pixel_values_local, local_grid_thw)
if isinstance(image_embeds_local, list): if isinstance(image_embeds_local, list):
image_embeds_local = torch.cat(image_embeds_local, dim=0) image_embeds_local = concat_or_single(image_embeds_local, dim=0)
else: else:
out_dim = getattr(vision_model.config, "hidden_size", None) out_dim = getattr(vision_model.config, "hidden_size", None)
image_embeds_local = torch.empty( image_embeds_local = torch.empty(
@@ -734,7 +740,7 @@ def run_dp_sharded_mrope_vision_model(
pixel_values_local, torch.tensor(local_grid_thw_list) pixel_values_local, torch.tensor(local_grid_thw_list)
) )
if isinstance(image_embeds_local, list): if isinstance(image_embeds_local, list):
image_embeds_local = torch.cat(image_embeds_local, dim=0) image_embeds_local = concat_or_single(image_embeds_local, dim=0)
else: else:
# Handle empty case # Handle empty case
out_dim = getattr(vision_model, "out_hidden_size", None) out_dim = getattr(vision_model, "out_hidden_size", None)
@@ -181,6 +181,9 @@ class BaseMultimodalProcessor(ABC):
gpu_image_decode = True # Enable GPU decoding by default gpu_image_decode = True # Enable GPU decoding by default
prefer_tokenized_input = False prefer_tokenized_input = False
precompute_hash_before_cpu_transfer = False precompute_hash_before_cpu_transfer = False
# Set by processors that already build input_ids from the request's own
# tokens, so the retokenize-avoidance rebuild below has nothing to add.
preserve_processor_input_ids = False
auto_mm_processor_worker_num = 1 auto_mm_processor_worker_num = 1
auto_mm_io_worker_num = 4 auto_mm_io_worker_num = 4
supports_mm_processor_concurrency = False supports_mm_processor_concurrency = False
@@ -1510,6 +1513,7 @@ class BaseMultimodalProcessor(ABC):
# Drift happens when Retokenization is not identity: Decode(X) => String => Re-tokenize => Y, X != Y. # Drift happens when Retokenization is not identity: Decode(X) => String => Re-tokenize => Y, X != Y.
if ( if (
envs.SGLANG_MM_AVOID_RETOKENIZE.get() envs.SGLANG_MM_AVOID_RETOKENIZE.get()
and not self.preserve_processor_input_ids
and base_output.input_ids is not None and base_output.input_ids is not None
and input_ids is not None and input_ids is not None
and raw_images and raw_images
@@ -3,7 +3,7 @@
Shared by KimiVLImageProcessor and KimiK2_5VLImageProcessor. Shared by KimiVLImageProcessor and KimiK2_5VLImageProcessor.
""" """
from typing import Union from typing import Optional, Union
import numpy as np import numpy as np
import torch import torch
@@ -38,6 +38,22 @@ class KimiGridMMDataMixin:
for image in images for image in images
] ]
@staticmethod
def count_image_placeholders(input_ids, image_token_id: int) -> Optional[int]:
"""Structural image tokens in a pre-tokenized prompt, None if it is text."""
if not isinstance(input_ids, (list, torch.Tensor)):
return None
token_ids = np.asarray(
(
input_ids.detach().flatten().cpu()
if isinstance(input_ids, torch.Tensor)
else input_ids
),
dtype=np.int64,
)
return int(np.count_nonzero(token_ids == image_token_id))
def _num_image_tokens_from_grid( def _num_image_tokens_from_grid(
self, grid_thw: Union[torch.Tensor, np.ndarray, list, tuple] self, grid_thw: Union[torch.Tensor, np.ndarray, list, tuple]
) -> int: ) -> int:
@@ -8,6 +8,7 @@ import torch
import torch.nn.functional as F import torch.nn.functional as F
from PIL import Image from PIL import Image
from sglang.kernels.ops.mm.process import normalize_and_patchify
from sglang.srt.managers.schedule_batch import ( from sglang.srt.managers.schedule_batch import (
MultimodalProcessorOutput, MultimodalProcessorOutput,
) )
@@ -79,6 +80,34 @@ def _get_image_dimensions(image: Union[torch.Tensor, Image.Image]) -> tuple[int,
return image.size # PIL returns (width, height) return image.size # PIL returns (width, height)
def _expand_image_token_ids(
input_ids: Union[List[int], torch.Tensor],
image_token_id: int,
image_token_counts: List[int],
) -> torch.Tensor:
"""Expand one placeholder per image without tokenizing the media string again.
Same rebuild as ``BaseMultimodalProcessor._expand_input_ids``, but staying in
the array domain skips a list round trip on the way to the output tensor.
test_kimi_k25.py pins the two together.
"""
if isinstance(input_ids, torch.Tensor):
input_ids = input_ids.detach().flatten().cpu().numpy()
input_ids = np.asarray(input_ids, dtype=np.int64)
placeholder_mask = input_ids == image_token_id
placeholder_count = np.count_nonzero(placeholder_mask)
if placeholder_count != len(image_token_counts):
raise ValueError(
f"Expected {len(image_token_counts)} image placeholder token(s), "
f"found {placeholder_count}."
)
repeats = np.ones(input_ids.shape, dtype=np.int64)
repeats[placeholder_mask] = image_token_counts
return torch.from_numpy(np.repeat(input_ids, repeats)).unsqueeze(0)
def _pil_to_cuda_chw(image: Image.Image) -> torch.Tensor: def _pil_to_cuda_chw(image: Image.Image) -> torch.Tensor:
"""Convert PIL Image to (C, H, W) uint8 CUDA tensor.""" """Convert PIL Image to (C, H, W) uint8 CUDA tensor."""
arr = np.asarray(image.convert("RGB")) arr = np.asarray(image.convert("RGB"))
@@ -95,10 +124,16 @@ def _ensure_chw_rgb(image: torch.Tensor) -> torch.Tensor:
consistent channel dimension. Normalize every tensor to 3 channels. consistent channel dimension. Normalize every tensor to 3 channels.
Also move the tensor to the GPU (matching _pil_to_cuda_chw) so a CPU Also move the tensor to the GPU (matching _pil_to_cuda_chw) so a CPU
input does not trip a device mismatch against the CUDA image_mean / input does not trip a device mismatch against the CUDA normalization
image_std_inv normalization constants downstream. No-op if already on constants downstream. No-op if already on the device.
the device.
""" """
if image.dtype != torch.uint8:
# Raw 0-255 is load-bearing downstream: the resize rounds to integers
# and the normalization folds in a 1/255 scale, so a normalized float
# image would collapse to 0/1 and then be rescaled.
raise ValueError(
f"Kimi GPU preprocessing expects raw uint8 pixels, got {image.dtype}"
)
image = image.cuda() image = image.cuda()
if image.dim() == 2: # (H, W) grayscale -> (1, H, W) if image.dim() == 2: # (H, W) grayscale -> (1, H, W)
image = image.unsqueeze(0) image = image.unsqueeze(0)
@@ -111,39 +146,67 @@ def _ensure_chw_rgb(image: torch.Tensor) -> torch.Tensor:
return image[:3] return image[:3]
def _resize_bicubic_if_needed(
image: torch.Tensor, target_height: int, target_width: int
) -> torch.Tensor:
"""Track the checkpoint processor's ``PIL.Image.resize(..., BICUBIC)``.
NaViT only ever downscales, and PIL's bicubic widens its kernel support by
the scale factor -- it always antialiases, which ``F.interpolate`` only does
under ``antialias=True``. PIL also returns uint8, so round and clip back to
integer pixels; the bicubic overshoot would otherwise survive normalization.
Close but not exact: PIL evaluates uint8 resizes in fixed point, so a few
8-bit levels of residual remain -- against PIL's float path we agree to 5e-3,
i.e. the kernel matches and only the arithmetic differs.
"""
image = image.float()
if image.shape[-2:] == (target_height, target_width):
return image
return (
F.interpolate(
image,
size=(target_height, target_width),
mode="bicubic",
align_corners=False,
antialias=True,
)
.round_()
.clamp_(0.0, 255.0)
)
def _grid_thw_from_resize_config(config: dict, patch_size: int) -> tuple[int, int, int]:
height = config["new_height"] + config["pad_height"]
width = config["new_width"] + config["pad_width"]
return 1, height // patch_size, width // patch_size
def _to_cuda_chw(image: Union[torch.Tensor, Image.Image]) -> torch.Tensor:
if isinstance(image, Image.Image):
return _pil_to_cuda_chw(image)
return _ensure_chw_rgb(image)
def _process_single_image( def _process_single_image(
image: Union[torch.Tensor, Image.Image], image: Union[torch.Tensor, Image.Image],
config: dict, config: dict,
image_mean: torch.Tensor, image_scale: torch.Tensor,
image_std_inv: torch.Tensor, image_bias: torch.Tensor,
patch_size: int, patch_size: int,
) -> tuple[torch.Tensor, torch.Tensor]: ) -> torch.Tensor:
"""Process a single image on GPU: resize -> pad -> normalize -> patchify.""" """Process a single image on GPU: resize -> pad -> normalize -> patchify."""
if isinstance(image, Image.Image): image = _to_cuda_chw(image)
image = _pil_to_cuda_chw(image)
else:
image = _ensure_chw_rgb(image)
new_h, new_w = config["new_height"], config["new_width"] new_h, new_w = config["new_height"], config["new_width"]
pad_h, pad_w = config["pad_height"], config["pad_width"] padded_h = new_h + config["pad_height"]
padded_w = new_w + config["pad_width"]
x = image.unsqueeze(0).float() x = _resize_bicubic_if_needed(image.unsqueeze(0), new_h, new_w)
x = F.interpolate(x, size=(new_h, new_w), mode="bicubic", align_corners=False)
if pad_h > 0 or pad_w > 0: return normalize_and_patchify(
x = F.pad(x, (0, pad_w, 0, pad_h), value=0.0) x, image_scale, image_bias, patch_size, padded_h, padded_w
).squeeze(0)
x = x / 255.0
x = (x - image_mean) * image_std_inv
_, C, H, W = x.shape
T = 1
gh, gw = H // patch_size, W // patch_size
x = x.view(T, C, gh, patch_size, gw, patch_size)
x = x.permute(0, 2, 4, 1, 3, 5).reshape(-1, C, patch_size, patch_size)
grid_thw = torch.tensor([T, gh, gw], dtype=torch.int64, device=x.device)
return x, grid_thw
def _resize_images_by_source_shape( def _resize_images_by_source_shape(
@@ -166,22 +229,14 @@ def _resize_images_by_source_shape(
for images in by_source_shape.values(): for images in by_source_shape.values():
if len(images) == 1: if len(images) == 1:
index, image = images[0] index, image = images[0]
resized_by_index[index] = F.interpolate( resized_by_index[index] = _resize_bicubic_if_needed(
image.unsqueeze(0).float(), image.unsqueeze(0), target_height, target_width
size=(target_height, target_width),
mode="bicubic",
align_corners=False,
) )
continue continue
source_batch = torch.cat( source_batch = torch.cat([image.unsqueeze(0) for _, image in images], dim=0)
[image.unsqueeze(0) for _, image in images], dim=0 resized_batch = _resize_bicubic_if_needed(
).float() source_batch, target_height, target_width
resized_batch = F.interpolate(
source_batch,
size=(target_height, target_width),
mode="bicubic",
align_corners=False,
) )
for local_index, (index, _) in enumerate(images): for local_index, (index, _) in enumerate(images):
resized_by_index[index] = resized_batch[local_index : local_index + 1] resized_by_index[index] = resized_batch[local_index : local_index + 1]
@@ -192,8 +247,8 @@ def _resize_images_by_source_shape(
def _gpu_preprocess_images( def _gpu_preprocess_images(
images: list[Union[torch.Tensor, Image.Image]], images: list[Union[torch.Tensor, Image.Image]],
resize_configs: list[dict], resize_configs: list[dict],
image_mean: torch.Tensor, image_scale: torch.Tensor,
image_std_inv: torch.Tensor, image_bias: torch.Tensor,
patch_size: int, patch_size: int,
) -> tuple[torch.Tensor, torch.Tensor]: ) -> tuple[torch.Tensor, torch.Tensor]:
"""GPU preprocessing pipeline for a batch of images. """GPU preprocessing pipeline for a batch of images.
@@ -202,10 +257,10 @@ def _gpu_preprocess_images(
""" """
n = len(images) n = len(images)
if n == 0: if n == 0:
device = image_mean.device device = image_scale.device
return ( return (
torch.empty(0, 3, patch_size, patch_size, device=device), torch.empty(0, 3, patch_size, patch_size, device=device),
torch.empty(0, 3, dtype=torch.int64, device=device), torch.empty(0, 3, dtype=torch.int64),
) )
groups = defaultdict(list) groups = defaultdict(list)
@@ -222,19 +277,13 @@ def _gpu_preprocess_images(
for (target_h, target_w, padded_h, padded_w), group in groups.items(): for (target_h, target_w, padded_h, padded_w), group in groups.items():
if len(group) == 1: if len(group) == 1:
idx, image, config = group[0] idx, image, config = group[0]
patches, grid = _process_single_image( patches = _process_single_image(
image, config, image_mean, image_std_inv, patch_size image, config, image_scale, image_bias, patch_size
) )
all_patches[idx] = patches all_patches[idx] = patches
all_grids[idx] = grid all_grids[idx] = _grid_thw_from_resize_config(config, patch_size)
else: else:
indexed_images = [] indexed_images = [(idx, _to_cuda_chw(image)) for idx, image, _ in group]
for idx, image, _ in group:
if isinstance(image, Image.Image):
image = _pil_to_cuda_chw(image)
else:
image = _ensure_chw_rgb(image)
indexed_images.append((idx, image))
# One NaViT target group can include several original resolutions. # One NaViT target group can include several original resolutions.
# Batch only source-compatible images, which removes redundant # Batch only source-compatible images, which removes redundant
@@ -245,29 +294,24 @@ def _gpu_preprocess_images(
dim=0, dim=0,
) )
pad_h = padded_h - target_h
pad_w = padded_w - target_w
if pad_h > 0 or pad_w > 0:
batch = F.pad(batch, (0, pad_w, 0, pad_h), value=0.0)
batch = batch / 255.0
batch = (batch - image_mean) * image_std_inv
B, C, H, W = batch.shape
T = 1 T = 1
gh, gw = H // patch_size, W // patch_size gh, gw = padded_h // patch_size, padded_w // patch_size
batch = batch.view(B, C, gh, patch_size, gw, patch_size) batch = normalize_and_patchify(
batch = batch.permute(0, 2, 4, 1, 3, 5).reshape( batch,
B, -1, C, patch_size, patch_size image_scale,
image_bias,
patch_size,
padded_h,
padded_w,
) )
grid = torch.tensor([T, gh, gw], dtype=torch.int64, device=batch.device) grid = (T, gh, gw)
for i, (idx, _, _) in enumerate(group): for i, (idx, _, _) in enumerate(group):
all_patches[idx] = batch[i] all_patches[idx] = batch[i]
all_grids[idx] = grid all_grids[idx] = grid
pixel_values = torch.cat(all_patches, dim=0) pixel_values = torch.cat(all_patches, dim=0)
grid_thws = torch.stack(all_grids, dim=0) grid_thws = torch.tensor(all_grids, dtype=torch.int64)
return pixel_values, grid_thws return pixel_values, grid_thws
@@ -290,6 +334,7 @@ class KimiGPUProcessorWrapper:
self, self,
hf_processor, hf_processor,
image_token, image_token,
image_token_id,
patch_size, patch_size,
merge_kernel_size, merge_kernel_size,
in_patch_limit, in_patch_limit,
@@ -300,6 +345,7 @@ class KimiGPUProcessorWrapper:
): ):
self._hf_processor = hf_processor self._hf_processor = hf_processor
self._image_token = image_token self._image_token = image_token
self._image_token_id = image_token_id
self._patch_size = patch_size self._patch_size = patch_size
self._merge_kernel_size = merge_kernel_size self._merge_kernel_size = merge_kernel_size
self._in_patch_limit = in_patch_limit self._in_patch_limit = in_patch_limit
@@ -320,12 +366,30 @@ class KimiGPUProcessorWrapper:
def __call__(self, text=None, images=None, **kwargs): def __call__(self, text=None, images=None, **kwargs):
# process_mm_data passes images via kwargs["images"] # process_mm_data passes images via kwargs["images"]
images = images or kwargs.pop("images", None) images = images or kwargs.pop("images", None)
original_input_ids = kwargs.pop("sglang_original_input_ids", None)
if images and torch.cuda.is_available(): if images and torch.cuda.is_available():
return self._gpu_call(text, images) return self._gpu_call(text, images, original_input_ids)
return self._cpu_call(text, images, **kwargs) return self._cpu_call(text, images, original_input_ids, **kwargs)
def _gpu_call(self, text, images): def _prepare_input_ids(self, input_text, resize_configs, original_input_ids):
if original_input_ids is not None:
return _expand_image_token_ids(
original_input_ids,
self._image_token_id,
[config["num_tokens"] for config in resize_configs],
)
parts = input_text.split(self._image_token)
result = [parts[0]]
for config, part in zip(resize_configs, parts[1:]):
result.append(self._image_token * config["num_tokens"] + part)
expanded_text = "".join(result)
return self._hf_processor.tokenizer(expanded_text, return_tensors="pt")[
"input_ids"
]
def _gpu_call(self, text, images, original_input_ids=None):
"""Bypass HF KimiK25VisionProcessor.preprocess entirely -- use GPU ops.""" """Bypass HF KimiK25VisionProcessor.preprocess entirely -- use GPU ops."""
input_text = text[0] if isinstance(text, list) else text input_text = text[0] if isinstance(text, list) else text
@@ -345,44 +409,44 @@ class KimiGPUProcessorWrapper:
) )
) )
# 2. Expand image tokens # 2. Reuse the request's tokenization when available: expanding the
parts = input_text.split(self._image_token) # placeholders is exact, and skips tokenizing thousands of repeated
result = [parts[0]] # ``<|media_pad|>`` strings.
for config, part in zip(resize_configs, parts[1:]): input_ids = self._prepare_input_ids(
result.append(self._image_token * config["num_tokens"] + part) input_text, resize_configs, original_input_ids
input_text = "".join(result)
# 3. Tokenize
text_inputs = self._hf_processor.tokenizer(input_text, return_tensors="pt")
# 4. GPU image preprocessing
image_mean, image_std_inv = self._get_gpu_norm_tensors()
pixel_values, grid_thws = _gpu_preprocess_images(
images, resize_configs, image_mean, image_std_inv, self._patch_size
) )
grid_thws = grid_thws.cpu() # 3. GPU image preprocessing
image_scale, image_bias = self._get_gpu_norm_tensors()
pixel_values, grid_thws = _gpu_preprocess_images(
images, resize_configs, image_scale, image_bias, self._patch_size
)
return { return {
"input_ids": text_inputs["input_ids"], "input_ids": input_ids,
"pixel_values": pixel_values, "pixel_values": pixel_values,
# Use SGL-standard key so get_new_expanded_mm_items() can split # Use SGL-standard key so get_new_expanded_mm_items() can split
# per-image for cache granularity (it looks up 'image_grid_thw'). # per-image for cache granularity (it looks up 'image_grid_thw').
"image_grid_thw": grid_thws, "image_grid_thw": grid_thws,
} }
def _cpu_call(self, text, images, **kwargs): def _cpu_call(self, text, images, original_input_ids=None, **kwargs):
"""Fallback: token expansion + medias kwarg -> original HF processor.""" """Fallback: token expansion + medias kwarg -> original HF processor."""
input_text = text[0] if isinstance(text, list) else text input_text = text[0] if isinstance(text, list) else text
if images: if images:
# Token expansion via media_tokens_calculator # Token expansion via media_tokens_calculator
image_token_counts = [
int(
self._hf_processor.media_processor.media_tokens_calculator(
{"type": "image", "image": image}
)
)
for image in images
]
parts = input_text.split(self._image_token) parts = input_text.split(self._image_token)
result = [parts[0]] result = [parts[0]]
for image, part in zip(images, parts[1:]): for num_tokens, part in zip(image_token_counts, parts[1:]):
num_tokens = self._hf_processor.media_processor.media_tokens_calculator(
{"type": "image", "image": image}
)
result.append(self._image_token * num_tokens + part) result.append(self._image_token * num_tokens + part)
input_text = "".join(result) input_text = "".join(result)
@@ -390,6 +454,12 @@ class KimiGPUProcessorWrapper:
kwargs["medias"] = [{"type": "image", "image": img} for img in images] kwargs["medias"] = [{"type": "image", "image": img} for img in images]
out = self._hf_processor(text=[input_text], **kwargs) out = self._hf_processor(text=[input_text], **kwargs)
if images and original_input_ids is not None:
# preserve_processor_input_ids turns off the base class rebuild, so
# this path has to keep the request's own tokens itself.
out["input_ids"] = _expand_image_token_ids(
original_input_ids, self._image_token_id, image_token_counts
)
grid_thws = out.pop("grid_thws", None) grid_thws = out.pop("grid_thws", None)
if grid_thws is not None: if grid_thws is not None:
out["image_grid_thw"] = grid_thws out["image_grid_thw"] = grid_thws
@@ -397,13 +467,17 @@ class KimiGPUProcessorWrapper:
def _get_gpu_norm_tensors(self, device="cuda"): def _get_gpu_norm_tensors(self, device="cuda"):
if self._gpu_norm_tensors is None: if self._gpu_norm_tensors is None:
image_mean = torch.tensor( image_scale = torch.tensor(
self._image_mean, device=device, dtype=torch.float32 [1.0 / (255.0 * std) for std in self._image_std],
device=device,
dtype=torch.float32,
).view(1, 3, 1, 1) ).view(1, 3, 1, 1)
image_std_inv = ( image_bias = torch.tensor(
1.0 / torch.tensor(self._image_std, device=device, dtype=torch.float32) [-mean / std for mean, std in zip(self._image_mean, self._image_std)],
device=device,
dtype=torch.float32,
).view(1, 3, 1, 1) ).view(1, 3, 1, 1)
self._gpu_norm_tensors = (image_mean, image_std_inv) self._gpu_norm_tensors = (image_scale, image_bias)
return self._gpu_norm_tensors return self._gpu_norm_tensors
@@ -418,23 +492,25 @@ class KimiK2_5VLImageProcessor(KimiGridMMDataMixin, SGLangBaseProcessor):
gpu_image_decode = True # nvJPEG for JPEG, PIL fallback for others gpu_image_decode = True # nvJPEG for JPEG, PIL fallback for others
prefer_tokenized_input = True prefer_tokenized_input = True
precompute_hash_before_cpu_transfer = True precompute_hash_before_cpu_transfer = True
# The GPU wrapper expands placeholders from the request's own token IDs.
preserve_processor_input_ids = True
auto_mm_processor_worker_num = 2
auto_mm_io_worker_num = 16
supports_mm_processor_concurrency = True
def __init__(self, hf_config, server_args, _processor, *args, **kwargs): def __init__(self, hf_config, server_args, _processor, *args, **kwargs):
super().__init__(hf_config, server_args, _processor, *args, **kwargs) mm_tokens = MultimodalSpecialTokens(
self.mm_tokens = MultimodalSpecialTokens(
image_token="<|media_pad|>", image_token="<|media_pad|>",
# TODO: could we convert in MultimodalSpecialTokens? # TODO: could we convert in MultimodalSpecialTokens?
image_token_id=hf_config.media_placeholder_token_id, image_token_id=hf_config.media_placeholder_token_id,
image_token_regex=re.compile(r"(?:<\|media_pad\|>)+"), image_token_regex=re.compile(r"(?:<\|media_pad\|>)+"),
).build(_processor) ).build(_processor)
# Extract media processing config from HF processor
media_proc_cfg = _processor.media_processor.media_proc_cfg media_proc_cfg = _processor.media_processor.media_proc_cfg
processor = KimiGPUProcessorWrapper(
# Replace with GPU-capable wrapper
self._processor = KimiGPUProcessorWrapper(
_processor, _processor,
image_token=self.mm_tokens.image_token, image_token=mm_tokens.image_token,
image_token_id=mm_tokens.image_token_id,
patch_size=media_proc_cfg["patch_size"], patch_size=media_proc_cfg["patch_size"],
merge_kernel_size=media_proc_cfg["merge_kernel_size"], merge_kernel_size=media_proc_cfg["merge_kernel_size"],
in_patch_limit=media_proc_cfg["in_patch_limit"], in_patch_limit=media_proc_cfg["in_patch_limit"],
@@ -443,6 +519,10 @@ class KimiK2_5VLImageProcessor(KimiGridMMDataMixin, SGLangBaseProcessor):
image_mean=media_proc_cfg["image_mean"], image_mean=media_proc_cfg["image_mean"],
image_std=media_proc_cfg["image_std"], image_std=media_proc_cfg["image_std"],
) )
# Initialize the executor from the final GPU wrapper. Cloning the raw
# HF processor here would silently bypass Kimi's GPU preprocessing.
super().__init__(hf_config, server_args, processor, *args, **kwargs)
self.mm_tokens = mm_tokens
async def process_mm_data_async( async def process_mm_data_async(
self, self,
@@ -452,14 +532,43 @@ class KimiK2_5VLImageProcessor(KimiGridMMDataMixin, SGLangBaseProcessor):
*args, *args,
**kwargs, **kwargs,
): ):
base_output = await self.load_mm_data( expected_image_count = len(image_data or [])
prompt=input_text, placeholder_count = self.count_image_placeholders(
image_data=image_data, input_text, self.mm_tokens.image_token_id
multimodal_tokens=self.mm_tokens,
) )
if placeholder_count is not None:
if placeholder_count != expected_image_count:
raise ValueError(
"Kimi image placeholders must map one-to-one to image data: "
f"expected {expected_image_count}, found {placeholder_count} token(s)"
)
base_output = await self.fast_load_mm_data(
prompt=input_text,
image_data=image_data,
multimodal_tokens=self.mm_tokens,
# fast_load_mm_data, unlike load_mm_data, does not derive
# input_ids from the prompt; without this the wrapper falls back
# to re-tokenizing the expanded string.
input_ids=input_text,
)
else:
base_output = await self.load_mm_data(
prompt=input_text,
image_data=image_data,
multimodal_tokens=self.mm_tokens,
)
# Only the text-scanning loader can come back with a different
# count; fast_load_mm_data fills one slot per image_data entry.
if len(base_output.images) != expected_image_count:
raise ValueError(
"Kimi image placeholders must map one-to-one to image data: "
f"expected {expected_image_count}, loaded {len(base_output.images)}"
)
mm_items, input_ids, _ = self.process_and_combine_mm_data( mm_items, input_ids, _ = await self.process_and_combine_mm_data_async(
base_output, self.mm_tokens base_output,
self.mm_tokens,
sglang_original_input_ids=base_output.input_ids,
) )
# K2.5/K2.7 encoder-DP assigns an image to exactly one TP rank. Keep # K2.5/K2.7 encoder-DP assigns an image to exactly one TP rank. Keep
@@ -39,6 +39,12 @@ class KimiVLImageProcessor(KimiGridMMDataMixin, SGLangBaseProcessor):
image_data=image_data, image_data=image_data,
multimodal_tokens=self.mm_tokens, multimodal_tokens=self.mm_tokens,
) )
expected_image_count = len(image_data or [])
if len(base_output.images) != expected_image_count:
raise ValueError(
"Kimi image placeholders must map one-to-one to image data: "
f"expected {expected_image_count}, loaded {len(base_output.images)}"
)
mm_items, input_ids, _ = self.process_and_combine_mm_data( mm_items, input_ids, _ = self.process_and_combine_mm_data(
base_output, self.mm_tokens base_output, self.mm_tokens
+145
View File
@@ -0,0 +1,145 @@
"""Prove the two GPU-only rewrites in the K2.5 port are equivalent to main.
1. normalize_and_patchify(scale/bias) == pad -> /255 -> (x-mean)*inv_std -> patchify
2. apply_fused_qk_complex_rope_inplace == the torch complex reference
"""
import torch
import torch.nn.functional as F
from sglang.kernels.ops.attention.vision_rope import (
apply_fused_qk_complex_rope_inplace,
prepare_fused_qk_complex_rope_inplace,
)
from sglang.kernels.ops.mm.process import normalize_and_patchify
MEAN = [0.5, 0.5, 0.5]
STD = [0.5, 0.5, 0.5]
ASYM_MEAN = [0.481, 0.457, 0.408]
ASYM_STD = [0.268, 0.261, 0.275]
def reference_preprocess(batch_u8, mean, std, patch_size, padded_h, padded_w):
"""Exactly what main does, in main's order."""
image_mean = torch.tensor(mean, device="cuda", dtype=torch.float32).view(1, 3, 1, 1)
image_std_inv = (1.0 / torch.tensor(std, device="cuda", dtype=torch.float32)).view(
1, 3, 1, 1
)
x = batch_u8.float()
pad_h = padded_h - x.shape[-2]
pad_w = padded_w - x.shape[-1]
if pad_h > 0 or pad_w > 0:
x = F.pad(x, (0, pad_w, 0, pad_h), value=0.0)
x = x / 255.0
x = (x - image_mean) * image_std_inv
B, C, H, W = x.shape
gh, gw = H // patch_size, W // patch_size
x = x.view(B, C, gh, patch_size, gw, patch_size)
return x.permute(0, 2, 4, 1, 3, 5).reshape(B, -1, C, patch_size, patch_size)
def check_patchify():
print("== normalize_and_patchify vs main's pad/normalize/patchify ==")
torch.manual_seed(0)
cases = [
# (H, W, padded_h, padded_w, patch, mean, std, label)
(32, 24, 32, 24, 8, MEAN, STD, "no padding, symmetric norm"),
(30, 22, 32, 24, 8, MEAN, STD, "padded, symmetric norm"),
(30, 22, 32, 24, 8, ASYM_MEAN, ASYM_STD, "padded, per-channel norm"),
(64, 64, 64, 64, 16, ASYM_MEAN, ASYM_STD, "large patch"),
]
ok = True
for h, w, ph, pw, patch, mean, std, label in cases:
raw = torch.randint(0, 256, (3, 3, h, w), dtype=torch.uint8, device="cuda")
ref = reference_preprocess(raw, mean, std, patch, ph, pw)
scale = torch.tensor(
[1.0 / (255.0 * s) for s in std], device="cuda", dtype=torch.float32
).view(1, 3, 1, 1)
bias = torch.tensor(
[-m / s for m, s in zip(mean, std)], device="cuda", dtype=torch.float32
).view(1, 3, 1, 1)
got = normalize_and_patchify(raw.float(), scale, bias, patch, ph, pw)
max_abs = (got - ref).abs().max().item()
# The padded rows must carry -mean/std, not zero.
pad_ok = True
if ph > h or pw > w:
pad_ok = torch.allclose(
got.flatten()[(got - ref).abs().argmax()],
ref.flatten()[(got - ref).abs().argmax()],
atol=1e-5,
)
good = max_abs < 1e-5 and pad_ok
ok &= good
print(f" {'PASS' if good else 'FAIL'} {label:32s} max|d|={max_abs:.3e}")
return ok
def check_padded_value_is_not_zero():
"""The old pipeline padded in raw space, so pad cells become -mean/std."""
print("== padded cells carry -mean/std, not 0 ==")
raw = torch.full((1, 3, 8, 8), 128, dtype=torch.uint8, device="cuda")
scale = torch.tensor(
[1.0 / (255.0 * s) for s in ASYM_STD], device="cuda", dtype=torch.float32
).view(1, 3, 1, 1)
bias = torch.tensor(
[-m / s for m, s in zip(ASYM_MEAN, ASYM_STD)],
device="cuda",
dtype=torch.float32,
).view(1, 3, 1, 1)
got = normalize_and_patchify(raw.float(), scale, bias, 8, 16, 16)
# patch index 1 is the (row 0, col 1) patch -- entirely padding.
pad_patch = got[0, 1]
expected = bias.view(3, 1, 1).expand(3, 8, 8)
good = torch.allclose(pad_patch, expected, atol=1e-6)
print(
f" {'PASS' if good else 'FAIL'} pad cell = {pad_patch[0, 0, 0].item():.6f}, "
f"expected -mean/std = {expected[0, 0, 0].item():.6f}"
)
return good
def reference_rope(xq, xk, freqs_cis):
freqs_cis = freqs_cis.unsqueeze(-2)
xq_ = torch.view_as_complex(xq.float().view(*xq.shape[:-1], -1, 2))
xk_ = torch.view_as_complex(xk.float().view(*xk.shape[:-1], -1, 2))
xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(-2)
xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(-2)
return xq_out.type_as(xq), xk_out.type_as(xk)
def check_rope():
print("== fused vision RoPE vs the torch complex reference ==")
torch.manual_seed(0)
ok = True
for dtype, tol in ((torch.bfloat16, 8e-3), (torch.float16, 2e-3)):
for tokens, heads, head_dim in ((1024, 16, 72), (4096, 8, 128), (37, 4, 64)):
xq = torch.randn(tokens, heads, head_dim, device="cuda", dtype=dtype)
xk = torch.randn(tokens, heads, head_dim, device="cuda", dtype=dtype)
angle = torch.randn(tokens, head_dim // 2, device="cuda")
freqs_cis = torch.polar(torch.ones_like(angle), angle)
ref_q, ref_k = reference_rope(xq, xk, freqs_cis)
prepared = prepare_fused_qk_complex_rope_inplace(freqs_cis)
got_q, got_k = apply_fused_qk_complex_rope_inplace(
xq.clone(), xk.clone(), prepared
)
dq = (got_q.float() - ref_q.float()).abs().max().item()
dk = (got_k.float() - ref_k.float()).abs().max().item()
good = dq < tol and dk < tol
ok &= good
print(
f" {'PASS' if good else 'FAIL'} {str(dtype):16s} "
f"t={tokens:5d} h={heads:2d} d={head_dim:3d} "
f"max|dq|={dq:.2e} max|dk|={dk:.2e}"
)
return ok
if __name__ == "__main__":
results = [check_patchify(), check_padded_value_is_not_zero(), check_rope()]
print()
print("ALL PASS" if all(results) else "SOME CHECKS FAILED")
raise SystemExit(0 if all(results) else 1)
@@ -20,6 +20,8 @@ from sglang.kernels.ops.attention.set_mla_kv_concat_q import (
from sglang.kernels.ops.attention.utils import concat_mla_absorb_q_general from sglang.kernels.ops.attention.utils import concat_mla_absorb_q_general
from sglang.kernels.ops.attention.vision_rope import ( from sglang.kernels.ops.attention.vision_rope import (
apply_fused_qk_complex_rope, apply_fused_qk_complex_rope,
apply_fused_qk_complex_rope_inplace,
prepare_fused_qk_complex_rope_inplace,
) )
from sglang.kernels.ops.elementwise import add3 from sglang.kernels.ops.elementwise import add3
from sglang.kernels.ops.gemm.tiny_gemm import ( from sglang.kernels.ops.gemm.tiny_gemm import (
@@ -427,6 +429,33 @@ class TestKimiK3PrerequisiteOps(CustomTestCase):
torch.testing.assert_close(actual_q, reference(q), rtol=0, atol=atol) torch.testing.assert_close(actual_q, reference(q), rtol=0, atol=atol)
torch.testing.assert_close(actual_k, reference(k), rtol=0, atol=atol) torch.testing.assert_close(actual_k, reference(k), rtol=0, atol=atol)
def test_vision_rope_inplace(self):
# VisionAttention hands the applier contiguous q/k, which is what the
# in-place kernel requires; mirror that rather than qkv.unbind views.
for dtype in (torch.bfloat16, torch.float16):
torch.manual_seed(4)
q = torch.randn(480, 12, 128, device="cuda", dtype=dtype)
k = torch.randn(480, 12, 128, device="cuda", dtype=dtype)
angles = torch.randn(480, 64, device="cuda")
freqs = torch.polar(torch.ones_like(angles), angles)
freqs_expanded = freqs.unsqueeze(-2)
def reference(x):
value = torch.view_as_complex(x.float().view(*x.shape[:-1], -1, 2))
return torch.view_as_real(value * freqs_expanded).flatten(-2).type_as(x)
expected_q, expected_k = reference(q), reference(k)
prepared = prepare_fused_qk_complex_rope_inplace(freqs)
actual_q, actual_k = apply_fused_qk_complex_rope_inplace(q, k, prepared)
atol = 2 * torch.finfo(dtype).eps
torch.testing.assert_close(actual_q, expected_q, rtol=0, atol=atol)
torch.testing.assert_close(actual_k, expected_k, rtol=0, atol=atol)
def test_vision_rope_inplace_rejects_non_complex_frequencies(self):
with self.assertRaises(ValueError):
prepare_fused_qk_complex_rope_inplace(torch.randn(8, 64, device="cuda"))
def test_normalize_and_patchify(self): def test_normalize_and_patchify(self):
torch.manual_seed(5) torch.manual_seed(5)
image = torch.randn(2, 3, 17, 19, device="cuda") image = torch.randn(2, 3, 17, 19, device="cuda")
+242 -7
View File
@@ -3,10 +3,12 @@
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import Mock, patch from unittest.mock import Mock, patch
import numpy as np
import pytest import pytest
import torch import torch
import torch.nn as nn import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from PIL import Image
from sglang.srt.managers.schedule_batch import ( from sglang.srt.managers.schedule_batch import (
Modality, Modality,
@@ -14,9 +16,19 @@ from sglang.srt.managers.schedule_batch import (
MultimodalInputs, MultimodalInputs,
MultimodalProcessorOutput, MultimodalProcessorOutput,
) )
from sglang.srt.models.kimi_k25 import KimiK25ForConditionalGeneration from sglang.srt.models.kimi_k25 import (
KimiK25ForConditionalGeneration,
mm_projection_auto,
)
from sglang.srt.models.kimi_vl_moonvit import tpool_patch_merger
from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model
from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor
from sglang.srt.multimodal.processors.kimi_common import KimiGridMMDataMixin
from sglang.srt.multimodal.processors.kimi_k25 import ( from sglang.srt.multimodal.processors.kimi_k25 import (
KimiGPUProcessorWrapper,
_ensure_chw_rgb,
_expand_image_token_ids,
_resize_bicubic_if_needed,
_resize_images_by_source_shape, _resize_images_by_source_shape,
) )
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel
@@ -63,14 +75,12 @@ def _image_item(feature, grid_thw):
def test_kimi_gpu_preprocess_batches_only_source_compatible_images(): def test_kimi_gpu_preprocess_batches_only_source_compatible_images():
torch.manual_seed(0) torch.manual_seed(0)
indexed_images = [ indexed_images = [
(0, torch.randn(3, 32, 24)), (0, torch.randint(0, 256, (3, 32, 24), dtype=torch.uint8)),
(1, torch.randn(3, 32, 24)), (1, torch.randint(0, 256, (3, 32, 24), dtype=torch.uint8)),
(2, torch.randn(3, 28, 20)), (2, torch.randint(0, 256, (3, 28, 20), dtype=torch.uint8)),
] ]
expected = [ expected = [
F.interpolate( _resize_bicubic_if_needed(image.unsqueeze(0), 16, 12)
image.unsqueeze(0), size=(16, 12), mode="bicubic", align_corners=False
)
for _, image in indexed_images for _, image in indexed_images
] ]
real_interpolate = F.interpolate real_interpolate = F.interpolate
@@ -92,6 +102,209 @@ def test_kimi_gpu_preprocess_batches_only_source_compatible_images():
torch.testing.assert_close(result, reference) torch.testing.assert_close(result, reference)
def test_kimi_resize_tracks_the_checkpoint_processors_pil_bicubic():
# Plain F.interpolate skips PIL's implicit antialiasing on downscale and
# drifts far outside 8-bit rounding; photo-like content, not pure noise.
rng = np.random.default_rng(0)
yy, xx = np.mgrid[0:512, 0:512].astype(np.float32)
plane = np.clip(
128
+ 90 * np.sin(xx / 40) * np.cos(yy / 55)
+ 40 * ((xx // 37 + yy // 41) % 2)
+ rng.normal(0, 6, (512, 512)),
0,
255,
)
array = np.stack([plane, np.roll(plane, 7, 0), np.roll(plane, 13, 1)], -1).astype(
np.uint8
)
pil = torch.from_numpy(
np.asarray(Image.fromarray(array).resize((252, 252), Image.BICUBIC)).astype(
np.float32
)
).permute(2, 0, 1)
source = torch.from_numpy(array).permute(2, 0, 1).unsqueeze(0)
resized = _resize_bicubic_if_needed(source, 252, 252)
assert resized.shape == (1, 3, 252, 252)
torch.testing.assert_close(resized, resized.round())
assert resized.min() >= 0.0 and resized.max() <= 255.0
# Within a couple of 8-bit levels of PIL; the non-antialiased resize is off
# by an order of magnitude more, which is the regression this guards.
assert (resized[0] - pil).abs().max() <= 4.0
naive = F.interpolate(
source.float(), size=(252, 252), mode="bicubic", align_corners=False
)
assert (naive[0] - pil).abs().max() > 20.0
def test_kimi_resize_is_a_dtype_only_cast_when_already_at_target():
image = torch.randint(0, 256, (1, 3, 16, 12), dtype=torch.uint8)
resized = _resize_bicubic_if_needed(image, 16, 12)
assert resized.dtype == torch.float32
torch.testing.assert_close(resized, image.float())
def test_kimi_expands_one_placeholder_per_image_from_existing_ids():
# 7 is the placeholder; the two images claim 3 and 2 tokens.
input_ids = [1, 7, 2, 7, 3]
expanded = _expand_image_token_ids(
input_ids, image_token_id=7, image_token_counts=[3, 2]
)
assert expanded.tolist() == [[1, 7, 7, 7, 2, 7, 7, 3]]
def test_kimi_expansion_rejects_a_placeholder_count_mismatch():
with pytest.raises(ValueError, match="placeholder"):
_expand_image_token_ids([1, 7, 2], image_token_id=7, image_token_counts=[3, 2])
def test_kimi_expansion_matches_the_base_retokenize_avoidance_rebuild():
# preserve_processor_input_ids skips the base rebuild, which is only safe
# while both produce the same sequence. Reference is the original loop.
def reference(original_ids, counts, placeholder):
rebuilt, next_image = [], 0
for token_id in original_ids:
if token_id == placeholder:
rebuilt.extend([placeholder] * counts[next_image])
next_image += 1
else:
rebuilt.append(token_id)
return rebuilt
rng = np.random.default_rng(0)
for n_images in (1, 3, 8):
# Placeholder 7 is below the random range, so only the inserted
# positions count as placeholders.
ids = rng.integers(100, 5000, 400).tolist()
for slot in range(n_images):
ids.insert(slot * 37 + 5, 7)
counts = rng.integers(1, 400, n_images).tolist()
expected = reference(ids, counts, 7)
assert BaseMultimodalProcessor._expand_input_ids(ids, counts, 7) == expected
wrapper = _expand_image_token_ids(
ids, image_token_id=7, image_token_counts=counts
)
assert wrapper.flatten().tolist() == expected
def test_kimi_cpu_fallback_keeps_the_request_tokens():
# preserve_processor_input_ids disables the base rebuild on every path.
hf_processor = Mock()
hf_processor.media_processor.media_tokens_calculator = Mock(return_value=3)
hf_processor.return_value = {"input_ids": torch.tensor([[99, 99, 99]])}
wrapper = KimiGPUProcessorWrapper.__new__(KimiGPUProcessorWrapper)
wrapper._hf_processor = hf_processor
wrapper._image_token = "<|media_pad|>"
wrapper._image_token_id = 7
out = wrapper._cpu_call(
"a<|media_pad|>b", ["img"], original_input_ids=[1, 7, 2], medias=None
)
# Not the [99, 99, 99] the HF processor returned.
assert out["input_ids"].flatten().tolist() == [1, 7, 7, 7, 2]
def test_kimi_cpu_fallback_falls_back_to_the_hf_tokens_without_request_ids():
hf_processor = Mock()
hf_processor.media_processor.media_tokens_calculator = Mock(return_value=3)
hf_processor.return_value = {"input_ids": torch.tensor([[99, 99, 99]])}
wrapper = KimiGPUProcessorWrapper.__new__(KimiGPUProcessorWrapper)
wrapper._hf_processor = hf_processor
wrapper._image_token = "<|media_pad|>"
wrapper._image_token_id = 7
out = wrapper._cpu_call("a<|media_pad|>b", ["img"], medias=None)
assert out["input_ids"].flatten().tolist() == [99, 99, 99]
def test_kimi_refuses_already_normalized_float_pixels():
with pytest.raises(ValueError, match="uint8"):
_ensure_chw_rgb(torch.rand(3, 8, 8))
def test_kimi_placeholder_count_only_reads_real_token_ids():
count = KimiGridMMDataMixin.count_image_placeholders
assert count([1, 7, 2, 7], 7) == 2
assert count(torch.tensor([[1, 7, 2]]), 7) == 1
assert count([1, 2, 3], 7) == 0
# A prompt string carries no token IDs, so the caller must not take the
# tokenized fast path.
assert count("<|media_pad|>", 7) is None
def test_kimi_single_frame_pool_matches_the_temporal_mean():
torch.manual_seed(0)
x = torch.randn(1 * 4 * 4, 8)
grid_thws = torch.tensor([[1, 4, 4]])
(merged,) = tpool_patch_merger(x, grid_thws)
# t == 1 skips the mean; it must stay bit-identical to averaging one frame.
reference = (
x.view(1, 2, 2, 2, 2, 8).permute(0, 1, 3, 2, 4, 5).contiguous().mean(dim=0)
)
assert torch.equal(merged, reference.view(4, 4, 8))
def test_kimi_multi_frame_pool_still_averages_across_frames():
torch.manual_seed(0)
x = torch.randn(3 * 4 * 4, 8)
grid_thws = torch.tensor([[3, 4, 4]])
(merged,) = tpool_patch_merger(x, grid_thws)
reference = (
x.view(3, 2, 2, 2, 2, 8).permute(0, 1, 3, 2, 4, 5).contiguous().mean(dim=0)
)
assert merged.shape == (4, 4, 8)
torch.testing.assert_close(merged, reference.view(4, 4, 8))
class _IdentityProjector(nn.Module):
"""Stands in for K2VLMultiModalProjector, which is never None in production."""
def __init__(self):
super().__init__()
self.seen = None
def forward(self, x):
self.seen = x
return x
def test_kimi_projection_returns_one_flattened_feature_tensor():
torch.manual_seed(0)
per_image = [torch.randn(4, 2, 8), torch.randn(6, 2, 8)]
packed = mm_projection_auto(_IdentityProjector(), per_image)
assert packed.shape == (20, 8)
torch.testing.assert_close(packed, torch.cat(per_image, dim=0).reshape(-1, 8))
def test_kimi_projection_does_not_copy_a_single_image():
single = torch.randn(4, 2, 8)
projector = _IdentityProjector()
packed = mm_projection_auto(projector, [single])
# The projector must receive the tensor itself, not a one-element cat of it.
assert projector.seen.data_ptr() == single.data_ptr()
assert packed.data_ptr() == single.data_ptr()
def test_dp_helper_supports_moonvit3d_packed_embeddings_on_tp1(): def test_dp_helper_supports_moonvit3d_packed_embeddings_on_tp1():
tower = _MoonViT3dTower() tower = _MoonViT3dTower()
pixel_values = torch.randn(4, 2) pixel_values = torch.randn(4, 2)
@@ -248,6 +461,28 @@ def test_kimi_k25_encoder_dp_selects_packed_moonvit_contract():
assert callable(run_dp.call_args.kwargs["load_local_pixel_values"]) assert callable(run_dp.call_args.kwargs["load_local_pixel_values"])
def test_kimi_non_dp_keeps_grid_thws_on_the_host():
model = KimiK25ForConditionalGeneration.__new__(KimiK25ForConditionalGeneration)
nn.Module.__init__(model)
model.use_data_parallel = False
model.vision_tower = _MoonViT3dTower()
# Not the host, so a stray .to(tower.device) shows up without a GPU.
model.vision_tower.device = torch.device("meta")
model.mm_projector = _IdentityProjector()
items = [_image_item(torch.randn(4, 2), [[1, 2, 2]])]
with get_parallel().override(
tp_size=1, tp_rank=0, attn_tp_size=1, attn_tp_rank=0
), patch(
"sglang.srt.models.kimi_k25.get_server_args",
return_value=SimpleNamespace(tp_size=1),
):
model.get_image_feature(items)
# A device copy would cost one sync per .tolist() inside MoonViT3d.
assert model.vision_tower.grid_thws.device.type == "cpu"
def test_kimi_lazy_ipc_feature_skips_scheduler_reconstruction(): def test_kimi_lazy_ipc_feature_skips_scheduler_reconstruction():
proxy = CudaIpcTensorTransportProxy.__new__(CudaIpcTensorTransportProxy) proxy = CudaIpcTensorTransportProxy.__new__(CudaIpcTensorTransportProxy)
proxy.reconstruct_on_target_device = Mock() proxy.reconstruct_on_target_device = Mock()