diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py index 5646b79f5..a1fa66e60 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py @@ -762,6 +762,49 @@ class ImagePipelineConfig(PipelineConfig): return latents, batch_size, channels, height, width +@dataclass +class SpatialImagePipelineConfig(ImagePipelineConfig): + """Base config for spatial image pipelines (e.g. GLM-Image) with 4D latents (B, C, H', W'). + + Overrides shard_latents_for_sp / gather_latents_for_sp to shard along the height dimension + so that each SP rank gets (B, C, H'_local, W') instead of using the token-style (B, S, C) path. + """ + + def shard_latents_for_sp(self, batch, latents): + # 4D latents (B, C, H', W') -> shard along H' (dim=2); otherwise fall back to base (B, S, C) + sp_world_size = get_sp_world_size() + if sp_world_size <= 1: + return latents, False + if latents.dim() != 4: + return super().shard_latents_for_sp(batch, latents) + + # (B, C, H', W') + _, _, h_lat, w_lat = latents.shape + if h_lat % sp_world_size != 0: + pad_len = sp_world_size - (h_lat % sp_world_size) + pad = torch.zeros( + (latents.shape[0], latents.shape[1], pad_len, latents.shape[3]), + dtype=latents.dtype, + device=latents.device, + ) + latents = torch.cat([latents, pad], dim=2) + h_lat = latents.shape[2] + rank_in_sp_group = get_sp_parallel_rank() + chunk_size = h_lat // sp_world_size + h0 = rank_in_sp_group * chunk_size + h1 = h0 + chunk_size + sharded = latents[:, :, h0:h1, :].contiguous() + return sharded, True + + def gather_latents_for_sp(self, latents): + if get_sp_world_size() <= 1: + return latents + if latents.dim() != 4: + return super().gather_latents_for_sp(latents) + # Gather along dim=2 (H') to match shard_latents_for_sp + return sequence_model_parallel_all_gather(latents, dim=2) + + @dataclass class SlidingTileAttnConfig(PipelineConfig): """Configuration for sliding tile attention.""" diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/glm_image.py b/python/sglang/multimodal_gen/configs/pipeline_configs/glm_image.py index 95ba40333..2801eeb37 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/glm_image.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/glm_image.py @@ -9,13 +9,13 @@ from sglang.multimodal_gen.configs.models.encoders.base import EncoderConfig from sglang.multimodal_gen.configs.models.encoders.t5 import T5Config from sglang.multimodal_gen.configs.models.vaes.glmimage import GlmImageVAEConfig from sglang.multimodal_gen.configs.pipeline_configs.base import ( - ImagePipelineConfig, ModelTaskType, + SpatialImagePipelineConfig, ) @dataclass -class GlmImagePipelineConfig(ImagePipelineConfig): +class GlmImagePipelineConfig(SpatialImagePipelineConfig): """Configuration for the GlmImage pipeline.""" vae_precision: str = "bf16" @@ -35,6 +35,12 @@ class GlmImagePipelineConfig(ImagePipelineConfig): # VAE vae_config: VAEConfig = field(default_factory=GlmImageVAEConfig) + # GLM-Image uses T5 text encoder; base default is EncoderConfig() which lacks + # parallel_folding and causes AttributeError + fallback to native T5 with missing weights. + text_encoder_configs: tuple[EncoderConfig, ...] = field( + default_factory=lambda: (T5Config(),) + ) + enable_autocast: bool = False def __post_init__(self): diff --git a/python/sglang/multimodal_gen/runtime/models/dits/glm_image.py b/python/sglang/multimodal_gen/runtime/models/dits/glm_image.py index 1bae8c855..e6ab13a1c 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/glm_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/glm_image.py @@ -20,6 +20,10 @@ import torch.nn.functional as F from diffusers.models.attention import FeedForward from sglang.multimodal_gen.configs.models.dits.glmimage import GlmImageDitConfig +from sglang.multimodal_gen.runtime.distributed.parallel_state import ( + get_sp_parallel_rank, + get_sp_world_size, +) from sglang.multimodal_gen.runtime.layers.attention import USPAttention from sglang.multimodal_gen.runtime.layers.layernorm import ( ScaleResidualLayerNormScaleShift, @@ -806,6 +810,18 @@ class GlmImageTransformer2DModel(CachableDiT, OffloadableDiTMixin): prior_embedding = self.prior_token_embedding(prior_token_id) prior_embedding[prior_token_drop] *= 0.0 prior_hidden_states = self.prior_projector(prior_embedding) + # SP: when latents are H-sharded, hidden_states has fewer patches than prior_hidden_states. + # Shard prior_hidden_states along seq dim to match (prior is row-major, same as latent patches). + if ( + get_sp_world_size() > 1 + and prior_hidden_states.shape[1] != hidden_states.shape[1] + ): + rank = get_sp_parallel_rank() + sp_world_size = get_sp_world_size() + chunk = prior_hidden_states.shape[1] // sp_world_size + prior_hidden_states = prior_hidden_states[ + :, rank * chunk : (rank + 1) * chunk, : + ] hidden_states = hidden_states + prior_hidden_states temb = self.time_condition_embed(