fix(diffusion): unshard FSDP root group for custom encoder entry points (#34575)

This commit is contained in:
triple-mu
2026-08-13 23:24:00 +08:00
committed by GitHub
parent ebca0bbde4
commit ea7a6e0e99
7 changed files with 88 additions and 3 deletions
@@ -27,7 +27,10 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import (
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
ComponentLoader,
)
from sglang.multimodal_gen.runtime.loader.fsdp_load import shard_model
from sglang.multimodal_gen.runtime.loader.fsdp_load import (
register_fsdp_entrypoints,
shard_model,
)
from sglang.multimodal_gen.runtime.loader.utils import (
set_default_torch_dtype,
skip_init_modules,
@@ -499,6 +502,7 @@ class TextEncoderLoader(ComponentLoader):
or getattr(model, "_fsdp_shard_conditions", None),
pin_cpu_memory=server_args.pin_cpu_memory,
)
register_fsdp_entrypoints(model)
else:
model = model.to("cpu")
else:
@@ -207,6 +207,20 @@ def _maybe_dequantize_fp8(
return full_tensor
def register_fsdp_entrypoints(model: torch.nn.Module) -> None:
"""Let FSDP2 unshard around forward passes that bypass ``__call__``.
FSDP2 only unshards around the wrapped module's own ``forward``. Parameters
the shard conditions did not match stay in the catch-all root group, whose
hook therefore never fires for a model driven through a custom method, and
the first op mixing them with a plain tensor fails. Models declare those
entry points in ``_fsdp_forward_methods``, which every model loaded through
FSDP must define; ``BaseDiT`` and ``TextEncoder`` default it to ``()``.
"""
for name in model._fsdp_forward_methods:
register_fsdp_forward_method(model, name)
# TODO(PY): add compile option
def maybe_load_fsdp_model(
model_cls: type[nn.Module],
@@ -226,6 +240,9 @@ def maybe_load_fsdp_model(
) -> torch.nn.Module:
"""Load a model with optional FSDP (Fully Sharded Data Parallel) support.
``model_cls`` must declare ``_fsdp_forward_methods``, the entry points FSDP2
has to unshard around (empty when the model is driven through ``__call__``).
Args:
param_dtype: Data type for model parameters, also used for:
- Model initialization context (set_default_torch_dtype)
@@ -316,8 +333,7 @@ def maybe_load_fsdp_model(
fsdp_shard_conditions=getattr(model, "_fsdp_shard_conditions", None),
pin_cpu_memory=pin_cpu_memory,
)
if callable(getattr(model, "refine_prompt_embeds", None)):
register_fsdp_forward_method(model, "refine_prompt_embeds")
register_fsdp_entrypoints(model)
param_names_mapping_fn = get_param_names_mapping(model.param_names_mapping)
@@ -27,6 +27,11 @@ class BaseDiT(nn.Module, ABC):
# execution semantics support only a subset of the available backends.
_fsdp_shard_conditions: list = []
_compile_conditions: list = []
# Methods that drive a forward pass without going through __call__. FSDP2
# only unshards around the wrapped module's own forward, so anything the
# shard conditions left in the root group stays sharded unless the entry
# point is registered; loaders read this and register each name.
_fsdp_forward_methods: tuple[str, ...] = ()
param_names_mapping: dict
reverse_param_names_mapping: dict
hidden_size: int
@@ -1026,6 +1026,8 @@ class MiniMaxH3FinalLayer(nn.Module):
class MiniMaxH3DiTModel(BaseDiT, LayerwiseOffloadableModuleMixin):
_fsdp_shard_conditions = [is_block]
# refine_prompt_embeds drives a forward pass outside __call__.
_fsdp_forward_methods = ("refine_prompt_embeds",)
# parameters mix fp32 (patch projections, timestep embedder, and output
# heads) with bf16 blocks; FSDP must gather in each parameter's own dtype
_fsdp_mixed_dtype_params = True
@@ -163,6 +163,11 @@ class TextEncoder(nn.Module, ABC, LayerwiseOffloadableModuleMixin):
"model.language_model.layers",
]
_fsdp_shard_conditions: list = field(default_factory=lambda: [])
# Methods that drive a forward pass without going through __call__. FSDP2
# only unshards around the wrapped module's own forward, so anything the
# shard conditions left in the root group stays sharded unless the entry
# point is registered; loaders read this and register each name.
_fsdp_forward_methods: tuple[str, ...] = ()
_stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=list)
_supported_attention_backends: set[AttentionBackendEnum] = (
TextEncoderConfig()._supported_attention_backends
@@ -41,6 +41,10 @@ class MiniMaxH3Qwen3VLEncoder(TextEncoder):
eight otherwise-idle ranks during encoding.
"""
# encode_ids drives the forward pass; __call__ is never used, so FSDP2
# needs it registered or the root group (the vision tower) stays sharded.
_fsdp_forward_methods = ("encode_ids",)
supports_dp_encode = True
@staticmethod
@@ -16,6 +16,9 @@ from sglang.multimodal_gen.runtime.loader.weight_load_plan import WeightLoadPlan
class _UniformDtypeModel(nn.Module):
param_names_mapping = {}
# Every model the FSDP loader accepts declares its custom forward entry
# points, as BaseDiT and TextEncoder do; none here drive one.
_fsdp_forward_methods: tuple[str, ...] = ()
def __init__(self) -> None:
super().__init__()
@@ -34,6 +37,13 @@ class _ReplicatedLinearModel(_UniformDtypeModel):
self.proj = ReplicatedLinear(4, 4, bias=False)
class _CustomEntrypointModel(_UniformDtypeModel):
_fsdp_forward_methods = ("refine_prompt_embeds",)
def refine_prompt_embeds(self) -> None:
pass
class TestFSDPMixedPrecisionPolicy(unittest.TestCase):
def _load_and_capture_policy(
self,
@@ -105,6 +115,45 @@ class TestFSDPMixedPrecisionPolicy(unittest.TestCase):
shard_model.assert_not_called()
class TestFSDPEntrypointRegistration(unittest.TestCase):
def _load_and_capture_registrations(self, model_cls: type[nn.Module]):
with (
patch.object(fsdp_load.current_platform, "is_mps", return_value=False),
patch.object(fsdp_load, "init_device_mesh", return_value=object()),
patch.object(fsdp_load, "shard_model"),
patch.object(
fsdp_load,
"safetensors_weights_iterator",
return_value=iter(()),
),
patch.object(fsdp_load, "load_model_from_full_model_state_dict"),
patch.object(fsdp_load, "register_fsdp_forward_method") as register,
):
model = fsdp_load.maybe_load_fsdp_model(
model_cls=model_cls,
init_params={},
weight_dir_list=[],
device=torch.device("cpu"),
hsdp_replicate_dim=1,
hsdp_shard_dim=1,
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
fsdp_inference=True,
)
return model, register
def test_declared_entry_points_are_registered(self):
model, register = self._load_and_capture_registrations(_CustomEntrypointModel)
register.assert_called_once_with(model, "refine_prompt_embeds")
def test_model_without_entry_points_registers_nothing(self):
_, register = self._load_and_capture_registrations(_UniformDtypeModel)
register.assert_not_called()
class TestOrdinaryWeightLoading(unittest.TestCase):
def test_direct_device_loading_skips_rank_local_cpu_checkpoint(self):
load_plan = WeightLoadPlan(