[diffusion] feat: compose third-party component bundles safely (#37816)

This commit is contained in:
Mick
2026-09-04 09:15:05 +08:00
committed by GitHub
parent 516cfbd362
commit 795dd7abce
10 changed files with 408 additions and 13 deletions
@@ -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())