Add TeaCache start/end step

This commit is contained in:
kijai
2025-05-09 11:41:31 +03:00
parent 64b5d31765
commit e18afda414
4 changed files with 28 additions and 8 deletions
+3 -2
View File
@@ -7,8 +7,9 @@ def fp8_linear_forward(cls, original_dtype, input):
weight_dtype = cls.weight.dtype weight_dtype = cls.weight.dtype
if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]: if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
if len(input.shape) == 3: if len(input.shape) == 3:
target_dtype = torch.float8_e5m2 if weight_dtype == torch.float8_e4m3fn else torch.float8_e4m3fn #target_dtype = torch.float8_e5m2 if weight_dtype == torch.float8_e4m3fn else torch.float8_e4m3fn
inn = input.reshape(-1, input.shape[2]).to(target_dtype) #inn = input.reshape(-1, input.shape[2]).to(target_dtype)
inn = input.reshape(-1, input.shape[2]).to(weight_dtype)
w = cls.weight.t() w = cls.weight.t()
scale = torch.ones((1), device=input.device, dtype=torch.float32) scale = torch.ones((1), device=input.device, dtype=torch.float32)
@@ -876,6 +876,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
stg_block_idx=stg_block_idx, stg_block_idx=stg_block_idx,
stg_mode=stg_mode, stg_mode=stg_mode,
return_dict=True, return_dict=True,
ref_latents=ref_latents
)["x"] )["x"]
window_mask = torch.ones_like(noise_pred_context) window_mask = torch.ones_like(noise_pred_context)
@@ -913,7 +914,8 @@ class HunyuanVideoPipeline(DiffusionPipeline):
stg_mode=stg_mode, stg_mode=stg_mode,
return_dict=True, return_dict=True,
ref_latents=ref_latents, ref_latents=ref_latents,
is_uncond = False is_uncond = False,
current_step = i,
)["x"] )["x"]
else: else:
uncond = self.transformer( uncond = self.transformer(
@@ -929,7 +931,8 @@ class HunyuanVideoPipeline(DiffusionPipeline):
stg_mode=stg_mode, stg_mode=stg_mode,
return_dict=True, return_dict=True,
ref_latents=uncond_ref_latents, ref_latents=uncond_ref_latents,
is_uncond = True is_uncond = True,
current_step = i
)["x"] )["x"]
cond = self.transformer( cond = self.transformer(
latent_model_input[1].unsqueeze(0), latent_model_input[1].unsqueeze(0),
@@ -944,7 +947,8 @@ class HunyuanVideoPipeline(DiffusionPipeline):
stg_mode=stg_mode, stg_mode=stg_mode,
return_dict=True, return_dict=True,
ref_latents=ref_latents, ref_latents=ref_latents,
is_uncond = False is_uncond = False,
current_step = i
)["x"] )["x"]
# perform guidance # perform guidance
+4 -1
View File
@@ -754,6 +754,8 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
self.num_steps = 0 self.num_steps = 0
self.teacache_skipped_steps_cond = 0 self.teacache_skipped_steps_cond = 0
self.teacache_skipped_steps_uncond = 0 self.teacache_skipped_steps_uncond = 0
self.teacache_start_step = 0
self.teacache_end_step = 100
self.rel_l1_thresh = 0.15 self.rel_l1_thresh = 0.15
self.accumulated_rel_l1_distance = 0 self.accumulated_rel_l1_distance = 0
self.previous_modulated_input = None self.previous_modulated_input = None
@@ -953,6 +955,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
return_dict: bool = True, return_dict: bool = True,
ref_latents: torch.Tensor = None, ref_latents: torch.Tensor = None,
is_uncond = False, is_uncond = False,
current_step: int = 0,
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]: ) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
def _process_double_blocks(img, txt, vec, block_args): def _process_double_blocks(img, txt, vec, block_args):
@@ -1079,7 +1082,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
] ]
#tea_cache #tea_cache
if self.enable_teacache: if self.enable_teacache and self.teacache_start_step <= current_step <= self.teacache_end_step:
inp = img.clone() inp = img.clone()
vec_ = vec.clone() vec_ = vec.clone()
txt_ = txt.clone() txt_ = txt.clone()
+14 -2
View File
@@ -221,6 +221,8 @@ class HyVideoTeaCache:
"rel_l1_thresh": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.01, "rel_l1_thresh": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.01,
"tooltip": "Higher values will make TeaCache more aggressive, faster, but may cause artifacts"}), "tooltip": "Higher values will make TeaCache more aggressive, faster, but may cause artifacts"}),
"cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}), "cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}),
"start_step": ("INT", {"default": 0, "min": 0, "max": 100, "step": 1, "tooltip": "Start step to apply TeaCache"}),
"end_step": ("INT", {"default": -1, "min": -1, "max": 100, "step": 1, "tooltip": "End step to apply TeaCache"}),
}, },
} }
@@ -230,14 +232,16 @@ class HyVideoTeaCache:
CATEGORY = "HunyuanVideoWrapper" CATEGORY = "HunyuanVideoWrapper"
DESCRIPTION = "TeaCache settings for HunyuanVideo to speed up inference" DESCRIPTION = "TeaCache settings for HunyuanVideo to speed up inference"
def process(self, rel_l1_thresh, cache_device): def process(self, rel_l1_thresh, cache_device, start_step, end_step):
if cache_device == "main_device": if cache_device == "main_device":
teacache_device = mm.get_torch_device() teacache_device = mm.get_torch_device()
else: else:
teacache_device = mm.unet_offload_device() teacache_device = mm.unet_offload_device()
teacache_args = { teacache_args = {
"rel_l1_thresh": rel_l1_thresh, "rel_l1_thresh": rel_l1_thresh,
"cache_device": teacache_device "cache_device": teacache_device,
"start_step": start_step,
"end_step": end_step
} }
return (teacache_args,) return (teacache_args,)
@@ -449,6 +453,7 @@ class HyVideoModelLoader:
if quantization == "fp8_e4m3fn_fast": if quantization == "fp8_e4m3fn_fast":
from .fp8_optimization import convert_fp8_linear from .fp8_optimization import convert_fp8_linear
params_to_keep.update({"mlp", "modulation", "mod"})
convert_fp8_linear(patcher.model.diffusion_model, base_dtype, params_to_keep=params_to_keep) convert_fp8_linear(patcher.model.diffusion_model, base_dtype, params_to_keep=params_to_keep)
elif quantization == "fp8_scaled": elif quantization == "fp8_scaled":
from .hyvideo.modules.fp8_optimization import convert_fp8_linear from .hyvideo.modules.fp8_optimization import convert_fp8_linear
@@ -1378,10 +1383,17 @@ class HyVideoSampler:
transformer.last_dimensions = (height, width, num_frames) transformer.last_dimensions = (height, width, num_frames)
transformer.last_frame_count = num_frames transformer.last_frame_count = num_frames
transformer.teacache_device = device transformer.teacache_device = device
transformer.teacache_start_step = 0
transformer.teacache_end_step = steps - 1
transformer.enable_teacache = True transformer.enable_teacache = True
transformer.num_steps = steps transformer.num_steps = steps
transformer.rel_l1_thresh = teacache_args["rel_l1_thresh"] transformer.rel_l1_thresh = teacache_args["rel_l1_thresh"]
transformer.teacache_start_step = teacache_args["start_step"]
teacache_end_step = teacache_args["end_step"]
if teacache_end_step < 0:
teacache_end_step = steps - 1
transformer.teacache_end_step = teacache_end_step
else: else:
transformer.enable_teacache = False transformer.enable_teacache = False