[diffusion] refactor: reuse plain state-dict loading without per-model classes (#38127)
This commit is contained in:
@@ -8,7 +8,7 @@ import pkgutil
|
||||
import traceback
|
||||
from abc import ABC
|
||||
from collections.abc import Callable, Iterator
|
||||
from typing import Any, Type
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import transformers
|
||||
@@ -328,6 +328,7 @@ class ComponentLoader(ABC):
|
||||
*,
|
||||
component_attn_backend: Any = None,
|
||||
component_attn_name: str | None = None,
|
||||
allow_native_fallback: bool = True,
|
||||
) -> tuple[AutoModel, float]:
|
||||
"""
|
||||
Template method that standardizes logging around the core load implementation.
|
||||
@@ -399,7 +400,7 @@ class ComponentLoader(ABC):
|
||||
):
|
||||
raise
|
||||
except Exception as e:
|
||||
if require_backend_selection:
|
||||
if require_backend_selection or not allow_native_fallback:
|
||||
raise
|
||||
native_loader_required = isinstance(e, NativeComponentLoaderRequired)
|
||||
if native_loader_required and component_weight_override is not None:
|
||||
@@ -649,6 +650,8 @@ class ComponentLoader(ABC):
|
||||
component_type: str,
|
||||
transformers_or_diffusers: str,
|
||||
component_architecture: str | None = None,
|
||||
*,
|
||||
loader_cls: type["ComponentLoader"] | None = None,
|
||||
) -> "ComponentLoader":
|
||||
"""
|
||||
Factory method to create a component loader for a specific component type.
|
||||
@@ -667,10 +670,9 @@ class ComponentLoader(ABC):
|
||||
transformers_or_diffusers, loader_type
|
||||
)
|
||||
|
||||
if loader_type in component_name_to_loader_cls:
|
||||
loader_cls: Type[ComponentLoader] = component_name_to_loader_cls[
|
||||
loader_type
|
||||
]
|
||||
if loader_cls is None:
|
||||
loader_cls = component_name_to_loader_cls.get(loader_type)
|
||||
if loader_cls is not None:
|
||||
expected_library = loader_cls.expected_library
|
||||
# Assert that the library matches what's expected for this component type
|
||||
assert transformers_or_diffusers == expected_library, (
|
||||
@@ -1013,6 +1015,7 @@ class PipelineComponentLoader:
|
||||
component_attn_backend: Any = None,
|
||||
component_attn_name: str | None = None,
|
||||
component_type: str | None = None,
|
||||
loader_cls: type[ComponentLoader] | None = None,
|
||||
):
|
||||
"""
|
||||
Load a pipeline component.
|
||||
@@ -1023,6 +1026,7 @@ class PipelineComponentLoader:
|
||||
transformers_or_diffusers: Whether the component is from transformers or diffusers
|
||||
component_architecture: the class name of the module
|
||||
component_type: structural config slot when it differs from the exact key
|
||||
loader_cls: explicit pipeline-local loader, with no native fallback
|
||||
"""
|
||||
|
||||
# Get the appropriate loader for this component type
|
||||
@@ -1030,6 +1034,7 @@ class PipelineComponentLoader:
|
||||
component_type or component_name,
|
||||
transformers_or_diffusers,
|
||||
component_architecture,
|
||||
loader_cls=loader_cls,
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -1040,6 +1045,7 @@ class PipelineComponentLoader:
|
||||
transformers_or_diffusers,
|
||||
component_attn_backend=component_attn_backend,
|
||||
component_attn_name=component_attn_name,
|
||||
allow_native_fallback=loader_cls is None,
|
||||
)
|
||||
except Exception:
|
||||
logger.error(
|
||||
|
||||
@@ -9,7 +9,7 @@ This module defines the base class for pipelines that are composed of multiple s
|
||||
|
||||
import os
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Callable, Iterator, Literal, cast
|
||||
from typing import Any, Callable, ClassVar, Iterator, Literal, cast
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
@@ -19,6 +19,7 @@ from sglang.multimodal_gen.runtime.disaggregation.roles import (
|
||||
filter_modules_for_role,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
||||
ComponentLoader,
|
||||
PipelineComponentLoader,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.loader.utils import _normalize_component_type
|
||||
@@ -76,6 +77,8 @@ class ComposedPipelineBase(ABC):
|
||||
_required_config_modules: list[str] = []
|
||||
_unfiltered_required_config_modules: tuple[str, ...] = ()
|
||||
_extra_config_module_map: dict[str, str] = {}
|
||||
# Exact module keys; unspecified components retain the default loader dispatch.
|
||||
component_loaders: ClassVar[dict[str, type[ComponentLoader]]] = {}
|
||||
server_args: ServerArgs | None = None
|
||||
modules: dict[str, Any] = {}
|
||||
executor: PipelineExecutor | None = None
|
||||
@@ -616,6 +619,7 @@ class ComposedPipelineBase(ABC):
|
||||
module, memory_usage = PipelineComponentLoader.load_component(
|
||||
component_name=module_name,
|
||||
component_type=load_module_name,
|
||||
loader_cls=self.component_loaders.get(module_name),
|
||||
component_model_path=component_model_path,
|
||||
transformers_or_diffusers=transformers_or_diffusers,
|
||||
server_args=server_args,
|
||||
|
||||
@@ -106,6 +106,7 @@ class TestComponentLoaderIdentity(unittest.TestCase):
|
||||
load_component.assert_called_once_with(
|
||||
component_name="auxiliary_head",
|
||||
component_type="duration_head_2",
|
||||
loader_cls=None,
|
||||
component_model_path="/model/duration_head_2",
|
||||
transformers_or_diffusers="ltx2",
|
||||
server_args=server_args,
|
||||
|
||||
@@ -0,0 +1,196 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import json
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from safetensors.torch import save_file
|
||||
from torch import nn
|
||||
|
||||
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
|
||||
from sglang.multimodal_gen.runtime.loader.component_loaders.adapter_loader import (
|
||||
AdapterLoader,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
||||
ComponentCheckpointUnsupportedError,
|
||||
ComponentLoader,
|
||||
GenericComponentLoader,
|
||||
PipelineComponentLoader,
|
||||
PlainStateDictComponentLoader,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency import (
|
||||
COMPONENT_OFFLOAD,
|
||||
RESIDENT,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
|
||||
|
||||
class _ProjectionPipeline(ComposedPipelineBase):
|
||||
_required_config_modules = ["projection"]
|
||||
component_loaders = {"projection": PlainStateDictComponentLoader}
|
||||
|
||||
def initialize_pipeline(self, server_args):
|
||||
pass
|
||||
|
||||
def create_pipeline_stages(self, server_args):
|
||||
pass
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def checkpoint(tmp_path):
|
||||
component = tmp_path / "projection"
|
||||
component.mkdir()
|
||||
config = {
|
||||
"_class_name": "PlainLoaderTestProjection",
|
||||
"_diffusers_version": "0",
|
||||
"_name_or_path": "unused",
|
||||
"in_features": 4,
|
||||
"out_features": 3,
|
||||
}
|
||||
(component / "config.json").write_text(json.dumps(config))
|
||||
(tmp_path / "model_index.json").write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"_class_name": "ProjectionPipeline",
|
||||
"_diffusers_version": "0",
|
||||
"projection": ["diffusers", "PlainLoaderTestProjection"],
|
||||
"scheduler": None,
|
||||
}
|
||||
)
|
||||
)
|
||||
weights = {"weight": torch.arange(12).reshape(3, 4).float(), "bias": torch.ones(3)}
|
||||
save_file(weights, component / "model.safetensors")
|
||||
with patch.dict(ModelRegistry.registered_models):
|
||||
ModelRegistry.register_model("PlainLoaderTestProjection", nn.Linear)
|
||||
yield component, weights
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def server_args():
|
||||
return ServerArgs(
|
||||
model_path="x",
|
||||
component_residency={"projection": COMPONENT_OFFLOAD},
|
||||
)
|
||||
|
||||
|
||||
def _load(component, server_args, **kwargs):
|
||||
with patch.object(current_platform, "get_available_gpu_memory", return_value=16.0):
|
||||
model, _ = PipelineComponentLoader.load_component(
|
||||
"projection",
|
||||
str(component),
|
||||
"diffusers",
|
||||
server_args,
|
||||
loader_cls=PlainStateDictComponentLoader,
|
||||
**kwargs,
|
||||
)
|
||||
return model
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sharded", [False, True])
|
||||
@pytest.mark.parametrize("precision", [None, "fp32"])
|
||||
def test_pipeline_loads_strict_state_dict(checkpoint, server_args, sharded, precision):
|
||||
component, weights = checkpoint
|
||||
if sharded:
|
||||
for key, weight in weights.items():
|
||||
save_file({key: weight}, component / f"{key}.safetensors")
|
||||
(component / "diffusion_pytorch_model.safetensors.index.json").write_text(
|
||||
json.dumps({"weight_map": {key: f"{key}.safetensors" for key in weights}})
|
||||
)
|
||||
# The index must win over this unrelated full checkpoint.
|
||||
save_file({"wrong": torch.zeros(1)}, component / "model.safetensors")
|
||||
if precision is not None:
|
||||
server_args.component_precisions["projection"] = precision
|
||||
|
||||
pipeline = object.__new__(_ProjectionPipeline)
|
||||
pipeline.model_path = str(component.parent)
|
||||
pipeline.server_args = server_args
|
||||
pipeline._disagg_role = RoleType.MONOLITHIC
|
||||
pipeline.memory_usages = {}
|
||||
with patch.object(current_platform, "get_available_gpu_memory", return_value=16.0):
|
||||
model = pipeline.load_modules(server_args)["projection"]
|
||||
|
||||
expected_dtype = torch.float32 if precision else torch.bfloat16
|
||||
assert not model.training
|
||||
assert model.weight.device.type == "cpu"
|
||||
assert model.weight.dtype == expected_dtype
|
||||
for key, tensor in model.state_dict().items():
|
||||
torch.testing.assert_close(tensor, weights[key].to(expected_dtype))
|
||||
assert server_args.model_paths == {"projection": str(component)}
|
||||
|
||||
|
||||
def test_weight_override_and_architecture_fallback(checkpoint, server_args, tmp_path):
|
||||
component, weights = checkpoint
|
||||
config_path = component / "config.json"
|
||||
config = json.loads(config_path.read_text())
|
||||
config.pop("_class_name")
|
||||
config_path.write_text(json.dumps(config))
|
||||
override = tmp_path / "override.safetensors"
|
||||
replacement = {key: value + 1 for key, value in weights.items()}
|
||||
save_file(replacement, override)
|
||||
server_args.component_weights_paths["projection"] = str(override)
|
||||
model = _load(
|
||||
component, server_args, component_architecture="PlainLoaderTestProjection"
|
||||
)
|
||||
torch.testing.assert_close(model.weight, replacement["weight"].bfloat16())
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"failure", ["missing", "unexpected", "shape", "config", "quantized"]
|
||||
)
|
||||
def test_explicit_loader_never_falls_back(checkpoint, server_args, failure):
|
||||
component, weights = checkpoint
|
||||
if failure == "missing":
|
||||
weights.pop("bias")
|
||||
elif failure == "unexpected":
|
||||
weights["extra"] = torch.zeros(1)
|
||||
elif failure == "shape":
|
||||
weights["bias"] = torch.zeros(4)
|
||||
else:
|
||||
config_path = component / "config.json"
|
||||
config = json.loads(config_path.read_text())
|
||||
if failure == "config":
|
||||
config["unsupported_argument"] = True
|
||||
else:
|
||||
config["quantization_config"] = {"quant_method": "fp8"}
|
||||
config_path.write_text(json.dumps(config))
|
||||
save_file(weights, component / "model.safetensors")
|
||||
with patch.object(ComponentLoader, "load_native") as native:
|
||||
with pytest.raises(
|
||||
(RuntimeError, TypeError, ComponentCheckpointUnsupportedError)
|
||||
):
|
||||
_load(component, server_args)
|
||||
native.assert_not_called()
|
||||
|
||||
|
||||
def test_explicit_selection_does_not_change_other_pipelines():
|
||||
selected = ComponentLoader.for_component_type(
|
||||
"duration_head_2", "ltx2", loader_cls=PlainStateDictComponentLoader
|
||||
)
|
||||
assert type(selected) is PlainStateDictComponentLoader
|
||||
assert selected.component_type == "duration_head_2"
|
||||
assert isinstance(
|
||||
ComponentLoader.for_component_type("duration_head_2", "ltx2"), AdapterLoader
|
||||
)
|
||||
assert isinstance(
|
||||
ComponentLoader.for_component_type("projection", "diffusers"),
|
||||
GenericComponentLoader,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required")
|
||||
def test_cuda_residency_and_forward(checkpoint, server_args):
|
||||
component, weights = checkpoint
|
||||
server_args.component_residency["projection"] = RESIDENT
|
||||
model = _load(component, server_args)
|
||||
assert model.weight.device.type == "cuda"
|
||||
inputs = torch.ones(2, 4, device=model.weight.device, dtype=torch.bfloat16)
|
||||
expected = nn.functional.linear(
|
||||
inputs, weights["weight"].to(inputs), weights["bias"].to(inputs)
|
||||
)
|
||||
torch.testing.assert_close(model(inputs), expected, rtol=0, atol=0)
|
||||
Reference in New Issue
Block a user