Add device option for teacache
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user