diff --git a/docs_new/docs/sglang-diffusion/progressive_resolution.mdx b/docs_new/docs/sglang-diffusion/progressive_resolution.mdx index f38fc8d24..2bc23323e 100644 --- a/docs_new/docs/sglang-diffusion/progressive_resolution.mdx +++ b/docs_new/docs/sglang-diffusion/progressive_resolution.mdx @@ -4,7 +4,7 @@ description: "Experimental spectral progressive resolution growing for selected tag: "approx" --- -Progressive resolution growing is an experimental feature for selected SGLang Diffusion pipelines. It runs early denoising steps at a coarser latent resolution and spectrally upsamples the latent before the full-resolution steps. On the benchmark setup below, this reduces the quadratic attention cost of the DiT transformer and yields up to **1.63× speedup on FLUX.1**, **1.93× speedup on FLUX.2**, **2.33× speedup on Z-Image**, **2.78× speedup on Wan 2.1 T2V**, and **1.69× speedup on Qwen-Image**. +Progressive resolution growing is an experimental feature for selected SGLang Diffusion pipelines. It runs early denoising steps at a coarser latent resolution and spectrally upsamples the latent before the full-resolution steps. On the benchmark setup below, this reduces the quadratic attention cost of the DiT transformer and yields up to **1.63× speedup on FLUX.1**, **1.93× speedup on FLUX.2**, **2.33× speedup on Z-Image**, **2.78× speedup on Wan 2.1 T2V**, **1.69× speedup on Qwen-Image**, and **1.56× speedup on Ideogram 4**. Based on [Spectral Progressive Diffusion (arXiv 2605.18736)](https://arxiv.org/abs/2605.18736). @@ -20,6 +20,7 @@ The transition point — how many steps to run at each resolution — is compute | FLUX.2 1024×1024 | 4,096 | 1,024 | 4.0× | | Z-Image 1024×1024 | 4,096 | 1,024 | 4.0× | | Wan 2.1 T2V 480×832 (81 frames) | 6,240 | 1,560 | 4.0× | +| Ideogram 4 1024×1024 | 4,096 | 1,024 | 4.0× | ## Parameters @@ -280,6 +281,86 @@ Hardware: RTX A6000 48 GB, `--dit-cpu-offload false`. Timing = denoising loop on | dct_rewind L1 δ=0.10 | 16@64² + 14@128² | 33.86 s | **1.27×** | | dct_rewind L1 δ=0.20 | 19@64² + 11@128² | 25.40 s | **1.69×** | +## Ideogram 4 + +Supports `ideogram-ai/ideogram-4`. Ideogram 4 uses a **dual-transformer architecture**: a conditional transformer (text + image tokens) and a separately-weighted unconditional transformer (image tokens only, zero LLM features). Both transformers shrink at coarse resolution, providing the same token-ratio benefit as single-transformer models. + +> **Note:** Ideogram 4's logit-normal noise schedule (`std=1.75`, `mu=0`) concentrates steps near the mid-sigma range. Fewer steps fall in the high-sigma coarse-eligible region compared to FLUX, which limits the achievable speedup at a given δ. + +### Usage + +**20-step (V4_DEFAULT_20 preset)** + +```bash +sglang generate \ + --model-path ideogram-ai/ideogram-4 \ + --prompt "A serene mountain lake at golden hour, photorealistic" \ + --height 1024 --width 1024 \ + --num-inference-steps 20 \ + --dit-cpu-offload false \ + --progressive-mode dct_rewind \ + --progressive-levels 1 \ + --progressive-delta 0.05 +``` + +**48-step (V4_QUALITY_48 preset)** + +```bash +sglang generate \ + --model-path ideogram-ai/ideogram-4 \ + --prompt "A serene mountain lake at golden hour, photorealistic" \ + --height 1024 --width 1024 \ + --num-inference-steps 48 \ + --dit-cpu-offload false \ + --progressive-mode dct_rewind \ + --progressive-levels 1 \ + --progressive-delta 0.05 +``` + +### Benchmark + +Hardware: RTX A6000 48 GB, `torch_sdpa`, `--dit-cpu-offload false`. Timing = denoising loop only. + +**20-step (V4_DEFAULT_20)** + +| Config | Stage split | Denoise | Speedup | +|--------|-------------|---------|---------| +| Fullres (baseline) | 20 @ 64² | 53.99 s | 1.00× | +| dct_rewind L1 δ=0.01 | 6 @ 32² + 14 @ 64² | 43.47 s | **1.24×** | +| dct_rewind L1 δ=0.05 | 9 @ 32² + 11 @ 64² | 38.14 s | **1.42×** | +| dct_rewind L1 δ=0.10 | 11 @ 32² + 9 @ 64² | 34.60 s | **1.56×** | + +**48-step (V4_QUALITY_48)** + +| Config | Stage split | Denoise | Speedup | +|--------|-------------|---------|---------| +| Fullres (baseline) | 48 @ 64² | 130.92 s | 1.00× | +| dct_rewind L1 δ=0.01 | 12 @ 32² + 36 @ 64² | 109.79 s | **1.19×** | +| dct_rewind L1 δ=0.05 | 21 @ 32² + 27 @ 64² | 93.83 s | **1.40×** | +| dct_rewind L1 δ=0.10 | 26 @ 32² + 22 @ 64² | 84.94 s | **1.54×** | + +### Python API + +```python +from sglang.multimodal_gen import DiffGenerator + +gen = DiffGenerator.from_pretrained( + model_path="ideogram-ai/ideogram-4", + dit_cpu_offload=False, +) +result = gen.generate(sampling_params_kwargs={ + "prompt": "A serene mountain lake at golden hour, photorealistic", + "num_inference_steps": 48, + "height": 1024, + "width": 1024, + "progressive_mode": "dct_rewind", + "progressive_levels": 1, + "progressive_delta": 0.05, +}) +``` + +--- + ## Limitations - **Sequence parallelism incompatible.** Cannot be combined with `--ulysses-degree` or `--ring-degree`. The stage raises a `RuntimeError` if SP is enabled. diff --git a/python/sglang/multimodal_gen/runtime/pipelines/ideogram.py b/python/sglang/multimodal_gen/runtime/pipelines/ideogram.py index 92c7c369c..d1f8e1a7b 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines/ideogram.py +++ b/python/sglang/multimodal_gen/runtime/pipelines/ideogram.py @@ -18,6 +18,12 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.i Ideogram4DenoisingStage, Ideogram4TextEncodingStage, ) +from sglang.multimodal_gen.runtime.pipelines_core.stages.progressive_resolution.denoising import ( + ProgressiveDenoisingStageRouter, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.progressive_resolution.ideogram import ( + Ideogram4ProgressiveDenoisingStage, +) from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import ( maybe_download_model, @@ -115,6 +121,22 @@ class Ideogram4Pipeline(LoRAPipeline, ComposedPipelineBase): "scheduler", ] + def _create_denoising_stage(self): + transformer = self.get_module("transformer") + unconditional_transformer = self.get_module("unconditional_transformer") + return ProgressiveDenoisingStageRouter( + standard_stage=Ideogram4DenoisingStage( + transformer=transformer, + unconditional_transformer=unconditional_transformer, + pipeline=self, + ), + progressive_stage_factory=lambda: Ideogram4ProgressiveDenoisingStage( + transformer=transformer, + unconditional_transformer=unconditional_transformer, + pipeline=self, + ), + ) + def create_pipeline_stages(self, server_args: ServerArgs): self.add_stage(InputValidationStage()) self.add_stage_factory( @@ -128,11 +150,7 @@ class Ideogram4Pipeline(LoRAPipeline, ComposedPipelineBase): self.add_standard_latent_preparation_stage() self.add_stage_factory( RoleType.DENOISER, - lambda: Ideogram4DenoisingStage( - transformer=self.get_module("transformer"), - unconditional_transformer=self.get_module("unconditional_transformer"), - pipeline=self, - ), + self._create_denoising_stage, "ideogram4_denoising_stage", ) self.add_stage_factory( diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/progressive_resolution/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/progressive_resolution/denoising.py index cc9db15f3..6a3a32b21 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/progressive_resolution/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/progressive_resolution/denoising.py @@ -258,6 +258,18 @@ class ProgressiveDenoisingStage(DenoisingStage): """ return server_args.pipeline_config.vae_config.arch_config.vae_scale_factor + def _spectrum_latent_dims( + self, batch: Req, server_args: ServerArgs, H_lat: int, W_lat: int + ) -> tuple[int, int]: + """Physical spatial-latent dims for the Nyquist-frequency calculation. + + By default these equal the grid dims returned by _latent_scale_factor. + Override for models (e.g. Ideogram 4) where patch packing causes the + grid dimension to be smaller than the true spatial-latent dimension, + so that the spectrum threshold is computed at the correct scale. + """ + return H_lat, W_lat + def _unpack_latent( self, latent: torch.Tensor, h_lat: int, w_lat: int ) -> torch.Tensor: @@ -286,6 +298,29 @@ class ProgressiveDenoisingStage(DenoisingStage): """Called after each stage transition. Update resolution-dependent state.""" pass + def _refresh_cache_dit_context( + self, n_remaining: int, scm_preset: str | None + ) -> None: + """Refresh cache-dit activations and step counter at a stage transition. + + Override in model-specific subclasses that use more than one transformer + (e.g. models with a separate unconditional branch). + """ + if self.transformer_2 is not None: + n_high = n_remaining // 2 + n_low = n_remaining - n_high + refresh_context_on_dual_transformer( + self.transformer, + self.transformer_2, + n_high, + n_low, + scm_preset=scm_preset, + ) + else: + refresh_context_on_transformer( + self.transformer, n_remaining, scm_preset=scm_preset + ) + # ------------------------------------------------------------------ # Helpers # ------------------------------------------------------------------ @@ -441,9 +476,12 @@ class ProgressiveDenoisingStage(DenoisingStage): init_h_lat = H_lat // downsample init_w_lat = W_lat // downsample - # Compute stage transitions from the power-law spectrum + # Compute stage transitions from the power-law spectrum. + # Use physical spatial-latent dims (may differ from grid dims for + # patch-packed models like Ideogram 4). + H_spec, W_spec = self._spectrum_latent_dims(batch, server_args, H_lat, W_lat) stage_sigmas = compute_stage_transitions( - delta, levels, self._spectrum_A, self._spectrum_beta, H_lat, W_lat + delta, levels, self._spectrum_A, self._spectrum_beta, H_spec, W_spec ) num_stages = len(stage_sigmas) @@ -572,23 +610,7 @@ class ProgressiveDenoisingStage(DenoisingStage): # residual-diff decision for the first full-res steps. if self._cache_dit_enabled: n_remaining = n_steps - stage_end - scm_preset = _get_scm_preset() - if self.transformer_2 is not None: - n_high = n_remaining // 2 - n_low = n_remaining - n_high - refresh_context_on_dual_transformer( - self.transformer, - self.transformer_2, - n_high, - n_low, - scm_preset=scm_preset, - ) - else: - refresh_context_on_transformer( - self.transformer, - n_remaining, - scm_preset=scm_preset, - ) + self._refresh_cache_dit_context(n_remaining, _get_scm_preset()) logger.info( "cache-dit context refreshed at stage transition " "(step %d, %d steps remaining)", diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/progressive_resolution/flux.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/progressive_resolution/flux.py index df4ac601c..232531c2d 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/progressive_resolution/flux.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/progressive_resolution/flux.py @@ -30,6 +30,10 @@ logger = init_logger(__name__) FLUX_SPECTRUM_A: float = 203.615097 FLUX_SPECTRUM_BETA: float = 1.915461 +# Module-level aliases for test discoverability and external reuse. +_flux_unpack = unpack_2x2_latent +_flux_pack = pack_2x2_latent + class FluxProgressiveDenoisingStage(ProgressiveDenoisingStage): """FLUX-specific progressive denoising stage. 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 new file mode 100644 index 000000000..f4abb0300 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/progressive_resolution/ideogram.py @@ -0,0 +1,462 @@ +# SPDX-License-Identifier: Apache-2.0 +""" +Ideogram 4 progressive-resolution denoising stage. + +Ideogram 4 latent layout: + packed: [B, grid_h * grid_w, in_channels] (row-major, same as FLUX.2) + spatial: [B, in_channels, grid_h, grid_w] + +where grid_h = height // (patch_size * ae_scale_factor) = height // 16 + grid_w = width // (patch_size * ae_scale_factor) = width // 16 + +On each stage transition _on_resolution_change rebuilds the position_ids, +segment_ids, indicator, attention masks, and the zero neg_llm_features tensor +that Ideogram4DenoisingStage reads from batch.extra["ideogram4"] and ctx.extra. +""" + +from __future__ import annotations + +import torch +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, +) +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.pipelines_core.schedule_batch import Req +from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import ( + DenoisingContext, + DenoisingStage, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ideogram import ( + IMAGE_POSITION_OFFSET, + OUTPUT_IMAGE_INDICATOR, + Ideogram4DenoisingStage, + Ideogram4Scheduler, + get_schedule_for_resolution, + make_step_intervals, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.progressive_resolution.denoising import ( + ProgressiveDenoisingStage, +) +from sglang.multimodal_gen.runtime.server_args import ServerArgs +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +logger = init_logger(__name__) + +# Power-law spectrum constants. +# Using FLUX.1-dev VAE values as a placeholder until Ideogram-specific +# coefficients are fitted on a representative latent dataset. +IDEOGRAM_SPECTRUM_A: float = 203.615097 +IDEOGRAM_SPECTRUM_BETA: float = 1.915461 + + +def _adapt_llm_features( + llm_features: torch.Tensor, + max_text_tokens: int, + new_num_image_tokens: int, +) -> torch.Tensor: + """Adapt LLM features to a different image-token count. + + The text encoder produces [text_tokens | image_tokens] features at full-res. + For progressive denoising the grid size changes between stages, so the + image portion must be resized. Following the reference Ideogram inference + code (run_experiment.py), we keep the text portion intact and use zeros for + image positions at the new resolution — the unconditional transformer always + receives zero image features anyway, so the model is designed to handle this. + """ + existing_image_tokens = llm_features.shape[1] - max_text_tokens + if existing_image_tokens == new_num_image_tokens: + return llm_features + B, _T, D = llm_features.shape + text_feat = llm_features[:, :max_text_tokens] + image_feat_new = llm_features.new_zeros(B, new_num_image_tokens, D) + return torch.cat([text_feat, image_feat_new], dim=1) + + +def _ideogram4_unpack(latent: torch.Tensor, h_lat: int, w_lat: int) -> torch.Tensor: + """Packed [B, grid_h*grid_w, C] → spatial [B, C, grid_h, grid_w] (row-major).""" + B, _S, C = latent.shape + return latent.permute(0, 2, 1).reshape(B, C, h_lat, w_lat) + + +def _ideogram4_pack(x: torch.Tensor) -> torch.Tensor: + """Spatial [B, C, grid_h, grid_w] → packed [B, grid_h*grid_w, C] (row-major).""" + B, C, H, W = x.shape + return x.reshape(B, C, H * W).permute(0, 2, 1) + + +def _build_ideogram4_seq_tensors( + data: dict, + grid_h: int, + grid_w: int, + batch_size: int, + device: torch.device, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Build full-sequence position_ids, segment_ids, indicator for a grid. + + Keeps the text prefix from data unchanged and appends new image-token rows. + Returns (position_ids, segment_ids, indicator), each shaped [B, T+S, ...]. + """ + max_text_tokens = data["max_text_tokens"] + num_image_tokens = grid_h * grid_w + h_idx = ( + torch.arange(grid_h, device=device) + .view(-1, 1) + .expand(grid_h, grid_w) + .reshape(-1) + ) + w_idx = ( + torch.arange(grid_w, device=device) + .view(1, -1) + .expand(grid_h, grid_w) + .reshape(-1) + ) + image_pos = ( + torch.stack([torch.zeros_like(h_idx), h_idx, w_idx], dim=1) + + IMAGE_POSITION_OFFSET + ) + position_ids = torch.cat( + [ + data["position_ids"][:, :max_text_tokens], + image_pos.unsqueeze(0).expand(batch_size, -1, -1), + ], + dim=1, + ) + segment_ids = torch.cat( + [ + data["segment_ids"][:, :max_text_tokens], + torch.ones(batch_size, num_image_tokens, dtype=torch.long, device=device), + ], + dim=1, + ) + indicator = torch.cat( + [ + data["indicator"][:, :max_text_tokens], + torch.full( + (batch_size, num_image_tokens), + OUTPUT_IMAGE_INDICATOR, + dtype=torch.long, + device=device, + ), + ], + dim=1, + ) + return position_ids, segment_ids, indicator + + +class Ideogram4ProgressiveDenoisingStage( + ProgressiveDenoisingStage, Ideogram4DenoisingStage +): + """Progressive-resolution denoising stage for Ideogram 4. + + Inherits the progressive loop from ProgressiveDenoisingStage and the + Ideogram-specific dual-transformer forward pass from Ideogram4DenoisingStage + via MRO. __init__ calls DenoisingStage directly to avoid cooperative-init + incompatibility between the two parent signatures. + + MRO for method resolution: + Ideogram4ProgressiveDenoisingStage + → ProgressiveDenoisingStage (forward / _run_stage_steps / latent hooks) + → Ideogram4DenoisingStage (_prepare_denoising_loop / _run_denoising_step) + → DenoisingStage (shared infrastructure) + → PipelineStage + """ + + def __init__( + self, + transformer, + unconditional_transformer, + pipeline=None, + ) -> None: + # Bypass cooperative __init__: the two parents have incompatible + # signatures (ProgressiveDenoisingStage takes scheduler/spectrum args; + # Ideogram4DenoisingStage takes unconditional_transformer). + # Initialise DenoisingStage — the common ancestor — directly, then + # set the attributes each parent __init__ would have added. + DenoisingStage.__init__( + self, + transformer=transformer, + scheduler=Ideogram4Scheduler(), + pipeline=pipeline, + ) + # ProgressiveDenoisingStage spectrum constants + self._spectrum_A = IDEOGRAM_SPECTRUM_A + self._spectrum_beta = IDEOGRAM_SPECTRUM_BETA + # Ideogram4DenoisingStage extra transformer + self.unconditional_transformer = unconditional_transformer + self._maybe_enable_torch_compile(self.unconditional_transformer) + + # ------------------------------------------------------------------ + # Latent scale factor + # ------------------------------------------------------------------ + + def _latent_scale_factor(self, server_args: ServerArgs) -> int: + # pixel → latent-grid: divide by patch_size (2) × ae_scale_factor (8) = 16 + cfg = server_args.pipeline_config + return cfg.patch_size * cfg.ae_scale_factor + + def _spectrum_latent_dims(self, batch, server_args, H_lat: int, W_lat: int): + # Ideogram 4 packs 2×2 latent patches per grid token. H_lat here is + # the grid dimension (= image_h // 16 = 64 for 1024-px input); the + # physical spatial-latent dimension is grid × patch_size (= 128). + # The Nyquist calculation must use physical dims to match the reference. + patch = server_args.pipeline_config.patch_size # 2 + return H_lat * patch, W_lat * patch + + # ------------------------------------------------------------------ + # Pack / Unpack + # ------------------------------------------------------------------ + + def _unpack_latent( + self, latent: torch.Tensor, h_lat: int, w_lat: int + ) -> torch.Tensor: + return _ideogram4_unpack(latent, h_lat, w_lat) + + def _repack_latent( + self, + x_spatial: torch.Tensor, + h_lat: int, + w_lat: int, + batch: Req, + server_args: ServerArgs, + ) -> torch.Tensor: + return _ideogram4_pack(x_spatial) + + # ------------------------------------------------------------------ + # Initial noise generation + # ------------------------------------------------------------------ + + def _generate_initial_noise( + self, + batch: Req, + server_args: ServerArgs, + h_lat: int, + w_lat: int, + seed, + ) -> torch.Tensor: + """Generate low-res packed noise for the initial progressive stage. + + Uses in_channels directly (no //4) because the spatial latent already + incorporates the patchification channel expansion (same as FLUX.2). + Ideogram denoising steps cast latents to fp32 internally, so we + generate fp32 noise to match. + """ + device = get_local_torch_device() + C = server_args.pipeline_config.dit_config.arch_config.in_channels + noise_spatial = randn_tensor( + (self._initial_noise_batch_size(batch), C, h_lat, w_lat), + generator=self._get_initial_noise_generator(batch, seed, device), + device=device, + dtype=torch.float32, + ) + return _ideogram4_pack(noise_spatial) + + # ------------------------------------------------------------------ + # Resolution-change hook + # ------------------------------------------------------------------ + + def _on_resolution_change( + self, + ctx: DenoisingContext, + batch: Req, + server_args: ServerArgs, + new_h_pixel: int, + new_w_pixel: int, + ) -> None: + """Rebuild Ideogram position IDs and attention masks for the new grid. + + Called after ctx.latents and batch.height/width are already updated to + the upsampled resolution. Patches batch.extra["ideogram4"] and ctx.extra + in-place so that _run_denoising_step sees correctly-sized tensors. + """ + if ctx.cfg_policy is None: + return + + cfg = server_args.pipeline_config + patch = cfg.patch_size * cfg.ae_scale_factor # 16 + grid_h = new_h_pixel // patch + grid_w = new_w_pixel // patch + num_image_tokens = grid_h * grid_w + + data = batch.extra["ideogram4"] + max_text_tokens = data["max_text_tokens"] + batch_size = ctx.latents.shape[0] + device = ctx.latents.device + + new_position_ids, new_segment_ids, new_indicator = _build_ideogram4_seq_tensors( + data, grid_h, grid_w, batch_size, device + ) + new_attn_mask = new_segment_ids > 0 + + # Negative (unconditional) tensors span the image tokens only. + neg_position_ids = new_position_ids[:, max_text_tokens:] + neg_segment_ids = new_segment_ids[:, max_text_tokens:] + neg_indicator = new_indicator[:, max_text_tokens:] + neg_attn_mask = neg_segment_ids > 0 + + llm_dim = ctx.extra["ideogram4_neg_llm_features"].shape[-1] + neg_llm_features = ctx.extra["ideogram4_neg_llm_features"].new_zeros( + batch_size, num_image_tokens, llm_dim + ) + + # Update batch.extra["ideogram4"] in-place. + data["position_ids"] = new_position_ids + data["segment_ids"] = new_segment_ids + data["indicator"] = new_indicator + data["num_image_tokens"] = num_image_tokens + data["grid_h"] = grid_h + data["grid_w"] = grid_w + + # Update ctx.extra in-place. + ctx.extra.update( + { + "ideogram4_attn_mask": new_attn_mask, + "ideogram4_attn_mask_meta": build_varlen_mask_meta(new_attn_mask), + "ideogram4_neg_position_ids": neg_position_ids, + "ideogram4_neg_segment_ids": neg_segment_ids, + "ideogram4_neg_indicator": neg_indicator, + "ideogram4_neg_attn_mask": neg_attn_mask, + "ideogram4_neg_attn_mask_meta": build_varlen_mask_meta(neg_attn_mask), + "ideogram4_neg_llm_features": neg_llm_features, + } + ) + + # Adapt LLM features to the new grid: keep text portion, use zeros for + # image positions at the new resolution (mirrors run_experiment.py). + full_res_llm = ctx.extra.get("ideogram4_full_res_llm_features") + if full_res_llm is not None: + batch.prompt_embeds[0] = _adapt_llm_features( + full_res_llm, max_text_tokens, num_image_tokens + ) + + logger.info( + "Updated position_ids / attn_masks / llm_features for %dx%d latent grid " + "(%d image tokens) across %d batch item(s)", + grid_h, + grid_w, + num_image_tokens, + batch_size, + ) + + # ------------------------------------------------------------------ + # Denoising-loop preparation + # ------------------------------------------------------------------ + + def _prepare_denoising_loop( + self, batch: Req, server_args: ServerArgs + ) -> DenoisingContext: + # ProgressiveDenoisingStage.forward() overrides batch.height/width to + # the initial low-res pixel dimensions before calling this method. + # batch.extra["ideogram4"] was built by the text-encoding stage at the + # full-res grid size, so we must resize position_ids / segment_ids / + # indicator / num_image_tokens to match the low-res grid BEFORE calling + # Ideogram4DenoisingStage._prepare_denoising_loop, which reads them to + # build the negative tensors, attn masks, and neg_llm_features. + cfg = server_args.pipeline_config + patch = cfg.patch_size * cfg.ae_scale_factor # 16 + grid_h = batch.height // patch + grid_w = batch.width // patch + num_image_tokens = grid_h * grid_w + + data = batch.extra["ideogram4"] + max_text_tokens = data["max_text_tokens"] + device = data["position_ids"].device + batch_size = data["position_ids"].shape[0] + + data["position_ids"], data["segment_ids"], data["indicator"] = ( + _build_ideogram4_seq_tensors(data, grid_h, grid_w, batch_size, device) + ) + data["num_image_tokens"] = num_image_tokens + data["grid_h"] = grid_h + data["grid_w"] = grid_w + + # The text encoder ran at full-res, so batch.prompt_embeds[0] has shape + # [B, max_text_tokens + full_res_image_tokens, d_llm]. The DiT forward + # expects llm_features and x=pos_z to share the same sequence length, so + # we bilinearly resize the image portion of llm_features to the low-res grid. + # ProgressiveDenoisingStage.forward() sets batch.height = init_h_pixel = + # (H_lat // downsample) * latent_scale before calling us, so: + # orig_grid_{h,w} = grid_{h,w} * downsample (downsample = 2^levels) + full_res_llm = batch.prompt_embeds[0] + batch.prompt_embeds[0] = _adapt_llm_features( + full_res_llm, max_text_tokens, num_image_tokens + ) + + ctx = Ideogram4DenoisingStage._prepare_denoising_loop(self, batch, server_args) + # Persist original LLM features so _on_resolution_change can restore + # them (with re-adapted image-token count) when the latent is upsampled. + ctx.extra["ideogram4_full_res_llm_features"] = full_res_llm + + # _prepare_denoising_loop (parent) reads batch.height = init_h_pixel (low-res) + # and computes the schedule at that resolution. The reference always uses the + # TARGET full-resolution schedule throughout denoising, even during coarse steps. + # Recompute schedule_values and schedule_deltas at the original full-res. + levels = int(getattr(batch, "progressive_levels", 1)) + orig_h = batch.height * (2**levels) + orig_w = batch.width * (2**levels) + preset = getattr(batch, "preset", "V4_DEFAULT_20") + preset_cfg = IDEOGRAM4_PRESETS[preset] + full_res_schedule = get_schedule_for_resolution( + (orig_h, orig_w), + known_mean=float(preset_cfg["mu"]), + std=float(preset_cfg["std"]), + ) + device = ctx.extra["ideogram4_schedule_values"].device + step_intervals = make_step_intervals(int(preset_cfg["num_steps"])).to(device) + full_res_schedule_values = full_res_schedule(step_intervals) + ctx.extra["ideogram4_schedule_values"] = full_res_schedule_values + ctx.extra["ideogram4_schedule_deltas"] = ( + full_res_schedule_values[:-1] - full_res_schedule_values[1:] + ) + + # Expose a sigma_NOISE tensor for stage-transition logic (find_transition_steps) + # and DWT upsample (sigma_t = sigmas[stage_end]). + # + # schedule_values convention: sigma_clean, index 0 = clean end (≈1), + # index N = noisy end (≈0). Step step_index k uses internal index i = N-1-k, + # so sigma_NOISE at step k = 1 - schedule_values[N-k]. + # flip(1 - schedule_values) gives sigmas[k] = 1 - schedule_values[N-k], + # which decreases from ≈1 (noisy, step 0) to ≈0 (clean, step N). + ctx.scheduler.sigmas = torch.flip(1.0 - full_res_schedule_values, [0]) + return ctx + + # ------------------------------------------------------------------ + # Denoising step override + # ------------------------------------------------------------------ + + def _run_denoising_step( + self, + ctx: DenoisingContext, + step, + batch: Req, + server_args: ServerArgs, + ) -> None: + # Ideogram4DenoisingStage uses step.t_int as a step index [0..N-1] into + # schedule_values. set_timesteps(N) produces timesteps = [N-1, N-2, ..., 0], + # so the correct mapping is t_int = N-1-step_index. + # + # In dct_rewind mode the progressive base patches + # ctx.timesteps[transition_step] = t_eff * 1000 (FLUX-convention noise level), + # which corrupts int(timesteps[step_index]) to ~950 and causes an IndexError. + # We bypass ctx.timesteps entirely and reconstruct from step_index. + num_steps = len(ctx.timesteps) + step.t_int = num_steps - 1 - step.step_index + Ideogram4DenoisingStage._run_denoising_step(self, ctx, step, batch, server_args) + + # ------------------------------------------------------------------ + # Cache-DiT refresh override + # ------------------------------------------------------------------ + + def _refresh_cache_dit_context( + 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 + ) + refresh_context_on_transformer( + self.unconditional_transformer, n_remaining, scm_preset=scm_preset + ) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/progressive_resolution/qwen_image.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/progressive_resolution/qwen_image.py index 5c6c2cb07..0304e6891 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/progressive_resolution/qwen_image.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/progressive_resolution/qwen_image.py @@ -37,6 +37,10 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) +# Module-level aliases for test discoverability and external reuse. +_qwen_image_unpack = unpack_2x2_latent +_qwen_image_pack = pack_2x2_latent + # Power-law spectrum constants P(ω) = A·|ω|^{-β} for Qwen-Image VAE latents. # TODO: fit these from Qwen-Image VAE latent statistics on a representative # dataset (e.g. Aesthetics-Train-V2). Using FLUX.1-dev fitted values as diff --git a/python/sglang/multimodal_gen/test/unit/progressive_resolution/test_progressive.py b/python/sglang/multimodal_gen/test/unit/progressive_resolution/test_progressive.py index d0aac40c8..ade149a8e 100644 --- a/python/sglang/multimodal_gen/test/unit/progressive_resolution/test_progressive.py +++ b/python/sglang/multimodal_gen/test/unit/progressive_resolution/test_progressive.py @@ -26,6 +26,11 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.progressive_resolution. _flux2_pack, _flux2_unpack, ) +from sglang.multimodal_gen.runtime.pipelines_core.stages.progressive_resolution.ideogram import ( + Ideogram4ProgressiveDenoisingStage, + _ideogram4_pack, + _ideogram4_unpack, +) from sglang.multimodal_gen.runtime.pipelines_core.stages.progressive_resolution.qwen_image import ( _qwen_image_pack, _qwen_image_unpack, @@ -293,6 +298,7 @@ class TestProgressiveStageHelpers(unittest.TestCase): def test_model_specific_latent_scale_factors(self): flux2_stage = object.__new__(Flux2ProgressiveDenoisingStage) wan_stage = object.__new__(WanProgressiveDenoisingStage) + ideogram_stage = object.__new__(Ideogram4ProgressiveDenoisingStage) flux2_args = SimpleNamespace( pipeline_config=SimpleNamespace( @@ -308,9 +314,13 @@ class TestProgressiveStageHelpers(unittest.TestCase): ) ) ) + ideogram_args = SimpleNamespace( + pipeline_config=SimpleNamespace(patch_size=2, ae_scale_factor=8) + ) self.assertEqual(flux2_stage._latent_scale_factor(flux2_args), 16) self.assertEqual(wan_stage._latent_scale_factor(wan_args), 8) + self.assertEqual(ideogram_stage._latent_scale_factor(ideogram_args), 16) class TestLatentAdapters(unittest.TestCase): @@ -354,5 +364,265 @@ class TestLatentAdapters(unittest.TestCase): ) +class TestIdeogram4LatentAdapters(unittest.TestCase): + def test_row_major_roundtrip(self): + # [B, C, H, W] → pack → unpack → original + x = torch.arange(2 * 128 * 4 * 6, dtype=torch.float32).reshape(2, 128, 4, 6) + + packed = _ideogram4_pack(x) + + self.assertEqual(packed.shape, (2, 4 * 6, 128)) + torch.testing.assert_close(_ideogram4_unpack(packed, 4, 6), x) + + def test_pack_preserves_row_major_order(self): + # Each packed token at position [b, row*W + col] should equal x[b, :, row, col] + x = torch.arange(1 * 4 * 3 * 5, dtype=torch.float32).reshape(1, 4, 3, 5) + packed = _ideogram4_pack(x) + + # token at spatial position (row=1, col=2) → flat index 1*5+2 = 7 + torch.testing.assert_close(packed[0, 7], x[0, :, 1, 2]) + + def test_ideogram_matches_flux2_pack_unpack(self): + # Ideogram and FLUX.2 share the same row-major token layout + from sglang.multimodal_gen.runtime.pipelines_core.stages.progressive_resolution.flux_2 import ( + _flux2_pack, + _flux2_unpack, + ) + + x = torch.randn(2, 128, 5, 7) + torch.testing.assert_close(_ideogram4_pack(x), _flux2_pack(x)) + packed = _ideogram4_pack(x) + torch.testing.assert_close( + _ideogram4_unpack(packed, 5, 7), _flux2_unpack(packed, 5, 7) + ) + + +class TestIdeogram4OnResolutionChange(unittest.TestCase): + """CPU-only test: _on_resolution_change rebuilds batch/ctx tensors correctly.""" + + _PATCH = 2 + _AE = 8 + _SCALE = 16 # patch * ae = 16 + _IN_C = 128 + _LLM_DIM = 64 # small stand-in for llm_features_dim + + def _make_ideogram_extra(self, batch_size, grid_h, grid_w, max_text_tokens): + """Build a minimal batch.extra["ideogram4"] dict matching _prepare_denoising_loop.""" + from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ideogram import ( + IMAGE_POSITION_OFFSET, + LLM_TOKEN_INDICATOR, + OUTPUT_IMAGE_INDICATOR, + SEQUENCE_PADDING_INDICATOR, + ) + + num_image_tokens = grid_h * grid_w + total_seq_len = max_text_tokens + num_image_tokens + + h_idx = torch.arange(grid_h).view(-1, 1).expand(grid_h, grid_w).reshape(-1) + w_idx = torch.arange(grid_w).view(1, -1).expand(grid_h, grid_w).reshape(-1) + image_pos = ( + torch.stack([torch.zeros_like(h_idx), h_idx, w_idx], dim=1) + + IMAGE_POSITION_OFFSET + ) + + position_ids = torch.zeros(batch_size, total_seq_len, 3, dtype=torch.long) + segment_ids = torch.full( + (batch_size, total_seq_len), SEQUENCE_PADDING_INDICATOR, dtype=torch.long + ) + indicator = torch.zeros(batch_size, total_seq_len, dtype=torch.long) + + for b in range(batch_size): + # simulate a single item with no text padding for simplicity + position_ids[b, :max_text_tokens] = ( + torch.arange(max_text_tokens).unsqueeze(-1).expand(-1, 3) + ) + position_ids[b, max_text_tokens:] = image_pos + segment_ids[b] = 1 + indicator[b, :max_text_tokens] = LLM_TOKEN_INDICATOR + indicator[b, max_text_tokens:] = OUTPUT_IMAGE_INDICATOR + + return { + "position_ids": position_ids, + "segment_ids": segment_ids, + "indicator": indicator, + "num_image_tokens": num_image_tokens, + "grid_h": grid_h, + "grid_w": grid_w, + "max_text_tokens": max_text_tokens, + } + + def _make_ctx_extra(self, batch_size, num_image_tokens, max_text_tokens): + from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ideogram import ( + OUTPUT_IMAGE_INDICATOR, + ) + + neg_llm_features = torch.zeros(batch_size, num_image_tokens, self._LLM_DIM) + attn_mask = torch.ones( + batch_size, max_text_tokens + num_image_tokens, dtype=torch.bool + ) + neg_attn_mask = torch.ones(batch_size, num_image_tokens, dtype=torch.bool) + return { + "ideogram4_attn_mask": attn_mask, + "ideogram4_attn_mask_meta": None, + "ideogram4_neg_position_ids": torch.zeros( + batch_size, num_image_tokens, 3, dtype=torch.long + ), + "ideogram4_neg_segment_ids": torch.ones( + batch_size, num_image_tokens, dtype=torch.long + ), + "ideogram4_neg_indicator": torch.full( + (batch_size, num_image_tokens), OUTPUT_IMAGE_INDICATOR, dtype=torch.long + ), + "ideogram4_neg_attn_mask": neg_attn_mask, + "ideogram4_neg_attn_mask_meta": None, + "ideogram4_neg_llm_features": neg_llm_features, + } + + def test_resolution_change_doubles_image_tokens(self): + B = 2 + old_grid_h, old_grid_w = 4, 4 + new_grid_h, new_grid_w = 8, 8 + max_text_tokens = 10 + + stage = object.__new__(Ideogram4ProgressiveDenoisingStage) + server_args = SimpleNamespace( + pipeline_config=SimpleNamespace( + patch_size=self._PATCH, ae_scale_factor=self._AE + ) + ) + + old_num_img = old_grid_h * old_grid_w # 16 + new_num_img = new_grid_h * new_grid_w # 64 + new_h_pixel = new_grid_h * self._SCALE + new_w_pixel = new_grid_w * self._SCALE + + # Build fake ctx and batch + ctx = SimpleNamespace( + latents=torch.zeros(B, new_num_img, self._IN_C), + extra=self._make_ctx_extra(B, old_num_img, max_text_tokens), + ) + batch = SimpleNamespace( + extra={ + "ideogram4": self._make_ideogram_extra( + B, old_grid_h, old_grid_w, max_text_tokens + ) + } + ) + + stage._on_resolution_change(ctx, batch, server_args, new_h_pixel, new_w_pixel) + + data = batch.extra["ideogram4"] + self.assertEqual(data["num_image_tokens"], new_num_img) + self.assertEqual(data["grid_h"], new_grid_h) + self.assertEqual(data["grid_w"], new_grid_w) + self.assertEqual( + data["position_ids"].shape, (B, max_text_tokens + new_num_img, 3) + ) + self.assertEqual(data["segment_ids"].shape, (B, max_text_tokens + new_num_img)) + self.assertEqual(data["indicator"].shape, (B, max_text_tokens + new_num_img)) + + # ctx.extra tensors updated to new sizes + self.assertEqual( + ctx.extra["ideogram4_attn_mask"].shape, + (B, max_text_tokens + new_num_img), + ) + self.assertEqual( + ctx.extra["ideogram4_neg_position_ids"].shape, (B, new_num_img, 3) + ) + self.assertEqual( + ctx.extra["ideogram4_neg_llm_features"].shape, + (B, new_num_img, self._LLM_DIM), + ) + + def test_text_portion_is_unchanged_after_resolution_change(self): + B = 1 + old_grid_h, old_grid_w = 4, 4 + new_grid_h, new_grid_w = 8, 8 + max_text_tokens = 6 + + stage = object.__new__(Ideogram4ProgressiveDenoisingStage) + server_args = SimpleNamespace( + pipeline_config=SimpleNamespace( + patch_size=self._PATCH, ae_scale_factor=self._AE + ) + ) + + old_data = self._make_ideogram_extra(B, old_grid_h, old_grid_w, max_text_tokens) + old_text_position_ids = old_data["position_ids"][:, :max_text_tokens].clone() + old_text_segment_ids = old_data["segment_ids"][:, :max_text_tokens].clone() + old_text_indicator = old_data["indicator"][:, :max_text_tokens].clone() + + ctx = SimpleNamespace( + latents=torch.zeros(B, new_grid_h * new_grid_w, self._IN_C), + extra=self._make_ctx_extra(B, old_grid_h * old_grid_w, max_text_tokens), + ) + batch = SimpleNamespace(extra={"ideogram4": old_data}) + + stage._on_resolution_change( + ctx, + batch, + server_args, + new_grid_h * self._SCALE, + new_grid_w * self._SCALE, + ) + + data = batch.extra["ideogram4"] + torch.testing.assert_close( + data["position_ids"][:, :max_text_tokens], old_text_position_ids + ) + torch.testing.assert_close( + data["segment_ids"][:, :max_text_tokens], old_text_segment_ids + ) + torch.testing.assert_close( + data["indicator"][:, :max_text_tokens], old_text_indicator + ) + + def test_image_position_ids_use_grid_coordinates(self): + B = 1 + grid_h, grid_w = 4, 6 + max_text_tokens = 4 + scale = self._SCALE + + stage = object.__new__(Ideogram4ProgressiveDenoisingStage) + server_args = SimpleNamespace( + pipeline_config=SimpleNamespace( + patch_size=self._PATCH, ae_scale_factor=self._AE + ) + ) + + old_num_img = 2 * 3 # half the new grid + ctx = SimpleNamespace( + latents=torch.zeros(B, grid_h * grid_w, self._IN_C), + extra=self._make_ctx_extra(B, old_num_img, max_text_tokens), + ) + batch = SimpleNamespace( + extra={ + "ideogram4": self._make_ideogram_extra( + B, grid_h // 2, grid_w // 2, max_text_tokens + ) + } + ) + + stage._on_resolution_change( + ctx, batch, server_args, grid_h * scale, grid_w * scale + ) + + img_pos = batch.extra["ideogram4"]["position_ids"][0, max_text_tokens:] + self.assertEqual(img_pos.shape, (grid_h * grid_w, 3)) + # All t-coordinates (dim 0) should be IMAGE_POSITION_OFFSET (the t=0 term) + from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ideogram import ( + IMAGE_POSITION_OFFSET, + ) + + # t dimension (index 0) should be IMAGE_POSITION_OFFSET + 0 + self.assertTrue((img_pos[:, 0] == IMAGE_POSITION_OFFSET).all()) + # h dimension at row-major index row*grid_w + col should be IMAGE_POSITION_OFFSET + row + for row in range(grid_h): + for col in range(grid_w): + idx = row * grid_w + col + self.assertEqual(img_pos[idx, 1].item(), IMAGE_POSITION_OFFSET + row) + self.assertEqual(img_pos[idx, 2].item(), IMAGE_POSITION_OFFSET + col) + + if __name__ == "__main__": unittest.main()