[diffusion] fix: fix Wan2.2-I2V-A14B video max size issue(#21390)

Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>
Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
R0CKSTAR
2026-03-31 21:49:53 +08:00
committed by GitHub
co-authored by Mick
parent 7790645b82
commit e67dbf257a
7 changed files with 243 additions and 52 deletions
@@ -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
@@ -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
@@ -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,
@@ -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))
@@ -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": [
@@ -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": {
@@ -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()