[diffusion] feat: compose third-party component bundles safely (#37816)
This commit is contained in:
@@ -27,6 +27,7 @@ from sglang.multimodal_gen.runtime.layers.linear import (
|
|||||||
QKVParallelLinear,
|
QKVParallelLinear,
|
||||||
ReplicatedLinear,
|
ReplicatedLinear,
|
||||||
RowParallelLinear,
|
RowParallelLinear,
|
||||||
|
UnquantizedLinearMethod,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.vocab_parallel_embedding import (
|
from sglang.multimodal_gen.runtime.layers.vocab_parallel_embedding import (
|
||||||
VocabParallelEmbedding,
|
VocabParallelEmbedding,
|
||||||
@@ -116,6 +117,16 @@ class BaseLayerWithLoRA(nn.Module):
|
|||||||
def bias(self):
|
def bias(self):
|
||||||
return getattr(self.base_layer, "bias", None)
|
return getattr(self.base_layer, "bias", None)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def can_merge_base_weight(self) -> bool:
|
||||||
|
"""Whether a LoRA delta may safely replace the stored base weight."""
|
||||||
|
weight = self.weight
|
||||||
|
if not (weight.dtype.is_floating_point or weight.dtype.is_complex):
|
||||||
|
return False
|
||||||
|
if isinstance(self.base_layer, LinearBase):
|
||||||
|
return isinstance(self.base_layer.quant_method, UnquantizedLinearMethod)
|
||||||
|
return True
|
||||||
|
|
||||||
@torch.compile()
|
@torch.compile()
|
||||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
lora_A = self.lora_A
|
lora_A = self.lora_A
|
||||||
|
|||||||
@@ -1,13 +1,17 @@
|
|||||||
import hashlib
|
import hashlib
|
||||||
import importlib.util
|
import importlib.util
|
||||||
import os
|
import os
|
||||||
|
from collections.abc import Iterable
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from safetensors.torch import load_file as safetensors_load_file
|
from safetensors.torch import load_file as safetensors_load_file
|
||||||
|
from safetensors.torch import safe_open
|
||||||
from safetensors.torch import save_file as safetensors_save_file
|
from safetensors.torch import save_file as safetensors_save_file
|
||||||
|
from torch.nn.utils import parametrize
|
||||||
|
|
||||||
from sglang.multimodal_gen import envs
|
from sglang.multimodal_gen import envs
|
||||||
|
from sglang.multimodal_gen.configs.models.vaes.base import VAEConfig
|
||||||
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 (
|
||||||
QwenImagePipelineConfig,
|
QwenImagePipelineConfig,
|
||||||
@@ -52,6 +56,7 @@ from sglang.srt.model_loader.checkpoint_quantization import (
|
|||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
VAE_CHANNELS_LAST_3D_ENV = "SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D"
|
VAE_CHANNELS_LAST_3D_ENV = "SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D"
|
||||||
|
_VAE_CHECKPOINT_ARCH_METADATA = ("latents_mean", "latents_std")
|
||||||
|
|
||||||
|
|
||||||
def _require_native_loader_for_quantized_vae(
|
def _require_native_loader_for_quantized_vae(
|
||||||
@@ -293,6 +298,101 @@ def _match_checkpoint_dtypes(loaded: dict, target_state: dict) -> dict:
|
|||||||
return loaded
|
return loaded
|
||||||
|
|
||||||
|
|
||||||
|
def _adopt_plain_weight_norm_state(
|
||||||
|
module: nn.Module, loaded_names: Iterable[str]
|
||||||
|
) -> int:
|
||||||
|
"""Make deparameterized checkpoint weights native module state.
|
||||||
|
|
||||||
|
PyTorch's weight-norm load hook accepts legacy ``weight_g``/``weight_v``
|
||||||
|
tensors, while inference exports commonly fold those tensors into one
|
||||||
|
plain ``weight``. Removing only the matching parametrizations preserves
|
||||||
|
that already-computed weight exactly and leaves every other parameterized
|
||||||
|
module untouched.
|
||||||
|
"""
|
||||||
|
state_names = set(module.state_dict())
|
||||||
|
module_by_name = dict(module.named_modules())
|
||||||
|
owners: set[str] = set()
|
||||||
|
for name in loaded_names:
|
||||||
|
if name == "weight":
|
||||||
|
owner_name = ""
|
||||||
|
elif name.endswith(".weight"):
|
||||||
|
owner_name = name.removesuffix(".weight")
|
||||||
|
else:
|
||||||
|
continue
|
||||||
|
state_prefix = f"{owner_name}." if owner_name else ""
|
||||||
|
if {
|
||||||
|
f"{state_prefix}parametrizations.weight.original0",
|
||||||
|
f"{state_prefix}parametrizations.weight.original1",
|
||||||
|
}.issubset(state_names):
|
||||||
|
owners.add(owner_name)
|
||||||
|
|
||||||
|
for owner_name in sorted(owners):
|
||||||
|
parametrize.remove_parametrizations(
|
||||||
|
module_by_name[owner_name], "weight", leave_parametrized=True
|
||||||
|
)
|
||||||
|
return len(owners)
|
||||||
|
|
||||||
|
|
||||||
|
def _vae_checkpoint_arch_metadata_names(
|
||||||
|
vae_config: VAEConfig,
|
||||||
|
target_state: dict[str, torch.Tensor],
|
||||||
|
) -> tuple[str, ...]:
|
||||||
|
arch_values = vars(vae_config.arch_config)
|
||||||
|
return tuple(
|
||||||
|
name
|
||||||
|
for name in _VAE_CHECKPOINT_ARCH_METADATA
|
||||||
|
if name not in target_state and name in arch_values
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _consume_vae_checkpoint_arch_metadata(
|
||||||
|
loaded: dict[str, torch.Tensor],
|
||||||
|
vae_config: VAEConfig,
|
||||||
|
target_state: dict[str, torch.Tensor],
|
||||||
|
) -> tuple[str, ...]:
|
||||||
|
"""Move checkpoint-carried latent statistics into the VAE config."""
|
||||||
|
arch_values = vars(vae_config.arch_config)
|
||||||
|
consumed = []
|
||||||
|
for name in _vae_checkpoint_arch_metadata_names(vae_config, target_state):
|
||||||
|
tensor = loaded.get(name)
|
||||||
|
if tensor is None:
|
||||||
|
continue
|
||||||
|
if tensor.ndim != 1:
|
||||||
|
raise ValueError(
|
||||||
|
f"VAE checkpoint metadata {name!r} must be one-dimensional, "
|
||||||
|
f"got shape {tuple(tensor.shape)}"
|
||||||
|
)
|
||||||
|
arch_values[name] = tensor.tolist()
|
||||||
|
del loaded[name]
|
||||||
|
consumed.append(name)
|
||||||
|
if consumed:
|
||||||
|
vae_config.post_init()
|
||||||
|
return tuple(consumed)
|
||||||
|
|
||||||
|
|
||||||
|
def _vae_checkpoint_tensor_names(weight_files: list[str]) -> set[str]:
|
||||||
|
names: set[str] = set()
|
||||||
|
for path in weight_files:
|
||||||
|
with safe_open(path, framework="pt", device="cpu") as checkpoint:
|
||||||
|
names.update(checkpoint.keys())
|
||||||
|
return names
|
||||||
|
|
||||||
|
|
||||||
|
def _log_vae_checkpoint_adaptations(
|
||||||
|
num_deparameterized: int, consumed_metadata: tuple[str, ...]
|
||||||
|
) -> None:
|
||||||
|
if num_deparameterized:
|
||||||
|
logger.info(
|
||||||
|
"VAE: adopted %d deparameterized weight-normalized layers",
|
||||||
|
num_deparameterized,
|
||||||
|
)
|
||||||
|
if consumed_metadata:
|
||||||
|
logger.info(
|
||||||
|
"VAE: loaded architecture metadata from checkpoint: %s",
|
||||||
|
", ".join(consumed_metadata),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _direct_gpu_vae_state_slots(
|
def _direct_gpu_vae_state_slots(
|
||||||
vae: nn.Module, component_name: str
|
vae: nn.Module, component_name: str
|
||||||
) -> tuple[dict[str, torch.Tensor], dict[str, tuple[nn.Module, str, bool]]]:
|
) -> tuple[dict[str, torch.Tensor], dict[str, tuple[nn.Module, str, bool]]]:
|
||||||
@@ -341,15 +441,24 @@ def _assign_direct_gpu_vae_state(
|
|||||||
*,
|
*,
|
||||||
component_name: str,
|
component_name: str,
|
||||||
device: torch.device,
|
device: torch.device,
|
||||||
) -> None:
|
vae_config: VAEConfig,
|
||||||
|
) -> tuple[int, tuple[str, ...]]:
|
||||||
"""Stream a complete standard VAE state directly onto its target device."""
|
"""Stream a complete standard VAE state directly onto its target device."""
|
||||||
|
num_deparameterized = _adopt_plain_weight_norm_state(
|
||||||
|
vae, _vae_checkpoint_tensor_names(weight_files)
|
||||||
|
)
|
||||||
target_state, slots = _direct_gpu_vae_state_slots(vae, component_name)
|
target_state, slots = _direct_gpu_vae_state_slots(vae, component_name)
|
||||||
|
metadata_names = _vae_checkpoint_arch_metadata_names(vae_config, target_state)
|
||||||
loaded_names: set[str] = set()
|
loaded_names: set[str] = set()
|
||||||
|
metadata: dict[str, torch.Tensor] = {}
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
for raw_name, tensor in safetensors_weights_iterator(
|
for raw_name, tensor in safetensors_weights_iterator(
|
||||||
weight_files, to_cpu=device.type == "cpu"
|
weight_files, to_cpu=device.type == "cpu"
|
||||||
):
|
):
|
||||||
name = raw_name
|
name = raw_name
|
||||||
|
if name in metadata_names:
|
||||||
|
metadata[name] = tensor
|
||||||
|
continue
|
||||||
if name in loaded_names:
|
if name in loaded_names:
|
||||||
raise ComponentCheckpointUnsupportedError(
|
raise ComponentCheckpointUnsupportedError(
|
||||||
f"Direct GPU VAE checkpoint maps multiple tensors to {name!r}"
|
f"Direct GPU VAE checkpoint maps multiple tensors to {name!r}"
|
||||||
@@ -378,6 +487,9 @@ def _assign_direct_gpu_vae_state(
|
|||||||
module._buffers[local_name] = tensor
|
module._buffers[local_name] = tensor
|
||||||
loaded_names.add(name)
|
loaded_names.add(name)
|
||||||
|
|
||||||
|
consumed_metadata = _consume_vae_checkpoint_arch_metadata(
|
||||||
|
metadata, vae_config, target_state
|
||||||
|
)
|
||||||
missing = sorted(set(slots) - loaded_names)
|
missing = sorted(set(slots) - loaded_names)
|
||||||
if missing:
|
if missing:
|
||||||
raise ComponentCheckpointUnsupportedError(
|
raise ComponentCheckpointUnsupportedError(
|
||||||
@@ -390,6 +502,7 @@ def _assign_direct_gpu_vae_state(
|
|||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"Direct GPU VAE loading left meta tensors: {remaining_meta}"
|
f"Direct GPU VAE loading left meta tensors: {remaining_meta}"
|
||||||
)
|
)
|
||||||
|
return num_deparameterized, consumed_metadata
|
||||||
|
|
||||||
|
|
||||||
class VAELoader(WeightOverrideComponentLoader):
|
class VAELoader(WeightOverrideComponentLoader):
|
||||||
@@ -506,12 +619,16 @@ class VAELoader(WeightOverrideComponentLoader):
|
|||||||
|
|
||||||
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:
|
if direct_gpu_weight_loading and auto_model_map and not native_only:
|
||||||
raise ComponentCheckpointUnsupportedError(
|
raise ComponentCheckpointUnsupportedError(
|
||||||
f"Direct GPU loading for {component_name!r} requires a native "
|
f"Direct GPU loading for {component_name!r} requires a native "
|
||||||
"ModelRegistry VAE; custom Diffusers auto_map code is unsupported"
|
"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 not native_only
|
||||||
|
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 "
|
||||||
"consume a weights-only override"
|
"consume a weights-only override"
|
||||||
@@ -591,12 +708,14 @@ class VAELoader(WeightOverrideComponentLoader):
|
|||||||
f"Found no safetensors files in {component_weights_path}"
|
f"Found no safetensors files in {component_weights_path}"
|
||||||
)
|
)
|
||||||
if direct_gpu_weight_loading:
|
if direct_gpu_weight_loading:
|
||||||
_assign_direct_gpu_vae_state(
|
adaptations = _assign_direct_gpu_vae_state(
|
||||||
vae,
|
vae,
|
||||||
safetensors_list,
|
safetensors_list,
|
||||||
component_name=component_name,
|
component_name=component_name,
|
||||||
device=target_device,
|
device=target_device,
|
||||||
|
vae_config=vae_config,
|
||||||
)
|
)
|
||||||
|
_log_vae_checkpoint_adaptations(*adaptations)
|
||||||
if _should_use_channels_last_3d(server_args, component_name):
|
if _should_use_channels_last_3d(server_args, component_name):
|
||||||
n = _convert_conv3d_weights_to_channels_last_3d(vae)
|
n = _convert_conv3d_weights_to_channels_last_3d(vae)
|
||||||
if n > 0:
|
if n > 0:
|
||||||
@@ -610,6 +729,12 @@ class VAELoader(WeightOverrideComponentLoader):
|
|||||||
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))
|
||||||
_backfill_ltx2_audio_vae_latent_stats(loaded, component_type)
|
_backfill_ltx2_audio_vae_latent_stats(loaded, component_type)
|
||||||
|
num_deparameterized = _adopt_plain_weight_norm_state(vae, loaded)
|
||||||
|
target_state = vae.state_dict()
|
||||||
|
consumed_metadata = _consume_vae_checkpoint_arch_metadata(
|
||||||
|
loaded, vae_config, target_state
|
||||||
|
)
|
||||||
|
_log_vae_checkpoint_adaptations(num_deparameterized, consumed_metadata)
|
||||||
strict_load = native_only
|
strict_load = native_only
|
||||||
# `loaded` holds views into the safetensors mapping. When the component
|
# `loaded` holds views into the safetensors mapping. When the component
|
||||||
# starts on the CPU and the host cannot afford copies of the whole
|
# starts on the CPU and the host cannot afford copies of the whole
|
||||||
@@ -635,7 +760,7 @@ class VAELoader(WeightOverrideComponentLoader):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
if keep_mapping:
|
if keep_mapping:
|
||||||
_match_checkpoint_dtypes(loaded, vae.state_dict())
|
_match_checkpoint_dtypes(loaded, target_state)
|
||||||
vae.load_state_dict(
|
vae.load_state_dict(
|
||||||
loaded,
|
loaded,
|
||||||
strict=strict_load,
|
strict=strict_load,
|
||||||
|
|||||||
@@ -85,13 +85,14 @@ def validate_minimax_h3_checkpoint_variant(
|
|||||||
checkpoint_paths: list[str], selected_variant: str
|
checkpoint_paths: list[str], selected_variant: str
|
||||||
) -> None:
|
) -> None:
|
||||||
names = " ".join(path.lower() for path in checkpoint_paths)
|
names = " ".join(path.lower() for path in checkpoint_paths)
|
||||||
checkpoint_variant = next(
|
checkpoint_variants = {
|
||||||
(variant for variant in ("fl2va", "ref2va") if variant in names), None
|
variant for variant in ("fl2va", "ref2va") if variant in names
|
||||||
)
|
}
|
||||||
if (
|
if (
|
||||||
checkpoint_variant is not None
|
len(checkpoint_variants) == 1
|
||||||
and checkpoint_variant != selected_variant.lower()
|
and selected_variant.lower() not in checkpoint_variants
|
||||||
):
|
):
|
||||||
|
(checkpoint_variant,) = checkpoint_variants
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"MiniMax-H3 checkpoint variant {checkpoint_variant!r} does not match "
|
f"MiniMax-H3 checkpoint variant {checkpoint_variant!r} does not match "
|
||||||
f"--model-variant {selected_variant!r}"
|
f"--model-variant {selected_variant!r}"
|
||||||
|
|||||||
@@ -600,6 +600,9 @@ class LoRAPipeline(ComposedPipelineBase):
|
|||||||
if merge_mode == "dynamic":
|
if merge_mode == "dynamic":
|
||||||
return False
|
return False
|
||||||
uses_dtensor_weights = self._uses_dtensor_weights(lora_layers)
|
uses_dtensor_weights = self._uses_dtensor_weights(lora_layers)
|
||||||
|
has_unmergeable_weights = any(
|
||||||
|
not layer.can_merge_base_weight for layer in lora_layers.values()
|
||||||
|
)
|
||||||
if merge_mode == "auto":
|
if merge_mode == "auto":
|
||||||
if uses_dtensor_weights:
|
if uses_dtensor_weights:
|
||||||
logger.info(
|
logger.info(
|
||||||
@@ -607,7 +610,18 @@ class LoRAPipeline(ComposedPipelineBase):
|
|||||||
module_name,
|
module_name,
|
||||||
)
|
)
|
||||||
return False
|
return False
|
||||||
|
if has_unmergeable_weights:
|
||||||
|
logger.info(
|
||||||
|
"Using dynamic LoRA for %s because its quantized weights cannot be merged in place.",
|
||||||
|
module_name,
|
||||||
|
)
|
||||||
|
return False
|
||||||
return True
|
return True
|
||||||
|
if has_unmergeable_weights:
|
||||||
|
raise ValueError(
|
||||||
|
f"LoRA merge mode is unavailable for {module_name} because its "
|
||||||
|
"quantized weights cannot be updated in place; use merge mode 'dynamic'"
|
||||||
|
)
|
||||||
if uses_dtensor_weights:
|
if uses_dtensor_weights:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Merging LoRA for %s with FSDP-sharded weights may require full-gather and can OOM.",
|
"Merging LoRA for %s with FSDP-sharded weights may require full-gather and can OOM.",
|
||||||
|
|||||||
@@ -956,7 +956,28 @@ def maybe_download_model(
|
|||||||
f"Cached ref for {model_name_or_path} is corrupt: resolved to the "
|
f"Cached ref for {model_name_or_path} is corrupt: resolved to the "
|
||||||
f"snapshots parent {local_path!r} instead of a revision directory."
|
f"snapshots parent {local_path!r} instead of a revision directory."
|
||||||
)
|
)
|
||||||
if not force_diffusers_model:
|
required_files = [
|
||||||
|
pattern for pattern in allow_patterns or () if not glob.has_magic(pattern)
|
||||||
|
]
|
||||||
|
missing_required_files = [
|
||||||
|
path
|
||||||
|
for path in required_files
|
||||||
|
if not os.path.isfile(os.path.join(local_path, path))
|
||||||
|
]
|
||||||
|
if missing_required_files:
|
||||||
|
if not download:
|
||||||
|
raise ValueError(
|
||||||
|
f"Model {model_name_or_path} is cached but is missing requested "
|
||||||
|
f"files: {missing_required_files}."
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
"Cached snapshot for %s is missing requested files %s; "
|
||||||
|
"will download them from %s",
|
||||||
|
model_name_or_path,
|
||||||
|
missing_required_files,
|
||||||
|
_model_hub_name(),
|
||||||
|
)
|
||||||
|
elif not force_diffusers_model:
|
||||||
# maybe_download_model_index's model_index.json fetch materializes a full
|
# maybe_download_model_index's model_index.json fetch materializes a full
|
||||||
# cache entry, so this resolve reports that stub as a hit; returning it
|
# cache entry, so this resolve reports that stub as a hit; returning it
|
||||||
# would skip the download. LoRA repos declare no components.
|
# would skip the download. LoRA repos declare no components.
|
||||||
|
|||||||
@@ -2713,7 +2713,7 @@
|
|||||||
"expected_e2e_ms": 5648.49,
|
"expected_e2e_ms": 5648.49,
|
||||||
"expected_avg_denoise_ms": 477.49,
|
"expected_avg_denoise_ms": 477.49,
|
||||||
"expected_median_denoise_ms": 56.74,
|
"expected_median_denoise_ms": 56.74,
|
||||||
"load_peak_vram_mb": 32892.0,
|
"load_peak_vram_mb": 34022.0,
|
||||||
"runtime_peak_vram_mb": 61190.0,
|
"runtime_peak_vram_mb": 61190.0,
|
||||||
"estimated_full_test_time_s": 153.1
|
"estimated_full_test_time_s": 153.1
|
||||||
},
|
},
|
||||||
@@ -2731,7 +2731,7 @@
|
|||||||
"expected_e2e_ms": 9086.81,
|
"expected_e2e_ms": 9086.81,
|
||||||
"expected_avg_denoise_ms": 492.02,
|
"expected_avg_denoise_ms": 492.02,
|
||||||
"expected_median_denoise_ms": 151.9,
|
"expected_median_denoise_ms": 151.9,
|
||||||
"load_peak_vram_mb": 32892.0,
|
"load_peak_vram_mb": 34022.0,
|
||||||
"runtime_peak_vram_mb": 62510.0,
|
"runtime_peak_vram_mb": 62510.0,
|
||||||
"estimated_full_test_time_s": 149.4
|
"estimated_full_test_time_s": 149.4
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -298,6 +298,46 @@ def test_metadata_only_cached_lora_snapshot_is_a_usable_hit(
|
|||||||
assert calls == ["probe"]
|
assert calls == ["probe"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_cached_lora_snapshot_downloads_missing_selected_weight(monkeypatch, tmp_path):
|
||||||
|
calls = []
|
||||||
|
selected_file = "loras/adapter.safetensors"
|
||||||
|
|
||||||
|
def fake_snapshot_download(**kwargs):
|
||||||
|
calls.append("probe" if kwargs.get("local_files_only") else "download")
|
||||||
|
if not kwargs.get("local_files_only"):
|
||||||
|
target = tmp_path / selected_file
|
||||||
|
target.parent.mkdir()
|
||||||
|
target.write_bytes(b"weights")
|
||||||
|
return str(tmp_path)
|
||||||
|
|
||||||
|
monkeypatch.setattr(hf_diffusers_utils, "snapshot_download", fake_snapshot_download)
|
||||||
|
|
||||||
|
result = maybe_download_model(
|
||||||
|
"org/repo",
|
||||||
|
is_lora=True,
|
||||||
|
allow_patterns=["*.json", selected_file],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result == str(tmp_path)
|
||||||
|
assert calls == ["probe", "download"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_cached_lora_snapshot_reports_missing_selected_weight_offline(
|
||||||
|
recording_snapshot_download, tmp_path
|
||||||
|
):
|
||||||
|
calls = recording_snapshot_download(tmp_path)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="loras/adapter.safetensors"):
|
||||||
|
maybe_download_model(
|
||||||
|
"org/repo",
|
||||||
|
download=False,
|
||||||
|
is_lora=True,
|
||||||
|
allow_patterns=["*.json", "loras/adapter.safetensors"],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert calls == ["probe"]
|
||||||
|
|
||||||
|
|
||||||
def test_force_diffusers_model_stub_keeps_its_existing_path(
|
def test_force_diffusers_model_stub_keeps_its_existing_path(
|
||||||
recording_snapshot_download, tmp_path
|
recording_snapshot_download, tmp_path
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -3,13 +3,16 @@ from contextlib import contextmanager, nullcontext
|
|||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import patch
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
|
||||||
from sglang.multimodal_gen.runtime.layers.lora.linear import (
|
from sglang.multimodal_gen.runtime.layers.lora.linear import (
|
||||||
BaseLayerWithLoRA,
|
BaseLayerWithLoRA,
|
||||||
_use_owned_base_snapshot,
|
_use_owned_base_snapshot,
|
||||||
wrap_with_lora_layer,
|
wrap_with_lora_layer,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization.fp8 import Fp8Config
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import LoRAPipeline
|
from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import LoRAPipeline
|
||||||
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_lora
|
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_lora
|
||||||
|
|
||||||
@@ -94,6 +97,30 @@ def test_zero_copy_snapshot_is_limited_to_cpu_backed_layers():
|
|||||||
assert meta_layer._base_is_view
|
assert meta_layer._base_is_view
|
||||||
|
|
||||||
|
|
||||||
|
def test_quantized_base_uses_dynamic_lora_in_auto_mode():
|
||||||
|
with patch(
|
||||||
|
"sglang.multimodal_gen.runtime.layers.quantization.fp8."
|
||||||
|
"get_tensor_model_parallel_world_size",
|
||||||
|
return_value=1,
|
||||||
|
):
|
||||||
|
base_layer = ReplicatedLinear(
|
||||||
|
2,
|
||||||
|
2,
|
||||||
|
bias=False,
|
||||||
|
quant_config=Fp8Config(is_checkpoint_fp8_serialized=True),
|
||||||
|
)
|
||||||
|
layer = BaseLayerWithLoRA(base_layer)
|
||||||
|
pipeline = _make_pipeline(layer)
|
||||||
|
|
||||||
|
assert not pipeline._should_merge_lora_for_layers(
|
||||||
|
"transformer", {"linear": layer}, "auto"
|
||||||
|
)
|
||||||
|
with pytest.raises(ValueError, match="use merge mode 'dynamic'"):
|
||||||
|
pipeline._should_merge_lora_for_layers(
|
||||||
|
"transformer", {"linear": layer}, "merge"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_dynamic_lora_reactivates_cached_layers_without_weight_update_context():
|
def test_dynamic_lora_reactivates_cached_layers_without_weight_update_context():
|
||||||
layer = _make_layer()
|
layer = _make_layer()
|
||||||
pipeline = _make_pipeline(layer)
|
pipeline = _make_pipeline(layer)
|
||||||
|
|||||||
@@ -93,6 +93,7 @@ from sglang.multimodal_gen.runtime.loader.component_loaders.transformer_loader i
|
|||||||
from sglang.multimodal_gen.runtime.loader.minimax_h3_weights import (
|
from sglang.multimodal_gen.runtime.loader.minimax_h3_weights import (
|
||||||
inspect_minimax_h3_safetensors,
|
inspect_minimax_h3_safetensors,
|
||||||
resolve_minimax_h3_checkpoint_quantization,
|
resolve_minimax_h3_checkpoint_quantization,
|
||||||
|
validate_minimax_h3_checkpoint_variant,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.transformer_load_utils import (
|
from sglang.multimodal_gen.runtime.loader.transformer_load_utils import (
|
||||||
TransformerQuantLoadSpec,
|
TransformerQuantLoadSpec,
|
||||||
@@ -691,6 +692,18 @@ class TestTransformerQuantHelpers(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_minimax_h3_hybrid_checkpoint_accepts_selected_partition(self):
|
||||||
|
checkpoint = "/cache/minimax_h3_hybrid_fl2va_ref2va_b25-49.safetensors"
|
||||||
|
|
||||||
|
validate_minimax_h3_checkpoint_variant([checkpoint], "fl2va")
|
||||||
|
validate_minimax_h3_checkpoint_variant([checkpoint], "ref2va")
|
||||||
|
|
||||||
|
def test_minimax_h3_single_partition_checkpoint_rejects_mismatch(self):
|
||||||
|
with self.assertRaisesRegex(ValueError, "does not match"):
|
||||||
|
validate_minimax_h3_checkpoint_variant(
|
||||||
|
["/cache/minimax_h3_fl2va.safetensors"], "ref2va"
|
||||||
|
)
|
||||||
|
|
||||||
def test_inspect_minimax_h3_safetensors_detects_curve_and_comfy_format(self):
|
def test_inspect_minimax_h3_safetensors_detects_curve_and_comfy_format(self):
|
||||||
marker = json.dumps(
|
marker = json.dumps(
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -7,6 +7,9 @@ import torch
|
|||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from safetensors.torch import save_file as safetensors_save_file
|
from safetensors.torch import save_file as safetensors_save_file
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.configs.models.vaes.minimax_h3_audio import (
|
||||||
|
MiniMaxH3AudioVAEConfig,
|
||||||
|
)
|
||||||
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 (
|
||||||
QwenImagePipelineConfig,
|
QwenImagePipelineConfig,
|
||||||
@@ -24,8 +27,10 @@ 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 (
|
||||||
|
_adopt_plain_weight_norm_state,
|
||||||
_assign_direct_gpu_vae_state,
|
_assign_direct_gpu_vae_state,
|
||||||
_backfill_ltx2_audio_vae_latent_stats,
|
_backfill_ltx2_audio_vae_latent_stats,
|
||||||
|
_consume_vae_checkpoint_arch_metadata,
|
||||||
_direct_gpu_vae_state_slots,
|
_direct_gpu_vae_state_slots,
|
||||||
_match_checkpoint_dtypes,
|
_match_checkpoint_dtypes,
|
||||||
_require_native_loader_for_quantized_vae,
|
_require_native_loader_for_quantized_vae,
|
||||||
@@ -147,6 +152,54 @@ class TestMatchCheckpointDtypes(unittest.TestCase):
|
|||||||
self.assertIs(loaded["extra"], before)
|
self.assertIs(loaded["extra"], before)
|
||||||
|
|
||||||
|
|
||||||
|
class TestPlainWeightNormCheckpoint(unittest.TestCase):
|
||||||
|
def test_adopts_a_folded_weight_without_reconstructing_it(self):
|
||||||
|
module = nn.Sequential(
|
||||||
|
torch.nn.utils.parametrizations.weight_norm(
|
||||||
|
nn.Conv1d(2, 3, kernel_size=3, bias=False)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
expected = torch.arange(18, dtype=torch.float32).reshape(3, 2, 3) / 19
|
||||||
|
loaded = {"0.weight": expected}
|
||||||
|
|
||||||
|
self.assertEqual(_adopt_plain_weight_norm_state(module, loaded), 1)
|
||||||
|
module.load_state_dict(loaded, strict=True)
|
||||||
|
|
||||||
|
self.assertEqual(set(module.state_dict()), {"0.weight"})
|
||||||
|
self.assertTrue(torch.equal(module[0].weight, expected))
|
||||||
|
|
||||||
|
def test_keeps_legacy_weight_norm_state_parameterized(self):
|
||||||
|
module = nn.Sequential(
|
||||||
|
torch.nn.utils.parametrizations.weight_norm(
|
||||||
|
nn.Conv1d(2, 3, kernel_size=3, bias=False)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
original_state = module.state_dict()
|
||||||
|
loaded = {
|
||||||
|
"0.weight_g": original_state["0.parametrizations.weight.original0"].clone(),
|
||||||
|
"0.weight_v": original_state["0.parametrizations.weight.original1"].clone(),
|
||||||
|
}
|
||||||
|
|
||||||
|
self.assertEqual(_adopt_plain_weight_norm_state(module, loaded), 0)
|
||||||
|
module.load_state_dict(loaded, strict=True)
|
||||||
|
|
||||||
|
self.assertIn("0.parametrizations.weight.original0", module.state_dict())
|
||||||
|
|
||||||
|
def test_moves_checkpoint_latent_stats_into_arch_config(self):
|
||||||
|
config = MiniMaxH3AudioVAEConfig()
|
||||||
|
loaded = {
|
||||||
|
"latents_mean": torch.arange(32, dtype=torch.float32),
|
||||||
|
"latents_std": torch.arange(1, 33, dtype=torch.float32),
|
||||||
|
}
|
||||||
|
|
||||||
|
consumed = _consume_vae_checkpoint_arch_metadata(loaded, config, {})
|
||||||
|
|
||||||
|
self.assertEqual(consumed, ("latents_mean", "latents_std"))
|
||||||
|
self.assertEqual(config.arch_config.latents_mean, list(range(32)))
|
||||||
|
self.assertEqual(config.arch_config.latents_std, list(range(1, 33)))
|
||||||
|
self.assertEqual(loaded, {})
|
||||||
|
|
||||||
|
|
||||||
class TestDirectGPUVAEState(unittest.TestCase):
|
class TestDirectGPUVAEState(unittest.TestCase):
|
||||||
class _StandardVAE(nn.Module):
|
class _StandardVAE(nn.Module):
|
||||||
def __init__(self, *_args, **_kwargs):
|
def __init__(self, *_args, **_kwargs):
|
||||||
@@ -169,12 +222,49 @@ class TestDirectGPUVAEState(unittest.TestCase):
|
|||||||
[str(checkpoint)],
|
[str(checkpoint)],
|
||||||
component_name="vae",
|
component_name="vae",
|
||||||
device=torch.device("cpu"),
|
device=torch.device("cpu"),
|
||||||
|
vae_config=QwenImagePipelineConfig().vae_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.assertTrue(torch.equal(vae.proj.weight, expected_weight))
|
self.assertTrue(torch.equal(vae.proj.weight, expected_weight))
|
||||||
self.assertTrue(torch.equal(vae.scale, expected_scale))
|
self.assertTrue(torch.equal(vae.scale, expected_scale))
|
||||||
self.assertFalse(any(tensor.is_meta for tensor in vae.state_dict().values()))
|
self.assertFalse(any(tensor.is_meta for tensor in vae.state_dict().values()))
|
||||||
|
|
||||||
|
def test_direct_load_adopts_folded_weight_norm_and_checkpoint_metadata(self):
|
||||||
|
class _WeightNormVAE(nn.Module):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
self.proj = torch.nn.utils.parametrizations.weight_norm(
|
||||||
|
nn.Conv1d(2, 3, kernel_size=3, bias=False)
|
||||||
|
)
|
||||||
|
|
||||||
|
expected = torch.arange(18, dtype=torch.float32).reshape(3, 2, 3) / 19
|
||||||
|
config = MiniMaxH3AudioVAEConfig()
|
||||||
|
with TemporaryDirectory() as root:
|
||||||
|
checkpoint = pathlib.Path(root) / "model.safetensors"
|
||||||
|
safetensors_save_file(
|
||||||
|
{
|
||||||
|
"proj.weight": expected,
|
||||||
|
"latents_mean": torch.arange(32, dtype=torch.float32),
|
||||||
|
"latents_std": torch.arange(1, 33, dtype=torch.float32),
|
||||||
|
},
|
||||||
|
checkpoint,
|
||||||
|
)
|
||||||
|
with torch.device("meta"):
|
||||||
|
vae = _WeightNormVAE()
|
||||||
|
adaptations = _assign_direct_gpu_vae_state(
|
||||||
|
vae,
|
||||||
|
[str(checkpoint)],
|
||||||
|
component_name="audio_vae",
|
||||||
|
device=torch.device("cpu"),
|
||||||
|
vae_config=config,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(adaptations, (1, ("latents_mean", "latents_std")))
|
||||||
|
self.assertEqual(set(vae.state_dict()), {"proj.weight"})
|
||||||
|
self.assertTrue(torch.equal(vae.proj.weight, expected))
|
||||||
|
self.assertEqual(config.arch_config.latents_mean, list(range(32)))
|
||||||
|
self.assertEqual(config.arch_config.latents_std, list(range(1, 33)))
|
||||||
|
|
||||||
def test_rejects_nonstandard_state_lifecycle(self):
|
def test_rejects_nonstandard_state_lifecycle(self):
|
||||||
class _CustomVAE(self._StandardVAE):
|
class _CustomVAE(self._StandardVAE):
|
||||||
def state_dict(self, *args, **kwargs):
|
def state_dict(self, *args, **kwargs):
|
||||||
@@ -243,6 +333,59 @@ class TestDirectGPUVAEState(unittest.TestCase):
|
|||||||
self.assertTrue(torch.equal(loaded.proj.weight, expected_weight))
|
self.assertTrue(torch.equal(loaded.proj.weight, expected_weight))
|
||||||
self.assertTrue(torch.equal(loaded.scale, expected_scale))
|
self.assertTrue(torch.equal(loaded.scale, expected_scale))
|
||||||
|
|
||||||
|
def test_native_vae_ignores_diffusers_auto_map_for_weight_override(self):
|
||||||
|
loader = vae_loader.VAELoader()
|
||||||
|
expected_weight = torch.arange(4, dtype=torch.bfloat16).reshape(2, 2)
|
||||||
|
expected_scale = torch.tensor([3.0], dtype=torch.bfloat16)
|
||||||
|
|
||||||
|
with TemporaryDirectory() as root:
|
||||||
|
checkpoint = pathlib.Path(root) / "override.safetensors"
|
||||||
|
safetensors_save_file(
|
||||||
|
{"proj.weight": expected_weight, "scale": expected_scale}, checkpoint
|
||||||
|
)
|
||||||
|
|
||||||
|
for direct_gpu_loading in (False, True):
|
||||||
|
with self.subTest(direct_gpu_loading=direct_gpu_loading):
|
||||||
|
pipeline_config = QwenImagePipelineConfig()
|
||||||
|
pipeline_config.native_only_components = ("vae",)
|
||||||
|
server_args = _FakeServerArgs(pipeline_config)
|
||||||
|
server_args.component_direct_gpu_weight_loading = {
|
||||||
|
"vae": direct_gpu_loading
|
||||||
|
}
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(
|
||||||
|
vae_loader,
|
||||||
|
"get_diffusers_component_config",
|
||||||
|
return_value={
|
||||||
|
"_class_name": "TestVAE",
|
||||||
|
"auto_map": {"AutoModel": "custom.TestVAE"},
|
||||||
|
},
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
loader,
|
||||||
|
"resolve_component_weights_path",
|
||||||
|
return_value=str(checkpoint),
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
vae_loader.ModelRegistry,
|
||||||
|
"resolve_model_cls",
|
||||||
|
return_value=(self._StandardVAE, None),
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
loader, "target_device", return_value=torch.device("cpu")
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
vae_loader.current_platform,
|
||||||
|
"optimize_vae",
|
||||||
|
side_effect=lambda vae: vae,
|
||||||
|
),
|
||||||
|
):
|
||||||
|
loaded = loader.load_customized(root, server_args, "vae")
|
||||||
|
|
||||||
|
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):
|
def test_quantized_checkpoint_does_not_fall_back_from_direct_loading(self):
|
||||||
loader = vae_loader.VAELoader()
|
loader = vae_loader.VAELoader()
|
||||||
server_args = _FakeServerArgs(QwenImagePipelineConfig())
|
server_args = _FakeServerArgs(QwenImagePipelineConfig())
|
||||||
|
|||||||
Reference in New Issue
Block a user