mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-08-23 23:51:23 +08:00
Add subsampling for heuristic inputs
This commit is contained in:
parent
0784085cc5
commit
7f0380f8a1
@ -16,7 +16,7 @@ def easycache_forward_wrapper(executor, *args, **kwargs):
|
|||||||
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
|
||||||
next_x_prev = x.clone()
|
next_x_prev = x
|
||||||
do_easycache = easycache.should_do_easycache(sigmas)
|
do_easycache = easycache.should_do_easycache(sigmas)
|
||||||
logging.info(f"easycache_wrapper: do_easycache: {do_easycache}")
|
logging.info(f"easycache_wrapper: do_easycache: {do_easycache}")
|
||||||
output_prev_norm = None
|
output_prev_norm = None
|
||||||
@ -26,9 +26,9 @@ def easycache_forward_wrapper(executor, *args, **kwargs):
|
|||||||
easycache.first_cond_uuid = transformer_options["uuids"][0]
|
easycache.first_cond_uuid = transformer_options["uuids"][0]
|
||||||
easycache.initial_step = False
|
easycache.initial_step = False
|
||||||
if easycache.has_x_prev():
|
if easycache.has_x_prev():
|
||||||
input_change = (x - easycache.x_prev).flatten().abs().mean()
|
input_change = (easycache.subsample(x, clone=False) - easycache.x_prev_subsampled).flatten().abs().mean()
|
||||||
if easycache.has_output_prev() and easycache.has_relative_transformation_rate():
|
if easycache.has_output_prev() and easycache.has_relative_transformation_rate():
|
||||||
output_prev_norm = easycache.output_prev.flatten().abs().mean()
|
output_prev_norm = easycache.output_prev_norm
|
||||||
approx_output_change_rate = (easycache.relative_transformation_rate * input_change) / output_prev_norm
|
approx_output_change_rate = (easycache.relative_transformation_rate * input_change) / 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:
|
||||||
@ -43,7 +43,7 @@ def easycache_forward_wrapper(executor, *args, **kwargs):
|
|||||||
if easycache.has_output_prev():
|
if easycache.has_output_prev():
|
||||||
output_change = (output - easycache.output_prev).flatten().abs().mean()
|
output_change = (output - easycache.output_prev).flatten().abs().mean()
|
||||||
if output_prev_norm is None:
|
if output_prev_norm is None:
|
||||||
output_prev_norm = easycache.output_prev.flatten().abs().mean()
|
output_prev_norm = easycache.output_prev_norm
|
||||||
output_change_rate = output_change / 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():
|
||||||
@ -54,7 +54,9 @@ def easycache_forward_wrapper(executor, *args, **kwargs):
|
|||||||
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}")
|
logging.info(f"easycache_wrapper: output_change_rate: {output_change_rate}")
|
||||||
easycache.cache_diff = output - next_x_prev
|
easycache.cache_diff = output - next_x_prev
|
||||||
easycache.x_prev = 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 = output.clone()
|
easycache.output_prev = output.clone()
|
||||||
return output
|
return output
|
||||||
|
|
||||||
@ -82,10 +84,11 @@ def easycache_sample_wrapper(executor, *args, **kwargs):
|
|||||||
|
|
||||||
|
|
||||||
class EasyCacheHolder:
|
class EasyCacheHolder:
|
||||||
def __init__(self, reuse_threshold: float, start_percent: float, end_percent: float):
|
def __init__(self, reuse_threshold: float, start_percent: float, end_percent: float, subsample_factor: int):
|
||||||
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
|
||||||
# timestep values
|
# timestep values
|
||||||
self.start_t = 0.0
|
self.start_t = 0.0
|
||||||
self.end_t = 0.0
|
self.end_t = 0.0
|
||||||
@ -96,8 +99,9 @@ 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 = None
|
self.x_prev_subsampled = None
|
||||||
self.output_prev = None
|
self.output_prev = None
|
||||||
|
self.output_prev_norm = None
|
||||||
self.cache_diff = None
|
self.cache_diff = None
|
||||||
self.output_change_rates = []
|
self.output_change_rates = []
|
||||||
self.approx_output_change_rates = []
|
self.approx_output_change_rates = []
|
||||||
@ -109,11 +113,14 @@ class EasyCacheHolder:
|
|||||||
return (timestep[0] <= self.start_t).item()
|
return (timestep[0] <= self.start_t).item()
|
||||||
|
|
||||||
def has_x_prev(self) -> bool:
|
def has_x_prev(self) -> bool:
|
||||||
return self.x_prev is not None
|
return self.x_prev_subsampled is not None
|
||||||
|
|
||||||
def has_output_prev(self) -> bool:
|
def has_output_prev(self) -> bool:
|
||||||
return self.output_prev is not None
|
return self.output_prev is not None
|
||||||
|
|
||||||
|
def has_output_prev_norm(self) -> bool:
|
||||||
|
return self.output_prev_norm is not None
|
||||||
|
|
||||||
def has_cache_diff(self) -> bool:
|
def has_cache_diff(self) -> bool:
|
||||||
return self.cache_diff is not None
|
return self.cache_diff is not None
|
||||||
|
|
||||||
@ -125,6 +132,16 @@ 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:
|
||||||
|
if self.subsample_factor > 1:
|
||||||
|
to_return = x[..., ::self.subsample_factor, ::self.subsample_factor]
|
||||||
|
if clone:
|
||||||
|
return to_return.clone()
|
||||||
|
return to_return
|
||||||
|
if clone:
|
||||||
|
return x.clone()
|
||||||
|
return x
|
||||||
|
|
||||||
def apply_cache(self):
|
def apply_cache(self):
|
||||||
...
|
...
|
||||||
|
|
||||||
@ -138,16 +155,18 @@ class EasyCacheHolder:
|
|||||||
self.skip_current_step = False
|
self.skip_current_step = False
|
||||||
self.output_change_rates = []
|
self.output_change_rates = []
|
||||||
self.first_cond_uuid = None
|
self.first_cond_uuid = None
|
||||||
del self.x_prev
|
del self.x_prev_subsampled
|
||||||
self.x_prev = None
|
self.x_prev_subsampled = None
|
||||||
del self.output_prev
|
del self.output_prev
|
||||||
self.output_prev = None
|
self.output_prev = None
|
||||||
|
del self.output_prev_norm
|
||||||
|
self.output_prev_norm = None
|
||||||
del self.cache_diff
|
del self.cache_diff
|
||||||
self.cache_diff = None
|
self.cache_diff = None
|
||||||
return self
|
return self
|
||||||
|
|
||||||
def clone(self):
|
def clone(self):
|
||||||
return EasyCacheHolder(self.reuse_threshold, self.start_percent, self.end_percent)
|
return EasyCacheHolder(self.reuse_threshold, self.start_percent, self.end_percent, self.subsample_factor)
|
||||||
|
|
||||||
|
|
||||||
class EasyCacheNode(io.ComfyNode):
|
class EasyCacheNode(io.ComfyNode):
|
||||||
@ -163,6 +182,7 @@ class EasyCacheNode(io.ComfyNode):
|
|||||||
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.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("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("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.Int.Input("subsample_factor", min=1, default=8, max=128, step=1, tooltip="The factor to subsample latents to cache by."),
|
||||||
],
|
],
|
||||||
outputs=[
|
outputs=[
|
||||||
io.Model.Output(tooltip="The model with EasyCache."),
|
io.Model.Output(tooltip="The model with EasyCache."),
|
||||||
@ -170,9 +190,9 @@ class EasyCacheNode(io.ComfyNode):
|
|||||||
)
|
)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def execute(cls, model: io.Model.Type, reuse_threshold: float, start_percent: float, end_percent: float) -> io.NodeOutput:
|
def execute(cls, model: io.Model.Type, reuse_threshold: float, start_percent: float, end_percent: float, subsample_factor: int) -> io.NodeOutput:
|
||||||
model = model.clone()
|
model = model.clone()
|
||||||
model.model_options["transformer_options"]["easycache"] = EasyCacheHolder(reuse_threshold, start_percent, end_percent)
|
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.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)
|
||||||
return io.NodeOutput(model)
|
return io.NodeOutput(model)
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user