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
|
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
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user