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)
| 5470 | return resolved_backend, schedule |
| 5471 | |
| 5472 | def compile(self, |
| 5473 | backend=get_accelerator().get_compile_backend(), |
| 5474 | compile_kwargs={}, |
| 5475 | schedule=None, |
| 5476 | compiled_autograd_enabled=False) -> None: |
| 5477 | """Compile the module using the specified backend and kwargs. |
| 5478 | If a compiler_fn is set, it will be used instead of torch.compile(). |
| 5479 | """ |
| 5480 | # Avoid graph breaks |
| 5481 | deepspeed.utils.nvtx.enable_nvtx = False |
| 5482 | |
| 5483 | if not is_compile_supported(): |
| 5484 | raise RuntimeError("compile is not supported in your version of PyTorch.") |
| 5485 | |
| 5486 | if self.is_compiled: |
| 5487 | return |
| 5488 | |
| 5489 | if 'backend' in compile_kwargs: |
| 5490 | logger.warning("The `backend` in `compile_kwargs` will be overridden. Use the `backend` argument instead.") |
| 5491 | |
| 5492 | logger.info(f"Compiling deepcompile={self.is_deepcompile_enabled()} backend={backend}") |
| 5493 | |
| 5494 | resolved_backend = None |
| 5495 | if self.is_deepcompile_enabled(): |
| 5496 | resolved_backend, schedule = self.get_deepspeed_compile_backend(backend, compile_kwargs, schedule) |
| 5497 | |
| 5498 | is_deepspeed_compile_backend = resolved_backend is not None |
| 5499 | |
| 5500 | # default to torch.compiler backend if deepspeed config validation fails |
| 5501 | backend = resolved_backend or backend |
| 5502 | |
| 5503 | # Hook state must align with whether DeepCompile is active. |
| 5504 | self._set_deepcompile_active(is_deepspeed_compile_backend) |
| 5505 | |
| 5506 | # create new dict to avoid modifying original dict |
| 5507 | try: |
| 5508 | self.module.compile(**{**compile_kwargs, 'backend': backend}) |
| 5509 | except BaseException: |
| 5510 | if is_deepspeed_compile_backend: |
| 5511 | # Restore default hooks if compilation fails before completing. |
| 5512 | self._set_deepcompile_active(False) |
| 5513 | raise |
| 5514 | |
| 5515 | self._is_compiled = True |
| 5516 | self._compile_kwargs = compile_kwargs |
| 5517 | if compiled_autograd_enabled: |
| 5518 | if not self._deepcompile_active: |
| 5519 | self._is_compiled_autograd_enabled = compiled_autograd_enabled |
| 5520 | else: |
| 5521 | logger.warning("Compiled autograd is not compatible with DeepCompile, disabling compiled autograd.") |
| 5522 | self._is_compiled_autograd_enabled = False |
| 5523 | |
| 5524 | def _set_deepcompile_active(self, active: bool) -> None: |
| 5525 | """Toggle DeepCompile runtime state and manage forward hooks accordingly.""" |