Compile the module using the specified backend and kwargs. If a compiler_fn is set, it will be used instead of torch.compile().
(self,
backend=get_accelerator().get_compile_backend(),
compile_kwargs={},
schedule=None,
compiled_autograd_enabled=False)
| 5601 | return resolved_backend, schedule |
| 5602 | |
| 5603 | def compile(self, |
| 5604 | backend=get_accelerator().get_compile_backend(), |
| 5605 | compile_kwargs={}, |
| 5606 | schedule=None, |
| 5607 | compiled_autograd_enabled=False) -> None: |
| 5608 | """Compile the module using the specified backend and kwargs. |
| 5609 | If a compiler_fn is set, it will be used instead of torch.compile(). |
| 5610 | """ |
| 5611 | # Avoid graph breaks |
| 5612 | deepspeed.utils.nvtx.enable_nvtx = False |
| 5613 | |
| 5614 | if not is_compile_supported(): |
| 5615 | raise RuntimeError("compile is not supported in your version of PyTorch.") |
| 5616 | |
| 5617 | if self.is_compiled: |
| 5618 | return |
| 5619 | |
| 5620 | if 'backend' in compile_kwargs: |
| 5621 | logger.warning("The `backend` in `compile_kwargs` will be overridden. Use the `backend` argument instead.") |
| 5622 | |
| 5623 | logger.info(f"Compiling deepcompile={self.is_deepcompile_enabled()} backend={backend}") |
| 5624 | |
| 5625 | resolved_backend = None |
| 5626 | if self.is_deepcompile_enabled(): |
| 5627 | resolved_backend, schedule = self.get_deepspeed_compile_backend(backend, compile_kwargs, schedule) |
| 5628 | |
| 5629 | is_deepspeed_compile_backend = resolved_backend is not None |
| 5630 | |
| 5631 | # default to torch.compiler backend if deepspeed config validation fails |
| 5632 | backend = resolved_backend or backend |
| 5633 | |
| 5634 | # Hook state must align with whether DeepCompile is active. |
| 5635 | self._set_deepcompile_active(is_deepspeed_compile_backend) |
| 5636 | |
| 5637 | # create new dict to avoid modifying original dict |
| 5638 | try: |
| 5639 | self.module.compile(**{**compile_kwargs, 'backend': backend}) |
| 5640 | except BaseException: |
| 5641 | if is_deepspeed_compile_backend: |
| 5642 | # Restore default hooks if compilation fails before completing. |
| 5643 | self._set_deepcompile_active(False) |
| 5644 | raise |
| 5645 | |
| 5646 | self._is_compiled = True |
| 5647 | self._compile_kwargs = compile_kwargs |
| 5648 | if compiled_autograd_enabled: |
| 5649 | if not self._deepcompile_active: |
| 5650 | self._is_compiled_autograd_enabled = compiled_autograd_enabled |
| 5651 | else: |
| 5652 | logger.warning("Compiled autograd is not compatible with DeepCompile, disabling compiled autograd.") |
| 5653 | self._is_compiled_autograd_enabled = False |
| 5654 | |
| 5655 | def _set_deepcompile_active(self, active: bool) -> None: |
| 5656 | """Toggle DeepCompile runtime state and manage forward hooks accordingly.""" |