[diffusion] feat: generalize layer-wise-offload to all supported models (#16150)

This commit is contained in:
Mick
2025-12-30 22:06:57 +08:00
committed by GitHub
parent b3817fa93b
commit 3449806727
10 changed files with 146 additions and 139 deletions
@@ -43,9 +43,7 @@ from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
get_diffusers_component_config,
get_hf_config,
)
from sglang.multimodal_gen.runtime.utils.layerwise_offload import (
LayerwiseOffloadManager,
)
from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
@@ -740,23 +738,14 @@ class TransformerLoader(ComponentLoader):
model = model.eval()
if server_args.dit_layerwise_offload and hasattr(model, "dit_module_names"):
# TODO(will): support multiple module names
module_name = getattr(model, "dit_module_names", ["transformer_blocks"])[0]
try:
num_layers = len(getattr(model, module_name))
except Exception:
num_layers = None
if isinstance(num_layers, int) and num_layers > 0:
mgr = LayerwiseOffloadManager(
model,
module_list_attr=module_name,
num_layers=num_layers,
enabled=True,
pin_cpu_memory=server_args.pin_cpu_memory,
auto_initialize=True,
if server_args.dit_layerwise_offload:
# enable layerwise offload if possible
if isinstance(model, OffloadableDiTMixin):
model.configure_layerwise_offload(server_args)
else:
logger.info(
"Disabling layerwise offload since current model does not support this feature"
)
setattr(model, "_layerwise_offload_manager", mgr)
return model