[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_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,46 +543,22 @@ 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)
|
for block in self.transformer_blocks:
|
||||||
if offload_mgr is not None and getattr(offload_mgr, "enabled", False):
|
encoder_hidden_states, hidden_states = block(
|
||||||
for i, block in enumerate(self.transformer_blocks):
|
hidden_states=hidden_states,
|
||||||
with offload_mgr.layer_scope(
|
encoder_hidden_states=encoder_hidden_states,
|
||||||
prefetch_layer_idx=i + 1,
|
temb=temb,
|
||||||
release_layer_idx=i,
|
freqs_cis=freqs_cis,
|
||||||
non_blocking=True,
|
joint_attention_kwargs=joint_attention_kwargs,
|
||||||
):
|
)
|
||||||
encoder_hidden_states, hidden_states = block(
|
for block in self.single_transformer_blocks:
|
||||||
hidden_states=hidden_states,
|
encoder_hidden_states, hidden_states = block(
|
||||||
encoder_hidden_states=encoder_hidden_states,
|
hidden_states=hidden_states,
|
||||||
temb=temb,
|
encoder_hidden_states=encoder_hidden_states,
|
||||||
freqs_cis=freqs_cis,
|
temb=temb,
|
||||||
joint_attention_kwargs=joint_attention_kwargs,
|
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,
|
|
||||||
)
|
|
||||||
|
|
||||||
hidden_states = self.norm_out(hidden_states, temb)
|
hidden_states = self.norm_out(hidden_states, temb)
|
||||||
|
|
||||||
|
|||||||
@@ -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,25 +808,10 @@ 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)
|
for block in self.blocks:
|
||||||
if offload_mgr is not None and getattr(offload_mgr, "enabled", False):
|
hidden_states = block(
|
||||||
for i, block in enumerate(self.blocks):
|
hidden_states, encoder_hidden_states, timestep_proj, freqs_cis
|
||||||
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
|
|
||||||
)
|
|
||||||
# if teacache is enabled, we need to cache the original hidden states
|
# if teacache is enabled, we need to cache the original hidden states
|
||||||
if enable_teacache:
|
if enable_teacache:
|
||||||
self.maybe_cache_states(hidden_states, original_hidden_states)
|
self.maybe_cache_states(hidden_states, original_hidden_states)
|
||||||
|
|||||||
@@ -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,8 +61,7 @@ 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:
|
||||||
m = self._layer_name_re.search(name)
|
m = self._layer_name_re.search(name)
|
||||||
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user