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
if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
if len(input.shape) == 3:
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)
#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(weight_dtype)
w = cls.weight.t()
scale = torch.ones((1), device=input.device, dtype=torch.float32)
@@ -876,6 +876,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
stg_block_idx=stg_block_idx,
stg_mode=stg_mode,
return_dict=True,
ref_latents=ref_latents
)["x"]
window_mask = torch.ones_like(noise_pred_context)
@@ -913,7 +914,8 @@ class HunyuanVideoPipeline(DiffusionPipeline):
stg_mode=stg_mode,
return_dict=True,
ref_latents=ref_latents,
is_uncond = False
is_uncond = False,
current_step = i,
)["x"]
else:
uncond = self.transformer(
@@ -929,7 +931,8 @@ class HunyuanVideoPipeline(DiffusionPipeline):
stg_mode=stg_mode,
return_dict=True,
ref_latents=uncond_ref_latents,
is_uncond = True
is_uncond = True,
current_step = i
)["x"]
cond = self.transformer(
latent_model_input[1].unsqueeze(0),
@@ -944,7 +947,8 @@ class HunyuanVideoPipeline(DiffusionPipeline):
stg_mode=stg_mode,
return_dict=True,
ref_latents=ref_latents,
is_uncond = False
is_uncond = False,
current_step = i
)["x"]
# perform guidance
+4 -1
View File
@@ -754,6 +754,8 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
self.num_steps = 0
self.teacache_skipped_steps_cond = 0
self.teacache_skipped_steps_uncond = 0
self.teacache_start_step = 0
self.teacache_end_step = 100
self.rel_l1_thresh = 0.15
self.accumulated_rel_l1_distance = 0
self.previous_modulated_input = None
@@ -953,6 +955,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
return_dict: bool = True,
ref_latents: torch.Tensor = None,
is_uncond = False,
current_step: int = 0,
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
def _process_double_blocks(img, txt, vec, block_args):
@@ -1079,7 +1082,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
]
#tea_cache
if self.enable_teacache:
if self.enable_teacache and self.teacache_start_step <= current_step <= self.teacache_end_step:
inp = img.clone()
vec_ = vec.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,
"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"}),
"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"
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":
teacache_device = mm.get_torch_device()
else:
teacache_device = mm.unet_offload_device()
teacache_args = {
"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,)
@@ -449,6 +453,7 @@ class HyVideoModelLoader:
if quantization == "fp8_e4m3fn_fast":
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)
elif quantization == "fp8_scaled":
from .hyvideo.modules.fp8_optimization import convert_fp8_linear
@@ -1378,10 +1383,17 @@ class HyVideoSampler:
transformer.last_dimensions = (height, width, num_frames)
transformer.last_frame_count = num_frames
transformer.teacache_device = device
transformer.teacache_start_step = 0
transformer.teacache_end_step = steps - 1
transformer.enable_teacache = True
transformer.num_steps = steps
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:
transformer.enable_teacache = False