Skip to content

vllm_omni.diffusion.cache.cachedit.model_specific

Model-specific Cache-DiT adapters and enablers.

logger module-attribute

logger = init_logger(__name__)

BagelCachedAdapter

Bases: CachedAdapter

Custom CachedAdapter for Bagel that uses BagelCachedContextManager and BagelCachedBlocks.

collect_unified_blocks classmethod

collect_unified_blocks(
    block_adapter: BlockAdapter, contexts_kwargs: list[dict]
) -> list[dict[str, ModuleList]]

create_context classmethod

create_context(
    block_adapter: BlockAdapter, **context_kwargs
) -> tuple[list[str], list[dict[str, Any]]]

BagelCachedBlocks

Bases: CachedBlocks_Pattern_0_1_2

Custom CachedBlocks for Bagel that safely handles NaiveCache objects by adding isinstance checks in call_Mn_blocks and compute_or_prune.

call_Mn_blocks

call_Mn_blocks(
    hidden_states: Tensor,
    encoder_hidden_states: Tensor,
    *args,
    **kwargs,
)

compute_or_prune

compute_or_prune(
    block_id: int,
    block,
    hidden_states: Tensor,
    encoder_hidden_states: Tensor,
    *args,
    **kwargs,
)

BagelCachedContextManager

Bases: CachedContextManager

Custom CachedContextManager for Bagel that safely handles NaiveCache objects (mapped to encoder_hidden_states) by skipping tensor operations on them.

apply_cache

apply_cache(
    hidden_states: Tensor,
    encoder_hidden_states: Tensor = None,
    prefix: str = "Bn",
    encoder_prefix: str = "Bn_encoder",
) -> tuple[Tensor, Tensor | None]

SensenovaCachedAdapter

Bases: CachedAdapter

Custom CachedAdapter for SenseNova-U1 that uses SensenovaCachedBlocks.

collect_unified_blocks classmethod

collect_unified_blocks(
    block_adapter: BlockAdapter, contexts_kwargs: list[dict]
) -> list[dict[str, ModuleList]]

SensenovaCachedBlocks

Bases: CachedBlocks_Pattern_3_4_5

Custom CachedBlocks for SenseNova-U1 that only caches image-token hidden states during denoising.

forward

forward(hidden_states: Tensor, *args, **kwargs)

Wan22S2VCachedAdapter

Bases: CachedAdapter

CacheDiT adapter that uses Wan22S2VCachedBlocks for S2V audio injection.

Only overrides collect_unified_blocks to use Wan22S2VCachedBlocks (which calls after_transformer_block per-layer internally). The base class mock_transformer handles the forward wrapping — after_transformer_block is permanently replaced with a no-op in enable_cache_for_wan22_s2v() to prevent double injection.

collect_unified_blocks classmethod

collect_unified_blocks(
    block_adapter: BlockAdapter, contexts_kwargs: list[dict]
) -> list[dict[str, ModuleList]]

Wan22S2VCachedBlocks

Bases: CachedBlocks_Pattern_3_4_5

CacheDiT blocks wrapper that preserves S2V per-layer audio injection.

call_Bn_blocks

call_Bn_blocks(hidden_states: Tensor, *args, **kwargs)

call_Fn_blocks

call_Fn_blocks(hidden_states: Tensor, *args, **kwargs)

call_Mn_blocks

call_Mn_blocks(hidden_states: Tensor, *args, **kwargs)

call_blocks

call_blocks(hidden_states: Tensor, *args, **kwargs)

enable_cache_for_cosmos3

enable_cache_for_cosmos3(
    pipeline: Any, cache_config: Any
) -> RefreshCacheContextFunc

Enable cache-dit for Cosmos3.

Cosmos3 has a dual-pathway architecture (UND + GEN) but only the GEN pathway (gen_layers) runs at every denoising step. The UND pathway computes once and its K/V are cached by the pipeline itself; no cache-dit needed there. We wrap only gen_layers via BlockAdapter.

Parameters:

Name Type Description Default
pipeline Any

The Cosmos3 pipeline instance.

required
cache_config Any

DiffusionCacheConfig instance with cache configuration.

required

Returns:

Type Description
RefreshCacheContextFunc

A refresh function that can be called to update cache context with new num_inference_steps.

