Skip to content
Open
30 changes: 23 additions & 7 deletions deepspeed/comm/comm.py
Original file line number Diff line number Diff line change
Expand Up @@ -104,16 +104,21 @@ def configure(

# Logging wrapper for timing ops
def timed_op(func):
default_log_name = get_default_args(func).get('log_name', func.__name__)

def log_wrapper(*args, **kwargs):
should_profile = False
# Add enabled flag so that overhead to each comm op is two if conditions at most
if comms_logger.enabled:
if ('prof' in kwargs
and kwargs['prof']) or comms_logger.prof_all or ('log_name' in kwargs
and kwargs['log_name'] in comms_logger.prof_ops):
selected_log_name = kwargs.get('log_name', default_log_name)
should_profile = (('prof' in kwargs and kwargs['prof']) or comms_logger.prof_all
or selected_log_name in comms_logger.prof_ops)
if should_profile:
# Need func args for their defaults
func_args = get_default_args(func)
func_args.update(kwargs)
# Ops that do not declare a log_name are logged under their own name
func_args['log_name'] = selected_log_name
msg_size = get_msg_size_from_args(func, *args, **kwargs)
log_name = get_debug_log_name(func_args, comms_logger.debug)
timers(log_name).start()
Expand All @@ -127,8 +132,7 @@ def log_wrapper(*args, **kwargs):
# If we're using MPI, we can't simply sync the stream
if cdb.using_mpi:
cdb.barrier()
if ('prof' in kwargs and kwargs['prof']) or comms_logger.prof_all or (
'log_name' in kwargs and kwargs['log_name'] in comms_logger.prof_ops):
if should_profile:
log_name = get_debug_log_name(func_args, comms_logger.debug)
raw_name = func.__name__
timers(log_name).stop()
Expand Down Expand Up @@ -230,7 +234,13 @@ def broadcast(tensor, src, group=None, async_op=False, prof=False, log_name='bro


@timed_op
def broadcast_object_list(object_list, src, group=None, device=None):
def broadcast_object_list(object_list,
src,
group=None,
device=None,
prof=False,
log_name='broadcast_object_list',
debug=get_caller_func()):
global cdb
return cdb.broadcast_object_list(object_list=object_list, src=src, group=group, device=device)

Expand Down Expand Up @@ -364,7 +374,13 @@ def all_to_all_single(output,


@timed_op
def all_to_all(output_tensor_list, input_tensor_list, group=None, async_op=False):
def all_to_all(output_tensor_list,
input_tensor_list,
group=None,
async_op=False,
prof=False,
log_name='all_to_all',
debug=get_caller_func()):
global cdb
return cdb.all_to_all(output_tensor_list, input_tensor_list, group=group, async_op=async_op)

Expand Down
45 changes: 45 additions & 0 deletions tests/unit/comm/test_comms_logger.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,8 @@

# DeepSpeed Team

from types import SimpleNamespace

from deepspeed.utils.comms_logging import CommsLogger


Expand Down Expand Up @@ -49,3 +51,46 @@ def test_trim_mean_does_not_mutate_its_argument():
data = [3.0, 1.0, 2.0]
assert trim_mean(data, 0.1) == 2.0
assert data == [3.0, 1.0, 2.0]


def test_timed_op_falls_back_to_the_op_name_when_log_name_is_missing(monkeypatch):
# timed_op looks up func_args['log_name'], so an op whose signature does not
# declare log_name used to raise KeyError as soon as profiling was turned on.
# Such an op must still be logged, under its own name.
from deepspeed.comm import comm

monkeypatch.setattr(comm, 'comms_logger', CommsLogger())
monkeypatch.setattr(
comm, 'cdb', SimpleNamespace(using_mpi=False, is_initialized=lambda: True,
get_world_size=lambda group=None: 1))
monkeypatch.setattr(comm, 'get_accelerator', lambda: SimpleNamespace(synchronize=lambda: None))

@comm.timed_op
def barrier():
return 'done'

comm.comms_logger.enabled = True
comm.comms_logger.start_profiling_comms()

assert barrier() == 'done'
assert 'barrier' in comm.comms_logger.comms_dict


def test_timed_op_profiles_default_log_name_with_prof_ops(monkeypatch):
from deepspeed.comm import comm

monkeypatch.setattr(comm, 'comms_logger', CommsLogger())
monkeypatch.setattr(
comm, 'cdb', SimpleNamespace(using_mpi=False, is_initialized=lambda: True,
get_world_size=lambda group=None: 1))
monkeypatch.setattr(comm, 'get_accelerator', lambda: SimpleNamespace(synchronize=lambda: None))

@comm.timed_op
def barrier(log_name='barrier'):
return 'done'

comm.comms_logger.enabled = True
comm.comms_logger.prof_ops = ['barrier']

assert barrier() == 'done'
assert 'barrier' in comm.comms_logger.comms_dict
Loading