[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_diffusers_component_config,
get_hf_config, get_hf_config,
) )
from sglang.multimodal_gen.runtime.utils.layerwise_offload import ( from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin
LayerwiseOffloadManager,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
@@ -740,23 +738,14 @@ class TransformerLoader(ComponentLoader):
model = model.eval() model = model.eval()
if server_args.dit_layerwise_offload and hasattr(model, "dit_module_names"): if server_args.dit_layerwise_offload:
# TODO(will): support multiple module names # enable layerwise offload if possible
module_name = getattr(model, "dit_module_names", ["transformer_blocks"])[0] if isinstance(model, OffloadableDiTMixin):
try: model.configure_layerwise_offload(server_args)
num_layers = len(getattr(model, module_name)) else:
except Exception: logger.info(
num_layers = None "Disabling layerwise offload since current model does not support this feature"
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,
) )
setattr(model, "_layerwise_offload_manager", mgr)
return model return model
@@ -13,6 +13,8 @@ from torch.nn.attention.flex_attention import (
flex_attention, 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 # 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 # see https://github.com/pytorch/pytorch/issues/133254
# change to default for other models # change to default for other models
@@ -421,7 +423,7 @@ class CausalWanTransformerBlock(nn.Module):
return hidden_states return hidden_states
class CausalWanTransformer3DModel(BaseDiT): class CausalWanTransformer3DModel(BaseDiT, OffloadableDiTMixin):
_fsdp_shard_conditions = WanVideoConfig()._fsdp_shard_conditions _fsdp_shard_conditions = WanVideoConfig()._fsdp_shard_conditions
_compile_conditions = WanVideoConfig()._compile_conditions _compile_conditions = WanVideoConfig()._compile_conditions
_supported_attention_backends = WanVideoConfig()._supported_attention_backends _supported_attention_backends = WanVideoConfig()._supported_attention_backends
@@ -505,6 +507,10 @@ class CausalWanTransformer3DModel(BaseDiT):
self.__post_init__() self.__post_init__()
self.layer_names = [
"blocks",
]
@staticmethod @staticmethod
def _prepare_blockwise_causal_attn_mask( def _prepare_blockwise_causal_attn_mask(
device: torch.device | str, device: torch.device | str,
@@ -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.models.dits.base import CachableDiT
from sglang.multimodal_gen.runtime.platforms import current_platform 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 from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__) # pylint: disable=invalid-name 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() return freqs_cos.contiguous().float(), freqs_sin.contiguous().float()
class FluxTransformer2DModel(CachableDiT): class FluxTransformer2DModel(CachableDiT, OffloadableDiTMixin):
""" """
The Transformer model introduced in Flux. The Transformer model introduced in Flux.
@@ -426,10 +427,6 @@ class FluxTransformer2DModel(CachableDiT):
self.inner_dim = ( self.inner_dim = (
self.config.num_attention_heads * self.config.attention_head_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) self.rotary_emb = FluxPosEmbed(theta=10000, axes_dim=self.config.axes_dims_rope)
@@ -484,6 +481,11 @@ class FluxTransformer2DModel(CachableDiT):
gather_output=True, gather_output=True,
) )
self.layer_names = [
"transformer_blocks",
"single_transformer_blocks",
]
def forward( def forward(
self, self,
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
@@ -541,30 +543,6 @@ class FluxTransformer2DModel(CachableDiT):
ip_hidden_states = self.encoder_hid_proj(ip_adapter_image_embeds) ip_hidden_states = self.encoder_hid_proj(ip_adapter_image_embeds)
joint_attention_kwargs.update({"ip_hidden_states": ip_hidden_states}) 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: for block in self.transformer_blocks:
encoder_hidden_states, hidden_states = block( encoder_hidden_states, hidden_states = block(
hidden_states=hidden_states, hidden_states=hidden_states,
@@ -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.models.dits.base import CachableDiT
from sglang.multimodal_gen.runtime.platforms import current_platform 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 from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__) # pylint: disable=invalid-name 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() return freqs_cos.contiguous().float(), freqs_sin.contiguous().float()
class Flux2Transformer2DModel(CachableDiT): class Flux2Transformer2DModel(CachableDiT, OffloadableDiTMixin):
""" """
The Transformer model introduced in Flux 2. 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.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( def forward(
self, self,
@@ -40,6 +40,7 @@ from sglang.multimodal_gen.runtime.platforms import (
AttentionBackendEnum, AttentionBackendEnum,
current_platform, current_platform,
) )
from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin
class MMDoubleStreamBlock(nn.Module): class MMDoubleStreamBlock(nn.Module):
@@ -386,7 +387,7 @@ class MMSingleStreamBlock(nn.Module):
return self.output_residual(x, output, mod_gate) return self.output_residual(x, output, mod_gate)
class HunyuanVideoTransformer3DModel(CachableDiT): class HunyuanVideoTransformer3DModel(CachableDiT, OffloadableDiTMixin):
""" """
HunyuanVideo Transformer backbone adapted for distributed training. HunyuanVideo Transformer backbone adapted for distributed training.
@@ -508,7 +509,7 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
mlp_ratio=config.mlp_ratio, mlp_ratio=config.mlp_ratio,
dtype=config.dtype, dtype=config.dtype,
supported_attention_backends=self._supported_attention_backends, 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) for i in range(config.num_single_layers)
] ]
@@ -524,6 +525,8 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
self.__post_init__() self.__post_init__()
self.layer_names = ["double_blocks", "single_blocks"]
# TODO: change the input the FORWARD_BATCH Dict # TODO: change the input the FORWARD_BATCH Dict
# TODO: change output to a dict # TODO: change output to a dict
def forward( def forward(
@@ -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.models.dits.base import CachableDiT
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum 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 from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__) # pylint: disable=invalid-name logger = init_logger(__name__) # pylint: disable=invalid-name
@@ -798,7 +799,7 @@ def to_hashable(obj):
return obj return obj
class QwenImageTransformer2DModel(CachableDiT): class QwenImageTransformer2DModel(CachableDiT, OffloadableDiTMixin):
""" """
The Transformer model introduced in Qwen. The Transformer model introduced in Qwen.
@@ -878,6 +879,8 @@ class QwenImageTransformer2DModel(CachableDiT):
(1,), dtype=torch.int, device=get_local_torch_device() (1,), dtype=torch.int, device=get_local_torch_device()
) )
self.layer_names = ["transformer_blocks"]
@functools.lru_cache(maxsize=50) @functools.lru_cache(maxsize=50)
def build_modulate_index(self, img_shapes: tuple[int, int, int], device): def build_modulate_index(self, img_shapes: tuple[int, int, int], device):
modulate_index_list = [] modulate_index_list = []
@@ -43,6 +43,7 @@ from sglang.multimodal_gen.runtime.platforms import (
current_platform, current_platform,
) )
from sglang.multimodal_gen.runtime.server_args import get_global_server_args 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 from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__) logger = init_logger(__name__)
@@ -602,7 +603,7 @@ class WanTransformerBlock_VSA(nn.Module):
return hidden_states return hidden_states
class WanTransformer3DModel(CachableDiT): class WanTransformer3DModel(CachableDiT, OffloadableDiTMixin):
_fsdp_shard_conditions = WanVideoConfig()._fsdp_shard_conditions _fsdp_shard_conditions = WanVideoConfig()._fsdp_shard_conditions
_compile_conditions = WanVideoConfig()._compile_conditions _compile_conditions = WanVideoConfig()._compile_conditions
_supported_attention_backends = WanVideoConfig()._supported_attention_backends _supported_attention_backends = WanVideoConfig()._supported_attention_backends
@@ -621,7 +622,6 @@ class WanTransformer3DModel(CachableDiT):
self.num_channels_latents = config.num_channels_latents self.num_channels_latents = config.num_channels_latents
self.patch_size = config.patch_size self.patch_size = config.patch_size
self.text_len = config.text_len self.text_len = config.text_len
self.dit_module_names = ["blocks"]
# 1. Patch & position embedding # 1. Patch & position embedding
self.patch_embedding = PatchEmbed( self.patch_embedding = PatchEmbed(
@@ -708,6 +708,8 @@ class WanTransformer3DModel(CachableDiT):
dtype=torch.float32 if current_platform.is_mps() else torch.float64, dtype=torch.float32 if current_platform.is_mps() else torch.float64,
) )
self.layer_names = ["blocks"]
def forward( def forward(
self, self,
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
@@ -806,21 +808,6 @@ class WanTransformer3DModel(CachableDiT):
if enable_teacache: if enable_teacache:
original_hidden_states = hidden_states.clone() 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: for block in self.blocks:
hidden_states = block( hidden_states = block(
hidden_states, encoder_hidden_states, timestep_proj, freqs_cis hidden_states, encoder_hidden_states, timestep_proj, freqs_cis
@@ -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.layers.rotary_embedding import _apply_rotary_emb
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
from sglang.multimodal_gen.runtime.platforms import current_platform 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 from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__) logger = init_logger(__name__)
@@ -350,7 +351,7 @@ class RopeEmbedder:
return torch.cat(cos_out, dim=-1), torch.cat(sin_out, dim=-1) return torch.cat(cos_out, dim=-1), torch.cat(sin_out, dim=-1)
class ZImageTransformer2DModel(CachableDiT): class ZImageTransformer2DModel(CachableDiT, OffloadableDiTMixin):
_supports_gradient_checkpointing = True _supports_gradient_checkpointing = True
_no_split_modules = ["ZImageTransformerBlock"] _no_split_modules = ["ZImageTransformerBlock"]
param_names_mapping = ZImageDitConfig().arch_config.param_names_mapping param_names_mapping = ZImageDitConfig().arch_config.param_names_mapping
@@ -465,6 +466,7 @@ class ZImageTransformer2DModel(CachableDiT):
self.rotary_emb = RopeEmbedder( self.rotary_emb = RopeEmbedder(
theta=self.rope_theta, axes_dims=self.axes_dims, axes_lens=self.axes_lens theta=self.rope_theta, axes_dims=self.axes_dims, axes_lens=self.axes_lens
) )
self.layer_names = ["layers"]
def unpatchify( def unpatchify(
self, x: List[torch.Tensor], size: List[Tuple], patch_size, f_patch_size self, x: List[torch.Tensor], size: List[Tuple], patch_size, f_patch_size
@@ -12,7 +12,7 @@ import time
import weakref import weakref
from collections.abc import Iterable from collections.abc import Iterable
from functools import lru_cache from functools import lru_cache
from typing import Any, Optional from typing import Any
import torch import torch
from einops import rearrange from einops import rearrange
@@ -61,9 +61,7 @@ from sglang.multimodal_gen.runtime.platforms import (
current_platform, current_platform,
) )
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.layerwise_offload import ( from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin
LayerwiseOffloadManager,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger 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.perf_logger import StageProfiler
from sglang.multimodal_gen.runtime.utils.profiler import SGLDiffusionProfiler from sglang.multimodal_gen.runtime.utils.profiler import SGLDiffusionProfiler
@@ -725,13 +723,10 @@ class DenoisingStage(PipelineStage):
torch.mps.current_allocated_memory(), torch.mps.current_allocated_memory(),
) )
# reset offload manager with prefetching first layer for next forward # reset offload managers with prefetching first layer for next forward
offload_mgr: Optional[LayerwiseOffloadManager] = None for dit in filter(None, [self.transformer, self.transformer_2]):
for transformer in filter(None, [self.transformer, self.transformer_2]): if isinstance(dit, OffloadableDiTMixin):
if ( dit.prepare_for_next_denoise()
offload_mgr := getattr(transformer, "_layerwise_offload_manager", None)
) is not None:
offload_mgr.prepare_for_next_denoise(non_blocking=True)
def _preprocess_sp_latents(self, batch: Req, server_args: ServerArgs): def _preprocess_sp_latents(self, batch: Req, server_args: ServerArgs):
"""Shard latents for Sequence Parallelism if applicable.""" """Shard latents for Sequence Parallelism if applicable."""
@@ -1,10 +1,14 @@
import re import re
from contextlib import contextmanager
from itertools import chain from itertools import chain
from typing import Any, Dict, List, Set, Tuple from typing import Any, Dict, List, Set, Tuple
import torch 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 # Adapted from skywork AI Infra diffusion optimize
class LayerwiseOffloadManager: class LayerwiseOffloadManager:
@@ -25,25 +29,24 @@ class LayerwiseOffloadManager:
self, self,
model: torch.nn.Module, model: torch.nn.Module,
*, *,
module_list_attr: str, layers_attr_str: str,
num_layers: int, num_layers: int,
enabled: bool, enabled: bool,
pin_cpu_memory: bool = True, pin_cpu_memory: bool = True,
auto_initialize: bool = False,
) -> None: ) -> None:
self.model = model self.model = model
self.module_list_attr = module_list_attr self.layers_attr_str = layers_attr_str
self.num_layers = num_layers self.num_layers = num_layers
self.pin_cpu_memory = pin_cpu_memory self.pin_cpu_memory = pin_cpu_memory
self.enabled = bool(enabled and torch.cuda.is_available()) self.enabled = bool(enabled and torch.cuda.is_available())
self.device = ( if not self.enabled:
torch.device("cuda", torch.cuda.current_device()) if self.enabled else None return
) self.device = torch.device("cuda", torch.cuda.current_device())
self.copy_stream = torch.cuda.Stream() if self.enabled else None self.copy_stream = torch.cuda.Stream()
self._layer_name_re = re.compile( 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} # layer_idx -> {dtype: consolidated_pinned_cpu_tensor}
@@ -58,7 +61,6 @@ class LayerwiseOffloadManager:
self._named_parameters: Dict[str, torch.nn.Parameter] = {} self._named_parameters: Dict[str, torch.nn.Parameter] = {}
self._named_buffers: Dict[str, torch.Tensor] = {} self._named_buffers: Dict[str, torch.Tensor] = {}
if auto_initialize:
self._initialize() self._initialize()
def _match_layer_idx(self, name: str) -> int | None: def _match_layer_idx(self, name: str) -> int | None:
@@ -125,6 +127,9 @@ class LayerwiseOffloadManager:
# prefetch the first layer for warm-up # prefetch the first layer for warm-up
self.prepare_for_next_denoise(non_blocking=False) 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): def prepare_for_next_denoise(self, non_blocking=True):
self.prefetch_layer(0, non_blocking=non_blocking) self.prefetch_layer(0, non_blocking=non_blocking)
if not non_blocking and self.copy_stream is not None: if not non_blocking and self.copy_stream is not None:
@@ -148,7 +153,6 @@ class LayerwiseOffloadManager:
return return
if layer_idx not in self._consolidated_cpu_weights: if layer_idx not in self._consolidated_cpu_weights:
return return
self.copy_stream.wait_stream(torch.cuda.current_stream()) self.copy_stream.wait_stream(torch.cuda.current_stream())
# create gpu buffer and load from CPU buffer # create gpu buffer and load from CPU buffer
@@ -174,30 +178,6 @@ class LayerwiseOffloadManager:
self._gpu_layers.add(layer_idx) 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 @torch.compiler.disable
def release_layer(self, layer_idx: int) -> None: def release_layer(self, layer_idx: int) -> None:
if not self.enabled or self.device is None: if not self.enabled or self.device is None:
@@ -223,3 +203,66 @@ class LayerwiseOffloadManager:
for layer_idx in list(self._gpu_layers): for layer_idx in list(self._gpu_layers):
self.release_layer(layer_idx) 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)