diff --git a/deepspeed/comm/comm.py b/deepspeed/comm/comm.py index 468166e35bba..e40fcc628ee4 100755 --- a/deepspeed/comm/comm.py +++ b/deepspeed/comm/comm.py @@ -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() @@ -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() @@ -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) @@ -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) diff --git a/tests/unit/comm/test_comms_logger.py b/tests/unit/comm/test_comms_logger.py index 7573fcbf43c7..06b0bc6b602d 100644 --- a/tests/unit/comm/test_comms_logger.py +++ b/tests/unit/comm/test_comms_logger.py @@ -3,6 +3,8 @@ # DeepSpeed Team +from types import SimpleNamespace + from deepspeed.utils.comms_logging import CommsLogger @@ -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