[diffusion] chore: re-home decode-dtype vae weights to a file-backed store (#35986)

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Mick
2026-08-23 09:32:04 +08:00
committed by GitHub
co-authored by Claude Fable 5
parent a36c0746b9
commit 59cdc9de7d
3 changed files with 235 additions and 7 deletions
+6
View File
@@ -22,6 +22,7 @@ if TYPE_CHECKING:
SGLANG_DIFFUSION_LOGGING_PREFIX: str = ""
SGLANG_DIFFUSION_TRACE_FUNCTION: int = 0
SGLANG_DIFFUSION_DISABLE_EARLY_VAE_DECODER_CAST: bool = False
SGLANG_DIFFUSION_DISABLE_VAE_DECODER_STORE: bool = False
SGLANG_DIFFUSION_WORKER_MULTIPROC_METHOD: str = "fork"
SGLANG_DIFFUSION_TARGET_DEVICE: str = "cuda"
SGLANG_DIFFUSION_PLATFORM_OVERRIDE: str = ""
@@ -272,6 +273,11 @@ environment_variables: dict[str, Callable[[], Any]] = {
"SGLANG_DIFFUSION_DISABLE_EARLY_VAE_DECODER_CAST": _lazy_bool(
"SGLANG_DIFFUSION_DISABLE_EARLY_VAE_DECODER_CAST"
),
# Kill-switch: keep the decode-dtype VAE decoder weights in anonymous host
# memory instead of a file-backed cache mapping the page cache can drop.
"SGLANG_DIFFUSION_DISABLE_VAE_DECODER_STORE": _lazy_bool(
"SGLANG_DIFFUSION_DISABLE_VAE_DECODER_STORE"
),
# ================== cache-dit Env Vars ==================
# Enable cache-dit acceleration for DiT inference
# CUDA-IPC transport for 2-rank Ulysses all-to-all (NVLink same-node)
@@ -1,9 +1,11 @@
import hashlib
import importlib.util
import os
import torch
import torch.nn as nn
from safetensors.torch import load_file as safetensors_load_file
from safetensors.torch import save_file as safetensors_save_file
from sglang.multimodal_gen import envs
from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import LTX2PipelineConfig
@@ -131,8 +133,85 @@ def _should_use_channels_last_3d(
return False
def _decode_dtype_store_path(
component_model_path: str, component_name: str, dtype: torch.dtype
) -> str:
key = hashlib.sha1(
f"{os.path.realpath(component_model_path)}|{component_name}|{dtype}".encode()
).hexdigest()[:16]
return os.path.join(
envs.SGLANG_DIFFUSION_CACHE_ROOT, "decode_dtype_store", f"{key}.safetensors"
)
def _assign_matching_store(vae, mapped: dict, dtype: torch.dtype) -> bool:
"""Adopt a decode-dtype store if it matches the module, else refuse."""
state = vae.state_dict()
for name, tensor in mapped.items():
param = state.get(name)
if param is None or param.shape != tensor.shape or tensor.dtype != dtype:
return False
vae.load_state_dict(mapped, strict=False, assign=True)
return True
def _rehome_cast_weights_to_file(
vae, dtype: torch.dtype, component_model_path: str, component_name: str, prepare
) -> tuple[int, bool]:
"""Hold the decode-dtype weights in a file-backed mapping.
The cast copies are anonymous host memory the kernel cannot reclaim, and
on a budgeted host every one of those bytes comes out of the pin budget
the stepped components live on. Written once to a cache file and mapped
back, the same bytes become page cache: droppable under pressure, free to
re-fault, and absent from the anonymous accounting. safetensors round-trips
tensors byte-exactly, so the mapping holds the identical rounded values —
and a later start adopts the store without paying the cast at all.
Returns (weights held, file-backed?).
"""
path = _decode_dtype_store_path(component_model_path, component_name, dtype)
try:
if os.path.exists(path):
mapped = safetensors_load_file(path)
if mapped and _assign_matching_store(vae, mapped, dtype):
return len(mapped), True
raise ValueError("existing decode-dtype store does not match the module")
converted = prepare(dtype)
if not converted:
return 0, False
cast_state = {
name: tensor
for name, tensor in vae.state_dict().items()
if tensor.dtype == dtype and tensor.device.type == "cpu"
}
os.makedirs(os.path.dirname(path), exist_ok=True)
tmp = f"{path}.tmp.{os.getpid()}"
safetensors_save_file({k: v.contiguous() for k, v in cast_state.items()}, tmp)
os.replace(tmp, path)
mapped = safetensors_load_file(path)
if set(mapped) != set(cast_state):
raise ValueError("decode-dtype store does not match the cast weights")
vae.load_state_dict(mapped, strict=False, assign=True)
return converted, True
except Exception as exc:
logger.warning(
"VAE: could not re-home %s decode-dtype weights to %s (%s); "
"keeping in-memory copies",
component_name,
path,
exc,
)
try:
if os.path.exists(path):
os.remove(path)
except OSError:
pass
return prepare(dtype), False
def _hold_decoder_weights_in_decode_dtype(
vae, server_args: ServerArgs, component_name: str
vae, server_args: ServerArgs, component_name: str, component_model_path: str = ""
) -> None:
"""Round decoder weights to their decode compute dtype at load.
@@ -156,13 +235,19 @@ def _hold_decoder_weights_in_decode_dtype(
return
if not autocast_enabled(dtype, server_args.disable_autocast):
return
converted = prepare(dtype)
if converted:
if component_model_path and not envs.SGLANG_DIFFUSION_DISABLE_VAE_DECODER_STORE:
held, file_backed = _rehome_cast_weights_to_file(
vae, dtype, component_model_path, component_name, prepare
)
else:
held, file_backed = prepare(dtype), False
if held:
logger.info(
"VAE: %s holds %d decoder weights in %s from load",
"VAE: %s holds %d decoder weights in %s from load (%s)",
component_name,
converted,
held,
dtype,
"file-backed" if file_backed else "anonymous host memory",
)
@@ -275,7 +360,9 @@ class VAELoader(ComponentLoader):
logger.info(
"VAE: converted %d Conv3d weights to channels_last_3d", n
)
_hold_decoder_weights_in_decode_dtype(vae, server_args, component_name)
_hold_decoder_weights_in_decode_dtype(
vae, server_args, component_name, component_model_path
)
vae = current_platform.optimize_vae(vae)
return vae
@@ -349,6 +436,8 @@ class VAELoader(ComponentLoader):
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)
_hold_decoder_weights_in_decode_dtype(
vae, server_args, component_name, component_model_path
)
vae = current_platform.optimize_vae(vae)
return vae
@@ -0,0 +1,133 @@
"""The decode-dtype weights live in a file-backed mapping, not anonymous memory.
What matters: the store round-trips the exact rounded bytes, a later start
adopts it without paying the cast, a mismatched store is discarded rather than
adopted, and the kill switch keeps everything in memory.
"""
from types import SimpleNamespace
import pytest
import torch
import torch.nn as nn
from sglang.multimodal_gen.runtime.loader.component_loaders.vae_loader import (
_decode_dtype_store_path,
_hold_decoder_weights_in_decode_dtype,
)
class _TinyVAE(nn.Module):
def __init__(self):
super().__init__()
self.blocks = nn.ModuleList([nn.Linear(8, 8) for _ in range(3)])
self.head = nn.Linear(8, 8) # stays fp32, like the output projection
self.prepare_calls = 0
def prepare_decoder_autocast_weights(self, dtype) -> int:
self.prepare_calls += 1
converted = 0
for block in self.blocks:
if block.weight.dtype != dtype:
block.to(dtype=dtype)
converted += 1
return converted
def _server_args():
return SimpleNamespace(
pipeline_config=SimpleNamespace(
vae_decode_precision="fp16",
vae_precision="fp32",
),
disable_autocast=False,
)
@pytest.fixture(autouse=True)
def _env(monkeypatch, tmp_path):
from sglang.multimodal_gen.runtime.utils import precision
monkeypatch.setattr(precision.current_platform, "is_amp_supported", lambda: True)
monkeypatch.setenv("SGLANG_DIFFUSION_CACHE_ROOT", str(tmp_path / "cache"))
def test_the_cast_weights_end_up_file_backed(tmp_path):
vae = _TinyVAE()
model_path = tmp_path / "ckpt"
model_path.mkdir()
_hold_decoder_weights_in_decode_dtype(
vae, _server_args(), "video_vae", str(model_path)
)
path = _decode_dtype_store_path(str(model_path), "video_vae", torch.float16)
import os
assert os.path.exists(path)
assert all(b.weight.dtype == torch.float16 for b in vae.blocks)
assert vae.head.weight.dtype == torch.float32
from safetensors.torch import load_file
stored = load_file(path)
for name, tensor in stored.items():
assert torch.equal(tensor, vae.state_dict()[name])
def test_a_second_start_adopts_the_store_without_casting(tmp_path):
model_path = tmp_path / "ckpt"
model_path.mkdir()
first = _TinyVAE()
_hold_decoder_weights_in_decode_dtype(
first, _server_args(), "video_vae", str(model_path)
)
second = _TinyVAE()
second.load_state_dict(
{
k: v.to(torch.float32) if v.dtype == torch.float16 else v
for k, v in first.state_dict().items()
}
)
_hold_decoder_weights_in_decode_dtype(
second, _server_args(), "video_vae", str(model_path)
)
assert second.prepare_calls == 0
for name in first.state_dict():
assert torch.equal(second.state_dict()[name], first.state_dict()[name])
def test_a_mismatched_store_is_discarded_and_the_cast_kept(tmp_path):
model_path = tmp_path / "ckpt"
model_path.mkdir()
path = _decode_dtype_store_path(str(model_path), "video_vae", torch.float16)
import os
os.makedirs(os.path.dirname(path), exist_ok=True)
from safetensors.torch import save_file
save_file({"blocks.0.weight": torch.zeros(4, 4, dtype=torch.float16)}, path)
vae = _TinyVAE()
_hold_decoder_weights_in_decode_dtype(
vae, _server_args(), "video_vae", str(model_path)
)
assert vae.prepare_calls == 1
assert all(b.weight.dtype == torch.float16 for b in vae.blocks)
assert not os.path.exists(path)
def test_the_store_kill_switch_keeps_the_copies_in_memory(monkeypatch, tmp_path):
monkeypatch.setenv("SGLANG_DIFFUSION_DISABLE_VAE_DECODER_STORE", "1")
model_path = tmp_path / "ckpt"
model_path.mkdir()
vae = _TinyVAE()
_hold_decoder_weights_in_decode_dtype(
vae, _server_args(), "video_vae", str(model_path)
)
path = _decode_dtype_store_path(str(model_path), "video_vae", torch.float16)
import os
assert not os.path.exists(path)
assert all(b.weight.dtype == torch.float16 for b in vae.blocks)