From 453b98c490d7a7956903d2f3a603034d7cc0780e Mon Sep 17 00:00:00 2001 From: Mick Date: Sat, 22 Aug 2026 23:09:35 +0800 Subject: [PATCH] [diffusion] feat: support single-file component weight overrides (#35979) --- docs/docs/sglang-diffusion/api/cli.mdx | 3 + docs/docs/sglang-diffusion/quantization.mdx | 2 +- .../component_loaders/image_encoder_loader.py | 8 +- .../component_loaders/text_encoder_loader.py | 80 +++++++++++++++++-- .../component_loaders/transformer_loader.py | 13 +-- .../runtime/pipelines/ideogram.py | 14 +--- .../runtime/server_args/server_args.py | 35 ++++++-- .../multimodal_gen/runtime/weights/source.py | 27 ++++++- .../test/unit/test_ideogram4.py | 4 +- .../test/unit/test_image_encoder_loader.py | 1 + .../test/unit/test_server_args.py | 17 ++++ .../test/unit/test_text_encoder_loader.py | 26 ++++++ .../test/unit/test_weight_source.py | 20 +++++ 13 files changed, 214 insertions(+), 36 deletions(-) diff --git a/docs/docs/sglang-diffusion/api/cli.mdx b/docs/docs/sglang-diffusion/api/cli.mdx index 1a9667433..70b7df159 100644 --- a/docs/docs/sglang-diffusion/api/cli.mdx +++ b/docs/docs/sglang-diffusion/api/cli.mdx @@ -150,6 +150,9 @@ for the current component matrix. Component overrides accept a local component directory, a standalone Hub repository, or a Hub component subfolder written as `owner/repo/subfolder`. +For transformer and native encoder loaders, an explicit weight filename keeps +the base component configuration and replaces only its weights; the selected +file format must still be supported by that loader. For supported realtime causal video models, `--kv-cache-quant {off|int4|int2}` compresses completed KV-cache chunks independently of transformer weight diff --git a/docs/docs/sglang-diffusion/quantization.mdx b/docs/docs/sglang-diffusion/quantization.mdx index b75f1dd96..6f029cbdf 100644 --- a/docs/docs/sglang-diffusion/quantization.mdx +++ b/docs/docs/sglang-diffusion/quantization.mdx @@ -16,7 +16,7 @@ precision are resolved independently. Use these paths: - `--model-path`: the base or original model -- `--component-paths.` / `---path`: replace a component from `model_index.json` or a native registered module with an independent repo or local directory +- `--component-paths.` / `---path`: replace a component from `model_index.json` or a native registered module with an independent repo/local directory; for transformer and native encoder loaders, an explicit weight filename retains the base component config - `--transformer-path`: a quantized transformers-style transformer component directory that already contains its own `config.json` - `--transformer-weights-path`: replacement transformer weights in safetensors format (file, directory, or Hub repository/file) or a supported GGUF file diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/image_encoder_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/image_encoder_loader.py index f4a9fbb54..a0e53672d 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/image_encoder_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/image_encoder_loader.py @@ -23,6 +23,11 @@ class ImageEncoderLoader(TextEncoderLoader): component_name: str = "image_encoder", ): """Load the text encoders based on the model path, and inference args.""" + component_weights_path = self.resolve_model_weights_path( + component_model_path, + server_args, + component_name, + ) # model_config: PretrainedConfig = get_hf_config( # model=model_path, # trust_remote_code=server_args.trust_remote_code, @@ -39,6 +44,7 @@ class ImageEncoderLoader(TextEncoderLoader): encoder_config, model_config, component_model_path, + component_weights_path, component_name, ) # real dims are populated now; resolve fold vs replicate @@ -50,7 +56,7 @@ class ImageEncoderLoader(TextEncoderLoader): # Always start with local device; load_model will adjust for offload if needed # TODO(will): add support for other dtypes return self.load_model( - component_model_path, + component_weights_path, encoder_config, server_args, server_args.pipeline_config.image_encoder_precision, diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py index 80b5ced24..61f30643f 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py @@ -62,7 +62,14 @@ from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import ( load_dict, ) from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger -from sglang.multimodal_gen.runtime.utils.quantization_utils import get_quant_config +from sglang.multimodal_gen.runtime.utils.quantization_utils import ( + get_quant_config, + get_quant_config_from_safetensors_metadata, +) +from sglang.multimodal_gen.runtime.weights.source import ( + materialize_weight, + resolve_weight, +) from sglang.multimodal_gen.utils import PRECISION_TO_TYPE from sglang.srt.environ import envs @@ -93,11 +100,29 @@ def _delegate_standard_bnb4_to_transformers( ) +def _get_encoder_quant_config( + component_config: dict, + component_model_path: str, + component_weights_path: str, +): + quant_config = get_quant_config(component_config, component_model_path) + if ( + quant_config is None + and component_weights_path != component_model_path + and component_weights_path.endswith(".safetensors") + ): + quant_config = get_quant_config_from_safetensors_metadata( + component_weights_path + ) + return quant_config + + def _configure_encoder_quantization( model_config: EncoderConfig, model_cls: type[nn.Module], component_config: dict, component_model_path: str, + component_weights_path: str, component_name: str, ) -> None: if getattr(model_cls, "manages_checkpoint_quantization", False): @@ -111,9 +136,10 @@ def _configure_encoder_quantization( component_name, ) try: - quant_config = get_quant_config( + quant_config = _get_encoder_quant_config( component_config, component_model_path, + component_weights_path, ) except (KeyError, TypeError, ValueError) as error: raise ComponentCheckpointUnsupportedError( @@ -134,6 +160,7 @@ def _resolve_and_configure_encoder_quantization( model_config: EncoderConfig, component_config: dict, component_model_path: str, + component_weights_path: str, component_name: str, ) -> type[nn.Module]: architectures = getattr(model_config, "architectures", []) @@ -145,7 +172,11 @@ def _resolve_and_configure_encoder_quantization( component_name, ) try: - quant_config = get_quant_config(component_config, component_model_path) + quant_config = _get_encoder_quant_config( + component_config, + component_model_path, + component_weights_path, + ) except Exception as quantization_error: raise ComponentCheckpointUnsupportedError( f"Cannot parse checkpoint quantization for {component_name!r}: " @@ -163,6 +194,7 @@ def _resolve_and_configure_encoder_quantization( model_cls, component_config, component_model_path, + component_weights_path, component_name, ) return model_cls @@ -242,6 +274,8 @@ def _require_quantized_encoder_layers( def _checkpoint_bytes(model_path: str) -> int: """On-disk size of a checkpoint, readable before any weight of it is.""" + if os.path.isfile(model_path): + return os.path.getsize(model_path) total = 0 for path in glob.glob( os.path.join(str(model_path), "**", "*.safetensors"), recursive=True @@ -274,6 +308,23 @@ class TextEncoderLoader(ComponentLoader): component_names = ["text_encoder"] expected_library = "transformers" + @staticmethod + def resolve_model_weights_path( + component_model_path: str, + server_args: ServerArgs, + component_name: str, + ) -> str: + weights_override = server_args.component_weights_paths.get(component_name) + if weights_override is None: + return component_model_path + model_weights_path = materialize_weight(resolve_weight(weights_override)) + logger.info( + "Using weight-file override for %s: %s", + component_name, + model_weights_path, + ) + return model_weights_path + @dataclasses.dataclass class Source: """A source for weights.""" @@ -320,8 +371,19 @@ class TextEncoderLoader(ComponentLoader): # model_name_or_path = (self._maybe_download_from_modelscope( # model_name_or_path, revision) or model_name_or_path) - is_local = os.path.isdir(model_name_or_path) - assert is_local, "Model path must be a local directory" + if os.path.isfile(model_name_or_path): + if model_name_or_path.endswith(".safetensors"): + return os.path.dirname(model_name_or_path), [model_name_or_path], True + if fall_back_to_pt and model_name_or_path.endswith((".bin", ".pt")): + return os.path.dirname(model_name_or_path), [model_name_or_path], False + raise ValueError( + "Native encoder weight overrides currently support one " + f"safetensors, bin, or pt file, got {model_name_or_path!r}" + ) + if not os.path.isdir(model_name_or_path): + raise ValueError( + f"Model path must be a local file or directory: {model_name_or_path!r}" + ) use_safetensors = False index_file = SAFE_WEIGHTS_INDEX_NAME @@ -447,6 +509,11 @@ class TextEncoderLoader(ComponentLoader): component_starts_on_cpu: bool | None = None, ): """Load the text encoders based on the model path, and inference args.""" + component_weights_path = self.resolve_model_weights_path( + component_model_path, + server_args, + component_name, + ) diffusers_pretrained_config = get_config( component_model_path, trust_remote_code=True ) @@ -478,6 +545,7 @@ class TextEncoderLoader(ComponentLoader): encoder_config, model_config, component_model_path, + component_weights_path, component_name, ) encoder_dp_group = get_encoder_data_parallel_group() @@ -499,7 +567,7 @@ class TextEncoderLoader(ComponentLoader): ] # TODO(will): add support for other dtypes return self.load_model( - component_model_path, + component_weights_path, encoder_config, server_args, encoder_dtype, diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py index 935ece08e..eee1ab2ba 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py @@ -99,15 +99,7 @@ def _server_args_for_transformer_component( server_args: ServerArgs, component_name: str ) -> ServerArgs: """Mask global quantized override flags for secondary transformer components.""" - if component_name not in ("transformer_2", "unconditional_transformer"): - return server_args - - # Some pipelines have secondary DiT components with their own quantized - # weight file. Keep the mapping model-owned and the loader generic. - component_weights_paths = getattr( - server_args, "component_transformer_weights_paths", {} - ) - component_weights_path = component_weights_paths.get(component_name) + component_weights_path = server_args.component_weights_paths.get(component_name) if component_weights_path is not None: component_server_args = copy.copy(server_args) component_server_args.transformer_weights_path = component_weights_path @@ -119,6 +111,9 @@ def _server_args_for_transformer_component( ) return component_server_args + if component_name not in ("transformer_2", "unconditional_transformer"): + return server_args + if ( server_args.transformer_weights_path is None and server_args.nunchaku_config is None diff --git a/python/sglang/multimodal_gen/runtime/pipelines/ideogram.py b/python/sglang/multimodal_gen/runtime/pipelines/ideogram.py index 4613f777c..f32afd668 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines/ideogram.py +++ b/python/sglang/multimodal_gen/runtime/pipelines/ideogram.py @@ -244,25 +244,19 @@ class Ideogram4Nvfp4Pipeline(Ideogram4Pipeline): # The loader treats transformer_weights_path as the base DiT override. # Route the sibling unconditional DiT weights through the generic # per-component override map instead of hard-coding Ideogram there. - component_transformer_weights_paths = dict( - getattr(server_args, "component_transformer_weights_paths", {}) - ) - component_transformer_weights_paths.setdefault( + component_weights_paths = dict(server_args.component_weights_paths) + component_weights_paths.setdefault( "unconditional_transformer", model_resolution.unconditional_transformer_weights_path, ) - server_args.component_transformer_weights_paths = ( - component_transformer_weights_paths - ) + server_args.component_weights_paths = component_weights_paths logger.info( "NVFP4 transformer weights: %s", model_resolution.transformer_weights_path, ) logger.info( "NVFP4 unconditional transformer weights: %s", - server_args.component_transformer_weights_paths.get( - "unconditional_transformer" - ), + server_args.component_weights_paths.get("unconditional_transformer"), ) return super().load_modules(server_args, loaded_modules) diff --git a/python/sglang/multimodal_gen/runtime/server_args/server_args.py b/python/sglang/multimodal_gen/runtime/server_args/server_args.py index d444537a1..ca8821472 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py @@ -77,6 +77,9 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import ( configure_logger, init_logger, ) +from sglang.multimodal_gen.runtime.weights.source import ( + is_explicit_weight_file_reference, +) from sglang.multimodal_gen.utils import ( FlexibleArgumentParser, StoreBoolean, @@ -301,6 +304,8 @@ class ServerArgs(DisaggServerArgsMixin): # Component path overrides (key = model_index.json component name, value = path) component_paths: dict[str, str] = field(default_factory=dict) + # Exact weight-file overrides retain the base component configuration. + component_weights_paths: dict[str, str] = field(default_factory=dict) # Optional LTX-2.5 decoder is large enough to load only when requested. load_diffusion_decoder: bool = False @@ -313,12 +318,6 @@ class ServerArgs(DisaggServerArgsMixin): # Widest timestep plan the rebuild slab is sized for; see # MINIMAX_H3_ADALN_MAX_PLAN_WIDTH. minimax_h3_adaln_plan_width: int = 4 - # Per-component transformer weight overrides (key = model_index.json component name). - # Pipelines use this when a checkpoint ships separate quantized weights for - # secondary DiT components; the generic loader consumes it without model-specific - # filename logic. - component_transformer_weights_paths: dict[str, str] = field(default_factory=dict) - # Explicit quantization method override (e.g. "mxfp8", "fp8", "modelslim"). # When set, the transformer loader uses it instead of auto-detection. quantization: str | None = None @@ -1699,6 +1698,30 @@ class ServerArgs(DisaggServerArgsMixin): # configure logger before use configure_logger(server_args=self) + component_paths: dict[str, str] = {} + component_weights_paths = dict(self.component_weights_paths) + for component, path in self.component_paths.items(): + supports_weight_file_override = ( + is_dit_component_name(component) + or is_text_encoder_component_name(component) + or is_image_encoder_component_name(component) + ) + if ( + not supports_weight_file_override + or not is_explicit_weight_file_reference(path) + ): + component_paths[component] = path + continue + existing = component_weights_paths.get(component) + if existing is not None and existing != path: + raise ValueError( + f"Conflicting weight overrides for component {component!r}: " + f"{existing!r} and {path!r}" + ) + component_weights_paths[component] = path + self.component_paths = component_paths + self.component_weights_paths = component_weights_paths + # Convert string disagg_role to enum (from CLI/config) if isinstance(self.disagg_role, str): self.disagg_role = RoleType.from_string(self.disagg_role) diff --git a/python/sglang/multimodal_gen/runtime/weights/source.py b/python/sglang/multimodal_gen/runtime/weights/source.py index 08ae5c289..d2a7ae3c3 100644 --- a/python/sglang/multimodal_gen/runtime/weights/source.py +++ b/python/sglang/multimodal_gen/runtime/weights/source.py @@ -8,7 +8,7 @@ from pathlib import Path, PurePosixPath from typing import Literal from urllib.parse import unquote, urlparse -from huggingface_hub import HfApi +from huggingface_hub import HfApi, hf_hub_download from huggingface_hub.utils import validate_repo_id WeightSourceKind = Literal["local", "huggingface"] @@ -40,6 +40,14 @@ class ResolvedWeight: selected_file: str +def is_explicit_weight_file_reference(source: str) -> bool: + """Whether a component override names one weight file, not a component root.""" + expanded = os.path.expanduser(source) + if os.path.isdir(expanded): + return False + return urlparse(source).path.lower().endswith(_WEIGHT_SUFFIXES) + + def _validate_relative_hub_path(path: str, field_name: str) -> str: normalized = str(PurePosixPath(path)) pure_path = PurePosixPath(normalized) @@ -264,3 +272,20 @@ def resolve_weight( inventory=inventory, selected_file=selected_file, ) + + +def materialize_weight(resolved: ResolvedWeight) -> str: + """Return the selected local file, downloading one pinned Hub file if needed.""" + source = resolved.inventory.source + if source.kind == "local": + assert source.local_path is not None + if os.path.isfile(source.local_path): + return source.local_path + return os.path.join(source.local_path, resolved.selected_file) + + assert source.repo_id is not None + return hf_hub_download( + repo_id=source.repo_id, + filename=resolved.selected_file, + revision=resolved.inventory.resolved_revision, + ) diff --git a/python/sglang/multimodal_gen/test/unit/test_ideogram4.py b/python/sglang/multimodal_gen/test/unit/test_ideogram4.py index e4a7f4e8f..2bfe875e6 100644 --- a/python/sglang/multimodal_gen/test/unit/test_ideogram4.py +++ b/python/sglang/multimodal_gen/test/unit/test_ideogram4.py @@ -509,7 +509,7 @@ class TestIdeogram4(unittest.TestCase): server_args = SimpleNamespace( transformer_weights_path="/unused/override.safetensors", nunchaku_config={"enabled": True}, - component_transformer_weights_paths={}, + component_weights_paths={}, ) component_args = _server_args_for_transformer_component( server_args, "unconditional_transformer" @@ -524,7 +524,7 @@ class TestIdeogram4(unittest.TestCase): "/ckpt/diffusion_models/ideogram4_nvfp4_mixed.safetensors" ), nunchaku_config={"enabled": True}, - component_transformer_weights_paths={ + component_weights_paths={ "unconditional_transformer": ( "/ckpt/diffusion_models/" "ideogram4_unconditional_nvfp4_mixed.safetensors" diff --git a/python/sglang/multimodal_gen/test/unit/test_image_encoder_loader.py b/python/sglang/multimodal_gen/test/unit/test_image_encoder_loader.py index cd1d9c6df..e092a6908 100644 --- a/python/sglang/multimodal_gen/test/unit/test_image_encoder_loader.py +++ b/python/sglang/multimodal_gen/test/unit/test_image_encoder_loader.py @@ -27,6 +27,7 @@ class TestImageEncoderQuantizationAdmission(unittest.TestCase): image_encoder_precision="bf16", native_only_components=(), ), + component_weights_paths={}, encoder_parallel="replicate", resolve_component_attention_backend=lambda _name: (None, None), ) diff --git a/python/sglang/multimodal_gen/test/unit/test_server_args.py b/python/sglang/multimodal_gen/test/unit/test_server_args.py index 8a5ec7804..7dd737f1b 100644 --- a/python/sglang/multimodal_gen/test/unit/test_server_args.py +++ b/python/sglang/multimodal_gen/test/unit/test_server_args.py @@ -163,6 +163,23 @@ class TestServerArgsPathExpansion(unittest.TestCase): args.component_paths["vae"], os.path.expanduser("~/fake/local/vae") ) + def test_component_weight_file_keeps_base_component_config(self): + args = self._from_dict_without_model_resolution( + { + "model_path": "/data/my-model", + "component_paths": { + "text_encoder": "owner/repo/text_encoder/model.safetensors", + "vae": "owner/repo/vae", + }, + } + ) + + self.assertEqual(args.component_paths, {"vae": "owner/repo/vae"}) + self.assertEqual( + args.component_weights_paths, + {"text_encoder": "owner/repo/text_encoder/model.safetensors"}, + ) + def test_component_attention_backends_are_normalized(self): args = self._from_dict_without_model_resolution( { diff --git a/python/sglang/multimodal_gen/test/unit/test_text_encoder_loader.py b/python/sglang/multimodal_gen/test/unit/test_text_encoder_loader.py index bd7380ea4..a1e6dc461 100644 --- a/python/sglang/multimodal_gen/test/unit/test_text_encoder_loader.py +++ b/python/sglang/multimodal_gen/test/unit/test_text_encoder_loader.py @@ -187,10 +187,31 @@ class TestTextEncoderQuantization(unittest.TestCase): TextEncoder, {}, "/model/text_encoder", + "/model/text_encoder", "text_encoder", ) self.assertIs(model_config.quant_config, self.serialized) + def test_weight_file_metadata_configures_native_encoder(self): + model_config = SimpleNamespace(quant_config=None) + self.get_quant_config.return_value = None + with mock.patch( + "sglang.multimodal_gen.runtime.loader.component_loaders." + "text_encoder_loader.get_quant_config_from_safetensors_metadata", + return_value=self.serialized, + ) as get_file_quant_config: + _configure_encoder_quantization( + model_config, + TextEncoder, + {}, + "/model/text_encoder", + "/weights/encoder.safetensors", + "text_encoder", + ) + + self.assertIs(model_config.quant_config, self.serialized) + get_file_quant_config.assert_called_once_with("/weights/encoder.safetensors") + def test_encoder_must_use_native_loader(self): model_config = SimpleNamespace(quant_config=None) with self.assertRaisesRegex( @@ -201,6 +222,7 @@ class TestTextEncoderQuantization(unittest.TestCase): nn.Module, {}, "/model/text_encoder", + "/model/text_encoder", "text_encoder", ) @@ -225,6 +247,7 @@ class TestTextEncoderQuantization(unittest.TestCase): SimpleNamespace(architectures=[architecture], quant_config=None), component_config, "/model/text_encoder", + "/model/text_encoder", "text_encoder", ) self.get_quant_config.assert_not_called() @@ -244,6 +267,7 @@ class TestTextEncoderQuantization(unittest.TestCase): } }, "/model/text_encoder", + "/model/text_encoder", "text_encoder", ) @@ -264,6 +288,7 @@ class TestTextEncoderQuantization(unittest.TestCase): } }, "/model/text_encoder", + "/model/text_encoder", "text_encoder", ) @@ -284,6 +309,7 @@ class TestTextEncoderQuantization(unittest.TestCase): } }, "/model/text_encoder", + "/model/text_encoder", "text_encoder", ) diff --git a/python/sglang/multimodal_gen/test/unit/test_weight_source.py b/python/sglang/multimodal_gen/test/unit/test_weight_source.py index fe7c77ec4..8332d6938 100644 --- a/python/sglang/multimodal_gen/test/unit/test_weight_source.py +++ b/python/sglang/multimodal_gen/test/unit/test_weight_source.py @@ -4,12 +4,25 @@ from unittest.mock import patch import pytest from sglang.multimodal_gen.runtime.weights.source import ( + is_explicit_weight_file_reference, + materialize_weight, parse_weight_source, resolve_weight, resolve_weight_inventory, ) +def test_explicit_weight_file_reference_does_not_claim_directories(tmp_path): + component = tmp_path / "component.safetensors" + component.mkdir() + + assert not is_explicit_weight_file_reference(str(component)) + assert is_explicit_weight_file_reference("owner/repo/model.safetensors") + assert is_explicit_weight_file_reference( + "https://huggingface.co/owner/repo/resolve/main/model.gguf?download=true" + ) + + def test_parse_weight_source_accepts_repo_subfolder_and_exact_url(): subfolder = parse_weight_source("owner/repo/text_encoder", revision="v1") repo_file = parse_weight_source("owner/repo/adapter.safetensors") @@ -80,3 +93,10 @@ def test_weight_source_rejects_ambiguous_files(tmp_path): with pytest.raises(ValueError, match="multiple independent weight files"): resolve_weight(str(tmp_path)) + + +def test_materialize_local_weight_returns_selected_file(tmp_path): + checkpoint = tmp_path / "model.safetensors" + checkpoint.write_bytes(b"fixture") + + assert materialize_weight(resolve_weight(str(checkpoint))) == str(checkpoint)