[diffusion] feat: support resolution check for video model (#14881)

Co-authored-by: Brain97 <Brain97@users.noreply.github.com>
Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
blahblah
2025-12-14 17:50:13 +08:00
committed by GitHub
co-authored by Brain97 Mick
parent 5c75907e62
commit 19c16748ce
11 changed files with 469 additions and 377 deletions
@@ -18,6 +18,24 @@ class HunyuanSamplingParams(SamplingParams):
guidance_scale: float = 1.0
# HunyuanVideo supported resolutions
supported_resolutions: list[tuple[int, int]] | None = field(
default_factory=lambda: [
# 540p resolutions
(960, 544), # 9:16
(544, 960), # 16:9
(832, 624), # 4:3
(624, 832), # 3:4
(720, 720), # 1:1
# 720p resolutions (recommended)
(1280, 720), # 9:16
(720, 1280), # 16:9
(832, 1104), # 4:3
(1104, 832), # 3:4
(960, 960), # 1:1
]
)
teacache_params: TeaCacheParams = field(
default_factory=lambda: TeaCacheParams(
teacache_thresh=0.15,
@@ -115,6 +115,11 @@ class SamplingParams:
width_not_provided: bool = False
fps: int = 24
# Resolution validation
supported_resolutions: list[tuple[int, int]] | None = (
None # None means all resolutions allowed
)
# Denoising parameters
num_inference_steps: int = None
guidance_scale: float = None
@@ -223,6 +228,27 @@ class SamplingParams:
f"num_frames={self.num_frames}"
)
# Validate resolution against pipeline-specific supported resolutions
if self.height is None and self.width is None:
if self.supported_resolutions is not None:
self.width, self.height = self.supported_resolutions[0]
logger.info(
f"Resolution unspecified, using default: {self.supported_resolutions[0]}"
)
if self.height is not None and self.width is not None:
if self.supported_resolutions is not None:
if (self.width, self.height) not in self.supported_resolutions:
supported_str = ", ".join(
[f"{w}x{h}" for w, h in self.supported_resolutions]
)
error_msg = (
f"Unsupported resolution: {self.width}x{self.height}. "
f"Supported resolutions: {supported_str}"
)
logger.error(error_msg)
raise ValueError(error_msg)
if pipeline_config.task_type.is_image_gen():
# settle num_frames
logger.debug(f"Setting num_frames to 1 because this is an image-gen model")
@@ -558,7 +584,9 @@ class SamplingParams:
args.width = 1280
args.height = 720
attrs = [attr.name for attr in dataclasses.fields(cls)]
sampling_params_fields = {attr.name for attr in dataclasses.fields(cls)}
args_attrs = set(vars(args).keys())
attrs = sampling_params_fields & args_attrs
args.height_not_provided = False
args.width_not_provided = False
return {attr: getattr(args, attr) for attr in attrs}
@@ -22,6 +22,14 @@ class WanT2V_1_3B_SamplingParams(SamplingParams):
)
num_inference_steps: int = 50
# Wan T2V 1.3B supported resolutions
supported_resolutions: list[tuple[int, int]] | None = field(
default_factory=lambda: [
(832, 480), # 16:9
(480, 832), # 9:16
]
)
teacache_params: WanTeaCacheParams = field(
default_factory=lambda: WanTeaCacheParams(
teacache_thresh=0.08,
@@ -58,6 +66,16 @@ class WanT2V_14B_SamplingParams(SamplingParams):
)
num_inference_steps: int = 50
# Wan T2V 14B supported resolutions
supported_resolutions: list[tuple[int, int]] | None = field(
default_factory=lambda: [
(1280, 720), # 16:9
(720, 1280), # 9:16
(832, 480), # 16:9
(480, 832), # 9:16
]
)
teacache_params: WanTeaCacheParams = field(
default_factory=lambda: WanTeaCacheParams(
teacache_thresh=0.20,
@@ -87,6 +105,14 @@ class WanI2V_14B_480P_SamplingParam(WanT2V_1_3B_SamplingParams):
num_inference_steps: int = 50
# num_inference_steps: int = 40
# Wan I2V 480P supported resolutions (override parent)
supported_resolutions: list[tuple[int, int]] | None = field(
default_factory=lambda: [
(832, 480), # 16:9
(480, 832), # 9:16
]
)
teacache_params: WanTeaCacheParams = field(
default_factory=lambda: WanTeaCacheParams(
teacache_thresh=0.26,
@@ -115,6 +141,16 @@ class WanI2V_14B_720P_SamplingParam(WanT2V_14B_SamplingParams):
num_inference_steps: int = 50
# num_inference_steps: int = 40
# Wan I2V 720P supported resolutions (override parent)
supported_resolutions: list[tuple[int, int]] | None = field(
default_factory=lambda: [
(1280, 720), # 16:9
(720, 1280), # 9:16
(832, 480), # 16:9
(480, 832), # 9:16
]
)
teacache_params: WanTeaCacheParams = field(
default_factory=lambda: WanTeaCacheParams(
teacache_thresh=0.3,
@@ -188,6 +224,14 @@ class Wan2_2_TI2V_5B_SamplingParam(Wan2_2_Base_SamplingParams):
guidance_scale: float = 5.0
num_inference_steps: int = 50
# Wan2.2 TI2V 5B supported resolutions
supported_resolutions: list[tuple[int, int]] | None = field(
default_factory=lambda: [
(1280, 704), # 16:9-ish
(704, 1280), # 9:16-ish
]
)
@dataclass
class Wan2_2_T2V_A14B_SamplingParam(Wan2_2_Base_SamplingParams):
@@ -198,6 +242,16 @@ class Wan2_2_T2V_A14B_SamplingParam(Wan2_2_Base_SamplingParams):
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
# can be overridden during sampling
# Wan2.2 T2V A14B supported resolutions
supported_resolutions: list[tuple[int, int]] | None = field(
default_factory=lambda: [
(1280, 720), # 16:9
(720, 1280), # 9:16
(832, 480), # 16:9
(480, 832), # 9:16
]
)
@dataclass
class Wan2_2_I2V_A14B_SamplingParam(Wan2_2_Base_SamplingParams):
@@ -208,6 +262,16 @@ class Wan2_2_I2V_A14B_SamplingParam(Wan2_2_Base_SamplingParams):
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
# can be overridden during sampling
# Wan2.2 I2V A14B supported resolutions
supported_resolutions: list[tuple[int, int]] | None = field(
default_factory=lambda: [
(1280, 720), # 16:9
(720, 1280), # 9:16
(832, 480), # 16:9
(480, 832), # 9:16
]
)
# =============================================
# ============= Causal Self-Forcing =============