Add device option for teacache

This commit is contained in:
kijai
2025-03-08 01:20:20 +02:00
parent e332935ad6
commit 02b875b717
2 changed files with 12 additions and 3 deletions
+3 -2
View File
@@ -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)
+9 -1
View File
@@ -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