[diffusion] feat: support streaming native vae weights directly to gpu (#37004)
This commit is contained in:
@@ -131,6 +131,7 @@ pipeline's registered module name:
|
|||||||
| --- | --- | --- | --- |
|
| --- | --- | --- | --- |
|
||||||
| Replace a component | `--component-paths.<component> {MODEL}` | `--<component>-path {MODEL}` | Load the replacement component's configuration and weights |
|
| Replace a component | `--component-paths.<component> {MODEL}` | `--<component>-path {MODEL}` | Load the replacement component's configuration and weights |
|
||||||
| Replace only its weights | `--component-weights-paths.<component> {WEIGHTS}` | `--<component>-weights-path {WEIGHTS}` | Retain the base component configuration and replace its weights |
|
| Replace only its weights | `--component-weights-paths.<component> {WEIGHTS}` | `--<component>-weights-path {WEIGHTS}` | Retain the base component configuration and replace its weights |
|
||||||
|
| Direct-load an eligible component | `--component-direct-gpu-weight-loading.<component>` | None | Use that component's audited direct-GPU loader; it must stay resident |
|
||||||
| Quantize an unquantized component online | `--component-quantizations.<component> {METHOD}` | `--<component>-quantization {METHOD}` | Apply a method supported by that component's native loader |
|
| Quantize an unquantized component online | `--component-quantizations.<component> {METHOD}` | `--<component>-quantization {METHOD}` | Apply a method supported by that component's native loader |
|
||||||
| Keep selected component layers unquantized | `--component-quantization-ignored-layers.<component> {PATTERN...}` | None | Pass component-local ignored-layer patterns to its online quantizer |
|
| Keep selected component layers unquantized | `--component-quantization-ignored-layers.<component> {PATTERN...}` | None | Pass component-local ignored-layer patterns to its online quantizer |
|
||||||
|
|
||||||
@@ -166,6 +167,14 @@ published, model-specific checkpoint examples; for example, all H3 sources and
|
|||||||
their exact overlays are kept in one
|
their exact overlays are kept in one
|
||||||
[MiniMax-H3 compatibility table](/cookbook/diffusion/MiniMax/MiniMax-H3#checkpoint-and-adapter-formats).
|
[MiniMax-H3 compatibility table](/cookbook/diffusion/MiniMax/MiniMax-H3#checkpoint-and-adapter-formats).
|
||||||
|
|
||||||
|
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:
|
||||||
|
it streams each safetensors tensor directly to the resident
|
||||||
|
module, rather than materializing a complete CPU state dict. It rejects custom
|
||||||
|
Diffusers `auto_map` classes, quantized checkpoints, tied state entries, and
|
||||||
|
any component-offload or layerwise-offload placement.
|
||||||
|
|
||||||
Component overrides accept a local component directory, a standalone Hub
|
Component overrides accept a local component directory, a standalone Hub
|
||||||
repository, or a Hub component subfolder written as `owner/repo/subfolder`.
|
repository, or a Hub component subfolder written as `owner/repo/subfolder`.
|
||||||
For transformer and native encoder loaders, an explicit weight filename keeps
|
For transformer and native encoder loaders, an explicit weight filename keeps
|
||||||
|
|||||||
@@ -137,6 +137,9 @@ class ComponentLoader(ABC):
|
|||||||
# Gates only --component-quantizations.<name>. Quantization declared by a
|
# Gates only --component-quantizations.<name>. Quantization declared by a
|
||||||
# checkpoint is discovered and admitted by the component's normal loader.
|
# checkpoint is discovered and admitted by the component's normal loader.
|
||||||
supports_online_quantization_override = False
|
supports_online_quantization_override = False
|
||||||
|
# Gates only --component-direct-gpu-weight-loading.<name>. The checkpoint
|
||||||
|
# source stays component-specific because its streaming ABI is loader-owned.
|
||||||
|
supports_direct_gpu_weight_loading = False
|
||||||
supports_fsdp_inference = False
|
supports_fsdp_inference = False
|
||||||
|
|
||||||
_loaders_registered = False
|
_loaders_registered = False
|
||||||
@@ -169,6 +172,11 @@ class ComponentLoader(ABC):
|
|||||||
) -> dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
return {}
|
return {}
|
||||||
|
|
||||||
|
def supports_direct_gpu_weight_loading_for_component(
|
||||||
|
self, _component_name: str
|
||||||
|
) -> bool:
|
||||||
|
return self.supports_direct_gpu_weight_loading
|
||||||
|
|
||||||
def should_raise_customized_load_error(
|
def should_raise_customized_load_error(
|
||||||
self, server_args: ServerArgs, component_name: str
|
self, server_args: ServerArgs, component_name: str
|
||||||
) -> bool:
|
) -> bool:
|
||||||
@@ -252,6 +260,12 @@ class ComponentLoader(ABC):
|
|||||||
|
|
||||||
"""
|
"""
|
||||||
self._native_load_manages_placement = False
|
self._native_load_manages_placement = False
|
||||||
|
if server_args.should_direct_gpu_weight_load_component(
|
||||||
|
component_name
|
||||||
|
) and not self.supports_direct_gpu_weight_loading_for_component(component_name):
|
||||||
|
raise ComponentCheckpointUnsupportedError(
|
||||||
|
f"{component_name!r} does not support direct GPU weight loading"
|
||||||
|
)
|
||||||
self.disable_unsupported_component_fsdp(server_args, component_name)
|
self.disable_unsupported_component_fsdp(server_args, component_name)
|
||||||
component_quantization = server_args.component_quantizations.get(component_name)
|
component_quantization = server_args.component_quantizations.get(component_name)
|
||||||
if (
|
if (
|
||||||
|
|||||||
@@ -25,6 +25,9 @@ from sglang.multimodal_gen.runtime.loader.utils import (
|
|||||||
set_default_torch_dtype,
|
set_default_torch_dtype,
|
||||||
skip_init_modules,
|
skip_init_modules,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.loader.weight_utils import (
|
||||||
|
safetensors_weights_iterator,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
|
from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
|
||||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
@@ -52,13 +55,27 @@ VAE_CHANNELS_LAST_3D_ENV = "SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D"
|
|||||||
|
|
||||||
|
|
||||||
def _require_native_loader_for_quantized_vae(
|
def _require_native_loader_for_quantized_vae(
|
||||||
config: dict, component_name: str, *, native_only: bool = False
|
config: dict,
|
||||||
|
component_name: str,
|
||||||
|
*,
|
||||||
|
native_only: bool = False,
|
||||||
|
direct_gpu_weight_loading: bool = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
try:
|
||||||
quant_spec = resolve_checkpoint_quant_spec(config)
|
quant_spec = resolve_checkpoint_quant_spec(config)
|
||||||
|
except (TypeError, ValueError) as error:
|
||||||
|
raise ComponentCheckpointUnsupportedError(
|
||||||
|
f"Cannot parse checkpoint quantization for {component_name!r}: {error}"
|
||||||
|
) from error
|
||||||
if quant_spec is None:
|
if quant_spec is None:
|
||||||
return
|
return
|
||||||
|
|
||||||
method = quant_spec.declared_method or "unspecified"
|
method = quant_spec.declared_method or "unspecified"
|
||||||
|
if direct_gpu_weight_loading:
|
||||||
|
raise ComponentCheckpointUnsupportedError(
|
||||||
|
f"Direct GPU loading for {component_name!r} cannot restore "
|
||||||
|
f"quant_method={method!r}"
|
||||||
|
)
|
||||||
if native_only:
|
if native_only:
|
||||||
raise ComponentCheckpointUnsupportedError(
|
raise ComponentCheckpointUnsupportedError(
|
||||||
f"{component_name} uses a native-only SGLang implementation that "
|
f"{component_name} uses a native-only SGLang implementation that "
|
||||||
@@ -270,11 +287,116 @@ def _match_checkpoint_dtypes(loaded: dict, target_state: dict) -> dict:
|
|||||||
return loaded
|
return loaded
|
||||||
|
|
||||||
|
|
||||||
|
def _direct_gpu_vae_state_slots(
|
||||||
|
vae: nn.Module, component_name: str
|
||||||
|
) -> tuple[dict[str, torch.Tensor], dict[str, tuple[nn.Module, str, bool]]]:
|
||||||
|
"""Return assignable parameter/buffer slots for a standard native VAE."""
|
||||||
|
if type(vae).state_dict is not nn.Module.state_dict:
|
||||||
|
raise ComponentCheckpointUnsupportedError(
|
||||||
|
f"Direct GPU loading for {component_name!r} requires the standard "
|
||||||
|
"torch.nn.Module state-dict ABI"
|
||||||
|
)
|
||||||
|
|
||||||
|
state = vae.state_dict(keep_vars=True)
|
||||||
|
slots: dict[str, tuple[nn.Module, str, bool]] = {}
|
||||||
|
object_names: dict[int, list[str]] = {}
|
||||||
|
for prefix, module in vae.named_modules():
|
||||||
|
for local_name, parameter in module._parameters.items():
|
||||||
|
if parameter is None:
|
||||||
|
continue
|
||||||
|
name = f"{prefix}.{local_name}" if prefix else local_name
|
||||||
|
slots[name] = (module, local_name, True)
|
||||||
|
object_names.setdefault(id(parameter), []).append(name)
|
||||||
|
for local_name, buffer in module._buffers.items():
|
||||||
|
if buffer is None or local_name in module._non_persistent_buffers_set:
|
||||||
|
continue
|
||||||
|
name = f"{prefix}.{local_name}" if prefix else local_name
|
||||||
|
slots[name] = (module, local_name, False)
|
||||||
|
object_names.setdefault(id(buffer), []).append(name)
|
||||||
|
|
||||||
|
if set(state) != set(slots):
|
||||||
|
unsupported = sorted(set(state) ^ set(slots))
|
||||||
|
raise ComponentCheckpointUnsupportedError(
|
||||||
|
f"Direct GPU loading for {component_name!r} cannot assign custom "
|
||||||
|
f"state entries: {unsupported}"
|
||||||
|
)
|
||||||
|
aliases = [names for names in object_names.values() if len(names) > 1]
|
||||||
|
if aliases:
|
||||||
|
raise ComponentCheckpointUnsupportedError(
|
||||||
|
f"Direct GPU loading for {component_name!r} does not support tied "
|
||||||
|
f"state entries: {aliases}"
|
||||||
|
)
|
||||||
|
return state, slots
|
||||||
|
|
||||||
|
|
||||||
|
def _assign_direct_gpu_vae_state(
|
||||||
|
vae: nn.Module,
|
||||||
|
weight_files: list[str],
|
||||||
|
*,
|
||||||
|
component_name: str,
|
||||||
|
device: torch.device,
|
||||||
|
) -> None:
|
||||||
|
"""Stream a complete standard VAE state directly onto its target device."""
|
||||||
|
target_state, slots = _direct_gpu_vae_state_slots(vae, component_name)
|
||||||
|
loaded_names: set[str] = set()
|
||||||
|
with torch.no_grad():
|
||||||
|
for raw_name, tensor in safetensors_weights_iterator(
|
||||||
|
weight_files, to_cpu=device.type == "cpu"
|
||||||
|
):
|
||||||
|
name = raw_name
|
||||||
|
if name in loaded_names:
|
||||||
|
raise ComponentCheckpointUnsupportedError(
|
||||||
|
f"Direct GPU VAE checkpoint maps multiple tensors to {name!r}"
|
||||||
|
)
|
||||||
|
slot = slots.get(name)
|
||||||
|
if slot is None:
|
||||||
|
raise ComponentCheckpointUnsupportedError(
|
||||||
|
f"Direct GPU VAE checkpoint has unexpected tensor {raw_name!r}"
|
||||||
|
)
|
||||||
|
expected = target_state[name]
|
||||||
|
if tensor.shape != expected.shape:
|
||||||
|
raise ComponentCheckpointUnsupportedError(
|
||||||
|
f"Direct GPU VAE tensor {raw_name!r} has shape "
|
||||||
|
f"{tuple(tensor.shape)}, expected {tuple(expected.shape)}"
|
||||||
|
)
|
||||||
|
if tensor.dtype != expected.dtype:
|
||||||
|
tensor = tensor.to(dtype=expected.dtype)
|
||||||
|
|
||||||
|
module, local_name, is_parameter = slot
|
||||||
|
if is_parameter:
|
||||||
|
previous = module._parameters[local_name]
|
||||||
|
module._parameters[local_name] = nn.Parameter(
|
||||||
|
tensor, requires_grad=previous.requires_grad
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
module._buffers[local_name] = tensor
|
||||||
|
loaded_names.add(name)
|
||||||
|
|
||||||
|
missing = sorted(set(slots) - loaded_names)
|
||||||
|
if missing:
|
||||||
|
raise ComponentCheckpointUnsupportedError(
|
||||||
|
f"Direct GPU VAE checkpoint is missing tensors: {missing}"
|
||||||
|
)
|
||||||
|
remaining_meta = sorted(
|
||||||
|
name for name, tensor in vae.state_dict().items() if tensor.is_meta
|
||||||
|
)
|
||||||
|
if remaining_meta:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Direct GPU VAE loading left meta tensors: {remaining_meta}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class VAELoader(ComponentLoader):
|
class VAELoader(ComponentLoader):
|
||||||
"""Shared loader for (video/audio) VAE modules."""
|
"""Shared loader for (video/audio) VAE modules."""
|
||||||
|
|
||||||
component_names = ["vae", "audio_vae", "video_vae"]
|
component_names = ["vae", "audio_vae", "video_vae"]
|
||||||
expected_library = "diffusers"
|
expected_library = "diffusers"
|
||||||
|
supports_direct_gpu_weight_loading = True
|
||||||
|
|
||||||
|
def supports_direct_gpu_weight_loading_for_component(
|
||||||
|
self, component_name: str
|
||||||
|
) -> bool:
|
||||||
|
return component_name in ("vae", "video_vae")
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def resolve_model_weights_path(
|
def resolve_model_weights_path(
|
||||||
@@ -318,6 +440,13 @@ class VAELoader(ComponentLoader):
|
|||||||
cpu_offload_flag: bool = False,
|
cpu_offload_flag: bool = False,
|
||||||
):
|
):
|
||||||
"""Load the VAE based on the model path, and inference args."""
|
"""Load the VAE based on the model path, and inference args."""
|
||||||
|
direct_gpu_weight_loading = server_args.should_direct_gpu_weight_load_component(
|
||||||
|
component_name
|
||||||
|
)
|
||||||
|
if direct_gpu_weight_loading and component_name not in ("vae", "video_vae"):
|
||||||
|
raise ComponentCheckpointUnsupportedError(
|
||||||
|
f"Direct GPU loading is not implemented for {component_name!r}"
|
||||||
|
)
|
||||||
component_weights_path = self.resolve_model_weights_path(
|
component_weights_path = self.resolve_model_weights_path(
|
||||||
component_model_path,
|
component_model_path,
|
||||||
server_args,
|
server_args,
|
||||||
@@ -329,7 +458,10 @@ class VAELoader(ComponentLoader):
|
|||||||
server_args.pipeline_config, "native_only_components", ()
|
server_args.pipeline_config, "native_only_components", ()
|
||||||
)
|
)
|
||||||
_require_native_loader_for_quantized_vae(
|
_require_native_loader_for_quantized_vae(
|
||||||
config, component_name, native_only=native_only
|
config,
|
||||||
|
component_name,
|
||||||
|
native_only=native_only,
|
||||||
|
direct_gpu_weight_loading=direct_gpu_weight_loading,
|
||||||
)
|
)
|
||||||
|
|
||||||
class_name = config.pop("_class_name", None)
|
class_name = config.pop("_class_name", None)
|
||||||
@@ -368,6 +500,11 @@ class VAELoader(ComponentLoader):
|
|||||||
|
|
||||||
auto_map = config.get("auto_map", {})
|
auto_map = config.get("auto_map", {})
|
||||||
auto_model_map = auto_map.get("AutoModel")
|
auto_model_map = auto_map.get("AutoModel")
|
||||||
|
if direct_gpu_weight_loading and auto_model_map:
|
||||||
|
raise ComponentCheckpointUnsupportedError(
|
||||||
|
f"Direct GPU loading for {component_name!r} requires a native "
|
||||||
|
"ModelRegistry VAE; custom Diffusers auto_map code is unsupported"
|
||||||
|
)
|
||||||
if auto_model_map and component_weights_path != component_model_path:
|
if auto_model_map and component_weights_path != component_model_path:
|
||||||
raise ComponentCheckpointUnsupportedError(
|
raise ComponentCheckpointUnsupportedError(
|
||||||
f"{component_name!r} uses a custom Diffusers class that cannot "
|
f"{component_name!r} uses a custom Diffusers class that cannot "
|
||||||
@@ -400,6 +537,15 @@ class VAELoader(ComponentLoader):
|
|||||||
return vae
|
return vae
|
||||||
|
|
||||||
# Load from ModelRegistry (standard VAE classes)
|
# Load from ModelRegistry (standard VAE classes)
|
||||||
|
if direct_gpu_weight_loading:
|
||||||
|
with (
|
||||||
|
set_default_torch_dtype(vae_dtype),
|
||||||
|
skip_init_modules(),
|
||||||
|
torch.device("meta"),
|
||||||
|
):
|
||||||
|
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||||
|
vae = vae_cls(vae_config)
|
||||||
|
else:
|
||||||
with (
|
with (
|
||||||
set_default_torch_dtype(vae_dtype),
|
set_default_torch_dtype(vae_dtype),
|
||||||
skip_init_modules(),
|
skip_init_modules(),
|
||||||
@@ -426,6 +572,22 @@ class VAELoader(ComponentLoader):
|
|||||||
assert (
|
assert (
|
||||||
len(safetensors_list) >= 1
|
len(safetensors_list) >= 1
|
||||||
), f"Found no safetensors files in {component_weights_path}"
|
), f"Found no safetensors files in {component_weights_path}"
|
||||||
|
if direct_gpu_weight_loading:
|
||||||
|
_assign_direct_gpu_vae_state(
|
||||||
|
vae,
|
||||||
|
safetensors_list,
|
||||||
|
component_name=component_name,
|
||||||
|
device=target_device,
|
||||||
|
)
|
||||||
|
if _should_use_channels_last_3d(server_args, component_name):
|
||||||
|
n = _convert_conv3d_weights_to_channels_last_3d(vae)
|
||||||
|
if n > 0:
|
||||||
|
logger.info(
|
||||||
|
"VAE: converted %d Conv3d weights to channels_last_3d", n
|
||||||
|
)
|
||||||
|
_hold_decoder_weights_in_decode_dtype(vae, server_args, component_name)
|
||||||
|
return current_platform.optimize_vae(vae)
|
||||||
|
|
||||||
loaded = {}
|
loaded = {}
|
||||||
for sf_path in safetensors_list:
|
for sf_path in safetensors_list:
|
||||||
loaded.update(safetensors_load_file(sf_path))
|
loaded.update(safetensors_load_file(sf_path))
|
||||||
|
|||||||
@@ -366,6 +366,22 @@ class ComposedPipelineBase(ABC):
|
|||||||
logger.debug("Resolved component path: %s", component_model_path)
|
logger.debug("Resolved component path: %s", component_model_path)
|
||||||
return component_model_path
|
return component_model_path
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _validate_direct_gpu_component_selection(
|
||||||
|
model_index: dict[str, Any], server_args: ServerArgs
|
||||||
|
) -> None:
|
||||||
|
unavailable = sorted(
|
||||||
|
component_name
|
||||||
|
for component_name in server_args.component_direct_gpu_weight_loading
|
||||||
|
if component_name not in model_index or model_index[component_name] is None
|
||||||
|
)
|
||||||
|
if unavailable:
|
||||||
|
raise ValueError(
|
||||||
|
"--component-direct-gpu-weight-loading selects component(s) "
|
||||||
|
"that are not available in this pipeline: "
|
||||||
|
f"{', '.join(unavailable)}"
|
||||||
|
)
|
||||||
|
|
||||||
def load_modules(
|
def load_modules(
|
||||||
self,
|
self,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
@@ -432,6 +448,7 @@ class ComposedPipelineBase(ABC):
|
|||||||
model_index.pop("boundary_ratio", None)
|
model_index.pop("boundary_ratio", None)
|
||||||
# used by Wan2.2 ti2v
|
# used by Wan2.2 ti2v
|
||||||
model_index.pop("expand_timesteps", None)
|
model_index.pop("expand_timesteps", None)
|
||||||
|
self._validate_direct_gpu_component_selection(model_index, server_args)
|
||||||
|
|
||||||
# some sanity checks
|
# some sanity checks
|
||||||
assert (
|
assert (
|
||||||
|
|||||||
@@ -306,6 +306,9 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
component_paths: dict[str, str] = field(default_factory=dict)
|
component_paths: dict[str, str] = field(default_factory=dict)
|
||||||
# Exact weight-file overrides retain the base component configuration.
|
# Exact weight-file overrides retain the base component configuration.
|
||||||
component_weights_paths: dict[str, str] = field(default_factory=dict)
|
component_weights_paths: dict[str, str] = field(default_factory=dict)
|
||||||
|
# Opt in one component to a loader-specific direct-GPU weight path. The
|
||||||
|
# existing --direct-gpu-weight-loading remains the primary-DiT control.
|
||||||
|
component_direct_gpu_weight_loading: dict[str, bool] = field(default_factory=dict)
|
||||||
# Explicit quantization override for one component. Self-describing
|
# Explicit quantization override for one component. Self-describing
|
||||||
# checkpoints remain auto-detected and do not need this override.
|
# checkpoints remain auto-detected and do not need this override.
|
||||||
component_quantizations: dict[str, str] = field(default_factory=dict)
|
component_quantizations: dict[str, str] = field(default_factory=dict)
|
||||||
@@ -1752,6 +1755,16 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
component_weights_paths[component] = path
|
component_weights_paths[component] = path
|
||||||
self.component_paths = component_paths
|
self.component_paths = component_paths
|
||||||
self.component_weights_paths = component_weights_paths
|
self.component_weights_paths = component_weights_paths
|
||||||
|
normalized_direct_gpu_loading: dict[str, bool] = {}
|
||||||
|
for component, enabled in self.component_direct_gpu_weight_loading.items():
|
||||||
|
component_name = str(component).strip().replace("-", "_")
|
||||||
|
if not component_name or not isinstance(enabled, bool):
|
||||||
|
raise ValueError(
|
||||||
|
"Component direct GPU loading entries require a component and "
|
||||||
|
"a boolean value"
|
||||||
|
)
|
||||||
|
normalized_direct_gpu_loading[component_name] = enabled
|
||||||
|
self.component_direct_gpu_weight_loading = normalized_direct_gpu_loading
|
||||||
normalized_quantizations: dict[str, str] = {}
|
normalized_quantizations: dict[str, str] = {}
|
||||||
for component, quantization in self.component_quantizations.items():
|
for component, quantization in self.component_quantizations.items():
|
||||||
component = str(component).strip().replace("-", "_")
|
component = str(component).strip().replace("-", "_")
|
||||||
@@ -2914,6 +2927,52 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
alias_suffix="-weights-path",
|
alias_suffix="-weights-path",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _extract_component_direct_gpu_weight_loading(
|
||||||
|
unknown_args: list[str],
|
||||||
|
) -> tuple[dict[str, bool], list[str]]:
|
||||||
|
"""Extract exact component direct-GPU loading toggles.
|
||||||
|
|
||||||
|
Dynamic boolean flags mirror ``StoreBoolean``: an omitted value means
|
||||||
|
true, while an explicit ``true`` or ``false`` works with either an
|
||||||
|
equals form or a following argument.
|
||||||
|
"""
|
||||||
|
values: dict[str, bool] = {}
|
||||||
|
remaining: list[str] = []
|
||||||
|
prefixes = (
|
||||||
|
"--component-direct-gpu-weight-loading.",
|
||||||
|
"--component_direct_gpu_weight_loading.",
|
||||||
|
)
|
||||||
|
i = 0
|
||||||
|
while i < len(unknown_args):
|
||||||
|
arg = unknown_args[i]
|
||||||
|
key_part = arg.split("=", 1)[0] if "=" in arg else arg
|
||||||
|
prefix = next(
|
||||||
|
(candidate for candidate in prefixes if key_part.startswith(candidate)),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
if prefix is None:
|
||||||
|
remaining.append(arg)
|
||||||
|
i += 1
|
||||||
|
continue
|
||||||
|
|
||||||
|
component = key_part[len(prefix) :].replace("-", "_")
|
||||||
|
value = "true"
|
||||||
|
if "=" in arg:
|
||||||
|
value = arg.split("=", 1)[1]
|
||||||
|
elif i + 1 < len(unknown_args):
|
||||||
|
next_value = unknown_args[i + 1].lower()
|
||||||
|
if next_value in ("true", "false"):
|
||||||
|
i += 1
|
||||||
|
value = next_value
|
||||||
|
|
||||||
|
if not component or value.lower() not in ("true", "false"):
|
||||||
|
remaining.append(arg)
|
||||||
|
else:
|
||||||
|
values[component] = value.lower() == "true"
|
||||||
|
i += 1
|
||||||
|
return values, remaining
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _extract_component_quantizations(
|
def _extract_component_quantizations(
|
||||||
cls,
|
cls,
|
||||||
@@ -3019,6 +3078,9 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
dynamic_quantizations, remaining = cls._extract_component_quantizations(
|
dynamic_quantizations, remaining = cls._extract_component_quantizations(
|
||||||
unknown_args
|
unknown_args
|
||||||
)
|
)
|
||||||
|
dynamic_direct_gpu_loading, remaining = (
|
||||||
|
cls._extract_component_direct_gpu_weight_loading(remaining)
|
||||||
|
)
|
||||||
dynamic_ignored_layers, remaining = (
|
dynamic_ignored_layers, remaining = (
|
||||||
cls._extract_component_quantization_ignored_layers(remaining)
|
cls._extract_component_quantization_ignored_layers(remaining)
|
||||||
)
|
)
|
||||||
@@ -3062,6 +3124,13 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
existing.update(dynamic_quantizations)
|
existing.update(dynamic_quantizations)
|
||||||
provided_args["component_quantizations"] = existing
|
provided_args["component_quantizations"] = existing
|
||||||
explicit_arg_names.add("component_quantizations")
|
explicit_arg_names.add("component_quantizations")
|
||||||
|
if dynamic_direct_gpu_loading:
|
||||||
|
existing = dict(
|
||||||
|
provided_args.get("component_direct_gpu_weight_loading") or {}
|
||||||
|
)
|
||||||
|
existing.update(dynamic_direct_gpu_loading)
|
||||||
|
provided_args["component_direct_gpu_weight_loading"] = existing
|
||||||
|
explicit_arg_names.add("component_direct_gpu_weight_loading")
|
||||||
if dynamic_ignored_layers:
|
if dynamic_ignored_layers:
|
||||||
existing = dict(
|
existing = dict(
|
||||||
provided_args.get("component_quantization_ignored_layers") or {}
|
provided_args.get("component_quantization_ignored_layers") or {}
|
||||||
@@ -3371,8 +3440,7 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def _validate_direct_gpu_weight_loading(self) -> None:
|
def _validate_direct_gpu_weight_loading(self) -> None:
|
||||||
if not self.direct_gpu_weight_loading:
|
if self.direct_gpu_weight_loading:
|
||||||
return
|
|
||||||
if not current_platform.is_cuda():
|
if not current_platform.is_cuda():
|
||||||
raise ValueError("--direct-gpu-weight-loading requires CUDA")
|
raise ValueError("--direct-gpu-weight-loading requires CUDA")
|
||||||
if (
|
if (
|
||||||
@@ -3380,8 +3448,8 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
or self.residency_mode("transformer") == LAYERWISE_OFFLOAD
|
or self.residency_mode("transformer") == LAYERWISE_OFFLOAD
|
||||||
):
|
):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"--direct-gpu-weight-loading requires a GPU-resident DiT; disable "
|
"--direct-gpu-weight-loading requires a GPU-resident DiT; "
|
||||||
"DiT CPU and layerwise offload"
|
"disable DiT CPU and layerwise offload"
|
||||||
)
|
)
|
||||||
if self.use_fsdp_inference:
|
if self.use_fsdp_inference:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -3390,6 +3458,21 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
if self.tp_size != 1:
|
if self.tp_size != 1:
|
||||||
raise ValueError("--direct-gpu-weight-loading requires --tp-size 1")
|
raise ValueError("--direct-gpu-weight-loading requires --tp-size 1")
|
||||||
|
|
||||||
|
for component_name, enabled in self.component_direct_gpu_weight_loading.items():
|
||||||
|
if not enabled:
|
||||||
|
continue
|
||||||
|
if not current_platform.is_cuda():
|
||||||
|
raise ValueError("--component-direct-gpu-weight-loading requires CUDA")
|
||||||
|
if self.should_start_component_on_cpu(component_name):
|
||||||
|
raise ValueError(
|
||||||
|
"--component-direct-gpu-weight-loading requires "
|
||||||
|
f"{component_name!r} to be resident"
|
||||||
|
)
|
||||||
|
|
||||||
|
def should_direct_gpu_weight_load_component(self, component_name: str) -> bool:
|
||||||
|
"""Return whether an exact component opted into direct GPU loading."""
|
||||||
|
return self.component_direct_gpu_weight_loading.get(component_name, False)
|
||||||
|
|
||||||
def _validate_parallelism(self):
|
def _validate_parallelism(self):
|
||||||
if self.kv_gather_degree < 1:
|
if self.kv_gather_degree < 1:
|
||||||
raise ValueError("kv_gather_degree must be >= 1")
|
raise ValueError("kv_gather_degree must be >= 1")
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ from sglang.multimodal_gen.runtime.loader.component_loaders.bridge_loader import
|
|||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
||||||
ComponentCheckpointUnsupportedError,
|
ComponentCheckpointUnsupportedError,
|
||||||
|
ComponentLoader,
|
||||||
PlainStateDictComponentLoader,
|
PlainStateDictComponentLoader,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.component_loaders.diffusion_decoder_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.diffusion_decoder_loader import (
|
||||||
@@ -24,9 +25,13 @@ from sglang.multimodal_gen.runtime.loader.component_loaders.sound_tokenizer_load
|
|||||||
from sglang.multimodal_gen.runtime.loader.component_loaders.upsampler_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.upsampler_loader import (
|
||||||
UpsamplerLoader,
|
UpsamplerLoader,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.loader.component_loaders.vae_loader import VAELoader
|
||||||
from sglang.multimodal_gen.runtime.loader.component_loaders.vocoder_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.vocoder_loader import (
|
||||||
VocoderLoader,
|
VocoderLoader,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||||
|
ComposedPipelineBase,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class _TestLoader(PlainStateDictComponentLoader):
|
class _TestLoader(PlainStateDictComponentLoader):
|
||||||
@@ -34,6 +39,42 @@ class _TestLoader(PlainStateDictComponentLoader):
|
|||||||
|
|
||||||
|
|
||||||
class TestComponentQuantizationAdmission(unittest.TestCase):
|
class TestComponentQuantizationAdmission(unittest.TestCase):
|
||||||
|
def test_direct_gpu_selection_requires_a_declared_component(self):
|
||||||
|
server_args = SimpleNamespace(
|
||||||
|
component_direct_gpu_weight_loading={"missing_vae": True}
|
||||||
|
)
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(ValueError, "missing_vae"):
|
||||||
|
ComposedPipelineBase._validate_direct_gpu_component_selection(
|
||||||
|
{"vae": ["diffusers", "AutoencoderKL"]}, server_args
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_direct_gpu_selector_is_rejected_by_unqualified_loader(self):
|
||||||
|
server_args = SimpleNamespace(
|
||||||
|
component_quantizations={},
|
||||||
|
should_direct_gpu_weight_load_component=lambda component: component
|
||||||
|
== "vocoder",
|
||||||
|
)
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(
|
||||||
|
ComponentCheckpointUnsupportedError, "does not support direct GPU"
|
||||||
|
):
|
||||||
|
ComponentLoader().load(
|
||||||
|
"/model/vocoder", server_args, "vocoder", "diffusers"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_direct_gpu_selector_is_rejected_by_unqualified_component(self):
|
||||||
|
server_args = SimpleNamespace(
|
||||||
|
component_quantizations={},
|
||||||
|
should_direct_gpu_weight_load_component=lambda component: component
|
||||||
|
== "audio_vae",
|
||||||
|
)
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(
|
||||||
|
ComponentCheckpointUnsupportedError, "does not support direct GPU"
|
||||||
|
):
|
||||||
|
VAELoader().load("/model/audio_vae", server_args, "audio_vae", "diffusers")
|
||||||
|
|
||||||
def test_plain_loader_resolves_weights_separately_from_config(self):
|
def test_plain_loader_resolves_weights_separately_from_config(self):
|
||||||
server_args = SimpleNamespace(
|
server_args = SimpleNamespace(
|
||||||
component_weights_paths={"vocoder": "owner/repo/vocoder.safetensors"}
|
component_weights_paths={"vocoder": "owner/repo/vocoder.safetensors"}
|
||||||
|
|||||||
@@ -41,6 +41,7 @@ class TestImageEncoderQuantizationAdmission(unittest.TestCase):
|
|||||||
component_quantizations={},
|
component_quantizations={},
|
||||||
encoder_parallel="replicate",
|
encoder_parallel="replicate",
|
||||||
resolve_component_attention_backend=lambda _name: (None, None),
|
resolve_component_attention_backend=lambda _name: (None, None),
|
||||||
|
should_direct_gpu_weight_load_component=lambda _name: False,
|
||||||
should_use_fsdp_for_component=lambda _name: False,
|
should_use_fsdp_for_component=lambda _name: False,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -236,6 +237,7 @@ class TestImageEncoderNativeLoading(unittest.TestCase):
|
|||||||
require_component_resident=mock.Mock(),
|
require_component_resident=mock.Mock(),
|
||||||
should_use_fsdp_for_component=lambda _name: False,
|
should_use_fsdp_for_component=lambda _name: False,
|
||||||
should_start_component_on_cpu=lambda _name: False,
|
should_start_component_on_cpu=lambda _name: False,
|
||||||
|
should_direct_gpu_weight_load_component=lambda _name: False,
|
||||||
revision=None,
|
revision=None,
|
||||||
trust_remote_code=False,
|
trust_remote_code=False,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -3162,6 +3162,7 @@ class TestDirectGpuWeightLoading(unittest.TestCase):
|
|||||||
def _args(self) -> ServerArgs:
|
def _args(self) -> ServerArgs:
|
||||||
args = ServerArgs.__new__(ServerArgs)
|
args = ServerArgs.__new__(ServerArgs)
|
||||||
args.direct_gpu_weight_loading = True
|
args.direct_gpu_weight_loading = True
|
||||||
|
args.component_direct_gpu_weight_loading = {}
|
||||||
args.component_residency = None
|
args.component_residency = None
|
||||||
args.cpu_offload_components = None
|
args.cpu_offload_components = None
|
||||||
args.dit_cpu_offload = False
|
args.dit_cpu_offload = False
|
||||||
@@ -3189,6 +3190,28 @@ class TestDirectGpuWeightLoading(unittest.TestCase):
|
|||||||
self.assertFalse(default_args.direct_gpu_weight_loading)
|
self.assertFalse(default_args.direct_gpu_weight_loading)
|
||||||
self.assertTrue(enabled_args.direct_gpu_weight_loading)
|
self.assertTrue(enabled_args.direct_gpu_weight_loading)
|
||||||
|
|
||||||
|
def test_component_direct_gpu_parser_is_exact_and_boolean(self):
|
||||||
|
values, remaining = ServerArgs._extract_component_direct_gpu_weight_loading(
|
||||||
|
[
|
||||||
|
"--component-direct-gpu-weight-loading.video-vae",
|
||||||
|
"--component-direct-gpu-weight-loading.audio_vae=false",
|
||||||
|
"--other-flag",
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual({"video_vae": True, "audio_vae": False}, values)
|
||||||
|
self.assertEqual(["--other-flag"], remaining)
|
||||||
|
|
||||||
|
def test_component_direct_gpu_rejects_nonresident_component(self):
|
||||||
|
args = self._args()
|
||||||
|
args.direct_gpu_weight_loading = False
|
||||||
|
args.component_direct_gpu_weight_loading = {"video_vae": True}
|
||||||
|
args.vae_cpu_offload = True
|
||||||
|
|
||||||
|
with patch.object(current_platform, "is_cuda", return_value=True):
|
||||||
|
with self.assertRaisesRegex(ValueError, "'video_vae' to be resident"):
|
||||||
|
args._validate_direct_gpu_weight_loading()
|
||||||
|
|
||||||
def test_rejects_cpu_offload_fsdp_and_tp(self):
|
def test_rejects_cpu_offload_fsdp_and_tp(self):
|
||||||
cpu_offload_args = self._args()
|
cpu_offload_args = self._args()
|
||||||
cpu_offload_args.dit_cpu_offload = True
|
cpu_offload_args.dit_cpu_offload = True
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ from unittest.mock import patch
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
from safetensors.torch import save_file as safetensors_save_file
|
||||||
|
|
||||||
from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import LTX2PipelineConfig
|
from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import LTX2PipelineConfig
|
||||||
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
|
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
|
||||||
@@ -20,7 +21,9 @@ from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader imp
|
|||||||
ComponentCheckpointUnsupportedError,
|
ComponentCheckpointUnsupportedError,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.component_loaders.vae_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.vae_loader import (
|
||||||
|
_assign_direct_gpu_vae_state,
|
||||||
_backfill_ltx2_audio_vae_latent_stats,
|
_backfill_ltx2_audio_vae_latent_stats,
|
||||||
|
_direct_gpu_vae_state_slots,
|
||||||
_match_checkpoint_dtypes,
|
_match_checkpoint_dtypes,
|
||||||
_require_native_loader_for_quantized_vae,
|
_require_native_loader_for_quantized_vae,
|
||||||
_should_use_channels_last_3d,
|
_should_use_channels_last_3d,
|
||||||
@@ -44,6 +47,7 @@ class _FakeServerArgs:
|
|||||||
self.trust_remote_code = True
|
self.trust_remote_code = True
|
||||||
self.layerwise_components = set()
|
self.layerwise_components = set()
|
||||||
self.component_quantizations = {}
|
self.component_quantizations = {}
|
||||||
|
self.component_direct_gpu_weight_loading = {}
|
||||||
|
|
||||||
def resolve_component_attention_backend(self, _component_name):
|
def resolve_component_attention_backend(self, _component_name):
|
||||||
return None, None
|
return None, None
|
||||||
@@ -54,6 +58,9 @@ class _FakeServerArgs:
|
|||||||
def should_configure_layerwise_offload_for_lazy_component(self, component_name):
|
def should_configure_layerwise_offload_for_lazy_component(self, component_name):
|
||||||
return component_name in self.layerwise_components
|
return component_name in self.layerwise_components
|
||||||
|
|
||||||
|
def should_direct_gpu_weight_load_component(self, component_name):
|
||||||
|
return self.component_direct_gpu_weight_loading.get(component_name, False)
|
||||||
|
|
||||||
def should_use_fsdp_for_component(self, _component_name):
|
def should_use_fsdp_for_component(self, _component_name):
|
||||||
return False
|
return False
|
||||||
|
|
||||||
@@ -126,6 +133,131 @@ class TestMatchCheckpointDtypes(unittest.TestCase):
|
|||||||
self.assertIs(loaded["extra"], before)
|
self.assertIs(loaded["extra"], before)
|
||||||
|
|
||||||
|
|
||||||
|
class TestDirectGPUVAEState(unittest.TestCase):
|
||||||
|
class _StandardVAE(nn.Module):
|
||||||
|
def __init__(self, *_args, **_kwargs):
|
||||||
|
super().__init__()
|
||||||
|
self.proj = nn.Linear(2, 2, bias=False)
|
||||||
|
self.register_buffer("scale", torch.ones(1))
|
||||||
|
|
||||||
|
def test_assigns_one_complete_state_without_a_cpu_state_dict(self):
|
||||||
|
expected_weight = torch.arange(4, dtype=torch.float32).reshape(2, 2)
|
||||||
|
expected_scale = torch.tensor([3.0])
|
||||||
|
with TemporaryDirectory() as root:
|
||||||
|
checkpoint = pathlib.Path(root) / "model.safetensors"
|
||||||
|
safetensors_save_file(
|
||||||
|
{"proj.weight": expected_weight, "scale": expected_scale}, checkpoint
|
||||||
|
)
|
||||||
|
with torch.device("meta"):
|
||||||
|
vae = self._StandardVAE()
|
||||||
|
_assign_direct_gpu_vae_state(
|
||||||
|
vae,
|
||||||
|
[str(checkpoint)],
|
||||||
|
component_name="vae",
|
||||||
|
device=torch.device("cpu"),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertTrue(torch.equal(vae.proj.weight, expected_weight))
|
||||||
|
self.assertTrue(torch.equal(vae.scale, expected_scale))
|
||||||
|
self.assertFalse(any(tensor.is_meta for tensor in vae.state_dict().values()))
|
||||||
|
|
||||||
|
def test_rejects_nonstandard_state_lifecycle(self):
|
||||||
|
class _CustomVAE(self._StandardVAE):
|
||||||
|
def state_dict(self, *args, **kwargs):
|
||||||
|
return super().state_dict(*args, **kwargs)
|
||||||
|
|
||||||
|
with torch.device("meta"):
|
||||||
|
vae = _CustomVAE()
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(ComponentCheckpointUnsupportedError, "ABI"):
|
||||||
|
_direct_gpu_vae_state_slots(vae, "vae")
|
||||||
|
|
||||||
|
def test_ignores_nonpersistent_runtime_buffers(self):
|
||||||
|
class _VAEWithRuntimeBuffer(self._StandardVAE):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
self.register_buffer("runtime_cache", torch.zeros(1), persistent=False)
|
||||||
|
|
||||||
|
with torch.device("meta"):
|
||||||
|
vae = _VAEWithRuntimeBuffer()
|
||||||
|
|
||||||
|
state, slots = _direct_gpu_vae_state_slots(vae, "vae")
|
||||||
|
|
||||||
|
self.assertNotIn("runtime_cache", state)
|
||||||
|
self.assertNotIn("runtime_cache", slots)
|
||||||
|
|
||||||
|
def test_loader_streams_native_vae_without_the_legacy_cpu_state_dict(self):
|
||||||
|
loader = vae_loader.VAELoader()
|
||||||
|
server_args = _FakeServerArgs(QwenImagePipelineConfig())
|
||||||
|
server_args.component_weights_paths = {}
|
||||||
|
server_args.component_direct_gpu_weight_loading = {"vae": True}
|
||||||
|
expected_weight = torch.arange(4, dtype=torch.bfloat16).reshape(2, 2)
|
||||||
|
expected_scale = torch.tensor([3.0], dtype=torch.bfloat16)
|
||||||
|
|
||||||
|
with (
|
||||||
|
TemporaryDirectory() as root,
|
||||||
|
patch.object(
|
||||||
|
vae_loader,
|
||||||
|
"get_diffusers_component_config",
|
||||||
|
return_value={"_class_name": "TestVAE"},
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
vae_loader.ModelRegistry,
|
||||||
|
"resolve_model_cls",
|
||||||
|
return_value=(self._StandardVAE, None),
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
vae_loader,
|
||||||
|
"_list_safetensors_files",
|
||||||
|
return_value=[str(pathlib.Path(root) / "model.safetensors")],
|
||||||
|
),
|
||||||
|
patch.object(loader, "target_device", return_value=torch.device("cpu")),
|
||||||
|
patch.object(
|
||||||
|
vae_loader.current_platform,
|
||||||
|
"optimize_vae",
|
||||||
|
side_effect=lambda vae: vae,
|
||||||
|
),
|
||||||
|
patch.object(vae_loader, "safetensors_load_file") as legacy_load,
|
||||||
|
):
|
||||||
|
safetensors_save_file(
|
||||||
|
{"proj.weight": expected_weight, "scale": expected_scale},
|
||||||
|
pathlib.Path(root) / "model.safetensors",
|
||||||
|
)
|
||||||
|
loaded = loader.load_customized(root, server_args, "vae")
|
||||||
|
|
||||||
|
legacy_load.assert_not_called()
|
||||||
|
self.assertTrue(torch.equal(loaded.proj.weight, expected_weight))
|
||||||
|
self.assertTrue(torch.equal(loaded.scale, expected_scale))
|
||||||
|
|
||||||
|
def test_quantized_checkpoint_does_not_fall_back_from_direct_loading(self):
|
||||||
|
loader = vae_loader.VAELoader()
|
||||||
|
server_args = _FakeServerArgs(QwenImagePipelineConfig())
|
||||||
|
server_args.component_direct_gpu_weight_loading = {"vae": True}
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(
|
||||||
|
vae_loader,
|
||||||
|
"get_diffusers_component_config",
|
||||||
|
return_value={
|
||||||
|
"_class_name": "AutoencoderKL",
|
||||||
|
"quantization_config": {"quant_method": "bitsandbytes"},
|
||||||
|
},
|
||||||
|
),
|
||||||
|
patch("diffusers.AutoModel.from_pretrained") as native_load,
|
||||||
|
patch.object(
|
||||||
|
vae_loader.current_platform,
|
||||||
|
"get_available_gpu_memory",
|
||||||
|
return_value=10.0,
|
||||||
|
),
|
||||||
|
self.assertRaisesRegex(
|
||||||
|
ComponentCheckpointUnsupportedError, "Direct GPU loading"
|
||||||
|
),
|
||||||
|
):
|
||||||
|
loader.load("/quantized/vae", server_args, "vae", "diffusers")
|
||||||
|
|
||||||
|
native_load.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
class TestVAELoader(unittest.TestCase):
|
class TestVAELoader(unittest.TestCase):
|
||||||
def test_weights_override_keeps_base_component_config(self):
|
def test_weights_override_keeps_base_component_config(self):
|
||||||
loader = vae_loader.VAELoader()
|
loader = vae_loader.VAELoader()
|
||||||
|
|||||||
Reference in New Issue
Block a user