[diffusion] chore: reject incompatible transformer fallback (#36917)
This commit is contained in:
@@ -27,6 +27,7 @@ class BridgeLoader(PlainStateDictComponentLoader):
|
|||||||
|
|
||||||
component_names = ["dual_tower_bridge"]
|
component_names = ["dual_tower_bridge"]
|
||||||
expected_library = "diffusers"
|
expected_library = "diffusers"
|
||||||
|
supports_fsdp_inference = True
|
||||||
|
|
||||||
def load_customized(
|
def load_customized(
|
||||||
self, component_model_path: str, server_args: ServerArgs, component_name: str
|
self, component_model_path: str, server_args: ServerArgs, component_name: str
|
||||||
|
|||||||
@@ -137,6 +137,7 @@ class ComponentLoader(ABC):
|
|||||||
# Gates only --component-quantizations.<name>. Quantization declared by a
|
# Gates only --component-quantizations.<name>. Quantization declared by a
|
||||||
# checkpoint is discovered and admitted by the component's normal loader.
|
# checkpoint is discovered and admitted by the component's normal loader.
|
||||||
supports_online_quantization_override = False
|
supports_online_quantization_override = False
|
||||||
|
supports_fsdp_inference = False
|
||||||
|
|
||||||
_loaders_registered = False
|
_loaders_registered = False
|
||||||
|
|
||||||
@@ -176,6 +177,21 @@ class ComponentLoader(ABC):
|
|||||||
)
|
)
|
||||||
return component_name in native_only_components
|
return component_name in native_only_components
|
||||||
|
|
||||||
|
def validate_native_fallback(
|
||||||
|
self, _server_args: ServerArgs, _component_name: str
|
||||||
|
) -> None:
|
||||||
|
"""Validate that fallback preserves the exact component's runtime contract."""
|
||||||
|
pass
|
||||||
|
|
||||||
|
def disable_unsupported_component_fsdp(
|
||||||
|
self, server_args: ServerArgs, component_name: str
|
||||||
|
) -> None:
|
||||||
|
if (
|
||||||
|
not self.supports_fsdp_inference
|
||||||
|
and server_args.should_use_fsdp_for_component(component_name)
|
||||||
|
):
|
||||||
|
server_args.disable_fsdp_for_component(component_name)
|
||||||
|
|
||||||
def _load_customized_with_context(
|
def _load_customized_with_context(
|
||||||
self,
|
self,
|
||||||
component_model_path: str,
|
component_model_path: str,
|
||||||
@@ -236,6 +252,7 @@ class ComponentLoader(ABC):
|
|||||||
|
|
||||||
"""
|
"""
|
||||||
self._native_load_manages_placement = False
|
self._native_load_manages_placement = False
|
||||||
|
self.disable_unsupported_component_fsdp(server_args, component_name)
|
||||||
component_quantization = server_args.component_quantizations.get(component_name)
|
component_quantization = server_args.component_quantizations.get(component_name)
|
||||||
if (
|
if (
|
||||||
component_quantization is not None
|
component_quantization is not None
|
||||||
@@ -289,6 +306,7 @@ class ComponentLoader(ABC):
|
|||||||
f"Failed to load customized {component_name}; native fallback "
|
f"Failed to load customized {component_name}; native fallback "
|
||||||
"is disabled for this component configuration."
|
"is disabled for this component configuration."
|
||||||
) from e
|
) from e
|
||||||
|
self.validate_native_fallback(server_args, component_name)
|
||||||
if native_loader_required:
|
if native_loader_required:
|
||||||
logger.info("%s", e)
|
logger.info("%s", e)
|
||||||
elif "Unsupported model architecture" in str(e):
|
elif "Unsupported model architecture" in str(e):
|
||||||
|
|||||||
@@ -154,6 +154,7 @@ class TransformerLoader(ComponentLoader):
|
|||||||
|
|
||||||
allow_global_attention_backend_fallback = False
|
allow_global_attention_backend_fallback = False
|
||||||
supports_online_quantization_override = True
|
supports_online_quantization_override = True
|
||||||
|
supports_fsdp_inference = True
|
||||||
|
|
||||||
component_names = [
|
component_names = [
|
||||||
"transformer",
|
"transformer",
|
||||||
@@ -192,6 +193,33 @@ class TransformerLoader(ComponentLoader):
|
|||||||
or component_server_args.quantization is not None
|
or component_server_args.quantization is not None
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def validate_native_fallback(
|
||||||
|
self, server_args: ServerArgs, component_name: str
|
||||||
|
) -> None:
|
||||||
|
requested_distributed_execution = []
|
||||||
|
if server_args.tp_size is not None and server_args.tp_size > 1:
|
||||||
|
requested_distributed_execution.append(f"tp_size={server_args.tp_size}")
|
||||||
|
if server_args.sp_degree is not None and server_args.sp_degree > 1:
|
||||||
|
requested_distributed_execution.append(f"sp_degree={server_args.sp_degree}")
|
||||||
|
if server_args.ulysses_degree is not None and server_args.ulysses_degree > 1:
|
||||||
|
requested_distributed_execution.append(
|
||||||
|
f"ulysses_degree={server_args.ulysses_degree}"
|
||||||
|
)
|
||||||
|
if server_args.ring_degree is not None and server_args.ring_degree > 1:
|
||||||
|
requested_distributed_execution.append(
|
||||||
|
f"ring_degree={server_args.ring_degree}"
|
||||||
|
)
|
||||||
|
if server_args.should_use_fsdp_for_component(component_name):
|
||||||
|
requested_distributed_execution.append("FSDP")
|
||||||
|
if requested_distributed_execution:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Native Diffusers fallback for transformer component "
|
||||||
|
f"{component_name!r} cannot honor requested distributed execution: "
|
||||||
|
f"{', '.join(requested_distributed_execution)}. Use an SGLang-native "
|
||||||
|
"transformer implementation or set tp_size, sp_degree, "
|
||||||
|
"ulysses_degree, and ring_degree to 1 without FSDP."
|
||||||
|
)
|
||||||
|
|
||||||
def load_customized(
|
def load_customized(
|
||||||
self,
|
self,
|
||||||
component_model_path: str,
|
component_model_path: str,
|
||||||
|
|||||||
@@ -41,6 +41,7 @@ class TestImageEncoderQuantizationAdmission(unittest.TestCase):
|
|||||||
component_quantizations={},
|
component_quantizations={},
|
||||||
encoder_parallel="replicate",
|
encoder_parallel="replicate",
|
||||||
resolve_component_attention_backend=lambda _name: (None, None),
|
resolve_component_attention_backend=lambda _name: (None, None),
|
||||||
|
should_use_fsdp_for_component=lambda _name: False,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _component_config(self, architecture, *, quantized):
|
def _component_config(self, architecture, *, quantized):
|
||||||
|
|||||||
@@ -0,0 +1,77 @@
|
|||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.loader.component_loaders.bridge_loader import (
|
||||||
|
BridgeLoader,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
||||||
|
ComponentLoader,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.loader.component_loaders.transformer_loader import (
|
||||||
|
TransformerLoader,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency import (
|
||||||
|
RESIDENT,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
|
|
||||||
|
|
||||||
|
class TestTransformerLoaderFallbackAdmission(unittest.TestCase):
|
||||||
|
@staticmethod
|
||||||
|
def _server_args(*, fsdp_requested=False, **overrides):
|
||||||
|
values = {
|
||||||
|
"tp_size": 1,
|
||||||
|
"sp_degree": 1,
|
||||||
|
"ulysses_degree": 1,
|
||||||
|
"ring_degree": 1,
|
||||||
|
"should_use_fsdp_for_component": lambda _component: fsdp_requested,
|
||||||
|
}
|
||||||
|
values.update(overrides)
|
||||||
|
return SimpleNamespace(**values)
|
||||||
|
|
||||||
|
def test_parallel_execution_rejects_native_fallback(self):
|
||||||
|
cases = (
|
||||||
|
({"tp_size": 2}, "tp_size=2"),
|
||||||
|
({"sp_degree": 2}, "sp_degree=2"),
|
||||||
|
({"ulysses_degree": 2}, "ulysses_degree=2"),
|
||||||
|
({"ring_degree": 2}, "ring_degree=2"),
|
||||||
|
({"fsdp_requested": True}, "FSDP"),
|
||||||
|
)
|
||||||
|
|
||||||
|
for overrides, expected_error in cases:
|
||||||
|
with self.subTest(overrides=overrides):
|
||||||
|
with self.assertRaisesRegex(RuntimeError, expected_error):
|
||||||
|
TransformerLoader().validate_native_fallback(
|
||||||
|
self._server_args(**overrides), "transformer_2"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_replicated_execution_keeps_native_fallback_available(self):
|
||||||
|
self.assertIsNone(
|
||||||
|
TransformerLoader().validate_native_fallback(
|
||||||
|
self._server_args(), "transformer_2"
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_only_fsdp_materializers_keep_the_component_request(self):
|
||||||
|
server_args = ServerArgs.__new__(ServerArgs)
|
||||||
|
server_args.use_fsdp_inference = True
|
||||||
|
server_args._fsdp_disabled_components = set()
|
||||||
|
server_args.residency_mode = lambda _component: RESIDENT
|
||||||
|
|
||||||
|
ComponentLoader().disable_unsupported_component_fsdp(
|
||||||
|
server_args, "text_encoder"
|
||||||
|
)
|
||||||
|
self.assertFalse(server_args.should_use_fsdp_for_component("text_encoder"))
|
||||||
|
|
||||||
|
TransformerLoader().disable_unsupported_component_fsdp(
|
||||||
|
server_args, "transformer"
|
||||||
|
)
|
||||||
|
BridgeLoader().disable_unsupported_component_fsdp(
|
||||||
|
server_args, "dual_tower_bridge"
|
||||||
|
)
|
||||||
|
self.assertTrue(server_args.should_use_fsdp_for_component("transformer"))
|
||||||
|
self.assertTrue(server_args.should_use_fsdp_for_component("dual_tower_bridge"))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -54,6 +54,9 @@ class _FakeServerArgs:
|
|||||||
def should_configure_layerwise_offload_for_lazy_component(self, component_name):
|
def should_configure_layerwise_offload_for_lazy_component(self, component_name):
|
||||||
return component_name in self.layerwise_components
|
return component_name in self.layerwise_components
|
||||||
|
|
||||||
|
def should_use_fsdp_for_component(self, _component_name):
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
class TestDeploymentBytesRoot(unittest.TestCase):
|
class TestDeploymentBytesRoot(unittest.TestCase):
|
||||||
"""A hub repo id is not a directory; the component path always is."""
|
"""A hub repo id is not a directory; the component path always is."""
|
||||||
|
|||||||
Reference in New Issue
Block a user