[diffusion] feat: add exact component precision overrides (#36991)

This commit is contained in:
Mick
2026-08-31 11:12:07 +08:00
committed by GitHub
parent 4dc7dc8518
commit 881cbfe54c
24 changed files with 384 additions and 27 deletions
@@ -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:
@@ -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,
)
@@ -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)
@@ -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(
@@ -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
@@ -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,
@@ -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):
@@ -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))
]
@@ -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 {}
@@ -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"
@@ -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",
@@ -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
@@ -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,
):
@@ -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")
@@ -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,
)
@@ -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=(),
@@ -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")
)
@@ -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(
{
@@ -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(),
@@ -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()
@@ -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",
@@ -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())
@@ -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",