[diffusion] feat: support cache-dit for Ideogram 4 (#29631)

This commit is contained in:
Jzz1943
2026-07-03 19:58:48 +08:00
committed by GitHub
parent 42acfd1550
commit 1058d00fe0
4 changed files with 269 additions and 164 deletions
@@ -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,38 +533,50 @@ 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
transformer_blocks = getattr(transformer, "blocks", None)
transformer_2_blocks = getattr(transformer_2, "blocks", None)
if transformer_blocks is None or transformer_2_blocks is None:
raise ValueError(
"Dual transformers must have 'blocks' attribute for cache-dit. "
f"transformer has blocks: {transformer_blocks is not None}, "
f"transformer_2 has blocks: {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": if model_name == "wan2.2":
# Use Pattern_2 for Wan2.2 dual-transformer. We should check `model_name` transformer_blocks = getattr(transformer, "blocks", None)
# to ensure we only apply this for supported models. Different models transformer_2_blocks = getattr(transformer_2, "blocks", None)
# may require different ForwardPattern. blocks_name = None
cache_dit.enable_cache( forward_pattern = [ForwardPattern.Pattern_2, ForwardPattern.Pattern_2]
BlockAdapter( check_forward_pattern = True
transformer=[transformer, transformer_2], check_num_outputs = False
blocks=[transformer_blocks, transformer_2_blocks], has_separate_cfg = True
forward_pattern=[ForwardPattern.Pattern_2, ForwardPattern.Pattern_2], elif model_name == "ideogram4":
params_modifiers=[primary_modifier, secondary_modifier], transformer_blocks = getattr(transformer, "layers", None)
has_separate_cfg=True, transformer_2_blocks = getattr(transformer_2, "layers", None)
), blocks_name = ["layers", "layers"]
parallelism_config=None, forward_pattern = [ForwardPattern.Pattern_3, ForwardPattern.Pattern_3]
) check_forward_pattern = False
check_num_outputs = False
has_separate_cfg = False
else: else:
raise ValueError( raise ValueError(
f"Dual-transformer is not implemented for model {model_name} yet." f"Dual-transformer is not implemented for model {model_name} yet."
) )
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}"
)
cache_dit.enable_cache(
BlockAdapter(
transformer=[transformer, transformer_2],
blocks=[transformer_blocks, transformer_2_blocks],
blocks_name=blocks_name,
forward_pattern=forward_pattern,
params_modifiers=[primary_modifier, secondary_modifier],
check_forward_pattern=check_forward_pattern,
check_num_outputs=check_num_outputs,
has_separate_cfg=has_separate_cfg,
),
parallelism_config=None,
)
if parallelism_config is not None: if parallelism_config is not None:
for t in [transformer, transformer_2]: for t in [transformer, transformer_2]:
context_manager = getattr(t, "_context_manager", None) context_manager = getattr(t, "_context_manager", None)
@@ -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,
)
# generate SCM mask using cache-dit's steps_mask()
# 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,
) )
primary_config = self._build_cache_dit_config(
if isinstance(num_inference_steps, tuple): primary_num_steps,
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, steps_computation_mask=steps_computation_mask_2,
enable_taylorseer=envs.SGLANG_CACHE_DIT_SECONDARY_TAYLORSEER, scm_policy=scm_policy,
taylorseer_order=envs.SGLANG_CACHE_DIT_SECONDARY_TS_ORDER, secondary=True,
num_inference_steps=num_low_noise_steps, )
# SCM fields - shared with primary
steps_computation_mask=steps_computation_mask_2,
steps_computation_policy=scm_policy,
) )
# 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
@@ -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
@@ -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,
) )