supports forward offloading
| 182 | |
| 183 | |
| 184 | class 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 |
no outgoing calls
no test coverage detected