enable_cache_for_krea2

enable_cache_for_krea2(
    pipeline: Any, cache_config: Any
) -> RefreshCacheContextFunc

Enable cache-dit for Krea 2.

Krea 2 is a single-stream MMDiT: each Krea2TransformerBlock takes and returns only hidden_states (text is fused into the token stream), so the blocks follow ForwardPattern.Pattern_3.

has_separate_cfg is checkpoint-dependent, which is why this needs a custom enabler rather than a static _cache_dit_adapter_config: the distilled Turbo checkpoint runs no-CFG (a single transformer forward per denoise step), while the Raw checkpoint runs CFG as two separate forwards. cache-dit tells cond/uncond apart purely by transformer-forward parity, so the flag must match the actual per-step forward count — True only for the CFG (Raw) path. The pipeline exposes this via is_distilled (read from model_index.json).

Parameters:

Name Type Description Default
pipeline Any

The Krea2Pipeline instance.

required
cache_config Any

DiffusionCacheConfig instance with cache configuration.

required

Returns:

Type Description
RefreshCacheContextFunc

A refresh function that can be called to update cache context with new num_inference_steps.

enable_cache_for_magi2

enable_cache_for_magi2(
    pipeline: Any, cache_config: Any
) -> CacheDiTEnableResult

Cache only MAGI-2's repeated native transformer-layer stack.

The pre/post adapters still execute on every denoising step. MAGI-2 either packs both CFG branches into one transformer call or assigns one branch to each CFG-parallel rank. In both layouts each rank invokes this stack once per denoising step, so it remains a non-separate-CFG Pattern-3 stack.

enable_cache_for_mammothmoda2

enable_cache_for_mammothmoda2(
    pipeline: Any, cache_config: Any
) -> CacheDiTEnableResult

Cache only MammothModa2's repeated main DiT stack.

Transformer2DModel runs three Q-Former refiners (noise / ref-image / context) whose inputs change every denoise step, plus the layers stack that dominates per-step compute. Only layers is a repeated residual stack, so the BlockAdapter wraps it and the refiners stay outside the cached region. Blocks take hidden_states plus keyword step context and return only hidden states (Pattern_3).

The pipeline runs sequential CFG: each conditional forward is followed by an unconditional forward. cache-dit tells cond/uncond apart purely by transformer-forward parity (has_separate_cfg=True), so the cfg_range optimization that skips the unconditional pass outside the interval would desync that accounting. Like Cosmos3, we keep the passes paired and neutralize CFG via scale=1.0 outside the interval; the pipeline also disables hooks for no-CFG requests (single forward per step), whose parity the accounting cannot represent.

enable_cache_for_wan22

enable_cache_for_wan22(
    pipeline: Any, cache_config: Any
) -> RefreshCacheContextFunc

Enable cache-dit for Wan2.2 single or dual-transformer architecture.

Wan2.2 can use single or dual transformers (transformer and transformer_2) that need to be enabled using BlockAdapter.

Parameters:

Name Type Description Default
pipeline Any

The Wan2.2 pipeline instance.

required
cache_config Any

DiffusionCacheConfig instance with cache configuration.

required

Returns:

Type Description
RefreshCacheContextFunc

A refresh function that can be called to update cache context with new num_inference_steps.

enable_cache_for_wan22_s2v

enable_cache_for_wan22_s2v(
    pipeline: Any, cache_config: Any
) -> RefreshCacheContextFunc

Enable cache-dit for Wan2.2 S2V.

S2V uses a single transformer, but unlike the other Wan2.2 variants its block loop calls each block as block(hidden_states, **kwargs) and keeps the timestep modulation state in e rather than a second positional tensor. CacheDiT Pattern_3 matches that contract: cache hidden states only and pass the remaining conditioning through kwargs unchanged.

The S2V transformer has an after_transformer_block method that injects audio embeddings after specific layers. The cached blocks wrapper (Wan22S2VCachedBlocks._run_block) calls the original internally, so we permanently replace it with a no-op on the transformer to prevent double injection from the main forward loop.

register_custom_dit_enablers

register_custom_dit_enablers() -> None

Register model-specific Cache-DiT enablers.

This is called explicitly by the package initializer so registration does not depend on unrelated model-specific symbols being imported for their side effects.