[diffusion] feat: generalize layer-wise-offload to all supported models (#16150)
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user