Properly consider conds in EasyCache logic

This commit is contained in:
Jedrzej Kosinski 2025-08-20 17:27:14 -07:00
parent 129ad27062
commit 136654c48b

View File

@ -1,72 +1,89 @@
from __future__ import annotations
from typing import TYPE_CHECKING
from comfy_api.latest import io, ComfyExtension from comfy_api.latest import io, ComfyExtension
import comfy.patcher_extension import comfy.patcher_extension
import logging import logging
import torch import torch
import comfy.model_patcher import comfy.model_patcher
if TYPE_CHECKING:
from uuid import UUID
def easycache_forward_wrapper(executor, *args, **kwargs): def easycache_forward_wrapper(executor, *args, **kwargs):
# get values from args # get values from args
x: torch.Tensor = args[0] x: torch.Tensor = args[0]
transformer_options: dict[str] = args[-1] transformer_options: dict[str] = args[-1]
# x: torch.Tensor = args[0]
# timestep: torch.Tensor = args[4]
# transformer_options: dict[str] = args[-2]
easycache: EasyCacheHolder = transformer_options["easycache"] easycache: EasyCacheHolder = transformer_options["easycache"]
sigmas = transformer_options["sigmas"] sigmas = transformer_options["sigmas"]
uuids = transformer_options["uuids"]
if sigmas is not None and easycache.is_past_end_timestep(sigmas): if sigmas is not None and easycache.is_past_end_timestep(sigmas):
return executor(*args, **kwargs) return executor(*args, **kwargs)
# prepare next x_prev # prepare next x_prev
has_first_cond_uuid = easycache.has_first_cond_uuid(uuids)
next_x_prev = x next_x_prev = x
do_easycache = easycache.should_do_easycache(sigmas)
logging.info(f"easycache_wrapper: do_easycache: {do_easycache}")
output_prev_norm = None
input_change = None input_change = None
do_easycache = easycache.should_do_easycache(sigmas)
if do_easycache: if do_easycache:
# if first cond marked this step for skipping, skip it and use appropriate cached values
if easycache.skip_current_step:
return easycache.apply_cache_diff(x, uuids)
if easycache.initial_step: if easycache.initial_step:
easycache.first_cond_uuid = transformer_options["uuids"][0] easycache.first_cond_uuid = uuids[0]
has_first_cond_uuid = easycache.has_first_cond_uuid(uuids)
easycache.initial_step = False easycache.initial_step = False
if easycache.has_x_prev_subsampled(): if has_first_cond_uuid:
input_change = (easycache.subsample(x, clone=False) - easycache.x_prev_subsampled).flatten().abs().mean() if easycache.has_x_prev_subsampled():
if easycache.has_output_prev_norm() and easycache.has_relative_transformation_rate(): input_change = (easycache.subsample(x, uuids, clone=False) - easycache.x_prev_subsampled).flatten().abs().mean()
output_prev_norm = easycache.output_prev_norm if easycache.has_output_prev_norm() and easycache.has_relative_transformation_rate():
approx_output_change_rate = (easycache.relative_transformation_rate * input_change) / output_prev_norm approx_output_change_rate = (easycache.relative_transformation_rate * input_change) / easycache.output_prev_norm
easycache.cumulative_change_rate += approx_output_change_rate easycache.cumulative_change_rate += approx_output_change_rate
if easycache.cumulative_change_rate < easycache.reuse_threshold: if easycache.cumulative_change_rate < easycache.reuse_threshold:
logging.info(f"easycache_wrapper: skipping step; cumulative_change_rate: {easycache.cumulative_change_rate}, reuse_threshold: {easycache.reuse_threshold}") if easycache.verbose:
return x + easycache.cache_diff logging.info(f"EasyCache [verbose] - skipping step; cumulative_change_rate: {easycache.cumulative_change_rate}, reuse_threshold: {easycache.reuse_threshold}")
else: # other conds should also skip this step, and instead use their cached values
logging.info(f"easycache_wrapper: NOT skipping step; cumulative_change_rate: {easycache.cumulative_change_rate}, reuse_threshold: {easycache.reuse_threshold}") easycache.skip_current_step = True
logging.info(f"easycache_wrapper: approx_output_change_rate: {approx_output_change_rate}") return easycache.apply_cache_diff(x, uuids)
easycache.cumulative_change_rate = 0.0 else:
if easycache.verbose:
logging.info(f"EasyCache [verbose] - NOT skipping step; cumulative_change_rate: {easycache.cumulative_change_rate}, reuse_threshold: {easycache.reuse_threshold}")
easycache.cumulative_change_rate = 0.0
output: torch.Tensor = executor(*args, **kwargs) output: torch.Tensor = executor(*args, **kwargs)
if easycache.has_output_prev_norm(): if has_first_cond_uuid and easycache.has_output_prev_norm():
output_change = (easycache.subsample(output, clone=False) - easycache.output_prev_subsampled).flatten().abs().mean() output_change = (easycache.subsample(output, uuids, clone=False) - easycache.output_prev_subsampled).flatten().abs().mean()
if output_prev_norm is None: if easycache.verbose:
output_prev_norm = easycache.output_prev_norm output_change_rate = output_change / easycache.output_prev_norm
output_change_rate = output_change / output_prev_norm easycache.output_change_rates.append(output_change_rate.item())
easycache.output_change_rates.append(output_change_rate.item())
if easycache.has_relative_transformation_rate(): if easycache.has_relative_transformation_rate():
approx_output_change_rate = (easycache.relative_transformation_rate * input_change) / output_prev_norm approx_output_change_rate = (easycache.relative_transformation_rate * input_change) / easycache.output_prev_norm
easycache.approx_output_change_rates.append(approx_output_change_rate.item()) easycache.approx_output_change_rates.append(approx_output_change_rate.item())
logging.info(f"easycache_wrapper: approx_output_change_rate: {approx_output_change_rate}") if easycache.verbose:
logging.info(f"EasyCache [verbose] - approx_output_change_rate: {approx_output_change_rate}")
if input_change is not None: if input_change is not None:
easycache.relative_transformation_rate = output_change / input_change easycache.relative_transformation_rate = output_change / input_change
logging.info(f"easycache_wrapper: output_change_rate: {output_change_rate}") if easycache.verbose:
easycache.cache_diff = output - next_x_prev logging.info(f"EasyCache [verbose] - output_change_rate: {output_change_rate}")
easycache.x_prev_subsampled = easycache.subsample(next_x_prev) # TODO: allow cache_diff to be offloaded
logging.info(f"easycache_wrapper: x_prev_subsampled: {easycache.x_prev_subsampled.shape}") easycache.update_cache_diff(output, next_x_prev, uuids)
easycache.output_prev_norm = output.flatten().abs().mean() if has_first_cond_uuid:
easycache.output_prev_subsampled = easycache.subsample(output) easycache.x_prev_subsampled = easycache.subsample(next_x_prev, uuids)
easycache.output_prev_subsampled = easycache.subsample(output, uuids)
easycache.output_prev_norm = output.flatten().abs().mean()
if easycache.verbose:
logging.info(f"EasyCache [verbose] - x_prev_subsampled: {easycache.x_prev_subsampled.shape}")
return output return output
def easycache_calc_cond_batch_wrapper(executor, *args, **kwargs): def easycache_calc_cond_batch_wrapper(executor, *args, **kwargs):
model_options = args[-1] model_options = args[-1]
easycache: EasyCacheHolder = model_options["transformer_options"]["easycache"] easycache: EasyCacheHolder = model_options["transformer_options"]["easycache"]
easycache.skip_current_step = False easycache.skip_current_step = False
# TODO: check if first_cond_uuid is active at this timestep; otherwise, EasyCache needs to be partially reset
return executor(*args, **kwargs) return executor(*args, **kwargs)
def easycache_sample_wrapper(executor, *args, **kwargs): def easycache_sample_wrapper(executor, *args, **kwargs):
"""
This OUTER_SAMPLE wrapper makes sure easycache is prepped for current run, and all memory usage is cleared at the end.
"""
try: try:
guider = executor.class_obj guider = executor.class_obj
orig_model_options = guider.model_options orig_model_options = guider.model_options
@ -75,20 +92,26 @@ def easycache_sample_wrapper(executor, *args, **kwargs):
guider.model_options["transformer_options"]["easycache"] = guider.model_options["transformer_options"]["easycache"].clone().prepare_timesteps(guider.model_patcher.model.model_sampling) guider.model_options["transformer_options"]["easycache"] = guider.model_options["transformer_options"]["easycache"].clone().prepare_timesteps(guider.model_patcher.model.model_sampling)
return executor(*args, **kwargs) return executor(*args, **kwargs)
finally: finally:
output_change_rates = guider.model_options['transformer_options']['easycache'].output_change_rates easycache: EasyCacheHolder = guider.model_options['transformer_options']['easycache']
approx_output_change_rates = guider.model_options['transformer_options']['easycache'].approx_output_change_rates output_change_rates = easycache.output_change_rates
logging.info(f"easycache_sample_wrapper: output_change_rates {len(output_change_rates)}: {output_change_rates}") approx_output_change_rates = easycache.approx_output_change_rates
logging.info(f"easycache_sample_wrapper: approx_output_change_rates {len(approx_output_change_rates)}: {approx_output_change_rates}") if easycache.verbose:
guider.model_options["transformer_options"]["easycache"].reset() logging.info(f"EasyCache [verbose] - output_change_rates {len(output_change_rates)}: {output_change_rates}")
logging.info(f"EasyCache [verbose] - approx_output_change_rates {len(approx_output_change_rates)}: {approx_output_change_rates}")
total_steps = len(args[3])-1
logging.info(f"EasyCache - skipped {easycache.total_steps_skipped}/{total_steps} steps ({total_steps/(total_steps-easycache.total_steps_skipped):.2f}x speedup).")
easycache.reset()
guider.model_options = orig_model_options guider.model_options = orig_model_options
class EasyCacheHolder: class EasyCacheHolder:
def __init__(self, reuse_threshold: float, start_percent: float, end_percent: float, subsample_factor: int): def __init__(self, reuse_threshold: float, start_percent: float, end_percent: float, subsample_factor: int, offload_cache_diff: bool, verbose: bool=False):
self.reuse_threshold = reuse_threshold self.reuse_threshold = reuse_threshold
self.start_percent = start_percent self.start_percent = start_percent
self.end_percent = end_percent self.end_percent = end_percent
self.subsample_factor = subsample_factor self.subsample_factor = subsample_factor
self.offload_cache_diff = offload_cache_diff
self.verbose = verbose
# timestep values # timestep values
self.start_t = 0.0 self.start_t = 0.0
self.end_t = 0.0 self.end_t = 0.0
@ -99,12 +122,13 @@ class EasyCacheHolder:
self.skip_current_step = False self.skip_current_step = False
# cache values # cache values
self.first_cond_uuid = None self.first_cond_uuid = None
self.x_prev_subsampled = None self.x_prev_subsampled: torch.Tensor = None
self.output_prev_subsampled = None self.output_prev_subsampled: torch.Tensor = None
self.output_prev_norm = None self.output_prev_norm: torch.Tensor = None
self.cache_diff = None self.uuid_cache_diffs: dict[UUID, torch.Tensor] = {}
self.output_change_rates = [] self.output_change_rates = []
self.approx_output_change_rates = [] self.approx_output_change_rates = []
self.total_steps_skipped = 0
def is_past_end_timestep(self, timestep: float) -> bool: def is_past_end_timestep(self, timestep: float) -> bool:
return not (timestep[0] > self.end_t).item() return not (timestep[0] > self.end_t).item()
@ -121,9 +145,6 @@ class EasyCacheHolder:
def has_output_prev_norm(self) -> bool: def has_output_prev_norm(self) -> bool:
return self.output_prev_norm is not None return self.output_prev_norm is not None
def has_cache_diff(self) -> bool:
return self.cache_diff is not None
def has_relative_transformation_rate(self) -> bool: def has_relative_transformation_rate(self) -> bool:
return self.relative_transformation_rate is not None return self.relative_transformation_rate is not None
@ -132,21 +153,35 @@ class EasyCacheHolder:
self.end_t = model_sampling.percent_to_sigma(self.end_percent) self.end_t = model_sampling.percent_to_sigma(self.end_percent)
return self return self
def subsample(self, x: torch.Tensor, clone: bool = True) -> torch.Tensor: def subsample(self, x: torch.Tensor, uuids: list[UUID], clone: bool = True) -> torch.Tensor:
batch_offset = x.shape[0] // len(uuids)
uuid_idx = uuids.index(self.first_cond_uuid)
if self.subsample_factor > 1: if self.subsample_factor > 1:
to_return = x[..., ::self.subsample_factor, ::self.subsample_factor] to_return = x[uuid_idx*batch_offset:(uuid_idx+1)*batch_offset, ..., ::self.subsample_factor, ::self.subsample_factor]
if clone: if clone:
return to_return.clone() return to_return.clone()
return to_return return to_return
to_return = x[uuid_idx*batch_offset:(uuid_idx+1)*batch_offset, ...]
if clone: if clone:
return x.clone() return to_return.clone()
return to_return
def apply_cache_diff(self, x: torch.Tensor, uuids: list[UUID]):
if self.first_cond_uuid in uuids:
self.total_steps_skipped += 1
batch_offset = x.shape[0] // len(uuids)
for i, uuid in enumerate(uuids):
x[i*batch_offset:(i+1)*batch_offset, ...] += self.uuid_cache_diffs[uuid].to(x.device)
return x return x
def apply_cache(self): def update_cache_diff(self, output: torch.Tensor, x: torch.Tensor, uuids: list[UUID]):
... diff = output - x
batch_offset = diff.shape[0] // len(uuids)
for i, uuid in enumerate(uuids):
self.uuid_cache_diffs[uuid] = diff[i*batch_offset:(i+1)*batch_offset, ...]
def accumulate_change(self): def has_first_cond_uuid(self, uuids: list[UUID]) -> bool:
... return self.first_cond_uuid in uuids
def reset(self): def reset(self):
self.relative_transformation_rate = 0.0 self.relative_transformation_rate = 0.0
@ -161,12 +196,13 @@ class EasyCacheHolder:
self.output_prev_subsampled = None self.output_prev_subsampled = None
del self.output_prev_norm del self.output_prev_norm
self.output_prev_norm = None self.output_prev_norm = None
del self.cache_diff del self.uuid_cache_diffs
self.cache_diff = None self.uuid_cache_diffs = {}
self.total_steps_skipped = 0
return self return self
def clone(self): def clone(self):
return EasyCacheHolder(self.reuse_threshold, self.start_percent, self.end_percent, self.subsample_factor) return EasyCacheHolder(self.reuse_threshold, self.start_percent, self.end_percent, self.subsample_factor, self.offload_cache_diff, self.verbose)
class EasyCacheNode(io.ComfyNode): class EasyCacheNode(io.ComfyNode):
@ -174,15 +210,16 @@ class EasyCacheNode(io.ComfyNode):
def define_schema(cls) -> io.Schema: def define_schema(cls) -> io.Schema:
return io.Schema( return io.Schema(
node_id="EasyCache", node_id="EasyCache",
display_name="Easy Cache", display_name="EasyCache",
description="Easy Cache", description="Native EasyCache implementation.",
category="advanced/debug/model", category="advanced/debug/model",
inputs=[ inputs=[
io.Model.Input("model", tooltip="The model to add EasyCache to."), io.Model.Input("model", tooltip="The model to add EasyCache to."),
io.Float.Input("reuse_threshold", min=0.0, default=0.0, max=1.0, step=0.01, tooltip="The threshold for reusing cached steps."), io.Float.Input("reuse_threshold", min=0.0, default=0.2, max=1.0, step=0.01, tooltip="The threshold for reusing cached steps."),
io.Float.Input("start_percent", min=0.0, default=0.0, max=1.0, step=0.01, tooltip="The relative sampling step to begin use of EasyCache."), io.Float.Input("start_percent", min=0.0, default=0.15, max=1.0, step=0.01, tooltip="The relative sampling step to begin use of EasyCache."),
io.Float.Input("end_percent", min=0.0, default=1.0, max=1.0, step=0.01, tooltip="The relative sampling step to end use of EasyCache."), io.Float.Input("end_percent", min=0.0, default=0.95, max=1.0, step=0.01, tooltip="The relative sampling step to end use of EasyCache."),
io.Int.Input("subsample_factor", min=1, default=8, max=128, step=1, tooltip="The factor to subsample latents to cache by."), io.Int.Input("subsample_factor", min=1, default=8, max=128, step=1, tooltip="The factor to subsample latents to cache by."),
io.Boolean.Input("verbose", default=False, tooltip="Whether to log verbose information."),
], ],
outputs=[ outputs=[
io.Model.Output(tooltip="The model with EasyCache."), io.Model.Output(tooltip="The model with EasyCache."),
@ -190,11 +227,12 @@ class EasyCacheNode(io.ComfyNode):
) )
@classmethod @classmethod
def execute(cls, model: io.Model.Type, reuse_threshold: float, start_percent: float, end_percent: float, subsample_factor: int) -> io.NodeOutput: def execute(cls, model: io.Model.Type, reuse_threshold: float, start_percent: float, end_percent: float, subsample_factor: int, verbose: bool) -> io.NodeOutput:
model = model.clone() model = model.clone()
model.model_options["transformer_options"]["easycache"] = EasyCacheHolder(reuse_threshold, start_percent, end_percent, subsample_factor) model.model_options["transformer_options"]["easycache"] = EasyCacheHolder(reuse_threshold, start_percent, end_percent, subsample_factor, offload_cache_diff=False, verbose=verbose)
model.add_wrapper_with_key(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, "easycache", easycache_forward_wrapper)
model.add_wrapper_with_key(comfy.patcher_extension.WrappersMP.OUTER_SAMPLE, "easycache", easycache_sample_wrapper) model.add_wrapper_with_key(comfy.patcher_extension.WrappersMP.OUTER_SAMPLE, "easycache", easycache_sample_wrapper)
model.add_wrapper_with_key(comfy.patcher_extension.WrappersMP.CALC_COND_BATCH, "easycache", easycache_calc_cond_batch_wrapper)
model.add_wrapper_with_key(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, "easycache", easycache_forward_wrapper)
return io.NodeOutput(model) return io.NodeOutput(model)
class EasyCacheExtension(ComfyExtension): class EasyCacheExtension(ComfyExtension):