From 24d3de21260df126aab2a18af52168a8cd8b6ef4 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 14 Jul 2025 17:27:10 +0300 Subject: [PATCH] Add EasyCache --- nodes.py | 51 +++++++++++++++++++++++- wanvideo/modules/model.py | 84 ++++++++++++++++++++++++++++++++++++++- 2 files changed, 131 insertions(+), 4 deletions(-) diff --git a/nodes.py b/nodes.py index 23c123d..da23542 100644 --- a/nodes.py +++ b/nodes.py @@ -137,7 +137,6 @@ Official recommended values https://github.com/ali-vilab/TeaCache/tree/main/TeaC +-------------------+--------+---------+--------+ """ - EXPERIMENTAL = True def process(self, rel_l1_thresh, start_step, end_step, cache_device, use_coefficients, mode="e"): if cache_device == "main_device": @@ -189,6 +188,39 @@ class WanVideoMagCache: "cache_device": cache_device, } return (cache_args,) + +class WanVideoEasyCache: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "easycache_thresh": ("FLOAT", {"default": 0.015, "min": 0.0, "max": 0.3, "step": 0.001, "tooltip": "How strongly to cache the output of diffusion model. This value must be non-negative."}), + "start_step": ("INT", {"default": 1, "min": 10, "max": 9999, "step": 1, "tooltip": "Step to start applying EasyCache"}), + "end_step": ("INT", {"default": -1, "min": -1, "max": 9999, "step": 1, "tooltip": "Step to end applying EasyCache"}), + "cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}), + }, + } + RETURN_TYPES = ("CACHEARGS",) + RETURN_NAMES = ("cache_args",) + FUNCTION = "setargs" + CATEGORY = "WanVideoWrapper" + EXPERIMENTAL = True + DESCRIPTION = "EasyCache for WanVideoWrapper, source https://github.com/H-EmbodVis/EasyCache" + + def setargs(self, easycache_thresh, start_step, end_step, cache_device): + if cache_device == "main_device": + cache_device = mm.get_torch_device() + else: + cache_device = mm.unet_offload_device() + + cache_args = { + "cache_type": "EasyCache", + "easycache_thresh": easycache_thresh, + "start_step": start_step, + "end_step": end_step, + "cache_device": cache_device, + } + return (cache_args,) class WanVideoEnhanceAVideo: @classmethod @@ -2428,6 +2460,13 @@ class WanVideoSampler: transformer.magcache_end_step = len(timesteps)-1 if cache_args["end_step"] == -1 else cache_args["end_step"] transformer.magcache_thresh = cache_args["magcache_thresh"] transformer.magcache_K = cache_args["magcache_K"] + elif cache_args["cache_type"] == "EasyCache": + log.info(f"EasyCache: Using cache device: {transformer.cache_device}") + transformer.easycache_state.clear_all() + transformer.enable_easycache = True + transformer.easycache_start_step = cache_args["start_step"] + transformer.easycache_end_step = len(timesteps)-1 if cache_args["end_step"] == -1 else cache_args["end_step"] + transformer.easycache_thresh = cache_args["easycache_thresh"] if slg_args is not None: assert batched_cfg is not None, "Batched cfg is not supported with SLG" @@ -3430,7 +3469,12 @@ class WanVideoSampler: if cache_args is not None: cache_type = cache_args["cache_type"] - states = transformer.teacache_state.states if cache_type == "TeaCache" else transformer.magcache_state.states + states = ( + transformer.teacache_state.states if cache_type == "TeaCache" else + transformer.magcache_state.states if cache_type == "MagCache" else + transformer.easycache_state.states if cache_type == "EasyCache" else + None + ) state_names = { 0: "conditional", 1: "unconditional" @@ -3441,6 +3485,7 @@ class WanVideoSampler: log.info(f"{cache_type} skipped: {len(state['skipped_steps'])} {name} steps: {state['skipped_steps']}") transformer.teacache_state.clear_all() transformer.magcache_state.clear_all() + transformer.easycache_state.clear_all() del states if force_offload: @@ -3690,6 +3735,7 @@ NODE_CLASS_MAPPINGS = { "WanVideoContextOptions": WanVideoContextOptions, "WanVideoTeaCache": WanVideoTeaCache, "WanVideoMagCache": WanVideoMagCache, + "WanVideoEasyCache": WanVideoEasyCache, "WanVideoVRAMManagement": WanVideoVRAMManagement, "WanVideoTextEmbedBridge": WanVideoTextEmbedBridge, "WanVideoFlowEdit": WanVideoFlowEdit, @@ -3728,6 +3774,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoContextOptions": "WanVideo Context Options", "WanVideoTeaCache": "WanVideo TeaCache", "WanVideoMagCache": "WanVideo MagCache", + "WanVideoEasyCache": "WanVideo EasyCache", "WanVideoVRAMManagement": "WanVideo VRAM Management", "WanVideoTextEmbedBridge": "WanVideo TextEmbed Bridge", "WanVideoFlowEdit": "WanVideo FlowEdit", diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index f696191..5c08e27 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -1058,6 +1058,13 @@ class WanModel(ModelMixin, ConfigMixin): self.magcache_end_step = -1 self.magcache_ratios = magcache_ratios + #init EasyCache variables + self.enable_easycache = False + self.easycache_thresh = 0.1 + self.easycache_start_step = 0 + self.easycache_end_step = -1 + self.easycache_state = EasyCacheState(cache_device=self.cache_device) + self.slg_blocks = None self.slg_start_percent = 0.0 self.slg_end_percent = 1.0 @@ -1578,6 +1585,7 @@ class WanModel(ModelMixin, ConfigMixin): token_ref_target_masks = token_ref_target_masks.to(x.dtype).to(device) should_calc = True + #TeaCache if self.enable_teacache and self.teacache_start_step <= current_step <= self.teacache_end_step: accumulated_rel_l1_distance = torch.tensor(0.0, dtype=torch.float32, device=device) if pred_id is None: @@ -1619,7 +1627,7 @@ class WanModel(ModelMixin, ConfigMixin): ) self.teacache_state.get(pred_id)['skipped_steps'].append(current_step) - # enable magcache + # MagCache if self.enable_magcache and self.magcache_start_step <= current_step <= self.magcache_end_step: if pred_id is None: pred_id = self.magcache_state.new_prediction(cache_device=self.cache_device) @@ -1656,8 +1664,38 @@ class WanModel(ModelMixin, ConfigMixin): accumulated_err=0 ) + # EasyCache + if self.enable_easycache and self.easycache_start_step <= current_step <= self.easycache_end_step: + if pred_id is None: + pred_id = self.easycache_state.new_prediction(cache_device=self.cache_device) + should_calc = True + else: + state = self.easycache_state.get(pred_id) + previous_raw_input = state.get('previous_raw_input') + previous_raw_output = state.get('previous_raw_output') + cache = state.get('cache') + accumulated_error = state.get('accumulated_error') + + if previous_raw_input is not None and previous_raw_output is not None: + raw_input = x.clone() + # Calculate input change + raw_input_change = (raw_input - previous_raw_input.to(raw_input.device)).abs().mean() + + accumulated_error += raw_input_change + + # Predict output change + if accumulated_error < self.easycache_thresh: + should_calc = False + x = raw_input + cache.to(x.device) + self.easycache_state.get(pred_id)['skipped_steps'].append(current_step) + else: + should_calc = True + accumulated_error = 0.0 + else: + should_calc = True + if should_calc: - if self.enable_teacache or self.enable_magcache: + if self.enable_teacache or self.enable_magcache or self.enable_easycache: original_x = x.to(self.cache_device).clone() if hasattr(self, "dwpose_embedding") and unianim_data is not None: @@ -1756,6 +1794,14 @@ class WanModel(ModelMixin, ConfigMixin): pred_id, residual_cache=(x.to(original_x.device) - original_x) ) + elif self.enable_easycache and (self.easycache_start_step <= current_step <= self.easycache_end_step) and pred_id is not None: + self.easycache_state.update( + pred_id, + previous_raw_input=original_x, + previous_raw_output=x.clone(), + cache=x.to(original_x.device) - original_x, + accumulated_error=0.0 + ) if self.ref_conv is not None and fun_ref is not None: full_ref_length = fun_ref.size(1) @@ -1863,6 +1909,40 @@ class MagCacheState: self.states = {} self._next_pred_id = 0 +class EasyCacheState: + def __init__(self, cache_device='cpu'): + self.cache_device = cache_device + self.states = {} + self._next_pred_id = 0 + + def new_prediction(self, cache_device='cpu'): + """Create a new prediction state and return its ID.""" + self.cache_device = cache_device + pred_id = self._next_pred_id + self._next_pred_id += 1 + self.states[pred_id] = { + 'previous_raw_input': None, + 'previous_raw_output': None, + 'cache': None, + 'accumulated_error': 0.0, + 'skipped_steps': [], + } + return pred_id + + def update(self, pred_id, **kwargs): + """Update state for a specific prediction.""" + if pred_id not in self.states: + return None + for key, value in kwargs.items(): + self.states[pred_id][key] = value + + def get(self, pred_id): + return self.states.get(pred_id, {}) + + def clear_all(self): + self.states = {} + self._next_pred_id = 0 + def relative_l1_distance(last_tensor, current_tensor): l1_distance = torch.abs(last_tensor.to(current_tensor.device) - current_tensor).mean() norm = torch.abs(last_tensor).mean()