Skip to content

vllm_omni.diffusion.cache.seacache

Modules:

Name Description
backend
config
hook
sea_filter
state

SeaCacheBackend

Bases: CacheBackend

Backend for spectral-evolution-aware diffusion caching.

enable

enable(pipeline: Any) -> None

refresh

refresh(
    pipeline: Any,
    num_inference_steps: int,
    verbose: bool = True,
) -> None

SeaCacheConfig dataclass

Configuration for SeaCache.

Defaults are tuned for Cosmos3 and may require adjustment for other models.

max_consecutive_cached class-attribute instance-attribute

max_consecutive_cached: int = 2

power_exp class-attribute instance-attribute

power_exp: float = 3.0

residual_order class-attribute instance-attribute

residual_order: int = 1

threshold class-attribute instance-attribute

threshold: float = 0.25

SeaCacheRootHook

Bases: ModelHook

Drive SeaCache gating and transformer forward control.

config instance-attribute

config = config

current_sigma_callback instance-attribute

current_sigma_callback = current_sigma_callback

current_step_callback instance-attribute

current_step_callback = current_step_callback

extractor_fn instance-attribute

extractor_fn = extractor_fn

full_count instance-attribute

full_count = 0

num_inference_steps_callback instance-attribute

num_inference_steps_callback = num_inference_steps_callback

skip_count instance-attribute

skip_count = 0

state_manager instance-attribute

state_manager = StateManager(SeaCacheState)

cache_context

cache_context(name: str) -> Iterator[None]

initialize_hook

initialize_hook(module: Module) -> Module

new_forward

new_forward(
    module: Module, *args: Any, **kwargs: Any
) -> Any

refresh

refresh(module: Module) -> None

reset_state

reset_state(module: Module) -> Module

SeaCacheState dataclass

Per-cache-context trajectory state.

accumulated_distance class-attribute instance-attribute

accumulated_distance: float = 0.0

consecutive_cached class-attribute instance-attribute

consecutive_cached: int = 0

history class-attribute instance-attribute

history: list[tuple[int, Tensor]] = field(
    default_factory=list
)

last_step class-attribute instance-attribute

last_step: int | None = None

previous_indicator class-attribute instance-attribute

previous_indicator: list[Tensor] | None = None

reset

reset() -> None

apply_sea_cache_hook

apply_sea_cache_hook(
    module: Module,
    config: SeaCacheConfig,
    *,
    current_step_callback: Callable[[], int | Tensor | None]
    | None = None,
    current_sigma_callback: Callable[
        [], float | Tensor | None
    ]
    | None = None,
    num_inference_steps_callback: Callable[
        [], int | Tensor | None
    ]
    | None = None,
    extractor_fn: Callable[..., CacheContext] | None = None,
) -> SeaCacheRootHook