[diffusion] feat: add plain component weight overrides (#36086)
This commit is contained in:
@@ -44,6 +44,9 @@ class AdapterLoader(PlainStateDictComponentLoader):
|
|||||||
*args,
|
*args,
|
||||||
):
|
):
|
||||||
config = self.load_component_config(component_model_path, component_name)
|
config = self.load_component_config(component_model_path, component_name)
|
||||||
|
component_weights_path = self.resolve_component_weights_path(
|
||||||
|
component_model_path, server_args, component_name
|
||||||
|
)
|
||||||
|
|
||||||
cls_name = config.pop("_class_name", None)
|
cls_name = config.pop("_class_name", None)
|
||||||
if cls_name is None:
|
if cls_name is None:
|
||||||
@@ -74,7 +77,7 @@ class AdapterLoader(PlainStateDictComponentLoader):
|
|||||||
adapter_cfg.update_model_arch(config)
|
adapter_cfg.update_model_arch(config)
|
||||||
model = model_cls(adapter_cfg).to(device=target_device, dtype=default_dtype)
|
model = model_cls(adapter_cfg).to(device=target_device, dtype=default_dtype)
|
||||||
|
|
||||||
loaded = load_safetensors_state_dict(component_model_path)
|
loaded = load_safetensors_state_dict(component_weights_path)
|
||||||
mapping = adapter_cfg.arch_config.param_names_mapping
|
mapping = adapter_cfg.arch_config.param_names_mapping
|
||||||
loaded = {_remap_connector_key(k, mapping): v for k, v in loaded.items()}
|
loaded = {_remap_connector_key(k, mapping): v for k, v in loaded.items()}
|
||||||
|
|
||||||
@@ -84,7 +87,7 @@ class AdapterLoader(PlainStateDictComponentLoader):
|
|||||||
# else uninitialized would surface later as garbage embeddings.
|
# else uninitialized would surface later as garbage embeddings.
|
||||||
if missing or unexpected:
|
if missing or unexpected:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Adapter weights at '{component_model_path}' do not match the "
|
f"Adapter weights at '{component_weights_path}' do not match the "
|
||||||
f"instantiated {cls_name}. Missing: {sorted(missing)}. "
|
f"instantiated {cls_name}. Missing: {sorted(missing)}. "
|
||||||
f"Unexpected: {sorted(unexpected)}. This usually means the "
|
f"Unexpected: {sorted(unexpected)}. This usually means the "
|
||||||
"adapter config or its weight-name mapping is wrong."
|
"adapter config or its weight-name mapping is wrong."
|
||||||
|
|||||||
@@ -46,6 +46,10 @@ from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
|
|||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
from sglang.multimodal_gen.runtime.utils.precision import resolve_component_precision
|
from sglang.multimodal_gen.runtime.utils.precision import resolve_component_precision
|
||||||
|
from sglang.multimodal_gen.runtime.weights.source import (
|
||||||
|
materialize_weight,
|
||||||
|
resolve_weight,
|
||||||
|
)
|
||||||
from sglang.srt.model_loader.checkpoint_quantization import (
|
from sglang.srt.model_loader.checkpoint_quantization import (
|
||||||
resolve_checkpoint_quant_spec,
|
resolve_checkpoint_quant_spec,
|
||||||
)
|
)
|
||||||
@@ -518,6 +522,19 @@ class PlainStateDictComponentLoader(ComponentLoader):
|
|||||||
self.ensure_plain_state_dict_checkpoint(config, component_name)
|
self.ensure_plain_state_dict_checkpoint(config, component_name)
|
||||||
return config
|
return config
|
||||||
|
|
||||||
|
def resolve_component_weights_path(
|
||||||
|
self,
|
||||||
|
component_model_path: str,
|
||||||
|
server_args: ServerArgs,
|
||||||
|
component_name: str,
|
||||||
|
) -> str:
|
||||||
|
override = server_args.component_weights_paths.get(component_name)
|
||||||
|
if override is None:
|
||||||
|
return component_model_path
|
||||||
|
weights_path = materialize_weight(resolve_weight(override))
|
||||||
|
logger.info("Using weight override for %s: %s", component_name, weights_path)
|
||||||
|
return weights_path
|
||||||
|
|
||||||
|
|
||||||
class ImageProcessorLoader(ComponentLoader):
|
class ImageProcessorLoader(ComponentLoader):
|
||||||
"""Loader for image processor."""
|
"""Loader for image processor."""
|
||||||
|
|||||||
+4
-1
@@ -30,6 +30,9 @@ class DiffusionDecoderLoader(PlainStateDictComponentLoader):
|
|||||||
*args,
|
*args,
|
||||||
):
|
):
|
||||||
config = self.load_component_config(component_model_path, component_name)
|
config = self.load_component_config(component_model_path, component_name)
|
||||||
|
component_weights_path = self.resolve_component_weights_path(
|
||||||
|
component_model_path, server_args, component_name
|
||||||
|
)
|
||||||
class_name = config.pop("_class_name", None)
|
class_name = config.pop("_class_name", None)
|
||||||
if class_name is None:
|
if class_name is None:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -54,6 +57,6 @@ class DiffusionDecoderLoader(PlainStateDictComponentLoader):
|
|||||||
model = model_cls(decoder_config).to(device=target_device, dtype=dtype)
|
model = model_cls(decoder_config).to(device=target_device, dtype=dtype)
|
||||||
|
|
||||||
model.load_state_dict(
|
model.load_state_dict(
|
||||||
load_safetensors_state_dict(component_model_path), strict=True
|
load_safetensors_state_dict(component_weights_path), strict=True
|
||||||
)
|
)
|
||||||
return model
|
return model
|
||||||
|
|||||||
+5
-7
@@ -1,11 +1,10 @@
|
|||||||
# SPDX-License-Identifier: Apache-2.0
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
from safetensors.torch import load_file as safetensors_load_file
|
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
||||||
PlainStateDictComponentLoader,
|
PlainStateDictComponentLoader,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.utils import (
|
from sglang.multimodal_gen.runtime.loader.utils import (
|
||||||
_list_safetensors_files,
|
load_safetensors_state_dict,
|
||||||
set_default_torch_dtype,
|
set_default_torch_dtype,
|
||||||
skip_init_modules,
|
skip_init_modules,
|
||||||
)
|
)
|
||||||
@@ -25,6 +24,9 @@ class SoundTokenizerLoader(PlainStateDictComponentLoader):
|
|||||||
self, component_model_path: str, server_args: ServerArgs, component_name: str
|
self, component_model_path: str, server_args: ServerArgs, component_name: str
|
||||||
):
|
):
|
||||||
config = self.load_component_config(component_model_path, component_name)
|
config = self.load_component_config(component_model_path, component_name)
|
||||||
|
component_weights_path = self.resolve_component_weights_path(
|
||||||
|
component_model_path, server_args, component_name
|
||||||
|
)
|
||||||
class_name = config.pop("_class_name", None) or self.component_architecture
|
class_name = config.pop("_class_name", None) or self.component_architecture
|
||||||
assert (
|
assert (
|
||||||
class_name is not None
|
class_name is not None
|
||||||
@@ -45,11 +47,7 @@ class SoundTokenizerLoader(PlainStateDictComponentLoader):
|
|||||||
model_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
model_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||||
model = model_cls(config).to(target_device)
|
model = model_cls(config).to(target_device)
|
||||||
|
|
||||||
safetensors_list = _list_safetensors_files(component_model_path)
|
loaded = load_safetensors_state_dict(component_weights_path)
|
||||||
assert (
|
|
||||||
len(safetensors_list) == 1
|
|
||||||
), f"Found {len(safetensors_list)} safetensors files in {component_model_path}"
|
|
||||||
loaded = safetensors_load_file(safetensors_list[0])
|
|
||||||
incompatible = model.load_state_dict(loaded, strict=False)
|
incompatible = model.load_state_dict(loaded, strict=False)
|
||||||
missing = getattr(incompatible, "missing_keys", [])
|
missing = getattr(incompatible, "missing_keys", [])
|
||||||
# The tokenizer is decoder-only; the checkpoint's encoder weights are
|
# The tokenizer is decoder-only; the checkpoint's encoder weights are
|
||||||
|
|||||||
@@ -1,12 +1,10 @@
|
|||||||
import re
|
import re
|
||||||
|
|
||||||
from safetensors.torch import load_file as safetensors_load_file
|
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
||||||
PlainStateDictComponentLoader,
|
PlainStateDictComponentLoader,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.utils import (
|
from sglang.multimodal_gen.runtime.loader.utils import (
|
||||||
_list_safetensors_files,
|
load_safetensors_state_dict,
|
||||||
set_default_torch_dtype,
|
set_default_torch_dtype,
|
||||||
skip_init_modules,
|
skip_init_modules,
|
||||||
)
|
)
|
||||||
@@ -27,6 +25,9 @@ class VocoderLoader(PlainStateDictComponentLoader):
|
|||||||
self, component_model_path: str, server_args: ServerArgs, component_name: str
|
self, component_model_path: str, server_args: ServerArgs, component_name: str
|
||||||
):
|
):
|
||||||
config = self.load_component_config(component_model_path, component_name)
|
config = self.load_component_config(component_model_path, component_name)
|
||||||
|
component_weights_path = self.resolve_component_weights_path(
|
||||||
|
component_model_path, server_args, component_name
|
||||||
|
)
|
||||||
class_name = config.pop("_class_name", None) or self.component_architecture
|
class_name = config.pop("_class_name", None) or self.component_architecture
|
||||||
assert (
|
assert (
|
||||||
class_name is not None
|
class_name is not None
|
||||||
@@ -57,11 +58,7 @@ class VocoderLoader(PlainStateDictComponentLoader):
|
|||||||
vocoder_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
vocoder_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||||
vocoder = vocoder_cls(vocoder_config).to(target_device)
|
vocoder = vocoder_cls(vocoder_config).to(target_device)
|
||||||
|
|
||||||
safetensors_list = _list_safetensors_files(component_model_path)
|
loaded = load_safetensors_state_dict(component_weights_path)
|
||||||
assert (
|
|
||||||
len(safetensors_list) == 1
|
|
||||||
), f"Found {len(safetensors_list)} safetensors files in {component_model_path}"
|
|
||||||
loaded = safetensors_load_file(safetensors_list[0])
|
|
||||||
mapping = vocoder_config.arch_config.param_names_mapping
|
mapping = vocoder_config.arch_config.param_names_mapping
|
||||||
loaded = {_remap_vocoder_key(k, mapping): v for k, v in loaded.items()}
|
loaded = {_remap_vocoder_key(k, mapping): v for k, v in loaded.items()}
|
||||||
|
|
||||||
@@ -69,7 +66,7 @@ class VocoderLoader(PlainStateDictComponentLoader):
|
|||||||
# A half-loaded vocoder produces plausible but wrong audio.
|
# A half-loaded vocoder produces plausible but wrong audio.
|
||||||
if missing_keys or unexpected_keys:
|
if missing_keys or unexpected_keys:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Vocoder weights at '{component_model_path}' do not match the "
|
f"Vocoder weights at '{component_weights_path}' do not match the "
|
||||||
f"instantiated {class_name}. Missing: {sorted(missing_keys)}. "
|
f"instantiated {class_name}. Missing: {sorted(missing_keys)}. "
|
||||||
f"Unexpected: {sorted(unexpected_keys)}."
|
f"Unexpected: {sorted(unexpected_keys)}."
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -297,6 +297,9 @@ def _list_safetensors_files(model_path: str) -> list[str]:
|
|||||||
automatically via HuggingFace Hub (if the path is an HF cache entry);
|
automatically via HuggingFace Hub (if the path is an HF cache entry);
|
||||||
if repair fails a clear RuntimeError is raised.
|
if repair fails a clear RuntimeError is raised.
|
||||||
"""
|
"""
|
||||||
|
if os.path.isfile(model_path):
|
||||||
|
return [str(model_path)] if str(model_path).endswith(".safetensors") else []
|
||||||
|
|
||||||
found = sorted(glob.glob(os.path.join(str(model_path), "*.safetensors")))
|
found = sorted(glob.glob(os.path.join(str(model_path), "*.safetensors")))
|
||||||
|
|
||||||
index_path = os.path.join(
|
index_path = os.path.join(
|
||||||
|
|||||||
@@ -2,6 +2,7 @@
|
|||||||
|
|
||||||
import re
|
import re
|
||||||
import unittest
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
from unittest.mock import patch
|
from unittest.mock import patch
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.loader.component_loaders.adapter_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.adapter_loader import (
|
||||||
@@ -33,6 +34,29 @@ class _TestLoader(PlainStateDictComponentLoader):
|
|||||||
|
|
||||||
|
|
||||||
class TestComponentQuantizationAdmission(unittest.TestCase):
|
class TestComponentQuantizationAdmission(unittest.TestCase):
|
||||||
|
def test_plain_loader_resolves_weights_separately_from_config(self):
|
||||||
|
server_args = SimpleNamespace(
|
||||||
|
component_weights_paths={"vocoder": "owner/repo/vocoder.safetensors"}
|
||||||
|
)
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.loader.component_loaders."
|
||||||
|
"component_loader.resolve_weight",
|
||||||
|
return_value="resolved",
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.loader.component_loaders."
|
||||||
|
"component_loader.materialize_weight",
|
||||||
|
return_value="/cache/vocoder.safetensors",
|
||||||
|
),
|
||||||
|
):
|
||||||
|
self.assertEqual(
|
||||||
|
_TestLoader().resolve_component_weights_path(
|
||||||
|
"/base/vocoder", server_args, "vocoder"
|
||||||
|
),
|
||||||
|
"/cache/vocoder.safetensors",
|
||||||
|
)
|
||||||
|
|
||||||
def test_plain_checkpoint_config_is_accepted(self):
|
def test_plain_checkpoint_config_is_accepted(self):
|
||||||
config = {"_class_name": "TestModel"}
|
config = {"_class_name": "TestModel"}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user