[diffusion] feat: enable passing Cache‑DiT config for diffusers backend (#16662)
Signed-off-by: Chi <chixie.mcisaac@gmail.com> Signed-off-by: qimcis <chixie.mcisaac@gmail.com> Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
@@ -32,6 +32,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.executors.sync_executor import
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages import PipelineStage
|
||||
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_model
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
@@ -410,7 +411,7 @@ class DiffusersPipeline(ComposedPipelineBase):
|
||||
load_kwargs["device_map"] = device_map
|
||||
|
||||
# Add quantization config if provided (e.g., BitsAndBytesConfig for 4/8-bit)
|
||||
config = getattr(server_args, "pipeline_config", None)
|
||||
config = server_args.pipeline_config
|
||||
if config is not None:
|
||||
quant_config = getattr(config, "quantization_config", None)
|
||||
if quant_config is not None:
|
||||
@@ -470,12 +471,15 @@ class DiffusersPipeline(ComposedPipelineBase):
|
||||
# Apply attention backend if specified
|
||||
self._apply_attention_backend(pipe, server_args)
|
||||
|
||||
# Apply cache-dit acceleration if configured
|
||||
pipe = self._apply_cache_dit(pipe, server_args)
|
||||
|
||||
logger.info("Loaded diffusers pipeline: %s", pipe.__class__.__name__)
|
||||
return pipe
|
||||
|
||||
def _apply_vae_optimizations(self, pipe: Any, server_args: ServerArgs) -> None:
|
||||
"""Apply VAE memory optimizations (tiling, slicing) from pipeline config."""
|
||||
config = getattr(server_args, "pipeline_config", None)
|
||||
config = server_args.pipeline_config
|
||||
if config is None:
|
||||
return
|
||||
|
||||
@@ -499,16 +503,30 @@ class DiffusersPipeline(ComposedPipelineBase):
|
||||
See: https://huggingface.co/docs/diffusers/main/en/optimization/attention_backends
|
||||
Available backends: flash, _flash_3_hub, sage, xformers, native, etc.
|
||||
"""
|
||||
backend = getattr(server_args, "diffusers_attention_backend", None)
|
||||
backend = server_args.attention_backend
|
||||
|
||||
if backend is None:
|
||||
config = getattr(server_args, "pipeline_config", None)
|
||||
config = server_args.pipeline_config
|
||||
if config is not None:
|
||||
backend = getattr(config, "diffusers_attention_backend", None)
|
||||
|
||||
if backend is None:
|
||||
return
|
||||
|
||||
backend = backend.lower()
|
||||
sglang_backends = {e.name.lower() for e in AttentionBackendEnum} | {
|
||||
"fa3",
|
||||
"fa4",
|
||||
}
|
||||
if backend in sglang_backends:
|
||||
logger.debug(
|
||||
"Skipping diffusers attention backend '%s' because it matches a "
|
||||
"SGLang backend name. Use diffusers backend names when running "
|
||||
"the diffusers backend.",
|
||||
backend,
|
||||
)
|
||||
return
|
||||
|
||||
for component_name in ["transformer", "unet"]:
|
||||
component = getattr(pipe, component_name, None)
|
||||
if component is not None and hasattr(component, "set_attention_backend"):
|
||||
@@ -525,6 +543,44 @@ class DiffusersPipeline(ComposedPipelineBase):
|
||||
e,
|
||||
)
|
||||
|
||||
def _apply_cache_dit(self, pipe: Any, server_args: ServerArgs) -> Any:
|
||||
"""Enable cache-dit for diffusers pipeline if configured."""
|
||||
cache_dit_config = server_args.cache_dit_config
|
||||
if not cache_dit_config:
|
||||
return pipe
|
||||
|
||||
try:
|
||||
import cache_dit
|
||||
except ImportError as e:
|
||||
raise RuntimeError(
|
||||
"cache-dit is required for --cache-dit-config. "
|
||||
"Install it with `pip install cache-dit`."
|
||||
) from e
|
||||
|
||||
if not hasattr(cache_dit, "load_configs"):
|
||||
raise RuntimeError(
|
||||
"cache-dit>=1.2.0 is required for --cache-dit-config. "
|
||||
"Please upgrade cache-dit."
|
||||
)
|
||||
|
||||
try:
|
||||
cache_options = cache_dit.load_configs(cache_dit_config)
|
||||
except Exception as e:
|
||||
raise ValueError(
|
||||
"Failed to load cache-dit config. Provide a YAML/JSON path (or a dict "
|
||||
"supported by cache-dit>=1.2.0)."
|
||||
) from e
|
||||
|
||||
try:
|
||||
pipe = cache_dit.enable_cache(pipe, **cache_options)
|
||||
except Exception:
|
||||
# cache-dit is an external integration and can raise a variety of errors.
|
||||
logger.exception("Failed to enable cache-dit for diffusers pipeline")
|
||||
raise
|
||||
|
||||
logger.info("Enabled cache-dit for diffusers pipeline")
|
||||
return pipe
|
||||
|
||||
def _get_device_map(self, server_args: ServerArgs) -> str | None:
|
||||
"""
|
||||
Determine device_map for pipeline loading.
|
||||
@@ -540,7 +596,7 @@ class DiffusersPipeline(ComposedPipelineBase):
|
||||
dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16
|
||||
|
||||
if hasattr(server_args, "pipeline_config") and server_args.pipeline_config:
|
||||
dit_precision = getattr(server_args.pipeline_config, "dit_precision", None)
|
||||
dit_precision = server_args.pipeline_config.dit_precision
|
||||
if dit_precision == "fp16":
|
||||
dtype = torch.float16
|
||||
elif dit_precision == "bf16":
|
||||
|
||||
@@ -235,7 +235,9 @@ class ServerArgs:
|
||||
|
||||
# Attention
|
||||
attention_backend: str = None
|
||||
diffusers_attention_backend: str = None # for diffusers backend only
|
||||
cache_dit_config: str | dict[str, Any] | None = (
|
||||
None # cache-dit config for diffusers
|
||||
)
|
||||
|
||||
# Distributed executor backend
|
||||
nccl_port: Optional[int] = None
|
||||
@@ -452,15 +454,25 @@ class ServerArgs:
|
||||
"--attention-backend",
|
||||
type=str,
|
||||
default=None,
|
||||
choices=[e.name.lower() for e in AttentionBackendEnum] + ["fa3", "fa4"],
|
||||
help="The attention backend to use. If not specified, the backend is automatically selected based on hardware and installed packages.",
|
||||
help=(
|
||||
"The attention backend to use. For SGLang-native pipelines, use "
|
||||
"values like fa, torch_sdpa, sage_attn, etc. For diffusers pipelines, "
|
||||
"use diffusers attention backend names such as flash, _flash_3_hub, "
|
||||
"sage, or xformers."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--diffusers-attention-backend",
|
||||
type=str,
|
||||
dest="attention_backend",
|
||||
default=None,
|
||||
help="Attention backend for diffusers pipelines (e.g., flash, _flash_3_hub, sage, xformers). "
|
||||
"See: https://huggingface.co/docs/diffusers/main/en/optimization/attention_backends",
|
||||
help=argparse.SUPPRESS,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--cache-dit-config",
|
||||
type=str,
|
||||
default=ServerArgs.cache_dit_config,
|
||||
help="Path to a Cache-DiT YAML/JSON config. Enables cache-dit for diffusers backend.",
|
||||
)
|
||||
|
||||
# HuggingFace specific parameters
|
||||
@@ -999,7 +1011,7 @@ class ServerArgs:
|
||||
raise ValueError("pipeline_config is not set in ServerArgs")
|
||||
|
||||
self.pipeline_config.check_pipeline_config()
|
||||
if self.attention_backend is None:
|
||||
if self.attention_backend is None and self.backend != Backend.DIFFUSERS:
|
||||
self._set_default_attention_backend()
|
||||
|
||||
# parallelism
|
||||
|
||||
Reference in New Issue
Block a user