diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loader.py index 100bdc321..3ae15135d 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loader.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/models/dits/causal_wanvideo.py b/python/sglang/multimodal_gen/runtime/models/dits/causal_wanvideo.py index 1f9e9cb50..e4f4940e1 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/causal_wanvideo.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/causal_wanvideo.py @@ -13,6 +13,8 @@ from torch.nn.attention.flex_attention import ( flex_attention, ) +from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin + # wan 1.3B model has a weird channel / head configurations and require max-autotune to work with flexattention # see https://github.com/pytorch/pytorch/issues/133254 # change to default for other models @@ -421,7 +423,7 @@ class CausalWanTransformerBlock(nn.Module): return hidden_states -class CausalWanTransformer3DModel(BaseDiT): +class CausalWanTransformer3DModel(BaseDiT, OffloadableDiTMixin): _fsdp_shard_conditions = WanVideoConfig()._fsdp_shard_conditions _compile_conditions = WanVideoConfig()._compile_conditions _supported_attention_backends = WanVideoConfig()._supported_attention_backends @@ -505,6 +507,10 @@ class CausalWanTransformer3DModel(BaseDiT): self.__post_init__() + self.layer_names = [ + "blocks", + ] + @staticmethod def _prepare_blockwise_causal_attn_mask( device: torch.device | str, diff --git a/python/sglang/multimodal_gen/runtime/models/dits/flux.py b/python/sglang/multimodal_gen/runtime/models/dits/flux.py index d7ee2fd6d..b678a633a 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/flux.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/flux.py @@ -44,6 +44,7 @@ from sglang.multimodal_gen.runtime.layers.rotary_embedding import ( ) from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) # pylint: disable=invalid-name @@ -407,7 +408,7 @@ class FluxPosEmbed(nn.Module): return freqs_cos.contiguous().float(), freqs_sin.contiguous().float() -class FluxTransformer2DModel(CachableDiT): +class FluxTransformer2DModel(CachableDiT, OffloadableDiTMixin): """ The Transformer model introduced in Flux. @@ -426,10 +427,6 @@ class FluxTransformer2DModel(CachableDiT): self.inner_dim = ( self.config.num_attention_heads * self.config.attention_head_dim ) - self.dit_module_names = [ - "transformer_blocks", - "single_transformer_blocks", - ] self.rotary_emb = FluxPosEmbed(theta=10000, axes_dim=self.config.axes_dims_rope) @@ -484,6 +481,11 @@ class FluxTransformer2DModel(CachableDiT): gather_output=True, ) + self.layer_names = [ + "transformer_blocks", + "single_transformer_blocks", + ] + def forward( self, hidden_states: torch.Tensor, @@ -541,46 +543,22 @@ class FluxTransformer2DModel(CachableDiT): ip_hidden_states = self.encoder_hid_proj(ip_adapter_image_embeds) joint_attention_kwargs.update({"ip_hidden_states": ip_hidden_states}) - offload_mgr = getattr(self, "_layerwise_offload_manager", None) - if offload_mgr is not None and getattr(offload_mgr, "enabled", False): - for i, block in enumerate(self.transformer_blocks): - with offload_mgr.layer_scope( - prefetch_layer_idx=i + 1, - release_layer_idx=i, - non_blocking=True, - ): - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - freqs_cis=freqs_cis, - joint_attention_kwargs=joint_attention_kwargs, - ) - for block in self.single_transformer_blocks: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - freqs_cis=freqs_cis, - joint_attention_kwargs=joint_attention_kwargs, - ) - else: - for block in self.transformer_blocks: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - freqs_cis=freqs_cis, - joint_attention_kwargs=joint_attention_kwargs, - ) - for block in self.single_transformer_blocks: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - freqs_cis=freqs_cis, - joint_attention_kwargs=joint_attention_kwargs, - ) + for block in self.transformer_blocks: + encoder_hidden_states, hidden_states = block( + hidden_states=hidden_states, + encoder_hidden_states=encoder_hidden_states, + temb=temb, + freqs_cis=freqs_cis, + joint_attention_kwargs=joint_attention_kwargs, + ) + for block in self.single_transformer_blocks: + encoder_hidden_states, hidden_states = block( + hidden_states=hidden_states, + encoder_hidden_states=encoder_hidden_states, + temb=temb, + freqs_cis=freqs_cis, + joint_attention_kwargs=joint_attention_kwargs, + ) hidden_states = self.norm_out(hidden_states, temb) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py b/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py index 232af8fca..3b88d99bd 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py @@ -30,6 +30,7 @@ from sglang.multimodal_gen.runtime.layers.rotary_embedding import ( ) from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) # pylint: disable=invalid-name @@ -593,7 +594,7 @@ class Flux2PosEmbed(nn.Module): return freqs_cos.contiguous().float(), freqs_sin.contiguous().float() -class Flux2Transformer2DModel(CachableDiT): +class Flux2Transformer2DModel(CachableDiT, OffloadableDiTMixin): """ The Transformer model introduced in Flux 2. @@ -692,7 +693,7 @@ class Flux2Transformer2DModel(CachableDiT): self.inner_dim, patch_size * patch_size * self.out_channels, bias=False ) - self.gradient_checkpointing = False + self.layer_names = ["transformer_blocks", "single_transformer_blocks"] def forward( self, diff --git a/python/sglang/multimodal_gen/runtime/models/dits/hunyuanvideo.py b/python/sglang/multimodal_gen/runtime/models/dits/hunyuanvideo.py index 76ec6c471..e8dc063cc 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/hunyuanvideo.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/hunyuanvideo.py @@ -40,6 +40,7 @@ from sglang.multimodal_gen.runtime.platforms import ( AttentionBackendEnum, current_platform, ) +from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin class MMDoubleStreamBlock(nn.Module): @@ -386,7 +387,7 @@ class MMSingleStreamBlock(nn.Module): return self.output_residual(x, output, mod_gate) -class HunyuanVideoTransformer3DModel(CachableDiT): +class HunyuanVideoTransformer3DModel(CachableDiT, OffloadableDiTMixin): """ HunyuanVideo Transformer backbone adapted for distributed training. @@ -508,7 +509,7 @@ class HunyuanVideoTransformer3DModel(CachableDiT): mlp_ratio=config.mlp_ratio, dtype=config.dtype, supported_attention_backends=self._supported_attention_backends, - prefix=f"{config.prefix}.single_blocks.{i+config.num_layers}", + prefix=f"{config.prefix}.single_blocks.{i + config.num_layers}", ) for i in range(config.num_single_layers) ] @@ -524,6 +525,8 @@ class HunyuanVideoTransformer3DModel(CachableDiT): self.__post_init__() + self.layer_names = ["double_blocks", "single_blocks"] + # TODO: change the input the FORWARD_BATCH Dict # TODO: change output to a dict def forward( diff --git a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py index c3a2f2406..73d478dcb 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py @@ -32,6 +32,7 @@ from sglang.multimodal_gen.runtime.layers.triton_ops import ( ) from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum +from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) # pylint: disable=invalid-name @@ -798,7 +799,7 @@ def to_hashable(obj): return obj -class QwenImageTransformer2DModel(CachableDiT): +class QwenImageTransformer2DModel(CachableDiT, OffloadableDiTMixin): """ The Transformer model introduced in Qwen. @@ -878,6 +879,8 @@ class QwenImageTransformer2DModel(CachableDiT): (1,), dtype=torch.int, device=get_local_torch_device() ) + self.layer_names = ["transformer_blocks"] + @functools.lru_cache(maxsize=50) def build_modulate_index(self, img_shapes: tuple[int, int, int], device): modulate_index_list = [] diff --git a/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py b/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py index 80f19d035..60b98eac4 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py @@ -43,6 +43,7 @@ from sglang.multimodal_gen.runtime.platforms import ( current_platform, ) from sglang.multimodal_gen.runtime.server_args import get_global_server_args +from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) @@ -602,7 +603,7 @@ class WanTransformerBlock_VSA(nn.Module): return hidden_states -class WanTransformer3DModel(CachableDiT): +class WanTransformer3DModel(CachableDiT, OffloadableDiTMixin): _fsdp_shard_conditions = WanVideoConfig()._fsdp_shard_conditions _compile_conditions = WanVideoConfig()._compile_conditions _supported_attention_backends = WanVideoConfig()._supported_attention_backends @@ -621,7 +622,6 @@ class WanTransformer3DModel(CachableDiT): self.num_channels_latents = config.num_channels_latents self.patch_size = config.patch_size self.text_len = config.text_len - self.dit_module_names = ["blocks"] # 1. Patch & position embedding self.patch_embedding = PatchEmbed( @@ -708,6 +708,8 @@ class WanTransformer3DModel(CachableDiT): dtype=torch.float32 if current_platform.is_mps() else torch.float64, ) + self.layer_names = ["blocks"] + def forward( self, hidden_states: torch.Tensor, @@ -806,25 +808,10 @@ class WanTransformer3DModel(CachableDiT): if enable_teacache: original_hidden_states = hidden_states.clone() - offload_mgr = getattr(self, "_layerwise_offload_manager", None) - if offload_mgr is not None and getattr(offload_mgr, "enabled", False): - for i, block in enumerate(self.blocks): - with offload_mgr.layer_scope( - prefetch_layer_idx=i + 1, - release_layer_idx=i, - non_blocking=True, - ): - hidden_states = block( - hidden_states, - encoder_hidden_states, - timestep_proj, - freqs_cis, - ) - else: - for block in self.blocks: - hidden_states = block( - hidden_states, encoder_hidden_states, timestep_proj, freqs_cis - ) + for block in self.blocks: + hidden_states = block( + hidden_states, encoder_hidden_states, timestep_proj, freqs_cis + ) # if teacache is enabled, we need to cache the original hidden states if enable_teacache: self.maybe_cache_states(hidden_states, original_hidden_states) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/zimage.py b/python/sglang/multimodal_gen/runtime/models/dits/zimage.py index 9fb166b69..f9743cebf 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/zimage.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/zimage.py @@ -17,6 +17,7 @@ from sglang.multimodal_gen.runtime.layers.linear import ( from sglang.multimodal_gen.runtime.layers.rotary_embedding import _apply_rotary_emb from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) @@ -350,7 +351,7 @@ class RopeEmbedder: return torch.cat(cos_out, dim=-1), torch.cat(sin_out, dim=-1) -class ZImageTransformer2DModel(CachableDiT): +class ZImageTransformer2DModel(CachableDiT, OffloadableDiTMixin): _supports_gradient_checkpointing = True _no_split_modules = ["ZImageTransformerBlock"] param_names_mapping = ZImageDitConfig().arch_config.param_names_mapping @@ -465,6 +466,7 @@ class ZImageTransformer2DModel(CachableDiT): self.rotary_emb = RopeEmbedder( theta=self.rope_theta, axes_dims=self.axes_dims, axes_lens=self.axes_lens ) + self.layer_names = ["layers"] def unpatchify( self, x: List[torch.Tensor], size: List[Tuple], patch_size, f_patch_size diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py index 8f71617aa..e126af0c2 100755 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -12,7 +12,7 @@ import time import weakref from collections.abc import Iterable from functools import lru_cache -from typing import Any, Optional +from typing import Any import torch from einops import rearrange @@ -61,9 +61,7 @@ from sglang.multimodal_gen.runtime.platforms import ( current_platform, ) from sglang.multimodal_gen.runtime.server_args import ServerArgs -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.runtime.utils.perf_logger import StageProfiler from sglang.multimodal_gen.runtime.utils.profiler import SGLDiffusionProfiler @@ -725,13 +723,10 @@ class DenoisingStage(PipelineStage): torch.mps.current_allocated_memory(), ) - # reset offload manager with prefetching first layer for next forward - offload_mgr: Optional[LayerwiseOffloadManager] = None - for transformer in filter(None, [self.transformer, self.transformer_2]): - if ( - offload_mgr := getattr(transformer, "_layerwise_offload_manager", None) - ) is not None: - offload_mgr.prepare_for_next_denoise(non_blocking=True) + # reset offload managers with prefetching first layer for next forward + for dit in filter(None, [self.transformer, self.transformer_2]): + if isinstance(dit, OffloadableDiTMixin): + dit.prepare_for_next_denoise() def _preprocess_sp_latents(self, batch: Req, server_args: ServerArgs): """Shard latents for Sequence Parallelism if applicable.""" diff --git a/python/sglang/multimodal_gen/runtime/utils/layerwise_offload.py b/python/sglang/multimodal_gen/runtime/utils/layerwise_offload.py index 75fdbf265..0cb14ce4d 100644 --- a/python/sglang/multimodal_gen/runtime/utils/layerwise_offload.py +++ b/python/sglang/multimodal_gen/runtime/utils/layerwise_offload.py @@ -1,10 +1,14 @@ import re -from contextlib import contextmanager from itertools import chain from typing import Any, Dict, List, Set, Tuple import torch +from sglang.multimodal_gen.runtime.server_args import ServerArgs +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +logger = init_logger(__name__) + # Adapted from skywork AI Infra diffusion optimize class LayerwiseOffloadManager: @@ -25,25 +29,24 @@ class LayerwiseOffloadManager: self, model: torch.nn.Module, *, - module_list_attr: str, + layers_attr_str: str, num_layers: int, enabled: bool, pin_cpu_memory: bool = True, - auto_initialize: bool = False, ) -> None: self.model = model - self.module_list_attr = module_list_attr + self.layers_attr_str = layers_attr_str self.num_layers = num_layers self.pin_cpu_memory = pin_cpu_memory self.enabled = bool(enabled and torch.cuda.is_available()) - self.device = ( - torch.device("cuda", torch.cuda.current_device()) if self.enabled else None - ) - self.copy_stream = torch.cuda.Stream() if self.enabled else None + if not self.enabled: + return + self.device = torch.device("cuda", torch.cuda.current_device()) + self.copy_stream = torch.cuda.Stream() self._layer_name_re = re.compile( - rf"(^|\.){re.escape(module_list_attr)}\.(\d+)(\.|$)" + rf"(^|\.){re.escape(layers_attr_str)}\.(\d+)(\.|$)" ) # layer_idx -> {dtype: consolidated_pinned_cpu_tensor} @@ -58,8 +61,7 @@ class LayerwiseOffloadManager: self._named_parameters: Dict[str, torch.nn.Parameter] = {} self._named_buffers: Dict[str, torch.Tensor] = {} - if auto_initialize: - self._initialize() + self._initialize() def _match_layer_idx(self, name: str) -> int | None: m = self._layer_name_re.search(name) @@ -125,6 +127,9 @@ class LayerwiseOffloadManager: # prefetch the first layer for warm-up self.prepare_for_next_denoise(non_blocking=False) + self.register_forward_hooks() + logger.info("LayerwiseOffloadManager initialized") + def prepare_for_next_denoise(self, non_blocking=True): self.prefetch_layer(0, non_blocking=non_blocking) if not non_blocking and self.copy_stream is not None: @@ -148,7 +153,6 @@ class LayerwiseOffloadManager: return if layer_idx not in self._consolidated_cpu_weights: return - self.copy_stream.wait_stream(torch.cuda.current_stream()) # create gpu buffer and load from CPU buffer @@ -174,30 +178,6 @@ class LayerwiseOffloadManager: self._gpu_layers.add(layer_idx) - @contextmanager - def layer_scope( - self, - *, - prefetch_layer_idx: int | None, - release_layer_idx: int | None, - non_blocking: bool = True, - ): - """A helper context manager to improve readability at call sites. - - It optionally prefetches ``prefetch_layer_idx`` before entering the - context, and waits for the copy stream then releases - ``release_layer_idx`` on exit. - """ - if self.enabled and prefetch_layer_idx is not None: - self.prefetch_layer(prefetch_layer_idx, non_blocking=non_blocking) - try: - yield - finally: - if self.enabled and self.copy_stream is not None: - torch.cuda.current_stream().wait_stream(self.copy_stream) - if self.enabled and release_layer_idx is not None: - self.release_layer(release_layer_idx) - @torch.compiler.disable def release_layer(self, layer_idx: int) -> None: if not self.enabled or self.device is None: @@ -223,3 +203,66 @@ class LayerwiseOffloadManager: for layer_idx in list(self._gpu_layers): self.release_layer(layer_idx) + + def register_forward_hooks(self) -> None: + if not self.enabled: + return + + layers = getattr(self.model, self.layers_attr_str) + + def make_pre_hook(i): + def hook(module, input): + self.prefetch_layer(i + 1, non_blocking=True) + + return hook + + def make_post_hook(i): + def hook(module, input, output): + if self.copy_stream is not None: + torch.cuda.current_stream().wait_stream(self.copy_stream) + self.release_layer(i) + + return hook + + # register prefetch & release hooks for each layer + for i, layer in enumerate(layers): + layer.register_forward_pre_hook(make_pre_hook(i)) + layer.register_forward_hook(make_post_hook(i)) + + +class OffloadableDiTMixin: + """ + A mixin that registers forward hooks for a DiT to enable layerwise offload + """ + + # the list of names of a DiT's layers/blocks + layer_names: List[str] + layerwise_offload_managers: list[LayerwiseOffloadManager] | None = None + + def configure_layerwise_offload(self, server_args: ServerArgs): + self.layerwise_offload_managers = [] + for layer_name in self.layer_names: + # a manager per layer-list + module_list = getattr(self, layer_name, None) + if module_list is None or not isinstance(module_list, torch.nn.ModuleList): + continue + + num_layers = len(module_list) + manager = LayerwiseOffloadManager( + model=self, + layers_attr_str=layer_name, + num_layers=num_layers, + enabled=True, + pin_cpu_memory=server_args.pin_cpu_memory, + ) + self.layerwise_offload_managers.append(manager) + + logger.info( + f"Enabled layerwise offload for {self.__class__.__name__} on modules: {self.layer_names}" + ) + + def prepare_for_next_denoise(self): + if self.layerwise_offload_managers is None: + return + for manager in self.layerwise_offload_managers: + manager.prepare_for_next_denoise(non_blocking=True)