diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/wan.py b/python/sglang/multimodal_gen/configs/pipeline_configs/wan.py index 6a824e678..c08851e93 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/wan.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/wan.py @@ -214,7 +214,7 @@ class Wan2_2_T2V_A14B_Config(WanT2V480PConfig): @dataclass -class Wan2_2_I2V_A14B_Config(WanI2V480PConfig): +class Wan2_2_I2V_A14B_Config(WanI2V720PConfig): flow_shift: float | None = 5.0 boundary_ratio: float | None = 0.900 diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py index 9831d368a..b326a2950 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py @@ -90,6 +90,8 @@ class VideoGenerationsRequest(BaseModel): seed: Optional[int] = 1024 generator_device: Optional[str] = "cuda" # SGLang extensions + width: Optional[int] = None + height: Optional[int] = None num_inference_steps: Optional[int] = None guidance_scale: Optional[float] = None guidance_scale_2: Optional[float] = None diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py index 2a444111d..abccf31bb 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py @@ -60,6 +60,8 @@ def _build_video_sampling_params(request_id: str, request: VideoGenerationsReque request_id, prompt=request.prompt, size=request.size, + width=request.width, + height=request.height, num_frames=num_frames, fps=fps, image_path=request.input_reference, diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/input_validation.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/input_validation.py index aa0b686b8..fc27f657f 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/input_validation.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/input_validation.py @@ -191,8 +191,30 @@ class InputValidationStage(PipelineStage): server_args.pipeline_config.vae_config.arch_config.scale_factor_spatial * server_args.pipeline_config.dit_config.arch_config.patch_size[1] ) + + # User-specified width/height controls the target area (scale), + # capped by max_area. Aspect ratio always comes from the + # condition image for I2V. + if batch.width is not None or batch.height is not None: + # If one dimension is provided, calculate the other based on the image's aspect ratio. + if batch.width is None: + batch.width = round(batch.height / aspect_ratio) + elif batch.height is None: + batch.height = round(batch.width * aspect_ratio) + + target_area = min(batch.width * batch.height, max_area) + if batch.width * batch.height > max_area: + logger.warning( + "Requested resolution %dx%d exceeds max_area %d, " + "clamping to max_area", + batch.width, + batch.height, + max_area, + ) + else: + target_area = max_area width, height = self._calculate_dimensions_from_area( - max_area, aspect_ratio, mod_value + target_area, aspect_ratio, mod_value ) batch.condition_image = batch.condition_image.resize((width, height)) diff --git a/python/sglang/multimodal_gen/test/run_suite.py b/python/sglang/multimodal_gen/test/run_suite.py index f7182b9bb..b01120c02 100644 --- a/python/sglang/multimodal_gen/test/run_suite.py +++ b/python/sglang/multimodal_gen/test/run_suite.py @@ -35,6 +35,7 @@ SUITES = { "../unit/test_storage.py", "../unit/test_lora_format_adapter.py", "../unit/test_server_args.py", + "../unit/test_input_validation.py", # add new unit tests here ], "1-gpu": [ diff --git a/python/sglang/multimodal_gen/test/server/perf_baselines.json b/python/sglang/multimodal_gen/test/server/perf_baselines.json index 68ee16827..195ae846a 100644 --- a/python/sglang/multimodal_gen/test/server/perf_baselines.json +++ b/python/sglang/multimodal_gen/test/server/perf_baselines.json @@ -1317,60 +1317,60 @@ }, "wan2_2_i2v_a14b_2gpu": { "stages_ms": { - "InputValidationStage": 18.45, - "TextEncodingStage": 3337.77, - "TimestepPreparationStage": 2.9, - "LatentPreparationStage": 1.25, - "ImageVAEEncodingStage": 1655.89, - "DenoisingStage": 106972.82, - "DecodingStage": 1355.52, + "InputValidationStage": 27.74, + "TextEncodingStage": 1121.93, + "ImageVAEEncodingStage": 1889.26, + "LatentPreparationStage": 0.44, + "TimestepPreparationStage": 6.39, + "DenoisingStage": 137454.52, + "DecodingStage": 2287.25, "per_frame_generation": null }, "denoise_step_ms": { - "0": 1525.6, - "1": 1582.6, - "2": 1597.84, - "3": 1601.34, - "4": 1600.86, - "5": 1598.32, - "6": 1600.93, - "7": 1599.88, - "8": 1600.0, - "9": 1600.55, - "10": 1599.27, - "11": 1600.59, - "12": 1600.17, - "13": 1599.72, - "14": 1599.76, - "15": 24098.85, - "16": 1601.29, - "17": 1598.89, - "18": 1600.12, - "19": 1600.52, - "20": 1599.59, - "21": 1600.37, - "22": 1600.35, - "23": 1599.7, - "24": 1599.92, - "25": 1599.75, - "26": 1600.2, - "27": 1600.06, - "28": 1600.41, - "29": 1599.35, - "30": 1600.69, - "31": 1600.15, - "32": 1599.33, - "33": 1599.86, - "34": 1600.52, - "35": 1599.84, - "36": 1600.38, - "37": 1599.23, - "38": 1600.27, - "39": 1599.78 + "0": 2231.66, + "1": 3489.95, + "2": 3436.66, + "3": 3407.31, + "4": 3422.63, + "5": 3417.48, + "6": 3425.34, + "7": 3423.93, + "8": 3429.36, + "9": 3431.95, + "10": 3435.35, + "11": 3430.29, + "12": 3435.09, + "13": 3436.59, + "14": 3436.94, + "15": 4835.04, + "16": 3416.6, + "17": 3427.03, + "18": 3421.59, + "19": 3427.95, + "20": 3427.21, + "21": 3428.96, + "22": 3430.96, + "23": 3431.29, + "24": 3430.44, + "25": 3430.09, + "26": 3432.23, + "27": 3430.61, + "28": 3430.51, + "29": 3427.92, + "30": 3429.01, + "31": 3430.05, + "32": 3429.63, + "33": 3426.97, + "34": 3426.71, + "35": 3428.44, + "36": 3427.1, + "37": 3425.52, + "38": 3422.81, + "39": 3403.77 }, - "expected_e2e_ms": 123182.9887, - "expected_avg_denoise_ms": 2831.0, - "expected_median_denoise_ms": 1600.09 + "expected_e2e_ms": 144621.32, + "expected_avg_denoise_ms": 3434.22, + "expected_median_denoise_ms": 3428.99 }, "turbo_wan2_2_i2v_a14b_2gpu": { "stages_ms": { diff --git a/python/sglang/multimodal_gen/test/unit/test_input_validation.py b/python/sglang/multimodal_gen/test/unit/test_input_validation.py new file mode 100644 index 000000000..75bd30bf7 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_input_validation.py @@ -0,0 +1,164 @@ +"""Unit tests for InputValidationStage.preprocess_condition_image resolution logic.""" + +import unittest +from unittest.mock import MagicMock, patch + +from PIL import Image + +from sglang.multimodal_gen.configs.pipeline_configs.wan import ( + WanI2V480PConfig, + WanI2V720PConfig, +) +from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams +from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req +from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import ( + InputValidationStage, +) + +# Patch path for get_global_server_args used by Stage.__init__ +_GLOBAL_ARGS_PATCH = ( + "sglang.multimodal_gen.runtime.pipelines_core.stages.base.get_global_server_args" +) + + +def _make_batch(condition_image: Image.Image, width=None, height=None) -> Req: + """Create a minimal Req with a condition image and optional user dimensions.""" + sp = SamplingParams( + seed=42, + num_outputs_per_prompt=1, + width=width, + height=height, + ) + batch = Req(sampling_params=sp, condition_image=condition_image) + return batch + + +def _make_server_args(pipeline_config): + """Create a mock ServerArgs with the given pipeline config.""" + sa = MagicMock() + sa.pipeline_config = pipeline_config + return sa + + +class TestCalculateDimensionsFromArea(unittest.TestCase): + """Tests for InputValidationStage._calculate_dimensions_from_area.""" + + def test_square_aspect_ratio(self): + # area=921600, aspect=1.0, mod=16 → sqrt(921600)=~960 + w, h = InputValidationStage._calculate_dimensions_from_area(921600, 1.0, 16) + self.assertEqual(w % 16, 0) + self.assertEqual(h % 16, 0) + self.assertEqual((w, h), (960, 960)) + + def test_16_9_aspect_ratio(self): + # aspect = 720/1280 = 0.5625 + w, h = InputValidationStage._calculate_dimensions_from_area(921600, 9 / 16, 16) + self.assertEqual(w % 16, 0) + self.assertEqual(h % 16, 0) + self.assertEqual((w, h), (1280, 720)) + + def test_9_16_aspect_ratio(self): + w, h = InputValidationStage._calculate_dimensions_from_area(921600, 16 / 9, 16) + self.assertEqual(w % 16, 0) + self.assertEqual(h % 16, 0) + self.assertEqual((w, h), (720, 1280)) + + def test_mod_alignment(self): + # Ensure dimensions are always multiples of mod_value + w, h = InputValidationStage._calculate_dimensions_from_area(500000, 1.3, 16) + self.assertEqual(w % 16, 0) + self.assertEqual(h % 16, 0) + + +class TestPreprocessConditionImageResolution(unittest.TestCase): + """Tests for the WanI2V480PConfig branch of preprocess_condition_image. + + Verifies that: + - Aspect ratio always comes from the condition image + - User-specified width/height controls target area (scale) + - Output is clamped to max_area when user dimensions exceed it + - Dimensions are always mod-aligned + """ + + def setUp(self): + with patch(_GLOBAL_ARGS_PATCH, return_value=MagicMock()): + self.stage = InputValidationStage() + + def _run(self, config, img_w, img_h, user_w=None, user_h=None): + """Run preprocess_condition_image and return (batch.width, batch.height).""" + img = Image.new("RGB", (img_w, img_h), color="red") + batch = _make_batch(img, width=user_w, height=user_h) + server_args = _make_server_args(config) + self.stage.preprocess_condition_image(batch, server_args, img_w, img_h) + return batch.width, batch.height + + def test_720p_no_user_dims_16_9_image(self): + """16:9 image, no user dims → 1280×720.""" + w, h = self._run(WanI2V720PConfig(), 1920, 1080) + self.assertEqual((w, h), (1280, 720)) + + def test_720p_no_user_dims_9_16_image(self): + """9:16 image, no user dims → 720×1280.""" + w, h = self._run(WanI2V720PConfig(), 1080, 1920) + self.assertEqual((w, h), (720, 1280)) + + def test_720p_no_user_dims_square_image(self): + """Square image, no user dims → ~960×960 (max_area=921600, sqrt≈960).""" + w, h = self._run(WanI2V720PConfig(), 1024, 1024) + self.assertEqual((w, h), (960, 960)) + self.assertEqual(w % 16, 0) + + def test_720p_user_dims_equal_max_area_16_9_image(self): + """16:9 image + user 1280×720 (=max_area) → 1280×720.""" + w, h = self._run(WanI2V720PConfig(), 1920, 1080, 1280, 720) + self.assertEqual((w, h), (1280, 720)) + + def test_720p_user_dims_equal_max_area_square_image(self): + """Square image + user 1280×720 → still square (~960×960) because + aspect ratio comes from image, not from user dimensions.""" + w, h = self._run(WanI2V720PConfig(), 1024, 1024, 1280, 720) + self.assertEqual((w, h), (960, 960)) + + def test_720p_user_dims_smaller_area(self): + """Square image + user 832×480 → smaller square (target_area=399360).""" + w, h = self._run(WanI2V720PConfig(), 1024, 1024, 832, 480) + self.assertEqual((w, h), (624, 624)) + self.assertEqual(w % 16, 0) + + def test_720p_user_dims_exceed_max_area(self): + """4K request clamped to max_area.""" + w, h = self._run(WanI2V720PConfig(), 1920, 1080, 3840, 2160) + self.assertEqual(w % 16, 0) + self.assertEqual(h % 16, 0) + self.assertEqual((w, h), (1280, 720)) + + def test_480p_no_user_dims_16_9_image(self): + """480p config, 16:9 image → area-based calc from max_area=399360.""" + w, h = self._run(WanI2V480PConfig(), 1920, 1080) + # max_area=480*832=399360, aspect=9/16 → (832, 464) due to rounding + self.assertEqual(w % 16, 0) + self.assertEqual(h % 16, 0) + self.assertEqual((w, h), (832, 464)) + + def test_condition_image_resized_to_output_dims(self): + """Condition image is resized to match output dimensions.""" + img = Image.new("RGB", (1920, 1080), color="blue") + batch = _make_batch(img) + server_args = _make_server_args(WanI2V720PConfig()) + self.stage.preprocess_condition_image(batch, server_args, 1920, 1080) + self.assertEqual(batch.condition_image.size, (batch.width, batch.height)) + + def test_list_condition_image_takes_first(self): + """List of condition images → uses first one.""" + img1 = Image.new("RGB", (1920, 1080), color="red") + img2 = Image.new("RGB", (800, 600), color="green") + batch = _make_batch(img1) + batch.condition_image = [img1, img2] + server_args = _make_server_args(WanI2V720PConfig()) + self.stage.preprocess_condition_image(batch, server_args, 1920, 1080) + self.assertIsInstance(batch.condition_image, Image.Image) + self.assertEqual((batch.width, batch.height), (1280, 720)) + + +if __name__ == "__main__": + unittest.main()