[diffusion] model: support Ideogram 4 FP8 (#27279)

Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
xly
2026-06-05 23:54:48 +08:00
committed by GitHub
co-authored by Mick
parent e7f94d0d40
commit d01cf27b7d
34 changed files with 2663 additions and 182 deletions
@@ -4,6 +4,7 @@ from sglang.multimodal_gen.configs.models.dits.cosmos3video import Cosmos3VideoC
from sglang.multimodal_gen.configs.models.dits.helios import HeliosConfig from sglang.multimodal_gen.configs.models.dits.helios import HeliosConfig
from sglang.multimodal_gen.configs.models.dits.hunyuan3d import Hunyuan3DDiTConfig from sglang.multimodal_gen.configs.models.dits.hunyuan3d import Hunyuan3DDiTConfig
from sglang.multimodal_gen.configs.models.dits.hunyuanvideo import HunyuanVideoConfig from sglang.multimodal_gen.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
from sglang.multimodal_gen.configs.models.dits.ideogram import Ideogram4DiTConfig
from sglang.multimodal_gen.configs.models.dits.lingbot_world import ( from sglang.multimodal_gen.configs.models.dits.lingbot_world import (
LingBotWorldVideoConfig, LingBotWorldVideoConfig,
) )
@@ -18,6 +19,7 @@ __all__ = [
"Cosmos3VideoConfig", "Cosmos3VideoConfig",
"HeliosConfig", "HeliosConfig",
"HunyuanVideoConfig", "HunyuanVideoConfig",
"Ideogram4DiTConfig",
"LingBotWorldVideoConfig", "LingBotWorldVideoConfig",
"WanVideoConfig", "WanVideoConfig",
"Hunyuan3DDiTConfig", "Hunyuan3DDiTConfig",
@@ -0,0 +1,39 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig
from sglang.multimodal_gen.configs.models.fsdp import is_layer
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
@dataclass
class Ideogram4DiTArchConfig(DiTArchConfig):
adaln_dim: int = 512
attention_head_dim: int = 256
in_channels: int = 128
intermediate_size: int = 12288
llm_features_dim: int = 53248
mrope_section: tuple[int, int, int] | list[int] = (24, 20, 20)
norm_eps: float = 1e-5
num_attention_heads: int = 18
num_layers: int = 34
rope_theta: int = 5_000_000
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_layer])
_supported_attention_backends: set[AttentionBackendEnum] = field(
default_factory=lambda: {
AttentionBackendEnum.FA,
AttentionBackendEnum.TORCH_SDPA,
}
)
def __post_init__(self) -> None:
super().__post_init__()
self.hidden_size = self.num_attention_heads * self.attention_head_dim
self.num_channels_latents = self.in_channels
@dataclass
class Ideogram4DiTConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=Ideogram4DiTArchConfig)
prefix: str = "ideogram4"
@@ -17,6 +17,9 @@ from sglang.multimodal_gen.configs.models.encoders.flux_2 import (
) )
from sglang.multimodal_gen.configs.models.encoders.gemma2 import Gemma2Config from sglang.multimodal_gen.configs.models.encoders.gemma2 import Gemma2Config
from sglang.multimodal_gen.configs.models.encoders.gemma_3 import Gemma3Config from sglang.multimodal_gen.configs.models.encoders.gemma_3 import Gemma3Config
from sglang.multimodal_gen.configs.models.encoders.ideogram import (
Ideogram4TextEncoderConfig,
)
from sglang.multimodal_gen.configs.models.encoders.llama import LlamaConfig from sglang.multimodal_gen.configs.models.encoders.llama import LlamaConfig
from sglang.multimodal_gen.configs.models.encoders.qwen3 import Qwen3TextConfig from sglang.multimodal_gen.configs.models.encoders.qwen3 import Qwen3TextConfig
from sglang.multimodal_gen.configs.models.encoders.qwen3vl import Qwen3VLConfig from sglang.multimodal_gen.configs.models.encoders.qwen3vl import Qwen3VLConfig
@@ -38,4 +41,5 @@ __all__ = [
"T5Config", "T5Config",
"Gemma2Config", "Gemma2Config",
"Gemma3Config", "Gemma3Config",
"Ideogram4TextEncoderConfig",
] ]
@@ -0,0 +1,26 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from sglang.multimodal_gen.configs.models.encoders.qwen3vl import (
Qwen3VLArchConfig,
Qwen3VLConfig,
)
@dataclass
class Ideogram4TextEncoderConfig(Qwen3VLConfig):
"""Use the local Ideogram text_encoder as a language-only Qwen3-VL encoder."""
def update_model_arch(self, source_model_dict):
super().update_model_arch(source_model_dict)
self.post_diffusers_config_update()
def post_diffusers_config_update(self):
self.arch_config.architectures = ["IdeogramQwen3VLTextEncoder"]
self.arch_config.ideogram_fp8_weight_only = True
def finalize_model_arch(self):
self.post_diffusers_config_update()
arch_config: Qwen3VLArchConfig = field(default_factory=Qwen3VLArchConfig)
@@ -28,6 +28,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.hunyuan import (
from sglang.multimodal_gen.configs.pipeline_configs.hunyuan3d import ( from sglang.multimodal_gen.configs.pipeline_configs.hunyuan3d import (
Hunyuan3D2PipelineConfig, Hunyuan3D2PipelineConfig,
) )
from sglang.multimodal_gen.configs.pipeline_configs.ideogram import (
Ideogram4PipelineConfig,
)
from sglang.multimodal_gen.configs.pipeline_configs.lingbot_world import ( from sglang.multimodal_gen.configs.pipeline_configs.lingbot_world import (
LingBotWorldCausalDMDConfig, LingBotWorldCausalDMDConfig,
) )
@@ -55,6 +58,7 @@ __all__ = [
"HunyuanConfig", "HunyuanConfig",
"FastHunyuanConfig", "FastHunyuanConfig",
"Hunyuan3D2PipelineConfig", "Hunyuan3D2PipelineConfig",
"Ideogram4PipelineConfig",
"FluxPipelineConfig", "FluxPipelineConfig",
"Flux2PipelineConfig", "Flux2PipelineConfig",
"Flux2KleinPipelineConfig", "Flux2KleinPipelineConfig",
@@ -0,0 +1,302 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from sglang.multimodal_gen.configs.models import DiTConfig, EncoderConfig, VAEConfig
from sglang.multimodal_gen.configs.models.dits.ideogram import Ideogram4DiTConfig
from sglang.multimodal_gen.configs.models.encoders.ideogram import (
Ideogram4TextEncoderConfig,
)
from sglang.multimodal_gen.configs.models.vaes.flux import Flux2VAEConfig
from sglang.multimodal_gen.configs.pipeline_configs.base import (
ImagePipelineConfig,
ModelTaskType,
)
LATENT_SHIFT = (
0.01984364,
0.10149707,
0.29689495,
0.27188619,
-0.21445648,
-0.15979549,
0.05021099,
-0.15083604,
-0.15360136,
-0.20131799,
0.01922352,
0.0622626,
0.10140969,
-0.06739428,
0.3758261,
-0.233712,
0.35164491,
-0.02590912,
-0.0271935,
-0.10833897,
-0.1476848,
-0.01130957,
-0.2298372,
0.23526423,
-0.10893522,
0.11957631,
0.04047799,
0.3134589,
-0.17225064,
-0.18646109,
-0.34691978,
-0.03571246,
0.02583857,
0.10190072,
0.28402294,
0.26952152,
-0.21634675,
-0.17938656,
0.04358909,
-0.15007621,
-0.1548502,
-0.18971131,
0.02710861,
0.05609494,
0.10697846,
-0.06854968,
0.38167698,
-0.24269937,
0.35705471,
-0.03063305,
-0.02946109,
-0.11244286,
-0.14336038,
-0.01362137,
-0.21863696,
0.23228983,
-0.11739769,
0.11693044,
0.02563311,
0.31356594,
-0.17420591,
-0.19006285,
-0.34905377,
-0.04025005,
0.01924137,
0.07652984,
0.2995608,
0.2628057,
-0.22011674,
-0.12715361,
0.04879879,
-0.14075719,
-0.15935895,
-0.2123584,
0.01974813,
0.05523547,
0.10011992,
-0.06428964,
0.37781868,
-0.21491644,
0.34254215,
-0.03153528,
-0.0310082,
-0.10761415,
-0.14730405,
-0.02475182,
-0.2285588,
0.2515081,
-0.10445128,
0.12446,
0.07062869,
0.30880162,
-0.18016875,
-0.18869164,
-0.34533499,
-0.0129177,
0.02578168,
0.07993659,
0.28642181,
0.26038408,
-0.22459419,
-0.14820155,
0.04059549,
-0.14043529,
-0.16111187,
-0.2020305,
0.02602069,
0.04852717,
0.10432153,
-0.06309942,
0.38402443,
-0.22397003,
0.34814481,
-0.03774432,
-0.03381438,
-0.11245691,
-0.14128767,
-0.02853208,
-0.21752016,
0.24872463,
-0.11399775,
0.1222687,
0.05620835,
0.309178,
-0.18065738,
-0.19401479,
-0.34495114,
-0.01760592,
)
LATENT_SCALE = (
1.63933691,
1.70204478,
1.73642566,
1.90004803,
1.6675316,
1.69059584,
1.56853198,
1.62314944,
1.89106626,
1.58086668,
1.60822129,
1.60962993,
1.63322129,
1.56074359,
1.73419528,
1.7919265,
1.64040632,
1.66802808,
1.60390303,
1.75480492,
1.63187587,
1.64334594,
1.61722884,
1.60146046,
1.63459219,
1.55291476,
1.68771497,
1.68415657,
1.78966054,
1.66631641,
1.65626686,
1.65976433,
1.63487607,
1.69513249,
1.72933756,
1.91310663,
1.67035057,
1.72286863,
1.56719251,
1.61934825,
1.88628859,
1.56911539,
1.59455129,
1.60829869,
1.62470611,
1.56052853,
1.73677003,
1.77563606,
1.63732541,
1.66370527,
1.59508952,
1.75153949,
1.63029275,
1.64517667,
1.61659342,
1.59722044,
1.64103121,
1.5408531,
1.68610394,
1.67772755,
1.78998563,
1.66621713,
1.65458955,
1.66041308,
1.64710857,
1.68163503,
1.74000294,
1.92784786,
1.67411194,
1.67395548,
1.57406532,
1.62199356,
1.87618195,
1.5584375,
1.57438785,
1.61711053,
1.63094305,
1.55644029,
1.73124302,
1.80666627,
1.6463621,
1.65932006,
1.60816188,
1.75682671,
1.64695873,
1.63121722,
1.61380832,
1.60478651,
1.63396035,
1.53505068,
1.65534289,
1.67132281,
1.80317197,
1.6767314,
1.65700938,
1.68426259,
1.65339716,
1.67540638,
1.73298504,
1.94067348,
1.67893609,
1.70635117,
1.5730906,
1.61928553,
1.87148809,
1.56244866,
1.56697152,
1.61584394,
1.62759496,
1.55480378,
1.73484107,
1.79055143,
1.64688773,
1.66121492,
1.60135887,
1.75254572,
1.64798332,
1.62989921,
1.61381592,
1.60792883,
1.63939668,
1.53075757,
1.65371318,
1.66801185,
1.80029087,
1.67591476,
1.65655173,
1.68533454,
)
@dataclass
class Ideogram4PipelineConfig(ImagePipelineConfig):
task_type: ModelTaskType = ModelTaskType.T2I
should_use_guidance: bool = False
vae_precision: str = "bf16"
dit_precision: str = "bf16"
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16",))
dit_config: DiTConfig = field(default_factory=Ideogram4DiTConfig)
vae_config: VAEConfig = field(default_factory=Flux2VAEConfig)
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (Ideogram4TextEncoderConfig(),)
)
text_encoder_extra_args: list[dict] = field(default_factory=lambda: [{}])
preprocess_text_funcs: tuple = field(default_factory=lambda: (None,))
postprocess_text_funcs: tuple = field(default_factory=lambda: (None,))
patch_size: int = 2
ae_scale_factor: int = 8
max_text_tokens: int = 2048
def prepare_latent_shape(self, batch, batch_size, num_frames):
patch = self.patch_size * self.ae_scale_factor
grid_h = batch.height // patch
grid_w = batch.width // patch
return (batch_size, grid_h * grid_w, self.dit_config.arch_config.in_channels)
@@ -3,6 +3,11 @@
from sglang.multimodal_gen.configs.sample.diffusers_generic import ( from sglang.multimodal_gen.configs.sample.diffusers_generic import (
DiffusersGenericSamplingParams, DiffusersGenericSamplingParams,
) )
from sglang.multimodal_gen.configs.sample.ideogram import Ideogram4SamplingParams
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
__all__ = ["SamplingParams", "DiffusersGenericSamplingParams"] __all__ = [
"SamplingParams",
"DiffusersGenericSamplingParams",
"Ideogram4SamplingParams",
]
@@ -0,0 +1,78 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from sglang.multimodal_gen.configs.sample.sampling_params import (
DataType,
SamplingParams,
)
IDEOGRAM4_PRESETS: dict[str, dict[str, object]] = {
"V4_QUALITY_48": {
"num_steps": 48,
"guidance_schedule": (3.0,) * 3 + (7.0,) * 45,
"mu": 0.0,
"std": 1.5,
},
"V4_DEFAULT_20": {
"num_steps": 20,
"guidance_schedule": (3.0,) * 2 + (7.0,) * 18,
"mu": 0.0,
"std": 1.75,
},
"V4_TURBO_12": {
"num_steps": 12,
"guidance_schedule": (3.0,) * 1 + (7.0,) * 11,
"mu": 0.5,
"std": 1.75,
},
}
@dataclass
class Ideogram4SamplingParams(SamplingParams):
data_type: DataType = DataType.IMAGE
prompt: str = " "
negative_prompt: str = " "
height: int = 1024
width: int = 1024
num_frames: int = 1
num_inference_steps: int | None = None
guidance_scale: float | None = None
preset: str = "V4_DEFAULT_20"
def __post_init__(self) -> None:
if self.preset not in IDEOGRAM4_PRESETS:
raise ValueError(
f"Unknown Ideogram 4 preset {self.preset!r}; "
f"expected one of {sorted(IDEOGRAM4_PRESETS)}"
)
preset_cfg = IDEOGRAM4_PRESETS[self.preset]
preset_steps = int(preset_cfg["num_steps"])
explicit_fields = getattr(self, "_explicit_fields", None)
num_steps_is_explicit = (
explicit_fields is None or "num_inference_steps" in explicit_fields
)
guidance_is_explicit = (
explicit_fields is None or "guidance_scale" in explicit_fields
)
if (
self.num_inference_steps is not None
and self.num_inference_steps != preset_steps
and num_steps_is_explicit
):
raise ValueError(
"Ideogram 4 derives num_inference_steps from preset "
f"{self.preset!r}; got {self.num_inference_steps}, expected "
f"{preset_steps}."
)
if self.guidance_scale is not None and guidance_is_explicit:
preset_guidance = float(preset_cfg["guidance_schedule"][-1])
if self.guidance_scale != preset_guidance:
raise ValueError(
"Ideogram 4 derives guidance from the preset guidance_schedule; "
"guidance_scale cannot be set directly."
)
self.num_inference_steps = preset_steps
self.guidance_scale = float(preset_cfg["guidance_schedule"][-1])
super().__post_init__()
@@ -356,7 +356,7 @@ class SamplingParams:
or self.seed < 0 or self.seed < 0
): ):
raise ValueError( raise ValueError(
"seed must be a non-negative int or list of ints, " f"got {self.seed!r}" f"seed must be a non-negative int or list of ints, got {self.seed!r}"
) )
# Used by seconds() and video writer; fps <= 0 is always invalid. # Used by seconds() and video writer; fps <= 0 is always invalid.
@@ -536,13 +536,13 @@ class SamplingParams:
if self.enable_sequence_shard: if self.enable_sequence_shard:
self.adjust_frames = False self.adjust_frames = False
logger.info( logger.info(
f"Sequence dimension shard is enabled, disabling frame adjustment for better performance" "Sequence dimension shard is enabled, disabling frame adjustment for better performance"
) )
if pipeline_config.task_type.is_image_gen(): if pipeline_config.task_type.is_image_gen():
# settle num_frames # settle num_frames
if not server_args.pipeline_config.allow_set_num_frames(): if not server_args.pipeline_config.allow_set_num_frames():
logger.debug(f"Setting `num_frames` to 1 for image generation model") logger.debug("Setting `num_frames` to 1 for image generation model")
self.num_frames = 1 self.num_frames = 1
else: else:
@@ -1045,26 +1045,29 @@ class SamplingParams:
# global switch: if True, allow overriding protected fields # global switch: if True, allow overriding protected fields
allow_override_protected = not user_params.no_override_protected_fields allow_override_protected = not user_params.no_override_protected_fields
for field in dataclasses.fields(user_params): for field_info in dataclasses.fields(user_params):
field_name = field.name field_name = field_info.name
user_value = getattr(user_params, field_name) user_value = getattr(user_params, field_name)
if hasattr(SamplingParams, field_name): if hasattr(SamplingParams, field_name):
default_class_value = getattr(SamplingParams, field_name) default_class_value = getattr(SamplingParams, field_name)
elif field.default is not dataclasses.MISSING: elif field_info.default is not dataclasses.MISSING:
default_class_value = field.default default_class_value = field_info.default
elif field.default_factory is not dataclasses.MISSING: elif field_info.default_factory is not dataclasses.MISSING:
default_class_value = field.default_factory() default_class_value = field_info.default_factory()
else: else:
default_class_value = dataclasses.MISSING default_class_value = dataclasses.MISSING
is_user_modified = user_value != default_class_value or ( if explicit_fields is not None:
explicit_fields is not None and field_name in explicit_fields is_user_modified = field_name in explicit_fields
) else:
is_user_modified = user_value != default_class_value
is_protected_field = field_name in predefined_fields is_protected_field = field_name in predefined_fields
if is_user_modified and ( if is_user_modified and (
allow_override_protected or not is_protected_field allow_override_protected or not is_protected_field
): ):
setattr(self, field_name, user_value) setattr(self, field_name, user_value)
if explicit_fields is not None:
self._explicit_fields = set(explicit_fields)
self.__post_init__() self.__post_init__()
@property @property
+46 -21
View File
@@ -58,6 +58,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.glm_image import (
from sglang.multimodal_gen.configs.pipeline_configs.hunyuan3d import ( from sglang.multimodal_gen.configs.pipeline_configs.hunyuan3d import (
Hunyuan3D2PipelineConfig, Hunyuan3D2PipelineConfig,
) )
from sglang.multimodal_gen.configs.pipeline_configs.ideogram import (
Ideogram4PipelineConfig,
)
from sglang.multimodal_gen.configs.pipeline_configs.joy_image import ( from sglang.multimodal_gen.configs.pipeline_configs.joy_image import (
JoyImageEditPipelineConfig, JoyImageEditPipelineConfig,
) )
@@ -105,6 +108,7 @@ from sglang.multimodal_gen.configs.sample.hunyuan import (
HunyuanSamplingParams, HunyuanSamplingParams,
) )
from sglang.multimodal_gen.configs.sample.hunyuan3d import Hunyuan3DSamplingParams from sglang.multimodal_gen.configs.sample.hunyuan3d import Hunyuan3DSamplingParams
from sglang.multimodal_gen.configs.sample.ideogram import Ideogram4SamplingParams
from sglang.multimodal_gen.configs.sample.joy_image import ( from sglang.multimodal_gen.configs.sample.joy_image import (
JoyImageEditSamplingParams, JoyImageEditSamplingParams,
) )
@@ -404,7 +408,7 @@ def _get_config_info(
if len(matched_model_names) >= 1: if len(matched_model_names) >= 1:
if len(matched_model_names) > 1: if len(matched_model_names) > 1:
logger.warning( logger.warning(
f"More than one model name is matched, using the first matched" "More than one model name is matched, using the first matched"
) )
model_id = matched_model_names[0] model_id = matched_model_names[0]
return _CONFIG_REGISTRY.get(model_id) return _CONFIG_REGISTRY.get(model_id)
@@ -805,9 +809,9 @@ def _register_configs():
], ],
model_detectors=[ model_detectors=[
lambda hf_id: ( lambda hf_id: (
"flux.2-klein" in hf_id.lower() or "flux2-klein" in hf_id.lower() ("flux.2-klein" in hf_id.lower() or "flux2-klein" in hf_id.lower())
and "base" not in hf_id.lower()
) )
and "base" not in hf_id.lower()
], ],
) )
register_configs( register_configs(
@@ -819,9 +823,9 @@ def _register_configs():
], ],
model_detectors=[ model_detectors=[
lambda hf_id: ( lambda hf_id: (
"flux.2-klein" in hf_id.lower() or "flux2-klein" in hf_id.lower() ("flux.2-klein" in hf_id.lower() or "flux2-klein" in hf_id.lower())
and "base" in hf_id.lower()
) )
and "base" in hf_id.lower()
], ],
) )
register_configs( register_configs(
@@ -859,10 +863,12 @@ def _register_configs():
pipeline_config_cls=QwenImagePipelineConfig, pipeline_config_cls=QwenImagePipelineConfig,
hf_model_paths=["Qwen/Qwen-Image"], hf_model_paths=["Qwen/Qwen-Image"],
model_detectors=[ model_detectors=[
lambda hf_id: "qwen-image" in hf_id.lower() lambda hf_id: (
and "edit" not in hf_id.lower() "qwen-image" in hf_id.lower()
and "layered" not in hf_id.lower() and "edit" not in hf_id.lower()
and "2512" not in hf_id.lower() and "layered" not in hf_id.lower()
and "2512" not in hf_id.lower()
)
], ],
) )
register_configs( register_configs(
@@ -876,9 +882,11 @@ def _register_configs():
pipeline_config_cls=QwenImageEditPipelineConfig, pipeline_config_cls=QwenImageEditPipelineConfig,
hf_model_paths=["Qwen/Qwen-Image-Edit"], hf_model_paths=["Qwen/Qwen-Image-Edit"],
model_detectors=[ model_detectors=[
lambda hf_id: "qwen-image-edit" in hf_id.lower() lambda hf_id: (
and "2509" not in hf_id.lower() "qwen-image-edit" in hf_id.lower()
and "2511" not in hf_id.lower() and "2509" not in hf_id.lower()
and "2511" not in hf_id.lower()
)
], ],
) )
@@ -914,12 +922,14 @@ def _register_configs():
"stabilityai/stable-diffusion-3.5-large-diffusers", "stabilityai/stable-diffusion-3.5-large-diffusers",
], ],
model_detectors=[ model_detectors=[
lambda hf_id: "stable-diffusion-3-medium" in hf_id.lower() lambda hf_id: (
or "stable-diffusion-3.5-medium" in hf_id.lower() "stable-diffusion-3-medium" in hf_id.lower()
or "stable-diffusion-3.5-large" in hf_id.lower() or "stable-diffusion-3.5-medium" in hf_id.lower()
or "sd3-medium" in hf_id.lower() or "stable-diffusion-3.5-large" in hf_id.lower()
or "sd3.5-medium" in hf_id.lower() or "sd3-medium" in hf_id.lower()
or "sd3.5-large" in hf_id.lower() or "sd3.5-medium" in hf_id.lower()
or "sd3.5-large" in hf_id.lower()
)
], ],
) )
@@ -945,9 +955,11 @@ def _register_configs():
"BestWishYsh/Helios-Base", "BestWishYsh/Helios-Base",
], ],
model_detectors=[ model_detectors=[
lambda hf_id: "helios" in hf_id.lower() lambda hf_id: (
and "mid" not in hf_id.lower() "helios" in hf_id.lower()
and "distill" not in hf_id.lower() and "mid" not in hf_id.lower()
and "distill" not in hf_id.lower()
)
], ],
) )
register_configs( register_configs(
@@ -1031,6 +1043,19 @@ def _register_configs():
], ],
) )
# Ideogram 4
register_configs(
sampling_param_cls=Ideogram4SamplingParams,
pipeline_config_cls=Ideogram4PipelineConfig,
hf_model_paths=[
"ideogram-ai/ideogram-4-fp8",
],
model_detectors=[
lambda hf_id: "ideogram4pipeline" in hf_id.lower(),
lambda hf_id: "ideogram-4-fp8" in hf_id.lower(),
],
)
_register_configs() _register_configs()
@@ -49,6 +49,7 @@ def get_module_role(module_name: str) -> "RoleType | None":
denoising_prefixes = ( denoising_prefixes = (
"transformer", "transformer",
"unconditional_transformer",
"video_dit", "video_dit",
"audio_dit", "audio_dit",
"dual_tower_bridge", "dual_tower_bridge",
@@ -122,6 +122,10 @@ def get_world_group() -> GroupCoordinator:
return _WORLD return _WORLD
def world_group_is_initialized() -> bool:
return _WORLD is not None
def init_world_group( def init_world_group(
ranks: list[int], local_rank: int, backend: str ranks: list[int], local_rank: int, backend: str
) -> GroupCoordinator: ) -> GroupCoordinator:
@@ -200,6 +200,7 @@ async def generations(
upscaling_scale=request.upscaling_scale, upscaling_scale=request.upscaling_scale,
perf_dump_path=request.perf_dump_path, perf_dump_path=request.perf_dump_path,
use_pe=_get_extra_field(request, "use_pe"), use_pe=_get_extra_field(request, "use_pe"),
preset=_get_extra_field(request, "preset"),
) )
trace_headers = extract_trace_headers(raw_request.headers) trace_headers = extract_trace_headers(raw_request.headers)
batch = prepare_request( batch = prepare_request(
@@ -0,0 +1,87 @@
# SPDX-License-Identifier: Apache-2.0
import torch
import torch.nn as nn
import torch.nn.functional as F
from sglang.multimodal_gen.runtime.models.utils import set_weight_attrs
FP8_WEIGHT_DTYPE = torch.float8_e4m3fn
def dequantize_rowwise_fp8_weight(
weight: torch.Tensor,
weight_scale: torch.Tensor,
dtype: torch.dtype,
) -> torch.Tensor:
if weight.ndim != 2:
raise ValueError(f"FP8 linear weight must be 2-D, got shape {weight.shape}")
if weight_scale.ndim != 1 or weight_scale.shape[0] != weight.shape[0]:
raise ValueError(
"FP8 row-wise scale must have shape (out_features,), "
f"got weight={tuple(weight.shape)} scale={tuple(weight_scale.shape)}"
)
return weight.to(dtype) * weight_scale.to(dtype).unsqueeze(1)
class WeightOnlyFP8Linear(nn.Module):
"""Storage-only e4m3 FP8 linear with row-wise dequantization before matmul."""
def __init__(
self,
in_features: int,
out_features: int,
bias: bool = True,
compute_dtype: torch.dtype | None = None,
) -> None:
super().__init__()
self.in_features = in_features
self.out_features = out_features
self.compute_dtype = compute_dtype
self.weight = nn.Parameter(
torch.empty(out_features, in_features, dtype=FP8_WEIGHT_DTYPE),
requires_grad=False,
)
self.weight_scale = nn.Parameter(
torch.empty(out_features, dtype=torch.float32),
requires_grad=False,
)
set_weight_attrs(self.weight_scale, {"missing_param_init": "error"})
if bias:
self.bias = nn.Parameter(
torch.empty(
out_features, dtype=compute_dtype or torch.get_default_dtype()
),
requires_grad=False,
)
else:
self.register_parameter("bias", None)
def forward(self, x: torch.Tensor) -> torch.Tensor:
compute_dtype = self.compute_dtype or x.dtype
weight = dequantize_rowwise_fp8_weight(
self.weight, self.weight_scale, compute_dtype
)
bias = self.bias.to(compute_dtype) if self.bias is not None else None
return F.linear(x.to(compute_dtype), weight, bias)
def swap_linears_to_weight_only_fp8(module: nn.Module) -> None:
"""Recursively replace nn.Linear with WeightOnlyFP8Linear.
Ideogram FP8 checkpoints provide ``<linear>.weight_scale`` for every
quantized linear. Swapping before load lets strict state-dict checks verify
both the FP8 weight and its row-wise scale.
"""
for name, child in list(module.named_children()):
if isinstance(child, nn.Linear):
replacement = WeightOnlyFP8Linear(
child.in_features,
child.out_features,
bias=child.bias is not None,
compute_dtype=child.weight.dtype,
)
setattr(module, name, replacement)
else:
swap_linears_to_weight_only_fp8(child)
@@ -28,7 +28,11 @@
from .base import RotaryEmbedding from .base import RotaryEmbedding
from .factory import get_rope, get_rotary_pos_embed from .factory import get_rope, get_rotary_pos_embed
from .mrope import NDRotaryEmbedding from .mrope import (
NDRotaryEmbedding,
Qwen3VLTextRotaryEmbedding,
qwen3_apply_rotary_pos_emb,
)
from .utils import ( from .utils import (
_apply_rotary_emb, _apply_rotary_emb,
apply_flashinfer_rope_qk_inplace, apply_flashinfer_rope_qk_inplace,
@@ -42,6 +46,8 @@ __all__ = [
"RotaryEmbedding", "RotaryEmbedding",
# _mrope # _mrope
"NDRotaryEmbedding", "NDRotaryEmbedding",
"Qwen3VLTextRotaryEmbedding",
"qwen3_apply_rotary_pos_emb",
# _factory # _factory
"get_rope", "get_rope",
"get_rotary_pos_embed", "get_rotary_pos_embed",
@@ -69,6 +69,92 @@ def get_1d_rotary_pos_embed(
return freqs_cos, freqs_sin return freqs_cos, freqs_sin
def qwen3_apply_rotary_pos_emb(
q: torch.Tensor,
k: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Apply Qwen3-style RoPE to q/k tensors shaped [B, S, H, D]."""
half = q.shape[-1] // 2
q1 = q[..., :half]
q2 = q[..., half:]
q_embed = torch.empty_like(q)
q_embed[..., :half] = q1 * cos[..., :half] - q2 * sin[..., :half]
q_embed[..., half:] = q2 * cos[..., half:] + q1 * sin[..., half:]
half = k.shape[-1] // 2
k1 = k[..., :half]
k2 = k[..., half:]
k_embed = torch.empty_like(k)
k_embed[..., :half] = k1 * cos[..., :half] - k2 * sin[..., :half]
k_embed[..., half:] = k2 * cos[..., half:] + k1 * sin[..., half:]
return q_embed, k_embed
class Qwen3VLTextRotaryEmbedding(torch.nn.Module):
"""Qwen3-VL multi-dimensional rotary embedding with interleaved mRoPE."""
def __init__(
self,
head_dim: int = 128,
rope_theta: float = 5_000_000.0,
mrope_section: tuple[int, int, int] | list[int] = (24, 20, 20),
):
super().__init__()
self.rope_type = "default"
self.max_seq_len_cached = 262144
self.mrope_section = list(mrope_section)
self.head_dim = head_dim
inv_freq = 1.0 / (
rope_theta ** (torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim)
)
self.register_buffer("inv_freq", inv_freq, persistent=False)
self.attention_scaling = 1.0
def apply_interleaved_mrope(
self, freqs: torch.Tensor, mrope_section: list[int]
) -> torch.Tensor:
freqs_t = freqs[0].clone()
for dim, offset in enumerate((1, 2), start=1):
length = mrope_section[dim] * 3
idx = slice(offset, length, 3)
freqs_t[..., idx] = freqs[dim, ..., idx]
return freqs_t
@torch.no_grad()
def forward(
self, x: torch.Tensor, position_ids: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor]:
"""Return cos/sin for position IDs shaped [3, B, S], [B, S, 3], or [B, S]."""
if position_ids.ndim == 3 and position_ids.shape[-1] == 3:
position_ids = position_ids.permute(2, 0, 1)
elif position_ids.ndim == 2:
position_ids = position_ids[None, ...].expand(3, position_ids.shape[0], -1)
elif position_ids.ndim != 3 or position_ids.shape[0] != 3:
raise ValueError(
"Qwen3 mRoPE position_ids must have shape [3, B, S], [B, S, 3], "
f"or [B, S], got {tuple(position_ids.shape)}"
)
inv_freq_expanded = (
self.inv_freq[None, None, :, None]
.float()
.expand(3, position_ids.shape[1], -1, 1)
.to(position_ids.device)
)
position_ids_expanded = position_ids[:, :, None, :].float()
freqs = (inv_freq_expanded @ position_ids_expanded).transpose(2, 3)
freqs = self.apply_interleaved_mrope(freqs, self.mrope_section)
emb = torch.cat((freqs, freqs), dim=-1)
cos = emb.cos() * self.attention_scaling
sin = emb.sin() * self.attention_scaling
return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype)
class OneDRotaryEmbedding(torch.nn.Module): class OneDRotaryEmbedding(torch.nn.Module):
"""1D rotary positional embedding with caching.""" """1D rotary positional embedding with caching."""
@@ -246,6 +246,11 @@ class TextEncoderLoader(ComponentLoader):
if encoder_index == 0: if encoder_index == 0:
for key, value in diffusers_pretrained_config.__dict__.items(): for key, value in diffusers_pretrained_config.__dict__.items():
setattr(encoder_config.arch_config, key, value) setattr(encoder_config.arch_config, key, value)
post_diffusers_config_update = getattr(
encoder_config, "post_diffusers_config_update", None
)
if post_diffusers_config_update is not None:
post_diffusers_config_update()
encoder_dtype = server_args.pipeline_config.text_encoder_precisions[ encoder_dtype = server_args.pipeline_config.text_encoder_precisions[
encoder_index encoder_index
] ]
@@ -31,7 +31,7 @@ def _server_args_for_transformer_component(
server_args: ServerArgs, component_name: str server_args: ServerArgs, component_name: str
) -> ServerArgs: ) -> ServerArgs:
"""Mask global quantized override flags for secondary transformer components.""" """Mask global quantized override flags for secondary transformer components."""
if component_name != "transformer_2": if component_name not in ("transformer_2", "unconditional_transformer"):
return server_args return server_args
if ( if (
@@ -54,7 +54,12 @@ def _server_args_for_transformer_component(
class TransformerLoader(ComponentLoader): class TransformerLoader(ComponentLoader):
"""Shared loader for (video/audio) DiT transformers.""" """Shared loader for (video/audio) DiT transformers."""
component_names = ["transformer", "audio_dit", "video_dit"] component_names = [
"transformer",
"unconditional_transformer",
"audio_dit",
"video_dit",
]
expected_library = "diffusers" expected_library = "diffusers"
def load_customized( def load_customized(
@@ -76,7 +81,7 @@ class TransformerLoader(ComponentLoader):
# Config from Diffusers supersedes sgl_diffusion's model config # Config from Diffusers supersedes sgl_diffusion's model config
component_name = _normalize_component_type(component_name) component_name = _normalize_component_type(component_name)
server_args.model_paths[component_name] = component_model_path server_args.model_paths[component_name] = component_model_path
if component_name in ("transformer", "video_dit"): if component_name in ("transformer", "unconditional_transformer", "video_dit"):
pipeline_dit_config_attr = "dit_config" pipeline_dit_config_attr = "dit_config"
elif component_name in ("audio_dit",): elif component_name in ("audio_dit",):
pipeline_dit_config_attr = "audio_dit_config" pipeline_dit_config_attr = "audio_dit_config"
@@ -115,7 +120,7 @@ class TransformerLoader(ComponentLoader):
and component_server_args.transformer_weights_path is not None and component_server_args.transformer_weights_path is not None
): ):
logger.warning( logger.warning(
f"transformer_weights_path provided, but quantization config not resolved, which is unexpected and likely to cause errors" "transformer_weights_path provided, but quantization config not resolved, which is unexpected and likely to cause errors"
) )
else: else:
logger.debug("quantization config: %s", init_params["quant_config"]) logger.debug("quantization config: %s", init_params["quant_config"])
@@ -629,6 +629,13 @@ def load_model_from_full_model_state_dict(
else None else None
) )
if missing_param_init == "error":
raise ValueError(
f"Required checkpoint parameter '{new_param_name}' was not loaded. "
"This usually indicates a checkpoint/model-arch mismatch or a "
"broken weight-name mapping."
)
if missing_param_init is None and not any( if missing_param_init is None and not any(
pattern in new_param_name for pattern in LEGACY_ALLOWED_NEW_PARAM_PATTERNS pattern in new_param_name for pattern in LEGACY_ALLOWED_NEW_PARAM_PATTERNS
): ):
@@ -373,12 +373,15 @@ def resolve_transformer_quant_load_spec(
model_cls: type[nn.Module], model_cls: type[nn.Module],
cls_name: str, cls_name: str,
) -> TransformerQuantLoadSpec: ) -> TransformerQuantLoadSpec:
quant_config = _resolve_quant_config( if getattr(model_cls, "handles_checkpoint_quantization", False):
hf_config=hf_config, quant_config = None
server_args=server_args, else:
safetensors_list=safetensors_list, quant_config = _resolve_quant_config(
component_model_path=component_model_path, hf_config=hf_config,
) server_args=server_args,
safetensors_list=safetensors_list,
component_model_path=component_model_path,
)
if quant_config is not None: if quant_config is not None:
packed = getattr(model_cls, "packed_modules_mapping", None) packed = getattr(model_cls, "packed_modules_mapping", None)
@@ -32,6 +32,10 @@ from sglang.multimodal_gen.runtime.layers.linear import (
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import ( from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
QuantizationConfig, QuantizationConfig,
) )
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
Qwen3VLTextRotaryEmbedding,
qwen3_apply_rotary_pos_emb,
)
from sglang.multimodal_gen.runtime.layers.visual_embedding import timestep_embedding from sglang.multimodal_gen.runtime.layers.visual_embedding import timestep_embedding
from sglang.multimodal_gen.runtime.layers.vocab_parallel_embedding import ( from sglang.multimodal_gen.runtime.layers.vocab_parallel_embedding import (
VocabParallelEmbedding, VocabParallelEmbedding,
@@ -127,129 +131,6 @@ def compute_mrope_position_ids_vision(
return mrope_ids, next_offset return mrope_ids, next_offset
# -----------------------------------------------------------------------------
# Qwen3-style RoPE functions
# -----------------------------------------------------------------------------
def qwen3_apply_rotary_pos_emb(
q: torch.Tensor,
k: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Qwen3-style RoPE: (x * cos) + (rotate_half(x) * sin).
Args:
q: [B, S, H, D]
k: [B, S, H_kv, D]
cos: [1, S, 1, D] or broadcastable
sin: [1, S, 1, D] or broadcastable
"""
half = q.shape[-1] // 2
q1 = q[..., :half]
q2 = q[..., half:]
q_embed = torch.empty_like(q)
q_embed[..., :half] = q1 * cos[..., :half] - q2 * sin[..., :half]
q_embed[..., half:] = q2 * cos[..., half:] + q1 * sin[..., half:]
half = k.shape[-1] // 2
k1 = k[..., :half]
k2 = k[..., half:]
k_embed = torch.empty_like(k)
k_embed[..., :half] = k1 * cos[..., :half] - k2 * sin[..., :half]
k_embed[..., half:] = k2 * cos[..., half:] + k1 * sin[..., half:]
return q_embed, k_embed
# -----------------------------------------------------------------------------
# Qwen3VL-style Rotary Embedding
# -----------------------------------------------------------------------------
class Qwen3VLTextRotaryEmbedding(nn.Module):
"""Qwen3VL-style multi-dimensional rotary embedding."""
def __init__(
self,
head_dim: int = 128,
rope_theta: float = 5000000.0,
mrope_section: tuple[int, int, int] = (24, 20, 20),
):
super().__init__()
self.rope_type = "default"
self.max_seq_len_cached = 262144
self.mrope_section = list(mrope_section)
self.head_dim = head_dim
# Compute inverse frequencies
dim = head_dim
inv_freq = 1.0 / (
rope_theta ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim)
)
self.register_buffer("inv_freq", inv_freq, persistent=False)
self.attention_scaling = 1.0
def apply_interleaved_mrope(
self, freqs: torch.Tensor, mrope_section: list[int]
) -> torch.Tensor:
"""Apply interleaved MRoPE to 3D rotary embeddings.
Reorganizes frequency layout from chunked [TTT...HHH...WWW] to
interleaved [THTHWHTHW...TT], preserving frequency continuity.
Args:
freqs: (3, bs, seq_len, head_dim // 2)
mrope_section: (3,) section sizes
Returns:
freqs_t: (bs, seq_len, head_dim // 2)
"""
freqs_t = freqs[0].clone()
for dim, offset in enumerate((1, 2), start=1): # H, W
length = mrope_section[dim] * 3
idx = slice(offset, length, 3)
freqs_t[..., idx] = freqs[dim, ..., idx]
return freqs_t
@torch.no_grad()
def forward(
self, x: torch.Tensor, position_ids: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor]:
"""Compute cos and sin for rotary embeddings.
Args:
x: dummy tensor for dtype
position_ids: [3, B, S] or [B, S] position IDs
Returns:
(cos, sin) each of shape [B, S, D]
"""
if position_ids.ndim == 2:
position_ids = position_ids[None, ...].expand(3, position_ids.shape[0], -1)
# Expand inv_freq: [3, B, D//2, 1]
inv_freq_expanded = (
self.inv_freq[None, None, :, None]
.float()
.expand(3, position_ids.shape[1], -1, 1)
.to(position_ids.device)
)
# position_ids_expanded: [3, B, 1, S]
position_ids_expanded = position_ids[:, :, None, :].float()
# freqs: [3, B, D//2, S] -> transpose -> [3, B, S, D//2]
freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(
2, 3
)
freqs = self.apply_interleaved_mrope(freqs, self.mrope_section)
emb = torch.cat((freqs, freqs), dim=-1)
cos = emb.cos() * self.attention_scaling
sin = emb.sin() * self.attention_scaling
return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype)
# ----------------------------------------------------------------------------- # -----------------------------------------------------------------------------
# Cosmos3 Timestep Embedder # Cosmos3 Timestep Embedder
# ----------------------------------------------------------------------------- # -----------------------------------------------------------------------------
@@ -0,0 +1,298 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Any
import torch
import torch.nn as nn
import torch.nn.functional as F
from sglang.multimodal_gen.configs.models.dits.ideogram import Ideogram4DiTConfig
from sglang.multimodal_gen.runtime.layers.attention import (
USPAttention,
build_varlen_mask_meta,
)
from sglang.multimodal_gen.runtime.layers.quantization.weight_only_fp8 import (
WeightOnlyFP8Linear,
)
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
Qwen3VLTextRotaryEmbedding,
qwen3_apply_rotary_pos_emb,
)
from sglang.multimodal_gen.runtime.models.dits.base import BaseDiT
OUTPUT_IMAGE_INDICATOR = 2
LLM_TOKEN_INDICATOR = 3
class Ideogram4RMSNorm(nn.Module):
def __init__(self, dim: int, eps: float = 1e-6) -> None:
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
self.eps = eps
def forward(self, x: torch.Tensor) -> torch.Tensor:
return F.rms_norm(x, self.weight.shape, self.weight, self.eps)
def _linear(in_features: int, out_features: int, bias: bool = True):
return WeightOnlyFP8Linear(in_features, out_features, bias=bias)
class Ideogram4Attention(nn.Module):
def __init__(
self,
hidden_size: int,
num_heads: int,
eps: float,
supported_attention_backends,
) -> None:
super().__init__()
self.hidden_size = hidden_size
self.num_heads = num_heads
self.head_dim = hidden_size // num_heads
self.qkv = _linear(hidden_size, hidden_size * 3, bias=False)
self.norm_q = Ideogram4RMSNorm(self.head_dim, eps=eps)
self.norm_k = Ideogram4RMSNorm(self.head_dim, eps=eps)
self.attn = USPAttention(
num_heads=num_heads,
head_size=self.head_dim,
dropout_rate=0,
softmax_scale=None,
causal=False,
supported_attention_backends=supported_attention_backends,
)
self.o = _linear(hidden_size, hidden_size, bias=False)
def forward(self, x, cos, sin, attn_mask, attn_mask_meta):
batch_size, seq_len, _ = x.shape
qkv = self.qkv(x).view(batch_size, seq_len, 3, self.num_heads, self.head_dim)
q, k, v = qkv.unbind(dim=2)
q = self.norm_q(q)
k = self.norm_k(k)
q, k = qwen3_apply_rotary_pos_emb(q, k, cos, sin)
out = self.attn(q, k, v, attn_mask=attn_mask, attn_mask_meta=attn_mask_meta)
out = out.reshape(batch_size, seq_len, self.hidden_size)
return self.o(out)
class Ideogram4MLP(nn.Module):
def __init__(self, dim: int, hidden_dim: int) -> None:
super().__init__()
self.w1 = _linear(dim, hidden_dim, bias=False)
self.w2 = _linear(hidden_dim, dim, bias=False)
self.w3 = _linear(dim, hidden_dim, bias=False)
def forward(self, x):
return self.w2(F.silu(self.w1(x)) * self.w3(x))
class Ideogram4TransformerBlock(nn.Module):
def __init__(
self,
hidden_size,
intermediate_size,
num_heads,
norm_eps,
adaln_dim,
supported_attention_backends,
):
super().__init__()
self.attention = Ideogram4Attention(
hidden_size,
num_heads,
eps=1e-5,
supported_attention_backends=supported_attention_backends,
)
self.feed_forward = Ideogram4MLP(hidden_size, intermediate_size)
self.attention_norm1 = Ideogram4RMSNorm(hidden_size, eps=norm_eps)
self.ffn_norm1 = Ideogram4RMSNorm(hidden_size, eps=norm_eps)
self.attention_norm2 = Ideogram4RMSNorm(hidden_size, eps=norm_eps)
self.ffn_norm2 = Ideogram4RMSNorm(hidden_size, eps=norm_eps)
self.adaln_modulation = _linear(adaln_dim, 4 * hidden_size, bias=True)
def forward(self, x, cos, sin, adaln_input, attn_mask, attn_mask_meta):
scale_msa, gate_msa, scale_mlp, gate_mlp = self.adaln_modulation(
adaln_input
).chunk(4, dim=-1)
gate_msa = torch.tanh(gate_msa)
gate_mlp = torch.tanh(gate_mlp)
attn_out = self.attention(
self.attention_norm1(x) * (1.0 + scale_msa),
cos=cos,
sin=sin,
attn_mask=attn_mask,
attn_mask_meta=attn_mask_meta,
)
x = x + gate_msa * self.attention_norm2(attn_out)
x = x + gate_mlp * self.ffn_norm2(
self.feed_forward(self.ffn_norm1(x) * (1.0 + scale_mlp))
)
return x
def _sinusoidal_embedding(t: torch.Tensor, dim: int, scale: float = 1e4):
t = t.to(torch.float32)
half = dim // 2
freq = math.log(scale) / (half - 1)
freq = torch.exp(torch.arange(half, dtype=torch.float32, device=t.device) * -freq)
emb = t.unsqueeze(-1) * freq
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1)
if dim % 2 == 1:
emb = F.pad(emb, (0, 1))
return emb
class Ideogram4EmbedScalar(nn.Module):
def __init__(self, dim: int, input_range: tuple[float, float]) -> None:
super().__init__()
self.dim = dim
self.range_min, self.range_max = input_range
self.mlp_in = _linear(dim, dim, bias=True)
self.mlp_out = _linear(dim, dim, bias=True)
def forward(self, x):
compute_dtype = x.dtype
x = x.to(torch.float32)
scaled = 1e4 * (x - self.range_min) / (self.range_max - self.range_min)
emb = _sinusoidal_embedding(scaled, self.dim).to(compute_dtype)
return self.mlp_out(F.silu(self.mlp_in(emb)))
class Ideogram4FinalLayer(nn.Module):
def __init__(self, hidden_size: int, out_channels: int, adaln_dim: int) -> None:
super().__init__()
self.norm_final = nn.LayerNorm(hidden_size, eps=1e-6, elementwise_affine=False)
self.linear = _linear(hidden_size, out_channels, bias=True)
self.adaln_modulation = _linear(adaln_dim, hidden_size, bias=True)
def forward(self, x, c):
scale = 1.0 + self.adaln_modulation(F.silu(c))
return self.linear(self.norm_final(x) * scale)
class Ideogram4Transformer2DModel(BaseDiT):
_repeated_blocks = ["Ideogram4TransformerBlock"]
_fsdp_shard_conditions = Ideogram4DiTConfig().arch_config._fsdp_shard_conditions
_compile_conditions = Ideogram4DiTConfig().arch_config._compile_conditions
_supported_attention_backends = (
Ideogram4DiTConfig().arch_config._supported_attention_backends
)
param_names_mapping = {}
reverse_param_names_mapping = {}
handles_checkpoint_quantization = True
def __init__(
self,
config: Ideogram4DiTConfig,
hf_config: dict[str, Any],
**kwargs,
) -> None:
super().__init__(config, hf_config, **kwargs)
cfg = config.arch_config
self._supported_attention_backends = cfg._supported_attention_backends
hidden_size = cfg.num_attention_heads * cfg.attention_head_dim
self.hidden_size = hidden_size
self.num_attention_heads = cfg.num_attention_heads
self.num_channels_latents = cfg.in_channels
self.input_proj = _linear(cfg.in_channels, hidden_size, bias=True)
self.llm_cond_norm = Ideogram4RMSNorm(cfg.llm_features_dim, eps=1e-6)
self.llm_cond_proj = _linear(cfg.llm_features_dim, hidden_size, bias=True)
self.t_embedding = Ideogram4EmbedScalar(hidden_size, input_range=(0.0, 1.0))
self.adaln_proj = _linear(hidden_size, cfg.adaln_dim, bias=True)
self.embed_image_indicator = nn.Embedding(2, hidden_size)
self.rotary_emb = Qwen3VLTextRotaryEmbedding(
head_dim=cfg.attention_head_dim,
rope_theta=cfg.rope_theta,
mrope_section=cfg.mrope_section,
)
self.layers = nn.ModuleList(
[
Ideogram4TransformerBlock(
hidden_size=hidden_size,
intermediate_size=cfg.intermediate_size,
num_heads=cfg.num_attention_heads,
norm_eps=cfg.norm_eps,
adaln_dim=cfg.adaln_dim,
supported_attention_backends=self._supported_attention_backends,
)
for _ in range(cfg.num_layers)
]
)
self.final_layer = Ideogram4FinalLayer(
hidden_size=hidden_size,
out_channels=cfg.in_channels,
adaln_dim=cfg.adaln_dim,
)
def post_load_weights(self) -> None:
if not self.rotary_emb.inv_freq.is_meta:
return
cfg = self.config.arch_config
inv_freq = 1.0 / (
cfg.rope_theta
** (
torch.arange(
0,
cfg.attention_head_dim,
2,
dtype=torch.float32,
device=self.input_proj.weight.device,
)
/ cfg.attention_head_dim
)
)
self.rotary_emb.register_buffer("inv_freq", inv_freq, persistent=False)
def forward(
self,
*,
llm_features: torch.Tensor,
x: torch.Tensor,
t: torch.Tensor,
position_ids: torch.Tensor,
segment_ids: torch.Tensor,
indicator: torch.Tensor,
**kwargs,
) -> torch.Tensor:
param_dtype = self.embed_image_indicator.weight.dtype
x = x.to(param_dtype)
t = t.to(param_dtype)
llm_features = llm_features.to(param_dtype)
indicator = indicator.to(torch.long)
llm_token_mask = (indicator == LLM_TOKEN_INDICATOR).to(x.dtype).unsqueeze(-1)
output_image_mask = (
(indicator == OUTPUT_IMAGE_INDICATOR).to(x.dtype).unsqueeze(-1)
)
llm_features = llm_features * llm_token_mask
x = x * output_image_mask
x = self.input_proj(x) * output_image_mask
t_cond = self.t_embedding(t)
if t.dim() == 1:
t_cond = t_cond.unsqueeze(1)
adaln_input = F.silu(self.adaln_proj(t_cond))
llm_features = self.llm_cond_proj(self.llm_cond_norm(llm_features))
llm_features = llm_features * llm_token_mask
h = x + llm_features
h = h + self.embed_image_indicator(
(indicator == OUTPUT_IMAGE_INDICATOR).to(torch.long)
)
cos, sin = self.rotary_emb(h, position_ids)
cos = cos.unsqueeze(2)
sin = sin.unsqueeze(2)
# ideogram uses -1 padding; varlen meta enables fa packed attention
attn_mask = segment_ids > 0
attn_mask_meta = build_varlen_mask_meta(attn_mask)
for layer in self.layers:
h = layer(
h,
cos=cos,
sin=sin,
adaln_input=adaln_input,
attn_mask=attn_mask,
attn_mask_meta=attn_mask_meta,
)
return self.final_layer(h, c=adaln_input).to(torch.float32)
EntryClass = Ideogram4Transformer2DModel
@@ -0,0 +1,121 @@
# SPDX-License-Identifier: Apache-2.0
from collections.abc import Iterable
from typing import Tuple
import torch
from transformers.models.qwen3_vl.configuration_qwen3_vl import Qwen3VLTextConfig
from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput
from sglang.multimodal_gen.configs.models.encoders.ideogram import (
Ideogram4TextEncoderConfig,
)
from sglang.multimodal_gen.runtime.layers.quantization.weight_only_fp8 import (
swap_linears_to_weight_only_fp8,
)
from sglang.multimodal_gen.runtime.loader.weight_utils import default_weight_loader
from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context
from sglang.multimodal_gen.runtime.models.encoders.base import TextEncoder
from sglang.multimodal_gen.runtime.models.encoders.qwen3vl import Qwen3VLTextModel
class IdeogramQwen3VLTextEncoder(TextEncoder):
"""Language-only Qwen3-VL text encoder stored inside Ideogram checkpoints."""
_activation_layers = (0, 3, 6, 9, 12, 15, 18, 21, 24, 27, 30, 33, 35)
def __init__(self, config: Ideogram4TextEncoderConfig) -> None:
super().__init__(config)
arch_config = config.arch_config
text_config = getattr(arch_config, "text_config")
if isinstance(text_config, dict):
text_config = Qwen3VLTextConfig(**text_config)
self.language_model = Qwen3VLTextModel(text_config)
if getattr(arch_config, "ideogram_fp8_weight_only", False):
swap_linears_to_weight_only_fp8(self.language_model)
@torch.no_grad()
def forward(
self,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
**kwargs,
) -> BaseEncoderOutput:
outputs = self.language_model(
input_ids=input_ids,
attention_mask=attention_mask,
position_ids=position_ids,
inputs_embeds=inputs_embeds,
use_cache=False,
output_hidden_states=output_hidden_states,
return_dict=True,
)
return BaseEncoderOutput(
last_hidden_state=outputs.last_hidden_state,
hidden_states=outputs.hidden_states,
attentions=outputs.attentions,
)
def encode_ideogram_features(
self,
token_ids: torch.Tensor,
text_position_ids: torch.Tensor,
indicator: torch.Tensor,
llm_token_indicator: int,
) -> torch.Tensor:
batch_size, seq_len = token_ids.shape
hidden_size = self.language_model.config.hidden_size
out_dim = hidden_size * len(self._activation_layers)
features = torch.zeros(
batch_size,
seq_len,
out_dim,
dtype=torch.float32,
device=token_ids.device,
)
for batch_idx in range(batch_size):
text_mask = indicator[batch_idx] == llm_token_indicator
cur_token_ids = token_ids[batch_idx, text_mask].unsqueeze(0)
if cur_token_ids.numel() == 0:
continue
pos_2d = text_position_ids[batch_idx, text_mask, 0].unsqueeze(0)
position_ids = pos_2d[None, ...].expand(4, 1, -1)
attention_mask = torch.ones_like(cur_token_ids)
with set_forward_context(current_timestep=0, attn_metadata=None):
outputs = self.forward(
input_ids=cur_token_ids,
position_ids=position_ids,
attention_mask=attention_mask,
output_hidden_states=True,
)
assert outputs.hidden_states is not None
selected = [outputs.hidden_states[i] for i in self._activation_layers]
stacked = torch.stack(selected, dim=0).permute(1, 2, 3, 0)
features[batch_idx, text_mask] = stacked.reshape(
1, cur_token_ids.shape[1], -1
)[0].to(torch.float32)
return features
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
loaded_params: set[str] = set()
params_dict = dict(self.named_parameters(remove_duplicate=False))
for name, loaded_weight in weights:
if name.startswith("visual."):
continue
if "rotary_emb.inv_freq" in name:
continue
param = params_dict.get(name)
if param is None:
raise KeyError(
f"Unexpected weight name while loading Ideogram text encoder: {name}"
)
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight.to(param.dtype))
loaded_params.add(name)
return loaded_params
EntryClass = IdeogramQwen3VLTextEncoder
@@ -0,0 +1,58 @@
# SPDX-License-Identifier: Apache-2.0
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import (
InputValidationStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ideogram import (
Ideogram4DecodingStage,
Ideogram4DenoisingStage,
Ideogram4TextEncodingStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
class Ideogram4Pipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "Ideogram4Pipeline"
_required_config_modules = [
"text_encoder",
"tokenizer",
"vae",
"transformer",
"unconditional_transformer",
"scheduler",
]
def create_pipeline_stages(self, server_args: ServerArgs):
self.add_stage(InputValidationStage())
self.add_stage_factory(
RoleType.ENCODER,
lambda: Ideogram4TextEncodingStage(
text_encoder=self.get_module("text_encoder"),
tokenizer=self.get_module("tokenizer"),
),
"ideogram4_text_encoding_stage",
)
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,
),
"ideogram4_denoising_stage",
)
self.add_stage_factory(
RoleType.DECODER,
lambda: Ideogram4DecodingStage(vae=self.get_module("vae")),
"ideogram4_decoding_stage",
)
EntryClass = Ideogram4Pipeline
@@ -18,7 +18,10 @@ import torch
from tqdm.auto import tqdm from tqdm.auto import tqdm
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
from sglang.multimodal_gen.runtime.distributed.parallel_state import get_world_rank from sglang.multimodal_gen.runtime.distributed.parallel_state import (
get_world_rank,
world_group_is_initialized,
)
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import ( from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
ComponentUse, ComponentUse,
) )
@@ -97,10 +100,11 @@ class PipelineStage(StageDedupMixin, ABC):
disable: bool = False, disable: bool = False,
**kwargs, **kwargs,
) -> tqdm: ) -> tqdm:
is_main_rank = not world_group_is_initialized() or get_world_rank() == 0
return tqdm( return tqdm(
iterable=iterable, iterable=iterable,
total=total, total=total,
disable=disable or get_world_rank() != 0, disable=disable or not is_main_rank,
**kwargs, **kwargs,
) )
@@ -42,6 +42,7 @@ from sglang.multimodal_gen.runtime.distributed import (
get_tp_group, get_tp_group,
get_world_group, get_world_group,
get_world_size, get_world_size,
model_parallel_is_initialized,
) )
from sglang.multimodal_gen.runtime.distributed.cfg_parallel_utils import ( from sglang.multimodal_gen.runtime.distributed.cfg_parallel_utils import (
run_cfg_parallel, run_cfg_parallel,
@@ -57,6 +58,7 @@ from sglang.multimodal_gen.runtime.distributed.communication_op import (
) )
from sglang.multimodal_gen.runtime.distributed.parallel_state import ( from sglang.multimodal_gen.runtime.distributed.parallel_state import (
get_classifier_free_guidance_world_size, get_classifier_free_guidance_world_size,
world_group_is_initialized,
) )
from sglang.multimodal_gen.runtime.layers.attention.selector import get_attn_backend from sglang.multimodal_gen.runtime.layers.attention.selector import get_attn_backend
from sglang.multimodal_gen.runtime.layers.attention.STA_configuration import ( from sglang.multimodal_gen.runtime.layers.attention.STA_configuration import (
@@ -852,7 +854,12 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
"invalidations": 0, "invalidations": 0,
} }
if ctx.is_warmup or get_world_group().local_rank != 0: if not (active or requested):
return
if ctx.is_warmup or (
world_group_is_initialized() and get_world_group().local_rank != 0
):
return return
if active: if active:
@@ -1074,7 +1081,7 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
not state not state
or not state["requested"] or not state["requested"]
or ctx.is_warmup or ctx.is_warmup
or get_world_group().local_rank != 0 or (world_group_is_initialized() and get_world_group().local_rank != 0)
): ):
return return
@@ -1117,7 +1124,7 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
# Gather noise_pred if using sequence parallelism # Gather noise_pred if using sequence parallelism
# noise_pred has the same shape as latents (sharded along sequence dimension) # noise_pred has the same shape as latents (sharded along sequence dimension)
if ( if (
get_sp_world_size() > 1 self._sp_world_size() > 1
and getattr(batch, "did_sp_shard_latents", False) and getattr(batch, "did_sp_shard_latents", False)
and server_args.comfyui_mode and server_args.comfyui_mode
and hasattr(batch, "noise_pred") and hasattr(batch, "noise_pred")
@@ -1166,7 +1173,7 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
def _preprocess_sp_latents(self, batch: Req, server_args: ServerArgs): def _preprocess_sp_latents(self, batch: Req, server_args: ServerArgs):
"""Shard latents for Sequence Parallelism if applicable.""" """Shard latents for Sequence Parallelism if applicable."""
if get_sp_world_size() <= 1: if self._sp_world_size() <= 1:
return return
if batch.latents is not None: if batch.latents is not None:
@@ -1204,7 +1211,7 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
trajectory_tensor: torch.Tensor | None, trajectory_tensor: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor | None]: ) -> tuple[torch.Tensor, torch.Tensor | None]:
"""Gather latents after Sequence Parallelism if they were sharded.""" """Gather latents after Sequence Parallelism if they were sharded."""
if get_sp_world_size() > 1 and getattr(batch, "did_sp_shard_latents", False): if self._sp_world_size() > 1 and getattr(batch, "did_sp_shard_latents", False):
latents = self.server_args.pipeline_config.gather_latents_for_sp( latents = self.server_args.pipeline_config.gather_latents_for_sp(
latents, batch=batch latents, batch=batch
) )
@@ -1230,6 +1237,11 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
trajectory_tensor = trajectory_tensor[:, :, :orig_s, :] trajectory_tensor = trajectory_tensor[:, :, :orig_s, :]
return latents, trajectory_tensor return latents, trajectory_tensor
def _sp_world_size(self) -> int:
if not model_parallel_is_initialized():
return 1
return get_sp_world_size()
def step_profile(self): def step_profile(self):
profiler = SGLDiffusionProfiler.get_instance() profiler = SGLDiffusionProfiler.get_instance()
if profiler: if profiler:
@@ -0,0 +1,525 @@
# SPDX-License-Identifier: Apache-2.0
import math
from dataclasses import dataclass
import torch
from sglang.multimodal_gen.configs.pipeline_configs.ideogram import (
LATENT_SCALE,
LATENT_SHIFT,
)
from sglang.multimodal_gen.configs.sample.ideogram import IDEOGRAM4_PRESETS
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
ComponentUse,
)
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch, Req
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
from sglang.multimodal_gen.runtime.pipelines_core.stages.decoding import (
_ensure_tensor_decode_output,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import (
DenoisingContext,
DenoisingStage,
DenoisingStepState,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.text_encoding import (
TextEncodingStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
StageValidators as V,
)
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.nvtx_pytorch_hooks import maybe_nvtx_range
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
SEQUENCE_PADDING_INDICATOR = -1
OUTPUT_IMAGE_INDICATOR = 2
LLM_TOKEN_INDICATOR = 3
IMAGE_POSITION_OFFSET = 65536
@dataclass(frozen=True)
class LogitNormalSchedule:
mean: float
std: float = 1.0
logsnr_min: float = -15.0
logsnr_max: float = 18.0
def __call__(self, t: torch.Tensor) -> torch.Tensor:
t = t.to(torch.float64)
z = torch.special.ndtri(t)
y = self.mean + self.std * z
t_ = 1 - torch.special.expit(y)
t_min = 1.0 / (1 + math.exp(0.5 * self.logsnr_max))
t_max = 1.0 / (1 + math.exp(0.5 * self.logsnr_min))
return t_.clamp(t_min, t_max).to(torch.float32)
@dataclass(frozen=True)
class Ideogram4TextEncodingFingerprint:
prompt: object
height: int
width: int
num_outputs_per_prompt: int
max_text_tokens: int
patch_size: int
ae_scale_factor: int
def get_schedule_for_resolution(image_resolution, known_mean: float, std: float):
num_pixels = image_resolution[0] * image_resolution[1]
known_pixels = 512 * 512
mean = known_mean + 0.5 * math.log(num_pixels / known_pixels)
return LogitNormalSchedule(mean=mean, std=std)
def make_step_intervals(num_steps: int) -> torch.Tensor:
return torch.linspace(0.0, 1.0, num_steps + 1, dtype=torch.float32)
class Ideogram4Scheduler:
order = 1
init_noise_sigma = 1.0
num_train_timesteps = 1
def __init__(self) -> None:
self.timesteps = torch.empty(0, dtype=torch.float32)
self._begin_index = None
def set_begin_index(self, begin_index: int) -> None:
self._begin_index = begin_index
def set_timesteps(self, num_inference_steps: int, device=None) -> None:
self.timesteps = torch.arange(
num_inference_steps - 1,
-1,
-1,
dtype=torch.float32,
device=device or get_local_torch_device(),
)
def scale_model_input(self, sample: torch.Tensor, timestep=None) -> torch.Tensor:
return sample
def step(self, model_output, timestep, sample, return_dict=False, **kwargs):
raise RuntimeError("Ideogram4DenoisingStage applies its custom scheduler step")
class Ideogram4TextEncodingStage(TextEncodingStage):
deduplicated_extra_tensor_tree_output_keys = ("ideogram4",)
def __init__(self, text_encoder, tokenizer) -> None:
super().__init__([text_encoder], [tokenizer])
def _tokenize(self, prompt: str, max_text_tokens: int):
messages = [{"role": "user", "content": [{"type": "text", "text": prompt}]}]
text = self.tokenizers[0].apply_chat_template(
messages, add_generation_prompt=True, tokenize=False
)
encoded = self.tokenizers[0](
text, return_tensors="pt", add_special_tokens=False
)
token_ids = encoded["input_ids"][0]
num_text_tokens = int(token_ids.shape[0])
if num_text_tokens > max_text_tokens:
raise ValueError(
f"prompt has {num_text_tokens} tokens, exceeds max_text_tokens={max_text_tokens}"
)
return token_ids, num_text_tokens
def _build_inputs(self, prompts: list[str], height: int, width: int, server_args):
cfg = server_args.pipeline_config
tokenized = [self._tokenize(p, cfg.max_text_tokens) for p in prompts]
batch_size = len(prompts)
patch = cfg.patch_size * cfg.ae_scale_factor
if height < 256 or height > 2048 or width < 256 or width > 2048:
raise ValueError("height/width must be between 256 and 2048")
if height % patch != 0 or width % patch != 0:
raise ValueError(
f"height/width must be divisible by patch_size*ae_scale_factor={patch}"
)
grid_h = height // patch
grid_w = width // patch
num_image_tokens = grid_h * grid_w
max_text_tokens = max(num_text for _, num_text in tokenized)
total_seq_len = max_text_tokens + num_image_tokens
device = get_local_torch_device()
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)
t_idx = torch.zeros_like(h_idx)
image_pos = torch.stack([t_idx, h_idx, w_idx], dim=1) + IMAGE_POSITION_OFFSET
token_ids = torch.zeros(batch_size, total_seq_len, dtype=torch.long)
text_position_ids = torch.zeros(batch_size, total_seq_len, 3, dtype=torch.long)
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, (toks, num_text) in enumerate(tokenized):
pad_len = max_text_tokens - num_text
total_unpadded = num_text + num_image_tokens
offset = pad_len
token_ids[b, offset : offset + num_text] = toks
text_pos = torch.arange(num_text)
text_pos_3d = torch.stack([text_pos, text_pos, text_pos], dim=1)
text_position_ids[b, offset : offset + num_text] = text_pos_3d
position_ids[b, offset : offset + num_text] = text_pos_3d
position_ids[b, offset + num_text :] = image_pos
indicator[b, offset : offset + num_text] = LLM_TOKEN_INDICATOR
indicator[b, offset + num_text :] = OUTPUT_IMAGE_INDICATOR
segment_ids[b, offset : offset + total_unpadded] = 1
return {
"token_ids": token_ids.to(device),
"text_position_ids": text_position_ids.to(device),
"position_ids": position_ids.to(device),
"segment_ids": segment_ids.to(device),
"indicator": indicator.to(device),
"num_image_tokens": num_image_tokens,
"grid_h": grid_h,
"grid_w": grid_w,
"max_text_tokens": max_text_tokens,
}
@torch.no_grad()
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
prompts = batch.prompt if isinstance(batch.prompt, list) else [batch.prompt]
prompts = [p or " " for p in prompts]
if batch.num_outputs_per_prompt > 1:
prompts = [
prompt
for prompt in prompts
for _ in range(batch.num_outputs_per_prompt)
]
inputs = self._build_inputs(prompts, batch.height, batch.width, server_args)
with self.use_declared_component(
component_name="text_encoder", module=self.text_encoders[0]
) as text_encoder:
llm_features = text_encoder.encode_ideogram_features(
inputs["token_ids"],
inputs["text_position_ids"],
inputs["indicator"],
LLM_TOKEN_INDICATOR,
)
batch.prompt_embeds = [llm_features]
batch.prompt_embeds_mask = [
(inputs["indicator"] == LLM_TOKEN_INDICATOR).to(torch.bool)
]
batch.extra["ideogram4"] = inputs
return batch
def verify_input(self, batch: Req, server_args: ServerArgs) -> VerificationResult:
result = VerificationResult()
result.add_check("prompt", batch.prompt, V.string_or_list_strings)
result.add_check("height", batch.height, V.positive_int)
result.add_check("width", batch.width, V.positive_int)
result.add_check(
"num_outputs_per_prompt", batch.num_outputs_per_prompt, V.positive_int
)
return result
def verify_output(self, batch: Req, server_args: ServerArgs) -> VerificationResult:
result = VerificationResult()
result.add_check(
"prompt_embeds", batch.prompt_embeds, V.list_of_tensors_min_dims(2)
)
result.add_check(
"prompt_embeds_mask",
batch.prompt_embeds_mask,
V.list_of_tensors_min_dims(2),
)
result.add_check(
"ideogram4_extra",
batch.extra.get("ideogram4"),
lambda x: isinstance(x, dict),
)
return result
def build_dedup_fingerprint(
self, batch: Req, server_args: ServerArgs
) -> Ideogram4TextEncodingFingerprint:
cfg = server_args.pipeline_config
return Ideogram4TextEncodingFingerprint(
prompt=self.freeze_for_dedup(batch.prompt),
height=int(batch.height),
width=int(batch.width),
num_outputs_per_prompt=int(batch.num_outputs_per_prompt),
max_text_tokens=int(cfg.max_text_tokens),
patch_size=int(cfg.patch_size),
ae_scale_factor=int(cfg.ae_scale_factor),
)
class Ideogram4DenoisingStage(DenoisingStage):
def __init__(self, transformer, unconditional_transformer, pipeline=None) -> None:
super().__init__(
transformer=transformer,
scheduler=Ideogram4Scheduler(),
pipeline=pipeline,
)
self.unconditional_transformer = unconditional_transformer
self._maybe_enable_torch_compile(self.unconditional_transformer)
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 _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_enable_torch_compile(transformer)
def _manage_unconditional_transformer_use_site(self, batch: Req) -> None:
manager = self._component_residency_manager
if manager is None:
return
use = self._declared_component_use(
component_name="unconditional_transformer",
phase="unconditional_transformer",
)
manager.begin_use(use, module=self.unconditional_transformer)
def _manage_dit_use_site(
self,
current_model: torch.nn.Module,
current_phase: str,
batch: Req,
) -> None:
if self._component_residency_manager is None:
return
super()._manage_dit_use_site(current_model, current_phase, batch)
def _preprocess_sp_latents(self, batch: Req, server_args: ServerArgs):
batch.did_sp_shard_latents = False
def _postprocess_sp_latents(
self,
batch: Req,
latents: torch.Tensor,
trajectory_tensor: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor | None]:
return latents, trajectory_tensor
def verify_input(self, batch: Req, server_args: ServerArgs) -> VerificationResult:
return VerificationResult()
def _prepare_denoising_loop(
self, batch: Req, server_args: ServerArgs
) -> DenoisingContext:
preset = getattr(batch, "preset", "V4_DEFAULT_20")
if preset not in IDEOGRAM4_PRESETS:
raise ValueError(
f"Unknown Ideogram 4 preset {preset!r}; expected one of {sorted(IDEOGRAM4_PRESETS)}"
)
preset_cfg = IDEOGRAM4_PRESETS[preset]
num_steps = int(preset_cfg["num_steps"])
device = get_local_torch_device()
schedule = get_schedule_for_resolution(
(batch.height, batch.width),
known_mean=float(preset_cfg["mu"]),
std=float(preset_cfg["std"]),
)
step_intervals = make_step_intervals(num_steps).to(device)
guidance_schedule = torch.as_tensor(
preset_cfg["guidance_schedule"], dtype=torch.float32, device=device
)
self.scheduler.set_timesteps(num_steps, device=device)
batch.scheduler = self.scheduler
batch.timesteps = self.scheduler.timesteps
batch.num_inference_steps = num_steps
ctx = super()._prepare_denoising_loop(batch, server_args)
# ideogram fp8 denoising keeps explicit fp32 latent/scheduler math;
# wrapping the full loop in bf16 autocast collapses latent variance
ctx.autocast_enabled = False
data = batch.extra["ideogram4"]
z = ctx.latents.to(device, dtype=torch.float32)
llm_features = batch.prompt_embeds[0]
batch_size = z.shape[0]
max_text_tokens = data["max_text_tokens"]
num_image_tokens = data["num_image_tokens"]
latent_dim = z.shape[-1]
text_z_padding = torch.zeros(
batch_size,
max_text_tokens,
latent_dim,
dtype=torch.float32,
device=z.device,
)
neg_position_ids = data["position_ids"][:, max_text_tokens:]
neg_segment_ids = data["segment_ids"][:, max_text_tokens:]
neg_indicator = data["indicator"][:, max_text_tokens:]
neg_llm_features = torch.zeros(
batch_size,
num_image_tokens,
llm_features.shape[-1],
dtype=llm_features.dtype,
device=z.device,
)
ctx.latents = z
ctx.extra.update(
{
"ideogram4_schedule": schedule,
"ideogram4_step_intervals": step_intervals,
"ideogram4_guidance_schedule": guidance_schedule,
"ideogram4_text_z_padding": text_z_padding,
"ideogram4_neg_position_ids": neg_position_ids,
"ideogram4_neg_segment_ids": neg_segment_ids,
"ideogram4_neg_indicator": neg_indicator,
"ideogram4_neg_llm_features": neg_llm_features,
}
)
return ctx
def _run_denoising_step(
self,
ctx: DenoisingContext,
step: DenoisingStepState,
batch: Req,
server_args: ServerArgs,
) -> None:
data = batch.extra["ideogram4"]
z = ctx.latents.to(dtype=torch.float32)
llm_features = batch.prompt_embeds[0]
max_text_tokens = data["max_text_tokens"]
schedule = ctx.extra["ideogram4_schedule"]
step_intervals = ctx.extra["ideogram4_step_intervals"]
guidance_schedule = ctx.extra["ideogram4_guidance_schedule"]
i = step.t_int
t_val = float(schedule(step_intervals[i + 1].unsqueeze(0)).item())
s_val = float(schedule(step_intervals[i].unsqueeze(0)).item())
t = torch.full((z.shape[0],), t_val, dtype=torch.float32, device=z.device)
pos_z = torch.cat([ctx.extra["ideogram4_text_z_padding"], z], dim=1)
use_nvtx = self.current_use_nvtx
with maybe_nvtx_range("predict_noise", use_nvtx):
with set_forward_context(
current_timestep=i,
attn_metadata=step.attn_metadata,
forward_batch=batch,
):
pos_out = step.current_model(
llm_features=llm_features,
x=pos_z,
t=t,
position_ids=data["position_ids"],
segment_ids=data["segment_ids"],
indicator=data["indicator"],
)
pos_v = pos_out[:, max_text_tokens:]
self._manage_unconditional_transformer_use_site(batch)
with set_forward_context(
current_timestep=i,
attn_metadata=step.attn_metadata,
forward_batch=batch,
):
neg_v = self.unconditional_transformer(
llm_features=ctx.extra["ideogram4_neg_llm_features"],
x=z,
t=t,
position_ids=ctx.extra["ideogram4_neg_position_ids"],
segment_ids=ctx.extra["ideogram4_neg_segment_ids"],
indicator=ctx.extra["ideogram4_neg_indicator"],
)
with maybe_nvtx_range("scheduler_step", use_nvtx):
velocity = (
guidance_schedule[i] * pos_v + (1.0 - guidance_schedule[i]) * neg_v
)
ctx.latents = z + velocity * (s_val - t_val)
class Ideogram4DecodingStage(PipelineStage):
@property
def role_affinity(self):
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
return RoleType.DECODER
def __init__(self, vae) -> None:
super().__init__()
self.vae = vae
def component_uses(
self, server_args: ServerArgs, stage_name: str | None = None
) -> list[ComponentUse]:
return [
ComponentUse(
self._component_stage_name(stage_name),
"vae",
target_dtype=PRECISION_TO_TYPE[
server_args.pipeline_config.vae_precision
],
keep_ready_after_warmup=True,
)
]
@torch.no_grad()
def forward(self, batch: Req, server_args: ServerArgs) -> OutputBatch:
data = batch.extra["ideogram4"]
latents = batch.latents.to(get_local_torch_device())
cfg = server_args.pipeline_config
patch = cfg.patch_size
shift = torch.tensor(LATENT_SHIFT, device=latents.device, dtype=latents.dtype)
scale = torch.tensor(LATENT_SCALE, device=latents.device, dtype=latents.dtype)
z = latents * scale.to(latents.dtype) + shift.to(latents.dtype)
batch_size = z.shape[0]
grid_h = data["grid_h"]
grid_w = data["grid_w"]
ae_channels = z.shape[-1] // (patch * patch)
z = z.view(batch_size, grid_h, grid_w, patch, patch, ae_channels)
z = z.permute(0, 5, 1, 3, 2, 4).contiguous()
z = z.view(batch_size, ae_channels, grid_h * patch, grid_w * patch)
vae_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision]
with self.use_declared_component(component_name="vae", module=self.vae) as vae:
z = z.to(vae_dtype)
decoded = vae.decode(z)
frames = _ensure_tensor_decode_output(decoded)
frames = (frames / 2 + 0.5).clamp(0, 1)
return OutputBatch(
output=frames,
trajectory_timesteps=batch.trajectory_timesteps,
trajectory_latents=batch.trajectory_latents,
rollout_trajectory_data=batch.rollout_trajectory_data,
trajectory_decoded=None,
metrics=batch.metrics,
noise_pred=None,
)
@@ -1,3 +1,4 @@
import json
import os import os
import pytest import pytest
@@ -41,8 +42,6 @@ def _write_github_step_summary(content: str):
def _write_results_json(results: list, output_path: str = "diffusion-results.json"): def _write_results_json(results: list, output_path: str = "diffusion-results.json"):
"""Write performance results to JSON file for CI artifact collection.""" """Write performance results to JSON file for CI artifact collection."""
import json
try: try:
with open(output_path, "w") as f: with open(output_path, "w") as f:
json.dump(results, f, indent=2) json.dump(results, f, indent=2)
@@ -107,6 +106,8 @@ def pytest_sessionfinish(session):
print("[DEBUG] No results collected, skipping summary output") print("[DEBUG] No results collected, skipping summary output")
return return
sorted_results = sorted(results, key=lambda x: (x["class_name"], x["test_name"]))
# Print to stdout (existing behavior) # Print to stdout (existing behavior)
print("\n\n" + "=" * 35 + " Performance Summary " + "=" * 35) print("\n\n" + "=" * 35 + " Performance Summary " + "=" * 35)
print( print(
@@ -124,7 +125,7 @@ def pytest_sessionfinish(session):
+ "-" * 20 + "-" * 20
) )
for entry in sorted(results, key=lambda x: x["class_name"]): for entry in sorted_results:
print( print(
f"{entry['class_name']:<30} | {entry['test_name']:<20} | {entry['e2e_ms']:>12.2f} | " f"{entry['class_name']:<30} | {entry['test_name']:<20} | {entry['e2e_ms']:>12.2f} | "
f"{entry['avg_denoise_ms']:>18.2f} | {entry['median_denoise_ms']:>20.2f}" f"{entry['avg_denoise_ms']:>18.2f} | {entry['median_denoise_ms']:>20.2f}"
@@ -133,7 +134,7 @@ def pytest_sessionfinish(session):
print("=" * 91) print("=" * 91)
print("\n\n" + "=" * 36 + " Detailed Reports " + "=" * 37) print("\n\n" + "=" * 36 + " Detailed Reports " + "=" * 37)
for entry in sorted(results, key=lambda x: x["class_name"]): for entry in sorted_results:
print(f"\n--- Details for {entry['class_name']} / {entry['test_name']} ---") print(f"\n--- Details for {entry['class_name']} / {entry['test_name']} ---")
stage_report = ", ".join( stage_report = ", ".join(
f"{name}:{duration:.2f}ms" f"{name}:{duration:.2f}ms"
@@ -151,10 +152,14 @@ def pytest_sessionfinish(session):
print(f" Sampled Steps: {step_report}") print(f" Sampled Steps: {step_report}")
print("=" * 91) print("=" * 91)
print("\n\n" + "=" * 34 + " Performance Data JSON " + "=" * 34)
print(json.dumps(sorted_results, indent=2, sort_keys=True))
print("=" * 91)
# Write to GitHub Step Summary (new behavior for CI monitoring) # Write to GitHub Step Summary (new behavior for CI monitoring)
markdown_report = _generate_diffusion_markdown_report(results) markdown_report = _generate_diffusion_markdown_report(sorted_results)
if markdown_report: if markdown_report:
_write_github_step_summary(markdown_report) _write_github_step_summary(markdown_report)
# Write results to JSON file for CI artifact collection # Write results to JSON file for CI artifact collection
_write_results_json(results) _write_results_json(sorted_results)
@@ -19,6 +19,12 @@
"psnr_threshold": 28.0, "psnr_threshold": 28.0,
"mean_abs_diff_threshold": 8.0 "mean_abs_diff_threshold": 8.0
}, },
"ideogram4_fp8_t2i": {
"clip_threshold": 0.97,
"ssim_threshold": 0.78,
"psnr_threshold": 18.0,
"mean_abs_diff_threshold": 18.0
},
"flux_2_klein_image_t2i": { "flux_2_klein_image_t2i": {
"clip_threshold": 0.94, "clip_threshold": 0.94,
"ssim_threshold": 0.78, "ssim_threshold": 0.78,
@@ -18,6 +18,7 @@ from sglang.multimodal_gen.test.server.testcase_configs import (
DiffusionSamplingParams, DiffusionSamplingParams,
DiffusionServerArgs, DiffusionServerArgs,
DiffusionTestCase, DiffusionTestCase,
IDEOGRAM4_CI_sampling_params,
LINGBOT_WORLD_REALTIME_sampling_params, LINGBOT_WORLD_REALTIME_sampling_params,
MODELOPT_T2I_CI_sampling_params, MODELOPT_T2I_CI_sampling_params,
MODELOPT_T2V_CI_sampling_params, MODELOPT_T2V_CI_sampling_params,
@@ -95,6 +96,16 @@ ONE_GPU_CASES: list[DiffusionTestCase] = [
run_models_api_check=False, run_models_api_check=False,
run_t2v_input_reference_check=False, run_t2v_input_reference_check=False,
), ),
DiffusionTestCase(
"ideogram4_fp8_t2i",
DiffusionServerArgs(
model_path="ideogram-ai/ideogram-4-fp8",
),
IDEOGRAM4_CI_sampling_params,
run_perf_check=True,
run_consistency_check=True,
run_component_accuracy_check=False,
),
DiffusionTestCase( DiffusionTestCase(
"flux_image_t2i", "flux_image_t2i",
DiffusionServerArgs(model_path=DEFAULT_FLUX_1_DEV_MODEL_NAME_FOR_TEST), DiffusionServerArgs(model_path=DEFAULT_FLUX_1_DEV_MODEL_NAME_FOR_TEST),
@@ -167,6 +167,69 @@
"expected_median_denoise_ms": 190.02, "expected_median_denoise_ms": 190.02,
"estimated_full_test_time_s": 133.2 "estimated_full_test_time_s": 133.2
}, },
"ideogram4_fp8_t2i": {
"stages_ms": {
"InputValidationStage": 0.06,
"Ideogram4TextEncodingStage": 68.96,
"LatentPreparationStage": 0.17,
"Ideogram4DenoisingStage": 20813.98,
"Ideogram4DecodingStage": 95.09
},
"denoise_step_ms": {
"0": 347.26,
"1": 432.22,
"2": 436.95,
"3": 433.16,
"4": 435.54,
"5": 436.19,
"6": 434.96,
"7": 438.49,
"8": 434.57,
"9": 432.22,
"10": 433.62,
"11": 432.95,
"12": 434.85,
"13": 436.07,
"14": 435.5,
"15": 433.8,
"16": 435.74,
"17": 436.2,
"18": 435.16,
"19": 435.39,
"20": 433.49,
"21": 434.57,
"22": 434.71,
"23": 434.96,
"24": 436.35,
"25": 435.63,
"26": 435.37,
"27": 434.56,
"28": 434.36,
"29": 436.84,
"30": 437.8,
"31": 436.89,
"32": 434.56,
"33": 434.34,
"34": 436.35,
"35": 433.87,
"36": 435.89,
"37": 436.25,
"38": 435.15,
"39": 436.16,
"40": 436.25,
"41": 437.04,
"42": 435.25,
"43": 437.14,
"44": 434.85,
"45": 436.19,
"46": 436.2,
"47": 433.98
},
"expected_e2e_ms": 20982.91,
"expected_avg_denoise_ms": 433.46,
"expected_median_denoise_ms": 435.31,
"estimated_full_test_time_s": 120.0
},
"flux_image_t2i": { "flux_image_t2i": {
"stages_ms": { "stages_ms": {
"TimestepPreparationStage": 32.58, "TimestepPreparationStage": 32.58,
@@ -427,6 +427,61 @@ T2I_sampling_params = DiffusionSamplingParams(
output_size="1024x1024", output_size="1024x1024",
) )
IDEOGRAM4_CI_TEXT_PROMPT = "A cat sitting on a bench"
IDEOGRAM4_CI_PROMPT = json.dumps(
{
"high_level_description": IDEOGRAM4_CI_TEXT_PROMPT,
"style_description": {
"aesthetics": "warm, peaceful, vibrant",
"lighting": "bright afternoon sunlight, long soft shadows",
"photo": "shallow depth of field, eye-level, 85mm lens",
"medium": "photograph",
"color_palette": [
"#F5C542",
"#87CEEB",
"#4A4A4A",
"#FFFFFF",
"#2E8B57",
],
},
"compositional_deconstruction": {
"background": (
"A sunlit garden path with green hedges and a wooden bench. "
"Dappled light filters through overhead trees."
),
"elements": [
{
"type": "obj",
"bbox": [260, 260, 760, 780],
"desc": (
"A small tabby cat sitting calmly on a wooden bench, "
"looking toward the camera."
),
},
{
"type": "obj",
"bbox": [180, 580, 840, 840],
"desc": (
"A weathered wooden garden bench with soft sunlight "
"falling across the seat."
),
},
],
},
},
separators=(",", ":"),
ensure_ascii=False,
)
IDEOGRAM4_CI_sampling_params = replace(
T2I_sampling_params,
prompt=IDEOGRAM4_CI_PROMPT,
output_size="1024x1024",
output_format="png",
extras={"preset": "V4_QUALITY_48", "seed": 0},
)
MODELOPT_T2I_CI_sampling_params = DiffusionSamplingParams( MODELOPT_T2I_CI_sampling_params = DiffusionSamplingParams(
prompt="Doraemon is eating dorayaki", prompt="Doraemon is eating dorayaki",
output_size="768x768", output_size="768x768",
@@ -33,7 +33,7 @@ if TYPE_CHECKING:
logger = init_logger(__name__) logger = init_logger(__name__)
SGL_TEST_FILES_CI_DATA_REVISION = "ddaad3fca6eba761b0c9692972b8f22b7c463a4d" SGL_TEST_FILES_CI_DATA_REVISION = "50aa0d4d5d4d260302d74b80d97747efd0f0ae45"
SGL_TEST_FILES_CONSISTENCY_GT_ROOT = ( SGL_TEST_FILES_CONSISTENCY_GT_ROOT = (
"https://raw.githubusercontent.com/" "https://raw.githubusercontent.com/"
f"sgl-project/ci-data/{SGL_TEST_FILES_CI_DATA_REVISION}/" f"sgl-project/ci-data/{SGL_TEST_FILES_CI_DATA_REVISION}/"
@@ -0,0 +1,749 @@
import json
import os
import tempfile
import unittest
from types import SimpleNamespace
from unittest.mock import patch
import torch
import torch.nn.functional as F
from transformers.models.qwen3_vl.configuration_qwen3_vl import Qwen3VLTextConfig
from sglang.multimodal_gen.configs.models.dits.ideogram import Ideogram4DiTConfig
from sglang.multimodal_gen.configs.models.encoders.ideogram import (
Ideogram4TextEncoderConfig,
)
from sglang.multimodal_gen.configs.pipeline_configs.ideogram import (
Ideogram4PipelineConfig,
)
from sglang.multimodal_gen.configs.sample.ideogram import (
IDEOGRAM4_PRESETS,
Ideogram4SamplingParams,
)
from sglang.multimodal_gen.registry import _get_config_info, get_model_info
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType, get_module_role
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.layers.attention import USPAttention
from sglang.multimodal_gen.runtime.layers.quantization.weight_only_fp8 import (
FP8_WEIGHT_DTYPE,
WeightOnlyFP8Linear,
dequantize_rowwise_fp8_weight,
)
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
Qwen3VLTextRotaryEmbedding,
qwen3_apply_rotary_pos_emb,
)
from sglang.multimodal_gen.runtime.loader.component_loaders.transformer_loader import (
TransformerLoader,
_server_args_for_transformer_component,
)
from sglang.multimodal_gen.runtime.loader.fsdp_load import (
load_model_from_full_model_state_dict,
)
from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context
from sglang.multimodal_gen.runtime.models.dits.ideogram import (
Ideogram4Transformer2DModel,
)
from sglang.multimodal_gen.runtime.models.encoders.ideogram import (
IdeogramQwen3VLTextEncoder,
)
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import DenoisingStage
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ideogram import (
IMAGE_POSITION_OFFSET,
LLM_TOKEN_INDICATOR,
OUTPUT_IMAGE_INDICATOR,
Ideogram4DecodingStage,
Ideogram4DenoisingStage,
Ideogram4TextEncodingStage,
make_step_intervals,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.text_encoding import (
TextEncodingStage,
)
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
from sglang.multimodal_gen.runtime.server_args import set_global_server_args
def _reference_qwen3_mrope(position_ids, head_dim, rope_theta, mrope_section):
batch_size = position_ids.shape[0]
pos = position_ids.permute(2, 0, 1).to(dtype=torch.float32)
inv_freq = 1.0 / (
rope_theta ** (torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim)
)
inv_freq = inv_freq[None, None, :, None].expand(3, batch_size, -1, 1)
freqs = inv_freq @ pos.unsqueeze(2)
freqs = freqs.transpose(2, 3)
freqs_t = freqs[0].clone()
for axis, offset in ((1, 1), (2, 2)):
length = mrope_section[axis] * 3
idx = torch.arange(offset, length, 3, device=freqs_t.device)
freqs_t[..., idx] = freqs[axis][..., idx]
emb = torch.cat((freqs_t, freqs_t), dim=-1)
return emb.cos(), emb.sin()
class DummyTokenizer:
def apply_chat_template(self, messages, add_generation_prompt, tokenize):
return messages[0]["content"][0]["text"]
def __call__(self, text, return_tensors, add_special_tokens):
values = [int(x) for x in text.split()]
return {"input_ids": torch.tensor([values], dtype=torch.long)}
class FakeIdeogramTransformer(torch.nn.Module):
def forward(self, *, x, **kwargs):
return torch.zeros_like(x)
class FakeIdeogramVAE(torch.nn.Module):
def decode(self, z):
return z[:, :3]
class FakeIdeogramPipeline:
def __init__(self, transformer, unconditional_transformer):
self.modules = {
"transformer": transformer,
"unconditional_transformer": unconditional_transformer,
}
def _fake_server_args(cfg=None):
return SimpleNamespace(
pipeline_config=cfg or Ideogram4PipelineConfig(),
comfyui_mode=False,
enable_torch_compile=False,
attention_backend="torch_sdpa",
enable_layerwise_nvtx_marker=False,
model_loaded={"transformer": True},
model_paths={},
disable_autocast=False,
enable_cfg_parallel=False,
attention_backend_config=None,
)
def _fake_ideogram_pipeline(transformer, unconditional_transformer):
return FakeIdeogramPipeline(transformer, unconditional_transformer)
class TestIdeogram4(unittest.TestCase):
def test_registry_resolves_model_index_class_name(self):
get_model_info.cache_clear()
_get_config_info.cache_clear()
with tempfile.TemporaryDirectory() as tmpdir:
with open(f"{tmpdir}/model_index.json", "w", encoding="utf-8") as f:
json.dump(
{"_class_name": "Ideogram4Pipeline", "_diffusers_version": "0.0.0"},
f,
)
for subdir in (
"scheduler",
"text_encoder",
"tokenizer",
"transformer",
"unconditional_transformer",
"vae",
):
os.mkdir(f"{tmpdir}/{subdir}")
info = get_model_info(tmpdir, backend="sglang")
self.assertEqual(info.pipeline_cls.__name__, "Ideogram4Pipeline")
self.assertIs(info.pipeline_config_cls, Ideogram4PipelineConfig)
self.assertIs(info.sampling_param_cls, Ideogram4SamplingParams)
def test_rowwise_fp8_dequant_uses_output_channel_scale(self):
weight = torch.tensor(
[[1.0, 2.0, -3.0], [4.0, -5.0, 6.0]], dtype=FP8_WEIGHT_DTYPE
)
scale = torch.tensor([0.5, 2.0], dtype=torch.float32)
actual = dequantize_rowwise_fp8_weight(weight, scale, torch.float32)
expected = weight.to(torch.float32) * scale[:, None]
torch.testing.assert_close(actual, expected)
def test_shared_qwen3_mrope_matches_ideogram_reference_layout(self):
position_ids = torch.tensor(
[
[[0, 0, 0], [1, 1, 1], [65536, 65536, 65536]],
[[0, 0, 0], [0, 2, 3], [65536, 65537, 65538]],
],
dtype=torch.long,
)
head_dim = 8
rope_theta = 5_000_000.0
mrope_section = (2, 1, 1)
rotary_emb = Qwen3VLTextRotaryEmbedding(
head_dim=head_dim,
rope_theta=rope_theta,
mrope_section=mrope_section,
)
cos, sin = rotary_emb(torch.empty((), dtype=torch.float32), position_ids)
ref_cos, ref_sin = _reference_qwen3_mrope(
position_ids, head_dim, rope_theta, mrope_section
)
torch.testing.assert_close(cos.float(), ref_cos)
torch.testing.assert_close(sin.float(), ref_sin)
def test_usp_attention_key_mask_matches_segment_mask_for_valid_tokens(self):
import sglang.multimodal_gen.runtime.server_args as server_args_module
prev_args = server_args_module._global_server_args
try:
set_global_server_args(
SimpleNamespace(attention_backend="torch_sdpa", comfyui_mode=False)
)
torch.manual_seed(0)
batch_size, seq_len, num_heads, head_dim = 2, 5, 2, 8
q = torch.randn(batch_size, seq_len, num_heads, head_dim)
k = torch.randn(batch_size, seq_len, num_heads, head_dim)
v = torch.randn(batch_size, seq_len, num_heads, head_dim)
segment_ids = torch.tensor(
[[-1, -1, 1, 1, 1], [-1, 1, 1, 1, 1]], dtype=torch.long
)
position_ids = torch.stack(
[
torch.arange(seq_len).expand(batch_size, -1),
torch.arange(seq_len).expand(batch_size, -1) + 1,
torch.arange(seq_len).expand(batch_size, -1) + 2,
],
dim=-1,
)
rotary_emb = Qwen3VLTextRotaryEmbedding(
head_dim=head_dim, mrope_section=(2, 1, 1)
)
cos, sin = rotary_emb(q, position_ids)
q, k = qwen3_apply_rotary_pos_emb(q, k, cos.unsqueeze(2), sin.unsqueeze(2))
full_mask = (
segment_ids.unsqueeze(2) == segment_ids.unsqueeze(1)
).unsqueeze(1)
expected = F.scaled_dot_product_attention(
q.transpose(1, 2),
k.transpose(1, 2),
v.transpose(1, 2),
attn_mask=full_mask,
).transpose(1, 2)
with (
patch(
"sglang.multimodal_gen.runtime.layers.attention.layer.get_ring_parallel_world_size",
return_value=1,
),
patch(
"sglang.multimodal_gen.runtime.layers.attention.layer.get_sequence_parallel_world_size",
return_value=1,
),
):
attn = USPAttention(
num_heads=num_heads,
head_size=head_dim,
supported_attention_backends={AttentionBackendEnum.TORCH_SDPA},
)
with set_forward_context(current_timestep=0, attn_metadata=None):
actual = attn(q, k, v, attn_mask=segment_ids > 0)
valid = segment_ids > 0
torch.testing.assert_close(actual[valid], expected[valid])
finally:
set_global_server_args(prev_args)
def test_ideogram_preset_guidance_order(self):
turbo = IDEOGRAM4_PRESETS["V4_TURBO_12"]
default = IDEOGRAM4_PRESETS["V4_DEFAULT_20"]
self.assertEqual(turbo["num_steps"], 12)
self.assertEqual(default["num_steps"], 20)
self.assertEqual(turbo["guidance_schedule"][0], 3.0)
self.assertEqual(turbo["guidance_schedule"][-1], 7.0)
self.assertEqual(tuple(make_step_intervals(2).tolist()), (0.0, 0.5, 1.0))
def test_ideogram_sampling_params_sync_steps_with_preset(self):
params = Ideogram4SamplingParams(preset="V4_TURBO_12")
self.assertEqual(params.num_inference_steps, 12)
self.assertEqual(params.guidance_scale, 7.0)
same_steps = Ideogram4SamplingParams(
preset="V4_TURBO_12", num_inference_steps=12
)
self.assertEqual(same_steps.num_inference_steps, 12)
with self.assertRaisesRegex(ValueError, "derives num_inference_steps"):
Ideogram4SamplingParams(preset="V4_TURBO_12", num_inference_steps=20)
same_guidance = Ideogram4SamplingParams(
preset="V4_TURBO_12", guidance_scale=7.0
)
self.assertEqual(same_guidance.guidance_scale, 7.0)
with self.assertRaisesRegex(ValueError, "guidance_scale cannot be set"):
Ideogram4SamplingParams(preset="V4_TURBO_12", guidance_scale=6.0)
with self.assertRaisesRegex(ValueError, "Unknown Ideogram 4 preset"):
Ideogram4SamplingParams(preset="V4_FAST")
def test_ideogram_sampling_params_merge_recomputes_preset_fields(self):
target = Ideogram4SamplingParams()
user = Ideogram4SamplingParams(
preset="V4_TURBO_12",
height=256,
width=256,
)
target._merge_with_user_params(
user, explicit_fields={"preset", "height", "width"}
)
self.assertEqual(target.preset, "V4_TURBO_12")
self.assertEqual(target.num_inference_steps, 12)
self.assertEqual(target.guidance_scale, 7.0)
self.assertEqual(target.height, 256)
self.assertEqual(target.width, 256)
def test_unconditional_transformer_uses_denoiser_loader_path(self):
self.assertIn("unconditional_transformer", TransformerLoader.component_names)
self.assertEqual(
get_module_role("unconditional_transformer"), RoleType.DENOISER
)
server_args = SimpleNamespace(
transformer_weights_path="/unused/override.safetensors",
nunchaku_config={"enabled": True},
)
component_args = _server_args_for_transformer_component(
server_args, "unconditional_transformer"
)
self.assertIsNot(component_args, server_args)
self.assertIsNone(component_args.transformer_weights_path)
self.assertIsNone(component_args.nunchaku_config)
def test_ideogram_denoiser_does_not_request_dtype_cast(self):
import sglang.multimodal_gen.runtime.server_args as server_args_module
prev_args = server_args_module._global_server_args
try:
set_global_server_args(_fake_server_args())
transformer = FakeIdeogramTransformer()
unconditional_transformer = FakeIdeogramTransformer()
stage = Ideogram4DenoisingStage(
transformer=transformer,
unconditional_transformer=unconditional_transformer,
pipeline=_fake_ideogram_pipeline(
transformer, unconditional_transformer
),
)
uses = stage.component_uses(_fake_server_args(), "stage")
finally:
set_global_server_args(prev_args)
self.assertEqual(
[use.component_name for use in uses],
[
"transformer",
"unconditional_transformer",
],
)
self.assertTrue(all(use.target_dtype is None for use in uses))
def test_ideogram_stages_inherit_common_stage_bases(self):
import sglang.multimodal_gen.runtime.server_args as server_args_module
prev_args = server_args_module._global_server_args
try:
set_global_server_args(_fake_server_args())
text_stage = Ideogram4TextEncodingStage(
text_encoder=None, tokenizer=DummyTokenizer()
)
denoising_stage = Ideogram4DenoisingStage(
transformer=FakeIdeogramTransformer(),
unconditional_transformer=FakeIdeogramTransformer(),
)
finally:
set_global_server_args(prev_args)
self.assertIsInstance(text_stage, TextEncodingStage)
self.assertIsInstance(denoising_stage, DenoisingStage)
def test_ideogram_text_encoding_dedup_fingerprint_and_extra_copy(self):
import sglang.multimodal_gen.runtime.server_args as server_args_module
cfg = Ideogram4PipelineConfig()
args = _fake_server_args(cfg)
prev_args = server_args_module._global_server_args
try:
set_global_server_args(args)
stage = Ideogram4TextEncodingStage(
text_encoder=None, tokenizer=DummyTokenizer()
)
finally:
set_global_server_args(prev_args)
base = Req(
sampling_params=Ideogram4SamplingParams(
prompt="11 12",
height=256,
width=256,
num_outputs_per_prompt=1,
)
)
same = Req(
sampling_params=Ideogram4SamplingParams(
prompt="11 12",
height=256,
width=256,
num_outputs_per_prompt=1,
)
)
different_height = Req(
sampling_params=Ideogram4SamplingParams(
prompt="11 12",
height=512,
width=256,
num_outputs_per_prompt=1,
)
)
different_width = Req(
sampling_params=Ideogram4SamplingParams(
prompt="11 12",
height=256,
width=512,
num_outputs_per_prompt=1,
)
)
different_outputs = Req(
sampling_params=Ideogram4SamplingParams(
prompt="11 12",
height=256,
width=256,
num_outputs_per_prompt=2,
)
)
base_fingerprint = stage.build_dedup_fingerprint(base, args)
self.assertEqual(base_fingerprint, stage.build_dedup_fingerprint(same, args))
self.assertNotEqual(
base_fingerprint, stage.build_dedup_fingerprint(different_height, args)
)
self.assertNotEqual(
base_fingerprint, stage.build_dedup_fingerprint(different_width, args)
)
self.assertNotEqual(
base_fingerprint, stage.build_dedup_fingerprint(different_outputs, args)
)
base.prompt_embeds = [torch.tensor([1.0])]
base.prompt_embeds_mask = [torch.tensor([True])]
base.extra["ideogram4"] = {
"position_ids": torch.tensor([[1]]),
"metadata": {"grid_h": 16},
}
stage.copy_deduplicated_outputs(base, same)
self.assertIn("ideogram4", same.extra)
self.assertTrue(
torch.equal(
same.extra["ideogram4"]["position_ids"],
base.extra["ideogram4"]["position_ids"],
)
)
self.assertIsNot(
same.extra["ideogram4"]["position_ids"],
base.extra["ideogram4"]["position_ids"],
)
def test_ideogram_text_encoding_verifies_custom_outputs(self):
import sglang.multimodal_gen.runtime.server_args as server_args_module
args = _fake_server_args()
prev_args = server_args_module._global_server_args
try:
set_global_server_args(args)
stage = Ideogram4TextEncodingStage(
text_encoder=None, tokenizer=DummyTokenizer()
)
finally:
set_global_server_args(prev_args)
batch = Req(
sampling_params=Ideogram4SamplingParams(
prompt="11 12",
height=256,
width=256,
num_outputs_per_prompt=1,
)
)
self.assertTrue(stage.verify_input(batch, args).is_valid())
batch.do_classifier_free_guidance = True
batch.negative_prompt = []
batch.negative_prompt_embeds = []
batch.prompt_embeds = [torch.zeros(1, 4, 8)]
batch.prompt_embeds_mask = [torch.ones(1, 4, dtype=torch.bool)]
batch.extra["ideogram4"] = {"position_ids": torch.zeros(1, 4, 3)}
self.assertTrue(stage.verify_output(batch, args).is_valid())
def test_ideogram_denoising_component_names_from_pipeline_modules(self):
import sglang.multimodal_gen.runtime.server_args as server_args_module
transformer = FakeIdeogramTransformer()
unconditional_transformer = FakeIdeogramTransformer()
prev_args = server_args_module._global_server_args
try:
set_global_server_args(_fake_server_args())
stage = Ideogram4DenoisingStage(
transformer=transformer,
unconditional_transformer=unconditional_transformer,
pipeline=_fake_ideogram_pipeline(
transformer, unconditional_transformer
),
)
uses = stage.component_uses(_fake_server_args(), "stage")
finally:
set_global_server_args(prev_args)
self.assertEqual(
[use.component_name for use in uses],
["transformer", "unconditional_transformer"],
)
def test_ideogram_attention_backend_is_passed_from_config(self):
import sglang.multimodal_gen.runtime.server_args as server_args_module
config = Ideogram4DiTConfig()
self.assertEqual(
config.arch_config._supported_attention_backends,
{AttentionBackendEnum.FA, AttentionBackendEnum.TORCH_SDPA},
)
prev_args = server_args_module._global_server_args
try:
set_global_server_args(_fake_server_args())
with patch(
"sglang.multimodal_gen.runtime.layers.attention.layer.get_ring_parallel_world_size",
return_value=1,
):
with torch.device("meta"):
model = Ideogram4Transformer2DModel(config, {})
finally:
set_global_server_args(prev_args)
self.assertEqual(
model.supported_attention_backends,
config.arch_config._supported_attention_backends,
)
self.assertEqual(
model.layers[0].attention.attn.backend,
AttentionBackendEnum.TORCH_SDPA,
)
def test_ideogram_dit_meta_state_dict_matches_checkpoint_shapes(self):
import sglang.multimodal_gen.runtime.server_args as server_args_module
prev_args = server_args_module._global_server_args
try:
set_global_server_args(
SimpleNamespace(attention_backend="torch_sdpa", comfyui_mode=False)
)
with patch(
"sglang.multimodal_gen.runtime.layers.attention.layer.get_ring_parallel_world_size",
return_value=1,
):
with torch.device("meta"):
model = Ideogram4Transformer2DModel(Ideogram4DiTConfig(), {})
finally:
set_global_server_args(prev_args)
state = model.state_dict()
self.assertEqual(len(state), 669)
self.assertEqual(tuple(state["input_proj.weight"].shape), (4608, 128))
self.assertEqual(tuple(state["input_proj.weight_scale"].shape), (4608,))
self.assertEqual(
tuple(state["layers.0.attention.qkv.weight"].shape), (13824, 4608)
)
self.assertEqual(state["layers.0.attention.qkv.weight"].dtype, FP8_WEIGHT_DTYPE)
def test_missing_weight_only_fp8_scale_is_fatal(self):
with torch.device("meta"):
model = WeightOnlyFP8Linear(3, 2, bias=False)
weights = iter(
[
(
"weight",
torch.tensor(
[[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]],
dtype=FP8_WEIGHT_DTYPE,
),
)
]
)
with self.assertRaisesRegex(ValueError, "Required checkpoint parameter"):
load_model_from_full_model_state_dict(
model,
weights,
torch.device("cpu"),
param_dtype=None,
strict=False,
param_names_mapping=lambda name: (name, None, None),
)
def test_weight_only_fp8_load_accepts_explicit_scale(self):
with torch.device("meta"):
model = WeightOnlyFP8Linear(3, 2, bias=False)
weights = iter(
[
(
"weight",
torch.tensor(
[[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]],
dtype=FP8_WEIGHT_DTYPE,
),
),
("weight_scale", torch.tensor([0.5, 2.0], dtype=torch.float32)),
]
)
load_model_from_full_model_state_dict(
model,
weights,
torch.device("cpu"),
param_dtype=None,
strict=False,
param_names_mapping=lambda name: (name, None, None),
)
self.assertEqual(model.weight.dtype, FP8_WEIGHT_DTYPE)
self.assertEqual(model.weight_scale.dtype, torch.float32)
def test_ideogram_text_encoder_post_config_hook_preserves_local_arch(self):
config = Ideogram4TextEncoderConfig()
config.arch_config.architectures = ["RemoteQwen3VLTextModel"]
config.arch_config.ideogram_fp8_weight_only = False
config.post_diffusers_config_update()
self.assertEqual(
config.arch_config.architectures, ["IdeogramQwen3VLTextEncoder"]
)
self.assertTrue(config.arch_config.ideogram_fp8_weight_only)
def test_ideogram_text_encoder_swaps_linears_to_weight_only_fp8(self):
config = Ideogram4TextEncoderConfig()
config.post_diffusers_config_update()
config.arch_config.text_config = Qwen3VLTextConfig(
vocab_size=32,
hidden_size=16,
intermediate_size=32,
num_hidden_layers=1,
num_attention_heads=2,
num_key_value_heads=2,
head_dim=8,
max_position_embeddings=64,
pad_token_id=0,
)
import sglang.multimodal_gen.runtime.server_args as server_args_module
prev_args = server_args_module._global_server_args
try:
set_global_server_args(
SimpleNamespace(attention_backend="torch_sdpa", comfyui_mode=False)
)
with torch.device("meta"):
encoder = IdeogramQwen3VLTextEncoder(config)
finally:
set_global_server_args(prev_args)
self.assertTrue(
any(isinstance(module, WeightOnlyFP8Linear) for module in encoder.modules())
)
self.assertFalse(
any(isinstance(module, torch.nn.Linear) for module in encoder.modules())
)
def test_denoise_and_decode_shape_smoke(self):
import sglang.multimodal_gen.runtime.server_args as server_args_module
cfg = Ideogram4PipelineConfig()
args = _fake_server_args(cfg)
device = get_local_torch_device()
prev_args = server_args_module._global_server_args
try:
set_global_server_args(args)
transformer = FakeIdeogramTransformer()
unconditional_transformer = FakeIdeogramTransformer()
denoise_stage = Ideogram4DenoisingStage(
transformer=transformer,
unconditional_transformer=unconditional_transformer,
pipeline=_fake_ideogram_pipeline(
transformer, unconditional_transformer
),
)
decode_stage = Ideogram4DecodingStage(vae=FakeIdeogramVAE())
batch = Req(
sampling_params=Ideogram4SamplingParams(
prompt="11 12",
height=256,
width=256,
preset="V4_TURBO_12",
suppress_logs=True,
)
)
batch.latents = torch.zeros(1, 1, 128, device=device)
batch.raw_latent_shape = batch.latents.shape
batch.prompt_embeds = [torch.zeros(1, 2, 8, device=device)]
batch.extra["ideogram4"] = {
"max_text_tokens": 1,
"num_image_tokens": 1,
"position_ids": torch.zeros(1, 2, 3, dtype=torch.long, device=device),
"segment_ids": torch.ones(1, 2, dtype=torch.long, device=device),
"indicator": torch.tensor(
[[LLM_TOKEN_INDICATOR, OUTPUT_IMAGE_INDICATOR]],
dtype=torch.long,
device=device,
),
"grid_h": 1,
"grid_w": 1,
}
denoised = denoise_stage.forward(batch, args)
self.assertEqual(tuple(denoised.latents.shape), (1, 1, 128))
decoded = decode_stage.forward(denoised, args)
finally:
set_global_server_args(prev_args)
self.assertEqual(tuple(decoded.output.shape), (1, 3, 2, 2))
def test_text_input_builder_matches_official_layout(self):
prev_args = None
import sglang.multimodal_gen.runtime.server_args as server_args_module
prev_args = server_args_module._global_server_args
try:
cfg = Ideogram4PipelineConfig()
args = SimpleNamespace(pipeline_config=cfg, comfyui_mode=False)
set_global_server_args(args)
stage = Ideogram4TextEncodingStage(
text_encoder=None, tokenizer=DummyTokenizer()
)
inputs = stage._build_inputs(["11 12 13", "21"], 256, 256, args)
finally:
set_global_server_args(prev_args)
self.assertEqual(inputs["grid_h"], 16)
self.assertEqual(inputs["grid_w"], 16)
self.assertEqual(inputs["num_image_tokens"], 256)
self.assertEqual(inputs["max_text_tokens"], 3)
self.assertEqual(inputs["token_ids"][0, :3].tolist(), [11, 12, 13])
self.assertEqual(inputs["token_ids"][1, :2].tolist(), [0, 0])
self.assertTrue(
torch.all(inputs["indicator"][0, :3] == LLM_TOKEN_INDICATOR).item()
)
self.assertTrue(
torch.all(inputs["indicator"][0, 3:] == OUTPUT_IMAGE_INDICATOR).item()
)
self.assertEqual(inputs["position_ids"][0, 3, 0].item(), IMAGE_POSITION_OFFSET)
def test_text_input_builder_rejects_unsupported_resolution(self):
import sglang.multimodal_gen.runtime.server_args as server_args_module
prev_args = server_args_module._global_server_args
try:
cfg = Ideogram4PipelineConfig()
args = SimpleNamespace(pipeline_config=cfg, comfyui_mode=False)
set_global_server_args(args)
stage = Ideogram4TextEncodingStage(
text_encoder=None, tokenizer=DummyTokenizer()
)
with self.assertRaisesRegex(ValueError, "between 256 and 2048"):
stage._build_inputs(["11"], 128, 256, args)
finally:
set_global_server_args(prev_args)
if __name__ == "__main__":
unittest.main()