[diffusion] feat: support cache-dit for Ideogram 4 (#29631)
This commit is contained in:
@@ -415,7 +415,8 @@ def enable_cache_on_dual_transformer(
|
|||||||
tp_group: Tensor parallel process group.
|
tp_group: Tensor parallel process group.
|
||||||
"""
|
"""
|
||||||
_supported_dual_transformer_models = [
|
_supported_dual_transformer_models = [
|
||||||
"wan2.2", # Currently, only Wan2.2 will run into dual-transformer case
|
"wan2.2",
|
||||||
|
"ideogram4",
|
||||||
]
|
]
|
||||||
if model_name not in _supported_dual_transformer_models:
|
if model_name not in _supported_dual_transformer_models:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -494,7 +495,7 @@ def enable_cache_on_dual_transformer(
|
|||||||
primary_config.enable_taylorseer,
|
primary_config.enable_taylorseer,
|
||||||
)
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
" Secondary (transformer_2): Fn=%d, Bn=%d, W=%d, R=%.2f, MC=%d, TaylorSeer=%s",
|
" Secondary transformer: Fn=%d, Bn=%d, W=%d, R=%.2f, MC=%d, TaylorSeer=%s",
|
||||||
secondary_config.Fn_compute_blocks,
|
secondary_config.Fn_compute_blocks,
|
||||||
secondary_config.Bn_compute_blocks,
|
secondary_config.Bn_compute_blocks,
|
||||||
secondary_config.max_warmup_steps,
|
secondary_config.max_warmup_steps,
|
||||||
@@ -532,37 +533,49 @@ def enable_cache_on_dual_transformer(
|
|||||||
transformer_2, parallelism_config, sp_group, tp_group
|
transformer_2, parallelism_config, sp_group, tp_group
|
||||||
)
|
)
|
||||||
|
|
||||||
# Get blocks attribute - Wan transformers use 'blocks' attribute
|
if model_name == "wan2.2":
|
||||||
transformer_blocks = getattr(transformer, "blocks", None)
|
transformer_blocks = getattr(transformer, "blocks", None)
|
||||||
transformer_2_blocks = getattr(transformer_2, "blocks", None)
|
transformer_2_blocks = getattr(transformer_2, "blocks", None)
|
||||||
|
blocks_name = None
|
||||||
if transformer_blocks is None or transformer_2_blocks is None:
|
forward_pattern = [ForwardPattern.Pattern_2, ForwardPattern.Pattern_2]
|
||||||
|
check_forward_pattern = True
|
||||||
|
check_num_outputs = False
|
||||||
|
has_separate_cfg = True
|
||||||
|
elif model_name == "ideogram4":
|
||||||
|
transformer_blocks = getattr(transformer, "layers", None)
|
||||||
|
transformer_2_blocks = getattr(transformer_2, "layers", None)
|
||||||
|
blocks_name = ["layers", "layers"]
|
||||||
|
forward_pattern = [ForwardPattern.Pattern_3, ForwardPattern.Pattern_3]
|
||||||
|
check_forward_pattern = False
|
||||||
|
check_num_outputs = False
|
||||||
|
has_separate_cfg = False
|
||||||
|
else:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"Dual transformers must have 'blocks' attribute for cache-dit. "
|
f"Dual-transformer is not implemented for model {model_name} yet."
|
||||||
f"transformer has blocks: {transformer_blocks is not None}, "
|
)
|
||||||
f"transformer_2 has blocks: {transformer_2_blocks is not None}"
|
|
||||||
|
if transformer_blocks is None or transformer_2_blocks is None:
|
||||||
|
expected_attr = "layers" if model_name == "ideogram4" else "blocks"
|
||||||
|
raise ValueError(
|
||||||
|
f"Dual transformers for {model_name} must have '{expected_attr}' "
|
||||||
|
"attribute for cache-dit. "
|
||||||
|
f"transformer has {expected_attr}: {transformer_blocks is not None}, "
|
||||||
|
f"secondary transformer has {expected_attr}: {transformer_2_blocks is not None}"
|
||||||
)
|
)
|
||||||
|
|
||||||
# Enable cache-dit using BlockAdapter for both transformers simultaneously
|
|
||||||
# This is required for Wan2.2 and similar dual-transformer architectures
|
|
||||||
if model_name == "wan2.2":
|
|
||||||
# Use Pattern_2 for Wan2.2 dual-transformer. We should check `model_name`
|
|
||||||
# to ensure we only apply this for supported models. Different models
|
|
||||||
# may require different ForwardPattern.
|
|
||||||
cache_dit.enable_cache(
|
cache_dit.enable_cache(
|
||||||
BlockAdapter(
|
BlockAdapter(
|
||||||
transformer=[transformer, transformer_2],
|
transformer=[transformer, transformer_2],
|
||||||
blocks=[transformer_blocks, transformer_2_blocks],
|
blocks=[transformer_blocks, transformer_2_blocks],
|
||||||
forward_pattern=[ForwardPattern.Pattern_2, ForwardPattern.Pattern_2],
|
blocks_name=blocks_name,
|
||||||
|
forward_pattern=forward_pattern,
|
||||||
params_modifiers=[primary_modifier, secondary_modifier],
|
params_modifiers=[primary_modifier, secondary_modifier],
|
||||||
has_separate_cfg=True,
|
check_forward_pattern=check_forward_pattern,
|
||||||
|
check_num_outputs=check_num_outputs,
|
||||||
|
has_separate_cfg=has_separate_cfg,
|
||||||
),
|
),
|
||||||
parallelism_config=None,
|
parallelism_config=None,
|
||||||
)
|
)
|
||||||
else:
|
|
||||||
raise ValueError(
|
|
||||||
f"Dual-transformer is not implemented for model {model_name} yet."
|
|
||||||
)
|
|
||||||
|
|
||||||
if parallelism_config is not None:
|
if parallelism_config is not None:
|
||||||
for t in [transformer, transformer_2]:
|
for t in [transformer, transformer_2]:
|
||||||
@@ -612,23 +625,30 @@ def refresh_context_on_dual_transformer(
|
|||||||
num_low_noise_steps: int,
|
num_low_noise_steps: int,
|
||||||
scm_preset: str | None = None,
|
scm_preset: str | None = None,
|
||||||
verbose: bool = False,
|
verbose: bool = False,
|
||||||
|
steps_computation_mask: Optional[List[int]] = None,
|
||||||
|
steps_computation_mask_2: Optional[List[int]] = None,
|
||||||
|
steps_computation_policy: str | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Refresh cache-dit context for dual transformers."""
|
"""Refresh cache-dit context for dual transformers."""
|
||||||
high_noise_steps_computation_mask = None
|
high_noise_steps_computation_mask = steps_computation_mask
|
||||||
low_noise_steps_computation_mask = None
|
low_noise_steps_computation_mask = steps_computation_mask_2
|
||||||
if scm_preset is not None:
|
if high_noise_steps_computation_mask is None and scm_preset is not None:
|
||||||
high_noise_steps_computation_mask = cache_dit.steps_mask(
|
high_noise_steps_computation_mask = cache_dit.steps_mask(
|
||||||
mask_policy=scm_preset, total_steps=num_high_noise_steps
|
mask_policy=scm_preset, total_steps=num_high_noise_steps
|
||||||
)
|
)
|
||||||
|
if low_noise_steps_computation_mask is None and scm_preset is not None:
|
||||||
low_noise_steps_computation_mask = cache_dit.steps_mask(
|
low_noise_steps_computation_mask = cache_dit.steps_mask(
|
||||||
mask_policy=scm_preset, total_steps=num_low_noise_steps
|
mask_policy=scm_preset, total_steps=num_low_noise_steps
|
||||||
)
|
)
|
||||||
|
policy = (
|
||||||
|
steps_computation_policy if steps_computation_policy is not None else scm_preset
|
||||||
|
)
|
||||||
cache_dit.refresh_context(
|
cache_dit.refresh_context(
|
||||||
transformer,
|
transformer,
|
||||||
cache_config=DBCacheConfig().reset(
|
cache_config=DBCacheConfig().reset(
|
||||||
num_inference_steps=num_high_noise_steps,
|
num_inference_steps=num_high_noise_steps,
|
||||||
steps_computation_mask=high_noise_steps_computation_mask,
|
steps_computation_mask=high_noise_steps_computation_mask,
|
||||||
steps_computation_policy=scm_preset,
|
steps_computation_policy=policy,
|
||||||
),
|
),
|
||||||
verbose=verbose,
|
verbose=verbose,
|
||||||
)
|
)
|
||||||
@@ -637,7 +657,7 @@ def refresh_context_on_dual_transformer(
|
|||||||
cache_config=DBCacheConfig().reset(
|
cache_config=DBCacheConfig().reset(
|
||||||
num_inference_steps=num_low_noise_steps,
|
num_inference_steps=num_low_noise_steps,
|
||||||
steps_computation_mask=low_noise_steps_computation_mask,
|
steps_computation_mask=low_noise_steps_computation_mask,
|
||||||
steps_computation_policy=scm_preset,
|
steps_computation_policy=policy,
|
||||||
),
|
),
|
||||||
verbose=verbose,
|
verbose=verbose,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ import weakref
|
|||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from dataclasses import dataclass, field, fields
|
from dataclasses import dataclass, field, fields
|
||||||
|
from enum import Enum
|
||||||
from functools import lru_cache
|
from functools import lru_cache
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
@@ -176,6 +177,17 @@ class DenoisingStepState:
|
|||||||
attn_metadata: Any | None
|
attn_metadata: Any | None
|
||||||
|
|
||||||
|
|
||||||
|
class DualTransformerExecutionMode(str, Enum):
|
||||||
|
"""How a denoising stage uses a second DiT.
|
||||||
|
|
||||||
|
BOUNDARY_EXPERTS means one transformer is selected per timestep.
|
||||||
|
PAIRED_PER_STEP means both transformers participate in each denoising step.
|
||||||
|
"""
|
||||||
|
|
||||||
|
BOUNDARY_EXPERTS = "boundary_experts"
|
||||||
|
PAIRED_PER_STEP = "paired_per_step"
|
||||||
|
|
||||||
|
|
||||||
class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
||||||
"""
|
"""
|
||||||
Stage for running the denoising loop in diffusion pipelines.
|
Stage for running the denoising loop in diffusion pipelines.
|
||||||
@@ -417,6 +429,133 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
|||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
def _cache_dit_dual_model_name(self) -> str:
|
||||||
|
return "wan2.2"
|
||||||
|
|
||||||
|
def _cache_dit_secondary_uses_primary_config(self) -> bool:
|
||||||
|
return False
|
||||||
|
|
||||||
|
def _dual_transformer_execution_mode(self) -> DualTransformerExecutionMode | None:
|
||||||
|
if self.transformer_2 is None:
|
||||||
|
return None
|
||||||
|
return DualTransformerExecutionMode.BOUNDARY_EXPERTS
|
||||||
|
|
||||||
|
def _cache_dit_step_counts(
|
||||||
|
self, num_inference_steps: int | tuple[int, int]
|
||||||
|
) -> tuple[int, int | None]:
|
||||||
|
if isinstance(num_inference_steps, tuple):
|
||||||
|
primary_steps, secondary_steps = num_inference_steps
|
||||||
|
return int(primary_steps), int(secondary_steps)
|
||||||
|
steps = int(num_inference_steps)
|
||||||
|
mode = self._dual_transformer_execution_mode()
|
||||||
|
if mode is None:
|
||||||
|
return steps, None
|
||||||
|
if mode == DualTransformerExecutionMode.PAIRED_PER_STEP:
|
||||||
|
return steps, steps
|
||||||
|
raise ValueError("Boundary-expert dual transformers require split step counts.")
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _parse_cache_dit_scm_bins() -> tuple[list[int] | None, list[int] | None, str]:
|
||||||
|
scm_preset = envs.SGLANG_CACHE_DIT_SCM_PRESET
|
||||||
|
compute_bins_str = envs.SGLANG_CACHE_DIT_SCM_COMPUTE_BINS
|
||||||
|
cache_bins_str = envs.SGLANG_CACHE_DIT_SCM_CACHE_BINS
|
||||||
|
compute_bins = None
|
||||||
|
cache_bins = None
|
||||||
|
if compute_bins_str and cache_bins_str:
|
||||||
|
try:
|
||||||
|
compute_bins = [int(x.strip()) for x in compute_bins_str.split(",")]
|
||||||
|
cache_bins = [int(x.strip()) for x in cache_bins_str.split(",")]
|
||||||
|
except ValueError as exc:
|
||||||
|
logger.warning("Failed to parse SCM bins: %s. SCM disabled.", exc)
|
||||||
|
scm_preset = "none"
|
||||||
|
elif compute_bins_str or cache_bins_str:
|
||||||
|
logger.warning(
|
||||||
|
"SCM custom bins require both compute_bins and cache_bins. "
|
||||||
|
"Only one was provided (compute=%s, cache=%s). Falling back to preset '%s'.",
|
||||||
|
compute_bins_str,
|
||||||
|
cache_bins_str,
|
||||||
|
scm_preset,
|
||||||
|
)
|
||||||
|
return compute_bins, cache_bins, scm_preset
|
||||||
|
|
||||||
|
def _cache_dit_scm_masks(
|
||||||
|
self, primary_num_steps: int, secondary_num_steps: int | None = None
|
||||||
|
) -> tuple[str, str, list[int] | None, list[int] | None]:
|
||||||
|
scm_compute_bins, scm_cache_bins, scm_preset = self._parse_cache_dit_scm_bins()
|
||||||
|
scm_policy = envs.SGLANG_CACHE_DIT_SCM_POLICY
|
||||||
|
steps_computation_mask = get_scm_mask(
|
||||||
|
preset=scm_preset,
|
||||||
|
num_inference_steps=primary_num_steps,
|
||||||
|
compute_bins=scm_compute_bins,
|
||||||
|
cache_bins=scm_cache_bins,
|
||||||
|
)
|
||||||
|
|
||||||
|
steps_computation_mask_2 = None
|
||||||
|
if secondary_num_steps is not None:
|
||||||
|
if (
|
||||||
|
self._cache_dit_secondary_uses_primary_config()
|
||||||
|
and secondary_num_steps == primary_num_steps
|
||||||
|
):
|
||||||
|
steps_computation_mask_2 = steps_computation_mask
|
||||||
|
else:
|
||||||
|
steps_computation_mask_2 = get_scm_mask(
|
||||||
|
preset=scm_preset,
|
||||||
|
num_inference_steps=secondary_num_steps,
|
||||||
|
compute_bins=scm_compute_bins,
|
||||||
|
cache_bins=scm_cache_bins,
|
||||||
|
)
|
||||||
|
return scm_preset, scm_policy, steps_computation_mask, steps_computation_mask_2
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _build_cache_dit_config(
|
||||||
|
num_inference_steps: int,
|
||||||
|
steps_computation_mask: list[int] | None,
|
||||||
|
scm_policy: str,
|
||||||
|
*,
|
||||||
|
secondary: bool = False,
|
||||||
|
) -> CacheDitConfig:
|
||||||
|
return CacheDitConfig(
|
||||||
|
enabled=True,
|
||||||
|
Fn_compute_blocks=(
|
||||||
|
envs.SGLANG_CACHE_DIT_SECONDARY_FN
|
||||||
|
if secondary
|
||||||
|
else envs.SGLANG_CACHE_DIT_FN
|
||||||
|
),
|
||||||
|
Bn_compute_blocks=(
|
||||||
|
envs.SGLANG_CACHE_DIT_SECONDARY_BN
|
||||||
|
if secondary
|
||||||
|
else envs.SGLANG_CACHE_DIT_BN
|
||||||
|
),
|
||||||
|
max_warmup_steps=(
|
||||||
|
envs.SGLANG_CACHE_DIT_SECONDARY_WARMUP
|
||||||
|
if secondary
|
||||||
|
else envs.SGLANG_CACHE_DIT_WARMUP
|
||||||
|
),
|
||||||
|
residual_diff_threshold=(
|
||||||
|
envs.SGLANG_CACHE_DIT_SECONDARY_RDT
|
||||||
|
if secondary
|
||||||
|
else envs.SGLANG_CACHE_DIT_RDT
|
||||||
|
),
|
||||||
|
max_continuous_cached_steps=(
|
||||||
|
envs.SGLANG_CACHE_DIT_SECONDARY_MC
|
||||||
|
if secondary
|
||||||
|
else envs.SGLANG_CACHE_DIT_MC
|
||||||
|
),
|
||||||
|
enable_taylorseer=(
|
||||||
|
envs.SGLANG_CACHE_DIT_SECONDARY_TAYLORSEER
|
||||||
|
if secondary
|
||||||
|
else envs.SGLANG_CACHE_DIT_TAYLORSEER
|
||||||
|
),
|
||||||
|
taylorseer_order=(
|
||||||
|
envs.SGLANG_CACHE_DIT_SECONDARY_TS_ORDER
|
||||||
|
if secondary
|
||||||
|
else envs.SGLANG_CACHE_DIT_TS_ORDER
|
||||||
|
),
|
||||||
|
num_inference_steps=num_inference_steps,
|
||||||
|
steps_computation_mask=steps_computation_mask,
|
||||||
|
steps_computation_policy=scm_policy,
|
||||||
|
)
|
||||||
|
|
||||||
def _maybe_enable_cache_dit(
|
def _maybe_enable_cache_dit(
|
||||||
self, num_inference_steps: int | tuple[int, int], batch: Req
|
self, num_inference_steps: int | tuple[int, int], batch: Req
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -429,25 +568,30 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
|||||||
transformers with (potentially) different configurations.
|
transformers with (potentially) different configurations.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
if isinstance(num_inference_steps, tuple):
|
|
||||||
num_high_noise_steps, num_low_noise_steps = num_inference_steps
|
|
||||||
|
|
||||||
# NOTE: When a new request arrives, we need to refresh the cache-dit context.
|
# NOTE: When a new request arrives, we need to refresh the cache-dit context.
|
||||||
if self._cache_dit_enabled:
|
if self._cache_dit_enabled:
|
||||||
scm_preset = envs.SGLANG_CACHE_DIT_SCM_PRESET
|
primary_num_steps, secondary_num_steps = self._cache_dit_step_counts(
|
||||||
scm_preset = None if scm_preset == "none" else scm_preset
|
num_inference_steps
|
||||||
if isinstance(num_inference_steps, tuple):
|
)
|
||||||
|
scm_preset, scm_policy, steps_computation_mask, steps_computation_mask_2 = (
|
||||||
|
self._cache_dit_scm_masks(primary_num_steps, secondary_num_steps)
|
||||||
|
)
|
||||||
|
if self.transformer_2 is not None:
|
||||||
|
assert secondary_num_steps is not None
|
||||||
refresh_context_on_dual_transformer(
|
refresh_context_on_dual_transformer(
|
||||||
self.transformer,
|
self.transformer,
|
||||||
self.transformer_2,
|
self.transformer_2,
|
||||||
num_high_noise_steps,
|
primary_num_steps,
|
||||||
num_low_noise_steps,
|
secondary_num_steps,
|
||||||
scm_preset=scm_preset,
|
steps_computation_mask=steps_computation_mask,
|
||||||
|
steps_computation_mask_2=steps_computation_mask_2,
|
||||||
|
steps_computation_policy=scm_policy,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
scm_preset = None if scm_preset == "none" else scm_preset
|
||||||
refresh_context_on_transformer(
|
refresh_context_on_transformer(
|
||||||
self.transformer,
|
self.transformer,
|
||||||
num_inference_steps,
|
primary_num_steps,
|
||||||
scm_preset=scm_preset,
|
scm_preset=scm_preset,
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
@@ -459,6 +603,9 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
|||||||
if batch.is_warmup and not self.server_args.enable_torch_compile:
|
if batch.is_warmup and not self.server_args.enable_torch_compile:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
primary_num_steps, secondary_num_steps = self._cache_dit_step_counts(
|
||||||
|
num_inference_steps
|
||||||
|
)
|
||||||
world_size = get_world_size()
|
world_size = get_world_size()
|
||||||
parallelized = world_size > 1
|
parallelized = world_size > 1
|
||||||
|
|
||||||
@@ -483,93 +630,29 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
|||||||
has_sp,
|
has_sp,
|
||||||
has_tp,
|
has_tp,
|
||||||
)
|
)
|
||||||
# === Parse SCM configuration from envs ===
|
_, scm_policy, steps_computation_mask, steps_computation_mask_2 = (
|
||||||
# SCM is shared between primary and secondary transformers
|
self._cache_dit_scm_masks(primary_num_steps, secondary_num_steps)
|
||||||
scm_preset = envs.SGLANG_CACHE_DIT_SCM_PRESET
|
|
||||||
scm_compute_bins_str = envs.SGLANG_CACHE_DIT_SCM_COMPUTE_BINS
|
|
||||||
scm_cache_bins_str = envs.SGLANG_CACHE_DIT_SCM_CACHE_BINS
|
|
||||||
scm_policy = envs.SGLANG_CACHE_DIT_SCM_POLICY
|
|
||||||
|
|
||||||
# parse custom bins if provided (both must be set together)
|
|
||||||
scm_compute_bins = None
|
|
||||||
scm_cache_bins = None
|
|
||||||
if scm_compute_bins_str and scm_cache_bins_str:
|
|
||||||
try:
|
|
||||||
scm_compute_bins = [
|
|
||||||
int(x.strip()) for x in scm_compute_bins_str.split(",")
|
|
||||||
]
|
|
||||||
scm_cache_bins = [int(x.strip()) for x in scm_cache_bins_str.split(",")]
|
|
||||||
except ValueError as e:
|
|
||||||
logger.warning("Failed to parse SCM bins: %s. SCM disabled.", e)
|
|
||||||
scm_preset = "none"
|
|
||||||
elif scm_compute_bins_str or scm_cache_bins_str:
|
|
||||||
# Only one of the bins was provided - warn user
|
|
||||||
logger.warning(
|
|
||||||
"SCM custom bins require both compute_bins and cache_bins. "
|
|
||||||
"Only one was provided (compute=%s, cache=%s). Falling back to preset '%s'.",
|
|
||||||
scm_compute_bins_str,
|
|
||||||
scm_cache_bins_str,
|
|
||||||
scm_preset,
|
|
||||||
)
|
)
|
||||||
|
primary_config = self._build_cache_dit_config(
|
||||||
# generate SCM mask using cache-dit's steps_mask()
|
primary_num_steps,
|
||||||
# cache-dit handles step count validation and scaling internally
|
|
||||||
steps_computation_mask = get_scm_mask(
|
|
||||||
preset=scm_preset,
|
|
||||||
num_inference_steps=(
|
|
||||||
num_inference_steps
|
|
||||||
if isinstance(num_inference_steps, int)
|
|
||||||
else num_high_noise_steps
|
|
||||||
),
|
|
||||||
compute_bins=scm_compute_bins,
|
|
||||||
cache_bins=scm_cache_bins,
|
|
||||||
)
|
|
||||||
|
|
||||||
if isinstance(num_inference_steps, tuple):
|
|
||||||
steps_computation_mask_2 = get_scm_mask(
|
|
||||||
preset=scm_preset,
|
|
||||||
num_inference_steps=num_low_noise_steps,
|
|
||||||
compute_bins=scm_compute_bins,
|
|
||||||
cache_bins=scm_cache_bins,
|
|
||||||
)
|
|
||||||
|
|
||||||
# build config for primary transformer (high-noise expert)
|
|
||||||
primary_config = CacheDitConfig(
|
|
||||||
enabled=True,
|
|
||||||
Fn_compute_blocks=envs.SGLANG_CACHE_DIT_FN,
|
|
||||||
Bn_compute_blocks=envs.SGLANG_CACHE_DIT_BN,
|
|
||||||
max_warmup_steps=envs.SGLANG_CACHE_DIT_WARMUP,
|
|
||||||
residual_diff_threshold=envs.SGLANG_CACHE_DIT_RDT,
|
|
||||||
max_continuous_cached_steps=envs.SGLANG_CACHE_DIT_MC,
|
|
||||||
enable_taylorseer=envs.SGLANG_CACHE_DIT_TAYLORSEER,
|
|
||||||
taylorseer_order=envs.SGLANG_CACHE_DIT_TS_ORDER,
|
|
||||||
num_inference_steps=(
|
|
||||||
num_inference_steps
|
|
||||||
if isinstance(num_inference_steps, int)
|
|
||||||
else num_high_noise_steps
|
|
||||||
),
|
|
||||||
# SCM fields
|
|
||||||
steps_computation_mask=steps_computation_mask,
|
steps_computation_mask=steps_computation_mask,
|
||||||
steps_computation_policy=scm_policy,
|
scm_policy=scm_policy,
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.transformer_2 is not None:
|
if self.transformer_2 is not None:
|
||||||
# dual transformer
|
# dual transformer
|
||||||
# build config for secondary transformer (low-noise expert)
|
# build config for secondary transformer (low-noise expert)
|
||||||
# uses secondary parameters which inherit from primary if not explicitly set
|
# uses secondary parameters which inherit from primary if not explicitly set
|
||||||
secondary_config = CacheDitConfig(
|
assert secondary_num_steps is not None
|
||||||
enabled=True,
|
secondary_config = (
|
||||||
Fn_compute_blocks=envs.SGLANG_CACHE_DIT_SECONDARY_FN,
|
primary_config
|
||||||
Bn_compute_blocks=envs.SGLANG_CACHE_DIT_SECONDARY_BN,
|
if self._cache_dit_secondary_uses_primary_config()
|
||||||
max_warmup_steps=envs.SGLANG_CACHE_DIT_SECONDARY_WARMUP,
|
else self._build_cache_dit_config(
|
||||||
residual_diff_threshold=envs.SGLANG_CACHE_DIT_SECONDARY_RDT,
|
secondary_num_steps,
|
||||||
max_continuous_cached_steps=envs.SGLANG_CACHE_DIT_SECONDARY_MC,
|
|
||||||
enable_taylorseer=envs.SGLANG_CACHE_DIT_SECONDARY_TAYLORSEER,
|
|
||||||
taylorseer_order=envs.SGLANG_CACHE_DIT_SECONDARY_TS_ORDER,
|
|
||||||
num_inference_steps=num_low_noise_steps,
|
|
||||||
# SCM fields - shared with primary
|
|
||||||
steps_computation_mask=steps_computation_mask_2,
|
steps_computation_mask=steps_computation_mask_2,
|
||||||
steps_computation_policy=scm_policy,
|
scm_policy=scm_policy,
|
||||||
|
secondary=True,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
# for dual transformers, must use BlockAdapter to enable cache on both simultaneously.
|
# for dual transformers, must use BlockAdapter to enable cache on both simultaneously.
|
||||||
@@ -579,14 +662,14 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
|||||||
self.transformer_2,
|
self.transformer_2,
|
||||||
primary_config,
|
primary_config,
|
||||||
secondary_config,
|
secondary_config,
|
||||||
model_name="wan2.2",
|
model_name=self._cache_dit_dual_model_name(),
|
||||||
sp_group=sp_group,
|
sp_group=sp_group,
|
||||||
tp_group=tp_group,
|
tp_group=tp_group,
|
||||||
)
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
"cache-dit enabled on dual transformers (steps=%d, %d)",
|
"cache-dit enabled on dual transformers (steps=%d, %d)",
|
||||||
num_high_noise_steps,
|
primary_num_steps,
|
||||||
num_low_noise_steps,
|
secondary_num_steps,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
# single transformer
|
# single transformer
|
||||||
@@ -600,7 +683,7 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
|||||||
)
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
"cache-dit enabled on transformer (steps=%d, Fn=%d, Bn=%d, rdt=%.3f)",
|
"cache-dit enabled on transformer (steps=%d, Fn=%d, Bn=%d, rdt=%.3f)",
|
||||||
num_inference_steps,
|
primary_num_steps,
|
||||||
envs.SGLANG_CACHE_DIT_FN,
|
envs.SGLANG_CACHE_DIT_FN,
|
||||||
envs.SGLANG_CACHE_DIT_BN,
|
envs.SGLANG_CACHE_DIT_BN,
|
||||||
envs.SGLANG_CACHE_DIT_RDT,
|
envs.SGLANG_CACHE_DIT_RDT,
|
||||||
@@ -683,9 +766,13 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
|||||||
scheduler = batch.scheduler
|
scheduler = batch.scheduler
|
||||||
assert scheduler is not None
|
assert scheduler is not None
|
||||||
|
|
||||||
|
dual_transformer_mode = self._dual_transformer_execution_mode()
|
||||||
|
uses_boundary_transformer_2 = (
|
||||||
|
dual_transformer_mode == DualTransformerExecutionMode.BOUNDARY_EXPERTS
|
||||||
|
)
|
||||||
boundary_timestep = (
|
boundary_timestep = (
|
||||||
self._handle_boundary_ratio(server_args, batch, scheduler)
|
self._handle_boundary_ratio(server_args, batch, scheduler)
|
||||||
if self.transformer_2 is not None
|
if uses_boundary_transformer_2
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
# Get timesteps and calculate warmup steps
|
# Get timesteps and calculate warmup steps
|
||||||
@@ -693,7 +780,7 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
|||||||
num_inference_steps = batch.num_inference_steps
|
num_inference_steps = batch.num_inference_steps
|
||||||
num_warmup_steps = len(timesteps) - num_inference_steps * scheduler.order
|
num_warmup_steps = len(timesteps) - num_inference_steps * scheduler.order
|
||||||
|
|
||||||
if self.transformer_2 is not None:
|
if uses_boundary_transformer_2:
|
||||||
assert boundary_timestep is not None, "boundary_timestep must be provided"
|
assert boundary_timestep is not None, "boundary_timestep must be provided"
|
||||||
num_high_noise_steps = (timesteps >= boundary_timestep).sum().item()
|
num_high_noise_steps = (timesteps >= boundary_timestep).sum().item()
|
||||||
num_low_noise_steps = num_inference_steps - num_high_noise_steps
|
num_low_noise_steps = num_inference_steps - num_high_noise_steps
|
||||||
|
|||||||
+19
-29
@@ -25,6 +25,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import (
|
|||||||
DenoisingContext,
|
DenoisingContext,
|
||||||
DenoisingStage,
|
DenoisingStage,
|
||||||
DenoisingStepState,
|
DenoisingStepState,
|
||||||
|
DualTransformerExecutionMode,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.text_encoding import (
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.text_encoding import (
|
||||||
TextEncodingStage,
|
TextEncodingStage,
|
||||||
@@ -36,6 +37,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
|||||||
VerificationResult,
|
VerificationResult,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
from sglang.multimodal_gen.runtime.utils.nvtx_pytorch_hooks import maybe_nvtx_range
|
from sglang.multimodal_gen.runtime.utils.nvtx_pytorch_hooks import maybe_nvtx_range
|
||||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||||
|
|
||||||
@@ -44,6 +46,8 @@ OUTPUT_IMAGE_INDICATOR = 2
|
|||||||
LLM_TOKEN_INDICATOR = 3
|
LLM_TOKEN_INDICATOR = 3
|
||||||
IMAGE_POSITION_OFFSET = 65536
|
IMAGE_POSITION_OFFSET = 65536
|
||||||
|
|
||||||
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class LogitNormalSchedule:
|
class LogitNormalSchedule:
|
||||||
@@ -264,45 +268,31 @@ class Ideogram4DenoisingStage(DenoisingStage):
|
|||||||
def __init__(self, transformer, unconditional_transformer, pipeline=None) -> None:
|
def __init__(self, transformer, unconditional_transformer, pipeline=None) -> None:
|
||||||
super().__init__(
|
super().__init__(
|
||||||
transformer=transformer,
|
transformer=transformer,
|
||||||
|
transformer_2=unconditional_transformer,
|
||||||
scheduler=Ideogram4Scheduler(),
|
scheduler=Ideogram4Scheduler(),
|
||||||
pipeline=pipeline,
|
pipeline=pipeline,
|
||||||
)
|
)
|
||||||
self.unconditional_transformer = unconditional_transformer
|
self.unconditional_transformer = self.transformer_2
|
||||||
self._maybe_torch_compile(self.unconditional_transformer)
|
|
||||||
|
|
||||||
def _component_name_for_stage_module(self, module, default_name: str) -> str:
|
def _component_name_for_stage_module(self, module, default_name: str) -> str:
|
||||||
if module is self.unconditional_transformer:
|
if module is self.unconditional_transformer:
|
||||||
return "unconditional_transformer"
|
return "unconditional_transformer"
|
||||||
return super()._component_name_for_stage_module(module, default_name)
|
return super()._component_name_for_stage_module(module, default_name)
|
||||||
|
|
||||||
def component_uses(
|
def _cache_dit_dual_model_name(self) -> str:
|
||||||
self, server_args: ServerArgs, stage_name: str | None = None
|
return "ideogram4"
|
||||||
) -> list[ComponentUse]:
|
|
||||||
stage_name = self._component_stage_name(stage_name)
|
|
||||||
return [
|
|
||||||
ComponentUse(
|
|
||||||
stage_name=stage_name,
|
|
||||||
component_name="transformer",
|
|
||||||
phase="transformer",
|
|
||||||
preferred_ready_after_request=True,
|
|
||||||
memory_intensive=True,
|
|
||||||
),
|
|
||||||
ComponentUse(
|
|
||||||
stage_name=stage_name,
|
|
||||||
component_name="unconditional_transformer",
|
|
||||||
phase="unconditional_transformer",
|
|
||||||
memory_intensive=True,
|
|
||||||
),
|
|
||||||
]
|
|
||||||
|
|
||||||
def _maybe_enable_cache_dit_and_torch_compile(
|
def _dual_transformer_execution_mode(
|
||||||
self, num_inference_steps: int | tuple[int, int], batch: Req
|
self,
|
||||||
) -> None:
|
) -> DualTransformerExecutionMode | None:
|
||||||
self._maybe_enable_cache_dit(num_inference_steps, batch)
|
return DualTransformerExecutionMode.PAIRED_PER_STEP
|
||||||
for transformer in filter(
|
|
||||||
None, [self.transformer, self.unconditional_transformer]
|
def _cache_dit_secondary_uses_primary_config(self) -> bool:
|
||||||
):
|
return True
|
||||||
self._maybe_torch_compile(transformer)
|
|
||||||
|
def _maybe_enable_cache_dit(self, *args, **kwargs) -> None:
|
||||||
|
super()._maybe_enable_cache_dit(*args, **kwargs)
|
||||||
|
self.unconditional_transformer = self.transformer_2
|
||||||
|
|
||||||
def _manage_unconditional_transformer_use_site(self, batch: Req) -> None:
|
def _manage_unconditional_transformer_use_site(self, batch: Req) -> None:
|
||||||
manager = self._component_residency_manager
|
manager = self._component_residency_manager
|
||||||
|
|||||||
+15
-7
@@ -21,7 +21,7 @@ from diffusers.utils.torch_utils import randn_tensor
|
|||||||
|
|
||||||
from sglang.multimodal_gen.configs.sample.ideogram import IDEOGRAM4_PRESETS
|
from sglang.multimodal_gen.configs.sample.ideogram import IDEOGRAM4_PRESETS
|
||||||
from sglang.multimodal_gen.runtime.cache.cache_dit_integration import (
|
from sglang.multimodal_gen.runtime.cache.cache_dit_integration import (
|
||||||
refresh_context_on_transformer,
|
refresh_context_on_dual_transformer,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
|
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
|
||||||
from sglang.multimodal_gen.runtime.layers.attention import build_varlen_mask_meta
|
from sglang.multimodal_gen.runtime.layers.attention import build_varlen_mask_meta
|
||||||
@@ -179,6 +179,7 @@ class Ideogram4ProgressiveDenoisingStage(
|
|||||||
DenoisingStage.__init__(
|
DenoisingStage.__init__(
|
||||||
self,
|
self,
|
||||||
transformer=transformer,
|
transformer=transformer,
|
||||||
|
transformer_2=unconditional_transformer,
|
||||||
scheduler=Ideogram4Scheduler(),
|
scheduler=Ideogram4Scheduler(),
|
||||||
pipeline=pipeline,
|
pipeline=pipeline,
|
||||||
)
|
)
|
||||||
@@ -186,8 +187,7 @@ class Ideogram4ProgressiveDenoisingStage(
|
|||||||
self._spectrum_A = IDEOGRAM_SPECTRUM_A
|
self._spectrum_A = IDEOGRAM_SPECTRUM_A
|
||||||
self._spectrum_beta = IDEOGRAM_SPECTRUM_BETA
|
self._spectrum_beta = IDEOGRAM_SPECTRUM_BETA
|
||||||
# Ideogram4DenoisingStage extra transformer
|
# Ideogram4DenoisingStage extra transformer
|
||||||
self.unconditional_transformer = unconditional_transformer
|
self.unconditional_transformer = self.transformer_2
|
||||||
self._maybe_torch_compile(self.unconditional_transformer)
|
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
# Latent scale factor
|
# Latent scale factor
|
||||||
@@ -454,9 +454,17 @@ class Ideogram4ProgressiveDenoisingStage(
|
|||||||
self, n_remaining: int, scm_preset: str | None
|
self, n_remaining: int, scm_preset: str | None
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Refresh both conditional and unconditional transformers."""
|
"""Refresh both conditional and unconditional transformers."""
|
||||||
refresh_context_on_transformer(
|
# Recompute the full SCM config here so custom compute/cache bins stay
|
||||||
self.transformer, n_remaining, scm_preset=scm_preset
|
# active after a progressive-resolution stage transition.
|
||||||
|
_, scm_policy, steps_computation_mask, steps_computation_mask_2 = (
|
||||||
|
self._cache_dit_scm_masks(n_remaining, n_remaining)
|
||||||
)
|
)
|
||||||
refresh_context_on_transformer(
|
refresh_context_on_dual_transformer(
|
||||||
self.unconditional_transformer, n_remaining, scm_preset=scm_preset
|
self.transformer,
|
||||||
|
self.unconditional_transformer,
|
||||||
|
n_remaining,
|
||||||
|
n_remaining,
|
||||||
|
steps_computation_mask=steps_computation_mask,
|
||||||
|
steps_computation_mask_2=steps_computation_mask_2,
|
||||||
|
steps_computation_policy=scm_policy,
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user