Add EasyCache

This commit is contained in:
kijai
2025-07-14 17:27:10 +03:00
parent e1f5d185ee
commit 24d3de2126
2 changed files with 131 additions and 4 deletions
+49 -2
View File
@@ -137,7 +137,6 @@ Official recommended values https://github.com/ali-vilab/TeaCache/tree/main/TeaC
+-------------------+--------+---------+--------+
</pre>
"""
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",
+82 -2
View File
@@ -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()