| 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: |