From 2b1d3c935e5b811681e04482e75f6b7f740ab6f8 Mon Sep 17 00:00:00 2001 From: Ratish P <114130421+Ratish1@users.noreply.github.com> Date: Tue, 24 Mar 2026 07:45:33 +0530 Subject: [PATCH] [diffusion] fix Z-Image SP sharding for portrait and padded resolutions (#21042) Co-authored-by: Xiaoyu Zhang <35585791+BBuf@users.noreply.github.com> --- .../configs/pipeline_configs/base.py | 8 +- .../configs/pipeline_configs/ltx_2.py | 4 +- .../configs/pipeline_configs/zimage.py | 178 +++++++++++------- .../runtime/models/dits/zimage.py | 23 ++- .../pipelines_core/stages/denoising.py | 28 ++- .../test/server/testcase_configs.py | 14 ++ 6 files changed, 175 insertions(+), 80 deletions(-) diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py index 7303f6994..80b7758b8 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py @@ -359,7 +359,7 @@ class PipelineConfig: def preprocess_decoding(self, latents, server_args=None, vae=None): return latents - def gather_latents_for_sp(self, latents): + def gather_latents_for_sp(self, latents, batch=None): # For video latents [B, C, T_local, H, W], gather along time dim=2 latents = sequence_model_parallel_all_gather(latents, dim=2) return latents @@ -808,7 +808,7 @@ class ImagePipelineConfig(PipelineConfig): sharded_tensor = sharded_tensor[:, rank_in_sp_group, :, :] return sharded_tensor, True - def gather_latents_for_sp(self, latents): + def gather_latents_for_sp(self, latents, batch=None): # For image latents [B, S_local, D], gather along sequence dim=1 latents = sequence_model_parallel_all_gather(latents, dim=1) return latents @@ -862,11 +862,11 @@ class SpatialImagePipelineConfig(ImagePipelineConfig): sharded = latents[:, :, h0:h1, :].contiguous() return sharded, True - def gather_latents_for_sp(self, latents): + def gather_latents_for_sp(self, latents, batch=None): if get_sp_world_size() <= 1: return latents if latents.dim() != 4: - return super().gather_latents_for_sp(latents) + return super().gather_latents_for_sp(latents, batch=batch) # Gather along dim=2 (H') to match shard_latents_for_sp return sequence_model_parallel_all_gather(latents, dim=2) diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/ltx_2.py b/python/sglang/multimodal_gen/configs/pipeline_configs/ltx_2.py index b3438eab8..301423753 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/ltx_2.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/ltx_2.py @@ -350,13 +350,13 @@ class LTX2PipelineConfig(PipelineConfig): return latents, True - def gather_latents_for_sp(self, latents): + def gather_latents_for_sp(self, latents, batch=None): """Gather latents after SP. For packed token latents [B, S_local, D], gather on dim=1.""" if get_sp_world_size() <= 1: return latents if isinstance(latents, torch.Tensor) and latents.ndim == 3: return sequence_model_parallel_all_gather(latents.contiguous(), dim=1) - return super().gather_latents_for_sp(latents) + return super().gather_latents_for_sp(latents, batch=batch) def maybe_pack_audio_latents(self, latents, batch_size, batch): # If already packed (3D shape [B, T, C*F]), skip packing diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/zimage.py b/python/sglang/multimodal_gen/configs/pipeline_configs/zimage.py index 3bdbf6059..de08a27ef 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/zimage.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/zimage.py @@ -4,6 +4,7 @@ from dataclasses import dataclass, field from typing import Callable import torch +import torch.distributed as dist from sglang.multimodal_gen.configs.models import DiTConfig, EncoderConfig, VAEConfig from sglang.multimodal_gen.configs.models.dits.zimage import ZImageDitConfig @@ -14,10 +15,8 @@ from sglang.multimodal_gen.configs.pipeline_configs.base import ( ImagePipelineConfig, ModelTaskType, ) -from sglang.multimodal_gen.runtime.distributed.communication_op import ( - sequence_model_parallel_all_gather, -) from sglang.multimodal_gen.runtime.distributed.parallel_state import ( + get_sp_group, get_sp_parallel_rank, get_sp_world_size, ) @@ -86,8 +85,13 @@ class ZImagePipelineConfig(ImagePipelineConfig): return x return int(math.ceil(x / m) * m) + @staticmethod + def _split_evenly(total: int, parts: int) -> list[int]: + base, remainder = divmod(total, parts) + return [base + int(rank < remainder) for rank in range(parts)] + def _build_zimage_sp_plan(self, batch) -> dict: - """Build a minimal SP plan on batch for zimage spatial sharding.""" + """Build an SP plan that preserves native spatial layout for Z-Image.""" sp_size = get_sp_world_size() rank = get_sp_parallel_rank() @@ -103,32 +107,62 @@ class ZImagePipelineConfig(ImagePipelineConfig): batch.width // self.vae_config.arch_config.spatial_compression_ratio ) - # Rule: shard along the larger spatial dimension (W/H), implemented via optional H/W transpose. - # Choose the larger of H and W for sharding, so H_eff = max(H, W). - swap_hw = W > H - H_eff = W if swap_hw else H - W_eff = H if swap_hw else W + # ZImage patchifies [C, F, H, W] latents in native F/H/W order, so shard + # native H or W directly. + H_tok = H // self.PATCH_SIZE + W_tok = W // self.PATCH_SIZE - # ZImage uses PATCH_SIZE=2 for spatial patchify; shard in token space and convert back to latent rows. - H_tok = H_eff // self.PATCH_SIZE - W_tok = W_eff // self.PATCH_SIZE - H_tok_pad = self._ceil_to_multiple(H_tok, sp_size) - H_tok_local = H_tok_pad // sp_size - h0_tok = rank * H_tok_local + shard_options = [] + for shard_axis, axis_tok, other_tok, tie_break in ( + ("h", H_tok, W_tok, 0), + ("w", W_tok, H_tok, 1), + ): + axis_sizes = self._split_evenly(axis_tok, sp_size) + local_seq_lens = [axis_size * other_tok for axis_size in axis_sizes] + img_seq_target = self._ceil_to_multiple( + max(local_seq_lens), self.SEQ_LEN_MULTIPLE + ) + total_pad_tokens = img_seq_target * sp_size - (H_tok * W_tok) + shard_options.append( + ( + total_pad_tokens, + -axis_tok, + tie_break, + shard_axis, + axis_sizes, + img_seq_target, + ) + ) + + _, _, _, shard_axis, axis_sizes, img_seq_target = min(shard_options) + axis_start_tok = sum(axis_sizes[:rank]) + axis_local_tok = axis_sizes[rank] + + if shard_axis == "h": + h0_tok = axis_start_tok + w0_tok = 0 + local_h_tok = axis_local_tok + local_w_tok = W_tok + else: + h0_tok = 0 + w0_tok = axis_start_tok + local_h_tok = H_tok + local_w_tok = axis_local_tok plan = { "sp_size": sp_size, "rank": rank, - "swap_hw": swap_hw, "H": H, "W": W, - "H_eff": H_eff, - "W_eff": W_eff, "H_tok": H_tok, "W_tok": W_tok, - "H_tok_pad": H_tok_pad, - "H_tok_local": H_tok_local, + "shard_axis": shard_axis, + "shard_sizes_tok": axis_sizes, "h0_tok": h0_tok, + "w0_tok": w0_tok, + "local_h_tok": local_h_tok, + "local_w_tok": local_w_tok, + "img_seq_target": img_seq_target, } batch._zimage_sp_plan = plan return plan @@ -154,51 +188,55 @@ class ZImagePipelineConfig(ImagePipelineConfig): return latents, False plan = self._get_zimage_sp_plan(batch) + if plan["shard_axis"] == "h": + h0 = plan["h0_tok"] * self.PATCH_SIZE + h1 = (plan["h0_tok"] + plan["local_h_tok"]) * self.PATCH_SIZE + return latents[:, :, :, h0:h1, :].contiguous(), True - # Layout: [B, C, T, H, W]. Always shard on dim=3 by optionally swapping H/W. - if plan["swap_hw"]: - latents = latents.transpose(3, 4).contiguous() + w0 = plan["w0_tok"] * self.PATCH_SIZE + w1 = (plan["w0_tok"] + plan["local_w_tok"]) * self.PATCH_SIZE + return latents[:, :, :, :, w0:w1].contiguous(), True - # Pad on effective-H so that H_tok is divisible by sp. - H_eff = latents.size(3) - - H_tok = H_eff // self.PATCH_SIZE - pad_tok = plan["H_tok_pad"] - H_tok - pad_lat = pad_tok * self.PATCH_SIZE - if pad_lat > 0: - pad = latents[:, :, :, -1:, :].repeat(1, 1, 1, pad_lat, 1) - latents = torch.cat([latents, pad], dim=3) - h0 = plan["h0_tok"] * self.PATCH_SIZE - h1 = (plan["h0_tok"] + plan["H_tok_local"]) * self.PATCH_SIZE - latents = latents[:, :, :, h0:h1, :] - - batch._zimage_sp_swap_hw = plan["swap_hw"] - return latents, True - - def gather_latents_for_sp(self, latents): - # Gather on effective-H dim=3 (matches shard_latents_for_sp); swap-back is handled in post_denoising_loop. + def gather_latents_for_sp(self, latents, batch): + # Gather native H/W shards by padding to a common collective shape, then crop. latents = latents.contiguous() - if get_sp_world_size() <= 1 or latents.dim() != 5: + if get_sp_world_size() <= 1 or latents.dim() not in (4, 5, 6): return latents - return sequence_model_parallel_all_gather(latents, dim=3) + + assert batch is not None + plan = self._get_zimage_sp_plan(batch) + if latents.dim() == 4: + shard_dim = 2 if plan["shard_axis"] == "h" else 3 + elif latents.dim() == 5: + shard_dim = 3 if plan["shard_axis"] == "h" else 4 + else: + shard_dim = 4 if plan["shard_axis"] == "h" else 5 + max_axis_tok = max(plan["shard_sizes_tok"]) + max_axis_lat = max_axis_tok * self.PATCH_SIZE + + pad_shape = list(latents.shape) + pad_shape[shard_dim] = max_axis_lat + padded = latents.new_zeros(pad_shape) + axis_len = latents.shape[shard_dim] + padded_slices = [slice(None)] * latents.dim() + padded_slices[shard_dim] = slice(axis_len) + padded[tuple(padded_slices)] = latents + + gathered = [torch.empty_like(padded) for _ in range(plan["sp_size"])] + dist.all_gather(gathered, padded, group=get_sp_group().device_group) + + pieces = [] + for rank, tensor in enumerate(gathered): + axis_lat = plan["shard_sizes_tok"][rank] * self.PATCH_SIZE + gather_slices = [slice(None)] * latents.dim() + gather_slices[shard_dim] = slice(axis_lat) + pieces.append(tensor[tuple(gather_slices)]) + return torch.cat(pieces, dim=shard_dim) def gather_noise_pred_for_sp(self, batch, noise_pred): - # Z-Image shards 5D latents on the effective-H axis, but ComfyUI noise_pred is 4D [B, C, H_local, W]. - noise_pred = self.gather_latents_for_sp(noise_pred) - if noise_pred.dim() == 4: - # reconstruct the full spatial tensor - noise_pred = sequence_model_parallel_all_gather( - noise_pred.contiguous(), dim=2 - ) - # restore the original H/W orientation - if getattr(batch, "_zimage_sp_swap_hw", False): - noise_pred = noise_pred.transpose(2, 3).contiguous() - return noise_pred + return self.gather_latents_for_sp(noise_pred, batch=batch) def post_denoising_loop(self, latents, batch): - # Restore swapped H/W and crop padded spatial dims before final reshape. - if latents.dim() == 5 and getattr(batch, "_zimage_sp_swap_hw", False): - latents = latents.transpose(3, 4).contiguous() raw_latent_shape = getattr(batch, "raw_latent_shape", None) if raw_latent_shape is not None and latents.dim() == 5: latents = latents[:, :, :, : raw_latent_shape[3], : raw_latent_shape[4]] @@ -239,16 +277,20 @@ class ZImagePipelineConfig(ImagePipelineConfig): ).flatten(0, 2) cap_freqs_cis = rotary_emb(cap_pos_ids) - # image (local, effective H-shard), offset after the full caption. + # Build image positions for the local native shard. F_tokens = 1 - H_tokens_local = plan["H_tok_local"] - W_tokens = plan["W_tok"] + H_tokens_local = plan["local_h_tok"] + W_tokens_local = plan["local_w_tok"] img_pos_ids = create_coordinate_grid( - size=(F_tokens, H_tokens_local, W_tokens), - start=(cap_ori_len + cap_padding_len + 1, plan["h0_tok"], 0), + size=(F_tokens, H_tokens_local, W_tokens_local), + start=( + cap_ori_len + cap_padding_len + 1, + plan["h0_tok"], + plan["w0_tok"], + ), device=device, ).flatten(0, 2) - img_pad_len = (-img_pos_ids.shape[0]) % self.SEQ_LEN_MULTIPLE + img_pad_len = plan["img_seq_target"] - img_pos_ids.shape[0] if img_pad_len: pad_ids = create_coordinate_grid( size=(1, 1, 1), start=(0, 0, 0), device=device @@ -308,6 +350,11 @@ class ZImagePipelineConfig(ImagePipelineConfig): rotary_emb, batch, ), + "image_seq_len_target": ( + self._get_zimage_sp_plan(batch)["img_seq_target"] + if get_sp_world_size() > 1 + else None + ), } def prepare_neg_cond_kwargs(self, batch, device, rotary_emb, dtype): @@ -320,4 +367,9 @@ class ZImagePipelineConfig(ImagePipelineConfig): rotary_emb, batch, ), + "image_seq_len_target": ( + self._get_zimage_sp_plan(batch)["img_seq_target"] + if get_sp_world_size() > 1 + else None + ), } diff --git a/python/sglang/multimodal_gen/runtime/models/dits/zimage.py b/python/sglang/multimodal_gen/runtime/models/dits/zimage.py index ca191075f..3758e04e6 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/zimage.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/zimage.py @@ -718,12 +718,19 @@ class ZImageTransformer2DModel(CachableDiT, OffloadableDiTMixin): grids = torch.meshgrid(axes, indexing="ij") return torch.stack(grids, dim=-1) + @staticmethod + def _ceil_to_multiple(value: int, multiple: int) -> int: + if multiple <= 0: + return value + return int(math.ceil(value / multiple) * multiple) + def patchify_and_embed( self, all_image: List[torch.Tensor], all_cap_feats: List[torch.Tensor], patch_size: int, f_patch_size: int, + image_seq_len_target: int | None = None, ): assert len(all_image) == len(all_cap_feats) == 1 @@ -761,7 +768,12 @@ class ZImageTransformer2DModel(CachableDiT, OffloadableDiTMixin): F_tokens * H_tokens * W_tokens, pF * pH * pW * C ) image_ori_len = image.size(0) - image_padding_len = (-image_ori_len) % SEQ_MULTI_OF + min_image_seq_len = self._ceil_to_multiple(image_ori_len, SEQ_MULTI_OF) + if image_seq_len_target is None: + image_seq_len_target = min_image_seq_len + else: + image_seq_len_target = max(min_image_seq_len, image_seq_len_target) + image_padding_len = image_seq_len_target - image_ori_len # padded feature image_padded_feat = torch.cat( @@ -788,6 +800,7 @@ class ZImageTransformer2DModel(CachableDiT, OffloadableDiTMixin): patch_size=2, f_patch_size=1, freqs_cis=None, + image_seq_len_target: int | None = None, **kwargs, ): assert patch_size in self.all_patch_size @@ -807,7 +820,13 @@ class ZImageTransformer2DModel(CachableDiT, OffloadableDiTMixin): x_size, x_valid_lens, cap_valid_lens, - ) = self.patchify_and_embed(x, cap_feats, patch_size, f_patch_size) + ) = self.patchify_and_embed( + x, + cap_feats, + patch_size, + f_patch_size, + image_seq_len_target=image_seq_len_target, + ) x = torch.cat(x, dim=0) x, _ = self.all_x_embedder[f"{patch_size}-{f_patch_size}"](x) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py index 937446899..2611c2733 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -24,6 +24,7 @@ from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType, S from sglang.multimodal_gen.configs.pipeline_configs.wan import ( Wan2_2_TI2V_5B_Config, ) +from sglang.multimodal_gen.configs.pipeline_configs.zimage import ZImagePipelineConfig from sglang.multimodal_gen.runtime.cache.cache_dit_integration import ( CacheDitConfig, enable_cache_on_dual_transformer, @@ -809,20 +810,29 @@ class DenoisingStage(PipelineStage): ) -> tuple[torch.Tensor, torch.Tensor | None]: """Gather latents after Sequence Parallelism if they were sharded.""" if get_sp_world_size() > 1 and getattr(batch, "did_sp_shard_latents", False): - latents = self.server_args.pipeline_config.gather_latents_for_sp(latents) + latents = self.server_args.pipeline_config.gather_latents_for_sp( + latents, batch=batch + ) if trajectory_tensor is not None: # trajectory_tensor shapes: # - video: [b, num_steps, c, t_local, h, w] -> gather on dim=3 # - image: [b, num_steps, s_local, d] -> gather on dim=2 trajectory_tensor = trajectory_tensor.to(get_local_torch_device()) - gather_dim = 3 if trajectory_tensor.dim() >= 5 else 2 - trajectory_tensor = sequence_model_parallel_all_gather( - trajectory_tensor, dim=gather_dim - ) - if gather_dim == 2 and hasattr(batch, "raw_latent_shape"): - orig_s = batch.raw_latent_shape[1] - if trajectory_tensor.shape[2] > orig_s: - trajectory_tensor = trajectory_tensor[:, :, :orig_s, :] + if isinstance(self.server_args.pipeline_config, ZImagePipelineConfig): + trajectory_tensor = ( + self.server_args.pipeline_config.gather_latents_for_sp( + trajectory_tensor, batch=batch + ) + ) + else: + gather_dim = 3 if trajectory_tensor.dim() >= 5 else 2 + trajectory_tensor = sequence_model_parallel_all_gather( + trajectory_tensor, dim=gather_dim + ) + if gather_dim == 2 and hasattr(batch, "raw_latent_shape"): + orig_s = batch.raw_latent_shape[1] + if trajectory_tensor.shape[2] > orig_s: + trajectory_tensor = trajectory_tensor[:, :, :orig_s, :] return latents, trajectory_tensor def step_profile(self): diff --git a/python/sglang/multimodal_gen/test/server/testcase_configs.py b/python/sglang/multimodal_gen/test/server/testcase_configs.py index f7c356e56..f35ae77b7 100644 --- a/python/sglang/multimodal_gen/test/server/testcase_configs.py +++ b/python/sglang/multimodal_gen/test/server/testcase_configs.py @@ -933,6 +933,20 @@ TWO_GPU_CASES_B = [ ), T2I_sampling_params, ), + DiffusionTestCase( + "zimage_image_t2i_2_gpus_non_square", + DiffusionServerArgs( + model_path=DEFAULT_SMALL_MODEL_NAME_FOR_TEST, + modality="image", + num_gpus=2, + ulysses_degree=2, + ), + DiffusionSamplingParams( + prompt=T2I_sampling_params.prompt, + output_size="1280x720", + ), + run_perf_check=False, + ), DiffusionTestCase( "flux_image_t2i_2_gpus", DiffusionServerArgs(