[diffusion] Fix GLM-Image resolution alignment (#32999)
Co-authored-by: AuFlow <AuFlow@users.noreply.github.com>
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
+47
-15
@@ -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
|
||||
)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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"]
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user