From 657bd1af1c20bf93cdcc148f55967d96e318df83 Mon Sep 17 00:00:00 2001 From: iLeGend <824040212@qq.com> Date: Mon, 17 Aug 2026 23:55:11 +0800 Subject: [PATCH 1/4] Fix comms logger KeyError when log_name is omitted * add missing log_name for all_to_all & broadcast_object_list * add fallback log_name for time_op * add ut Signed-off-by: iLeGend <824040212@qq.com> --- deepspeed/comm/comm.py | 18 ++++++++++++++++-- tests/unit/comm/test_comms_logger.py | 25 +++++++++++++++++++++++++ 2 files changed, 41 insertions(+), 2 deletions(-) diff --git a/deepspeed/comm/comm.py b/deepspeed/comm/comm.py index 437754d715c5..9d3c4804014e 100755 --- a/deepspeed/comm/comm.py +++ b/deepspeed/comm/comm.py @@ -114,6 +114,8 @@ def log_wrapper(*args, **kwargs): # 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.setdefault('log_name', func.__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() @@ -230,7 +232,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 +372,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..0b956c6b3ed9 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,26 @@ 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 From 59e91c0ce8333811bba9dce70cb72701a063fcc5 Mon Sep 17 00:00:00 2001 From: iLeGend <824040212@qq.com> Date: Mon, 24 Aug 2026 07:54:04 +0000 Subject: [PATCH 2/4] Fix comms logger prof_ops with default log_name Co-authored-by: Ehsan Barkhordar Signed-off-by: iLeGend <824040212@qq.com> --- deepspeed/comm/comm.py | 13 +++++++------ tests/unit/comm/test_comms_logger.py | 20 ++++++++++++++++++++ 2 files changed, 27 insertions(+), 6 deletions(-) diff --git a/deepspeed/comm/comm.py b/deepspeed/comm/comm.py index 7ab8db5517ff..d84e1cfb83af 100755 --- a/deepspeed/comm/comm.py +++ b/deepspeed/comm/comm.py @@ -104,18 +104,20 @@ 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): + 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) # 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): + 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.setdefault('log_name', func.__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() @@ -129,8 +131,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() diff --git a/tests/unit/comm/test_comms_logger.py b/tests/unit/comm/test_comms_logger.py index 0b956c6b3ed9..06b0bc6b602d 100644 --- a/tests/unit/comm/test_comms_logger.py +++ b/tests/unit/comm/test_comms_logger.py @@ -74,3 +74,23 @@ def barrier(): 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 From 9a92148c4b8a0b807421403c2639c269b5512d30 Mon Sep 17 00:00:00 2001 From: iLeGend <824040212@qq.com> Date: Wed, 26 Aug 2026 21:18:30 +0000 Subject: [PATCH 3/4] refactor to avoid perf regression Signed-off-by: iLeGend <824040212@qq.com> --- deepspeed/comm/comm.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/deepspeed/comm/comm.py b/deepspeed/comm/comm.py index a5a7f946e56d..e40fcc628ee4 100755 --- a/deepspeed/comm/comm.py +++ b/deepspeed/comm/comm.py @@ -107,11 +107,12 @@ def timed_op(func): default_log_name = get_default_args(func).get('log_name', func.__name__) def log_wrapper(*args, **kwargs): - 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) + should_profile = False # Add enabled flag so that overhead to each comm op is two if conditions at most if comms_logger.enabled: + 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) From 09ec337f0b2838706291f2942ccdf3851da27150 Mon Sep 17 00:00:00 2001 From: iLeGend <824040212@qq.com> Date: Tue, 8 Sep 2026 11:49:11 +0000 Subject: [PATCH 4/4] Fix collective profiling arguments and bandwidth logging Honor positional prof and log_name arguments using a cached signature without changing collective APIs or the disabled-logging fast path. Reuse bound arguments when recording operations. Add broadcast_object_list and all_to_all bandwidth calculations, with regression coverage for real wrappers, profiling selection, debug names, and world-size scaling. Signed-off-by: iLeGend <824040212@qq.com> Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- deepspeed/comm/comm.py | 13 ++-- deepspeed/utils/comms_logging.py | 4 +- tests/unit/comm/test_comms_logger.py | 104 +++++++++++++++++++++++++++ 3 files changed, 114 insertions(+), 7 deletions(-) diff --git a/deepspeed/comm/comm.py b/deepspeed/comm/comm.py index e40fcc628ee4..383d2e23d53a 100755 --- a/deepspeed/comm/comm.py +++ b/deepspeed/comm/comm.py @@ -23,6 +23,7 @@ import torch from torch.distributed import GradBucket # noqa: F401 +import inspect import os from typing import Any, Optional, TYPE_CHECKING @@ -105,18 +106,20 @@ def configure( # Logging wrapper for timing ops def timed_op(func): default_log_name = get_default_args(func).get('log_name', func.__name__) + # Cache the signature to avoid inspecting it on every communication call. + func_signature = inspect.signature(func) 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: - selected_log_name = kwargs.get('log_name', default_log_name) - should_profile = (('prof' in kwargs and kwargs['prof']) or comms_logger.prof_all + bound_args = func_signature.bind_partial(*args, **kwargs) + bound_args.apply_defaults() + func_args = bound_args.arguments + selected_log_name = func_args.get('log_name', default_log_name) + should_profile = (func_args.get('prof', False) 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) diff --git a/deepspeed/utils/comms_logging.py b/deepspeed/utils/comms_logging.py index 32ecff8bf9cc..ab76a55b8853 100644 --- a/deepspeed/utils/comms_logging.py +++ b/deepspeed/utils/comms_logging.py @@ -37,7 +37,7 @@ def calc_bw_log(comm_op, size, duration): n = dist.get_world_size() tput = 0 busbw = 0 - if comm_op == "all_to_all_single": + if comm_op == "all_to_all_single" or comm_op == "all_to_all": tput = (size / duration) busbw = (size / duration) * ((n - 1) / n) elif comm_op == "all_gather" or comm_op == "all_gather_into_tensor" or comm_op == "reduce_scatter" or comm_op == "reduce_scatter_tensor": @@ -47,7 +47,7 @@ def calc_bw_log(comm_op, size, duration): elif comm_op == "all_reduce" or comm_op == "all_reduce_coalesced" or comm_op == "inference_all_reduce": tput = (size * 2 / duration) busbw = (size / duration) * (2 * (n - 1) / n) - elif comm_op == "send" or comm_op == "recv" or comm_op == "isend" or comm_op == "irecv" or comm_op == "broadcast" or comm_op == "reduce" or comm_op == "gather" or comm_op == "scatter" or comm_op == "barrier": + elif comm_op == "send" or comm_op == "recv" or comm_op == "isend" or comm_op == "irecv" or comm_op == "broadcast" or comm_op == "broadcast_object_list" or comm_op == "reduce" or comm_op == "gather" or comm_op == "scatter" or comm_op == "barrier": tput = (size / duration) busbw = tput else: diff --git a/tests/unit/comm/test_comms_logger.py b/tests/unit/comm/test_comms_logger.py index 06b0bc6b602d..2e85864e3022 100644 --- a/tests/unit/comm/test_comms_logger.py +++ b/tests/unit/comm/test_comms_logger.py @@ -4,6 +4,10 @@ # DeepSpeed Team from types import SimpleNamespace +from unittest.mock import Mock + +import pytest +import torch from deepspeed.utils.comms_logging import CommsLogger @@ -94,3 +98,103 @@ def barrier(log_name='barrier'): assert barrier() == 'done' assert 'barrier' in comm.comms_logger.comms_dict + + +@pytest.mark.parametrize('op_name', ['broadcast_object_list', 'all_to_all']) +@pytest.mark.parametrize('positional', [False, True]) +@pytest.mark.parametrize('debug', [False, True]) +@pytest.mark.parametrize('profile_mode', ['prof_all', 'prof', 'prof_ops_default', 'prof_ops_custom', 'unselected']) +def test_collective_profiling(monkeypatch, op_name, positional, debug, profile_mode): + from deepspeed.comm import comm + + backend_op = Mock(return_value='done') + backend = SimpleNamespace(using_mpi=False, is_initialized=lambda: True, get_world_size=lambda group=None: 2) + setattr(backend, op_name, backend_op) + monkeypatch.setattr(comm, 'cdb', backend) + monkeypatch.setattr(comm, 'get_accelerator', lambda: SimpleNamespace(synchronize=lambda: None)) + op_timer = Mock() + op_timer.elapsed.return_value = 1.0 + timers = Mock(return_value=op_timer) + monkeypatch.setattr(comm, 'timers', timers) + + monkeypatch.setattr(comm, 'comms_logger', CommsLogger()) + comm.comms_logger.enabled = True + comm.comms_logger.debug = debug + comm.comms_logger.prof_all = profile_mode == 'prof_all' + log_name = 'custom_collective' if profile_mode in ('prof', 'prof_ops_custom') else op_name + comm.comms_logger.prof_ops = [log_name] if profile_mode.startswith('prof_ops') else [] + + if op_name == 'broadcast_object_list': + object_list = [{'value': 1}] + args = (object_list, 0, None, None) + expected_size = 0 + else: + input_list = [torch.ones(4), torch.ones(4)] + output_list = [torch.empty_like(tensor) for tensor in input_list] + args = (output_list, input_list, None, False) + expected_size = sum(tensor.numel() * tensor.element_size() for tensor in input_list) + + kwargs = {} + if profile_mode in ('prof', 'prof_ops_custom'): + prof = profile_mode == 'prof' + if positional: + args += (prof, log_name) + else: + kwargs = {'prof': prof, 'log_name': log_name} + + assert getattr(comm, op_name)(*args, **kwargs) == 'done' + if op_name == 'broadcast_object_list': + backend_op.assert_called_once_with(object_list=object_list, src=0, group=None, device=None) + else: + backend_op.assert_called_once_with(output_list, input_list, group=None, async_op=False) + + if profile_mode == 'unselected': + assert comm.comms_logger.comms_dict == {} + timers.assert_not_called() + else: + record_name, = comm.comms_logger.comms_dict + if debug: + assert record_name.startswith(log_name + ' | [Caller Func: ') + else: + assert record_name == log_name + record = comm.comms_logger.comms_dict[record_name][expected_size] + assert record[0] == 1 + assert record[1] == [1.0] + op_timer.start.assert_called_once_with() + op_timer.stop.assert_called_once_with() + op_timer.elapsed.assert_called_once_with(reset=False) + assert all(call.args == (record_name, ) for call in timers.call_args_list) + + +def test_timed_op_disabled_does_not_access_profiling_state(monkeypatch): + from deepspeed.comm import comm + + @comm.timed_op + def barrier(prof=False, log_name='barrier'): + return 'done' + + # No other logger attributes or backend should be needed on the disabled path. + monkeypatch.setattr(comm, 'comms_logger', SimpleNamespace(enabled=False)) + monkeypatch.setattr(comm, 'cdb', None) + timers = Mock() + monkeypatch.setattr(comm, 'timers', timers) + + assert barrier(True, 'custom_barrier') == 'done' + timers.assert_not_called() + + +@pytest.mark.parametrize('world_size', [1, 2, 4]) +@pytest.mark.parametrize('comm_op', ['broadcast_object_list', 'all_to_all']) +def test_calc_bw_log_supports_object_and_list_collectives(monkeypatch, comm_op, world_size): + import deepspeed.comm as dist + from deepspeed.utils.comms_logging import calc_bw_log + + monkeypatch.setattr(dist, 'get_world_size', lambda group=None: world_size) + + tput, busbw = calc_bw_log(comm_op, 1024, 2.0) + expected_tput = 1024 / 2.0 * 8 / 1e6 + expected_busbw = expected_tput + if comm_op == 'all_to_all': + expected_busbw *= (world_size - 1) / world_size + assert tput == pytest.approx(expected_tput) + assert busbw == pytest.approx(expected_busbw)