[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:
@@ -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 =============
|
||||
|
||||
Reference in New Issue
Block a user