[diffusion] feat: generalize layerwise offloader to flux1 (#15633)
This commit is contained in:
@@ -712,16 +712,17 @@ class TransformerLoader(ComponentLoader):
|
|||||||
|
|
||||||
model = model.eval()
|
model = model.eval()
|
||||||
|
|
||||||
if server_args.dit_layerwise_offload and hasattr(model, "blocks"):
|
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:
|
try:
|
||||||
num_layers = len(getattr(model, "blocks"))
|
num_layers = len(getattr(model, module_name))
|
||||||
except Exception:
|
except Exception:
|
||||||
num_layers = None
|
num_layers = None
|
||||||
|
|
||||||
if isinstance(num_layers, int) and num_layers > 0:
|
if isinstance(num_layers, int) and num_layers > 0:
|
||||||
mgr = LayerwiseOffloadManager(
|
mgr = LayerwiseOffloadManager(
|
||||||
model,
|
model,
|
||||||
module_list_attr="blocks",
|
module_list_attr=module_name,
|
||||||
num_layers=num_layers,
|
num_layers=num_layers,
|
||||||
enabled=True,
|
enabled=True,
|
||||||
pin_cpu_memory=server_args.pin_cpu_memory,
|
pin_cpu_memory=server_args.pin_cpu_memory,
|
||||||
|
|||||||
@@ -391,6 +391,10 @@ 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)
|
||||||
|
|
||||||
@@ -496,7 +500,14 @@ 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})
|
||||||
|
|
||||||
for index_block, block in enumerate(self.transformer_blocks):
|
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(
|
encoder_hidden_states, hidden_states = block(
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
encoder_hidden_states=encoder_hidden_states,
|
encoder_hidden_states=encoder_hidden_states,
|
||||||
@@ -504,8 +515,24 @@ class FluxTransformer2DModel(CachableDiT):
|
|||||||
freqs_cis=freqs_cis,
|
freqs_cis=freqs_cis,
|
||||||
joint_attention_kwargs=joint_attention_kwargs,
|
joint_attention_kwargs=joint_attention_kwargs,
|
||||||
)
|
)
|
||||||
|
for block in self.single_transformer_blocks:
|
||||||
for index_block, block in enumerate(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(
|
encoder_hidden_states, hidden_states = block(
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
encoder_hidden_states=encoder_hidden_states,
|
encoder_hidden_states=encoder_hidden_states,
|
||||||
|
|||||||
@@ -610,6 +610,7 @@ 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(
|
||||||
|
|||||||
Reference in New Issue
Block a user