From 1058d00fe07b6d75758824cb5de98f08960d0cbf Mon Sep 17 00:00:00 2001 From: Jzz1943 Date: Fri, 3 Jul 2026 19:58:48 +0800 Subject: [PATCH] [diffusion] feat: support cache-dit for Ideogram 4 (#29631) --- .../runtime/cache/cache_dit_integration.py | 86 +++--- .../pipelines_core/stages/denoising.py | 277 ++++++++++++------ .../stages/model_specific_stages/ideogram.py | 48 ++- .../stages/progressive_resolution/ideogram.py | 22 +- 4 files changed, 269 insertions(+), 164 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/cache/cache_dit_integration.py b/python/sglang/multimodal_gen/runtime/cache/cache_dit_integration.py index 14efc7c1c..d8aa17926 100644 --- a/python/sglang/multimodal_gen/runtime/cache/cache_dit_integration.py +++ b/python/sglang/multimodal_gen/runtime/cache/cache_dit_integration.py @@ -415,7 +415,8 @@ def enable_cache_on_dual_transformer( tp_group: Tensor parallel process group. """ _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: raise ValueError( @@ -494,7 +495,7 @@ def enable_cache_on_dual_transformer( primary_config.enable_taylorseer, ) 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.Bn_compute_blocks, secondary_config.max_warmup_steps, @@ -532,38 +533,50 @@ def enable_cache_on_dual_transformer( 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": - # 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( - BlockAdapter( - transformer=[transformer, transformer_2], - blocks=[transformer_blocks, transformer_2_blocks], - forward_pattern=[ForwardPattern.Pattern_2, ForwardPattern.Pattern_2], - params_modifiers=[primary_modifier, secondary_modifier], - has_separate_cfg=True, - ), - parallelism_config=None, - ) + transformer_blocks = getattr(transformer, "blocks", None) + transformer_2_blocks = getattr(transformer_2, "blocks", None) + blocks_name = 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( 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: for t in [transformer, transformer_2]: context_manager = getattr(t, "_context_manager", None) @@ -612,23 +625,30 @@ def refresh_context_on_dual_transformer( num_low_noise_steps: int, scm_preset: str | None = None, 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: """Refresh cache-dit context for dual transformers.""" - high_noise_steps_computation_mask = None - low_noise_steps_computation_mask = None - if scm_preset is not None: + high_noise_steps_computation_mask = steps_computation_mask + low_noise_steps_computation_mask = steps_computation_mask_2 + if high_noise_steps_computation_mask is None and scm_preset is not None: high_noise_steps_computation_mask = cache_dit.steps_mask( 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( 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( transformer, cache_config=DBCacheConfig().reset( num_inference_steps=num_high_noise_steps, steps_computation_mask=high_noise_steps_computation_mask, - steps_computation_policy=scm_preset, + steps_computation_policy=policy, ), verbose=verbose, ) @@ -637,7 +657,7 @@ def refresh_context_on_dual_transformer( cache_config=DBCacheConfig().reset( num_inference_steps=num_low_noise_steps, steps_computation_mask=low_noise_steps_computation_mask, - steps_computation_policy=scm_preset, + steps_computation_policy=policy, ), verbose=verbose, ) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py index bcee12b2b..56ddb8b54 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -13,6 +13,7 @@ import weakref from collections.abc import Callable from contextlib import contextmanager from dataclasses import dataclass, field, fields +from enum import Enum from functools import lru_cache from typing import Any @@ -176,6 +177,17 @@ class DenoisingStepState: 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): """ Stage for running the denoising loop in diffusion pipelines. @@ -417,6 +429,133 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): return True 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( self, num_inference_steps: int | tuple[int, int], batch: Req ) -> None: @@ -429,25 +568,30 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): 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. if self._cache_dit_enabled: - scm_preset = envs.SGLANG_CACHE_DIT_SCM_PRESET - scm_preset = None if scm_preset == "none" else scm_preset - if isinstance(num_inference_steps, tuple): + primary_num_steps, secondary_num_steps = self._cache_dit_step_counts( + num_inference_steps + ) + 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( self.transformer, self.transformer_2, - num_high_noise_steps, - num_low_noise_steps, - scm_preset=scm_preset, + primary_num_steps, + secondary_num_steps, + steps_computation_mask=steps_computation_mask, + steps_computation_mask_2=steps_computation_mask_2, + steps_computation_policy=scm_policy, ) else: + scm_preset = None if scm_preset == "none" else scm_preset refresh_context_on_transformer( self.transformer, - num_inference_steps, + primary_num_steps, scm_preset=scm_preset, ) return @@ -459,6 +603,9 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): if batch.is_warmup and not self.server_args.enable_torch_compile: return + primary_num_steps, secondary_num_steps = self._cache_dit_step_counts( + num_inference_steps + ) world_size = get_world_size() parallelized = world_size > 1 @@ -483,93 +630,29 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): has_sp, has_tp, ) - # === Parse SCM configuration from envs === - # SCM is shared between primary and secondary transformers - 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, + _, scm_policy, steps_computation_mask, steps_computation_mask_2 = ( + self._cache_dit_scm_masks(primary_num_steps, secondary_num_steps) ) - - 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 + primary_config = self._build_cache_dit_config( + primary_num_steps, steps_computation_mask=steps_computation_mask, - steps_computation_policy=scm_policy, + scm_policy=scm_policy, ) if self.transformer_2 is not None: # dual transformer # build config for secondary transformer (low-noise expert) # uses secondary parameters which inherit from primary if not explicitly set - secondary_config = CacheDitConfig( - enabled=True, - Fn_compute_blocks=envs.SGLANG_CACHE_DIT_SECONDARY_FN, - Bn_compute_blocks=envs.SGLANG_CACHE_DIT_SECONDARY_BN, - max_warmup_steps=envs.SGLANG_CACHE_DIT_SECONDARY_WARMUP, - residual_diff_threshold=envs.SGLANG_CACHE_DIT_SECONDARY_RDT, - 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_policy=scm_policy, + assert secondary_num_steps is not None + secondary_config = ( + primary_config + if self._cache_dit_secondary_uses_primary_config() + else self._build_cache_dit_config( + secondary_num_steps, + steps_computation_mask=steps_computation_mask_2, + scm_policy=scm_policy, + secondary=True, + ) ) # for dual transformers, must use BlockAdapter to enable cache on both simultaneously. @@ -579,14 +662,14 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): self.transformer_2, primary_config, secondary_config, - model_name="wan2.2", + model_name=self._cache_dit_dual_model_name(), sp_group=sp_group, tp_group=tp_group, ) logger.info( "cache-dit enabled on dual transformers (steps=%d, %d)", - num_high_noise_steps, - num_low_noise_steps, + primary_num_steps, + secondary_num_steps, ) else: # single transformer @@ -600,7 +683,7 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): ) logger.info( "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_BN, envs.SGLANG_CACHE_DIT_RDT, @@ -683,9 +766,13 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): scheduler = batch.scheduler 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 = ( self._handle_boundary_ratio(server_args, batch, scheduler) - if self.transformer_2 is not None + if uses_boundary_transformer_2 else None ) # Get timesteps and calculate warmup steps @@ -693,7 +780,7 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): num_inference_steps = batch.num_inference_steps 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" num_high_noise_steps = (timesteps >= boundary_timestep).sum().item() num_low_noise_steps = num_inference_steps - num_high_noise_steps diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ideogram.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ideogram.py index 6eba96640..9611058df 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ideogram.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ideogram.py @@ -25,6 +25,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import ( DenoisingContext, DenoisingStage, DenoisingStepState, + DualTransformerExecutionMode, ) from sglang.multimodal_gen.runtime.pipelines_core.stages.text_encoding import ( TextEncodingStage, @@ -36,6 +37,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import ( VerificationResult, ) 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.utils import PRECISION_TO_TYPE @@ -44,6 +46,8 @@ OUTPUT_IMAGE_INDICATOR = 2 LLM_TOKEN_INDICATOR = 3 IMAGE_POSITION_OFFSET = 65536 +logger = init_logger(__name__) + @dataclass(frozen=True) class LogitNormalSchedule: @@ -264,45 +268,31 @@ class Ideogram4DenoisingStage(DenoisingStage): def __init__(self, transformer, unconditional_transformer, pipeline=None) -> None: super().__init__( transformer=transformer, + transformer_2=unconditional_transformer, scheduler=Ideogram4Scheduler(), pipeline=pipeline, ) - self.unconditional_transformer = unconditional_transformer - self._maybe_torch_compile(self.unconditional_transformer) + self.unconditional_transformer = self.transformer_2 def _component_name_for_stage_module(self, module, default_name: str) -> str: if module is self.unconditional_transformer: return "unconditional_transformer" return super()._component_name_for_stage_module(module, default_name) - def component_uses( - self, server_args: ServerArgs, stage_name: str | None = None - ) -> 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 _cache_dit_dual_model_name(self) -> str: + return "ideogram4" - def _maybe_enable_cache_dit_and_torch_compile( - self, num_inference_steps: int | tuple[int, int], batch: Req - ) -> None: - self._maybe_enable_cache_dit(num_inference_steps, batch) - for transformer in filter( - None, [self.transformer, self.unconditional_transformer] - ): - self._maybe_torch_compile(transformer) + def _dual_transformer_execution_mode( + self, + ) -> DualTransformerExecutionMode | None: + return DualTransformerExecutionMode.PAIRED_PER_STEP + + def _cache_dit_secondary_uses_primary_config(self) -> bool: + return True + + 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: manager = self._component_residency_manager diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/progressive_resolution/ideogram.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/progressive_resolution/ideogram.py index 94433bcf2..f73c7fc8f 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/progressive_resolution/ideogram.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/progressive_resolution/ideogram.py @@ -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.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.layers.attention import build_varlen_mask_meta @@ -179,6 +179,7 @@ class Ideogram4ProgressiveDenoisingStage( DenoisingStage.__init__( self, transformer=transformer, + transformer_2=unconditional_transformer, scheduler=Ideogram4Scheduler(), pipeline=pipeline, ) @@ -186,8 +187,7 @@ class Ideogram4ProgressiveDenoisingStage( self._spectrum_A = IDEOGRAM_SPECTRUM_A self._spectrum_beta = IDEOGRAM_SPECTRUM_BETA # Ideogram4DenoisingStage extra transformer - self.unconditional_transformer = unconditional_transformer - self._maybe_torch_compile(self.unconditional_transformer) + self.unconditional_transformer = self.transformer_2 # ------------------------------------------------------------------ # Latent scale factor @@ -454,9 +454,17 @@ class Ideogram4ProgressiveDenoisingStage( self, n_remaining: int, scm_preset: str | None ) -> None: """Refresh both conditional and unconditional transformers.""" - refresh_context_on_transformer( - self.transformer, n_remaining, scm_preset=scm_preset + # Recompute the full SCM config here so custom compute/cache bins stay + # 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( - self.unconditional_transformer, n_remaining, scm_preset=scm_preset + refresh_context_on_dual_transformer( + 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, )