Add EasyCache
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user