From 881cbfe54c98356cfa1eaa134aa4d0be702fc90f Mon Sep 17 00:00:00 2001 From: Mick Date: Mon, 31 Aug 2026 11:12:07 +0800 Subject: [PATCH] [diffusion] feat: add exact component precision overrides (#36991) --- docs/docs/sglang-diffusion/api/cli.mdx | 6 ++ .../component_loaders/component_loader.py | 18 ++++++ .../component_loaders/image_encoder_loader.py | 11 +++- .../sound_tokenizer_loader.py | 16 ++++-- .../component_loaders/text_encoder_loader.py | 15 ++++- .../component_loaders/upsampler_loader.py | 6 +- .../loader/component_loaders/vae_loader.py | 5 ++ .../pipelines_core/stages/image_encoding.py | 21 ++++++- .../pipelines_core/stages/text_encoding.py | 7 +++ .../runtime/server_args/server_args.py | 57 +++++++++++++++++-- .../multimodal_gen/runtime/utils/precision.py | 31 +++++++++- .../test_component_quantization_admission.py | 9 +++ .../test/unit/test_component_residency.py | 40 ++++++++++++- .../unit/test_decoding_stage_parallelism.py | 17 +++++- .../test/unit/test_disagg_roles.py | 3 +- .../test/unit/test_ideogram4.py | 1 + .../test/unit/test_image_encoder_loader.py | 19 +++++++ .../test/unit/test_precision_consistency.py | 31 +++++++++- .../test/unit/test_server_args.py | 22 +++++++ .../test/unit/test_text_encoder_loader.py | 1 + .../test/unit/test_text_encoding_cache.py | 19 +++++++ .../test/unit/test_vae_decoder_store.py | 1 + .../test/unit/test_vae_loader.py | 54 ++++++++++++++++++ .../test/unit/test_vae_loader_decode_dtype.py | 1 + 24 files changed, 384 insertions(+), 27 deletions(-) diff --git a/docs/docs/sglang-diffusion/api/cli.mdx b/docs/docs/sglang-diffusion/api/cli.mdx index e5c4426fe..56120ad6f 100644 --- a/docs/docs/sglang-diffusion/api/cli.mdx +++ b/docs/docs/sglang-diffusion/api/cli.mdx @@ -131,6 +131,7 @@ pipeline's registered module name: | --- | --- | --- | --- | | Replace a component | `--component-paths. {MODEL}` | `---path {MODEL}` | Load the replacement component's configuration and weights | | Replace only its weights | `--component-weights-paths. {WEIGHTS}` | `---weights-path {WEIGHTS}` | Retain the base component configuration and replace its weights | +| Select component precision | `--component-precisions. {DTYPE}` | — | Use a loader-supported exact parameter and execution dtype | | Direct-load an eligible component | `--component-direct-gpu-weight-loading.` | None | Use that component's audited direct-GPU loader; it must stay resident | | Quantize an unquantized component online | `--component-quantizations. {METHOD}` | `---quantization {METHOD}` | Apply a method supported by that component's native loader | | Keep selected component layers unquantized | `--component-quantization-ignored-layers. {PATTERN...}` | None | Pass component-local ignored-layer patterns to its online quantizer | @@ -167,6 +168,11 @@ published, model-specific checkpoint examples; for example, all H3 sources and their exact overlays are kept in one [MiniMax-H3 compatibility table](/cookbook/diffusion/MiniMax/MiniMax-H3#checkpoint-and-adapter-formats). +Exact precision overrides are capability-based. Native text and image encoders, +standard VAE components, and native plain-state components support them; +other component loaders reject the option instead of accepting a dtype that +their execution stage would not honor. + Direct-GPU loading is also capability-based. The existing `--direct-gpu-weight-loading` remains the primary DiT path. The component form currently supports standard native `vae` and `video_vae` state dicts on CUDA: diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py index 17e375875..c0b45afb8 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py @@ -172,6 +172,18 @@ class ComponentLoader(ABC): ) -> dict[str, Any]: return {} + def component_load_precision( + self, server_args: ServerArgs, component_name: str + ) -> str | None: + """Return an exact precision override or reject an unsupported one.""" + precision = server_args.component_precisions.get(component_name) + if precision is not None: + raise ComponentCheckpointUnsupportedError( + f"{component_name!r} does not support an exact component precision " + "override" + ) + return None + def supports_direct_gpu_weight_loading_for_component( self, _component_name: str ) -> bool: @@ -269,6 +281,7 @@ class ComponentLoader(ABC): """ self._native_load_manages_placement = False + self.component_load_precision(server_args, component_name) if server_args.should_direct_gpu_weight_load_component( component_name ) and not self.supports_direct_gpu_weight_loading_for_component(component_name): @@ -582,6 +595,11 @@ class ComponentLoader(ABC): class PlainStateDictComponentLoader(ComponentLoader): """Base for native loaders whose current materializer expects plain weights.""" + def component_load_precision( + self, server_args: ServerArgs, component_name: str + ) -> str | None: + return server_args.component_precisions.get(component_name) + @staticmethod def ensure_plain_state_dict_checkpoint(config: object, component_name: str) -> None: try: 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 d83a102ea..d4e32fb48 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 @@ -16,6 +16,13 @@ class ImageEncoderLoader(TextEncoderLoader): component_names = ["image_encoder"] expected_library = "transformers" + def component_load_precision( + self, server_args: ServerArgs, component_name: str + ) -> str | None: + return server_args.component_precisions.get( + component_name, server_args.pipeline_config.image_encoder_precision + ) + def load_customized( self, component_model_path: str, @@ -56,10 +63,12 @@ class ImageEncoderLoader(TextEncoderLoader): # Always start with local device; load_model will adjust for offload if needed # TODO(will): add support for other dtypes + image_encoder_dtype = self.component_load_precision(server_args, component_name) + assert image_encoder_dtype is not None return self.load_model( component_weights_path, encoder_config, server_args, - server_args.pipeline_config.image_encoder_precision, + image_encoder_dtype, component_name=component_name, ) diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/sound_tokenizer_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/sound_tokenizer_loader.py index 5d236b221..94b5d237f 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/sound_tokenizer_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/sound_tokenizer_loader.py @@ -1,5 +1,7 @@ # SPDX-License-Identifier: Apache-2.0 +import torch + from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( PlainStateDictComponentLoader, ) @@ -11,6 +13,7 @@ from sglang.multimodal_gen.runtime.loader.utils import ( from sglang.multimodal_gen.runtime.models.registry import ModelRegistry from sglang.multimodal_gen.runtime.server_args import ServerArgs 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.utils import PRECISION_TO_TYPE logger = init_logger(__name__) @@ -34,18 +37,19 @@ class SoundTokenizerLoader(PlainStateDictComponentLoader): server_args.model_paths[component_name] = component_model_path - try: - precision = server_args.pipeline_config.vae_precision - except AttributeError: - precision = "bf16" - dtype = PRECISION_TO_TYPE[precision] + dtype = resolve_component_precision(server_args, component_name) + if dtype is None: + try: + dtype = PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision] + except AttributeError: + dtype = torch.bfloat16 target_device = self.target_device( server_args.should_start_component_on_cpu(component_name) ) with set_default_torch_dtype(dtype), skip_init_modules(): model_cls, _ = ModelRegistry.resolve_model_cls(class_name) - model = model_cls(config).to(target_device) + model = model_cls(config).to(device=target_device, dtype=dtype) loaded = load_safetensors_state_dict(component_weights_path) incompatible = model.load_state_dict(loaded, strict=False) 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 c2f5eeb25..5914e19de 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 @@ -500,6 +500,16 @@ class TextEncoderLoader(ComponentLoader): expected_library = "transformers" supports_online_quantization_override = True + def component_load_precision( + self, server_args: ServerArgs, component_name: str + ) -> str | None: + override = server_args.component_precisions.get(component_name) + if override is not None: + return override + return server_args.pipeline_config.text_encoder_precisions[ + self._extract_encoder_index(component_name) + ] + def should_raise_customized_load_error( self, server_args: ServerArgs, component_name: str ) -> bool: @@ -781,9 +791,8 @@ class TextEncoderLoader(ComponentLoader): server_args.encoder_parallel, prefer_dp=prefer_dp, ) - encoder_dtype = server_args.pipeline_config.text_encoder_precisions[ - encoder_index - ] + encoder_dtype = self.component_load_precision(server_args, component_name) + assert encoder_dtype is not None # TODO(will): add support for other dtypes try: return self.load_model( diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/upsampler_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/upsampler_loader.py index 30647c760..4f266d411 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/upsampler_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/upsampler_loader.py @@ -16,6 +16,7 @@ from sglang.multimodal_gen.runtime.models.upsampler.latent_upsampler import ( from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_model from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.runtime.utils.precision import resolve_component_precision logger = init_logger(__name__) @@ -226,12 +227,15 @@ class UpsamplerLoader(PlainStateDictComponentLoader): component_name ) target_device = self.target_device(component_starts_on_cpu) + dtype = resolve_component_precision(server_args, component_name) + if dtype is None: + dtype = torch.bfloat16 with torch.device("meta"): model = LatentUpsampler(**config) model.load_state_dict(state_dict, assign=True) - model = model.to(device=target_device, dtype=torch.bfloat16).eval() + model = model.to(device=target_device, dtype=dtype).eval() logger.info("Loaded LatentUpsampler to %s", target_device) return model diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/vae_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/vae_loader.py index 82cafd6d3..89c9b6a51 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/vae_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/vae_loader.py @@ -398,6 +398,11 @@ class VAELoader(ComponentLoader): ) -> bool: return component_name in ("vae", "video_vae") + def component_load_precision( + self, server_args: ServerArgs, component_name: str + ) -> str | None: + return server_args.component_precisions.get(component_name) + @staticmethod def resolve_model_weights_path( component_model_path: str, diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py index 76bbc2801..ba14eb7eb 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py @@ -42,6 +42,7 @@ from sglang.multimodal_gen.runtime.utils.precision import ( align_tensor_to_module_dtype, autocast_context, autocast_enabled, + resolve_component_precision_override, resolve_precision, temporary_module_dtype, ) @@ -156,9 +157,25 @@ class ImageEncodingStage(PipelineStage): stage_name = self._component_stage_name(stage_name) uses = [] if self.image_encoder is not None: - uses.append(ComponentUse(stage_name, "image_encoder")) + uses.append( + ComponentUse( + stage_name, + "image_encoder", + target_dtype=resolve_component_precision_override( + server_args, "image_encoder" + ), + ) + ) if self.text_encoder is not None: - uses.append(ComponentUse(stage_name, "text_encoder")) + uses.append( + ComponentUse( + stage_name, + "text_encoder", + target_dtype=resolve_component_precision_override( + server_args, "text_encoder" + ), + ) + ) return uses def encoding_image_edit(self, outputs, image_inputs, pipeline_config): diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/text_encoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/text_encoding.py index 509126830..8662f2cfb 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/text_encoding.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/text_encoding.py @@ -41,6 +41,9 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import ( ) from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.runtime.utils.precision import ( + resolve_component_precision_override, +) logger = init_logger(__name__) @@ -175,6 +178,10 @@ class TextEncodingStage(ConditionEncodingStage): stage_name=stage_name, component_name="text_encoder" if i == 0 else f"text_encoder_{i + 1}", preferred_ready_after_request=i == 0, + target_dtype=resolve_component_precision_override( + server_args, + "text_encoder" if i == 0 else f"text_encoder_{i + 1}", + ), ) for i in range(len(self.text_encoders)) ] 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 71bae888a..0e7d92179 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py @@ -81,6 +81,7 @@ from sglang.multimodal_gen.runtime.weights.source import ( is_explicit_weight_file_reference, ) from sglang.multimodal_gen.utils import ( + PRECISION_TO_TYPE, FlexibleArgumentParser, StoreBoolean, expand_path_fields, @@ -110,6 +111,23 @@ def is_ltx2_two_stage_pipeline_name(pipeline_class_name: str | None) -> bool: return pipeline_class_name in LTX2_TWO_STAGE_PIPELINE_NAMES +def _normalize_component_precisions(value: object) -> dict[str, str]: + if not isinstance(value, dict): + raise ValueError("component_precisions must be a mapping") + + normalized: dict[str, str] = {} + for component, precision in value.items(): + component_name = str(component).strip().replace("-", "_") + precision_name = str(precision).strip().lower() + if not component_name or precision_name not in PRECISION_TO_TYPE: + raise ValueError( + "Component precision entries require a component and one of " + f"{sorted(PRECISION_TO_TYPE)}" + ) + normalized[component_name] = precision_name + return normalized + + class Backend(str, Enum): """ Enumeration for different model backends. @@ -312,6 +330,9 @@ class ServerArgs(DisaggServerArgsMixin): # Explicit quantization override for one component. Self-describing # checkpoints remain auto-detected and do not need this override. component_quantizations: dict[str, str] = field(default_factory=dict) + # Exact load and execution precision overrides for components whose native + # loader advertises this capability. + component_precisions: dict[str, str] = field(default_factory=dict) # Component-local layer name patterns to skip during online quantization. component_quantization_ignored_layers: dict[str, list[str]] = field( default_factory=dict @@ -1768,6 +1789,9 @@ class ServerArgs(DisaggServerArgsMixin): component_weights_paths[component] = path self.component_paths = component_paths self.component_weights_paths = component_weights_paths + self.component_precisions = _normalize_component_precisions( + self.component_precisions + ) normalized_direct_gpu_loading: dict[str, bool] = {} for component, enabled in self.component_direct_gpu_weight_loading.items(): component_name = str(component).strip().replace("-", "_") @@ -2873,7 +2897,8 @@ class ServerArgs(DisaggServerArgsMixin): unknown_args: list[str], *, option_prefixes: tuple[str, ...], - alias_suffix: str, + alias_suffix: str | None, + expand_values: bool = True, ) -> tuple[dict[str, str], list[str]]: component_values: dict[str, str] = {} remaining: list[str] = [] @@ -2888,6 +2913,7 @@ class ServerArgs(DisaggServerArgsMixin): break if ( component is None + and alias_suffix is not None and key_part.startswith("--") and key_part.endswith(alias_suffix) ): @@ -2909,10 +2935,12 @@ class ServerArgs(DisaggServerArgsMixin): remaining.append(arg) i += 1 - return { - component: os.path.expanduser(value) - for component, value in component_values.items() - }, remaining + if expand_values: + component_values = { + component: os.path.expanduser(value) + for component, value in component_values.items() + } + return component_values, remaining @classmethod def _extract_component_paths( @@ -3002,6 +3030,19 @@ class ServerArgs(DisaggServerArgsMixin): alias_suffix="-quantization", ) + @classmethod + def _extract_component_precisions( + cls, + unknown_args: list[str], + ) -> tuple[dict[str, str], list[str]]: + """Extract exact component precision overrides.""" + return cls._extract_dynamic_component_map( + unknown_args, + option_prefixes=("--component-precisions.", "--component_precisions."), + alias_suffix=None, + expand_values=False, + ) + @staticmethod def _extract_component_quantization_ignored_layers( unknown_args: list[str], @@ -3092,6 +3133,7 @@ class ServerArgs(DisaggServerArgsMixin): dynamic_quantizations, remaining = cls._extract_component_quantizations( unknown_args ) + dynamic_precisions, remaining = cls._extract_component_precisions(remaining) dynamic_direct_gpu_loading, remaining = ( cls._extract_component_direct_gpu_weight_loading(remaining) ) @@ -3138,6 +3180,11 @@ class ServerArgs(DisaggServerArgsMixin): existing.update(dynamic_quantizations) provided_args["component_quantizations"] = existing explicit_arg_names.add("component_quantizations") + if dynamic_precisions: + existing = dict(provided_args.get("component_precisions") or {}) + existing.update(dynamic_precisions) + provided_args["component_precisions"] = existing + explicit_arg_names.add("component_precisions") if dynamic_direct_gpu_loading: existing = dict( provided_args.get("component_direct_gpu_weight_loading") or {} diff --git a/python/sglang/multimodal_gen/runtime/utils/precision.py b/python/sglang/multimodal_gen/runtime/utils/precision.py index 6397d781f..0cb673505 100644 --- a/python/sglang/multimodal_gen/runtime/utils/precision.py +++ b/python/sglang/multimodal_gen/runtime/utils/precision.py @@ -24,6 +24,14 @@ def resolve_precision( precision_attr: Optional[str] = None, field_name: Optional[str] = None, ) -> torch.dtype: + component_precision = server_args.component_precisions.get( + component_or_precision_attr + ) + if component_precision is not None: + return precision_to_dtype( + component_precision, + f"component_precisions.{component_or_precision_attr}", + ) precision_attr = precision_attr or component_or_precision_attr precision = getattr(server_args.pipeline_config, precision_attr) return precision_to_dtype(precision, field_name or precision_attr) @@ -35,6 +43,12 @@ def resolve_decode_precision( *, quality: str | None = None, ) -> torch.dtype: + component_precision = server_args.component_precisions.get(component_name) + if component_precision is not None: + return precision_to_dtype( + component_precision, f"component_precisions.{component_name}" + ) + pipeline_config = server_args.pipeline_config if component_name in ("audio_vae", "vocoder"): return resolve_precision( @@ -58,10 +72,21 @@ def resolve_decode_precision( ) -def resolve_component_precision(server_args, module_name: str) -> Optional[torch.dtype]: - pipeline_config = getattr(server_args, "pipeline_config", None) - if pipeline_config is None: +def resolve_component_precision_override( + server_args, module_name: str +) -> Optional[torch.dtype]: + exact_precision = server_args.component_precisions.get(module_name) + if exact_precision is None: return None + return precision_to_dtype(exact_precision, f"component_precisions.{module_name}") + + +def resolve_component_precision(server_args, module_name: str) -> Optional[torch.dtype]: + exact_precision = resolve_component_precision_override(server_args, module_name) + if exact_precision is not None: + return exact_precision + + pipeline_config = server_args.pipeline_config if module_name in ("audio_vae", "vocoder"): precision_attr = "audio_vae_precision" diff --git a/python/sglang/multimodal_gen/test/unit/test_component_quantization_admission.py b/python/sglang/multimodal_gen/test/unit/test_component_quantization_admission.py index f03c39b86..a7299708c 100644 --- a/python/sglang/multimodal_gen/test/unit/test_component_quantization_admission.py +++ b/python/sglang/multimodal_gen/test/unit/test_component_quantization_admission.py @@ -39,6 +39,13 @@ class _TestLoader(PlainStateDictComponentLoader): class TestComponentQuantizationAdmission(unittest.TestCase): + def test_plain_loader_admits_its_exact_precision(self): + server_args = SimpleNamespace(component_precisions={"vocoder": "fp16"}) + + self.assertEqual( + _TestLoader().component_load_precision(server_args, "vocoder"), "fp16" + ) + def test_direct_gpu_selection_requires_a_declared_component(self): server_args = SimpleNamespace( component_direct_gpu_weight_loading={"missing_vae": True} @@ -51,6 +58,7 @@ class TestComponentQuantizationAdmission(unittest.TestCase): def test_direct_gpu_selector_is_rejected_by_unqualified_loader(self): server_args = SimpleNamespace( + component_precisions={}, component_quantizations={}, should_direct_gpu_weight_load_component=lambda component: component == "vocoder", @@ -65,6 +73,7 @@ class TestComponentQuantizationAdmission(unittest.TestCase): def test_direct_gpu_selector_is_rejected_by_unqualified_component(self): server_args = SimpleNamespace( + component_precisions={}, component_quantizations={}, should_direct_gpu_weight_load_component=lambda component: component == "audio_vae", diff --git a/python/sglang/multimodal_gen/test/unit/test_component_residency.py b/python/sglang/multimodal_gen/test/unit/test_component_residency.py index 518f56d60..c6cd76686 100644 --- a/python/sglang/multimodal_gen/test/unit/test_component_residency.py +++ b/python/sglang/multimodal_gen/test/unit/test_component_residency.py @@ -16,6 +16,9 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency_ ComponentOffloadStrategy, ResidentStrategy, ) +from sglang.multimodal_gen.runtime.pipelines_core.stages.image_encoding import ( + ImageEncodingStage, +) from sglang.multimodal_gen.runtime.pipelines_core.stages.realtime.text_encoding import ( RealtimeTextEncodingStage, ) @@ -253,13 +256,48 @@ def test_realtime_text_encoder_use_starts_at_call_site(): stage.text_encoders = [None] stage._registered_stage_name = None - uses = stage.component_uses(SimpleNamespace(), "RealtimeTextEncodingStage") + uses = stage.component_uses( + SimpleNamespace(component_precisions={}, pipeline_config=None), + "RealtimeTextEncodingStage", + ) assert len(uses) == 1 assert uses[0].component_name == "text_encoder" assert uses[0].start_at_stage_entry is False +def test_image_encoder_use_has_exact_precision(): + stage = ImageEncodingStage.__new__(ImageEncodingStage) + stage.image_encoder = object() + stage.text_encoder = None + stage._registered_stage_name = None + + uses = stage.component_uses( + SimpleNamespace(component_precisions={"image_encoder": "fp16"}), + "ImageEncodingStage", + ) + + assert [(use.component_name, use.target_dtype) for use in uses] == [ + ("image_encoder", torch.float16) + ] + + +def test_image_encoder_use_preserves_loaded_dtype_without_override(): + stage = ImageEncodingStage.__new__(ImageEncodingStage) + stage.image_encoder = object() + stage.text_encoder = None + stage._registered_stage_name = None + + uses = stage.component_uses( + SimpleNamespace(component_precisions={}), + "ImageEncodingStage", + ) + + assert [(use.component_name, use.target_dtype) for use in uses] == [ + ("image_encoder", None) + ] + + def test_qwen_layered_uses_loaded_text_encoder(monkeypatch): from sglang.multimodal_gen.runtime.pipelines import qwen_image diff --git a/python/sglang/multimodal_gen/test/unit/test_decoding_stage_parallelism.py b/python/sglang/multimodal_gen/test/unit/test_decoding_stage_parallelism.py index edc376c0e..f2f4f1166 100644 --- a/python/sglang/multimodal_gen/test/unit/test_decoding_stage_parallelism.py +++ b/python/sglang/multimodal_gen/test/unit/test_decoding_stage_parallelism.py @@ -23,16 +23,31 @@ class TestDecodingStageParallelism(unittest.TestCase): def test_component_use_honors_decode_precision_override(self): stage = DecodingStage(FakeVAE()) server_args = SimpleNamespace( + component_precisions={}, pipeline_config=SimpleNamespace( vae_precision="fp32", vae_decode_precision="bf16", - ) + ), ) [component_use] = stage.component_uses(server_args) self.assertEqual(component_use.target_dtype, torch.bfloat16) + def test_component_use_honors_exact_component_precision(self): + stage = DecodingStage(FakeVAE()) + server_args = SimpleNamespace( + component_precisions={"vae": "fp16"}, + pipeline_config=SimpleNamespace( + vae_precision="fp32", + vae_decode_precision="bf16", + ), + ) + + [component_use] = stage.component_uses(server_args) + + self.assertEqual(component_use.target_dtype, torch.float16) + def test_cfg_parallel_uses_replicated_decode_when_decode_group_has_multiple_ranks( self, ): diff --git a/python/sglang/multimodal_gen/test/unit/test_disagg_roles.py b/python/sglang/multimodal_gen/test/unit/test_disagg_roles.py index 81d66c0b0..1237b0720 100644 --- a/python/sglang/multimodal_gen/test/unit/test_disagg_roles.py +++ b/python/sglang/multimodal_gen/test/unit/test_disagg_roles.py @@ -520,7 +520,8 @@ class TestImageVAEEncodingStageComponentName(_GlobalStageArgsMixin, unittest.Tes def test_component_name_can_follow_non_default_vae_key(self): stage = ImageVAEEncodingStage(vae=object(), component_name="video_vae") server_args = SimpleNamespace( - pipeline_config=SimpleNamespace(vae_precision="bf16") + component_precisions={}, + pipeline_config=SimpleNamespace(vae_precision="bf16"), ) uses = stage.component_uses(server_args, "image_vae_encoding") diff --git a/python/sglang/multimodal_gen/test/unit/test_ideogram4.py b/python/sglang/multimodal_gen/test/unit/test_ideogram4.py index 64155e085..657c58840 100644 --- a/python/sglang/multimodal_gen/test/unit/test_ideogram4.py +++ b/python/sglang/multimodal_gen/test/unit/test_ideogram4.py @@ -186,6 +186,7 @@ def _fake_server_args(cfg=None): disable_autocast=False, enable_cfg_parallel=False, attention_backend_config=None, + component_precisions={}, kv_gather_degree=1, sp_split_auto=False, ) 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 b020c966b..f7c97e090 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 @@ -8,6 +8,7 @@ from torch import nn from sglang.multimodal_gen.configs.models.encoders.clip import CLIPVisionConfig from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import ( ComponentCheckpointUnsupportedError, + ComponentLoader, NativeComponentLoaderRequired, ) from sglang.multimodal_gen.runtime.loader.component_loaders.image_encoder_loader import ( @@ -39,6 +40,7 @@ class TestImageEncoderQuantizationAdmission(unittest.TestCase): ), component_weights_paths={}, component_quantizations={}, + component_precisions={}, encoder_parallel="replicate", resolve_component_attention_backend=lambda _name: (None, None), should_direct_gpu_weight_load_component=lambda _name: False, @@ -83,6 +85,20 @@ class TestImageEncoderQuantizationAdmission(unittest.TestCase): {"qkv_proj": ["q_proj", "k_proj", "v_proj"]}, ) + def test_exact_image_encoder_precision_override(self): + self.server_args.component_precisions["image_encoder"] = "fp16" + + self.assertEqual( + self.loader.component_load_precision(self.server_args, "image_encoder"), + "fp16", + ) + + def test_unadmitted_component_precision_fails_closed(self): + with self.assertRaises(ComponentCheckpointUnsupportedError): + ComponentLoader().component_load_precision( + SimpleNamespace(component_precisions={"vae": "fp16"}), "vae" + ) + def test_unknown_transformers_quantized_architecture_falls_back(self): config = self._component_config("UnknownVisionModel", quantized=True) with self._config_patch(config): @@ -131,6 +147,7 @@ class TestImageEncoderNativeLoading(unittest.TestCase): from_pretrained=mock.Mock(return_value=loaded_encoder) ) server_args = SimpleNamespace( + component_precisions={}, pipeline_config=SimpleNamespace(image_encoder_precision="bf16"), explicit_residency_mode=mock.Mock(return_value=None), require_component_resident=mock.Mock(), @@ -182,6 +199,7 @@ class TestImageEncoderNativeLoading(unittest.TestCase): ) model_class = SimpleNamespace(from_pretrained=mock.Mock()) server_args = SimpleNamespace( + component_precisions={}, pipeline_config=SimpleNamespace(image_encoder_precision="bf16"), explicit_residency_mode=mock.Mock(return_value=COMPONENT_OFFLOAD), require_component_resident=mock.Mock(), @@ -228,6 +246,7 @@ class TestImageEncoderNativeLoading(unittest.TestCase): ) server_args = SimpleNamespace( component_quantizations={}, + component_precisions={}, pipeline_config=SimpleNamespace( image_encoder_precision="bf16", native_only_components=(), diff --git a/python/sglang/multimodal_gen/test/unit/test_precision_consistency.py b/python/sglang/multimodal_gen/test/unit/test_precision_consistency.py index 25afca467..18de56c2a 100644 --- a/python/sglang/multimodal_gen/test/unit/test_precision_consistency.py +++ b/python/sglang/multimodal_gen/test/unit/test_precision_consistency.py @@ -113,7 +113,9 @@ class TestDiffusionPrecisionConsistency(unittest.TestCase): "text_encoder_precisions": ["fp16", "bf16"], } config.update(overrides) - return SimpleNamespace(pipeline_config=SimpleNamespace(**config)) + return SimpleNamespace( + component_precisions={}, pipeline_config=SimpleNamespace(**config) + ) def test_precision_lookup(self): server_args = self._server_args() @@ -163,8 +165,27 @@ class TestDiffusionPrecisionConsistency(unittest.TestCase): self._server_args(vae_decode_precision_high="fp8"), quality="high" ) + def test_exact_vae_precision_overrides_load_and_decode_defaults(self): + server_args = self._server_args(vae_decode_precision="bf16") + server_args.component_precisions["vae"] = "fp16" + server_args.component_precisions["video_vae"] = "bf16" + + self.assertEqual( + resolve_precision(server_args, "vae", precision_attr="vae_precision"), + torch.float16, + ) + self.assertEqual(resolve_decode_precision(server_args, "vae"), torch.float16) + self.assertEqual( + resolve_precision(server_args, "video_vae", precision_attr="vae_precision"), + torch.bfloat16, + ) + self.assertEqual( + resolve_decode_precision(server_args, "video_vae"), torch.bfloat16 + ) + def test_component_precision_mapping(self): server_args = self._server_args() + server_args.component_precisions["text_encoder_2"] = "fp32" expected = { "vae": torch.float16, "video_vae": torch.float16, @@ -178,7 +199,7 @@ class TestDiffusionPrecisionConsistency(unittest.TestCase): "dual_tower_bridge": torch.float32, "image_encoder": torch.float16, "text_encoder": torch.float16, - "text_encoder_2": torch.bfloat16, + "text_encoder_2": torch.float32, } for module_name, expected_dtype in expected.items(): @@ -188,7 +209,11 @@ class TestDiffusionPrecisionConsistency(unittest.TestCase): module_name, ) - self.assertIsNone(resolve_component_precision(SimpleNamespace(), "vae")) + self.assertIsNone( + resolve_component_precision( + SimpleNamespace(component_precisions={}, pipeline_config=None), "vae" + ) + ) self.assertIsNone( resolve_component_precision(server_args, "unregistered_component") ) 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 ea54ee9ae..3a76697a2 100644 --- a/python/sglang/multimodal_gen/test/unit/test_server_args.py +++ b/python/sglang/multimodal_gen/test/unit/test_server_args.py @@ -313,6 +313,28 @@ class TestServerArgsPathExpansion(unittest.TestCase): server_args.component_attention_backends, {"text_encoder": "torch_sdpa"} ) + def test_dynamic_component_precision_cli_args(self): + parser = FlexibleArgumentParser() + ServerArgs.add_cli_args(parser) + argv = [ + "--model-path", + "/fake", + "--component-precisions.text-encoder-2", + "fp32", + ] + + with ( + patch.object(sys, "argv", ["sglang"] + argv), + patch.object( + PipelineConfig, "from_kwargs", return_value=QwenImagePipelineConfig() + ), + _mock_cuda_platform(), + ): + args, unknown_args = parser.parse_known_args(argv) + server_args = ServerArgs.from_cli_args(args, unknown_args) + + self.assertEqual(server_args.component_precisions, {"text_encoder_2": "fp32"}) + def test_layerwise_offload_components_imply_layerwise(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 371876947..c4c4c2138 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 @@ -150,6 +150,7 @@ class TestTextEncoderClassResolution(unittest.TestCase): from_pretrained=mock.Mock(return_value=loaded_encoder) ) server_args = SimpleNamespace( + component_precisions={}, pipeline_config=SimpleNamespace(text_encoder_precisions=["bf16"]), explicit_residency_mode=mock.Mock(return_value=None), require_component_resident=mock.Mock(), diff --git a/python/sglang/multimodal_gen/test/unit/test_text_encoding_cache.py b/python/sglang/multimodal_gen/test/unit/test_text_encoding_cache.py index abcb84462..a0ed2c93b 100644 --- a/python/sglang/multimodal_gen/test/unit/test_text_encoding_cache.py +++ b/python/sglang/multimodal_gen/test/unit/test_text_encoding_cache.py @@ -78,6 +78,25 @@ def test_negative_text_cache_key_tracks_encode_options(): assert stage.calls == 3 +def test_component_uses_exact_encoder_precision(): + with patch(_GLOBAL_ARGS_PATCH) as mock_global_args: + mock_global_args.return_value = MagicMock() + stage = TextEncodingStage(text_encoders=[object(), object()], tokenizers=[]) + server_args = make_server_args( + component_precisions={"text_encoder_2": "fp32"}, + pipeline_config=SimpleNamespace( + text_encoder_configs=[], text_encoder_precisions=["bf16", "bf16"] + ), + ) + + uses = stage.component_uses(server_args) + + assert [(use.component_name, use.target_dtype) for use in uses] == [ + ("text_encoder", None), + ("text_encoder_2", torch.float32), + ] + + def test_negative_text_cache_skips_warmup(): stage = DummyTextEncodingStage() server_args = make_server_args() diff --git a/python/sglang/multimodal_gen/test/unit/test_vae_decoder_store.py b/python/sglang/multimodal_gen/test/unit/test_vae_decoder_store.py index 543d5a1a9..13b3b431c 100644 --- a/python/sglang/multimodal_gen/test/unit/test_vae_decoder_store.py +++ b/python/sglang/multimodal_gen/test/unit/test_vae_decoder_store.py @@ -36,6 +36,7 @@ class _TinyVAE(nn.Module): def _server_args(): return SimpleNamespace( + component_precisions={}, pipeline_config=SimpleNamespace( vae_decode_precision="fp16", vae_precision="fp32", diff --git a/python/sglang/multimodal_gen/test/unit/test_vae_loader.py b/python/sglang/multimodal_gen/test/unit/test_vae_loader.py index b4e26aedb..ee73b768b 100644 --- a/python/sglang/multimodal_gen/test/unit/test_vae_loader.py +++ b/python/sglang/multimodal_gen/test/unit/test_vae_loader.py @@ -36,6 +36,9 @@ from sglang.multimodal_gen.runtime.managers.memory_managers import ( host_memory_budget, ) from sglang.multimodal_gen.runtime.models.vaes import wanvae +from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ltx_2.decoding_av import ( + LTX2AVDecodingStage, +) class _FakeServerArgs: @@ -47,6 +50,7 @@ class _FakeServerArgs: self.trust_remote_code = True self.layerwise_components = set() self.component_quantizations = {} + self.component_precisions = {} self.component_direct_gpu_weight_loading = {} def resolve_component_attention_backend(self, _component_name): @@ -262,6 +266,56 @@ class TestDirectGPUVAEState(unittest.TestCase): class TestVAELoader(unittest.TestCase): + def test_exact_precision_is_admitted_for_every_vae_component(self): + loader = vae_loader.VAELoader() + server_args = _FakeServerArgs(QwenImagePipelineConfig()) + server_args.component_precisions = {"vae": "bf16"} + + self.assertEqual(loader.component_load_precision(server_args, "vae"), "bf16") + + server_args.component_precisions = {"audio_vae": "bf16"} + self.assertEqual( + loader.component_load_precision(server_args, "audio_vae"), "bf16" + ) + + def test_exact_audio_vae_precision_reaches_customized_loader(self): + loader = vae_loader.VAELoader() + server_args = _FakeServerArgs(LTX2PipelineConfig()) + server_args.component_precisions = {"audio_vae": "bf16"} + audio_vae = nn.Identity() + + with ( + patch.object(loader, "load_customized", return_value=audio_vae) as load, + patch.object( + vae_loader.current_platform, + "get_available_gpu_memory", + side_effect=[10.0, 9.0], + ), + patch( + "sglang.multimodal_gen.runtime.loader.component_loaders." + "component_loader.get_memory_usage_of_component", + return_value=1.0, + ), + ): + loaded, _ = loader.load( + "/component/audio_vae", server_args, "audio_vae", "diffusers" + ) + + self.assertIs(loaded, audio_vae) + load.assert_called_once_with("/component/audio_vae", server_args, "audio_vae") + + def test_ltx_audio_vae_use_honors_exact_component_precision(self): + stage = LTX2AVDecodingStage( + vae=torch.nn.Identity(), + audio_vae=torch.nn.Identity(), + vocoder=torch.nn.Identity(), + ) + server_args = _FakeServerArgs(LTX2PipelineConfig()) + server_args.component_precisions = {"audio_vae": "fp32"} + + uses = {use.component_name: use for use in stage.component_uses(server_args)} + self.assertEqual(uses["audio_vae"].target_dtype, torch.float32) + def test_weights_override_keeps_base_component_config(self): loader = vae_loader.VAELoader() server_args = _FakeServerArgs(QwenImagePipelineConfig()) diff --git a/python/sglang/multimodal_gen/test/unit/test_vae_loader_decode_dtype.py b/python/sglang/multimodal_gen/test/unit/test_vae_loader_decode_dtype.py index bf23481a0..9809b337e 100644 --- a/python/sglang/multimodal_gen/test/unit/test_vae_loader_decode_dtype.py +++ b/python/sglang/multimodal_gen/test/unit/test_vae_loader_decode_dtype.py @@ -27,6 +27,7 @@ class _RecordingVAE: def _server_args(decode_precision="fp16", disable_autocast=False): return SimpleNamespace( + component_precisions={}, pipeline_config=SimpleNamespace( vae_decode_precision=decode_precision, vae_precision="fp32",