From 02b875b7178a54731e30587144e31a05df1a28c2 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 8 Mar 2025 01:20:20 +0200 Subject: [PATCH] Add device option for teacache --- hyvideo/modules/models.py | 5 +++-- nodes.py | 10 +++++++++- 2 files changed, 12 insertions(+), 3 deletions(-) diff --git a/hyvideo/modules/models.py b/hyvideo/modules/models.py index cb61c84..364304f 100644 --- a/hyvideo/modules/models.py +++ b/hyvideo/modules/models.py @@ -759,6 +759,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): self.previous_residual = None self.last_dimensions = None self.last_frame_count = None + self.teacache_device = None # thanks @2kpr for the initial block swap code! def block_swap(self, double_blocks_to_swap, single_blocks_to_swap, offload_txt_in=False, offload_img_in=False): @@ -1108,7 +1109,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): self.teacache_skipped_steps += 1 # Verify tensor dimensions match before adding if img.shape == self.previous_residual.shape: - img = img + self.previous_residual + img = img + self.previous_residual.to(img.device) else: should_calc = True # Force recalculation if dimensions don't match @@ -1121,7 +1122,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): x = _process_single_blocks(x, vec, txt.shape[1], block_args, stg_mode, stg_block_idx) img = x[:, :img_seq_len, ...] - self.previous_residual = img - ori_img + self.previous_residual = (img - ori_img).to(self.teacache_device) else: # Pass through DiT blocks img, txt = _process_double_blocks(img, txt, vec, block_args) diff --git a/nodes.py b/nodes.py index 2f49d94..8907094 100644 --- a/nodes.py +++ b/nodes.py @@ -219,6 +219,8 @@ class HyVideoTeaCache: "required": { "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"}), + }, } RETURN_TYPES = ("TEACACHEARGS",) @@ -227,9 +229,14 @@ class HyVideoTeaCache: CATEGORY = "HunyuanVideoWrapper" DESCRIPTION = "TeaCache settings for HunyuanVideo to speed up inference" - def process(self, rel_l1_thresh): + def process(self, rel_l1_thresh, cache_device): + 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 } return (teacache_args,) @@ -1318,6 +1325,7 @@ class HyVideoSampler: transformer.previous_residual = None transformer.last_dimensions = (height, width, num_frames) transformer.last_frame_count = num_frames + transformer.teacache_device = device transformer.enable_teacache = True transformer.num_steps = steps