diff --git a/python/sglang/multimodal_gen/configs/sample/glmimage.py b/python/sglang/multimodal_gen/configs/sample/glmimage.py index 27ff3c741..7824d8858 100644 --- a/python/sglang/multimodal_gen/configs/sample/glmimage.py +++ b/python/sglang/multimodal_gen/configs/sample/glmimage.py @@ -1,6 +1,11 @@ from dataclasses import dataclass from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +logger = init_logger(__name__) + +GLM_IMAGE_RESOLUTION_ALIGNMENT = 32 @dataclass @@ -10,3 +15,39 @@ class GlmImageSamplingParams(SamplingParams): num_frames: int = 1 guidance_scale: float = 1.5 num_inference_steps: int = 30 + + def _adjust(self, server_args): + requested_width = self.width + requested_height = self.height + if self.width is not None and self.height is not None: + self.width, self.height = align_glm_image_resolution( + self.width, self.height + ) + if (self.width, self.height) != ( + requested_width, + requested_height, + ): + logger.warning( + "GLM-Image requires dimensions divisible by %s; adjusted " + "requested resolution from %sx%s to %sx%s", + GLM_IMAGE_RESOLUTION_ALIGNMENT, + requested_width, + requested_height, + self.width, + self.height, + ) + super()._adjust(server_args) + + +def align_glm_image_dimension(value: int) -> int: + """Round a GLM-Image dimension up to a supported multiple.""" + return max( + GLM_IMAGE_RESOLUTION_ALIGNMENT, + (value + GLM_IMAGE_RESOLUTION_ALIGNMENT - 1) + // GLM_IMAGE_RESOLUTION_ALIGNMENT + * GLM_IMAGE_RESOLUTION_ALIGNMENT, + ) + + +def align_glm_image_resolution(width: int, height: int) -> tuple[int, int]: + return align_glm_image_dimension(width), align_glm_image_dimension(height) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py index e523ab3aa..5c43879b2 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py @@ -19,8 +19,13 @@ from fastapi import ( UploadFile, ) from fastapi.responses import FileResponse +from PIL import Image -from sglang.multimodal_gen.configs.sample.sampling_params import generate_request_id +from sglang.multimodal_gen.configs.sample.glmimage import GlmImageSamplingParams +from sglang.multimodal_gen.configs.sample.sampling_params import ( + SamplingParams, + generate_request_id, +) from sglang.multimodal_gen.runtime.entrypoints.openai.protocol import ( ImageGenerationsRequest, ImageResponse, @@ -170,6 +175,7 @@ def _build_image_response_kwargs( fallback_url: str | None = None, fallback_urls: list[str] | None = None, is_persistent: bool = True, + resize: str | None = None, ) -> dict: """Build ImageResponse data list. @@ -186,6 +192,7 @@ def _build_image_response_kwargs( b64_json=b64, revised_prompt=prompt, file_path=os.path.abspath(path) if is_persistent else None, + resize=resize, ) for b64, path in zip(b64_list, save_file_path_list) ] @@ -210,6 +217,7 @@ def _build_image_response_kwargs( url=url, revised_prompt=prompt, file_path=os.path.abspath(path) if is_persistent else None, + resize=resize, ) ) @@ -231,6 +239,28 @@ def _build_image_response_kwargs( return ret +def _get_response_resize( + sampling_params: SamplingParams, output_path: str | None = None +) -> str | None: + """Return a generated GLM-Image output's actual size as WIDTHxHEIGHT.""" + if not isinstance(sampling_params, GlmImageSamplingParams): + return None + + if output_path is not None: + try: + with Image.open(output_path) as output_image: + width, height = output_image.size + return f"{width}x{height}" + except (OSError, ValueError): + # Fall back to the aligned sampling canvas if the output cannot be + # inspected (for example, for a custom output transport). + pass + + if sampling_params.width is None or sampling_params.height is None: + return None + return sampling_params.output_size_str() + + @router.post("/generations", response_model=ImageResponse) async def generations( request: ImageGenerationsRequest, @@ -308,6 +338,7 @@ async def generations( async_scheduler_client, batch ) save_file_path = save_file_path_list[0] + response_resize = _get_response_resize(sampling, save_file_path) resp_format = (request.response_format or "b64_json").lower() if ( is_cosmos3 @@ -359,6 +390,7 @@ async def generations( cloud_urls=cloud_urls, fallback_urls=fallback_urls, is_persistent=is_persistent, + resize=response_resize, ) return ImageResponse(**response_kwargs) @@ -462,6 +494,7 @@ async def edits( async_scheduler_client, batch ) save_file_path = save_file_path_list[0] + response_resize = _get_response_resize(sampling, save_file_path) resp_format = (response_format or "b64_json").lower() # read b64 before cloud upload may delete the local file @@ -510,6 +543,7 @@ async def edits( cloud_urls=cloud_urls, fallback_urls=fallback_urls, is_persistent=is_persistent, + resize=response_resize, ) return ImageResponse(**response_kwargs) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py index ff4fed8f1..ac6895035 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py @@ -13,6 +13,7 @@ class ImageResponseData(BaseModel): url: Optional[str] = None revised_prompt: Optional[str] = None file_path: Optional[str] = None + resize: Optional[str] = None class ImagePromptTokensDetails(BaseModel): diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/glm_image.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/glm_image.py index 01e4d84ca..a8ff3e4f5 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/glm_image.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/glm_image.py @@ -11,6 +11,10 @@ import torch from diffusers.image_processor import VaeImageProcessor from diffusers.utils.torch_utils import randn_tensor +from sglang.multimodal_gen.configs.sample.glmimage import ( + GLM_IMAGE_RESOLUTION_ALIGNMENT, + align_glm_image_resolution, +) from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import ( @@ -118,6 +122,26 @@ def image_path_to_list(image_path: Union[str, List[str]]) -> List[str]: return image_path if isinstance(image_path, list) else [image_path] +def resize_glm_image_to_alignment(image: PIL.Image.Image) -> PIL.Image.Image: + """Resize an image up so both dimensions use GLM-Image's D32 grid.""" + width, height = image.size + aligned_width, aligned_height = align_glm_image_resolution(width, height) + if (aligned_width, aligned_height) == (width, height): + return image + return image.resize((aligned_width, aligned_height), PIL.Image.Resampling.LANCZOS) + + +def _validate_glm_image_resolution_alignment(width: int, height: int) -> None: + if ( + height % GLM_IMAGE_RESOLUTION_ALIGNMENT != 0 + or width % GLM_IMAGE_RESOLUTION_ALIGNMENT != 0 + ): + raise ValueError( + "GLM-Image dimensions must be aligned before AR token generation, " + f"got {width}x{height}" + ) + + def pooled_image_features_to_tensor(image_features) -> torch.Tensor: pooler_output = getattr(image_features, "pooler_output", None) if pooler_output is not None: @@ -337,7 +361,6 @@ class GlmImageAR(PipelineStage): width: int, server_args: ServerArgs, image: Optional[List[PIL.Image.Image]] = None, - factor: int = 32, seed: Optional[int] = None, ) -> Tuple[torch.Tensor, Optional[List[torch.Tensor]], Optional[dict[str, int]]]: """ @@ -348,14 +371,11 @@ class GlmImageAR(PipelineStage): condition_images: Optional list of condition images for i2i Returns: - Tuple of (prior_token_ids, pixel_height, pixel_width) - - prior_token_ids: Upsampled to d16 format, shape [1, token_h*token_w*4] - - pixel_height: Image height in pixels - - pixel_width: Image width in pixels + Tuple of the D16 prior token IDs, optional source-image token IDs, + and optional usage statistics returned by an external AR server. """ device = get_local_torch_device() - height = (height // factor) * factor - width = (width // factor) * factor + _validate_glm_image_resolution_alignment(width, height) is_text_to_image = image is None or len(image) == 0 # Build messages for processor @@ -456,11 +476,9 @@ class GlmImageAR(PipelineStage): height: int, width: int, server_args: ServerArgs, - factor: int = 32, ) -> tuple[list[torch.Tensor], list[dict[str, int] | None]]: device = get_local_torch_device() - height = (height // factor) * factor - width = (width // factor) * factor + _validate_glm_image_resolution_alignment(width, height) input_ids = [] image_data = [] @@ -650,7 +668,7 @@ class GlmImageAR(PipelineStage): width = batch.width if batch.image_path is not None: ar_condition_images = [ - load_image(img_path) + resize_glm_image_to_alignment(load_image(img_path)) for img_path in image_path_to_list(batch.image_path) ] else: @@ -662,6 +680,20 @@ class GlmImageAR(PipelineStage): height = height or ar_condition_images[0].height width = width or ar_condition_images[0].width + requested_width = width + requested_height = height + width, height = align_glm_image_resolution(width, height) + if (width, height) != (requested_width, requested_height): + logger.warning( + "GLM-Image requires dimensions divisible by %s; adjusted " + "runtime resolution from %sx%s to %sx%s", + GLM_IMAGE_RESOLUTION_ALIGNMENT, + requested_width, + requested_height, + width, + height, + ) + time_start = time.time() num_outputs = _num_outputs_per_prompt(batch) seed = getattr(batch, "seed", None) @@ -1037,7 +1069,7 @@ class GlmImageBeforeDenoisingStage(PipelineStage): num_inference_steps = batch.num_inference_steps if batch.image_path is not None: ar_condition_images = [ - load_image(img_path) + resize_glm_image_to_alignment(load_image(img_path)) for img_path in image_path_to_list(batch.image_path) ] else: @@ -1108,9 +1140,9 @@ class GlmImageBeforeDenoisingStage(PipelineStage): if isinstance(img, PIL.Image.Image) else img.shape[:2] ) - multiple_of = self.vae_scale_factor * self.transformer.config.patch_size - image_height = (image_height // multiple_of) * multiple_of - image_width = (image_width // multiple_of) * multiple_of + image_width, image_height = align_glm_image_resolution( + image_width, image_height + ) img = self.image_processor.preprocess( img, height=image_height, width=image_width ) diff --git a/python/sglang/multimodal_gen/test/unit/test_glm_image_ar.py b/python/sglang/multimodal_gen/test/unit/test_glm_image_ar.py index 3661df8f1..2cf9dc980 100644 --- a/python/sglang/multimodal_gen/test/unit/test_glm_image_ar.py +++ b/python/sglang/multimodal_gen/test/unit/test_glm_image_ar.py @@ -1,13 +1,15 @@ import unittest from types import SimpleNamespace -from unittest.mock import patch +from unittest.mock import MagicMock, patch import torch +from PIL import Image +from sglang.multimodal_gen.configs.sample.glmimage import GlmImageSamplingParams from sglang.multimodal_gen.runtime.entrypoints.openai.image_api import ( _build_image_response_kwargs, ) -from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch +from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch, Req from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.glm_image import ( GlmImageAR, ) @@ -200,6 +202,120 @@ class TestGlmImageARSrtBackend(unittest.TestCase): server_args=self._server_args(), ) + @patch( + "sglang.multimodal_gen.runtime.pipelines_core.stages." + "model_specific_stages.glm_image.get_local_torch_device", + return_value=torch.device("cpu"), + ) + def test_forward_aligns_runtime_dimensions_before_ar_generation(self, _mock_device): + stage = GlmImageAR(processor=_FakeProcessor(), vision_language_encoder=None) + stage.generate_prior_tokens = MagicMock( + return_value=(torch.zeros((1, 1), dtype=torch.long), None, None) + ) + cases = [ + ((500, 500), (512, 512)), + ((550, 1009), (576, 1024)), + ((1280, 720), (1280, 736)), + ] + + for requested, expected in cases: + with self.subTest(requested=requested): + stage.generate_prior_tokens.reset_mock() + sampling = GlmImageSamplingParams( + prompt="A simple product sketch", + width=requested[0], + height=requested[1], + ) + sampling.seed = None + batch = Req(sampling_params=sampling) + + stage.forward(batch, self._server_args()) + + self.assertEqual((batch.width, batch.height), expected) + stage.generate_prior_tokens.assert_called_once_with( + prompt="A simple product sketch", + image=None, + height=expected[1], + width=expected[0], + server_args=self._server_args(), + ) + + @patch( + "sglang.multimodal_gen.runtime.pipelines_core.stages." + "model_specific_stages.glm_image.get_local_torch_device", + return_value=torch.device("cpu"), + ) + @patch( + "sglang.multimodal_gen.runtime.pipelines_core.stages." + "model_specific_stages.glm_image.load_image", + return_value=Image.new("RGB", (1280, 720)), + ) + def test_forward_resizes_edit_image_up_to_d32_grid( + self, _mock_load_image, _mock_device + ): + stage = GlmImageAR(processor=_FakeProcessor(), vision_language_encoder=None) + stage.generate_prior_tokens = MagicMock( + return_value=(torch.zeros((1, 1), dtype=torch.long), None, None) + ) + sampling = GlmImageSamplingParams( + prompt="Edit this image", + width=1280, + height=720, + image_path="input.png", + ) + sampling.seed = None + batch = Req(sampling_params=sampling) + + stage.forward(batch, self._server_args()) + + call_kwargs = stage.generate_prior_tokens.call_args.kwargs + self.assertEqual((batch.width, batch.height), (1280, 736)) + self.assertEqual(call_kwargs["image"][0].size, (1280, 736)) + self.assertEqual((call_kwargs["width"], call_kwargs["height"]), (1280, 736)) + + @patch( + "sglang.multimodal_gen.runtime.pipelines_core.stages." + "model_specific_stages.glm_image.get_local_torch_device", + return_value=torch.device("cpu"), + ) + def test_generate_prior_tokens_rejects_unaligned_internal_dimensions( + self, _mock_device + ): + stage = GlmImageAR(processor=_FakeProcessor(), vision_language_encoder=None) + + with self.assertRaisesRegex( + ValueError, + "GLM-Image dimensions must be aligned before AR token generation", + ): + stage.generate_prior_tokens( + prompt="A simple product sketch", + height=1024, + width=550, + server_args=self._server_args(), + ) + + @patch( + "sglang.multimodal_gen.runtime.pipelines_core.stages." + "model_specific_stages.glm_image.get_local_torch_device", + return_value=torch.device("cpu"), + ) + def test_generate_prior_tokens_batch_rejects_unaligned_internal_dimensions( + self, _mock_device + ): + stage = GlmImageAR(processor=_FakeProcessor(), vision_language_encoder=None) + + with self.assertRaisesRegex( + ValueError, + "GLM-Image dimensions must be aligned before AR token generation", + ): + stage.generate_prior_tokens_batch( + prompts=["A simple product sketch"], + seeds=[42], + height=1024, + width=550, + server_args=self._server_args(), + ) + def test_image_response_adds_image_count_to_usage(self): set_global_server_args(SimpleNamespace(enable_cache_report=False)) response = _build_image_response_kwargs( diff --git a/python/sglang/multimodal_gen/test/unit/test_openai_image_api.py b/python/sglang/multimodal_gen/test/unit/test_openai_image_api.py index abba9174f..07a7a09b0 100644 --- a/python/sglang/multimodal_gen/test/unit/test_openai_image_api.py +++ b/python/sglang/multimodal_gen/test/unit/test_openai_image_api.py @@ -1,10 +1,14 @@ import os from fastapi import HTTPException +from PIL import Image +from sglang.multimodal_gen.configs.sample.glmimage import GlmImageSamplingParams +from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams from sglang.multimodal_gen.runtime.entrypoints.openai.image_api import ( _build_image_response_kwargs, _fallback_image_urls, + _get_response_resize, _raise_if_image_variant_not_found, _select_image_variant_cloud_url, _select_image_variant_path, @@ -36,6 +40,46 @@ def test_url_response_returns_one_item_per_output_path(): ] +def test_image_response_includes_resize_for_every_output(): + response = _build_image_response_kwargs( + ["first.png", "second.png"], + "b64_json", + "a lantern", + "req-123", + OutputBatch(), + b64_list=["first", "second"], + resize="1280x736", + ) + + assert [item.resize for item in response["data"]] == [ + "1280x736", + "1280x736", + ] + + +def test_response_resize_is_only_populated_for_glm_image(): + glm_sampling = GlmImageSamplingParams(width=1280, height=736) + + assert _get_response_resize(glm_sampling) == "1280x736" + assert _get_response_resize(SamplingParams(width=1280, height=736)) is None + + +def test_response_resize_uses_actual_generated_image_size(tmp_path): + output_path = tmp_path / "output.png" + Image.new("RGB", (1280, 736)).save(output_path) + glm_sampling = GlmImageSamplingParams(image_path="input.png") + + assert _get_response_resize(glm_sampling, str(output_path)) == "1280x736" + + +def test_response_resize_prefers_final_output_over_sampling_canvas(tmp_path): + output_path = tmp_path / "upscaled.png" + Image.new("RGB", (2560, 1472)).save(output_path) + glm_sampling = GlmImageSamplingParams(width=1280, height=736) + + assert _get_response_resize(glm_sampling, str(output_path)) == "2560x1472" + + def test_url_response_uses_variant_fallback_urls_for_multiple_persistent_outputs(): paths = ["first.png", "second.png"] diff --git a/python/sglang/multimodal_gen/test/unit/test_openai_utils.py b/python/sglang/multimodal_gen/test/unit/test_openai_utils.py index 69f8ad058..91a507535 100644 --- a/python/sglang/multimodal_gen/test/unit/test_openai_utils.py +++ b/python/sglang/multimodal_gen/test/unit/test_openai_utils.py @@ -2,13 +2,17 @@ import asyncio import io +from types import SimpleNamespace from starlette.datastructures import UploadFile as StarletteUploadFile +from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams +from sglang.multimodal_gen.runtime.entrypoints.openai import utils as openai_utils from sglang.multimodal_gen.runtime.entrypoints.openai.utils import ( _parse_size_or_raise, _save_upload_to_path, _validate_positive_int, + build_sampling_params, ) @@ -57,3 +61,49 @@ def test_validate_positive_int_rejects_non_positive_sampling_fields(): assert "num_frames must be positive" in exc.detail else: raise AssertionError("expected bad request") + + +def test_build_sampling_params_resolves_size_and_explicit_dimensions(monkeypatch): + server_args = SimpleNamespace(model_path="zai-org/GLM-Image") + monkeypatch.setattr(openai_utils, "get_global_server_args", lambda: server_args) + captured = {} + + def fake_from_user_sampling_params_args(**kwargs): + captured.update(kwargs) + return SimpleNamespace() + + monkeypatch.setattr( + SamplingParams, + "from_user_sampling_params_args", + fake_from_user_sampling_params_args, + ) + + cases = [ + ( + {"size": "500x500", "width": None, "height": None}, + (500, 500), + ), + ( + {"size": "1024x1024", "width": None, "height": 600}, + (1024, 600), + ), + ( + {"size": "500x500", "width": None, "height": 600}, + (500, 600), + ), + ( + {"size": "500x500", "width": 600, "height": None}, + (600, 500), + ), + ( + {"size": "500x500", "width": 600, "height": 700}, + (600, 700), + ), + ] + + for request_fields, expected in cases: + captured.clear() + + build_sampling_params("request-id", **request_fields) + + assert (captured["width"], captured["height"]) == expected diff --git a/python/sglang/multimodal_gen/test/unit/test_sampling_params.py b/python/sglang/multimodal_gen/test/unit/test_sampling_params.py index fe1dad199..f8a5b0f3e 100644 --- a/python/sglang/multimodal_gen/test/unit/test_sampling_params.py +++ b/python/sglang/multimodal_gen/test/unit/test_sampling_params.py @@ -4,6 +4,9 @@ import unittest from types import SimpleNamespace from unittest.mock import MagicMock, patch +from sglang.multimodal_gen.configs.pipeline_configs.glm_image import ( + GlmImagePipelineConfig, +) from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import ( LTX2PipelineConfig, is_ltx23_native_variant, @@ -17,6 +20,10 @@ from sglang.multimodal_gen.configs.sample.flux import ( Flux2SamplingParams, FluxSamplingParams, ) +from sglang.multimodal_gen.configs.sample.glmimage import ( + GlmImageSamplingParams, + align_glm_image_dimension, +) from sglang.multimodal_gen.configs.sample.qwenimage import QwenImageSamplingParams from sglang.multimodal_gen.configs.sample.sampling_params import ( SamplingParams, @@ -100,6 +107,65 @@ class TestSamplingParamsValidate(unittest.TestCase): class TestSamplingParamsSubclass(unittest.TestCase): + def test_glm_image_rounds_resolution_up_to_multiple_of_32(self): + server_args = SimpleNamespace( + pipeline_config=GlmImagePipelineConfig(), + output_path=None, + comfyui_mode=True, + ) + cases = [ + ((500, 500), (512, 512)), + ((1024, 600), (1024, 608)), + ((500, 600), (512, 608)), + ((550, 1009), (576, 1024)), + ((1280, 720), (1280, 736)), + ] + + for requested, expected in cases: + with self.subTest(requested=requested): + params = GlmImageSamplingParams( + width=requested[0], + height=requested[1], + ) + + with patch( + "sglang.multimodal_gen.configs.sample.glmimage.logger.warning" + ) as mock_warning: + params._adjust(server_args) + + self.assertEqual((params.width, params.height), expected) + mock_warning.assert_called_once_with( + "GLM-Image requires dimensions divisible by %s; adjusted " + "requested resolution from %sx%s to %sx%s", + 32, + requested[0], + requested[1], + expected[0], + expected[1], + ) + + def test_glm_image_resolution_rounds_up(self): + self.assertEqual(align_glm_image_dimension(560), 576) + + def test_glm_image_resolution_keeps_minimum_alignment(self): + self.assertEqual(align_glm_image_dimension(0), 32) + self.assertEqual(align_glm_image_dimension(-1), 32) + + def test_glm_image_does_not_warn_for_aligned_resolution(self): + server_args = SimpleNamespace( + pipeline_config=GlmImagePipelineConfig(), + output_path=None, + comfyui_mode=True, + ) + params = GlmImageSamplingParams(width=1024, height=1024) + + with patch( + "sglang.multimodal_gen.configs.sample.glmimage.logger.warning" + ) as mock_warning: + params._adjust(server_args) + + mock_warning.assert_not_called() + def test_flux_defaults_resolution_when_not_provided(self): params = FluxSamplingParams()