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

Method create_backward_hook

utils/offloading.py:231–255  ·  view source on GitHub ↗
(self, block_index: int)

Source from the content-addressed store, hash-verified

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
242 block_idx_to_cuda = self.blocks_to_swap - num_blocks_propagated
243 block_idx_to_wait = block_index - 1
244
245 def backward_hook(module, grad_input, grad_output):
246 if self.debug:
247 print(f"Backward hook for block {block_index}")
248
249 if swapping:
250 self._submit_move_blocks(block_idx_to_cpu, block_idx_to_cuda)
251 if waiting:
252 self._wait_blocks_move(block_idx_to_wait)
253 return None
254
255 return backward_hook
256
257 def prepare_block_devices_before_forward(self):
258 if self.blocks_to_swap is None or self.blocks_to_swap == 0:

Callers 1

__init__Method · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected