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.")