From fbebfdec9a6486d43ae24e88e00f776c1149afe1 Mon Sep 17 00:00:00 2001 From: Mick Date: Wed, 6 May 2026 14:55:51 +0800 Subject: [PATCH] [diffusion] fix: fix diffusion FSDP sharding (#24431) --- .../configs/models/dits/ernie_image.py | 9 +- .../configs/models/dits/helios.py | 7 +- .../configs/models/dits/hunyuanvideo.py | 22 +--- .../configs/models/dits/joy_image.py | 9 +- .../configs/models/dits/ltx_2.py | 9 +- .../configs/models/dits/mova_audio.py | 7 +- .../configs/models/dits/mova_video.py | 7 +- .../configs/models/dits/qwenimage.py | 3 + .../configs/models/dits/wanvideo.py | 7 +- .../configs/models/dits/zimage.py | 12 +- .../configs/models/encoders/clip.py | 14 +- .../configs/models/encoders/flux_2.py | 8 +- .../configs/models/encoders/gemma2.py | 19 +-- .../configs/models/encoders/gemma_3.py | 19 +-- .../configs/models/encoders/llama.py | 19 +-- .../configs/models/encoders/mistral3.py | 19 +-- .../configs/models/encoders/qwen3.py | 19 +-- .../configs/models/encoders/qwen3vl.py | 19 +-- .../configs/models/encoders/qwen_image.py | 19 +-- .../configs/models/encoders/t5.py | 23 ++-- .../multimodal_gen/configs/models/fsdp.py | 80 ++++++++++++ .../loader/component_loaders/bridge_loader.py | 9 +- .../runtime/loader/fsdp_load.py | 120 +++++++++++++++++- .../runtime/models/dits/qwen_image.py | 1 + .../pipelines/comfyui_qwen_image_pipeline.py | 4 +- .../pipelines/comfyui_zimage_pipeline.py | 6 +- .../pipelines_core/stages/denoising.py | 71 +++++++++-- .../test/server/perf_baselines.json | 36 +++--- .../test/unit/test_server_args.py | 43 +++++++ 29 files changed, 407 insertions(+), 233 deletions(-) create mode 100644 python/sglang/multimodal_gen/configs/models/fsdp.py diff --git a/python/sglang/multimodal_gen/configs/models/dits/ernie_image.py b/python/sglang/multimodal_gen/configs/models/dits/ernie_image.py index 91aef8162..ee78fbe44 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/ernie_image.py +++ b/python/sglang/multimodal_gen/configs/models/dits/ernie_image.py @@ -4,10 +4,7 @@ from dataclasses import dataclass, field from typing import Tuple from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig - - -def _is_transformer_layer(n: str, m) -> bool: - return "layers" in n and str.isdigit(n.split(".")[-1]) +from sglang.multimodal_gen.configs.models.fsdp import is_layer @dataclass @@ -34,9 +31,7 @@ class ErnieImageArchConfig(DiTArchConfig): } ) - _fsdp_shard_conditions: list = field( - default_factory=lambda: [_is_transformer_layer] - ) + _fsdp_shard_conditions: list = field(default_factory=lambda: [is_layer]) def __post_init__(self): super().__post_init__() diff --git a/python/sglang/multimodal_gen/configs/models/dits/helios.py b/python/sglang/multimodal_gen/configs/models/dits/helios.py index 15f73fb02..f7d417450 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/helios.py +++ b/python/sglang/multimodal_gen/configs/models/dits/helios.py @@ -2,15 +2,12 @@ from dataclasses import dataclass, field from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig - - -def is_blocks(n: str, m) -> bool: - return "blocks" in n and str.isdigit(n.split(".")[-1]) +from sglang.multimodal_gen.configs.models.fsdp import is_block @dataclass class HeliosArchConfig(DiTArchConfig): - _fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks]) + _fsdp_shard_conditions: list = field(default_factory=lambda: [is_block]) param_names_mapping: dict = field( default_factory=lambda: { diff --git a/python/sglang/multimodal_gen/configs/models/dits/hunyuanvideo.py b/python/sglang/multimodal_gen/configs/models/dits/hunyuanvideo.py index 1cae921ff..b30d66dae 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/hunyuanvideo.py +++ b/python/sglang/multimodal_gen/configs/models/dits/hunyuanvideo.py @@ -6,22 +6,12 @@ from dataclasses import dataclass, field import torch from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig - - -def is_double_block(n: str, m) -> bool: - return "double" in n and str.isdigit(n.split(".")[-1]) - - -def is_single_block(n: str, m) -> bool: - return "single" in n and str.isdigit(n.split(".")[-1]) - - -def is_refiner_block(n: str, m) -> bool: - return "refiner" in n and str.isdigit(n.split(".")[-1]) - - -def is_txt_in(n: str, m) -> bool: - return n.split(".")[-1] == "txt_in" +from sglang.multimodal_gen.configs.models.fsdp import ( + is_double_block, + is_refiner_block, + is_single_block, + is_txt_in, +) @dataclass diff --git a/python/sglang/multimodal_gen/configs/models/dits/joy_image.py b/python/sglang/multimodal_gen/configs/models/dits/joy_image.py index 7e9ee6006..b51a44c00 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/joy_image.py +++ b/python/sglang/multimodal_gen/configs/models/dits/joy_image.py @@ -3,15 +3,14 @@ from dataclasses import dataclass, field from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig - - -def is_blocks(n: str, m) -> bool: - return "blocks" in n and str.isdigit(n.split(".")[-1]) +from sglang.multimodal_gen.configs.models.fsdp import is_blocks_or_double_blocks @dataclass class JoyImageArchConfig(DiTArchConfig): - _fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks]) + _fsdp_shard_conditions: list = field( + default_factory=lambda: [is_blocks_or_double_blocks] + ) param_names_mapping: dict = field( default_factory=lambda: { diff --git a/python/sglang/multimodal_gen/configs/models/dits/ltx_2.py b/python/sglang/multimodal_gen/configs/models/dits/ltx_2.py index 5378883a1..f0318a559 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/ltx_2.py +++ b/python/sglang/multimodal_gen/configs/models/dits/ltx_2.py @@ -3,6 +3,7 @@ from dataclasses import dataclass, field from enum import Enum from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig +from sglang.multimodal_gen.configs.models.fsdp import is_blocks_or_transformer_blocks class LTXModelType(Enum): @@ -47,15 +48,13 @@ class LTX2AttentionFunction(str, Enum): DEFAULT = "default" -def is_blocks(n: str, m) -> bool: - return "blocks" in n and str.isdigit(n.split(".")[-1]) - - @dataclass class LTX2ArchConfig(DiTArchConfig): """Architecture configuration for LTX-2 Video Transformer.""" - _fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks]) + _fsdp_shard_conditions: list = field( + default_factory=lambda: [is_blocks_or_transformer_blocks] + ) param_names_mapping: dict = field( default_factory=lambda: { diff --git a/python/sglang/multimodal_gen/configs/models/dits/mova_audio.py b/python/sglang/multimodal_gen/configs/models/dits/mova_audio.py index 1240d5f44..4f056d620 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/mova_audio.py +++ b/python/sglang/multimodal_gen/configs/models/dits/mova_audio.py @@ -4,15 +4,12 @@ from dataclasses import dataclass, field from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig - - -def _is_blocks(n: str, m) -> bool: - return "blocks" in n and str.isdigit(n.split(".")[-1]) +from sglang.multimodal_gen.configs.models.fsdp import is_block @dataclass class MOVAAudioArchConfig(DiTArchConfig): - _fsdp_shard_conditions: list = field(default_factory=lambda: [_is_blocks]) + _fsdp_shard_conditions: list = field(default_factory=lambda: [is_block]) param_names_mapping: dict = field( default_factory=lambda: { diff --git a/python/sglang/multimodal_gen/configs/models/dits/mova_video.py b/python/sglang/multimodal_gen/configs/models/dits/mova_video.py index 66156be4f..0606dbd55 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/mova_video.py +++ b/python/sglang/multimodal_gen/configs/models/dits/mova_video.py @@ -4,15 +4,12 @@ from dataclasses import dataclass, field from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig - - -def _is_blocks(n: str, m) -> bool: - return "blocks" in n and str.isdigit(n.split(".")[-1]) +from sglang.multimodal_gen.configs.models.fsdp import is_block @dataclass class MOVAVideoArchConfig(DiTArchConfig): - _fsdp_shard_conditions: list = field(default_factory=lambda: [_is_blocks]) + _fsdp_shard_conditions: list = field(default_factory=lambda: [is_block]) param_names_mapping: dict = field( default_factory=lambda: { diff --git a/python/sglang/multimodal_gen/configs/models/dits/qwenimage.py b/python/sglang/multimodal_gen/configs/models/dits/qwenimage.py index cb1b15b39..aaaf6e52e 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/qwenimage.py +++ b/python/sglang/multimodal_gen/configs/models/dits/qwenimage.py @@ -5,6 +5,7 @@ from dataclasses import dataclass, field from typing import Tuple from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig +from sglang.multimodal_gen.configs.models.fsdp import is_transformer_block @dataclass @@ -22,6 +23,8 @@ class QwenImageArchConfig(DiTArchConfig): axes_dims_rope: Tuple[int, int, int] = (16, 56, 56) zero_cond_t: bool = False + _fsdp_shard_conditions: list = field(default_factory=lambda: [is_transformer_block]) + stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=list) param_names_mapping: dict = field( diff --git a/python/sglang/multimodal_gen/configs/models/dits/wanvideo.py b/python/sglang/multimodal_gen/configs/models/dits/wanvideo.py index b5c8d6939..544240d8a 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/wanvideo.py +++ b/python/sglang/multimodal_gen/configs/models/dits/wanvideo.py @@ -4,15 +4,12 @@ from dataclasses import dataclass, field from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig - - -def is_blocks(n: str, m) -> bool: - return "blocks" in n and str.isdigit(n.split(".")[-1]) +from sglang.multimodal_gen.configs.models.fsdp import is_block @dataclass class WanVideoArchConfig(DiTArchConfig): - _fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks]) + _fsdp_shard_conditions: list = field(default_factory=lambda: [is_block]) param_names_mapping: dict = field( default_factory=lambda: { diff --git a/python/sglang/multimodal_gen/configs/models/dits/zimage.py b/python/sglang/multimodal_gen/configs/models/dits/zimage.py index 33c50e0cb..5398b07b4 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/zimage.py +++ b/python/sglang/multimodal_gen/configs/models/dits/zimage.py @@ -5,17 +5,7 @@ from dataclasses import dataclass, field from typing import Tuple from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig - - -def is_zimage_layer(n: str, m) -> bool: - """Returns if the module should be sharded for Z-Image model.""" - if "layers" in n and str.isdigit(n.split(".")[-1]): - return True - if ("noise_refiner" in n or "context_refiner" in n) and str.isdigit( - n.split(".")[-1] - ): - return True - return False +from sglang.multimodal_gen.configs.models.fsdp import is_zimage_layer @dataclass diff --git a/python/sglang/multimodal_gen/configs/models/encoders/clip.py b/python/sglang/multimodal_gen/configs/models/encoders/clip.py index ff9a90b32..7cd89b697 100644 --- a/python/sglang/multimodal_gen/configs/models/encoders/clip.py +++ b/python/sglang/multimodal_gen/configs/models/encoders/clip.py @@ -9,17 +9,13 @@ from sglang.multimodal_gen.configs.models.encoders.base import ( TextEncoderArchConfig, TextEncoderConfig, ) +from sglang.multimodal_gen.configs.models.fsdp import ( + is_embeddings, + is_layer, +) from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum -def _is_transformer_layer(n: str, m) -> bool: - return "layers" in n and str.isdigit(n.split(".")[-1]) - - -def _is_embeddings(n: str, m) -> bool: - return n.endswith("embeddings") - - @dataclass class CLIPTextArchConfig(TextEncoderArchConfig): vocab_size: int = 49408 @@ -53,7 +49,7 @@ class CLIPTextArchConfig(TextEncoderArchConfig): ] ) _fsdp_shard_conditions: list = field( - default_factory=lambda: [_is_transformer_layer, _is_embeddings] + default_factory=lambda: [is_layer, is_embeddings] ) diff --git a/python/sglang/multimodal_gen/configs/models/encoders/flux_2.py b/python/sglang/multimodal_gen/configs/models/encoders/flux_2.py index 1dd870f20..40ce6e3f1 100644 --- a/python/sglang/multimodal_gen/configs/models/encoders/flux_2.py +++ b/python/sglang/multimodal_gen/configs/models/encoders/flux_2.py @@ -7,9 +7,7 @@ from sglang.multimodal_gen.configs.models.encoders.base import ( TextEncoderArchConfig, TextEncoderConfig, ) -from sglang.multimodal_gen.configs.models.encoders.qwen_image import ( - _is_transformer_layer, -) +from sglang.multimodal_gen.configs.models.fsdp import is_layer FLUX_2_SYSTEM_MESSAGE = ( "You are an AI that reasons about image descriptions. You give structured responses focusing on object relationships, object\n" @@ -40,9 +38,7 @@ class Flux2MistralTextArchConfig(TextEncoderArchConfig): ("qkv_proj", "v_proj", "v"), ] ) - _fsdp_shard_conditions: list = field( - default_factory=lambda: [_is_transformer_layer] - ) + _fsdp_shard_conditions: list = field(default_factory=lambda: [is_layer]) def __post_init__(self) -> None: self.tokenizer_kwargs = { diff --git a/python/sglang/multimodal_gen/configs/models/encoders/gemma2.py b/python/sglang/multimodal_gen/configs/models/encoders/gemma2.py index 278479f56..0270ca771 100644 --- a/python/sglang/multimodal_gen/configs/models/encoders/gemma2.py +++ b/python/sglang/multimodal_gen/configs/models/encoders/gemma2.py @@ -14,18 +14,11 @@ from sglang.multimodal_gen.configs.models.encoders.base import ( TextEncoderArchConfig, TextEncoderConfig, ) - - -def _is_transformer_layer(n: str, m) -> bool: - return "layers" in n and str.isdigit(n.split(".")[-1]) - - -def _is_embeddings(n: str, m) -> bool: - return n.endswith("embed_tokens") - - -def _is_final_norm(n: str, m) -> bool: - return n.endswith("norm") +from sglang.multimodal_gen.configs.models.fsdp import ( + is_embed_tokens, + is_final_norm, + is_layer, +) @dataclass @@ -77,7 +70,7 @@ class Gemma2ArchConfig(TextEncoderArchConfig): ] ) _fsdp_shard_conditions: list = field( - default_factory=lambda: [_is_transformer_layer, _is_embeddings, _is_final_norm] + default_factory=lambda: [is_layer, is_embed_tokens, is_final_norm] ) diff --git a/python/sglang/multimodal_gen/configs/models/encoders/gemma_3.py b/python/sglang/multimodal_gen/configs/models/encoders/gemma_3.py index 64636985f..dc84c9964 100644 --- a/python/sglang/multimodal_gen/configs/models/encoders/gemma_3.py +++ b/python/sglang/multimodal_gen/configs/models/encoders/gemma_3.py @@ -8,18 +8,11 @@ from sglang.multimodal_gen.configs.models.encoders.base import ( TextEncoderArchConfig, TextEncoderConfig, ) - - -def _is_transformer_layer(n: str, m) -> bool: - return "layers" in n and str.isdigit(n.split(".")[-1]) - - -def _is_embeddings(n: str, m) -> bool: - return n.endswith("embed_tokens") - - -def _is_final_norm(n: str, m) -> bool: - return n.endswith("norm") +from sglang.multimodal_gen.configs.models.fsdp import ( + is_embed_tokens, + is_final_norm, + is_layer, +) @dataclass @@ -70,7 +63,7 @@ class Gemma3ArchConfig(TextEncoderArchConfig): ] ) _fsdp_shard_conditions: list = field( - default_factory=lambda: [_is_transformer_layer, _is_embeddings, _is_final_norm] + default_factory=lambda: [is_layer, is_embed_tokens, is_final_norm] ) diff --git a/python/sglang/multimodal_gen/configs/models/encoders/llama.py b/python/sglang/multimodal_gen/configs/models/encoders/llama.py index 41d98cab2..3f171bb86 100644 --- a/python/sglang/multimodal_gen/configs/models/encoders/llama.py +++ b/python/sglang/multimodal_gen/configs/models/encoders/llama.py @@ -7,18 +7,11 @@ from sglang.multimodal_gen.configs.models.encoders.base import ( TextEncoderArchConfig, TextEncoderConfig, ) - - -def _is_transformer_layer(n: str, m) -> bool: - return "layers" in n and str.isdigit(n.split(".")[-1]) - - -def _is_embeddings(n: str, m) -> bool: - return n.endswith("embed_tokens") - - -def _is_final_norm(n: str, m) -> bool: - return n.endswith("norm") +from sglang.multimodal_gen.configs.models.fsdp import ( + is_embed_tokens, + is_final_norm, + is_layer, +) @dataclass @@ -58,7 +51,7 @@ class LlamaArchConfig(TextEncoderArchConfig): ] ) _fsdp_shard_conditions: list = field( - default_factory=lambda: [_is_transformer_layer, _is_embeddings, _is_final_norm] + default_factory=lambda: [is_layer, is_embed_tokens, is_final_norm] ) diff --git a/python/sglang/multimodal_gen/configs/models/encoders/mistral3.py b/python/sglang/multimodal_gen/configs/models/encoders/mistral3.py index 472ca6df8..7c8f1c969 100644 --- a/python/sglang/multimodal_gen/configs/models/encoders/mistral3.py +++ b/python/sglang/multimodal_gen/configs/models/encoders/mistral3.py @@ -7,18 +7,11 @@ from sglang.multimodal_gen.configs.models.encoders.base import ( TextEncoderArchConfig, TextEncoderConfig, ) - - -def _is_transformer_layer(n: str, m) -> bool: - return "layers" in n and str.isdigit(n.split(".")[-1]) - - -def _is_embeddings(n: str, m) -> bool: - return n.endswith("embed_tokens") - - -def _is_final_norm(n: str, m) -> bool: - return n.endswith("norm") +from sglang.multimodal_gen.configs.models.fsdp import ( + is_embed_tokens, + is_final_norm, + is_layer, +) @dataclass @@ -57,7 +50,7 @@ class Mistral3EncoderArchConfig(TextEncoderArchConfig): ) _fsdp_shard_conditions: list = field( - default_factory=lambda: [_is_transformer_layer, _is_embeddings, _is_final_norm] + default_factory=lambda: [is_layer, is_embed_tokens, is_final_norm] ) def __post_init__(self): diff --git a/python/sglang/multimodal_gen/configs/models/encoders/qwen3.py b/python/sglang/multimodal_gen/configs/models/encoders/qwen3.py index ed48da96a..1909a15c6 100644 --- a/python/sglang/multimodal_gen/configs/models/encoders/qwen3.py +++ b/python/sglang/multimodal_gen/configs/models/encoders/qwen3.py @@ -7,18 +7,11 @@ from sglang.multimodal_gen.configs.models.encoders.base import ( TextEncoderArchConfig, TextEncoderConfig, ) - - -def _is_transformer_layer(n: str, m) -> bool: - return "layers" in n and str.isdigit(n.split(".")[-1]) - - -def _is_embeddings(n: str, m) -> bool: - return n.endswith("embed_tokens") - - -def _is_final_norm(n: str, m) -> bool: - return n.endswith("norm") +from sglang.multimodal_gen.configs.models.fsdp import ( + is_embed_tokens, + is_final_norm, + is_layer, +) @dataclass @@ -66,7 +59,7 @@ class Qwen3TextArchConfig(TextEncoderArchConfig): # FSDP sharding conditions for CPU offload _fsdp_shard_conditions: list = field( - default_factory=lambda: [_is_transformer_layer, _is_embeddings, _is_final_norm] + default_factory=lambda: [is_layer, is_embed_tokens, is_final_norm] ) def __post_init__(self) -> None: diff --git a/python/sglang/multimodal_gen/configs/models/encoders/qwen3vl.py b/python/sglang/multimodal_gen/configs/models/encoders/qwen3vl.py index d23d269fb..c2447235c 100644 --- a/python/sglang/multimodal_gen/configs/models/encoders/qwen3vl.py +++ b/python/sglang/multimodal_gen/configs/models/encoders/qwen3vl.py @@ -6,18 +6,11 @@ from sglang.multimodal_gen.configs.models.encoders.base import ( TextEncoderArchConfig, TextEncoderConfig, ) - - -def _is_transformer_layer(n: str, m) -> bool: - return "layers" in n and str.isdigit(n.split(".")[-1]) - - -def _is_embeddings(n: str, m) -> bool: - return n.endswith("embed_tokens") - - -def _is_final_norm(n: str, m) -> bool: - return n.endswith("norm") +from sglang.multimodal_gen.configs.models.fsdp import ( + is_embed_tokens, + is_final_norm, + is_layer, +) @dataclass @@ -63,7 +56,7 @@ class Qwen3VLArchConfig(TextEncoderArchConfig): ] ) _fsdp_shard_conditions: list = field( - default_factory=lambda: [_is_transformer_layer, _is_embeddings, _is_final_norm] + default_factory=lambda: [is_layer, is_embed_tokens, is_final_norm] ) # JoyImage specific settings diff --git a/python/sglang/multimodal_gen/configs/models/encoders/qwen_image.py b/python/sglang/multimodal_gen/configs/models/encoders/qwen_image.py index 98b1a5085..b22b2bf5f 100644 --- a/python/sglang/multimodal_gen/configs/models/encoders/qwen_image.py +++ b/python/sglang/multimodal_gen/configs/models/encoders/qwen_image.py @@ -7,18 +7,11 @@ from sglang.multimodal_gen.configs.models.encoders.base import ( TextEncoderArchConfig, TextEncoderConfig, ) - - -def _is_transformer_layer(n: str, m) -> bool: - return "layers" in n and str.isdigit(n.split(".")[-1]) - - -def _is_embeddings(n: str, m) -> bool: - return n.endswith("embed_tokens") - - -def _is_final_norm(n: str, m) -> bool: - return n.endswith("norm") +from sglang.multimodal_gen.configs.models.fsdp import ( + is_embed_tokens, + is_final_norm, + is_layer, +) @dataclass @@ -63,7 +56,7 @@ class QwenImageArchConfig(TextEncoderArchConfig): ] ) _fsdp_shard_conditions: list = field( - default_factory=lambda: [_is_transformer_layer, _is_embeddings, _is_final_norm] + default_factory=lambda: [is_layer, is_embed_tokens, is_final_norm] ) diff --git a/python/sglang/multimodal_gen/configs/models/encoders/t5.py b/python/sglang/multimodal_gen/configs/models/encoders/t5.py index 658569088..7de1aae3b 100644 --- a/python/sglang/multimodal_gen/configs/models/encoders/t5.py +++ b/python/sglang/multimodal_gen/configs/models/encoders/t5.py @@ -8,18 +8,11 @@ from sglang.multimodal_gen.configs.models.encoders.base import ( TextEncoderArchConfig, TextEncoderConfig, ) - - -def _is_transformer_layer(n: str, m) -> bool: - return "block" in n and str.isdigit(n.split(".")[-1]) - - -def _is_embeddings(n: str, m) -> bool: - return n.endswith("shared") - - -def _is_final_layernorm(n: str, m) -> bool: - return n.endswith("final_layer_norm") +from sglang.multimodal_gen.configs.models.fsdp import ( + is_final_layer_norm, + is_shared, + is_t5_block, +) @dataclass @@ -55,9 +48,9 @@ class T5ArchConfig(TextEncoderArchConfig): ) _fsdp_shard_conditions: list = field( default_factory=lambda: [ - _is_transformer_layer, - _is_embeddings, - _is_final_layernorm, + is_t5_block, + is_shared, + is_final_layer_norm, ] ) diff --git a/python/sglang/multimodal_gen/configs/models/fsdp.py b/python/sglang/multimodal_gen/configs/models/fsdp.py new file mode 100644 index 000000000..b3a0e5491 --- /dev/null +++ b/python/sglang/multimodal_gen/configs/models/fsdp.py @@ -0,0 +1,80 @@ +# SPDX-License-Identifier: Apache-2.0 + + +def is_module_list_entry(name: str, container_name: str) -> bool: + # Match only direct block entries, not their inner submodules. + parts = name.split(".") + return len(parts) >= 2 and parts[-2] == container_name and parts[-1].isdigit() + + +def is_module_list_entry_in(name: str, container_names: tuple[str, ...]) -> bool: + parts = name.split(".") + return len(parts) >= 2 and parts[-2] in container_names and parts[-1].isdigit() + + +def is_layer(name: str, module: object) -> bool: + return is_module_list_entry(name, "layers") + + +def is_block(name: str, module: object) -> bool: + return is_module_list_entry(name, "blocks") + + +def is_t5_block(name: str, module: object) -> bool: + return is_module_list_entry(name, "block") + + +def is_transformer_block(name: str, module: object) -> bool: + return is_module_list_entry(name, "transformer_blocks") + + +def is_double_block(name: str, module: object) -> bool: + return is_module_list_entry(name, "double_blocks") + + +def is_single_block(name: str, module: object) -> bool: + return is_module_list_entry(name, "single_blocks") + + +def is_refiner_block(name: str, module: object) -> bool: + return is_module_list_entry(name, "refiner_blocks") + + +def is_blocks_or_double_blocks(name: str, module: object) -> bool: + return is_module_list_entry_in(name, ("blocks", "double_blocks")) + + +def is_blocks_or_transformer_blocks(name: str, module: object) -> bool: + return is_module_list_entry_in(name, ("blocks", "transformer_blocks")) + + +def is_zimage_layer(name: str, module: object) -> bool: + last_part = name.split(".")[-1] + # Preserve Z-Image's finer historical FSDP granularity for perf. + return last_part.isdigit() and ( + "layers" in name or "noise_refiner" in name or "context_refiner" in name + ) + + +def is_embed_tokens(name: str, module: object) -> bool: + return name.endswith("embed_tokens") + + +def is_embeddings(name: str, module: object) -> bool: + return name.endswith("embeddings") + + +def is_final_norm(name: str, module: object) -> bool: + return name.endswith("norm") + + +def is_shared(name: str, module: object) -> bool: + return name.endswith("shared") + + +def is_final_layer_norm(name: str, module: object) -> bool: + return name.endswith("final_layer_norm") + + +def is_txt_in(name: str, module: object) -> bool: + return name.split(".")[-1] == "txt_in" diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/bridge_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/bridge_loader.py index 646c23283..bcc2203e5 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/bridge_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/bridge_loader.py @@ -71,11 +71,10 @@ class BridgeLoader(ComponentLoader): default_dtype, ) - # Check if FSDP loading is available - if ( - server_args.hsdp_shard_dim is not None - and hasattr(model_cls, "_fsdp_shard_conditions") - and model_cls._fsdp_shard_conditions + # Use the FSDP loader when FSDP is requested or shard rules are declared. + fsdp_shard_conditions = getattr(model_cls, "_fsdp_shard_conditions", None) + if server_args.use_fsdp_inference or ( + server_args.hsdp_shard_dim is not None and fsdp_shard_conditions ): # Load with FSDP support model = maybe_load_fsdp_model( diff --git a/python/sglang/multimodal_gen/runtime/loader/fsdp_load.py b/python/sglang/multimodal_gen/runtime/loader/fsdp_load.py index f9bdae93e..6a98d58fb 100644 --- a/python/sglang/multimodal_gen/runtime/loader/fsdp_load.py +++ b/python/sglang/multimodal_gen/runtime/loader/fsdp_load.py @@ -23,6 +23,7 @@ from torch.distributed.fsdp import ( ) from torch.nn.modules.module import _IncompatibleKeys +from sglang.multimodal_gen.configs.models.fsdp import is_module_list_entry_in from sglang.multimodal_gen.runtime.layers.linear import UnquantizedLinearMethod from sglang.multimodal_gen.runtime.loader.utils import ( get_param_names_mapping, @@ -79,6 +80,62 @@ def _make_param_like( return new_param +def _get_param_for_weight_loading( + model: torch.nn.Module, + param_dict: dict[str, torch.nn.Parameter], + param_name: str, +) -> torch.nn.Parameter | None: + actual_param = param_dict.get(param_name) + if actual_param is not None and getattr(actual_param, "weight_loader", None): + return actual_param + + pre_fsdp_weight_loader_params = getattr(model, "_pre_fsdp_weight_loader_params", {}) + pre_fsdp_param = pre_fsdp_weight_loader_params.get(param_name) + if pre_fsdp_param is not None: + return pre_fsdp_param + + return actual_param + + +def _make_class_name_shard_condition(class_names: set[str]): + def shard_condition(n: str, m: nn.Module) -> bool: + return type(m).__name__ in class_names + + return shard_condition + + +def _is_common_numbered_block(n: str, m: nn.Module) -> bool: + return is_module_list_entry_in( + n, + ( + "blocks", + "layers", + "double_blocks", + "single_blocks", + "refiner_blocks", + "noise_refiner", + "context_refiner", + "transformer_blocks", + "single_transformer_blocks", + ), + ) + + +def _resolve_fsdp_shard_conditions( + model: torch.nn.Module, + fsdp_shard_conditions: list[Callable[[str, nn.Module], bool]] | None, +) -> tuple[list[Callable[[str, nn.Module], bool]], str]: + if fsdp_shard_conditions: + return fsdp_shard_conditions, "explicit" + + block_class_names = set(getattr(model, "_repeated_blocks", []) or []) + block_class_names.update(getattr(model, "_no_split_modules", []) or []) + if block_class_names: + return [_make_class_name_shard_condition(block_class_names)], "block-class" + + return [_is_common_numbered_block], "common-numbered-block" + + def _maybe_dequantize_fp8( full_tensor: torch.Tensor, target_dtype: torch.dtype, @@ -161,6 +218,11 @@ def maybe_load_fsdp_model( logger.info("Disabling FSDP for MPS platform as it's not compatible") if use_fsdp: + model._pre_fsdp_weight_loader_params = { + n: p + for n, p in model.named_parameters() + if getattr(p, "weight_loader", None) + } world_size = hsdp_replicate_dim * hsdp_shard_dim if not fsdp_inference: hsdp_replicate_dim = world_size @@ -178,7 +240,7 @@ def maybe_load_fsdp_model( reshard_after_forward=True, mp_policy=mp_policy, mesh=device_mesh, - fsdp_shard_conditions=model._fsdp_shard_conditions, + fsdp_shard_conditions=getattr(model, "_fsdp_shard_conditions", None), pin_cpu_memory=pin_cpu_memory, ) @@ -224,7 +286,7 @@ def shard_model( reshard_after_forward: bool = True, mp_policy: MixedPrecisionPolicy | None = MixedPrecisionPolicy(), # noqa mesh: DeviceMesh | None = None, - fsdp_shard_conditions: list[Callable[[str, nn.Module], bool]] = [], # noqa + fsdp_shard_conditions: list[Callable[[str, nn.Module], bool]] | None = None, pin_cpu_memory: bool = True, ) -> None: """ @@ -247,12 +309,15 @@ def shard_model( pin_cpu_memory (bool): If set to True, FSDP will pin the CPU memory of the offloaded parameters. """ - if fsdp_shard_conditions is None or len(fsdp_shard_conditions) == 0: + fsdp_shard_conditions, condition_source = _resolve_fsdp_shard_conditions( + model, fsdp_shard_conditions + ) + if condition_source != "explicit": logger.warning( - "The FSDP shard condition list is empty or None. No modules will be sharded in %s", + "Using %s FSDP shard condition fallback for %s", + condition_source, type(model).__name__, ) - return fsdp_kwargs = { "reshard_after_forward": reshard_after_forward, @@ -274,11 +339,18 @@ def shard_model( if num_layers_sharded == 0: raise ValueError( - "No layer modules were sharded. Please check if shard conditions are working as expected." + f"No layer modules were sharded in {type(model).__name__}. " + f"FSDP shard condition source: {condition_source}." ) # Finally shard the entire model to account for any stragglers fully_shard(model, **fsdp_kwargs) + logger.info( + "Applied FSDP to %d submodules in %s using %s shard conditions", + num_layers_sharded, + type(model).__name__, + condition_source, + ) # TODO(mick): need refactor, to move out checkpoint-specific adjustments @@ -393,7 +465,9 @@ def load_model_from_full_model_state_dict( if not hasattr(meta_sharded_param, "device_mesh"): full_tensor = full_tensor.to(device=device, dtype=target_dtype) - actual_param = param_dict.get(target_param_name) + actual_param = _get_param_for_weight_loading( + model, param_dict, target_param_name + ) weight_loader = ( getattr(actual_param, "weight_loader", None) if actual_param is not None @@ -440,6 +514,38 @@ def load_model_from_full_model_state_dict( sharded_tensor = sharded_tensor.cpu() else: full_tensor = full_tensor.to(device=device, dtype=target_dtype) + actual_param = _get_param_for_weight_loading( + model, param_dict, target_param_name + ) + weight_loader = ( + getattr(actual_param, "weight_loader", None) + if actual_param is not None + else None + ) + if weight_loader is not None: + assert actual_param is not None + tp_sharded_tensor = torch.empty( + tuple(actual_param.shape), + device=device, + dtype=target_dtype, + ) + temp_param = _make_param_like(actual_param, tp_sharded_tensor) + if not ( + tp_sharded_tensor.is_floating_point() + or tp_sharded_tensor.is_complex() + ): + temp_param.requires_grad = False + try: + weight_loader(temp_param, full_tensor) + except AssertionError as exc: + raise AssertionError( + "Failed to TP-shard/load FSDP parameter " + f"{target_param_name}: full_tensor.shape={tuple(full_tensor.shape)}, " + f"meta_sharded_param.shape={tuple(meta_sharded_param.shape)}, " + f"temp_param.shape={tuple(temp_param.shape)}, " + f"param_cls={type(actual_param).__name__}" + ) from exc + full_tensor = temp_param.data sharded_tensor = distribute_tensor( full_tensor, meta_sharded_param.device_mesh, diff --git a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py index c5fe1bf12..3456d66cf 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py @@ -1165,6 +1165,7 @@ class QwenImageTransformer2DModel(CachableDiT, OffloadableDiTMixin): _repeated_blocks = ["QwenImageTransformerBlock"] param_names_mapping = QwenImageDitConfig().arch_config.param_names_mapping + _fsdp_shard_conditions = QwenImageDitConfig().arch_config._fsdp_shard_conditions @classmethod def get_nunchaku_quant_rules(cls) -> dict[str, list[str]]: diff --git a/python/sglang/multimodal_gen/runtime/pipelines/comfyui_qwen_image_pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines/comfyui_qwen_image_pipeline.py index 138f74b28..f1c92d79f 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines/comfyui_qwen_image_pipeline.py +++ b/python/sglang/multimodal_gen/runtime/pipelines/comfyui_qwen_image_pipeline.py @@ -247,7 +247,9 @@ class ComfyUIQwenImagePipelineBase(LoRAPipeline, ComposedPipelineBase): reshard_after_forward=True, mp_policy=mp_policy, mesh=device_mesh, - fsdp_shard_conditions=model._fsdp_shard_conditions, + fsdp_shard_conditions=getattr( + model, "_fsdp_shard_conditions", None + ), pin_cpu_memory=server_args.pin_cpu_memory, ) finally: diff --git a/python/sglang/multimodal_gen/runtime/pipelines/comfyui_zimage_pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines/comfyui_zimage_pipeline.py index efe5e7430..c86230cd7 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines/comfyui_zimage_pipeline.py +++ b/python/sglang/multimodal_gen/runtime/pipelines/comfyui_zimage_pipeline.py @@ -143,7 +143,6 @@ class ComfyUIZImagePipeline(LoRAPipeline, ComposedPipelineBase): head_dim = dim // num_heads q_size = dim k_size = head_dim * num_kv_heads - v_size = head_dim * num_kv_heads for name, tensor in weight_iterator: # Match qkv weights in layers, noise_refiner, or context_refiner @@ -311,7 +310,6 @@ class ComfyUIZImagePipeline(LoRAPipeline, ComposedPipelineBase): logger.info("Disabling FSDP for MPS platform as it's not compatible") if use_fsdp: - world_size = server_args.hsdp_replicate_dim * server_args.hsdp_shard_dim device_mesh = init_device_mesh( current_platform.device_type, mesh_shape=( @@ -326,7 +324,9 @@ class ComfyUIZImagePipeline(LoRAPipeline, ComposedPipelineBase): reshard_after_forward=True, mp_policy=mp_policy, mesh=device_mesh, - fsdp_shard_conditions=model._fsdp_shard_conditions, + fsdp_shard_conditions=getattr( + model, "_fsdp_shard_conditions", None + ), pin_cpu_memory=server_args.pin_cpu_memory, ) 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 b383f3c29..2901b50c7 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -195,6 +195,7 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): self._cache_dit_enabled = False self._cached_num_steps = None self._is_warmed_up = False + self._extra_func_kwarg_names_cache: dict[int, tuple[bool, frozenset[str]]] = {} def _infer_transformer_attention_backend(self) -> AttentionBackendEnum | None: backends = { @@ -699,7 +700,7 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): | server_args.pipeline_config.prepare_pos_cond_kwargs( batch, self.device, - getattr(self.transformer, "rotary_emb", None), + self._get_transformer_attr("rotary_emb"), dtype=target_dtype, ) | dict( @@ -720,7 +721,7 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): | server_args.pipeline_config.prepare_neg_cond_kwargs( batch, self.device, - getattr(self.transformer, "rotary_emb", None), + self._get_transformer_attr("rotary_emb"), dtype=target_dtype, ) | dict( @@ -778,6 +779,25 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): if hasattr(scheduler, "timestep_list"): scheduler.timestep_list = [None] * solver_order + def _get_transformer_attr(self, name: str) -> Any: + seen: set[int] = set() + stack = [self.transformer] + while stack: + module = stack.pop() + if module is None or id(module) in seen: + continue + seen.add(id(module)) + + value = getattr(module, name, None) + if value is not None: + return value + + for wrapper_attr in ("_fsdp_wrapped_module", "module", "_orig_mod"): + wrapped = getattr(module, wrapper_attr, None) + if wrapped is not None: + stack.append(wrapped) + return None + def _prepare_step_state( self, ctx: DenoisingContext, @@ -1270,15 +1290,7 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): self._finalize_denoising_loop(ctx, batch, server_args) return batch - # TODO: this will extends the preparation stage, should let subclass/passed-in variables decide which to prepare - def prepare_extra_func_kwargs(self, func, kwargs) -> dict[str, Any]: - """ - Prepare extra kwargs for the scheduler step / denoise step. - - Args: - func: The function to prepare kwargs for. - kwargs: The kwargs to prepare. - """ + def _get_extra_func_kwarg_names(self, func) -> tuple[bool, frozenset[str]]: import functools # Handle cache-dit's partial wrapping logic. @@ -1290,10 +1302,37 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): # Unwrap any decorators (e.g. functools.wraps) target_func = inspect.unwrap(func) + cache_target = ( + target_func.__func__ if inspect.ismethod(target_func) else target_func + ) + cache_key = id(cache_target) + cached = self._extra_func_kwarg_names_cache.get(cache_key) + if cached is not None: + return cached - # Filter kwargs based on the signature params = inspect.signature(target_func).parameters - return {k: v for k, v in kwargs.items() if k in params} + result = ( + any( + param.kind == inspect.Parameter.VAR_KEYWORD for param in params.values() + ), + frozenset(params), + ) + self._extra_func_kwarg_names_cache[cache_key] = result + return result + + # TODO: this will extends the preparation stage, should let subclass/passed-in variables decide which to prepare + def prepare_extra_func_kwargs(self, func, kwargs) -> dict[str, Any]: + """ + Prepare extra kwargs for the scheduler step / denoise step. + + Args: + func: The function to prepare kwargs for. + kwargs: The kwargs to prepare. + """ + accepts_var_kwargs, param_names = self._get_extra_func_kwarg_names(func) + if accepts_var_kwargs: + return kwargs + return {k: v for k, v in kwargs.items() if k in param_names} def progress_bar( self, iterable: Iterable | None = None, total: int | None = None @@ -1641,10 +1680,14 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): guidance: torch.Tensor, **kwargs, ): + guidance_kwargs = self.prepare_extra_func_kwargs( + getattr(current_model, "forward", current_model), + {"guidance": guidance}, + ) return current_model( hidden_states=latent_model_input, timestep=timestep, - guidance=guidance, + **guidance_kwargs, **kwargs, ) diff --git a/python/sglang/multimodal_gen/test/server/perf_baselines.json b/python/sglang/multimodal_gen/test/server/perf_baselines.json index ffcd75394..648b92510 100644 --- a/python/sglang/multimodal_gen/test/server/perf_baselines.json +++ b/python/sglang/multimodal_gen/test/server/perf_baselines.json @@ -2019,27 +2019,27 @@ }, "fsdp-inference": { "stages_ms": { - "InputValidationStage": 0.05, - "LatentPreparationStage": 0.15, - "TextEncodingStage": 297.4, - "TimestepPreparationStage": 57.84, - "DecodingStage": 8.17, - "DenoisingStage": 2142.47 + "InputValidationStage": 0.06, + "LatentPreparationStage": 0.16, + "TextEncodingStage": 305.97, + "TimestepPreparationStage": 57.19, + "DecodingStage": 16.88, + "DenoisingStage": 2422.53 }, "denoise_step_ms": { - "0": 80.53, - "1": 188.09, - "2": 219.79, - "3": 219.53, - "4": 218.6, - "5": 216.57, - "6": 215.51, - "7": 215.09, - "8": 215.14 + "0": 259.26, + "1": 284.25, + "2": 283.74, + "3": 270.48, + "4": 278.55, + "5": 271.58, + "6": 270.89, + "7": 277.75, + "8": 270.1 }, - "expected_e2e_ms": 2745.03, - "expected_avg_denoise_ms": 237.53, - "expected_median_denoise_ms": 260.88, + "expected_e2e_ms": 2775.88, + "expected_avg_denoise_ms": 268.55, + "expected_median_denoise_ms": 268.51, "estimated_full_test_time_s": 122.7 }, "hunyuan3d_shape_gen": { diff --git a/python/sglang/multimodal_gen/test/unit/test_server_args.py b/python/sglang/multimodal_gen/test/unit/test_server_args.py index 0d3e71554..d0550df1d 100644 --- a/python/sglang/multimodal_gen/test/unit/test_server_args.py +++ b/python/sglang/multimodal_gen/test/unit/test_server_args.py @@ -3,6 +3,11 @@ import sys import unittest from unittest.mock import patch +from sglang.multimodal_gen.configs.models.fsdp import ( + is_module_list_entry, + is_module_list_entry_in, + is_zimage_layer, +) from sglang.multimodal_gen.configs.pipeline_configs.base import ( ModelTaskType, PipelineConfig, @@ -11,6 +16,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import ( QwenImagePipelineConfig, ) from sglang.multimodal_gen.registry import _get_config_info +from sglang.multimodal_gen.runtime.models.dits.qwen_image import ( + QwenImageTransformer2DModel, +) from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.utils import FlexibleArgumentParser @@ -159,6 +167,41 @@ class TestOffloadDefaults(unittest.TestCase): self.assertTrue(args.vae_cpu_offload) +class TestFSDPShardConditions(unittest.TestCase): + def test_helpers_match_only_direct_block_entries(self): + self.assertTrue( + is_module_list_entry("transformer_blocks.0", "transformer_blocks") + ) + self.assertFalse( + is_module_list_entry("transformer_blocks.0.ff.net.0", "transformer_blocks") + ) + self.assertTrue( + is_module_list_entry_in( + "single_transformer_blocks.12", + ("transformer_blocks", "single_transformer_blocks"), + ) + ) + self.assertFalse( + is_module_list_entry_in( + "single_transformer_blocks.12.attn.to_out.0", + ("transformer_blocks", "single_transformer_blocks"), + ) + ) + + def test_qwen_dit_has_fsdp_shard_condition(self): + conditions = QwenImageTransformer2DModel._fsdp_shard_conditions + + self.assertTrue(conditions) + self.assertTrue(conditions[0]("transformer_blocks.0", None)) + self.assertFalse(conditions[0]("transformer_blocks.0.attn", None)) + self.assertFalse(conditions[0]("transformer_blocks.0.ff.net.0", None)) + + def test_zimage_condition_keeps_inner_numbered_modules(self): + self.assertTrue(is_zimage_layer("layers.0.mlp.0", None)) + self.assertTrue(is_zimage_layer("noise_refiner.0.attention.to_out.0", None)) + self.assertFalse(is_zimage_layer("transformer_blocks.0", None)) + + class TestModelIdResolution(unittest.TestCase): def setUp(self): _get_config_info.cache_clear()