[diffusion] fix: fix diffusion FSDP sharding (#24431)

This commit is contained in:
Mick
2026-05-06 14:55:51 +08:00
committed by GitHub
parent 660a77f221
commit fbebfdec9a
29 changed files with 407 additions and 233 deletions
@@ -4,10 +4,7 @@ from dataclasses import dataclass, field
from typing import Tuple from typing import Tuple
from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig
from sglang.multimodal_gen.configs.models.fsdp import is_layer
def _is_transformer_layer(n: str, m) -> bool:
return "layers" in n and str.isdigit(n.split(".")[-1])
@dataclass @dataclass
@@ -34,9 +31,7 @@ class ErnieImageArchConfig(DiTArchConfig):
} }
) )
_fsdp_shard_conditions: list = field( _fsdp_shard_conditions: list = field(default_factory=lambda: [is_layer])
default_factory=lambda: [_is_transformer_layer]
)
def __post_init__(self): def __post_init__(self):
super().__post_init__() super().__post_init__()
@@ -2,15 +2,12 @@
from dataclasses import dataclass, field from dataclasses import dataclass, field
from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig
from sglang.multimodal_gen.configs.models.fsdp import is_block
def is_blocks(n: str, m) -> bool:
return "blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass @dataclass
class HeliosArchConfig(DiTArchConfig): 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( param_names_mapping: dict = field(
default_factory=lambda: { default_factory=lambda: {
@@ -6,22 +6,12 @@ from dataclasses import dataclass, field
import torch import torch
from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig
from sglang.multimodal_gen.configs.models.fsdp import (
is_double_block,
def is_double_block(n: str, m) -> bool: is_refiner_block,
return "double" in n and str.isdigit(n.split(".")[-1]) is_single_block,
is_txt_in,
)
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"
@dataclass @dataclass
@@ -3,15 +3,14 @@
from dataclasses import dataclass, field from dataclasses import dataclass, field
from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig
from sglang.multimodal_gen.configs.models.fsdp import is_blocks_or_double_blocks
def is_blocks(n: str, m) -> bool:
return "blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass @dataclass
class JoyImageArchConfig(DiTArchConfig): 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( param_names_mapping: dict = field(
default_factory=lambda: { default_factory=lambda: {
@@ -3,6 +3,7 @@ from dataclasses import dataclass, field
from enum import Enum from enum import Enum
from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig 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): class LTXModelType(Enum):
@@ -47,15 +48,13 @@ class LTX2AttentionFunction(str, Enum):
DEFAULT = "default" DEFAULT = "default"
def is_blocks(n: str, m) -> bool:
return "blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass @dataclass
class LTX2ArchConfig(DiTArchConfig): class LTX2ArchConfig(DiTArchConfig):
"""Architecture configuration for LTX-2 Video Transformer.""" """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( param_names_mapping: dict = field(
default_factory=lambda: { default_factory=lambda: {
@@ -4,15 +4,12 @@
from dataclasses import dataclass, field from dataclasses import dataclass, field
from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig
from sglang.multimodal_gen.configs.models.fsdp import is_block
def _is_blocks(n: str, m) -> bool:
return "blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass @dataclass
class MOVAAudioArchConfig(DiTArchConfig): 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( param_names_mapping: dict = field(
default_factory=lambda: { default_factory=lambda: {
@@ -4,15 +4,12 @@
from dataclasses import dataclass, field from dataclasses import dataclass, field
from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig
from sglang.multimodal_gen.configs.models.fsdp import is_block
def _is_blocks(n: str, m) -> bool:
return "blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass @dataclass
class MOVAVideoArchConfig(DiTArchConfig): 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( param_names_mapping: dict = field(
default_factory=lambda: { default_factory=lambda: {
@@ -5,6 +5,7 @@ from dataclasses import dataclass, field
from typing import Tuple from typing import Tuple
from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig
from sglang.multimodal_gen.configs.models.fsdp import is_transformer_block
@dataclass @dataclass
@@ -22,6 +23,8 @@ class QwenImageArchConfig(DiTArchConfig):
axes_dims_rope: Tuple[int, int, int] = (16, 56, 56) axes_dims_rope: Tuple[int, int, int] = (16, 56, 56)
zero_cond_t: bool = False 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) stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=list)
param_names_mapping: dict = field( param_names_mapping: dict = field(
@@ -4,15 +4,12 @@
from dataclasses import dataclass, field from dataclasses import dataclass, field
from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig
from sglang.multimodal_gen.configs.models.fsdp import is_block
def is_blocks(n: str, m) -> bool:
return "blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass @dataclass
class WanVideoArchConfig(DiTArchConfig): 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( param_names_mapping: dict = field(
default_factory=lambda: { default_factory=lambda: {
@@ -5,17 +5,7 @@ from dataclasses import dataclass, field
from typing import Tuple from typing import Tuple
from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig
from sglang.multimodal_gen.configs.models.fsdp import is_zimage_layer
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
@dataclass @dataclass
@@ -9,17 +9,13 @@ from sglang.multimodal_gen.configs.models.encoders.base import (
TextEncoderArchConfig, TextEncoderArchConfig,
TextEncoderConfig, TextEncoderConfig,
) )
from sglang.multimodal_gen.configs.models.fsdp import (
is_embeddings,
is_layer,
)
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum 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 @dataclass
class CLIPTextArchConfig(TextEncoderArchConfig): class CLIPTextArchConfig(TextEncoderArchConfig):
vocab_size: int = 49408 vocab_size: int = 49408
@@ -53,7 +49,7 @@ class CLIPTextArchConfig(TextEncoderArchConfig):
] ]
) )
_fsdp_shard_conditions: list = field( _fsdp_shard_conditions: list = field(
default_factory=lambda: [_is_transformer_layer, _is_embeddings] default_factory=lambda: [is_layer, is_embeddings]
) )
@@ -7,9 +7,7 @@ from sglang.multimodal_gen.configs.models.encoders.base import (
TextEncoderArchConfig, TextEncoderArchConfig,
TextEncoderConfig, TextEncoderConfig,
) )
from sglang.multimodal_gen.configs.models.encoders.qwen_image import ( from sglang.multimodal_gen.configs.models.fsdp import is_layer
_is_transformer_layer,
)
FLUX_2_SYSTEM_MESSAGE = ( FLUX_2_SYSTEM_MESSAGE = (
"You are an AI that reasons about image descriptions. You give structured responses focusing on object relationships, object\n" "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"), ("qkv_proj", "v_proj", "v"),
] ]
) )
_fsdp_shard_conditions: list = field( _fsdp_shard_conditions: list = field(default_factory=lambda: [is_layer])
default_factory=lambda: [_is_transformer_layer]
)
def __post_init__(self) -> None: def __post_init__(self) -> None:
self.tokenizer_kwargs = { self.tokenizer_kwargs = {
@@ -14,18 +14,11 @@ from sglang.multimodal_gen.configs.models.encoders.base import (
TextEncoderArchConfig, TextEncoderArchConfig,
TextEncoderConfig, TextEncoderConfig,
) )
from sglang.multimodal_gen.configs.models.fsdp import (
is_embed_tokens,
def _is_transformer_layer(n: str, m) -> bool: is_final_norm,
return "layers" in n and str.isdigit(n.split(".")[-1]) is_layer,
)
def _is_embeddings(n: str, m) -> bool:
return n.endswith("embed_tokens")
def _is_final_norm(n: str, m) -> bool:
return n.endswith("norm")
@dataclass @dataclass
@@ -77,7 +70,7 @@ class Gemma2ArchConfig(TextEncoderArchConfig):
] ]
) )
_fsdp_shard_conditions: list = field( _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]
) )
@@ -8,18 +8,11 @@ from sglang.multimodal_gen.configs.models.encoders.base import (
TextEncoderArchConfig, TextEncoderArchConfig,
TextEncoderConfig, TextEncoderConfig,
) )
from sglang.multimodal_gen.configs.models.fsdp import (
is_embed_tokens,
def _is_transformer_layer(n: str, m) -> bool: is_final_norm,
return "layers" in n and str.isdigit(n.split(".")[-1]) is_layer,
)
def _is_embeddings(n: str, m) -> bool:
return n.endswith("embed_tokens")
def _is_final_norm(n: str, m) -> bool:
return n.endswith("norm")
@dataclass @dataclass
@@ -70,7 +63,7 @@ class Gemma3ArchConfig(TextEncoderArchConfig):
] ]
) )
_fsdp_shard_conditions: list = field( _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]
) )
@@ -7,18 +7,11 @@ from sglang.multimodal_gen.configs.models.encoders.base import (
TextEncoderArchConfig, TextEncoderArchConfig,
TextEncoderConfig, TextEncoderConfig,
) )
from sglang.multimodal_gen.configs.models.fsdp import (
is_embed_tokens,
def _is_transformer_layer(n: str, m) -> bool: is_final_norm,
return "layers" in n and str.isdigit(n.split(".")[-1]) is_layer,
)
def _is_embeddings(n: str, m) -> bool:
return n.endswith("embed_tokens")
def _is_final_norm(n: str, m) -> bool:
return n.endswith("norm")
@dataclass @dataclass
@@ -58,7 +51,7 @@ class LlamaArchConfig(TextEncoderArchConfig):
] ]
) )
_fsdp_shard_conditions: list = field( _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]
) )
@@ -7,18 +7,11 @@ from sglang.multimodal_gen.configs.models.encoders.base import (
TextEncoderArchConfig, TextEncoderArchConfig,
TextEncoderConfig, TextEncoderConfig,
) )
from sglang.multimodal_gen.configs.models.fsdp import (
is_embed_tokens,
def _is_transformer_layer(n: str, m) -> bool: is_final_norm,
return "layers" in n and str.isdigit(n.split(".")[-1]) is_layer,
)
def _is_embeddings(n: str, m) -> bool:
return n.endswith("embed_tokens")
def _is_final_norm(n: str, m) -> bool:
return n.endswith("norm")
@dataclass @dataclass
@@ -57,7 +50,7 @@ class Mistral3EncoderArchConfig(TextEncoderArchConfig):
) )
_fsdp_shard_conditions: list = field( _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): def __post_init__(self):
@@ -7,18 +7,11 @@ from sglang.multimodal_gen.configs.models.encoders.base import (
TextEncoderArchConfig, TextEncoderArchConfig,
TextEncoderConfig, TextEncoderConfig,
) )
from sglang.multimodal_gen.configs.models.fsdp import (
is_embed_tokens,
def _is_transformer_layer(n: str, m) -> bool: is_final_norm,
return "layers" in n and str.isdigit(n.split(".")[-1]) is_layer,
)
def _is_embeddings(n: str, m) -> bool:
return n.endswith("embed_tokens")
def _is_final_norm(n: str, m) -> bool:
return n.endswith("norm")
@dataclass @dataclass
@@ -66,7 +59,7 @@ class Qwen3TextArchConfig(TextEncoderArchConfig):
# FSDP sharding conditions for CPU offload # FSDP sharding conditions for CPU offload
_fsdp_shard_conditions: list = field( _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: def __post_init__(self) -> None:
@@ -6,18 +6,11 @@ from sglang.multimodal_gen.configs.models.encoders.base import (
TextEncoderArchConfig, TextEncoderArchConfig,
TextEncoderConfig, TextEncoderConfig,
) )
from sglang.multimodal_gen.configs.models.fsdp import (
is_embed_tokens,
def _is_transformer_layer(n: str, m) -> bool: is_final_norm,
return "layers" in n and str.isdigit(n.split(".")[-1]) is_layer,
)
def _is_embeddings(n: str, m) -> bool:
return n.endswith("embed_tokens")
def _is_final_norm(n: str, m) -> bool:
return n.endswith("norm")
@dataclass @dataclass
@@ -63,7 +56,7 @@ class Qwen3VLArchConfig(TextEncoderArchConfig):
] ]
) )
_fsdp_shard_conditions: list = field( _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 # JoyImage specific settings
@@ -7,18 +7,11 @@ from sglang.multimodal_gen.configs.models.encoders.base import (
TextEncoderArchConfig, TextEncoderArchConfig,
TextEncoderConfig, TextEncoderConfig,
) )
from sglang.multimodal_gen.configs.models.fsdp import (
is_embed_tokens,
def _is_transformer_layer(n: str, m) -> bool: is_final_norm,
return "layers" in n and str.isdigit(n.split(".")[-1]) is_layer,
)
def _is_embeddings(n: str, m) -> bool:
return n.endswith("embed_tokens")
def _is_final_norm(n: str, m) -> bool:
return n.endswith("norm")
@dataclass @dataclass
@@ -63,7 +56,7 @@ class QwenImageArchConfig(TextEncoderArchConfig):
] ]
) )
_fsdp_shard_conditions: list = field( _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]
) )
@@ -8,18 +8,11 @@ from sglang.multimodal_gen.configs.models.encoders.base import (
TextEncoderArchConfig, TextEncoderArchConfig,
TextEncoderConfig, TextEncoderConfig,
) )
from sglang.multimodal_gen.configs.models.fsdp import (
is_final_layer_norm,
def _is_transformer_layer(n: str, m) -> bool: is_shared,
return "block" in n and str.isdigit(n.split(".")[-1]) is_t5_block,
)
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")
@dataclass @dataclass
@@ -55,9 +48,9 @@ class T5ArchConfig(TextEncoderArchConfig):
) )
_fsdp_shard_conditions: list = field( _fsdp_shard_conditions: list = field(
default_factory=lambda: [ default_factory=lambda: [
_is_transformer_layer, is_t5_block,
_is_embeddings, is_shared,
_is_final_layernorm, is_final_layer_norm,
] ]
) )
@@ -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"
@@ -71,11 +71,10 @@ class BridgeLoader(ComponentLoader):
default_dtype, default_dtype,
) )
# Check if FSDP loading is available # Use the FSDP loader when FSDP is requested or shard rules are declared.
if ( fsdp_shard_conditions = getattr(model_cls, "_fsdp_shard_conditions", None)
server_args.hsdp_shard_dim is not None if server_args.use_fsdp_inference or (
and hasattr(model_cls, "_fsdp_shard_conditions") server_args.hsdp_shard_dim is not None and fsdp_shard_conditions
and model_cls._fsdp_shard_conditions
): ):
# Load with FSDP support # Load with FSDP support
model = maybe_load_fsdp_model( model = maybe_load_fsdp_model(
@@ -23,6 +23,7 @@ from torch.distributed.fsdp import (
) )
from torch.nn.modules.module import _IncompatibleKeys 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.layers.linear import UnquantizedLinearMethod
from sglang.multimodal_gen.runtime.loader.utils import ( from sglang.multimodal_gen.runtime.loader.utils import (
get_param_names_mapping, get_param_names_mapping,
@@ -79,6 +80,62 @@ def _make_param_like(
return new_param 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( def _maybe_dequantize_fp8(
full_tensor: torch.Tensor, full_tensor: torch.Tensor,
target_dtype: torch.dtype, 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") logger.info("Disabling FSDP for MPS platform as it's not compatible")
if use_fsdp: 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 world_size = hsdp_replicate_dim * hsdp_shard_dim
if not fsdp_inference: if not fsdp_inference:
hsdp_replicate_dim = world_size hsdp_replicate_dim = world_size
@@ -178,7 +240,7 @@ def maybe_load_fsdp_model(
reshard_after_forward=True, reshard_after_forward=True,
mp_policy=mp_policy, mp_policy=mp_policy,
mesh=device_mesh, 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, pin_cpu_memory=pin_cpu_memory,
) )
@@ -224,7 +286,7 @@ def shard_model(
reshard_after_forward: bool = True, reshard_after_forward: bool = True,
mp_policy: MixedPrecisionPolicy | None = MixedPrecisionPolicy(), # noqa mp_policy: MixedPrecisionPolicy | None = MixedPrecisionPolicy(), # noqa
mesh: DeviceMesh | None = None, 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, pin_cpu_memory: bool = True,
) -> None: ) -> 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. 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( 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__, type(model).__name__,
) )
return
fsdp_kwargs = { fsdp_kwargs = {
"reshard_after_forward": reshard_after_forward, "reshard_after_forward": reshard_after_forward,
@@ -274,11 +339,18 @@ def shard_model(
if num_layers_sharded == 0: if num_layers_sharded == 0:
raise ValueError( 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 # Finally shard the entire model to account for any stragglers
fully_shard(model, **fsdp_kwargs) 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 # 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"): if not hasattr(meta_sharded_param, "device_mesh"):
full_tensor = full_tensor.to(device=device, dtype=target_dtype) 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 = ( weight_loader = (
getattr(actual_param, "weight_loader", None) getattr(actual_param, "weight_loader", None)
if actual_param is not None if actual_param is not None
@@ -440,6 +514,38 @@ def load_model_from_full_model_state_dict(
sharded_tensor = sharded_tensor.cpu() sharded_tensor = sharded_tensor.cpu()
else: else:
full_tensor = full_tensor.to(device=device, dtype=target_dtype) 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( sharded_tensor = distribute_tensor(
full_tensor, full_tensor,
meta_sharded_param.device_mesh, meta_sharded_param.device_mesh,
@@ -1165,6 +1165,7 @@ class QwenImageTransformer2DModel(CachableDiT, OffloadableDiTMixin):
_repeated_blocks = ["QwenImageTransformerBlock"] _repeated_blocks = ["QwenImageTransformerBlock"]
param_names_mapping = QwenImageDitConfig().arch_config.param_names_mapping param_names_mapping = QwenImageDitConfig().arch_config.param_names_mapping
_fsdp_shard_conditions = QwenImageDitConfig().arch_config._fsdp_shard_conditions
@classmethod @classmethod
def get_nunchaku_quant_rules(cls) -> dict[str, list[str]]: def get_nunchaku_quant_rules(cls) -> dict[str, list[str]]:
@@ -247,7 +247,9 @@ class ComfyUIQwenImagePipelineBase(LoRAPipeline, ComposedPipelineBase):
reshard_after_forward=True, reshard_after_forward=True,
mp_policy=mp_policy, mp_policy=mp_policy,
mesh=device_mesh, 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, pin_cpu_memory=server_args.pin_cpu_memory,
) )
finally: finally:
@@ -143,7 +143,6 @@ class ComfyUIZImagePipeline(LoRAPipeline, ComposedPipelineBase):
head_dim = dim // num_heads head_dim = dim // num_heads
q_size = dim q_size = dim
k_size = head_dim * num_kv_heads k_size = head_dim * num_kv_heads
v_size = head_dim * num_kv_heads
for name, tensor in weight_iterator: for name, tensor in weight_iterator:
# Match qkv weights in layers, noise_refiner, or context_refiner # 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") logger.info("Disabling FSDP for MPS platform as it's not compatible")
if use_fsdp: if use_fsdp:
world_size = server_args.hsdp_replicate_dim * server_args.hsdp_shard_dim
device_mesh = init_device_mesh( device_mesh = init_device_mesh(
current_platform.device_type, current_platform.device_type,
mesh_shape=( mesh_shape=(
@@ -326,7 +324,9 @@ class ComfyUIZImagePipeline(LoRAPipeline, ComposedPipelineBase):
reshard_after_forward=True, reshard_after_forward=True,
mp_policy=mp_policy, mp_policy=mp_policy,
mesh=device_mesh, 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, pin_cpu_memory=server_args.pin_cpu_memory,
) )
@@ -195,6 +195,7 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
self._cache_dit_enabled = False self._cache_dit_enabled = False
self._cached_num_steps = None self._cached_num_steps = None
self._is_warmed_up = False 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: def _infer_transformer_attention_backend(self) -> AttentionBackendEnum | None:
backends = { backends = {
@@ -699,7 +700,7 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
| server_args.pipeline_config.prepare_pos_cond_kwargs( | server_args.pipeline_config.prepare_pos_cond_kwargs(
batch, batch,
self.device, self.device,
getattr(self.transformer, "rotary_emb", None), self._get_transformer_attr("rotary_emb"),
dtype=target_dtype, dtype=target_dtype,
) )
| dict( | dict(
@@ -720,7 +721,7 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
| server_args.pipeline_config.prepare_neg_cond_kwargs( | server_args.pipeline_config.prepare_neg_cond_kwargs(
batch, batch,
self.device, self.device,
getattr(self.transformer, "rotary_emb", None), self._get_transformer_attr("rotary_emb"),
dtype=target_dtype, dtype=target_dtype,
) )
| dict( | dict(
@@ -778,6 +779,25 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
if hasattr(scheduler, "timestep_list"): if hasattr(scheduler, "timestep_list"):
scheduler.timestep_list = [None] * solver_order 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( def _prepare_step_state(
self, self,
ctx: DenoisingContext, ctx: DenoisingContext,
@@ -1270,15 +1290,7 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
self._finalize_denoising_loop(ctx, batch, server_args) self._finalize_denoising_loop(ctx, batch, server_args)
return batch return batch
# TODO: this will extends the preparation stage, should let subclass/passed-in variables decide which to prepare def _get_extra_func_kwarg_names(self, func) -> tuple[bool, frozenset[str]]:
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.
"""
import functools import functools
# Handle cache-dit's partial wrapping logic. # Handle cache-dit's partial wrapping logic.
@@ -1290,10 +1302,37 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
# Unwrap any decorators (e.g. functools.wraps) # Unwrap any decorators (e.g. functools.wraps)
target_func = inspect.unwrap(func) 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 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( def progress_bar(
self, iterable: Iterable | None = None, total: int | None = None self, iterable: Iterable | None = None, total: int | None = None
@@ -1641,10 +1680,14 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
guidance: torch.Tensor, guidance: torch.Tensor,
**kwargs, **kwargs,
): ):
guidance_kwargs = self.prepare_extra_func_kwargs(
getattr(current_model, "forward", current_model),
{"guidance": guidance},
)
return current_model( return current_model(
hidden_states=latent_model_input, hidden_states=latent_model_input,
timestep=timestep, timestep=timestep,
guidance=guidance, **guidance_kwargs,
**kwargs, **kwargs,
) )
@@ -2019,27 +2019,27 @@
}, },
"fsdp-inference": { "fsdp-inference": {
"stages_ms": { "stages_ms": {
"InputValidationStage": 0.05, "InputValidationStage": 0.06,
"LatentPreparationStage": 0.15, "LatentPreparationStage": 0.16,
"TextEncodingStage": 297.4, "TextEncodingStage": 305.97,
"TimestepPreparationStage": 57.84, "TimestepPreparationStage": 57.19,
"DecodingStage": 8.17, "DecodingStage": 16.88,
"DenoisingStage": 2142.47 "DenoisingStage": 2422.53
}, },
"denoise_step_ms": { "denoise_step_ms": {
"0": 80.53, "0": 259.26,
"1": 188.09, "1": 284.25,
"2": 219.79, "2": 283.74,
"3": 219.53, "3": 270.48,
"4": 218.6, "4": 278.55,
"5": 216.57, "5": 271.58,
"6": 215.51, "6": 270.89,
"7": 215.09, "7": 277.75,
"8": 215.14 "8": 270.1
}, },
"expected_e2e_ms": 2745.03, "expected_e2e_ms": 2775.88,
"expected_avg_denoise_ms": 237.53, "expected_avg_denoise_ms": 268.55,
"expected_median_denoise_ms": 260.88, "expected_median_denoise_ms": 268.51,
"estimated_full_test_time_s": 122.7 "estimated_full_test_time_s": 122.7
}, },
"hunyuan3d_shape_gen": { "hunyuan3d_shape_gen": {
@@ -3,6 +3,11 @@ import sys
import unittest import unittest
from unittest.mock import patch 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 ( from sglang.multimodal_gen.configs.pipeline_configs.base import (
ModelTaskType, ModelTaskType,
PipelineConfig, PipelineConfig,
@@ -11,6 +16,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
QwenImagePipelineConfig, QwenImagePipelineConfig,
) )
from sglang.multimodal_gen.registry import _get_config_info 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.runtime.server_args import ServerArgs
from sglang.multimodal_gen.utils import FlexibleArgumentParser from sglang.multimodal_gen.utils import FlexibleArgumentParser
@@ -159,6 +167,41 @@ class TestOffloadDefaults(unittest.TestCase):
self.assertTrue(args.vae_cpu_offload) 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): class TestModelIdResolution(unittest.TestCase):
def setUp(self): def setUp(self):
_get_config_info.cache_clear() _get_config_info.cache_clear()