[diffusion] feat: progressive resolution growing for Ideogram 4 via GPU DCT upsampling with up to 1.56× speedup (#27736)
This commit is contained in:
@@ -4,7 +4,7 @@ description: "Experimental spectral progressive resolution growing for selected
|
|||||||
tag: "approx"
|
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).
|
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× |
|
| FLUX.2 1024×1024 | 4,096 | 1,024 | 4.0× |
|
||||||
| Z-Image 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× |
|
| 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
|
## 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.10 | 16@64² + 14@128² | 33.86 s | **1.27×** |
|
||||||
| dct_rewind L1 δ=0.20 | 19@64² + 11@128² | 25.40 s | **1.69×** |
|
| 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
|
## Limitations
|
||||||
|
|
||||||
- **Sequence parallelism incompatible.** Cannot be combined with `--ulysses-degree` or `--ring-degree`. The stage raises a `RuntimeError` if SP is enabled.
|
- **Sequence parallelism incompatible.** Cannot be combined with `--ulysses-degree` or `--ring-degree`. The stage raises a `RuntimeError` if SP is enabled.
|
||||||
|
|||||||
@@ -18,6 +18,12 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.i
|
|||||||
Ideogram4DenoisingStage,
|
Ideogram4DenoisingStage,
|
||||||
Ideogram4TextEncodingStage,
|
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.server_args import ServerArgs
|
||||||
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
|
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
|
||||||
maybe_download_model,
|
maybe_download_model,
|
||||||
@@ -115,6 +121,22 @@ class Ideogram4Pipeline(LoRAPipeline, ComposedPipelineBase):
|
|||||||
"scheduler",
|
"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):
|
def create_pipeline_stages(self, server_args: ServerArgs):
|
||||||
self.add_stage(InputValidationStage())
|
self.add_stage(InputValidationStage())
|
||||||
self.add_stage_factory(
|
self.add_stage_factory(
|
||||||
@@ -128,11 +150,7 @@ class Ideogram4Pipeline(LoRAPipeline, ComposedPipelineBase):
|
|||||||
self.add_standard_latent_preparation_stage()
|
self.add_standard_latent_preparation_stage()
|
||||||
self.add_stage_factory(
|
self.add_stage_factory(
|
||||||
RoleType.DENOISER,
|
RoleType.DENOISER,
|
||||||
lambda: Ideogram4DenoisingStage(
|
self._create_denoising_stage,
|
||||||
transformer=self.get_module("transformer"),
|
|
||||||
unconditional_transformer=self.get_module("unconditional_transformer"),
|
|
||||||
pipeline=self,
|
|
||||||
),
|
|
||||||
"ideogram4_denoising_stage",
|
"ideogram4_denoising_stage",
|
||||||
)
|
)
|
||||||
self.add_stage_factory(
|
self.add_stage_factory(
|
||||||
|
|||||||
+41
-19
@@ -258,6 +258,18 @@ class ProgressiveDenoisingStage(DenoisingStage):
|
|||||||
"""
|
"""
|
||||||
return server_args.pipeline_config.vae_config.arch_config.vae_scale_factor
|
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(
|
def _unpack_latent(
|
||||||
self, latent: torch.Tensor, h_lat: int, w_lat: int
|
self, latent: torch.Tensor, h_lat: int, w_lat: int
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
@@ -286,6 +298,29 @@ class ProgressiveDenoisingStage(DenoisingStage):
|
|||||||
"""Called after each stage transition. Update resolution-dependent state."""
|
"""Called after each stage transition. Update resolution-dependent state."""
|
||||||
pass
|
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
|
# Helpers
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
@@ -441,9 +476,12 @@ class ProgressiveDenoisingStage(DenoisingStage):
|
|||||||
init_h_lat = H_lat // downsample
|
init_h_lat = H_lat // downsample
|
||||||
init_w_lat = W_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(
|
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)
|
num_stages = len(stage_sigmas)
|
||||||
|
|
||||||
@@ -572,23 +610,7 @@ class ProgressiveDenoisingStage(DenoisingStage):
|
|||||||
# residual-diff decision for the first full-res steps.
|
# residual-diff decision for the first full-res steps.
|
||||||
if self._cache_dit_enabled:
|
if self._cache_dit_enabled:
|
||||||
n_remaining = n_steps - stage_end
|
n_remaining = n_steps - stage_end
|
||||||
scm_preset = _get_scm_preset()
|
self._refresh_cache_dit_context(n_remaining, _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,
|
|
||||||
)
|
|
||||||
logger.info(
|
logger.info(
|
||||||
"cache-dit context refreshed at stage transition "
|
"cache-dit context refreshed at stage transition "
|
||||||
"(step %d, %d steps remaining)",
|
"(step %d, %d steps remaining)",
|
||||||
|
|||||||
+4
@@ -30,6 +30,10 @@ logger = init_logger(__name__)
|
|||||||
FLUX_SPECTRUM_A: float = 203.615097
|
FLUX_SPECTRUM_A: float = 203.615097
|
||||||
FLUX_SPECTRUM_BETA: float = 1.915461
|
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):
|
class FluxProgressiveDenoisingStage(ProgressiveDenoisingStage):
|
||||||
"""FLUX-specific progressive denoising stage.
|
"""FLUX-specific progressive denoising stage.
|
||||||
|
|||||||
+462
@@ -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
|
||||||
|
)
|
||||||
+4
@@ -37,6 +37,10 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
|||||||
|
|
||||||
logger = init_logger(__name__)
|
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.
|
# Power-law spectrum constants P(ω) = A·|ω|^{-β} for Qwen-Image VAE latents.
|
||||||
# TODO: fit these from Qwen-Image VAE latent statistics on a representative
|
# 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
|
# dataset (e.g. Aesthetics-Train-V2). Using FLUX.1-dev fitted values as
|
||||||
|
|||||||
@@ -26,6 +26,11 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.progressive_resolution.
|
|||||||
_flux2_pack,
|
_flux2_pack,
|
||||||
_flux2_unpack,
|
_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 (
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.progressive_resolution.qwen_image import (
|
||||||
_qwen_image_pack,
|
_qwen_image_pack,
|
||||||
_qwen_image_unpack,
|
_qwen_image_unpack,
|
||||||
@@ -293,6 +298,7 @@ class TestProgressiveStageHelpers(unittest.TestCase):
|
|||||||
def test_model_specific_latent_scale_factors(self):
|
def test_model_specific_latent_scale_factors(self):
|
||||||
flux2_stage = object.__new__(Flux2ProgressiveDenoisingStage)
|
flux2_stage = object.__new__(Flux2ProgressiveDenoisingStage)
|
||||||
wan_stage = object.__new__(WanProgressiveDenoisingStage)
|
wan_stage = object.__new__(WanProgressiveDenoisingStage)
|
||||||
|
ideogram_stage = object.__new__(Ideogram4ProgressiveDenoisingStage)
|
||||||
|
|
||||||
flux2_args = SimpleNamespace(
|
flux2_args = SimpleNamespace(
|
||||||
pipeline_config=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(flux2_stage._latent_scale_factor(flux2_args), 16)
|
||||||
self.assertEqual(wan_stage._latent_scale_factor(wan_args), 8)
|
self.assertEqual(wan_stage._latent_scale_factor(wan_args), 8)
|
||||||
|
self.assertEqual(ideogram_stage._latent_scale_factor(ideogram_args), 16)
|
||||||
|
|
||||||
|
|
||||||
class TestLatentAdapters(unittest.TestCase):
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user