DeepSpeed
016fb769 - Skip ZeRO DeepCompile backward hooks for AutoSP (#8667)

Commit
13 days ago
Skip ZeRO DeepCompile backward hooks for AutoSP (#8667) ## What breaks `DeepSpeedEngine._backward_prologue()` / `_backward_epilogue()` enter the ZeRO DeepCompile backward hooks whenever DeepCompile is active, exempting only the AutoTP compile pass: ```python if self.is_deepcompile_active() and not self.compile_autotp(): deepcompile_backward_prologue(self.is_gradient_accumulation_boundary()) ``` With `compile.passes: ["autosp"]`, `engine.compile()` installs the AutoSP backend and marks DeepCompile active, so every backward calls `deepcompile_backward_prologue()`. That calls `get_deepcompile_handle()`, which loads (and JIT-builds on first use) the DeepCompile native extension, then calls `start_backward()` on it. AutoSP never initializes that runtime (`dc.init()` is only called on the ZeRO DeepCompile path), so on CUDA this silently loads an unused extension, and on accelerators that cannot build it, backward fails. `allreduce_gradients()` already uses `uses_parallelization_pass_only()` (AutoSP or AutoTP) for the same decision; the backward hooks were missed when that helper was introduced. ## Fix Use `uses_parallelization_pass_only()` in both the backward prologue and epilogue checks. The AutoTP behavior is unchanged. ## Tests Added `TestAutoSPEngineBackward::test_backward_skips_zero_deepcompile_hooks` in `tests/unit/v1/compile/test_compile_autosp.py`. It builds a ZeRO-0 engine with `compile.passes: ["autosp"]`, marks DeepCompile active as `engine.compile()` would, makes `get_deepcompile_handle()` and the post-backward hooks raise, and runs forward/backward/step. | | result | |---|---| | master | FAILED (`RuntimeError: AutoSP backward must not load the ZeRO DeepCompile runtime`) | | this PR | PASSED | Also ran locally: `tests/unit/compile/test_zero3_grad_dtype.py`, `tests/unit/compile/test_backend.py`, and the rest of `tests/unit/v1/compile/test_compile_autosp.py` (all IR-level tests pass). The end-to-end `TestAutoSPCompile` cases fail identically on master and this branch in my environment (torch 2.13, transformers 5.17) during Dynamo tracing of the forward pass, before backward is reached, so they are unrelated to this change. Signed-off-by: Jin, Youzhi <youzhi.jin@intel.com> Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Author
Parents
Loading