File size: 1,459 Bytes
ec0a9aa | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 | from __future__ import annotations
from methods.cache_strategy.FasterCache.processor import SVDFasterCacheAttnProcessor
def enable_fastercache(
unet,
*,
start_step: int = 0,
model_interval: int = 5,
block_interval: int = 3,
first_layers_fp: int = 2,
) -> None:
runtime_state = {
"start_step": start_step,
"model_interval": model_interval,
"block_interval": block_interval,
"first_layers_fp": first_layers_fp,
"current_step_idx": -1,
"resolved_start_step": 0,
"delta_lf": None,
"delta_hf": None,
"block_histories": {},
"block_alpha": 0.3,
}
processors = {}
for layer_idx, name in enumerate(unet.attn_processors.keys()):
processors[name] = SVDFasterCacheAttnProcessor(
runtime_state=runtime_state,
layer_idx=layer_idx,
first_layers_fp=first_layers_fp,
)
unet._fastercache_state = runtime_state
unet.set_attn_processor(processors)
print(
"[FasterCache] Enabled on "
f"{len(processors)} attention layers (start_step={start_step}, "
f"model_interval={model_interval}, block_interval={block_interval})"
)
def disable_fastercache(unet) -> None:
if hasattr(unet, "_fastercache_state"):
delattr(unet, "_fastercache_state")
unet.set_default_attn_processor()
print("[FasterCache] Disabled - restored default attention processors.")
|