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