[diffusion] refactor: remove useless lazy-import cache-dit codes (#17659)

This commit is contained in:
Mick
2026-01-25 22:43:22 +08:00
committed by GitHub
parent 7a1d7ab47b
commit b105dad5da
2 changed files with 10 additions and 34 deletions

View File

@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
"""
Cache acceleration module for SGLang multimodal generation.
Cache acceleration module for SGLang-diffusion
This module provides various caching strategies to accelerate
diffusion transformer (DiT) inference:
@@ -8,26 +8,13 @@ diffusion transformer (DiT) inference:
- TeaCache: Temporal similarity-based caching for diffusion models
- cache-dit integration: Block-level caching with DBCache and TaylorSeer
Usage:
# TeaCache (built into DiT models via TeaCacheMixin)
from sglang.multimodal_gen.runtime.cache import TeaCacheMixin, TeaCacheContext
# cache-dit integration (requires cache-dit package installed)
# These are lazy-loaded and will raise ImportError if cache-dit is not available
from sglang.multimodal_gen.runtime.cache import (
CacheDitConfig,
enable_cache_on_transformer,
enable_cache_on_dual_transformer,
get_scm_mask,
)
# Or import directly from the submodule:
from sglang.multimodal_gen.runtime.cache.cache_dit_integration import (
CacheDitConfig,
enable_cache_on_transformer,
)
"""
from sglang.multimodal_gen.runtime.cache.cache_dit_integration import (
CacheDitConfig,
enable_cache_on_dual_transformer,
enable_cache_on_transformer,
get_scm_mask,
)
from sglang.multimodal_gen.runtime.cache.teacache import TeaCacheContext, TeaCacheMixin
__all__ = [
@@ -40,17 +27,3 @@ __all__ = [
"enable_cache_on_dual_transformer",
"get_scm_mask",
]
# Lazy imports for cache-dit integration (requires cache-dit package)
def __getattr__(name):
if name in (
"CacheDitConfig",
"enable_cache_on_transformer",
"enable_cache_on_dual_transformer",
"get_scm_mask",
):
from sglang.multimodal_gen.runtime.cache import cache_dit_integration
return getattr(cache_dit_integration, name)
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")

View File

@@ -23,6 +23,8 @@ from typing import TYPE_CHECKING, Any
import numpy as np
import torch
from sglang.multimodal_gen.configs.models import DiTConfig
if TYPE_CHECKING:
from sglang.multimodal_gen.configs.sample.teacache import TeaCacheParams
@@ -125,6 +127,7 @@ class TeaCacheMixin:
# Models that support CFG cache separation (wan/hunyuan/zimage)
# Models not in this set (flux/qwen) auto-disable TeaCache when CFG is enabled
_CFG_SUPPORTED_PREFIXES: set[str] = {"wan", "hunyuan", "zimage"}
config: DiTConfig
def _init_teacache_state(self) -> None:
"""Initialize TeaCache state. Call this in subclass __init__."""