From 136654c48b6ed40fb7b7d00480ea3cbedcd64844 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 20 Aug 2025 17:27:14 -0700 Subject: [PATCH] Properly consider conds in EasyCache logic --- comfy_extras/nodes_easycache.py | 168 ++++++++++++++++++++------------ 1 file changed, 103 insertions(+), 65 deletions(-) diff --git a/comfy_extras/nodes_easycache.py b/comfy_extras/nodes_easycache.py index 61e16c992..acd5462f7 100644 --- a/comfy_extras/nodes_easycache.py +++ b/comfy_extras/nodes_easycache.py @@ -1,72 +1,89 @@ +from __future__ import annotations +from typing import TYPE_CHECKING from comfy_api.latest import io, ComfyExtension import comfy.patcher_extension import logging import torch import comfy.model_patcher +if TYPE_CHECKING: + from uuid import UUID + def easycache_forward_wrapper(executor, *args, **kwargs): # get values from args x: torch.Tensor = args[0] 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"] sigmas = transformer_options["sigmas"] + uuids = transformer_options["uuids"] if sigmas is not None and easycache.is_past_end_timestep(sigmas): return executor(*args, **kwargs) # prepare next x_prev + has_first_cond_uuid = easycache.has_first_cond_uuid(uuids) 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 + do_easycache = easycache.should_do_easycache(sigmas) 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: - 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 - if easycache.has_x_prev_subsampled(): - input_change = (easycache.subsample(x, clone=False) - easycache.x_prev_subsampled).flatten().abs().mean() - if easycache.has_output_prev_norm() and easycache.has_relative_transformation_rate(): - output_prev_norm = easycache.output_prev_norm - approx_output_change_rate = (easycache.relative_transformation_rate * input_change) / output_prev_norm - easycache.cumulative_change_rate += approx_output_change_rate - 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}") - return x + easycache.cache_diff - else: - logging.info(f"easycache_wrapper: NOT skipping step; cumulative_change_rate: {easycache.cumulative_change_rate}, reuse_threshold: {easycache.reuse_threshold}") - logging.info(f"easycache_wrapper: approx_output_change_rate: {approx_output_change_rate}") - easycache.cumulative_change_rate = 0.0 + if has_first_cond_uuid: + if easycache.has_x_prev_subsampled(): + input_change = (easycache.subsample(x, uuids, clone=False) - easycache.x_prev_subsampled).flatten().abs().mean() + if easycache.has_output_prev_norm() and easycache.has_relative_transformation_rate(): + approx_output_change_rate = (easycache.relative_transformation_rate * input_change) / easycache.output_prev_norm + easycache.cumulative_change_rate += approx_output_change_rate + if easycache.cumulative_change_rate < easycache.reuse_threshold: + if easycache.verbose: + logging.info(f"EasyCache [verbose] - skipping step; cumulative_change_rate: {easycache.cumulative_change_rate}, reuse_threshold: {easycache.reuse_threshold}") + # other conds should also skip this step, and instead use their cached values + easycache.skip_current_step = True + return easycache.apply_cache_diff(x, uuids) + 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) - if easycache.has_output_prev_norm(): - output_change = (easycache.subsample(output, clone=False) - easycache.output_prev_subsampled).flatten().abs().mean() - if output_prev_norm is None: - output_prev_norm = easycache.output_prev_norm - output_change_rate = output_change / output_prev_norm - easycache.output_change_rates.append(output_change_rate.item()) + if has_first_cond_uuid and easycache.has_output_prev_norm(): + output_change = (easycache.subsample(output, uuids, clone=False) - easycache.output_prev_subsampled).flatten().abs().mean() + if easycache.verbose: + output_change_rate = output_change / easycache.output_prev_norm + easycache.output_change_rates.append(output_change_rate.item()) 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()) - 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: easycache.relative_transformation_rate = output_change / input_change - logging.info(f"easycache_wrapper: output_change_rate: {output_change_rate}") - easycache.cache_diff = output - next_x_prev - easycache.x_prev_subsampled = easycache.subsample(next_x_prev) - logging.info(f"easycache_wrapper: x_prev_subsampled: {easycache.x_prev_subsampled.shape}") - easycache.output_prev_norm = output.flatten().abs().mean() - easycache.output_prev_subsampled = easycache.subsample(output) + if easycache.verbose: + logging.info(f"EasyCache [verbose] - output_change_rate: {output_change_rate}") + # TODO: allow cache_diff to be offloaded + easycache.update_cache_diff(output, next_x_prev, uuids) + if has_first_cond_uuid: + 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 def easycache_calc_cond_batch_wrapper(executor, *args, **kwargs): model_options = args[-1] easycache: EasyCacheHolder = model_options["transformer_options"]["easycache"] 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) 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: guider = executor.class_obj 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) return executor(*args, **kwargs) finally: - output_change_rates = guider.model_options['transformer_options']['easycache'].output_change_rates - approx_output_change_rates = guider.model_options['transformer_options']['easycache'].approx_output_change_rates - logging.info(f"easycache_sample_wrapper: output_change_rates {len(output_change_rates)}: {output_change_rates}") - logging.info(f"easycache_sample_wrapper: approx_output_change_rates {len(approx_output_change_rates)}: {approx_output_change_rates}") - guider.model_options["transformer_options"]["easycache"].reset() + easycache: EasyCacheHolder = guider.model_options['transformer_options']['easycache'] + output_change_rates = easycache.output_change_rates + approx_output_change_rates = easycache.approx_output_change_rates + if easycache.verbose: + 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 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.start_percent = start_percent self.end_percent = end_percent self.subsample_factor = subsample_factor + self.offload_cache_diff = offload_cache_diff + self.verbose = verbose # timestep values self.start_t = 0.0 self.end_t = 0.0 @@ -99,12 +122,13 @@ class EasyCacheHolder: self.skip_current_step = False # cache values self.first_cond_uuid = None - self.x_prev_subsampled = None - self.output_prev_subsampled = None - self.output_prev_norm = None - self.cache_diff = None + self.x_prev_subsampled: torch.Tensor = None + self.output_prev_subsampled: torch.Tensor = None + self.output_prev_norm: torch.Tensor = None + self.uuid_cache_diffs: dict[UUID, torch.Tensor] = {} self.output_change_rates = [] self.approx_output_change_rates = [] + self.total_steps_skipped = 0 def is_past_end_timestep(self, timestep: float) -> bool: return not (timestep[0] > self.end_t).item() @@ -121,9 +145,6 @@ class EasyCacheHolder: def has_output_prev_norm(self) -> bool: 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: 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) 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: - 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: return to_return.clone() return to_return + to_return = x[uuid_idx*batch_offset:(uuid_idx+1)*batch_offset, ...] 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 - 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): self.relative_transformation_rate = 0.0 @@ -161,12 +196,13 @@ class EasyCacheHolder: self.output_prev_subsampled = None del self.output_prev_norm self.output_prev_norm = None - del self.cache_diff - self.cache_diff = None + del self.uuid_cache_diffs + self.uuid_cache_diffs = {} + self.total_steps_skipped = 0 return 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): @@ -174,15 +210,16 @@ class EasyCacheNode(io.ComfyNode): def define_schema(cls) -> io.Schema: return io.Schema( node_id="EasyCache", - display_name="Easy Cache", - description="Easy Cache", + display_name="EasyCache", + description="Native EasyCache implementation.", category="advanced/debug/model", inputs=[ 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("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("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("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.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=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.Boolean.Input("verbose", default=False, tooltip="Whether to log verbose information."), ], outputs=[ io.Model.Output(tooltip="The model with EasyCache."), @@ -190,11 +227,12 @@ class EasyCacheNode(io.ComfyNode): ) @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_options["transformer_options"]["easycache"] = EasyCacheHolder(reuse_threshold, start_percent, end_percent, subsample_factor) - model.add_wrapper_with_key(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, "easycache", easycache_forward_wrapper) + 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.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) class EasyCacheExtension(ComfyExtension):