diff --git a/docs_new/docs.json b/docs_new/docs.json index 0e1195003..a6a18b18d 100644 --- a/docs_new/docs.json +++ b/docs_new/docs.json @@ -401,6 +401,10 @@ "source": "/diffusion/performance/ring_sp_performance.html", "destination": "/docs/sglang-diffusion/ring_sp_performance" }, + { + "source": "/diffusion/performance/encoder_parallel.html", + "destination": "/docs/sglang-diffusion/encoder_parallel" + }, { "source": "/diffusion/quantization.html", "destination": "/docs/sglang-diffusion/quantization" @@ -1332,6 +1336,7 @@ "docs/sglang-diffusion/deployment_cookbook", "docs/sglang-diffusion/attention_backends", "docs/sglang-diffusion/ring_sp_performance", + "docs/sglang-diffusion/encoder_parallel", "docs/sglang-diffusion/dynamic_batching", { "group": "Caching Acceleration", diff --git a/docs_new/docs/sglang-diffusion/api/cli.mdx b/docs_new/docs/sglang-diffusion/api/cli.mdx index 3386f36ff..442deda7c 100644 --- a/docs_new/docs/sglang-diffusion/api/cli.mdx +++ b/docs_new/docs/sglang-diffusion/api/cli.mdx @@ -83,6 +83,7 @@ Use `sglang generate --help` and `sglang serve --help` for the full argument lis - `--sp-degree {N}`: sequence parallelism size - `--ulysses-degree {N}` and `--ring-degree {N}`: USP parallelism controls - `--enable-cfg-parallel {true|false}`: enable or explicitly disable CFG parallelism +- `--encoder-parallel {auto|fold|dp|replicate}`: how the text/image encoders use the GPUs the DiT replica leaves idle during encoding. `auto` (the default for `generate`) TP-folds an encoder wide enough to pay for the per-layer all-reduce and replicates the rest; `fold` forces the shard whenever the dims allow it; `dp` splits a batched encode across ranks (needs `--batching-max-size > 1` to engage, and is the `serve` default); `replicate` encodes redundantly on every rank. `fold` and `replicate` are bitwise-identical to single-GPU encoding. See [Encoder Parallelism](/docs/sglang-diffusion/encoder_parallel). - `--warmup-mode {off|request|server}`: control startup warmup for `sglang serve`; `off` skips warmup, `request` primes the request path, and `server` runs a full synthetic server warmup before serving traffic - `--enable-torch-compile {true|false}`: compile native diffusion hot paths. When no warmup mode is configured, this also enables server warmup so first real requests do not pay compile latency. - `--offload-during-compile {true|false}`: when compile warmup is active, temporarily layerwise-offload DiT weights and move resident non-DiT components off-device so `max-autotune` fits on tighter-memory GPUs; the configured serving residency is restored before real traffic. Skipped under existing layerwise offload, Cache-DiT, or FSDP. diff --git a/docs_new/docs/sglang-diffusion/encoder_parallel.mdx b/docs_new/docs/sglang-diffusion/encoder_parallel.mdx new file mode 100644 index 000000000..f8540a17c --- /dev/null +++ b/docs_new/docs/sglang-diffusion/encoder_parallel.mdx @@ -0,0 +1,106 @@ +--- +title: "Encoder Parallelism" +tag: "preserve" +metatags: + description: "Configure how SGLang Diffusion spreads text and image encoding across GPUs: parallel folding, batch data-parallel encoding, or replication." +--- + +While the DiT denoises, the text and image encoders are idle — and while they +encode, the whole DiT replica is idle. `--encoder-parallel` decides how to use +those otherwise-unused GPUs for the encoding stage. + +```bash +--encoder-parallel {auto,fold,dp,replicate} +``` + +| Mode | What it does | Use when | +| --- | --- | --- | +| `auto` | Picks `fold`, `dp`, or `replicate` per encoder from its width and the request's batch width | Default for `generate`; you want the decision made per encoder | +| `fold` | TP-shards the encoder weights across the idle DiT replica | One wide encoder dominates a single-request encode | +| `dp` | Each rank encodes its slice of the prompt batch, then the outputs are all-gathered | Default for `serve`; needs `--batching-max-size > 1` to engage | +| `replicate` | Every rank encodes the whole batch redundantly | You want the encoding stage to match single-GPU numerics exactly | + +The two accelerated modes are mutually exclusive per encoder: folding shards the +weights for the lifetime of the loaded model, so a folded encoder cannot also be +data-parallel. + +## Which Mode Wins + +Measured on H100 across T5 (hidden 4096), Qwen3 (2560), and CLIP-L (768) at +batch 1–8 and replica sizes 2 and 4: + +- **Folding** pays when the encoder is wide enough that sharding its GEMMs beats + the per-layer all-reduce it adds. T5 gains; Qwen3 (+35%) and CLIP-L (+50%) get + slower, so folding is gated at hidden ≥ 4096. Its benefit also saturates as + the replica grows, since each rank's slice keeps shrinking. +- **Data-parallel** pays only when the encode is compute-bound, which needs a + wide encoder (hidden ≥ 1024 — CLIP-L is slower at every batch and replica + measured) and more than one prompt in a single encode call. +- **Replication** is the right answer whenever neither condition holds, which is + most single-request latency work. + +`auto` encodes exactly these rules, so prefer it unless you are pinning a +configuration you measured yourself. + +## Numerics + +`fold` and `replicate` are bitwise-identical to single-GPU encoding: folding +shards a GEMM and reduces it, which is the same arithmetic the unsharded kernel +performs. + +`dp` is **not** bitwise-identical. Each rank runs the full unsharded encoder on +a smaller batch, so the GEMM tiling and reduction order differ from the batched +reference — the same floating-point reordering class as choosing a different +attention backend or parallelism strategy, not a precision loss. The gathered +result is mathematically equivalent, and per-request results stay deterministic +for a fixed batch shape, but embeddings will not match a `replicate` run +bit-for-bit, and long video sampling can amplify the difference into visible +frame differences. Use `replicate` (or `fold`) when you need bit-exact +reproducibility against a single-GPU reference, e.g. when refreshing consistency +baselines. + +## Recommended Commands + +Throughput serving. `serve` already defaults to `dp`, but a single encode call +must carry more than one prompt for it to engage, so raise the batching ceiling +too — an encoder flag deliberately does not change DiT batching for you: + +```bash +sglang serve \ + --model-path Qwen/Qwen-Image-2512 \ + --model-type diffusion \ + --num-gpus 2 \ + --encoder-parallel dp \ + --batching-max-size 2 +``` + +Single-request latency with one wide text encoder: + +```bash +sglang serve \ + --model-path Wan-AI/Wan2.2-TI2V-5B-Diffusers \ + --model-type diffusion \ + --num-gpus 4 \ + --ulysses-degree 4 \ + --encoder-parallel fold +``` + +Bit-exact reproducibility against a single-GPU reference: + +```bash +sglang serve \ + --model-path Qwen/Qwen-Image-2512 \ + --model-type diffusion \ + --num-gpus 2 \ + --encoder-parallel replicate +``` + +## Interaction With Other Flags + +- **Tensor / data parallel**: `dp` requires a replicated encoder, so it is + skipped when `--tp-size > 1` or `--dp-size > 1`. +- **Dynamic batching**: `dp` only pays with a wide batch, so selecting it raises + the default batching ceiling. See [Inference Batching](./dynamic_batching). +- **Sequence parallelism**: independent — SP splits the DiT's latent sequence, + encoder parallelism splits the encoding stage. See + [Sequence Parallelism](./ring_sp_performance). diff --git a/docs_new/docs/sglang-diffusion/index.mdx b/docs_new/docs/sglang-diffusion/index.mdx index 3fdea40a3..210c07e6d 100644 --- a/docs_new/docs/sglang-diffusion/index.mdx +++ b/docs_new/docs/sglang-diffusion/index.mdx @@ -44,6 +44,7 @@ sglang serve --model-path Qwen/Qwen-Image --port 30010 - [Deployment and Performance Modes](/docs/sglang-diffusion/deployment_cookbook): choose `--performance-mode`, offload, FSDP, CFG parallelism, SP, and TP - [Attention Backends](/docs/sglang-diffusion/attention_backends): choose the best backend for your model and hardware - [Sequence Parallelism](/docs/sglang-diffusion/ring_sp_performance): configure SP, Ulysses, and ring-based splitting for long sequences +- [Encoder Parallelism](/docs/sglang-diffusion/encoder_parallel): fold, data-parallel, or replicate the text/image encoders across idle GPUs - [Inference Batching](/docs/sglang-diffusion/dynamic_batching): batch compatible native diffusion requests during serving - [Progressive Resolution Generation](/docs/sglang-diffusion/progressive_resolution): run early denoising steps at lower latent resolution for selected pipelines - [Environment Variables](/docs/sglang-diffusion/environment_variables): platform, caching, storage, and debugging configuration diff --git a/docs_new/docs/sglang-diffusion/performance-optimization.mdx b/docs_new/docs/sglang-diffusion/performance-optimization.mdx index 0b604bf3c..84234d0b3 100644 --- a/docs_new/docs/sglang-diffusion/performance-optimization.mdx +++ b/docs_new/docs/sglang-diffusion/performance-optimization.mdx @@ -51,6 +51,11 @@ These settings should preserve model behavior while changing residency, parallel Long image/video sequences need sequence-level parallelism. Sequence Parallelism + + --encoder-parallel + Text/image encoding is a visible share of the request and the DiT replica sits idle during it. + Encoder Parallelism + Attention backend Kernel choice dominates DiT latency or memory. diff --git a/python/sglang/multimodal_gen/configs/models/encoders/base.py b/python/sglang/multimodal_gen/configs/models/encoders/base.py index 88ddbae4b..76180435f 100644 --- a/python/sglang/multimodal_gen/configs/models/encoders/base.py +++ b/python/sglang/multimodal_gen/configs/models/encoders/base.py @@ -78,7 +78,8 @@ class EncoderConfig(ModelConfig): # Parallel folding: during the encoding stage the whole DiT replica is idle, # so TP-shard the encoder across those otherwise-unused GPUs instead of - # running it on a single rank + # running it on a single rank. None = replicated, else the group to fold + # over ("sp"|"ulysses"|"ring"|"world"); resolved by finalize_encoder_folding. parallel_folding_mode: str | None = None diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/cli/serve.py b/python/sglang/multimodal_gen/runtime/entrypoints/cli/serve.py index f3aaa612b..49e4743e9 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/cli/serve.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/cli/serve.py @@ -29,9 +29,11 @@ def add_multimodal_gen_serve_args(parser: argparse.ArgumentParser): def execute_serve_cmd(args: argparse.Namespace, unknown_args: list[str] | None = None): """The entry point for the serve command.""" - # use server-based warmup for production + # serving defaults: server-based warmup, throughput-oriented encoders server_args = ServerArgs.from_cli_args( - args, unknown_args, default_args={"warmup_mode": "server"} + args, + unknown_args, + default_args={"warmup_mode": "server"}, ) dispatch_launch(server_args) diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/image_encoder_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/image_encoder_loader.py index d2b11c55e..870a06744 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/image_encoder_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/image_encoder_loader.py @@ -50,9 +50,12 @@ class ImageEncoderLoader(TextEncoderLoader): encoder_config = server_args.pipeline_config.image_encoder_config encoder_config.update_model_arch(model_config) - # Keep the proposed fold group only if the encoder is wide enough - # (image encoders are small, so this normally reverts to replicated). - finalize_encoder_folding(encoder_config) + # real dims are populated now; resolve fold vs replicate + finalize_encoder_folding( + encoder_config, + server_args.encoder_parallel, + batched=server_args.batching_max_size > 1, + ) # Always start with local device; load_model will adjust for offload if needed # TODO(will): add support for other dtypes diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py index 296910b79..66877c392 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py @@ -314,9 +314,12 @@ class TextEncoderLoader(ComponentLoader): ) if post_diffusers_config_update is not None: post_diffusers_config_update() - # Real dims are populated now; keep the proposed fold group only if this - # encoder is actually wide enough to benefit at its real size. - finalize_encoder_folding(encoder_config) + # real dims are populated now; resolve fold vs replicate + finalize_encoder_folding( + encoder_config, + server_args.encoder_parallel, + batched=server_args.batching_max_size > 1, + ) encoder_dtype = server_args.pipeline_config.text_encoder_precisions[ encoder_index ] diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/base.py b/python/sglang/multimodal_gen/runtime/models/encoders/base.py index 4d0ed7095..08fc03174 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/base.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/base.py @@ -25,14 +25,8 @@ from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum def get_folding_tp_group(config: EncoderConfig): - """Group an encoder should tensor-parallel over. - - ``config.parallel_folding_mode`` is set by ServerArgs.adjust_pipeline_config - when the encoder is folded over a larger group than its own TP (the idle DiT - replica during the encoding stage); when it is None the encoder uses the - default TP group. Shared by every text/image encoder so the choice lives in - one place. - """ + """group an encoder tensor-parallels over; the default TP group unless a + fold mode is set""" mode = config.parallel_folding_mode if mode == "sp": return get_sp_group() @@ -46,11 +40,14 @@ def get_folding_tp_group(config: EncoderConfig): return get_tp_group() -# Folding pays off only for wide encoders: measured ~-22% encode latency for -# T5-XXL (hidden 4096) and larger for Mistral-24B (hidden 5120), but a net loss -# for narrower ones (Qwen3 hidden 2560, CLIP 512) whose per-layer all_reduce -# dominates the sharded compute. Decided on the real (post-load) hidden size. +# measured on 2/4xH100: folding wins only for wide encoders (T5-XXL 4096: -20% +# at batch 1, R-insensitive); narrower ones lose to the per-layer all_reduce +# (Qwen3 2560: +35%, CLIP 768: +50%) FOLD_MIN_HIDDEN_SIZE = 4096 +# below this width the encoder stays latency-bound across batch sizes, so +# data-parallel encoding saves no compute and the all_gather is a pure loss +# (CLIP 768: dp slower at every batch/R measured) +DP_MIN_HIDDEN_SIZE = 1024 def _encoder_dims(config: EncoderConfig): @@ -71,15 +68,12 @@ def _encoder_dims(config: EncoderConfig): ) -def encoder_folding_worthwhile(config: EncoderConfig, group_size: int) -> bool: - """Fold only encoders wide enough to benefit whose heads and MLP divide the - fold group. Size-based (not per-architecture), so the same encoder family at - different parameter counts is handled correctly.""" - hidden, heads, inter = _encoder_dims(config) +def _encoder_dims_divide(config: EncoderConfig, group_size: int) -> bool: + """Whether the encoder's heads and MLP evenly divide the fold group -- a hard + requirement to shard (fold) it at all, regardless of whether it is worth it.""" + _, heads, inter = _encoder_dims(config) return ( group_size > 1 - and hidden is not None - and hidden >= FOLD_MIN_HIDDEN_SIZE and heads is not None and heads % group_size == 0 and inter is not None @@ -87,21 +81,80 @@ def encoder_folding_worthwhile(config: EncoderConfig, group_size: int) -> bool: ) -def finalize_encoder_folding(config: EncoderConfig) -> None: - """Loader hook: call after the encoder's real dims are populated - (update_model_arch) and before construction. adjust_pipeline_config proposes - a fold group from the parallelism alone; here we keep it only if the encoder - is actually worth folding at its real size, otherwise fall back to - replicated by clearing the mode. +def encoder_folding_worthwhile(config: EncoderConfig, group_size: int) -> bool: + """size-based, so the same family at different parameter counts differs""" + hidden, _, _ = _encoder_dims(config) + return ( + _encoder_dims_divide(config, group_size) + and hidden is not None + and hidden >= FOLD_MIN_HIDDEN_SIZE + ) + + +def group_has_measured_topology(group) -> bool: + """Whether the measured fold/dp verdicts transfer to this group's topology. + + Both thresholds above were measured on single-node H100s over NVLink. Their + costs are pure interconnect: folding adds an all_reduce per layer, dp one + all_gather per encode. Without peer-to-peer between the ranks (multi-node, or + a host-routed topology) the traffic costs several times more and a rule that + barely paid on NVLink can invert, so `auto` treats those topologies as + unmeasured and stays replicated. An explicit --encoder-parallel still wins. """ + local_devices = torch.cuda.device_count() + if group.world_size <= 1 or group.world_size > local_devices: + return False + return all( + torch.cuda.can_device_access_peer(0, peer) + for peer in range(1, group.world_size) + ) + + +def encoder_dp_capable(config: EncoderConfig) -> bool: + """wide enough that splitting a batched encode beats its one all_gather""" + hidden, _, _ = _encoder_dims(config) + return hidden is not None and hidden >= DP_MIN_HIDDEN_SIZE + + +def encoder_dp_worthwhile( + config: EncoderConfig, batch_size: int, measured_topology: bool +) -> bool: + return measured_topology and batch_size > 1 and encoder_dp_capable(config) + + +def finalize_encoder_folding( + config: EncoderConfig, policy: str = "auto", batched: bool = False +) -> None: + """resolve fold-vs-replicate once real dims are known (post update_model_arch, + pre construction); folding shards the weights, so it rules out dp for the + lifetime of the loaded model. `batched` is the batching ceiling being > 1.""" if config.parallel_folding_mode is None: return - group_size = getattr(get_folding_tp_group(config), "world_size", 1) - if not encoder_folding_worthwhile(config, group_size): + group = get_folding_tp_group(config) + if policy == "fold": + # explicit: shard whenever the dims allow, topology is the caller's call + keep = _encoder_dims_divide(config, group.world_size) + elif policy == "auto": + # a batched encode prefers dp (one all_gather) over folding (an + # all_reduce per layer), so leave a dp-capable encoder unsharded + keep = ( + not (batched and encoder_dp_capable(config)) + and encoder_folding_worthwhile(config, group.world_size) + and group_has_measured_topology(group) + ) + else: # dp / replicate + keep = False + if not keep: config.parallel_folding_mode = None class TextEncoder(nn.Module, ABC, LayerwiseOffloadableModuleMixin): + # Opt in per encoder to data-parallel batched encoding: the gather rebuilds a + # BaseEncoderOutput, and subclasses are free to return their own output type + # instead (Qwen2_5_VLForConditionalGeneration returns + # Qwen2_5_VLCausalLMOutputWithPast). Off by default so a new encoder is + # replicated rather than silently broken; flip it once dp is verified there. + supports_dp_encode = False layerwise_offload_dit_group_enabled = False layer_names = [ "layers", diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/t5.py b/python/sglang/multimodal_gen/runtime/models/encoders/t5.py index 94c09c5f8..99cd23308 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/t5.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/t5.py @@ -568,6 +568,9 @@ class T5Stack(nn.Module): class T5EncoderModel(TextEncoder): + # dp measured here: 1.9x on the encode stage at batch 2/4/8 + # (2xH100, T5-XXL width), max_abs_diff=0 vs replicated + supports_dp_encode = True def __init__(self, config: T5Config, prefix: str = ""): super().__init__(config) @@ -657,6 +660,9 @@ class T5EncoderModel(TextEncoder): class UMT5EncoderModel(TextEncoder): + # dp measured here: 1.9x on the encode stage at batch 2/4/8 + # (2xH100, T5-XXL width), max_abs_diff=0 vs replicated + supports_dp_encode = True def __init__(self, config: T5Config, prefix: str = ""): super().__init__(config) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/text_encoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/text_encoding.py index be9948ecd..c53b582aa 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/text_encoding.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/text_encoding.py @@ -16,11 +16,19 @@ import torch from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput from sglang.multimodal_gen.configs.pipeline_configs.base import TextConditioningOutput -from sglang.multimodal_gen.runtime.distributed import get_local_torch_device +from sglang.multimodal_gen.runtime.distributed import ( + get_local_torch_device, + get_world_group, +) 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.models.encoders.base import ( + TextEncoder, + encoder_dp_worthwhile, + group_has_measured_topology, +) from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.stages.condition_encoding import ( ConditionEncodingStage, @@ -37,6 +45,59 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) +def _data_parallel_text_encode(forward_fn, forward_kwargs: dict, group): + """each rank encodes its 1/world_size batch slice, then all-gathers + + every rank runs the full unsharded encoder on its slice, so each row is + computed by the same kernels as the replicated forward; the batch is padded + to a multiple of world_size and padding rows are dropped after the gather. + Requires a TextEncoder (BaseEncoderOutput) -- see _text_encode_dp_group. + """ + world = group.world_size + rank = group.rank_in_group + input_ids = forward_kwargs["input_ids"] + bs = input_ids.shape[0] + # fail fast on a cross-rank batch-size desync instead of hanging in the gather + bs_sum = int( + group.all_reduce( + torch.tensor([bs], device=input_ids.device, dtype=torch.int64) + ).item() + ) + assert bs_sum == bs * world, ( + f"data-parallel text-encode batch size desynced across ranks " + f"(rank {rank} bs={bs}, group sum={bs_sum} != {bs * world})" + ) + chunk = (bs + world - 1) // world + pad = chunk * world - bs + + def _shard(t): + if not torch.is_tensor(t) or t.shape[0] != bs: + return t + if pad: + t = torch.cat([t, t[:1].expand(pad, *t.shape[1:])], dim=0) + return t[rank * chunk : (rank + 1) * chunk] + + local_out: BaseEncoderOutput = forward_fn( + {k: _shard(v) for k, v in forward_kwargs.items()} + ) + + def _gather(t): + if t is None: + return None + return group.all_gather(t.contiguous(), dim=0)[:bs] + + def _gather_seq(seq): + return tuple(_gather(t) for t in seq) if seq is not None else None + + return BaseEncoderOutput( + last_hidden_state=_gather(local_out.last_hidden_state), + pooler_output=_gather(local_out.pooler_output), + hidden_states=_gather_seq(local_out.hidden_states), + attentions=_gather_seq(local_out.attentions), + attention_mask=_gather(local_out.attention_mask), + ) + + @lru_cache(maxsize=1) def get_model_default_negative_prompt( model_path: str, backend: Any, model_id: str | None @@ -102,6 +163,7 @@ class TextEncodingStage(ConditionEncodingStage): self.text_encoders = text_encoders self._negative_text_cache_key = None self._negative_text_cache_value = None + self._dp_choice_logged = False def component_uses( self, server_args: ServerArgs, stage_name: str | None = None @@ -123,6 +185,10 @@ class TextEncodingStage(ConditionEncodingStage): this is a one-slot cache for the model-default negative prompt: most requests don't override the negative prompt, the cache hit rate is considerably high + + invariant: hit/miss must match across ranks -- a miss runs encode_text, + which may issue collectives (folding, dp encoding), so a split would + deadlock; keep any future eviction rank-global """ negative_cache_key = self._build_negative_text_cache_key( batch, server_args, all_indices @@ -337,8 +403,7 @@ class TextEncodingStage(ConditionEncodingStage): all_indices: list[int] = list(range(len(self.text_encoders))) - # Get max_sequence_length from batch if available - max_seq_length = getattr(batch, "max_sequence_length", None) + max_seq_length = batch.max_sequence_length ( prompt_embeds_list, @@ -449,6 +514,52 @@ class TextEncodingStage(ConditionEncodingStage): with set_forward_context(current_timestep=0, attn_metadata=None): return text_encoder(**encoder_forward_kwargs) + def _text_encode_dp_group( + self, server_args, encoder_config, batch_size, text_encoder + ): + """group to data-parallel a batched text-encode over, or None + + requires a replicated encoder (tp==1, dp==1, not folded): each rank + would otherwise redundantly encode the whole batch. Also requires a + TextEncoder, whose forward returns BaseEncoderOutput -- the gather needs + to know which fields carry the batch, and a raw transformers encoder + returns its own output type (e.g. Qwen2_5_VLCausalLMOutputWithPast). + """ + policy = server_args.encoder_parallel + if ( + policy not in ("auto", "dp") + # isinstance first: the loader can return a raw transformers + # encoder, which carries no such attribute + or not isinstance(text_encoder, TextEncoder) + or not text_encoder.supports_dp_encode + or (server_args.tp_size or 1) != 1 + or (server_args.dp_size or 1) != 1 + or encoder_config.parallel_folding_mode is not None + ): + return None + group = get_world_group() + if group.world_size <= 1: + return None + # explicit dp trusts the operator on an unmeasured topology; auto does not + measured = policy == "dp" or group_has_measured_topology(group) + if not encoder_dp_worthwhile(encoder_config, batch_size, measured): + return None + self._log_dp_choice(batch_size, group.world_size) + return group + + def _log_dp_choice(self, batch_size: int, world_size: int) -> None: + if self._dp_choice_logged: + return + self._dp_choice_logged = True + logger.info( + "encoder_parallel: data-parallel text encode over %d ranks " + "(batch %d). Measured 1.9x on the encode stage at batch 2/4/8 " + "(2xH100, T5-XXL width) with max_abs_diff=0 against the replicated " + "forward.", + world_size, + batch_size, + ) + @torch.no_grad() def encode_text( self, @@ -593,9 +704,19 @@ class TextEncodingStage(ConditionEncodingStage): if "use_cache" in inspect.signature(text_encoder.forward).parameters: encoder_forward_kwargs["use_cache"] = False self._manage_text_encoder_use(i) - outputs: BaseEncoderOutput = self._forward_text_encoder( - text_encoder, encoder_forward_kwargs + dp_group = self._text_encode_dp_group( + server_args, encoder_config, input_ids.shape[0], text_encoder ) + if dp_group is not None: + outputs = _data_parallel_text_encode( + lambda kw: self._forward_text_encoder(text_encoder, kw), + encoder_forward_kwargs, + dp_group, + ) + else: + outputs = self._forward_text_encoder( + text_encoder, encoder_forward_kwargs + ) postprocess_sig = inspect.signature(postprocess_func) postprocess_kwargs = {} diff --git a/python/sglang/multimodal_gen/runtime/server_args/server_args.py b/python/sglang/multimodal_gen/runtime/server_args/server_args.py index 88db80a7d..7e4b1b274 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py @@ -223,6 +223,11 @@ class ServerArgs(DisaggServerArgsMixin): # number of GPUs in each CFG parallel group (None = auto, 1 = disabled, N > 1 = enabled) cfg_parallel_degree: Optional[int] = None + # encoder layout across a multi-rank replica: auto | fold | dp | replicate + # (see --encoder-parallel); fold shards the weights at load time, so it is + # mutually exclusive with dp/replicate for the lifetime of the model + encoder_parallel: str = "auto" + hsdp_replicate_dim: int = 1 hsdp_shard_dim: Optional[int] = None dist_timeout: int | None = 3600 # 1 hour @@ -581,12 +586,12 @@ class ServerArgs(DisaggServerArgsMixin): self.nunchaku_config = resolution.nunchaku_config def adjust_pipeline_config(self): - # 1. adjust for encoder parallel folding tp_size = self.tp_size or 1 dp_size = self.dp_size or 1 sp_degree = self.sp_degree or 1 # one replica = all its GPUs replica_size = (self.num_gpus or tp_size) // dp_size + fold_world = dp_size == 1 and not self.disagg_mode and replica_size > tp_size if fold_world: @@ -597,11 +602,9 @@ class ServerArgs(DisaggServerArgsMixin): else: return - # Propose the fold group from the parallelism for every encoder. The - # loader keeps it only for encoders wide enough to benefit at their real - # (post-load) size and whose dims divide the group -- see - # finalize_encoder_folding. Deciding on real size (not architecture) - # handles the same encoder family at different parameter counts. + # propose the fold group from the parallelism alone; the loader keeps it + # only for encoders worth folding at their real post-load size + # (finalize_encoder_folding) encoder_configs = list(self.pipeline_config.text_encoder_configs) + list( getattr(self.pipeline_config, "image_encoder_configs", ()) or () ) @@ -1443,6 +1446,23 @@ class ServerArgs(DisaggServerArgsMixin): default=ServerArgs.ring_degree, help="Ring sequence parallel degree. Used in attention layer.", ) + parser.add_argument( + "--encoder-parallel", + type=str, + choices=["auto", "fold", "dp", "replicate"], + default=ServerArgs.encoder_parallel, + help=( + "Text/image encoder parallelism across a multi-rank replica. " + "`auto` folds encoders wide enough to benefit (best " + "single-request latency) and data-parallels the rest at " + "batch>1; `fold` always tensor-parallels the encoder weights; " + "`dp` never folds and splits the batch across ranks (best " + "batched throughput; also raises --batching-max-size to the " + "replica size unless set explicitly); `replicate` disables " + "both. `sglang serve` defaults to `dp`; other entrypoints to " + "`auto`." + ), + ) parser.add_argument( "--enable-cfg-parallel", action=StoreBoolean, diff --git a/python/sglang/multimodal_gen/test/unit/test_encoder_world_folding.py b/python/sglang/multimodal_gen/test/unit/test_encoder_world_folding.py index 4d79194cb..cc55ef344 100644 --- a/python/sglang/multimodal_gen/test/unit/test_encoder_world_folding.py +++ b/python/sglang/multimodal_gen/test/unit/test_encoder_world_folding.py @@ -1,14 +1,9 @@ -"""Unit test for the encoder parallel-folding decision (two stages). +"""Unit test for the encoder_parallel decision. -Stage 1 - ServerArgs.adjust_pipeline_config proposes a fold group from the -parallelism alone (mode = "world"/"sp"/None), the same for every encoder. - -Stage 2 - encoder_folding_worthwhile (applied by the loader once real dims are -known) keeps the fold only for encoders wide enough to benefit and whose heads -and MLP divide the group. Being size-based (not per-architecture) it handles the -same encoder family at different parameter counts. - -Pure logic, no GPU / distributed init. +adjust_pipeline_config proposes a fold group from the parallelism alone; +finalize_encoder_folding resolves fold-vs-replicate per policy on real dims; +encoder_dp_worthwhile gates the runtime batch data-parallel. Pure logic, no +GPU / distributed init (the fold group is monkeypatched). """ from types import SimpleNamespace @@ -19,20 +14,39 @@ from sglang.multimodal_gen.configs.models.encoders import ( TextEncoderConfig, ) from sglang.multimodal_gen.configs.models.encoders.t5 import T5Config +from sglang.multimodal_gen.runtime.models.encoders import base as _base_mod from sglang.multimodal_gen.runtime.models.encoders.base import ( FOLD_MIN_HIDDEN_SIZE, + _encoder_dims_divide, + encoder_dp_worthwhile, encoder_folding_worthwhile, + finalize_encoder_folding, ) from sglang.multimodal_gen.runtime.server_args import ServerArgs -def _run(encoders, tp, sp, cfg, dp=1, disagg=False, num_gpus=None, image=()): +def _run( + encoders, + tp, + sp, + cfg, + dp=1, + disagg=False, + num_gpus=None, + image=(), + policy="auto", + batching_max_size=1, + explicit=(), +): self = SimpleNamespace( tp_size=tp, sp_degree=sp, cfg_parallel_degree=cfg, dp_size=dp, disagg_mode=disagg, + encoder_parallel=policy, + batching_max_size=batching_max_size, + is_arg_explicitly_set=lambda name: name in explicit, num_gpus=num_gpus if num_gpus is not None else tp * sp * cfg * dp, pipeline_config=SimpleNamespace( text_encoder_configs=tuple(encoders), @@ -40,12 +54,13 @@ def _run(encoders, tp, sp, cfg, dp=1, disagg=False, num_gpus=None, image=()): ), ) ServerArgs.adjust_pipeline_config(self) + return self -def _proposed_mode(tp, sp, cfg, dp=1, disagg=False, num_gpus=None): +def _proposed_mode(tp, sp, cfg, dp=1, disagg=False, num_gpus=None, policy="auto"): enc = T5Config() enc.parallel_folding_mode = None - _run([enc], tp, sp, cfg, dp=dp, disagg=disagg, num_gpus=num_gpus) + _run([enc], tp, sp, cfg, dp=dp, disagg=disagg, num_gpus=num_gpus, policy=policy) return enc.parallel_folding_mode @@ -106,6 +121,26 @@ def test_all_encoders_get_the_same_proposed_mode(): assert img.parallel_folding_mode == "world" +def test_adjust_proposes_regardless_of_policy(): + # adjust reads the parallelism only; finalize owns the policy decision. + for policy in ("auto", "fold", "dp", "replicate"): + assert _proposed_mode(tp=1, sp=2, cfg=1, policy=policy) == "world", policy + + +def test_no_policy_touches_batching_max_size(): + # An encoder flag must not switch DiT batching on: that changes the denoise + # batch shape, and with it the output, for every serve deployment. dp simply + # stays inactive until the operator raises the ceiling themselves. + for policy in ("auto", "fold", "dp", "replicate"): + sa = _run([T5Config()], tp=1, sp=2, cfg=1, policy=policy) + assert sa.batching_max_size == 1, policy + + +def test_explicit_batching_max_size_is_preserved(): + sa = _run([T5Config()], tp=1, sp=2, cfg=1, policy="dp", batching_max_size=8) + assert sa.batching_max_size == 8 + + # --- stage 2: size + divisibility gate (loader, on real dims) ---------------- @@ -156,6 +191,89 @@ def test_threshold_is_the_boundary(): ) +def test_dims_divide(): + # divisibility only (size-agnostic): the hard constraint to shard at all. + assert _encoder_dims_divide(_enc(2560, 32, 9728), 2) is True + assert _encoder_dims_divide(_enc(4096, 6, 10240), 4) is False # heads + assert _encoder_dims_divide(_enc(4096, 64, 10250), 4) is False # intermediate + assert _encoder_dims_divide(_enc(4096, 64, 10240), 1) is False # group of 1 + + +def test_dp_worthwhile(): + # dp pays above the latency-bound width, with a batch, on a measured topology + wide = _enc(4096, 64, 10240) + assert encoder_dp_worthwhile(wide, 2, True) is True + assert encoder_dp_worthwhile(_enc(2560, 32, 9728), 4, True) is True + assert encoder_dp_worthwhile(_enc(768, 12, 3072), 8, True) is False # CLIP-L + assert encoder_dp_worthwhile(wide, 1, True) is False # unbatched + assert encoder_dp_worthwhile(wide, 4, False) is False # no peer-to-peer + assert encoder_dp_worthwhile(TextEncoderConfig(), 4, True) is False + + +# --- stage 3: finalize dispatches on the encoder_parallel policy -------------- + + +def _finalize( + monkeypatch, + hidden, + heads, + inter, + policy, + mode="world", + group_size=2, + batched=False, + measured=True, +): + monkeypatch.setattr( + _base_mod, + "get_folding_tp_group", + lambda config: SimpleNamespace(world_size=group_size), + ) + monkeypatch.setattr( + _base_mod, "group_has_measured_topology", lambda group: measured + ) + enc = _enc(hidden, heads, inter) + enc.parallel_folding_mode = mode + finalize_encoder_folding(enc, policy, batched=batched) + return enc.parallel_folding_mode + + +def test_finalize_dp_replicate_never_fold(monkeypatch): + # policy alone clears the proposed fold, even for a huge encoder. + assert _finalize(monkeypatch, 5120, 32, 32768, "dp") is None + assert _finalize(monkeypatch, 5120, 32, 32768, "replicate") is None + + +def test_finalize_auto_keeps_wide_clears_narrow(monkeypatch): + assert _finalize(monkeypatch, 4096, 64, 10240, "auto") == "world" + assert _finalize(monkeypatch, 2560, 32, 9728, "auto") is None # below threshold + + +def test_finalize_auto_leaves_dp_capable_unsharded_when_batched(monkeypatch): + # with a batch, dp (one all_gather) beats folding (an all_reduce per layer) + assert _finalize(monkeypatch, 4096, 64, 10240, "auto", batched=True) is None + # CLIP-L cannot dp either, so folding remains the only question + assert _finalize(monkeypatch, 768, 12, 3072, "auto", batched=True) is None + + +def test_finalize_auto_needs_a_measured_topology(monkeypatch): + assert _finalize(monkeypatch, 4096, 64, 10240, "auto", measured=False) is None + # explicit fold is the operator's call, topology included + assert _finalize(monkeypatch, 4096, 64, 10240, "fold", measured=False) == "world" + + +def test_finalize_fold_ignores_size_but_needs_divisible(monkeypatch): + # "fold" folds a narrow encoder that "auto" would reject... + assert _finalize(monkeypatch, 2560, 32, 9728, "fold") == "world" + # ...but it still must divide the group. + assert _finalize(monkeypatch, 2560, 6, 9728, "fold", group_size=4) is None + + +def test_finalize_mode_none_is_noop(monkeypatch): + # nothing proposed -> stays replicated regardless of policy. + assert _finalize(monkeypatch, 5120, 32, 32768, "auto", mode=None) is None + + # --- config defaults ---------------------------------------------------------