[diffusion] perf: tp-shard every text/image encoder across the full DiT replica (any parallelism) (#30086)

This commit is contained in:
Mick
2026-07-06 14:48:07 +08:00
committed by GitHub
parent de00b838c4
commit 5f98f62a8a
12 changed files with 396 additions and 51 deletions
@@ -76,16 +76,16 @@ class EncoderConfig(ModelConfig):
quant_config: QuantizationConfig | None = None quant_config: QuantizationConfig | None = None
lora_config: Any | None = None lora_config: Any | None = None
# 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
parallel_folding_mode: str | None = None
@dataclass @dataclass
class TextEncoderConfig(EncoderConfig): class TextEncoderConfig(EncoderConfig):
arch_config: ArchConfig = field(default_factory=TextEncoderArchConfig) arch_config: ArchConfig = field(default_factory=TextEncoderArchConfig)
# Use the SP Group of the transformer as the TP Group of T5.
parallel_folding: bool = False
# "sp" or "ulysses" or "ring"
parallel_folding_mode: str = "sp"
@dataclass @dataclass
class ImageEncoderConfig(EncoderConfig): class ImageEncoderConfig(EncoderConfig):
@@ -78,10 +78,6 @@ class T5Config(TextEncoderConfig):
arch_config: TextEncoderArchConfig = field(default_factory=T5ArchConfig) arch_config: TextEncoderArchConfig = field(default_factory=T5ArchConfig)
prefix: str = "t5" prefix: str = "t5"
# Use the SP Group of the transformer as the TP Group of T5.
parallel_folding: bool = False
# "sp" or "ulysses" or "ring"
parallel_folding_mode: str = "sp"
@staticmethod @staticmethod
def add_cli_args( def add_cli_args(
@@ -2,6 +2,7 @@ from sglang.multimodal_gen.configs.models import ModelConfig
from sglang.multimodal_gen.runtime.loader.component_loaders.text_encoder_loader import ( from sglang.multimodal_gen.runtime.loader.component_loaders.text_encoder_loader import (
TextEncoderLoader, TextEncoderLoader,
) )
from sglang.multimodal_gen.runtime.models.encoders.base import finalize_encoder_folding
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import ( from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
get_diffusers_component_config, get_diffusers_component_config,
@@ -49,6 +50,9 @@ class ImageEncoderLoader(TextEncoderLoader):
encoder_config = server_args.pipeline_config.image_encoder_config encoder_config = server_args.pipeline_config.image_encoder_config
encoder_config.update_model_arch(model_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)
# Always start with local device; load_model will adjust for offload if needed # Always start with local device; load_model will adjust for offload if needed
# TODO(will): add support for other dtypes # TODO(will): add support for other dtypes
@@ -3,6 +3,7 @@ import glob
import os import os
import re import re
from collections.abc import Generator, Iterable from collections.abc import Generator, Iterable
from contextlib import nullcontext
from typing import cast from typing import cast
import torch import torch
@@ -15,7 +16,14 @@ from sglang.multimodal_gen.configs.models import EncoderConfig, ModelConfig
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import ( from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
QwenImageEditPipelineConfig, QwenImageEditPipelineConfig,
) )
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.distributed import (
get_local_torch_device,
get_tp_group,
)
from sglang.multimodal_gen.runtime.distributed.group_coordinator import GroupCoordinator
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
patch_tensor_parallel_group,
)
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
ComponentLoader, ComponentLoader,
) )
@@ -30,6 +38,10 @@ from sglang.multimodal_gen.runtime.loader.weight_utils import (
pt_weights_iterator, pt_weights_iterator,
safetensors_weights_iterator, safetensors_weights_iterator,
) )
from sglang.multimodal_gen.runtime.models.encoders.base import (
finalize_encoder_folding,
get_folding_tp_group,
)
from sglang.multimodal_gen.runtime.models.registry import ModelRegistry from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.platforms import current_platform
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
@@ -302,6 +314,9 @@ class TextEncoderLoader(ComponentLoader):
) )
if post_diffusers_config_update is not None: if post_diffusers_config_update is not None:
post_diffusers_config_update() 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)
encoder_dtype = server_args.pipeline_config.text_encoder_precisions[ encoder_dtype = server_args.pipeline_config.text_encoder_precisions[
encoder_index encoder_index
] ]
@@ -374,7 +389,20 @@ class TextEncoderLoader(ComponentLoader):
else: else:
model_device = local_torch_device model_device = local_torch_device
with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]): # Parallel folding: build + shard the encoder over the folding group (the
# idle DiT replica during the encoding stage) instead of the default TP
# group, so every encoder folds without threading the group through each layer.
fold_ctx = nullcontext()
if getattr(model_config, "parallel_folding_mode", None) is not None:
folding_group = get_folding_tp_group(model_config)
if (
isinstance(folding_group, GroupCoordinator)
and folding_group is not get_tp_group()
):
fold_ctx = patch_tensor_parallel_group(folding_group)
# patch tp group with folding group to achieve TP among folding group
with fold_ctx, set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
with model_device, skip_init_modules(): with model_device, skip_init_modules():
architectures = getattr(model_config, "architectures", []) architectures = getattr(model_config, "architectures", [])
model_cls, _ = ModelRegistry.resolve_model_cls(architectures) model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
@@ -9,15 +9,98 @@ from torch import nn
from sglang.multimodal_gen.configs.models.encoders import ( from sglang.multimodal_gen.configs.models.encoders import (
BaseEncoderOutput, BaseEncoderOutput,
EncoderConfig,
ImageEncoderConfig, ImageEncoderConfig,
TextEncoderConfig, TextEncoderConfig,
) )
from sglang.multimodal_gen.runtime.distributed import (
get_sp_group,
get_tp_group,
get_world_group,
)
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
LayerwiseOffloadableModuleMixin, LayerwiseOffloadableModuleMixin,
) )
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum 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.
"""
mode = config.parallel_folding_mode
if mode == "sp":
return get_sp_group()
elif mode == "ulysses":
return get_sp_group().ulysses_group
elif mode == "ring":
return get_sp_group().ring_group
elif mode == "world":
# the whole single-replica DiT (all GPUs), regardless of tp/sp/cfg.
return get_world_group()
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.
FOLD_MIN_HIDDEN_SIZE = 4096
def _encoder_dims(config: EncoderConfig):
"""Best-effort (hidden, attention_heads, mlp_intermediate) from a config,
spelled differently across families (hidden_size/d_model, num_heads, d_ff)."""
def first(names):
for name in names:
value = getattr(config, name, None)
if isinstance(value, int) and value > 0:
return value
return None
return (
first(("hidden_size", "d_model")),
first(("num_attention_heads", "num_heads", "n_heads")),
first(("intermediate_size", "d_ff", "ffn_dim")),
)
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)
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
and inter % group_size == 0
)
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.
"""
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):
config.parallel_folding_mode = None
class TextEncoder(nn.Module, ABC, LayerwiseOffloadableModuleMixin): class TextEncoder(nn.Module, ABC, LayerwiseOffloadableModuleMixin):
layerwise_offload_dit_group_enabled = False layerwise_offload_dit_group_enabled = False
layer_names = [ layer_names = [
@@ -529,6 +529,9 @@ class SiglipAttention(nn.Module):
tp_size = get_tp_world_size() tp_size = get_tp_world_size()
self.head_dim = hidden_size // num_heads self.head_dim = hidden_size // num_heads
self.num_heads_per_partition = num_heads // tp_size self.num_heads_per_partition = num_heads // tp_size
# Cache the per-rank projection width so forward() does not re-read the
# global TP size (which is not patched to the folding group at run time).
self.embed_dim_per_partition = self.num_heads_per_partition * self.head_dim
self.scaling = self.head_dim**-0.5 self.scaling = self.head_dim**-0.5
self.qkv_proj = QKVParallelLinear( self.qkv_proj = QKVParallelLinear(
@@ -559,7 +562,7 @@ class SiglipAttention(nn.Module):
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
qkv, _ = self.qkv_proj(hidden_states) qkv, _ = self.qkv_proj(hidden_states)
q, k, v = qkv.split([self.hidden_size // get_tp_world_size()] * 3, dim=-1) q, k, v = qkv.split([self.embed_dim_per_partition] * 3, dim=-1)
batch_size, seq_len, _ = q.shape batch_size, seq_len, _ = q.shape
q = q.view(batch_size, seq_len, self.num_heads_per_partition, self.head_dim) q = q.view(batch_size, seq_len, self.num_heads_per_partition, self.head_dim)
@@ -569,7 +572,7 @@ class SiglipAttention(nn.Module):
attn_output = self.attn(q, k, v) attn_output = self.attn(q, k, v)
attn_output = attn_output.reshape( attn_output = attn_output.reshape(
batch_size, seq_len, self.hidden_size // get_tp_world_size() batch_size, seq_len, self.embed_dim_per_partition
) )
output, _ = self.out_proj(attn_output) output, _ = self.out_proj(attn_output)
@@ -30,7 +30,6 @@ import torch.nn.functional as F
from torch import nn from torch import nn
from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput, T5Config from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput, T5Config
from sglang.multimodal_gen.runtime.distributed import get_sp_group, get_tp_group
from sglang.multimodal_gen.runtime.layers.activation import get_act_fn from sglang.multimodal_gen.runtime.layers.activation import get_act_fn
from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm
from sglang.multimodal_gen.runtime.layers.linear import ( from sglang.multimodal_gen.runtime.layers.linear import (
@@ -44,23 +43,13 @@ from sglang.multimodal_gen.runtime.layers.vocab_parallel_embedding import (
VocabParallelEmbedding, VocabParallelEmbedding,
) )
from sglang.multimodal_gen.runtime.loader.weight_utils import default_weight_loader from sglang.multimodal_gen.runtime.loader.weight_utils import default_weight_loader
from sglang.multimodal_gen.runtime.models.encoders.base import TextEncoder from sglang.multimodal_gen.runtime.models.encoders.base import (
TextEncoder,
get_folding_tp_group,
)
from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.platforms import current_platform
def _get_folding_tp_group(
config: T5Config,
) -> torch.distributed.ProcessGroup | None:
if config.parallel_folding:
if config.parallel_folding_mode == "sp":
return get_sp_group()
elif config.parallel_folding_mode == "ulysses":
return get_sp_group().ulysses_group
elif config.parallel_folding_mode == "ring":
return get_sp_group().ring_group
return get_tp_group()
class AttentionType: class AttentionType:
""" """
Attention type. Attention type.
@@ -88,7 +77,7 @@ class T5DenseActDense(nn.Module):
self, config: T5Config, quant_config: QuantizationConfig | None = None self, config: T5Config, quant_config: QuantizationConfig | None = None
): ):
super().__init__() super().__init__()
tp_group = _get_folding_tp_group(config) tp_group = get_folding_tp_group(config)
self.wi = MergedColumnParallelLinear( self.wi = MergedColumnParallelLinear(
config.d_model, [config.d_ff], bias=False, tp_group=tp_group config.d_model, [config.d_ff], bias=False, tp_group=tp_group
) )
@@ -114,7 +103,7 @@ class T5DenseGatedActDense(nn.Module):
self, config: T5Config, quant_config: QuantizationConfig | None = None self, config: T5Config, quant_config: QuantizationConfig | None = None
): ):
super().__init__() super().__init__()
tp_group = _get_folding_tp_group(config) tp_group = get_folding_tp_group(config)
self.wi_0 = MergedColumnParallelLinear( self.wi_0 = MergedColumnParallelLinear(
config.d_model, config.d_model,
[config.d_ff], [config.d_ff],
@@ -210,7 +199,7 @@ class T5Attention(nn.Module):
self.total_num_heads = self.total_num_kv_heads = config.num_heads self.total_num_heads = self.total_num_kv_heads = config.num_heads
# Partition heads across multiple tensor parallel GPUs. # Partition heads across multiple tensor parallel GPUs.
self.tp_group = _get_folding_tp_group(config) self.tp_group = get_folding_tp_group(config)
self.tp_world_size = get_group_size(self.tp_group) self.tp_world_size = get_group_size(self.tp_group)
assert config.num_heads % self.tp_world_size == 0 assert config.num_heads % self.tp_world_size == 0
self.n_heads = config.num_heads // self.tp_world_size self.n_heads = config.num_heads // self.tp_world_size
@@ -584,7 +573,7 @@ class T5EncoderModel(TextEncoder):
super().__init__(config) super().__init__(config)
quant_config = None quant_config = None
tp_group = _get_folding_tp_group(config) tp_group = get_folding_tp_group(config)
self.shared = VocabParallelEmbedding( self.shared = VocabParallelEmbedding(
config.vocab_size, config.vocab_size,
config.d_model, config.d_model,
@@ -673,7 +662,7 @@ class UMT5EncoderModel(TextEncoder):
super().__init__(config) super().__init__(config)
quant_config = None quant_config = None
tp_group = _get_folding_tp_group(config) tp_group = get_folding_tp_group(config)
self.shared = VocabParallelEmbedding( self.shared = VocabParallelEmbedding(
config.vocab_size, config.vocab_size,
config.d_model, config.d_model,
@@ -20,7 +20,6 @@ import addict
import yaml import yaml
from sglang.multimodal_gen import envs from sglang.multimodal_gen import envs
from sglang.multimodal_gen.configs.models.encoders import T5Config
from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig
from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import ( from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
LTX2PipelineConfig, LTX2PipelineConfig,
@@ -442,23 +441,43 @@ class ServerArgs(DisaggServerArgsMixin):
self.nunchaku_config = resolution.nunchaku_config self.nunchaku_config = resolution.nunchaku_config
def adjust_pipeline_config(self): def adjust_pipeline_config(self):
# enable parallel folding when SP is enabled # 1. adjust for encoder parallel folding
if self.tp_size != 1 or self.sp_degree <= 1: 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:
mode = "world"
elif tp_size == 1 and sp_degree > 1:
# Preserve prior behavior for dp>1 / disaggregated SP runs.
mode = "sp"
else:
return return
enabled = False # Propose the fold group from the parallelism for every encoder. The
for text_encoder_config in self.pipeline_config.text_encoder_configs: # loader keeps it only for encoders wide enough to benefit at their real
if isinstance(text_encoder_config, T5Config): # (post-load) size and whose dims divide the group -- see
text_encoder_config.parallel_folding = True # finalize_encoder_folding. Deciding on real size (not architecture)
enabled = True # handles the same encoder family at different parameter counts.
text_encoder_config.parallel_folding_mode = "sp" encoder_configs = list(self.pipeline_config.text_encoder_configs) + list(
getattr(self.pipeline_config, "image_encoder_configs", ()) or ()
)
for encoder_config in encoder_configs:
encoder_config.parallel_folding_mode = mode
if enabled:
logger.info( logger.info(
"Enabled T5 text encoder parallel folding (mode=sp) for %s (tp_size=%s, sp_degree=%s).", "Proposed encoder parallel folding (mode=%s) for %s "
"(tp=%s sp=%s cfg=%s replica=%s); the loader keeps it for encoders "
"wide enough to benefit.",
mode,
self.__class__.__name__, self.__class__.__name__,
self.tp_size, tp_size,
self.sp_degree, sp_degree,
self.cfg_parallel_degree or 1,
replica_size,
) )
def _adjust_offload(self): def _adjust_offload(self):
@@ -66,6 +66,7 @@ from sglang.multimodal_gen.test.single_test_file.component_accuracy.utils import
generate_name_candidates, generate_name_candidates,
initialize_parallel_runtime, initialize_parallel_runtime,
load_checkpoint_weights, load_checkpoint_weights,
load_param_with_weight_loader,
materialize_module, materialize_module,
read_json_file, read_json_file,
resolve_text_encoder_forward_module, resolve_text_encoder_forward_module,
@@ -498,7 +499,13 @@ class AccuracyEngine:
shard_context.world_size if shard_context is not None else tp_world shard_context.world_size if shard_context is not None else tp_world
) )
shard_rank = shard_context.rank if shard_context is not None else rank shard_rank = shard_context.rank if shard_context is not None else rank
if copy_tensor(tensor, src_tensor, shard_world_size, shard_rank): # TP-sharded params must load via their own weight_loader; the
# generic narrow mis-slices fused QKV/gate_up projections.
if shard_world_size > 1 and load_param_with_weight_loader(
tensor, name, lookup, reverse_mapping
):
matched += 1
elif copy_tensor(tensor, src_tensor, shard_world_size, shard_rank):
matched += 1 matched += 1
else: else:
unmatched_details.append( unmatched_details.append(
@@ -866,3 +866,53 @@ def run_native_component_accuracy_case(
def run_text_encoder_accuracy_case(engine_cls: Any, case: Any, num_gpus: int) -> None: def run_text_encoder_accuracy_case(engine_cls: Any, case: Any, num_gpus: int) -> None:
_run_staged_text_encoder_accuracy_case(engine_cls, case, num_gpus) _run_staged_text_encoder_accuracy_case(engine_cls, case, num_gpus)
def _find_split_shard_parts(lookup, cand):
"""(shard_id, tensor) pairs when a fused target param ships split in the
source state dict (q/k/v or gate/up)."""
if "qkv_proj" in cand:
for repl in ("q_proj", "q"):
q_name = cand.replace("qkv_proj", repl)
k_name = q_name.replace(".q_proj", ".k_proj").replace(".q", ".k")
v_name = q_name.replace(".q_proj", ".v_proj").replace(".q", ".v")
if q_name in lookup and k_name in lookup and v_name in lookup:
return [
("q", lookup[q_name]),
("k", lookup[k_name]),
("v", lookup[v_name]),
]
if "gate_up_proj" in cand:
for gate_token, up_token in (("gate_proj", "up_proj"), ("wi_0", "wi_1")):
gate_name = cand.replace("gate_up_proj", gate_token)
up_name = cand.replace("gate_up_proj", up_token)
if gate_name in lookup and up_name in lookup:
return [(0, lookup[gate_name]), (1, lookup[up_name])]
return None
def load_param_with_weight_loader(param, name, lookup, reverse_mapping) -> bool:
"""Route the source tensor through the parameter's own ``weight_loader`` so
TP sharding matches production checkpoint loading. The generic narrow in
``copy_tensor`` mis-slices fused QKV/gate_up weights (it splits the fused
dim evenly instead of per-projection), which corrupts any TP/folded module.
Any failure falls back to the legacy path."""
loader = getattr(param, "weight_loader", None)
if loader is None or getattr(param, "device_mesh", None) is not None:
return False
try:
candidates = generate_name_candidates(name, reverse_mapping)
for cand in candidates:
parts = _find_split_shard_parts(lookup, cand)
if parts is not None:
for shard_id, tensor in parts:
loader(param, tensor.to(dtype=param.dtype), shard_id)
return True
for cand in candidates:
src = lookup.get(cand)
if src is not None:
loader(param, src.to(dtype=param.dtype))
return True
except Exception:
return False
return False
@@ -34,7 +34,7 @@ if TYPE_CHECKING:
logger = init_logger(__name__) logger = init_logger(__name__)
SGL_TEST_FILES_CI_DATA_REVISION = "46b9b53a429606cb6739c861f275c1277c314a10" SGL_TEST_FILES_CI_DATA_REVISION = "916cbff23aa4e89f78128397ede7ce29a73d6d8c"
if current_platform.is_npu(): if current_platform.is_npu():
SGL_TEST_FILES_CI_DATA_REVISION = "670d66a8a290b62c0c3c077b3e9b0f4a4d9a44e7" SGL_TEST_FILES_CI_DATA_REVISION = "670d66a8a290b62c0c3c077b3e9b0f4a4d9a44e7"
@@ -0,0 +1,166 @@
"""Unit test for the encoder parallel-folding decision (two stages).
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.
"""
from types import SimpleNamespace
from sglang.multimodal_gen.configs.models.encoders import (
EncoderConfig,
ImageEncoderConfig,
TextEncoderConfig,
)
from sglang.multimodal_gen.configs.models.encoders.t5 import T5Config
from sglang.multimodal_gen.runtime.models.encoders.base import (
FOLD_MIN_HIDDEN_SIZE,
encoder_folding_worthwhile,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
def _run(encoders, tp, sp, cfg, dp=1, disagg=False, num_gpus=None, image=()):
self = SimpleNamespace(
tp_size=tp,
sp_degree=sp,
cfg_parallel_degree=cfg,
dp_size=dp,
disagg_mode=disagg,
num_gpus=num_gpus if num_gpus is not None else tp * sp * cfg * dp,
pipeline_config=SimpleNamespace(
text_encoder_configs=tuple(encoders),
image_encoder_configs=tuple(image),
),
)
ServerArgs.adjust_pipeline_config(self)
def _proposed_mode(tp, sp, cfg, dp=1, disagg=False, num_gpus=None):
enc = T5Config()
enc.parallel_folding_mode = None
_run([enc], tp, sp, cfg, dp=dp, disagg=disagg, num_gpus=num_gpus)
return enc.parallel_folding_mode
# --- stage 1: adjust proposes a fold group from the parallelism --------------
def test_pure_tp_not_folded():
# replica == tp: encoder already uses every replica GPU; nothing to fold.
assert _proposed_mode(tp=2, sp=1, cfg=1) is None
def test_single_gpu_not_folded():
assert _proposed_mode(tp=1, sp=1, cfg=1) is None
def test_cfg_parallel_proposes_world():
assert _proposed_mode(tp=1, sp=1, cfg=2) == "world"
def test_sp_dp1_proposes_world():
assert _proposed_mode(tp=1, sp=2, cfg=1) == "world"
def test_tp_times_cfg_proposes_world():
assert _proposed_mode(tp=2, sp=1, cfg=2) == "world"
def test_sp_times_cfg_proposes_world():
assert _proposed_mode(tp=1, sp=2, cfg=2) == "world"
def test_dp_gt_1_keeps_sp():
# dp>1: world group spans replicas, so fold over the per-replica SP group.
assert _proposed_mode(tp=1, sp=2, cfg=1, dp=2) == "sp"
# dp>1 pure cfg has no SP group to fall back to -> nothing proposed.
assert _proposed_mode(tp=1, sp=1, cfg=2, dp=2) is None
def test_disagg_keeps_sp():
assert _proposed_mode(tp=1, sp=2, cfg=2, disagg=True) == "sp"
def test_num_gpus_is_authoritative_for_replica():
assert _proposed_mode(tp=2, sp=1, cfg=1, num_gpus=8) == "world"
def test_all_encoders_get_the_same_proposed_mode():
# adjust no longer size-gates: every encoder in the pipeline (text + image)
# gets the proposed group; the loader trims it later by real size.
t5 = T5Config()
clip = TextEncoderConfig()
img = ImageEncoderConfig()
for e in (t5, clip, img):
e.parallel_folding_mode = None
_run([t5, clip], tp=1, sp=2, cfg=1, image=[img])
assert t5.parallel_folding_mode == "world"
assert clip.parallel_folding_mode == "world"
assert img.parallel_folding_mode == "world"
# --- stage 2: size + divisibility gate (loader, on real dims) ----------------
def _enc(hidden, heads, inter):
enc = TextEncoderConfig()
enc.hidden_size = hidden
enc.num_attention_heads = heads
enc.intermediate_size = inter
return enc
def test_wide_encoder_worth_folding():
# T5-XXL / Mistral-24B class: hidden >= threshold and dims divide the group.
assert encoder_folding_worthwhile(_enc(4096, 64, 10240), group_size=2) is True
assert encoder_folding_worthwhile(_enc(5120, 32, 32768), group_size=2) is True
def test_narrow_encoder_not_worth_folding():
# Qwen3 (2560) measured a net loss -> below the bar.
assert encoder_folding_worthwhile(_enc(2560, 32, 9728), group_size=2) is False
def test_tiny_encoder_not_worth_folding():
# CLIP-L (512): far too small.
assert encoder_folding_worthwhile(_enc(512, 8, 2048), group_size=2) is False
def test_indivisible_dims_not_folded():
# wide enough but heads/intermediate do not divide the group -> cannot shard.
assert encoder_folding_worthwhile(_enc(4096, 6, 10240), group_size=4) is False
assert encoder_folding_worthwhile(_enc(4096, 64, 10250), group_size=4) is False
def test_group_size_one_not_folded():
assert encoder_folding_worthwhile(_enc(4096, 64, 10240), group_size=1) is False
def test_unknown_dims_not_folded():
# a bare encoder whose dims we cannot introspect is left replicated (safe).
assert encoder_folding_worthwhile(TextEncoderConfig(), group_size=2) is False
def test_threshold_is_the_boundary():
assert encoder_folding_worthwhile(_enc(FOLD_MIN_HIDDEN_SIZE, 8, 8192), 2) is True
assert (
encoder_folding_worthwhile(_enc(FOLD_MIN_HIDDEN_SIZE - 128, 8, 8192), 2)
is False
)
# --- config defaults ---------------------------------------------------------
def test_parallel_folding_mode_defaults_none():
assert EncoderConfig().parallel_folding_mode is None
assert TextEncoderConfig().parallel_folding_mode is None
assert ImageEncoderConfig().parallel_folding_mode is None
assert T5Config().parallel_folding_mode is None