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
3 changes: 3 additions & 0 deletions csrc/compile/z3.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -191,6 +191,9 @@ class Z3CustomOpExecutor : public CustomOpExecutor {
}
}
auto target_dtype = dtype ? dtype.value() : ds_tensor.scalar_type();
// Keep prefetch allocations in the same stream-local allocator pool as the
// all-gather that consumes them, reducing cross-stream reuse pressure.
at::cuda::CUDAStreamGuard guard(ag_stream_);
output_bufs[ds_id] =
torch::empty({padded_numel}, ds_tensor.options().dtype(target_dtype));
}
Expand Down
197 changes: 154 additions & 43 deletions deepspeed/compile/backend.py

Large diffs are not rendered by default.

97 changes: 70 additions & 27 deletions deepspeed/compile/inductor.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@

# DeepSpeed Team

from contextlib import nullcontext
from typing import Set

import torch
Expand All @@ -21,6 +22,37 @@
from .graph_param import DSGraphParamManager
from .partitioner import get_wrapped_partitioner

_DEEP_COMPILE_Z3_INDUCTOR_REDUCTION_CONFIG = {
"triton.mix_order_reduction": False,
"triton.persistent_reductions": False,
}


def deepcompile_z3_inductor_config_patch(enabled: bool):
"""Disable reduction heuristics that create oversized kernels for DeepCompile ZeRO-3 graphs."""
if not enabled:
return nullcontext()

inductor = getattr(torch, "_inductor", None)
config = getattr(inductor, "config", None)
if config is None or not hasattr(config, "patch"):
return nullcontext()

triton_config = getattr(config, "triton", None)
if triton_config is None:
return nullcontext()

overrides = {
config_name: value
for config_name, value in _DEEP_COMPILE_Z3_INDUCTOR_REDUCTION_CONFIG.items()
if hasattr(triton_config,
config_name.split(".", 1)[1])
}
if not overrides:
return nullcontext()

return config.patch(overrides)


def _get_graphsafe_run_with_rng_state():
try:
Expand All @@ -40,6 +72,7 @@ def _register_graphsafe_rng_state_no_reuse(register_fallback_no_reuse):


def patch_compiler(original_compiler, dc_compiler, z3_partition: bool, graph_id, graph_param_manager, bwd: bool):
"""Wrap an AOT compiler with DeepCompile rewrites and ZeRO-3 fake-shape repair."""

def wrapped_compiler(gm, fake_inputs):
mod_graph = dc_compiler(gm, fake_inputs)
Expand Down Expand Up @@ -74,7 +107,8 @@ def wrapped_compiler(gm, fake_inputs):
else:
patched_inputs = fake_inputs

return original_compiler(gm, patched_inputs)
with deepcompile_z3_inductor_config_patch(z3_partition):
return original_compiler(gm, patched_inputs)

return wrapped_compiler

Expand Down Expand Up @@ -138,36 +172,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. Recover first if a previous
# compile failed before reaching its restoration callback.
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():
Expand Down
9 changes: 7 additions & 2 deletions deepspeed/compile/init_z1.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
# DeepSpeed Team

import copy
from functools import partial

import torch

Expand Down Expand Up @@ -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)
92 changes: 89 additions & 3 deletions deepspeed/compile/init_z3.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,9 @@

# DeepSpeed Team

from functools import partial
from threading import Lock

import torch

from deepspeed import comm as dist
Expand All @@ -19,6 +22,70 @@
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 _deactivate_deepcompile_on_backend_failure(engine, backend_fn):

def backend_with_failure_cleanup(*args, **kwargs):
try:
return backend_fn(*args, **kwargs)
except Exception:
if engine.is_deepcompile_active():
try:
get_deepcompile_handle().cleanup()
finally:
engine._set_deepcompile_active(False)
raise

return backend_with_failure_cleanup


def _resolve_expected_grad_dtype(param):
Expand Down Expand Up @@ -102,9 +169,28 @@ 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

try:
backend_fn = make_backend(backend,
compile_config,
compile_kwargs=compile_kwargs,
process_group=engine.data_parallel_group,
owned_frames=engine._deepcompile_owned_frames)
except Exception:
if restore_dynamo_config is not None:
restore_dynamo_config()
del engine._deepcompile_dynamo_config_restore
raise
return _deactivate_deepcompile_on_backend_failure(engine, backend_fn)
Loading
Loading