[diffusion] optimize: precompute ideogram4 denoising metadata (#27443)

This commit is contained in:
Xiaoyu Zhang
2026-06-07 14:27:41 +08:00
committed by GitHub
parent 0ce3db3c0a
commit ff8b97406d
2 changed files with 26 additions and 10 deletions
@@ -393,6 +393,8 @@ class Ideogram4Transformer2DModel(BaseDiT):
position_ids: torch.Tensor, position_ids: torch.Tensor,
segment_ids: torch.Tensor, segment_ids: torch.Tensor,
indicator: torch.Tensor, indicator: torch.Tensor,
attn_mask: torch.Tensor | None = None,
attn_mask_meta: dict | None = None,
**kwargs, **kwargs,
) -> torch.Tensor: ) -> torch.Tensor:
param_dtype = self.embed_image_indicator.weight.dtype param_dtype = self.embed_image_indicator.weight.dtype
@@ -421,7 +423,9 @@ class Ideogram4Transformer2DModel(BaseDiT):
cos = cos.unsqueeze(2) cos = cos.unsqueeze(2)
sin = sin.unsqueeze(2) sin = sin.unsqueeze(2)
# ideogram uses -1 padding; varlen meta enables fa packed attention # ideogram uses -1 padding; varlen meta enables fa packed attention
if attn_mask is None:
attn_mask = segment_ids > 0 attn_mask = segment_ids > 0
if attn_mask_meta is None:
attn_mask_meta = build_varlen_mask_meta(attn_mask) attn_mask_meta = build_varlen_mask_meta(attn_mask)
for layer in self.layers: for layer in self.layers:
h = layer( h = layer(
@@ -11,6 +11,7 @@ from sglang.multimodal_gen.configs.pipeline_configs.ideogram import (
) )
from sglang.multimodal_gen.configs.sample.ideogram import IDEOGRAM4_PRESETS from sglang.multimodal_gen.configs.sample.ideogram import IDEOGRAM4_PRESETS
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.layers.attention import build_varlen_mask_meta
from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import ( from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
ComponentUse, ComponentUse,
@@ -357,6 +358,8 @@ class Ideogram4DenoisingStage(DenoisingStage):
guidance_schedule = torch.as_tensor( guidance_schedule = torch.as_tensor(
preset_cfg["guidance_schedule"], dtype=torch.float32, device=device preset_cfg["guidance_schedule"], dtype=torch.float32, device=device
) )
schedule_values = schedule(step_intervals)
schedule_deltas = schedule_values[:-1] - schedule_values[1:]
self.scheduler.set_timesteps(num_steps, device=device) self.scheduler.set_timesteps(num_steps, device=device)
batch.scheduler = self.scheduler batch.scheduler = self.scheduler
@@ -385,6 +388,8 @@ class Ideogram4DenoisingStage(DenoisingStage):
neg_position_ids = data["position_ids"][:, max_text_tokens:] neg_position_ids = data["position_ids"][:, max_text_tokens:]
neg_segment_ids = data["segment_ids"][:, max_text_tokens:] neg_segment_ids = data["segment_ids"][:, max_text_tokens:]
neg_indicator = data["indicator"][:, max_text_tokens:] neg_indicator = data["indicator"][:, max_text_tokens:]
attn_mask = data["segment_ids"] > 0
neg_attn_mask = neg_segment_ids > 0
neg_llm_features = torch.zeros( neg_llm_features = torch.zeros(
batch_size, batch_size,
num_image_tokens, num_image_tokens,
@@ -395,13 +400,17 @@ class Ideogram4DenoisingStage(DenoisingStage):
ctx.latents = z ctx.latents = z
ctx.extra.update( ctx.extra.update(
{ {
"ideogram4_schedule": schedule, "ideogram4_schedule_values": schedule_values,
"ideogram4_step_intervals": step_intervals, "ideogram4_schedule_deltas": schedule_deltas,
"ideogram4_guidance_schedule": guidance_schedule, "ideogram4_guidance_schedule": guidance_schedule,
"ideogram4_text_z_padding": text_z_padding, "ideogram4_text_z_padding": text_z_padding,
"ideogram4_attn_mask": attn_mask,
"ideogram4_attn_mask_meta": build_varlen_mask_meta(attn_mask),
"ideogram4_neg_position_ids": neg_position_ids, "ideogram4_neg_position_ids": neg_position_ids,
"ideogram4_neg_segment_ids": neg_segment_ids, "ideogram4_neg_segment_ids": neg_segment_ids,
"ideogram4_neg_indicator": neg_indicator, "ideogram4_neg_indicator": neg_indicator,
"ideogram4_neg_attn_mask": neg_attn_mask,
"ideogram4_neg_attn_mask_meta": build_varlen_mask_meta(neg_attn_mask),
"ideogram4_neg_llm_features": neg_llm_features, "ideogram4_neg_llm_features": neg_llm_features,
} }
) )
@@ -418,14 +427,13 @@ class Ideogram4DenoisingStage(DenoisingStage):
z = ctx.latents.to(dtype=torch.float32) z = ctx.latents.to(dtype=torch.float32)
llm_features = batch.prompt_embeds[0] llm_features = batch.prompt_embeds[0]
max_text_tokens = data["max_text_tokens"] max_text_tokens = data["max_text_tokens"]
schedule = ctx.extra["ideogram4_schedule"] schedule_values = ctx.extra["ideogram4_schedule_values"]
step_intervals = ctx.extra["ideogram4_step_intervals"] schedule_deltas = ctx.extra["ideogram4_schedule_deltas"]
guidance_schedule = ctx.extra["ideogram4_guidance_schedule"] guidance_schedule = ctx.extra["ideogram4_guidance_schedule"]
i = step.t_int i = step.t_int
t_val = float(schedule(step_intervals[i + 1].unsqueeze(0)).item()) t_val = schedule_values[i + 1]
s_val = float(schedule(step_intervals[i].unsqueeze(0)).item()) t = t_val.expand(z.shape[0])
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) pos_z = torch.cat([ctx.extra["ideogram4_text_z_padding"], z], dim=1)
use_nvtx = self.current_use_nvtx use_nvtx = self.current_use_nvtx
@@ -442,6 +450,8 @@ class Ideogram4DenoisingStage(DenoisingStage):
position_ids=data["position_ids"], position_ids=data["position_ids"],
segment_ids=data["segment_ids"], segment_ids=data["segment_ids"],
indicator=data["indicator"], indicator=data["indicator"],
attn_mask=ctx.extra["ideogram4_attn_mask"],
attn_mask_meta=ctx.extra["ideogram4_attn_mask_meta"],
) )
pos_v = pos_out[:, max_text_tokens:] pos_v = pos_out[:, max_text_tokens:]
@@ -458,13 +468,15 @@ class Ideogram4DenoisingStage(DenoisingStage):
position_ids=ctx.extra["ideogram4_neg_position_ids"], position_ids=ctx.extra["ideogram4_neg_position_ids"],
segment_ids=ctx.extra["ideogram4_neg_segment_ids"], segment_ids=ctx.extra["ideogram4_neg_segment_ids"],
indicator=ctx.extra["ideogram4_neg_indicator"], indicator=ctx.extra["ideogram4_neg_indicator"],
attn_mask=ctx.extra["ideogram4_neg_attn_mask"],
attn_mask_meta=ctx.extra["ideogram4_neg_attn_mask_meta"],
) )
with maybe_nvtx_range("scheduler_step", use_nvtx): with maybe_nvtx_range("scheduler_step", use_nvtx):
velocity = ( velocity = (
guidance_schedule[i] * pos_v + (1.0 - guidance_schedule[i]) * neg_v guidance_schedule[i] * pos_v + (1.0 - guidance_schedule[i]) * neg_v
) )
ctx.latents = z + velocity * (s_val - t_val) ctx.latents = z + velocity * schedule_deltas[i]
class Ideogram4DecodingStage(PipelineStage): class Ideogram4DecodingStage(PipelineStage):