Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions deepspeed/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -258,6 +258,8 @@ def initialize(
# Restore zero.Init context if necessary
zero.partition_parameters.restore_init_context()

engine._configure_python_gc()

return_items = [
engine,
engine.optimizer,
Expand Down
6 changes: 6 additions & 0 deletions deepspeed/module_inject/auto_ep_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,7 @@ def parse_autoep_config(param_dict: dict) -> AutoEPConfig:
config.comm_num_sm = param_dict.get("comm_num_sm", 12)
config.comm_qp_margin = param_dict.get("comm_qp_margin", 4)
config.comm_max_tokens_per_rank = param_dict.get("comm_max_tokens_per_rank", 0)
config.python_gc_policy = param_dict.get("python_gc_policy", "default")
config.num_expert_groups = param_dict.get("num_expert_groups", None)
config.num_limited_groups = param_dict.get("num_limited_groups", None)
config.score_func = param_dict.get("score_func", "auto")
Expand Down Expand Up @@ -119,6 +120,11 @@ def validate_autoep_config(
if not isinstance(config.validate_folding_routing, bool):
raise ValueError("expert_parallel.validate_folding_routing must be a boolean")

valid_python_gc_policies = ("default", "disable_during_training")
if config.python_gc_policy not in valid_python_gc_policies:
raise ValueError(f"python_gc_policy must be one of {valid_python_gc_policies}, "
f"got {config.python_gc_policy!r}")

if not config.enabled:
return

Expand Down
1 change: 1 addition & 0 deletions deepspeed/module_inject/auto_ep_presets/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -115,6 +115,7 @@ class AutoEPConfig:
comm_num_sm: int = 12
comm_qp_margin: int = 4
comm_max_tokens_per_rank: int = 0
python_gc_policy: Literal["default", "disable_during_training"] = "default"
num_expert_groups: int | None = None
num_limited_groups: int | None = None
score_func: Literal["auto", "softmax", "sigmoid"] = "auto"
Expand Down
59 changes: 39 additions & 20 deletions deepspeed/runtime/engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -394,6 +394,7 @@ def __init__(self,
self.mesh_device = mesh_device
self._autoep_folding_spec = None
self._autoep_folding_group_handles = None
self._python_gc_generation = None

# Flag to indicate that scale() was called before manual backward pass
self._manual_backward_expected = False
Expand Down Expand Up @@ -932,26 +933,44 @@ def __del__(self):
logger.debug("DeepSpeedEngine.__del__ cleanup skipped: %s", exc, exc_info=True)

def destroy(self):
# DeepEP buffers ask the library not to reclaim them, so they outlive
# the engine unless something releases them here. Only this engine's
# own buffers: another engine in the same process still needs its own.
module = getattr(self, "module", None)
if module is not None:
from deepspeed.module_inject.auto_ep_comm import destroy_exchanges
destroy_exchanges(module)

self._release_deepcompile_compiled_backward_state()
self._release_deepcompile_dynamo_config()
optimizer = getattr(self, "optimizer", None)
if optimizer is not None and hasattr(optimizer, 'destroy'):
optimizer.destroy()
if self.is_deepcompile_active():
get_deepcompile_handle().cleanup()
debug_clear_module_and_param_names()

checkpoint_engine = getattr(self, "checkpoint_engine", None)
if checkpoint_engine is not None and checkpoint_engine.is_decoupled():
checkpoint_engine.cleanup()
try:
# DeepEP buffers ask the library not to reclaim them, so they outlive
# the engine unless something releases them here. Only this engine's
# own buffers: another engine in the same process still needs its own.
module = getattr(self, "module", None)
if module is not None:
from deepspeed.module_inject.auto_ep_comm import destroy_exchanges
destroy_exchanges(module)

self._release_deepcompile_compiled_backward_state()
self._release_deepcompile_dynamo_config()
optimizer = getattr(self, "optimizer", None)
if optimizer is not None and hasattr(optimizer, 'destroy'):
optimizer.destroy()
if self.is_deepcompile_active():
get_deepcompile_handle().cleanup()
debug_clear_module_and_param_names()

checkpoint_engine = getattr(self, "checkpoint_engine", None)
if checkpoint_engine is not None and checkpoint_engine.is_decoupled():
checkpoint_engine.cleanup()
finally:
python_gc_generation = getattr(self, "_python_gc_generation", None)
if python_gc_generation is not None:
from deepspeed.runtime.python_gc import python_gc_manager
python_gc_manager.release(python_gc_generation)
self._python_gc_generation = None

def collect_python_gc(self):
"""Run Python cyclic GC at an application-selected safe boundary."""
from deepspeed.runtime.python_gc import python_gc_manager
return python_gc_manager.collect()

def _configure_python_gc(self):
autoep_config = self._config.expert_parallel_config
if autoep_config.enabled and autoep_config.python_gc_policy == "disable_during_training":
from deepspeed.runtime.python_gc import python_gc_manager
self._python_gc_generation = python_gc_manager.acquire()

def _get_model_parameters(self):
if self.autotuning_profile_model_info():
Expand Down
75 changes: 75 additions & 0 deletions deepspeed/runtime/python_gc.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,75 @@
# Copyright (c) Microsoft Corporation.
# SPDX-License-Identifier: Apache-2.0
# DeepSpeed Team
"""Process-wide Python cyclic garbage collection control."""

import gc
import os
from threading import RLock

from deepspeed.utils import logger


class PythonGCManager:
"""Coordinate process-wide GC state across multiple DeepSpeed engines."""

def __init__(self):
self._lock = RLock()
self._active_engines = 0
self._restore_enabled = False
self._generation = 0
if hasattr(os, "register_at_fork"):
os.register_at_fork(before=self._before_fork,
after_in_parent=self._after_fork_parent,
after_in_child=self._after_fork_child)

def _before_fork(self):
self._lock.acquire()

def _after_fork_parent(self):
self._lock.release()

def _after_fork_child(self):
if self._active_engines > 0 and self._restore_enabled and not gc.isenabled():
gc.enable()
self._active_engines = 0
self._restore_enabled = False
self._generation += 1
self._lock = RLock()

def acquire(self):
with self._lock:
if self._active_engines == 0:
self._restore_enabled = gc.isenabled()
if self._restore_enabled:
collected = gc.collect()
gc.disable()
logger.info("Disabled automatic Python cyclic GC after collecting %d objects", collected)
self._active_engines += 1
return self._generation

def release(self, generation):
with self._lock:
if generation != self._generation:
return
if self._active_engines == 0:
return
self._active_engines -= 1
restore_enabled = self._active_engines == 0 and self._restore_enabled
if self._active_engines == 0:
self._restore_enabled = False
if restore_enabled and not gc.isenabled():
gc.enable()
logger.info("Restored automatic Python cyclic GC")

def collect(self):
"""Run an explicit collection without changing the automatic-GC policy."""
return gc.collect()

@property
def active_engines(self):
with self._lock:
return self._active_engines


python_gc_manager = PythonGCManager()
26 changes: 26 additions & 0 deletions docs/code-docs/source/autoep.rst
Original file line number Diff line number Diff line change
Expand Up @@ -112,6 +112,32 @@ that set nothing keep the existing path unchanged.
DeepEP buffer is sized statically and must use the same capacity on every
rank. A batch that exceeds it is an error.

**Python cyclic GC policy (experimental):**

Large Python model graphs can accumulate cyclic objects during training. A
generation-2 collection pauses one rank's Python thread, and the pause can then
be exposed as collective wait time on every expert-parallel rank. AutoEP offers
an opt-in policy that collects once after engine initialization and disables
automatic cyclic collection until the engine is destroyed:

.. code-block:: json

{
"expert_parallel": {
"enabled": true,
"autoep_size": 8,
"python_gc_policy": "disable_during_training"
}
}

The default is ``"default"``, which leaves Python GC unchanged. The policy is
process-wide and reference-counted across DeepSpeed engines. Applications that
create cyclic Python objects during training should call
``engine.collect_python_gc()`` at a safe boundary such as after checkpointing.
Call ``engine.destroy()`` when the engine is no longer needed to restore the
process's original automatic-GC state; restoration does not rely on Python
finalization because disabled cyclic GC cannot reclaim engine reference cycles.

On 16 H100s across two nodes, replaying routing captured from real training,
DeepEP reduced payload AllToAll time from roughly 100 ms to 48 ms per step. A
full SFT step on Qwen3.5-MoE went from roughly 325 ms to 266 ms, a 1.2x speedup
Expand Down
159 changes: 159 additions & 0 deletions tests/unit/runtime/test_python_gc.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,159 @@
# Copyright (c) Microsoft Corporation.
# SPDX-License-Identifier: Apache-2.0
# DeepSpeed Team

from deepspeed.runtime.python_gc import PythonGCManager


def test_python_gc_manager_restores_enabled_state(monkeypatch):
calls = []
enabled = True

def isenabled():
return enabled

def collect():
calls.append("collect")
return 7

def disable():
nonlocal enabled
enabled = False
calls.append("disable")

def enable():
nonlocal enabled
enabled = True
calls.append("enable")

monkeypatch.setattr("deepspeed.runtime.python_gc.gc.isenabled", isenabled)
monkeypatch.setattr("deepspeed.runtime.python_gc.gc.collect", collect)
monkeypatch.setattr("deepspeed.runtime.python_gc.gc.disable", disable)
monkeypatch.setattr("deepspeed.runtime.python_gc.gc.enable", enable)

manager = PythonGCManager()
generation = manager.acquire()
assert manager.acquire() == generation
assert manager.active_engines == 2
assert calls == ["collect", "disable"]

manager.release(generation)
assert manager.active_engines == 1
assert calls == ["collect", "disable"]

manager.release(generation)
assert manager.active_engines == 0
assert calls == ["collect", "disable", "enable"]


def test_python_gc_manager_preserves_disabled_state(monkeypatch):
calls = []

monkeypatch.setattr("deepspeed.runtime.python_gc.gc.isenabled", lambda: False)
monkeypatch.setattr("deepspeed.runtime.python_gc.gc.collect", lambda: calls.append("collect"))
monkeypatch.setattr("deepspeed.runtime.python_gc.gc.disable", lambda: calls.append("disable"))
monkeypatch.setattr("deepspeed.runtime.python_gc.gc.enable", lambda: calls.append("enable"))

manager = PythonGCManager()
generation = manager.acquire()
manager.release(generation)

assert calls == []


def test_python_gc_manager_explicit_collection(monkeypatch):
monkeypatch.setattr("deepspeed.runtime.python_gc.gc.collect", lambda: 11)
assert PythonGCManager().collect() == 11


def test_python_gc_manager_restores_enabled_state_after_fork(monkeypatch):
calls = []
enabled = False

def isenabled():
return enabled

def enable():
nonlocal enabled
enabled = True
calls.append("enable")

monkeypatch.setattr("deepspeed.runtime.python_gc.gc.isenabled", isenabled)
monkeypatch.setattr("deepspeed.runtime.python_gc.gc.enable", enable)

manager = PythonGCManager()
manager._active_engines = 1
manager._restore_enabled = True
manager._after_fork_child()

assert calls == ["enable"]
assert manager.active_engines == 0


def test_python_gc_manager_preserves_disabled_state_after_fork(monkeypatch):
calls = []
monkeypatch.setattr("deepspeed.runtime.python_gc.gc.isenabled", lambda: False)
monkeypatch.setattr("deepspeed.runtime.python_gc.gc.enable", lambda: calls.append("enable"))

manager = PythonGCManager()
manager._active_engines = 1
manager._restore_enabled = False
manager._after_fork_child()

assert calls == []
assert manager.active_engines == 0


def test_python_gc_manager_ignores_pre_fork_release(monkeypatch):
enabled = True

def isenabled():
return enabled

def disable():
nonlocal enabled
enabled = False

def enable():
nonlocal enabled
enabled = True

monkeypatch.setattr("deepspeed.runtime.python_gc.gc.isenabled", isenabled)
monkeypatch.setattr("deepspeed.runtime.python_gc.gc.collect", lambda: 0)
monkeypatch.setattr("deepspeed.runtime.python_gc.gc.disable", disable)
monkeypatch.setattr("deepspeed.runtime.python_gc.gc.enable", enable)

manager = PythonGCManager()
parent_generation = manager.acquire()
manager._after_fork_child()
child_generation = manager.acquire()

manager.release(parent_generation)
assert manager.active_engines == 1
assert enabled is False

manager.release(child_generation)
assert manager.active_engines == 0
assert enabled is True


def test_python_gc_manager_collection_allows_reentrant_release(monkeypatch):
manager = PythonGCManager()
enabled = True

monkeypatch.setattr("deepspeed.runtime.python_gc.gc.isenabled", lambda: enabled)

def collect():
manager.release(-1)
return 0

def disable():
nonlocal enabled
enabled = False

monkeypatch.setattr("deepspeed.runtime.python_gc.gc.collect", collect)
monkeypatch.setattr("deepspeed.runtime.python_gc.gc.disable", disable)

generation = manager.acquire()
assert manager.active_engines == 1
manager.release(generation)
Loading