MCPcopy Create free account
hub / github.com/tdrussell/diffusion-pipe / ModelOffloader

Class ModelOffloader

utils/offloading.py:184–301  ·  view source on GitHub ↗

supports forward offloading

Source from the content-addressed store, hash-verified

182
183
184class ModelOffloader(Offloader):
185 """
186 supports forward offloading
187 """
188
189 def __init__(
190 self,
191 block_type: str,
192 blocks: list[nn.Module],
193 num_blocks: int,
194 blocks_to_swap: int,
195 supports_backward: bool,
196 device: torch.device,
197 reentrant_activation_checkpointing: bool,
198 debug: bool = False,
199 ):
200 super().__init__(block_type, blocks, num_blocks, blocks_to_swap, device, debug)
201
202 self.supports_backward = supports_backward
203 self.forward_only = not supports_backward # forward only offloading: can be changed to True for inference
204 self.reentrant_activation_checkpointing = reentrant_activation_checkpointing
205
206 if self.supports_backward:
207 # register backward hooks
208 self.remove_handles = []
209 for i, block in enumerate(blocks):
210 hook = self.create_backward_hook(i)
211 if hook is not None:
212 handle = block.register_full_backward_hook(hook)
213 self.remove_handles.append(handle)
214
215 def disable_block_swap(self):
216 self.blocks_to_swap_tmp = self.blocks_to_swap
217 self.blocks_to_swap = None
218
219 def enable_block_swap(self):
220 if self.blocks_to_swap_tmp is not None:
221 self.blocks_to_swap = self.blocks_to_swap_tmp
222
223 def set_forward_only(self, forward_only: bool):
224 self.forward_only = forward_only
225
226 def __del__(self):
227 if self.supports_backward:
228 for handle in self.remove_handles:
229 handle.remove()
230
231 def create_backward_hook(self, block_index: int) -> Optional[callable]:
232 # -1 for 0-based index
233 num_blocks_propagated = self.num_blocks - block_index - 1
234 swapping = num_blocks_propagated > 0 and num_blocks_propagated <= self.blocks_to_swap
235 waiting = block_index > 0 and block_index <= self.blocks_to_swap
236
237 if not swapping and not waiting:
238 return None
239
240 # create hook
241 block_idx_to_cpu = self.num_blocks - num_blocks_propagated

Callers 15

__init__Method · 0.90
enable_block_swapMethod · 0.90
__init__Method · 0.90
enable_block_swapMethod · 0.90
__init__Method · 0.90
enable_block_swapMethod · 0.90
__init__Method · 0.90
enable_block_swapMethod · 0.90
__init__Method · 0.90
enable_block_swapMethod · 0.90
__init__Method · 0.90
enable_block_swapMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected