Add TeaCache start/end step
This commit is contained in:
+3
-2
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user