[diffusion] feat: add an arg for controlling the number of prefetched layers in layerwise-offload (#17693)
This commit is contained in:
@@ -33,12 +33,13 @@ class LayerwiseOffloadManager:
|
||||
num_layers: int,
|
||||
enabled: bool,
|
||||
pin_cpu_memory: bool = True,
|
||||
prefetch_size: int = 1,
|
||||
) -> None:
|
||||
self.model = model
|
||||
self.layers_attr_str = layers_attr_str
|
||||
self.num_layers = num_layers
|
||||
self.pin_cpu_memory = pin_cpu_memory
|
||||
|
||||
self.prefetch_size = min(max(1, prefetch_size), self.num_layers)
|
||||
self.enabled = bool(enabled and torch.cuda.is_available())
|
||||
if not self.enabled:
|
||||
return
|
||||
@@ -57,6 +58,8 @@ class LayerwiseOffloadManager:
|
||||
self._weight_metadata: Dict[int, Dict[str, Dict[str, Any]]] = {}
|
||||
# layer indices that are already in gpu
|
||||
self._gpu_layers: Set[int] = set()
|
||||
# layer_idx -> torch.cuda.Event for fine-grained sync, to make sure the weight is resident in pre-hook
|
||||
self._prefetch_events: Dict[int, torch.cuda.Event] = {}
|
||||
|
||||
self._named_parameters: Dict[str, torch.nn.Parameter] = {}
|
||||
self._named_buffers: Dict[str, torch.Tensor] = {}
|
||||
@@ -127,13 +130,19 @@ class LayerwiseOffloadManager:
|
||||
self._consolidated_cpu_weights[layer_idx][dtype] = cpu_buffer
|
||||
|
||||
# prefetch the first layer for warm-up
|
||||
self.prepare_for_next_denoise(non_blocking=False)
|
||||
self.prepare_for_next_req(non_blocking=False)
|
||||
|
||||
self.register_forward_hooks()
|
||||
logger.info("LayerwiseOffloadManager initialized")
|
||||
logger.info(
|
||||
f"LayerwiseOffloadManager initialized with num prefetched layer: {self.prefetch_size}, total num layers: {self.num_layers}"
|
||||
)
|
||||
|
||||
def prepare_for_next_denoise(self, non_blocking=True):
|
||||
self.prefetch_layer(0, non_blocking=non_blocking)
|
||||
def prepare_for_next_req(self, non_blocking=True):
|
||||
"""
|
||||
Prepare for the next round of denoising loop with prefetching the necessary layers
|
||||
"""
|
||||
for i in range(self.prefetch_size):
|
||||
self.prefetch_layer(i, non_blocking=non_blocking)
|
||||
if not non_blocking and self.copy_stream is not None:
|
||||
torch.cuda.current_stream().wait_stream(self.copy_stream)
|
||||
|
||||
@@ -147,6 +156,9 @@ class LayerwiseOffloadManager:
|
||||
|
||||
@torch.compiler.disable
|
||||
def prefetch_layer(self, layer_idx: int, non_blocking: bool = True) -> None:
|
||||
"""
|
||||
idempotent
|
||||
"""
|
||||
if not self.enabled or self.device is None or self.copy_stream is None:
|
||||
return
|
||||
if layer_idx < 0 or layer_idx >= self.num_layers:
|
||||
@@ -167,6 +179,11 @@ class LayerwiseOffloadManager:
|
||||
gpu_buffer.copy_(cpu_buffer, non_blocking=non_blocking)
|
||||
gpu_buffers[dtype] = gpu_buffer
|
||||
|
||||
# record the prefetch event of this layer
|
||||
event = torch.cuda.Event()
|
||||
event.record(self.copy_stream)
|
||||
self._prefetch_events[layer_idx] = event
|
||||
|
||||
# restore model's weights by their metadata using gpu buffer
|
||||
for name, meta in self._weight_metadata[layer_idx].items():
|
||||
dtype = meta["dtype"]
|
||||
@@ -182,8 +199,16 @@ class LayerwiseOffloadManager:
|
||||
|
||||
@torch.compiler.disable
|
||||
def release_layer(self, layer_idx: int) -> None:
|
||||
"""
|
||||
lightweight release layer weights
|
||||
Basically set the reference count to the gpu weight tensor to zero. The weights on cpu is untouched
|
||||
"""
|
||||
if not self.enabled or self.device is None:
|
||||
return
|
||||
|
||||
# clear prefetch event, since it's useless and needs to be reset
|
||||
self._prefetch_events.pop(layer_idx, None)
|
||||
|
||||
if layer_idx <= 0:
|
||||
return
|
||||
|
||||
@@ -259,14 +284,24 @@ class LayerwiseOffloadManager:
|
||||
|
||||
def make_pre_hook(i):
|
||||
def hook(module, input):
|
||||
self.prefetch_layer(i + 1, non_blocking=True)
|
||||
# wait only for the current layer if it's being prefetched
|
||||
if i == 0:
|
||||
self.prepare_for_next_req(non_blocking=False)
|
||||
if i in self._prefetch_events:
|
||||
torch.cuda.current_stream().wait_event(self._prefetch_events[i])
|
||||
|
||||
# trigger batch prefetch (i + prefetch_size ~ i + 2 * prefetch_size) if needed
|
||||
if i % self.prefetch_size == 0:
|
||||
for j in range(i + self.prefetch_size, i + 2 * self.prefetch_size):
|
||||
layer_to_prefetch = j % self.num_layers
|
||||
self.prefetch_layer(layer_to_prefetch, 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)
|
||||
# previous, we wait here, until the copy stream for next layer is finished,
|
||||
# now with any prefetch_size, only wait for the copy stream, when the copy stream is for the next layer
|
||||
self.release_layer(i)
|
||||
|
||||
return hook
|
||||
@@ -292,7 +327,7 @@ class OffloadableDiTMixin:
|
||||
|
||||
# the list of names of a DiT's layers/blocks
|
||||
layer_names: List[str]
|
||||
layerwise_offload_managers: list[LayerwiseOffloadManager] | None = None
|
||||
layerwise_offload_managers: list[LayerwiseOffloadManager] = []
|
||||
|
||||
def configure_layerwise_offload(self, server_args: ServerArgs):
|
||||
self.layerwise_offload_managers = []
|
||||
@@ -303,12 +338,20 @@ class OffloadableDiTMixin:
|
||||
continue
|
||||
|
||||
num_layers = len(module_list)
|
||||
if server_args.dit_offload_prefetch_size < 1.0:
|
||||
prefetch_size = 1 + int(
|
||||
round(server_args.dit_offload_prefetch_size * (num_layers - 1))
|
||||
)
|
||||
else:
|
||||
prefetch_size = int(server_args.dit_offload_prefetch_size)
|
||||
|
||||
manager = LayerwiseOffloadManager(
|
||||
model=self,
|
||||
layers_attr_str=layer_name,
|
||||
num_layers=num_layers,
|
||||
enabled=True,
|
||||
pin_cpu_memory=server_args.pin_cpu_memory,
|
||||
prefetch_size=prefetch_size,
|
||||
)
|
||||
self.layerwise_offload_managers.append(manager)
|
||||
|
||||
@@ -316,11 +359,11 @@ class OffloadableDiTMixin:
|
||||
f"Enabled layerwise offload for {self.__class__.__name__} on modules: {self.layer_names}"
|
||||
)
|
||||
|
||||
def prepare_for_next_denoise(self):
|
||||
def prepare_for_next_req(self):
|
||||
if self.layerwise_offload_managers is None:
|
||||
return
|
||||
for manager in self.layerwise_offload_managers:
|
||||
manager.prepare_for_next_denoise(non_blocking=True)
|
||||
manager.prepare_for_next_req(non_blocking=True)
|
||||
|
||||
def disable_offload(self) -> None:
|
||||
"""Disable layerwise offload: load all layers to GPU and remove hooks."""
|
||||
@@ -338,7 +381,5 @@ class OffloadableDiTMixin:
|
||||
for manager in self.layerwise_offload_managers:
|
||||
if manager.enabled:
|
||||
manager.sync_all_layers_to_cpu()
|
||||
for layer_idx in list(manager._gpu_layers):
|
||||
if layer_idx > 0:
|
||||
manager.release_layer(layer_idx)
|
||||
manager.release_all()
|
||||
manager.register_forward_hooks()
|
||||
|
||||
Reference in New Issue
Block a user