From c20e99bd22d8ccfc6d76f9c4c2acc45782f6c4e7 Mon Sep 17 00:00:00 2001 From: Mick Date: Mon, 10 Aug 2026 09:42:52 +0800 Subject: [PATCH] fix(vlm): preserve Kimi-K3 GPU JPEG accuracy (#34163) --- docker/kimi_k3/kimi_k3_cu12.Dockerfile | 7 ++ docker/kimi_k3/kimi_k3_cu13.Dockerfile | 7 ++ .../srt/disaggregation/encode_server.py | 7 +- .../srt/multimodal/processors/kimi_k3.py | 5 +- python/sglang/srt/utils/common.py | 34 +++++-- python/sglang/srt/utils/nvjpeg_decoder.py | 86 ++++++++++++++++ .../test_kimi_k3_encoder_mode.py | 23 ++++- .../test_base_processor_image_decode.py | 98 +++++++++++++++++++ 8 files changed, 258 insertions(+), 9 deletions(-) create mode 100644 python/sglang/srt/utils/nvjpeg_decoder.py diff --git a/docker/kimi_k3/kimi_k3_cu12.Dockerfile b/docker/kimi_k3/kimi_k3_cu12.Dockerfile index 65e8ff864..5b57ad607 100644 --- a/docker/kimi_k3/kimi_k3_cu12.Dockerfile +++ b/docker/kimi_k3/kimi_k3_cu12.Dockerfile @@ -26,6 +26,7 @@ FROM lmsysorg/sglang:v0.5.16-cu129 AS base ARG SGL_DEEP_GEMM_VERSION="0.1.5.post2" +ARG NVIMGCODEC_VERSION="0.9.0.20" # Current Kimi-K3 source auto-discovers and builds its PyO3 extensions. ARG RUST_VERSION="1.90.0" @@ -89,6 +90,12 @@ RUN TORCH_CUDA_ARCH_LIST="${TORCH_CUDA_ARCH_LIST}" \ RUN python3 -m pip install --no-deps --force-reinstall \ "https://github.com/sgl-project/whl/releases/download/v${SGL_DEEP_GEMM_VERSION}/sgl_deep_gemm-${SGL_DEEP_GEMM_VERSION}+cu129-py3-none-manylinux2014_x86_64.whl" +# High-fidelity GPU JPEG decode. The K3 processor enables nvJPEG interpolated +# chroma upsampling through nvImageCodec and zero-copy DLPack handoff to Torch. +RUN python3 -m pip install \ + "nvidia-nvimgcodec-cu12[all]==${NVIMGCODEC_VERSION}" && \ + rm -rf /root/.cache/pip + # Install the pinned FlashInfer MXFP4 MoE runner cubin pool. ARG TRTLLM_GEN_MOE_CUBIN_URL="https://github.com/sgl-project/whl/releases/download/trtllm_gen_moe_cubin_20260617/trtllm_gen_moe_cubin_pool_20260617_v0613rc1.zip" ARG TRTLLM_GEN_MOE_CUBIN_SHA256="4900501cbe782a76b08a5858f9f07152287b97cb68114466dac286366b66c192" diff --git a/docker/kimi_k3/kimi_k3_cu13.Dockerfile b/docker/kimi_k3/kimi_k3_cu13.Dockerfile index 896dace3d..58b884ce4 100644 --- a/docker/kimi_k3/kimi_k3_cu13.Dockerfile +++ b/docker/kimi_k3/kimi_k3_cu13.Dockerfile @@ -27,6 +27,7 @@ FROM lmsysorg/sglang:v0.5.16 AS base ARG SGL_DEEP_GEMM_VERSION="0.1.5.post2" +ARG NVIMGCODEC_VERSION="0.9.0.20" # Current Kimi-K3 source auto-discovers and builds its PyO3 extensions. ARG RUST_VERSION="1.90.0" @@ -79,6 +80,12 @@ RUN TORCH_CUDA_ARCH_LIST="${TORCH_CUDA_ARCH_LIST}" \ RUN python3 -m pip install --no-deps --force-reinstall \ "sgl-deep-gemm==${SGL_DEEP_GEMM_VERSION}" +# High-fidelity GPU JPEG decode. The K3 processor enables nvJPEG interpolated +# chroma upsampling through nvImageCodec and zero-copy DLPack handoff to Torch. +RUN python3 -m pip install \ + "nvidia-nvimgcodec-cu13[all]==${NVIMGCODEC_VERSION}" && \ + rm -rf /root/.cache/pip + # Install the pinned FlashInfer MXFP4 MoE runner cubin pool. ARG TRTLLM_GEN_MOE_CUBIN_URL="https://github.com/sgl-project/whl/releases/download/trtllm_gen_moe_cubin_20260617/trtllm_gen_moe_cubin_pool_20260617_v0613rc1.zip" ARG TRTLLM_GEN_MOE_CUBIN_SHA256="4900501cbe782a76b08a5858f9f07152287b97cb68114466dac286366b66c192" diff --git a/python/sglang/srt/disaggregation/encode_server.py b/python/sglang/srt/disaggregation/encode_server.py index 3fbd16497..03e03bf4f 100644 --- a/python/sglang/srt/disaggregation/encode_server.py +++ b/python/sglang/srt/disaggregation/encode_server.py @@ -647,7 +647,12 @@ class MMEncoder: return data try: if modality == Modality.IMAGE: - img, _ = load_image(data, False) + gpu_image_decode = ( + "nvjpeg_fancy" + if self.use_image_processor_gpu and self.model_type == "kimi_k3" + else False + ) + img, _ = load_image(data, gpu_image_decode) if ( discard_alpha_channel and not isinstance(img, torch.Tensor) diff --git a/python/sglang/srt/multimodal/processors/kimi_k3.py b/python/sglang/srt/multimodal/processors/kimi_k3.py index e3735b622..7b1bc6f59 100644 --- a/python/sglang/srt/multimodal/processors/kimi_k3.py +++ b/python/sglang/srt/multimodal/processors/kimi_k3.py @@ -282,7 +282,10 @@ class KimiK3GPUProcessorWrapper(KimiGPUProcessorWrapper): class KimiK3ImageProcessor(KimiGridMMDataMixin, SGLangBaseProcessor): models = [KimiK3ForConditionalGeneration] - gpu_image_decode = True + # K3 accuracy is sensitive to the chroma upsampling used for common 4:2:0 + # JPEG inputs. This mode uses interpolated nvJPEG upsampling when the K3 + # image dependency is installed and otherwise falls back to PIL. + gpu_image_decode = "nvjpeg_fancy" prefer_tokenized_input = True precompute_hash_before_cpu_transfer = True auto_mm_processor_worker_num = 2 diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py index 0fac99139..17e298f4f 100644 --- a/python/sglang/srt/utils/common.py +++ b/python/sglang/srt/utils/common.py @@ -1643,9 +1643,12 @@ class VideoData: image_extension_names = (".png", ".jpg", ".jpeg", ".webp", ".gif") +GPUImageDecodeMode = Union[bool, Literal["nvjpeg_fancy"]] -def is_jpeg_with_cuda(image_bytes: bytes = b"", gpu_image_decode: bool = True) -> bool: +def is_jpeg_with_cuda( + image_bytes: bytes = b"", gpu_image_decode: GPUImageDecodeMode = True +) -> bool: """ Check three conditions: 1. whether CUDA is available. @@ -1659,10 +1662,19 @@ def is_jpeg_with_cuda(image_bytes: bytes = b"", gpu_image_decode: bool = True) - return False +@lru_cache(maxsize=16) +def _warn_fancy_jpeg_fallback(error: str) -> None: + logger.warning( + "High-fidelity GPU JPEG decode is unavailable; falling back to PIL. " + "Install the Kimi-K3 serving image or NVIDIA nvImageCodec. Error: %s", + error, + ) + + def _load_image( image_bytes: bytes = b"", image_file: str = "", - gpu_image_decode: bool = True, + gpu_image_decode: GPUImageDecodeMode = True, ) -> Union[torch.Tensor, Image.Image]: """ Try to decode JPEG with nvJPEG on GPU and return a torch device tensor, @@ -1673,19 +1685,29 @@ def _load_image( image_bytes = get_image_bytes(image_file) if is_jpeg_with_cuda(image_bytes, gpu_image_decode): try: + if gpu_image_decode == "nvjpeg_fancy": + from sglang.srt.utils.nvjpeg_decoder import ( + decode_jpeg_with_fancy_upsampling, + ) + + return decode_jpeg_with_fancy_upsampling(image_bytes) encoded_image = torch.frombuffer(image_bytes, dtype=torch.uint8) image_tensor = decode_jpeg(encoded_image, device="cuda") return image_tensor except Exception as e: - logger.warning( - f"Failed to decode JPEG on GPU, falling back to CPU. Error: {e}" - ) + if gpu_image_decode == "nvjpeg_fancy": + _warn_fancy_jpeg_fallback(f"{type(e).__name__}: {e}") + else: + logger.warning( + "Failed to decode JPEG on GPU, falling back to CPU. Error: %s", + e, + ) return Image.open(BytesIO(image_bytes)) def load_image( image_file: Union[Image.Image, str, ImageData, bytes], - gpu_image_decode: bool = True, + gpu_image_decode: GPUImageDecodeMode = True, ) -> tuple[Union[torch.Tensor, Image.Image], Optional[tuple[int, int]]]: """ Load image from multiple input formats, including: diff --git a/python/sglang/srt/utils/nvjpeg_decoder.py b/python/sglang/srt/utils/nvjpeg_decoder.py new file mode 100644 index 000000000..da3a6b221 --- /dev/null +++ b/python/sglang/srt/utils/nvjpeg_decoder.py @@ -0,0 +1,86 @@ +"""High-fidelity nvJPEG decoding for multimodal image processors.""" + +from __future__ import annotations + +import queue +import threading +from functools import lru_cache + +import torch + +# A decoder retains roughly 15-25 MiB of device-side scratch space. Two +# decoders are enough to overlap JPEG decode without mirroring the much larger +# I/O thread count into HBM usage. +_DECODER_POOL_SIZE = 2 +_DECODER_OPTIONS = ":num_cuda_streams=1 :fancy_upsampling=1" + + +class _NvJpegDecoderPool: + def __init__(self, device_id: int): + from nvidia import nvimgcodec + + self._nvimgcodec = nvimgcodec + self._device_id = device_id + self._decode_params = nvimgcodec.DecodeParams( + sample_format=nvimgcodec.SampleFormat.P_RGB, + apply_exif_orientation=False, + ) + self._decoders = queue.LifoQueue(maxsize=_DECODER_POOL_SIZE) + self._created = 0 + self._create_lock = threading.Lock() + + def _acquire(self): + try: + return self._decoders.get_nowait() + except queue.Empty: + pass + + with self._create_lock: + if self._created < _DECODER_POOL_SIZE: + decoder = self._nvimgcodec.Decoder( + device_id=self._device_id, + max_num_cpu_threads=1, + options=_DECODER_OPTIONS, + ) + self._created += 1 + return decoder + + return self._decoders.get() + + def decode(self, image_bytes: bytes) -> torch.Tensor: + decoder = self._acquire() + try: + stream = torch.cuda.current_stream(self._device_id) + image = decoder.decode( + image_bytes, + params=self._decode_params, + cuda_stream=stream.cuda_stream, + ) + if image is None: + raise RuntimeError("nvImageCodec could not decode the JPEG image") + return torch.from_dlpack(image.to_dlpack(cuda_stream=stream.cuda_stream)) + finally: + self._decoders.put(decoder) + + +@lru_cache(maxsize=None) +def _get_decoder_pool(device_id: int) -> _NvJpegDecoderPool: + return _NvJpegDecoderPool(device_id) + + +def decode_jpeg_with_fancy_upsampling(image_bytes: bytes) -> torch.Tensor: + """Decode a JPEG to contiguous CHW RGB uint8 on the current CUDA device. + + torchvision's CUDA JPEG decoder creates nvJPEG with its default flags, + which use nearest-neighbor chroma upsampling. nvImageCodec exposes nvJPEG's + interpolated ("fancy") upsampling and exports the result to PyTorch through + DLPack without copying it. + """ + device_id = torch.cuda.current_device() + image = _get_decoder_pool(device_id).decode(image_bytes) + if image.ndim != 3 or image.shape[0] != 3 or image.dtype != torch.uint8: + raise RuntimeError( + "nvImageCodec returned an invalid JPEG tensor: " + f"shape={tuple(image.shape)}, dtype={image.dtype}" + ) + return image diff --git a/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py b/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py index 2210d06b5..751b35aba 100644 --- a/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py +++ b/test/registered/unit/disaggregation/test_kimi_k3_encoder_mode.py @@ -6,7 +6,7 @@ import time from array import array from concurrent.futures import ThreadPoolExecutor from types import SimpleNamespace -from unittest.mock import AsyncMock +from unittest.mock import AsyncMock, patch import pytest import torch @@ -138,6 +138,27 @@ def test_kimi_k3_encoder_passes_media_dicts_to_image_processor(): assert kwargs == {"return_tensors": "pt"} +@pytest.mark.parametrize( + ("use_image_processor_gpu", "expected_decode_mode"), + [(False, False), (True, "nvjpeg_fancy")], +) +def test_kimi_k3_epd_selects_matching_jpeg_decode_mode( + use_image_processor_gpu, expected_decode_mode +): + expected = torch.zeros((3, 2, 3), dtype=torch.uint8) + encoder = _encoder() + encoder.use_image_processor_gpu = use_image_processor_gpu + + with patch( + "sglang.srt.disaggregation.encode_server.load_image", + return_value=(expected, None), + ) as load: + output = encoder._load_single_item(b"jpeg", Modality.IMAGE) + + assert output is expected + load.assert_called_once_with(b"jpeg", expected_decode_mode) + + def test_kimi_k3_epd_aggregates_original_image_sizes_in_part_order(): first = EmbeddingData( req_id="request", diff --git a/test/registered/unit/multimodal/test_base_processor_image_decode.py b/test/registered/unit/multimodal/test_base_processor_image_decode.py index e43b945d5..cc1fbaaba 100644 --- a/test/registered/unit/multimodal/test_base_processor_image_decode.py +++ b/test/registered/unit/multimodal/test_base_processor_image_decode.py @@ -16,15 +16,21 @@ register_cpu_ci(est_time=10, suite="base-a-test-cpu") import asyncio import concurrent.futures import io +import sys +import types import unittest +from types import SimpleNamespace from unittest.mock import Mock, patch import numpy as np import requests +import torch from PIL import Image from sglang.srt.managers.schedule_batch import Modality from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor +from sglang.srt.utils import common +from sglang.srt.utils.nvjpeg_decoder import _NvJpegDecoderPool from sglang.test.test_utils import CustomTestCase @@ -46,6 +52,13 @@ def _png_bytes(mode: str = "RGB", size=(8, 8)) -> bytes: return buf.getvalue() +def _jpeg_bytes(size=(8, 8)) -> bytes: + arr = (np.random.RandomState(0).rand(size[1], size[0], 3) * 255).astype("uint8") + buf = io.BytesIO() + Image.fromarray(arr, "RGB").save(buf, format="JPEG", quality=90, subsampling=2) + return buf.getvalue() + + def _is_decoded(img: Image.Image) -> bool: """A lazily-opened PIL image has no decoded core yet; ``load()`` populates it. PIL's ``.im`` property requires a completed load and raises otherwise.""" @@ -121,6 +134,91 @@ class TestLoadSingleItemImageDecode(CustomTestCase): with self.assertRaisesRegex(RuntimeError, "unexpected loader bug"): _StubProcessor._load_single_item(b"image", Modality.IMAGE) + def test_high_fidelity_gpu_jpeg_decoder_is_selected(self): + data = _jpeg_bytes() + expected = torch.zeros((3, 8, 8), dtype=torch.uint8) + with ( + patch.object(common, "is_cuda", return_value=True), + patch( + "sglang.srt.utils.nvjpeg_decoder.decode_jpeg_with_fancy_upsampling", + return_value=expected, + ) as decode, + ): + image, _ = common.load_image(data, gpu_image_decode="nvjpeg_fancy") + + self.assertIs(image, expected) + decode.assert_called_once_with(data) + + def test_high_fidelity_gpu_jpeg_decoder_falls_back_to_pil(self): + data = _jpeg_bytes() + common._warn_fancy_jpeg_fallback.cache_clear() + with ( + patch.object(common, "is_cuda", return_value=True), + patch( + "sglang.srt.utils.nvjpeg_decoder.decode_jpeg_with_fancy_upsampling", + side_effect=ImportError("nvImageCodec is unavailable"), + ), + ): + image, _ = common.load_image(data, gpu_image_decode="nvjpeg_fancy") + + self.assertIsInstance(image, Image.Image) + reference = Image.open(io.BytesIO(data)) + np.testing.assert_array_equal(np.asarray(image), np.asarray(reference)) + + def test_high_fidelity_decoder_uses_fancy_planar_rgb_and_reuses_pool(self): + expected = torch.zeros((3, 8, 8), dtype=torch.uint8) + fake_format = object() + + class FakeImage: + def to_dlpack(self, *, cuda_stream): + self.cuda_stream = cuda_stream + return object() + + class FakeDecoder: + instances = [] + + def __init__(self, **kwargs): + self.kwargs = kwargs + self.instances.append(self) + + def decode(self, data, *, params, cuda_stream): + self.call = (data, params, cuda_stream) + return FakeImage() + + class FakeDecodeParams: + def __init__(self, *, sample_format, apply_exif_orientation): + self.sample_format = sample_format + self.apply_exif_orientation = apply_exif_orientation + + fake_codec = SimpleNamespace( + DecodeParams=FakeDecodeParams, + Decoder=FakeDecoder, + SampleFormat=SimpleNamespace(P_RGB=fake_format), + ) + nvidia = types.ModuleType("nvidia") + nvidia.nvimgcodec = fake_codec + + with ( + patch.dict(sys.modules, {"nvidia": nvidia}), + patch.object( + torch.cuda, + "current_stream", + return_value=SimpleNamespace(cuda_stream=7), + ), + patch.object(torch, "from_dlpack", return_value=expected), + ): + pool = _NvJpegDecoderPool(device_id=2) + self.assertIs(pool.decode(b"jpeg"), expected) + self.assertIs(pool.decode(b"jpeg"), expected) + + self.assertEqual(len(FakeDecoder.instances), 1) + decoder = FakeDecoder.instances[0] + self.assertEqual(decoder.kwargs["device_id"], 2) + self.assertEqual(decoder.kwargs["max_num_cpu_threads"], 1) + self.assertIn(":fancy_upsampling=1", decoder.kwargs["options"]) + self.assertIs(pool._decode_params.sample_format, fake_format) + self.assertFalse(pool._decode_params.apply_exif_orientation) + if __name__ == "__main__": unittest.main()