fix(vlm): preserve Kimi-K3 GPU JPEG accuracy (#34163)

This commit is contained in:
Mick
2026-08-10 09:42:52 +08:00
committed by GitHub
parent 553dc0f936
commit c20e99bd22
8 changed files with 258 additions and 9 deletions
@@ -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",
@@ -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()