From 8f8f69327a7283d2b0dd57dfa9790581053509fe Mon Sep 17 00:00:00 2001 From: Masahiro Tanaka Date: Sun, 19 Jul 2026 14:56:14 -0700 Subject: [PATCH 1/2] Scope DeepCompile compiler lifecycle ownership Signed-off-by: Masahiro Tanaka --- deepspeed/compile/backend.py | 81 +++++++---- deepspeed/compile/inductor.py | 61 +++++---- deepspeed/compile/init_z1.py | 9 +- deepspeed/compile/init_z3.py | 67 ++++++++- deepspeed/compile/patch_compiled_func.py | 11 +- deepspeed/runtime/engine.py | 18 +++ tests/unit/compile/test_backend.py | 123 +++++++++++++++++ tests/unit/compile/test_zero3_grad_dtype.py | 143 +++++++++++++++++++- 8 files changed, 455 insertions(+), 58 deletions(-) create mode 100644 tests/unit/compile/test_backend.py diff --git a/deepspeed/compile/backend.py b/deepspeed/compile/backend.py index ea207f39f6a5..2ed005fd453e 100644 --- a/deepspeed/compile/backend.py +++ b/deepspeed/compile/backend.py @@ -77,6 +77,22 @@ def clear(self): fwd_real_inputs = [] +def cleanup_compiled_backward_state(frame_id=None, owned_frames=None): + """Release engine-owned process-global compiled-backward state.""" + if frame_id is None: + if owned_frames is None: + frames_needing_bwd.clear() + else: + frames_needing_bwd.difference_update(owned_frames) + owned_frames.clear() + else: + frames_needing_bwd.discard(frame_id) + if owned_frames is not None: + owned_frames.discard(frame_id) + if len(frames_needing_bwd) == 0: + unpatch_compiled_func() + + def register_compile_pass(name: str, opt_pass_fn, contract=None): from .passes.contract import register_pass_contract opt_passes[name] = opt_pass_fn @@ -97,7 +113,8 @@ def init_schedule(schedule): remaining_schedule = deque(schedule) -def launch_compile_passes(global_steps: int): +def launch_compile_passes(global_steps: int, owned_frames=None): + """Advance the pass schedule and discard state owned by the previous compile cycle.""" global next_pass_step, next_passes if len(remaining_schedule) > 0 and global_steps == remaining_schedule[0][0]: @@ -109,6 +126,7 @@ def launch_compile_passes(global_steps: int): graph_order_with_frame_id.clear() profiling_results.clear() param_manager.clear() + cleanup_compiled_backward_state(owned_frames=owned_frames) frames_partitioned.clear() @@ -201,6 +219,21 @@ def set_example_values_to_symints(real_inputs, param_indices=None): return tuple(real_inputs_ret) +def _get_fw_real_inputs(local_real_inputs, input_storage: InputStorage, graph_id: int, debug_log: bool = False): + """Resolve graph-local real inputs from the one-shot queue or persistent storage.""" + if local_real_inputs: + return local_real_inputs.popleft() + + if input_storage.has_data(): + if debug_log: + log_rank0(f"Retrieving real inputs from storage for graph_id={graph_id}", enable=True) + return input_storage.get() + + raise RuntimeError(f"No real inputs available for graph_id {graph_id}. " + f"Local queue size: {len(local_real_inputs)}, " + f"storage has data: {input_storage.has_data()}") + + def run_opt_passes(opt_passes: List[Callable], gm: GraphModule, graph_id: int, @@ -243,7 +276,7 @@ def run_opt_passes(opt_passes: List[Callable], get_accelerator().empty_cache() -def make_backend(backend, compile_config, compile_kwargs={}): +def make_backend(backend, compile_config, compile_kwargs={}, owned_frames=None): register_custom_ops() @@ -251,6 +284,10 @@ def make_backend(backend, compile_config, compile_kwargs={}): debug_log = compile_config.debug_log free_activation = compile_config.free_activation and not is_backend_inductor(backend) + if owned_frames is None: + owned_frames = set() + owner_token = object() + def backend_fn(gm: GraphModule, real_inputs): graph_id = id(gm.graph) @@ -263,6 +300,7 @@ def backend_fn(gm: GraphModule, real_inputs): # This check cannot be placed here because autograd creates the fw/bw compiler callables before graph # partitioning. It is thus postponed to the point where the fw compiler is called. frame_id = gm.meta["dynamo_compile_id"].frame_id + frame_key = (owner_token, frame_id) graph_order_with_frame_id.add_graph(graph_id, frame_id) z3_partition = any(hasattr(v, "ds_id") for v in real_inputs) @@ -275,17 +313,15 @@ def backend_fn(gm: GraphModule, real_inputs): param_indices = [(i, input_val.param_id, input_val.shape) for i, input_val in enumerate(real_inputs) if isinstance(input_val, torch.nn.Parameter)] - global fwd_real_inputs - # Create an InputStorage instance for this specific graph # It will be captured by the make_fw_graph closure, eliminating the need for graph ID management input_storage = InputStorage(keep_int_input_tensors=compile_config.keep_int_input_tensors, keep_all_input_tensors=compile_config.keep_all_input_tensors) - # Store in both list (for backward compatibility) and storage (for persistence) + # Store in a closure-local queue and storage (for persistence). # The input_storage keeps tensor metadata to handle cases where # backend_fn is called once but make_fw_graph is called multiple times - fwd_real_inputs.append(real_inputs) + local_fwd_real_inputs = deque([real_inputs]) input_storage.put(real_inputs) global profiling_results @@ -304,20 +340,10 @@ def make_fw_graph(gm, sample_inputs): if needs_backward: if len(frames_needing_bwd) == 0: patch_compiled_func() - frames_needing_bwd.add(frame_id) - - # Try to get real_inputs from the list first, then from storage - if fwd_real_inputs: - real_inputs = fwd_real_inputs.pop(0) - elif input_storage.has_data(): - # Note: input_storage is captured from the enclosing backend_fn scope - # Materialize tensors from storage when list is empty - log_rank0(f"Retrieving real inputs from storage for graph_id={graph_id}", enable=debug_log) - real_inputs = input_storage.get() - else: - raise RuntimeError(f"No real inputs available for graph_id {graph_id}. " - f"List size: {len(fwd_real_inputs)}, Storage has data: {input_storage.has_data()}") + frames_needing_bwd.add(frame_key) + owned_frames.add(frame_key) + real_inputs = _get_fw_real_inputs(local_fwd_real_inputs, input_storage, graph_id, debug_log=debug_log) real_inputs = set_example_values_to_symints(real_inputs) param_manager[graph_id] = DSGraphParamManager(gm.graph, real_inputs, param_indices) @@ -385,9 +411,7 @@ def make_bw_graph(gm, sample_inputs): add_free_activations(graph_id, gm.graph, get_activation_node_names(gm.graph, param_nodes_bw, non_param_input_names)) - frames_needing_bwd.remove(frame_id) - if len(frames_needing_bwd) == 0: - unpatch_compiled_func() + cleanup_compiled_backward_state(frame_key, owned_frames) log_rank0( f"Bwd end {graph_index} graph_id={graph_id} alloc_mem={get_accelerator().memory_allocated()} graph={gm.graph}", @@ -415,10 +439,15 @@ def compiler_fn(gm, sample_inputs): partition_fn=partition_fn) return torch._dynamo.optimize(**compile_kwargs)(aot_mod) elif backend == "inductor": - patch_create_aot_dispatcher_function(graph_id, z3_partition, make_fw_graph, make_bw_graph, real_inputs, - param_indices, param_manager, frame_id, frames_partitioned) - - return torch._inductor.compile(gm, real_inputs) + restore_aotautograd = patch_create_aot_dispatcher_function(graph_id, z3_partition, make_fw_graph, + make_bw_graph, real_inputs, param_indices, + param_manager, frame_id, frames_partitioned) + try: + return torch._inductor.compile(gm, real_inputs) + finally: + # AotAutograd.__init__ is process-global; never leak this + # graph-specific compiler wiring into a later compilation. + restore_aotautograd() raise ValueError(f"Unsupported backend {backend}") diff --git a/deepspeed/compile/inductor.py b/deepspeed/compile/inductor.py index 3c2fa02ba4f8..5bcbc791a313 100644 --- a/deepspeed/compile/inductor.py +++ b/deepspeed/compile/inductor.py @@ -144,36 +144,45 @@ def _patch_deepcompile_aot_kwargs(kwargs: dict, *, graph_id: int, z3_partition: def patch_create_aot_dispatcher_function(graph_id: int, z3_partition: bool, make_fw_graph, make_bw_graph, real_inputs, param_indices, param_manager, frame_id: int, frames_partitioned: Set[int]): + """Temporarily install graph-specific AOT compilers and return an idempotent restore callback.""" from torch._dynamo.backends.common import AotAutograd import functools - def patch_aotautograd(): - # Unpatch if it was already patched - if hasattr(AotAutograd, "__original_init"): - AotAutograd.__init__ = AotAutograd.__original_init - - original_init = AotAutograd.__init__ - - @functools.wraps(original_init) - def patched_init(self, **kwargs): - _patch_deepcompile_aot_kwargs(kwargs, - graph_id=graph_id, - z3_partition=z3_partition, - make_fw_graph=make_fw_graph, - make_bw_graph=make_bw_graph, - real_inputs=real_inputs, - param_indices=param_indices, - param_manager=param_manager, - frame_id=frame_id, - frames_partitioned=frames_partitioned) - - original_init(self, **kwargs) - - AotAutograd.__original_init = original_init - AotAutograd.__init__ = patched_init - - patch_aotautograd() + # The constructor patch is process-global. Replace the currently installed + # DeepCompile patch before taking ownership for this graph. + if hasattr(AotAutograd, "__original_init"): + AotAutograd.__init__ = AotAutograd.__original_init + delattr(AotAutograd, "__original_init") + + original_init = AotAutograd.__init__ + + @functools.wraps(original_init) + def patched_init(self, **kwargs): + _patch_deepcompile_aot_kwargs(kwargs, + graph_id=graph_id, + z3_partition=z3_partition, + make_fw_graph=make_fw_graph, + make_bw_graph=make_bw_graph, + real_inputs=real_inputs, + param_indices=param_indices, + param_manager=param_manager, + frame_id=frame_id, + frames_partitioned=frames_partitioned) + + original_init(self, **kwargs) + + AotAutograd.__original_init = original_init + AotAutograd.__init__ = patched_init + + def restore_aotautograd(): + """Restore only this invocation's patch without clobbering a newer owner.""" + if AotAutograd.__init__ is patched_init: + AotAutograd.__init__ = original_init + if getattr(AotAutograd, "__original_init", None) is original_init: + delattr(AotAutograd, "__original_init") + + return restore_aotautograd def register_custom_ops(): diff --git a/deepspeed/compile/init_z1.py b/deepspeed/compile/init_z1.py index 880494b94200..025e45d23f0d 100644 --- a/deepspeed/compile/init_z1.py +++ b/deepspeed/compile/init_z1.py @@ -4,6 +4,7 @@ # DeepSpeed Team import copy +from functools import partial import torch @@ -185,5 +186,9 @@ def release_grad_buffer(group_idx=None): init_schedule(schedule) - engine.launch_compile_passes = launch_compile_passes - return make_backend(backend, compile_config, compile_kwargs=compile_kwargs) + engine._deepcompile_owned_frames = set() + engine.launch_compile_passes = partial(launch_compile_passes, owned_frames=engine._deepcompile_owned_frames) + return make_backend(backend, + compile_config, + compile_kwargs=compile_kwargs, + owned_frames=engine._deepcompile_owned_frames) diff --git a/deepspeed/compile/init_z3.py b/deepspeed/compile/init_z3.py index 0e20c02ede2e..9cb1bcff339c 100644 --- a/deepspeed/compile/init_z3.py +++ b/deepspeed/compile/init_z3.py @@ -3,6 +3,9 @@ # DeepSpeed Team +from functools import partial +from threading import Lock + import torch from deepspeed import comm as dist @@ -19,6 +22,54 @@ WARMUP = 5 _MISSING = object() +_DYNAMO_CONFIG_NAMES = ("force_parameter_static_shapes", "force_nn_module_property_static_shapes") +_DYNAMO_CONFIG_OWNERS = {} +_DYNAMO_CONFIG_LOCK = Lock() + + +def _allow_dynamo_dynamic_parameter_shapes_for_z3(compile_kwargs): + """Acquire process-wide ZeRO-3 Dynamo config ownership and return its release callback.""" + dynamo = getattr(torch, "_dynamo", None) + if dynamo is None: + try: + import torch._dynamo as dynamo + except ImportError: + return None + + dynamo_config = getattr(dynamo, "config", None) + if dynamo_config is None: + return None + + owner_token = object() + config_key = id(dynamo_config) + with _DYNAMO_CONFIG_LOCK: + state = _DYNAMO_CONFIG_OWNERS.get(config_key) + if state is None or state["config"] is not dynamo_config: + previous_values = { + config_name: getattr(dynamo_config, config_name) + for config_name in _DYNAMO_CONFIG_NAMES if hasattr(dynamo_config, config_name) + } + if not previous_values: + return None + state = {"config": dynamo_config, "previous_values": previous_values, "owner_tokens": set()} + _DYNAMO_CONFIG_OWNERS[config_key] = state + state["owner_tokens"].add(owner_token) + for config_name in state["previous_values"]: + setattr(dynamo_config, config_name, False) + + def restore(): + with _DYNAMO_CONFIG_LOCK: + state = _DYNAMO_CONFIG_OWNERS.get(config_key) + if state is None or state["config"] is not dynamo_config or owner_token not in state["owner_tokens"]: + return + state["owner_tokens"].remove(owner_token) + if state["owner_tokens"]: + return + for config_name, previous_value in state["previous_values"].items(): + setattr(dynamo_config, config_name, previous_value) + del _DYNAMO_CONFIG_OWNERS[config_key] + + return restore def _resolve_expected_grad_dtype(param): @@ -102,9 +153,21 @@ def set_grad_buffer(_is_gradient_accumulation_boundary): if move_opt_states in passes or move_opt_states_sync in passes: init_offload_opt_states(optimizer, dc) - engine.launch_compile_passes = launch_compile_passes + engine._deepcompile_owned_frames = set() + engine.launch_compile_passes = partial(launch_compile_passes, owned_frames=engine._deepcompile_owned_frames) patch_fake_tensor() torch._inductor.config.size_asserts = False - return make_backend(backend, compile_config, compile_kwargs=compile_kwargs) + previous_restore = getattr(engine, "_deepcompile_dynamo_config_restore", None) + if previous_restore is not None: + previous_restore() + del engine._deepcompile_dynamo_config_restore + restore_dynamo_config = _allow_dynamo_dynamic_parameter_shapes_for_z3(compile_kwargs) + if restore_dynamo_config is not None: + engine._deepcompile_dynamo_config_restore = restore_dynamo_config + + return make_backend(backend, + compile_config, + compile_kwargs=compile_kwargs, + owned_frames=engine._deepcompile_owned_frames) diff --git a/deepspeed/compile/patch_compiled_func.py b/deepspeed/compile/patch_compiled_func.py index c77d529a64ac..2a0ffd593129 100644 --- a/deepspeed/compile/patch_compiled_func.py +++ b/deepspeed/compile/patch_compiled_func.py @@ -82,12 +82,21 @@ class PatchedFunction(torch.autograd.Function, metaclass=FunctionMeta): def unpatch_compiled_func(): + """Restore torch.autograd.Function and discard inputs captured for this compile cycle.""" global enabled_patched_func enabled_patched_func = False global original_grad_fn - torch.autograd.Function = original_grad_fn + if original_grad_fn is not None: + torch.autograd.Function = original_grad_fn + original_grad_fn = None + clear_backward_inputs() def get_backward_inputs(): return backward_inputs + + +def clear_backward_inputs(): + """Drop captured real backward inputs before the next graph compilation.""" + backward_inputs.clear() diff --git a/deepspeed/runtime/engine.py b/deepspeed/runtime/engine.py index 194dc3555387..bfb59b8c1163 100755 --- a/deepspeed/runtime/engine.py +++ b/deepspeed/runtime/engine.py @@ -772,6 +772,8 @@ def __del__(self): logger.debug("DeepSpeedEngine.__del__ cleanup skipped: %s", exc, exc_info=True) def destroy(self): + 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() @@ -5523,6 +5525,10 @@ def compile(self, def _set_deepcompile_active(self, active: bool) -> None: """Toggle DeepCompile runtime state and manage forward hooks accordingly.""" + if not active: + self._release_deepcompile_compiled_backward_state() + self._release_deepcompile_dynamo_config() + if self._deepcompile_active == active: return @@ -5541,6 +5547,18 @@ def _set_deepcompile_active(self, active: bool) -> None: self._deepcompile_active = active + def _release_deepcompile_compiled_backward_state(self) -> None: + owned_frames = getattr(self, "_deepcompile_owned_frames", None) + if owned_frames: + from deepspeed.compile.backend import cleanup_compiled_backward_state + cleanup_compiled_backward_state(owned_frames=owned_frames) + + def _release_deepcompile_dynamo_config(self) -> None: + restore_dynamo_config = getattr(self, "_deepcompile_dynamo_config_restore", None) + if restore_dynamo_config is not None: + restore_dynamo_config() + del self._deepcompile_dynamo_config_restore + def get_compile_time(self): from deepspeed.compile.backend import opt_pass_times return opt_pass_times diff --git a/tests/unit/compile/test_backend.py b/tests/unit/compile/test_backend.py new file mode 100644 index 000000000000..d3f7970e078e --- /dev/null +++ b/tests/unit/compile/test_backend.py @@ -0,0 +1,123 @@ +# Copyright (c) DeepSpeed Team. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +from collections import deque + +import torch + +from deepspeed.compile import backend as backend_mod +from deepspeed.compile.backend import _get_fw_real_inputs +from deepspeed.compile.inductor import patch_create_aot_dispatcher_function +from deepspeed.compile.input_storage import InputStorage +from deepspeed.compile.patch_compiled_func import (clear_backward_inputs, get_backward_inputs, patch_compiled_func, + unpatch_compiled_func) + + +def test_forward_real_inputs_are_graph_local(): + local_inputs = (torch.nn.Parameter(torch.ones(2, dtype=torch.float32)), ) + storage = InputStorage() + storage.put((torch.ones(1, dtype=torch.float32), )) + + selected = _get_fw_real_inputs(deque([local_inputs]), storage, graph_id=7) + + assert selected is local_inputs + + +def test_forward_real_inputs_fall_back_to_storage_when_local_queue_is_empty(): + storage = InputStorage() + storage.put((torch.ones(3, dtype=torch.float32), )) + + selected = _get_fw_real_inputs(deque(), storage, graph_id=7) + + assert len(selected) == 1 + assert selected[0].shape == torch.Size([3]) + assert selected[0].dtype is torch.float32 + + +def test_launch_compile_passes_clears_owned_compiled_backward_state(monkeypatch): + + class DummyDeepCompileHandle: + + def reset(self): + pass + + clear_backward_inputs() + backend_mod.frames_needing_bwd.clear() + unpatch_compiled_func() + original_autograd_function = torch.autograd.Function + owner = object() + owned_frames = {(owner, 17)} + backend_mod.frames_needing_bwd.update(owned_frames) + patch_compiled_func() + get_backward_inputs().append((torch.ones(1), )) + monkeypatch.setattr(backend_mod, "log_rank0", lambda *args, **kwargs: None) + monkeypatch.setattr(backend_mod, "get_deepcompile_handle", lambda: DummyDeepCompileHandle()) + + backend_mod.init_schedule([(0, [])]) + try: + backend_mod.launch_compile_passes(0, owned_frames=owned_frames) + + assert owned_frames == set() + assert backend_mod.frames_needing_bwd == set() + assert get_backward_inputs() == [] + assert torch.autograd.Function is original_autograd_function + finally: + backend_mod.frames_needing_bwd.clear() + unpatch_compiled_func() + + +def test_unpatch_compiled_func_clears_backward_inputs(): + clear_backward_inputs() + patch_compiled_func() + try: + get_backward_inputs().append((torch.ones(1), )) + unpatch_compiled_func() + assert get_backward_inputs() == [] + finally: + unpatch_compiled_func() + + +def _patch_aot_constructor(): + return patch_create_aot_dispatcher_function(graph_id=7, + z3_partition=False, + make_fw_graph=lambda gm, sample_inputs: gm.graph, + make_bw_graph=lambda gm, sample_inputs: gm.graph, + real_inputs=(torch.ones(1), ), + param_indices=[], + param_manager={}, + frame_id=0, + frames_partitioned=set()) + + +def test_inductor_aot_constructor_patch_is_restorable(): + from torch._dynamo.backends.common import AotAutograd + + original_init = AotAutograd.__init__ + restore = _patch_aot_constructor() + try: + assert AotAutograd.__init__ is not original_init + finally: + restore() + + assert AotAutograd.__init__ is original_init + assert not hasattr(AotAutograd, "__original_init") + + +def test_older_aot_restore_does_not_clobber_newer_patch(): + from torch._dynamo.backends.common import AotAutograd + + original_init = AotAutograd.__init__ + restore_first = _patch_aot_constructor() + restore_second = _patch_aot_constructor() + newer_init = AotAutograd.__init__ + try: + restore_first() + assert AotAutograd.__init__ is newer_init + assert hasattr(AotAutograd, "__original_init") + finally: + restore_second() + + assert AotAutograd.__init__ is original_init + assert not hasattr(AotAutograd, "__original_init") diff --git a/tests/unit/compile/test_zero3_grad_dtype.py b/tests/unit/compile/test_zero3_grad_dtype.py index a2fdaae89d9d..3201de5e419f 100644 --- a/tests/unit/compile/test_zero3_grad_dtype.py +++ b/tests/unit/compile/test_zero3_grad_dtype.py @@ -3,9 +3,12 @@ # DeepSpeed Team +import pytest import torch -from deepspeed.compile.init_z3 import _resolve_expected_grad_dtype +from deepspeed.compile import backend as backend_mod +from deepspeed.compile.init_z3 import _allow_dynamo_dynamic_parameter_shapes_for_z3, _resolve_expected_grad_dtype +from deepspeed.runtime.engine import DeepSpeedEngine def test_missing_grad_dtype_attribute_falls_back_to_param_dtype(): @@ -28,3 +31,141 @@ def test_explicit_grad_dtype_is_preserved(): param.grad_dtype = torch.float32 assert _resolve_expected_grad_dtype(param) is torch.float32 + + +def test_zero3_allows_dynamo_dynamic_parameter_shapes(monkeypatch): + + class FakeDynamoConfig: + force_parameter_static_shapes = True + force_nn_module_property_static_shapes = True + + class FakeDynamo: + config = FakeDynamoConfig() + + monkeypatch.setattr(torch, "_dynamo", FakeDynamo) + + restore = _allow_dynamo_dynamic_parameter_shapes_for_z3({}) + assert restore + try: + assert FakeDynamo.config.force_parameter_static_shapes is False + assert FakeDynamo.config.force_nn_module_property_static_shapes is False + finally: + restore() + + +@pytest.mark.parametrize("first_owner_to_restore", [0, 1]) +def test_zero3_dynamo_config_restores_after_last_overlapping_owner(monkeypatch, first_owner_to_restore): + + class FakeDynamoConfig: + force_parameter_static_shapes = True + force_nn_module_property_static_shapes = False + + class FakeDynamo: + config = FakeDynamoConfig() + + monkeypatch.setattr(torch, "_dynamo", FakeDynamo) + restores = [_allow_dynamo_dynamic_parameter_shapes_for_z3({}), _allow_dynamo_dynamic_parameter_shapes_for_z3({})] + + assert all(restores) + restores[first_owner_to_restore]() + assert FakeDynamo.config.force_parameter_static_shapes is False + restores[1 - first_owner_to_restore]() + assert FakeDynamo.config.force_parameter_static_shapes is True + assert FakeDynamo.config.force_nn_module_property_static_shapes is False + + +@pytest.mark.parametrize("first_owner_to_destroy", [0, 1]) +def test_zero3_dynamo_config_restores_when_overlapping_engines_are_destroyed(monkeypatch, first_owner_to_destroy): + + class FakeDynamoConfig: + force_parameter_static_shapes = True + force_nn_module_property_static_shapes = False + + class FakeDynamo: + config = FakeDynamoConfig() + + monkeypatch.setattr(torch, "_dynamo", FakeDynamo) + engines = [object.__new__(DeepSpeedEngine), object.__new__(DeepSpeedEngine)] + for engine in engines: + torch.nn.Module.__init__(engine) + engine._deepcompile_active = False + engine._deepcompile_dynamo_config_restore = _allow_dynamo_dynamic_parameter_shapes_for_z3({}) + + engines[first_owner_to_destroy].destroy() + assert FakeDynamo.config.force_parameter_static_shapes is False + engines[1 - first_owner_to_destroy].destroy() + assert FakeDynamo.config.force_parameter_static_shapes is True + assert FakeDynamo.config.force_nn_module_property_static_shapes is False + + +@pytest.mark.parametrize("first_owner_to_destroy", [0, 1]) +def test_destroy_releases_only_owner_with_overlapping_frame_ids(first_owner_to_destroy): + original_autograd_function = torch.autograd.Function + engines = [object.__new__(DeepSpeedEngine), object.__new__(DeepSpeedEngine)] + owners = [object(), object()] + frame_id = 17 + for owner, engine in zip(owners, engines): + torch.nn.Module.__init__(engine) + engine._deepcompile_active = False + engine._deepcompile_owned_frames = {(owner, frame_id)} + + backend_mod.frames_needing_bwd.clear() + backend_mod.frames_needing_bwd.update(((owners[0], frame_id), (owners[1], frame_id))) + backend_mod.patch_compiled_func() + backend_mod.get_backward_inputs().append((torch.ones(1), )) + + try: + engines[first_owner_to_destroy].destroy() + surviving_owner = owners[1 - first_owner_to_destroy] + assert backend_mod.frames_needing_bwd == {(surviving_owner, frame_id)} + assert len(backend_mod.get_backward_inputs()) == 1 + assert torch.autograd.Function is not original_autograd_function + + engines[1 - first_owner_to_destroy].destroy() + assert backend_mod.frames_needing_bwd == set() + assert backend_mod.get_backward_inputs() == [] + assert torch.autograd.Function is original_autograd_function + finally: + backend_mod.frames_needing_bwd.clear() + backend_mod.unpatch_compiled_func() + + +def test_deactivation_releases_only_the_engine_owned_state(monkeypatch): + + class FakeDynamoConfig: + force_parameter_static_shapes = True + force_nn_module_property_static_shapes = False + + class FakeDynamo: + config = FakeDynamoConfig() + + monkeypatch.setattr(torch, "_dynamo", FakeDynamo) + engine = object.__new__(DeepSpeedEngine) + torch.nn.Module.__init__(engine) + engine._deepcompile_active = True + engine.module_forward_pre_hook = object() + engine.module_forward_post_hook = object() + engine._deepcompile_dynamo_config_restore = _allow_dynamo_dynamic_parameter_shapes_for_z3({}) + original_autograd_function = torch.autograd.Function + owner = object() + other_owner = object() + backend_mod.frames_needing_bwd.clear() + backend_mod.frames_needing_bwd.update(((owner, 17), (other_owner, 18))) + engine._deepcompile_owned_frames = {(owner, 17)} + backend_mod.patch_compiled_func() + backend_mod.get_backward_inputs().append((torch.ones(1), )) + + try: + engine._set_deepcompile_active(False) + + assert FakeDynamo.config.force_parameter_static_shapes is True + assert FakeDynamo.config.force_nn_module_property_static_shapes is False + assert not hasattr(engine, "_deepcompile_dynamo_config_restore") + assert engine._deepcompile_owned_frames == set() + assert backend_mod.frames_needing_bwd == {(other_owner, 18)} + assert len(backend_mod.get_backward_inputs()) == 1 + assert torch.autograd.Function is not original_autograd_function + assert engine.is_deepcompile_active() is False + finally: + backend_mod.frames_needing_bwd.clear() + backend_mod.unpatch_compiled_func() From baca2be2a2608bb8c2a2522bcb85d4b33b18a510 Mon Sep 17 00:00:00 2001 From: Masahiro Tanaka Date: Thu, 23 Jul 2026 15:07:51 -0700 Subject: [PATCH 2/2] Fix DeepCompile backward input ownership Signed-off-by: Masahiro Tanaka --- deepspeed/compile/backend.py | 15 ++-- deepspeed/compile/patch_compiled_func.py | 55 +++++++++++--- tests/unit/compile/test_backend.py | 8 +-- tests/unit/compile/test_zero3_grad_dtype.py | 80 ++++++++++++++++++--- 4 files changed, 127 insertions(+), 31 deletions(-) diff --git a/deepspeed/compile/backend.py b/deepspeed/compile/backend.py index 2ed005fd453e..2030781c1dc9 100644 --- a/deepspeed/compile/backend.py +++ b/deepspeed/compile/backend.py @@ -28,7 +28,8 @@ from .graph_param import DSGraphParamManager from .profilers import ProfilingResult from .profilers.graph_profile import MemoryProfilingInterpreter -from .patch_compiled_func import patch_compiled_func, unpatch_compiled_func, get_backward_inputs +from .patch_compiled_func import (clear_backward_inputs, patch_compiled_func, pop_backward_input, + register_backward_frame, unpatch_compiled_func) from .util import get_input_nodes, get_activation_node_names, get_index_by_graph_id, get_deepcompile_handle, log_rank0, is_backend_inductor from .partitioner import get_wrapped_partitioner from .inductor import register_custom_ops, patch_create_aot_dispatcher_function @@ -81,14 +82,18 @@ def cleanup_compiled_backward_state(frame_id=None, owned_frames=None): """Release engine-owned process-global compiled-backward state.""" if frame_id is None: if owned_frames is None: + released_frames = set(frames_needing_bwd) frames_needing_bwd.clear() else: + released_frames = set(owned_frames) frames_needing_bwd.difference_update(owned_frames) owned_frames.clear() else: + released_frames = {frame_id} frames_needing_bwd.discard(frame_id) if owned_frames is not None: owned_frames.discard(frame_id) + clear_backward_inputs(released_frames) if len(frames_needing_bwd) == 0: unpatch_compiled_func() @@ -342,6 +347,7 @@ def make_fw_graph(gm, sample_inputs): patch_compiled_func() frames_needing_bwd.add(frame_key) owned_frames.add(frame_key) + register_backward_frame(frame_key) real_inputs = _get_fw_real_inputs(local_fwd_real_inputs, input_storage, graph_id, debug_log=debug_log) real_inputs = set_example_values_to_symints(real_inputs) @@ -377,10 +383,10 @@ def make_bw_graph(gm, sample_inputs): f"Bwd start {graph_index} graph_id={graph_id} alloc_mem={get_accelerator().memory_allocated()} graph={gm.graph}", enable=debug_log) - bwd_inputs_stack = get_backward_inputs() + bwd_real_inputs = pop_backward_input(frame_key) param_nodes_bw, _ = param_manager[graph_id].get_bwd_mapping(gm.graph) - if len(bwd_inputs_stack) == 0: + if bwd_real_inputs is None: # dynamo calls bw compiler ahead of time when symints are saved for backward. See the details for aot_dispatch_autograd in jit_compile_runtime_wrappers. # As we currently use actually bwd input values in bw compiler, we make dummy data for profiling. # Replace fake tensors with real parameters before calling set_example_values_to_symints @@ -388,9 +394,6 @@ def make_bw_graph(gm, sample_inputs): sample_inputs_with_real_params = param_manager[graph_id].replace_fake_tensors_with_real_params( sample_inputs, gm.graph) bwd_real_inputs = set_example_values_to_symints(sample_inputs_with_real_params) - else: - bwd_real_inputs = bwd_inputs_stack.pop() - run_opt_passes( opt_passes=next_passes, gm=gm, diff --git a/deepspeed/compile/patch_compiled_func.py b/deepspeed/compile/patch_compiled_func.py index 2a0ffd593129..8a83d0e25e09 100644 --- a/deepspeed/compile/patch_compiled_func.py +++ b/deepspeed/compile/patch_compiled_func.py @@ -3,10 +3,13 @@ # DeepSpeed Team +from collections import deque + import torch from deepspeed.utils.torch import required_torch_version -backward_inputs = [] +backward_inputs = {} +backward_frame_keys = deque() enabled_patched_func = False original_grad_fn = None @@ -19,12 +22,13 @@ class FunctionMeta(base_meta): def __new__(cls, name, bases, dct): if name == "CompiledFunction": original_backward_impl = dct.get("_backward_impl") + frame_key = backward_frame_keys.popleft() if backward_frame_keys else None def wrapped_backward_impl(ctx, all_args): assert original_backward_impl is not None - if enabled_patched_func: - backward_inputs.append(all_args) + if enabled_patched_func and frame_key is not None: + backward_inputs.setdefault(frame_key, []).append(all_args) wrapped_backward_impl.owner_class.compiled_bw = None return original_backward_impl(ctx, all_args) @@ -45,13 +49,14 @@ class FunctionMeta(base_meta): def __new__(cls, name, bases, dct): if name == "CompiledFunction": original_backward_prologue = dct.get("_backward_prologue") + frame_key = backward_frame_keys.popleft() if backward_frame_keys else None def wrapped_backward_prologue(ctx, *grad_outputs): assert original_backward_prologue is not None all_args = original_backward_prologue(ctx, *grad_outputs) - if enabled_patched_func: - backward_inputs.append(all_args) + if enabled_patched_func and frame_key is not None: + backward_inputs.setdefault(frame_key, []).append(all_args) wrapped_backward_prologue.owner_class.compiled_bw = None return all_args @@ -81,6 +86,11 @@ class PatchedFunction(torch.autograd.Function, metaclass=FunctionMeta): return backward_inputs +def register_backward_frame(frame_key): + """Associate the next AOT compiled function with its DeepCompile frame.""" + backward_frame_keys.append(frame_key) + + def unpatch_compiled_func(): """Restore torch.autograd.Function and discard inputs captured for this compile cycle.""" global enabled_patched_func @@ -93,10 +103,35 @@ def unpatch_compiled_func(): clear_backward_inputs() -def get_backward_inputs(): - return backward_inputs +def get_backward_inputs(frame_key=None): + if frame_key is None: + return backward_inputs + return backward_inputs.get(frame_key, []) + + +def pop_backward_input(frame_key): + """Pop captured real inputs for one DeepCompile frame.""" + frame_inputs = backward_inputs.get(frame_key) + if not frame_inputs: + return None + + inputs = frame_inputs.pop() + if not frame_inputs: + backward_inputs.pop(frame_key) + return inputs + + +def clear_backward_inputs(frame_keys=None): + """Drop captured inputs and pending capture registrations for selected frames.""" + if frame_keys is None: + backward_inputs.clear() + backward_frame_keys.clear() + return + frame_keys = set(frame_keys) + for frame_key in frame_keys: + backward_inputs.pop(frame_key, None) -def clear_backward_inputs(): - """Drop captured real backward inputs before the next graph compilation.""" - backward_inputs.clear() + retained_frame_keys = [frame_key for frame_key in backward_frame_keys if frame_key not in frame_keys] + backward_frame_keys.clear() + backward_frame_keys.extend(retained_frame_keys) diff --git a/tests/unit/compile/test_backend.py b/tests/unit/compile/test_backend.py index d3f7970e078e..538ea38eba45 100644 --- a/tests/unit/compile/test_backend.py +++ b/tests/unit/compile/test_backend.py @@ -51,7 +51,7 @@ def reset(self): owned_frames = {(owner, 17)} backend_mod.frames_needing_bwd.update(owned_frames) patch_compiled_func() - get_backward_inputs().append((torch.ones(1), )) + get_backward_inputs()[next(iter(owned_frames))] = [(torch.ones(1), )] monkeypatch.setattr(backend_mod, "log_rank0", lambda *args, **kwargs: None) monkeypatch.setattr(backend_mod, "get_deepcompile_handle", lambda: DummyDeepCompileHandle()) @@ -61,7 +61,7 @@ def reset(self): assert owned_frames == set() assert backend_mod.frames_needing_bwd == set() - assert get_backward_inputs() == [] + assert get_backward_inputs() == {} assert torch.autograd.Function is original_autograd_function finally: backend_mod.frames_needing_bwd.clear() @@ -72,9 +72,9 @@ def test_unpatch_compiled_func_clears_backward_inputs(): clear_backward_inputs() patch_compiled_func() try: - get_backward_inputs().append((torch.ones(1), )) + get_backward_inputs()[(object(), 17)] = [(torch.ones(1), )] unpatch_compiled_func() - assert get_backward_inputs() == [] + assert get_backward_inputs() == {} finally: unpatch_compiled_func() diff --git a/tests/unit/compile/test_zero3_grad_dtype.py b/tests/unit/compile/test_zero3_grad_dtype.py index 3201de5e419f..38552062a454 100644 --- a/tests/unit/compile/test_zero3_grad_dtype.py +++ b/tests/unit/compile/test_zero3_grad_dtype.py @@ -8,7 +8,9 @@ from deepspeed.compile import backend as backend_mod from deepspeed.compile.init_z3 import _allow_dynamo_dynamic_parameter_shapes_for_z3, _resolve_expected_grad_dtype +from deepspeed.compile.patch_compiled_func import (get_backward_inputs, pop_backward_input, register_backward_frame) from deepspeed.runtime.engine import DeepSpeedEngine +from deepspeed.utils.torch import required_torch_version def test_missing_grad_dtype_attribute_falls_back_to_param_dtype(): @@ -110,20 +112,73 @@ def test_destroy_releases_only_owner_with_overlapping_frame_ids(first_owner_to_d engine._deepcompile_owned_frames = {(owner, frame_id)} backend_mod.frames_needing_bwd.clear() - backend_mod.frames_needing_bwd.update(((owners[0], frame_id), (owners[1], frame_id))) + frame_keys = [(owners[0], frame_id), (owners[1], frame_id)] + backend_mod.frames_needing_bwd.update(frame_keys) backend_mod.patch_compiled_func() - backend_mod.get_backward_inputs().append((torch.ones(1), )) + for frame_key in frame_keys: + get_backward_inputs()[frame_key] = [(torch.ones(1), )] try: engines[first_owner_to_destroy].destroy() - surviving_owner = owners[1 - first_owner_to_destroy] - assert backend_mod.frames_needing_bwd == {(surviving_owner, frame_id)} - assert len(backend_mod.get_backward_inputs()) == 1 + surviving_frame = frame_keys[1 - first_owner_to_destroy] + assert backend_mod.frames_needing_bwd == {surviving_frame} + assert set(get_backward_inputs()) == {surviving_frame} assert torch.autograd.Function is not original_autograd_function engines[1 - first_owner_to_destroy].destroy() assert backend_mod.frames_needing_bwd == set() - assert backend_mod.get_backward_inputs() == [] + assert get_backward_inputs() == {} + assert torch.autograd.Function is original_autograd_function + finally: + backend_mod.frames_needing_bwd.clear() + backend_mod.unpatch_compiled_func() + + +def test_destroyed_owner_inputs_cannot_be_consumed_by_survivor(): + original_autograd_function = torch.autograd.Function + engines = [object.__new__(DeepSpeedEngine), object.__new__(DeepSpeedEngine)] + owners = [object(), object()] + frame_keys = [(owners[0], 17), (owners[1], 17)] + for frame_key, engine in zip(frame_keys, engines): + torch.nn.Module.__init__(engine) + engine._deepcompile_active = False + engine._deepcompile_owned_frames = {frame_key} + + backend_mod.frames_needing_bwd.clear() + backend_mod.frames_needing_bwd.update(frame_keys) + backend_mod.patch_compiled_func() + register_backward_frame(frame_keys[0]) + + class CompiledFunction(torch.autograd.Function): + compiled_bw = object() + + @staticmethod + def _backward_impl(ctx, all_args): + return all_args + + @staticmethod + def _backward_prologue(ctx, *grad_outputs): + return grad_outputs + + owner_a_inputs = (torch.ones(1), ) + if required_torch_version(min_version=2.7): + CompiledFunction._backward_impl(None, owner_a_inputs) + else: + CompiledFunction._backward_prologue(None, *owner_a_inputs) + + try: + assert len(get_backward_inputs(frame_keys[0])) == 1 + assert get_backward_inputs(frame_keys[1]) == [] + + engines[0].destroy() + + assert backend_mod.frames_needing_bwd == {frame_keys[1]} + assert get_backward_inputs(frame_keys[0]) == [] + assert pop_backward_input(frame_keys[1]) is None + assert torch.autograd.Function is not original_autograd_function + + engines[1].destroy() + assert get_backward_inputs() == {} assert torch.autograd.Function is original_autograd_function finally: backend_mod.frames_needing_bwd.clear() @@ -149,11 +204,14 @@ class FakeDynamo: original_autograd_function = torch.autograd.Function owner = object() other_owner = object() + frame_key = (owner, 17) + other_frame_key = (other_owner, 18) backend_mod.frames_needing_bwd.clear() - backend_mod.frames_needing_bwd.update(((owner, 17), (other_owner, 18))) - engine._deepcompile_owned_frames = {(owner, 17)} + backend_mod.frames_needing_bwd.update((frame_key, other_frame_key)) + engine._deepcompile_owned_frames = {frame_key} backend_mod.patch_compiled_func() - backend_mod.get_backward_inputs().append((torch.ones(1), )) + get_backward_inputs()[frame_key] = [(torch.ones(1), )] + get_backward_inputs()[other_frame_key] = [(torch.ones(1), )] try: engine._set_deepcompile_active(False) @@ -162,8 +220,8 @@ class FakeDynamo: assert FakeDynamo.config.force_nn_module_property_static_shapes is False assert not hasattr(engine, "_deepcompile_dynamo_config_restore") assert engine._deepcompile_owned_frames == set() - assert backend_mod.frames_needing_bwd == {(other_owner, 18)} - assert len(backend_mod.get_backward_inputs()) == 1 + assert backend_mod.frames_needing_bwd == {other_frame_key} + assert set(get_backward_inputs()) == {other_frame_key} assert torch.autograd.Function is not original_autograd_function assert engine.is_deepcompile_active() is False finally: