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)
| 5814 | return resolved_backend, schedule |
| 5815 | |
| 5816 | def compile(self, |
| 5817 | backend=get_accelerator().get_compile_backend(), |
| 5818 | compile_kwargs={}, |
| 5819 | schedule=None, |
| 5820 | compiled_autograd_enabled=False) -> None: |
| 5821 | """Compile the module using the specified backend and kwargs. |
| 5822 | If a compiler_fn is set, it will be used instead of torch.compile(). |
| 5823 | """ |
| 5824 | # Avoid graph breaks |
| 5825 | deepspeed.utils.nvtx.enable_nvtx = False |
| 5826 | |
| 5827 | if not is_compile_supported(): |
| 5828 | raise RuntimeError("compile is not supported in your version of PyTorch.") |
| 5829 | |
| 5830 | if self.is_compiled: |
| 5831 | return |
| 5832 | |
| 5833 | if 'backend' in compile_kwargs: |
| 5834 | logger.warning("The `backend` in `compile_kwargs` will be overridden. Use the `backend` argument instead.") |
| 5835 | |
| 5836 | logger.info(f"Compiling deepcompile={self.is_deepcompile_enabled()} backend={backend}") |
| 5837 | |
| 5838 | resolved_backend = None |
| 5839 | if self.is_deepcompile_enabled(): |
| 5840 | resolved_backend, schedule = self.get_deepspeed_compile_backend(backend, compile_kwargs, schedule) |
| 5841 | |
| 5842 | is_deepspeed_compile_backend = resolved_backend is not None |
| 5843 | |
| 5844 | # default to torch.compiler backend if deepspeed config validation fails |
| 5845 | backend = resolved_backend or backend |
| 5846 | |
| 5847 | # Hook state must align with whether DeepCompile is active. |
| 5848 | self._set_deepcompile_active(is_deepspeed_compile_backend) |
| 5849 | |
| 5850 | # create new dict to avoid modifying original dict |
| 5851 | try: |
| 5852 | self.module.compile(**{**compile_kwargs, 'backend': backend}) |
| 5853 | except BaseException: |
| 5854 | if is_deepspeed_compile_backend: |
| 5855 | # Restore default hooks if compilation fails before completing. |
| 5856 | self._set_deepcompile_active(False) |
| 5857 | raise |
| 5858 | |
| 5859 | self._is_compiled = True |
| 5860 | self._compile_kwargs = compile_kwargs |
| 5861 | if compiled_autograd_enabled: |
| 5862 | if not self._deepcompile_active: |
| 5863 | self._is_compiled_autograd_enabled = compiled_autograd_enabled |
| 5864 | else: |
| 5865 | logger.warning("Compiled autograd is not compatible with DeepCompile, disabling compiled autograd.") |
| 5866 | self._is_compiled_autograd_enabled = False |
| 5867 | |
| 5868 | def _set_deepcompile_active(self, active: bool) -> None: |
| 5869 | """Toggle DeepCompile runtime state and manage forward hooks accordingly.""" |