[diffusion] feat: keep a cpu-started vae weights on the checkpoint mapping (#35862)
Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
@@ -17,6 +17,8 @@ from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader imp
|
|||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.utils import (
|
from sglang.multimodal_gen.runtime.loader.utils import (
|
||||||
_list_safetensors_files,
|
_list_safetensors_files,
|
||||||
|
checkpoint_bytes,
|
||||||
|
keep_checkpoint_mapped,
|
||||||
set_default_torch_dtype,
|
set_default_torch_dtype,
|
||||||
skip_init_modules,
|
skip_init_modules,
|
||||||
)
|
)
|
||||||
@@ -124,6 +126,21 @@ def _should_use_channels_last_3d(
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _match_checkpoint_dtypes(loaded: dict, target_state: dict) -> dict:
|
||||||
|
"""Convert checkpoint tensors whose dtype differs from their parameter's.
|
||||||
|
|
||||||
|
Assignment replaces the parameter rather than writing through it, so a
|
||||||
|
mismatched dtype would silently change the module's. Converting makes a
|
||||||
|
copy, which is the point: only the tensors that already match can stay on
|
||||||
|
the mapping.
|
||||||
|
"""
|
||||||
|
for name, tensor in list(loaded.items()):
|
||||||
|
param = target_state.get(name)
|
||||||
|
if param is not None and param.dtype != tensor.dtype:
|
||||||
|
loaded[name] = tensor.to(dtype=param.dtype)
|
||||||
|
return loaded
|
||||||
|
|
||||||
|
|
||||||
class VAELoader(ComponentLoader):
|
class VAELoader(ComponentLoader):
|
||||||
"""Shared loader for (video/audio) VAE modules."""
|
"""Shared loader for (video/audio) VAE modules."""
|
||||||
|
|
||||||
@@ -242,10 +259,30 @@ class VAELoader(ComponentLoader):
|
|||||||
loaded.update(safetensors_load_file(sf_path))
|
loaded.update(safetensors_load_file(sf_path))
|
||||||
_backfill_ltx2_audio_vae_latent_stats(loaded, component_name)
|
_backfill_ltx2_audio_vae_latent_stats(loaded, component_name)
|
||||||
strict_load = native_only
|
strict_load = native_only
|
||||||
|
# `loaded` holds views into the safetensors mapping. When the component
|
||||||
|
# starts on the CPU and the host cannot afford copies of the whole
|
||||||
|
# deployment, assigning them keeps the weights file-backed instead of
|
||||||
|
# copying them into anonymous host memory: the page cache can drop and
|
||||||
|
# refetch file-backed bytes, and every anonymous byte here is a byte
|
||||||
|
# the stepped components' pin budget loses -- MiniMax-H3's video VAE is
|
||||||
|
# 9.70 GiB of a 32 GiB budget. On a host with room the copy stays the
|
||||||
|
# default, because its pages are resident where a mapping's first use
|
||||||
|
# pays a fault. MPS always assigns; the memory is unified. A tensor
|
||||||
|
# whose dtype differs from its parameter's is converted, which copies
|
||||||
|
# exactly the tensors that cannot stay.
|
||||||
|
keep_mapping = component_starts_on_cpu and (
|
||||||
|
current_platform.is_mps()
|
||||||
|
or keep_checkpoint_mapped(
|
||||||
|
weight_bytes=checkpoint_bytes(server_args.model_path),
|
||||||
|
component=f"{component_name or 'vae'} (VAE)",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if keep_mapping:
|
||||||
|
_match_checkpoint_dtypes(loaded, vae.state_dict())
|
||||||
vae.load_state_dict(
|
vae.load_state_dict(
|
||||||
loaded,
|
loaded,
|
||||||
strict=strict_load,
|
strict=strict_load,
|
||||||
assign=bool(cpu_offload_flag and current_platform.is_mps()),
|
assign=keep_mapping,
|
||||||
)
|
)
|
||||||
|
|
||||||
if not strict_load:
|
if not strict_load:
|
||||||
|
|||||||
@@ -249,6 +249,46 @@ def _try_redownload_missing_shards(model_path: str, missing: list[str]) -> bool:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def checkpoint_bytes(model_path: str) -> int:
|
||||||
|
"""On-disk size of every safetensors under a path, readable before any is."""
|
||||||
|
total = 0
|
||||||
|
for path in glob.glob(
|
||||||
|
os.path.join(str(model_path), "**", "*.safetensors"), recursive=True
|
||||||
|
):
|
||||||
|
try:
|
||||||
|
total += os.path.getsize(path)
|
||||||
|
except OSError:
|
||||||
|
continue
|
||||||
|
return total
|
||||||
|
|
||||||
|
|
||||||
|
def keep_checkpoint_mapped(*, weight_bytes: int, component: str) -> bool:
|
||||||
|
"""Whether a component's weights should stay on their file mapping.
|
||||||
|
|
||||||
|
Judged against the whole deployment rather than the one component: on a
|
||||||
|
host that cannot afford copies of everything it is about to serve, every
|
||||||
|
byte of anonymous memory a copy takes is a byte the pin budget for the
|
||||||
|
stepped components loses. On a host with room, the copy is the faster
|
||||||
|
choice -- its pages are resident, where a mapping's first use pays a fault.
|
||||||
|
"""
|
||||||
|
from sglang.multimodal_gen.runtime.managers.memory_managers.host_memory_budget import (
|
||||||
|
host_copies_would_not_fit,
|
||||||
|
host_memory_available_bytes,
|
||||||
|
)
|
||||||
|
|
||||||
|
if not host_copies_would_not_fit(weight_bytes):
|
||||||
|
return False
|
||||||
|
logger.info(
|
||||||
|
"%s stays on its checkpoint mapping: the deployment is %.2f GiB of "
|
||||||
|
"weights against %.2f GiB of host memory, so copies are host memory "
|
||||||
|
"the streamed components need more.",
|
||||||
|
component,
|
||||||
|
weight_bytes / 1024**3,
|
||||||
|
host_memory_available_bytes() / 1024**3,
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
def _list_safetensors_files(model_path: str) -> list[str]:
|
def _list_safetensors_files(model_path: str) -> list[str]:
|
||||||
"""List all .safetensors files under a directory.
|
"""List all .safetensors files under a directory.
|
||||||
|
|
||||||
|
|||||||
@@ -20,9 +20,14 @@ from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader imp
|
|||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.component_loaders.vae_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.vae_loader import (
|
||||||
_backfill_ltx2_audio_vae_latent_stats,
|
_backfill_ltx2_audio_vae_latent_stats,
|
||||||
|
_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,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.loader.utils import keep_checkpoint_mapped
|
||||||
|
from sglang.multimodal_gen.runtime.managers.memory_managers import (
|
||||||
|
host_memory_budget,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.models.vaes import wanvae
|
from sglang.multimodal_gen.runtime.models.vaes import wanvae
|
||||||
|
|
||||||
|
|
||||||
@@ -41,6 +46,51 @@ class _FakeServerArgs:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
class TestKeepCheckpointMapped(unittest.TestCase):
|
||||||
|
"""The mapping is for hosts that cannot afford the whole deployment."""
|
||||||
|
|
||||||
|
def test_a_small_deployment_on_a_roomy_host_copies(self):
|
||||||
|
with unittest.mock.patch.object(
|
||||||
|
host_memory_budget, "host_memory_available_bytes", lambda: 64 * 1024**3
|
||||||
|
):
|
||||||
|
self.assertFalse(
|
||||||
|
keep_checkpoint_mapped(weight_bytes=3 * 1024**3, component="vae (VAE)"),
|
||||||
|
"copies are the faster choice when the host has room: their "
|
||||||
|
"pages are resident where a mapping's first use pays a fault",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_a_deployment_larger_than_the_host_stays_mapped(self):
|
||||||
|
with unittest.mock.patch.object(
|
||||||
|
host_memory_budget, "host_memory_available_bytes", lambda: 19 * 1024**3
|
||||||
|
):
|
||||||
|
self.assertTrue(
|
||||||
|
keep_checkpoint_mapped(
|
||||||
|
weight_bytes=117 * 1024**3, component="vae (VAE)"
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestMatchCheckpointDtypes(unittest.TestCase):
|
||||||
|
"""Assignment replaces a parameter, so only matching dtypes may stay mapped."""
|
||||||
|
|
||||||
|
def test_a_matching_tensor_is_left_alone(self):
|
||||||
|
loaded = {"w": torch.zeros(4, dtype=torch.float32)}
|
||||||
|
before = loaded["w"]
|
||||||
|
_match_checkpoint_dtypes(loaded, {"w": torch.zeros(4, dtype=torch.float32)})
|
||||||
|
self.assertIs(loaded["w"], before)
|
||||||
|
|
||||||
|
def test_a_mismatched_tensor_is_converted(self):
|
||||||
|
loaded = {"w": torch.zeros(4, dtype=torch.float32)}
|
||||||
|
_match_checkpoint_dtypes(loaded, {"w": torch.zeros(4, dtype=torch.bfloat16)})
|
||||||
|
self.assertEqual(loaded["w"].dtype, torch.bfloat16)
|
||||||
|
|
||||||
|
def test_a_tensor_the_module_does_not_want_is_left_alone(self):
|
||||||
|
loaded = {"extra": torch.zeros(4, dtype=torch.float32)}
|
||||||
|
before = loaded["extra"]
|
||||||
|
_match_checkpoint_dtypes(loaded, {})
|
||||||
|
self.assertIs(loaded["extra"], before)
|
||||||
|
|
||||||
|
|
||||||
class TestVAELoader(unittest.TestCase):
|
class TestVAELoader(unittest.TestCase):
|
||||||
def test_quantized_vae_admission_leaves_plain_configs_unchanged(self):
|
def test_quantized_vae_admission_leaves_plain_configs_unchanged(self):
|
||||||
_require_native_loader_for_quantized_vae(
|
_require_native_loader_for_quantized_vae(
|
||||||
|
|||||||
Reference in New Issue
Block a user