Skip to content

Commit bd6e968

Browse files
jinyouzhisfc-gh-truwaseebarkhordarCopilottohtana
authored
Fix comms logger KeyError when log_name is omitted (#8267)
## Summary Fix a `KeyError: 'log_name'` raised by the DeepSpeed communication logger when a wrapped collective is called without an explicit `log_name`. ``` "comms_logger": { "enabled": true, "prof_all": true, "debug": true } ``` This is exposed by multi-rank AutoTP input consistency checks, which call `broadcast_object_list` without passing profiling metadata. Single-rank TP does not exercise this communication path. ## Changes - Add prof/log_name/debug for `broadcast_object_list` and `all_to_all` - Use the `func.__name__` as the default `log_name` to cover missing status - Add a regression test ## Validation - `python -m pytest -q tests/unit/comm/test_comms_logger.py` --------- Signed-off-by: iLeGend <824040212@qq.com> Co-authored-by: Olatunji Ruwase <tunji.ruwase@snowflake.com> Co-authored-by: Ehsan Barkhordar <realbarkhordar@gmail.com> Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Co-authored-by: Masahiro Tanaka <81312776+tohtana@users.noreply.github.com>
1 parent 1190946 commit bd6e968

3 files changed

Lines changed: 180 additions & 12 deletions

File tree

deepspeed/comm/comm.py

Lines changed: 29 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@
2323

2424
import torch
2525
from torch.distributed import GradBucket # noqa: F401
26+
import inspect
2627
import os
2728
from typing import Any, Optional, TYPE_CHECKING
2829

@@ -104,16 +105,23 @@ def configure(
104105

105106
# Logging wrapper for timing ops
106107
def timed_op(func):
108+
default_log_name = get_default_args(func).get('log_name', func.__name__)
109+
# Cache the signature to avoid inspecting it on every communication call.
110+
func_signature = inspect.signature(func)
107111

108112
def log_wrapper(*args, **kwargs):
113+
should_profile = False
109114
# Add enabled flag so that overhead to each comm op is two if conditions at most
110115
if comms_logger.enabled:
111-
if ('prof' in kwargs
112-
and kwargs['prof']) or comms_logger.prof_all or ('log_name' in kwargs
113-
and kwargs['log_name'] in comms_logger.prof_ops):
114-
# Need func args for their defaults
115-
func_args = get_default_args(func)
116-
func_args.update(kwargs)
116+
bound_args = func_signature.bind_partial(*args, **kwargs)
117+
bound_args.apply_defaults()
118+
func_args = bound_args.arguments
119+
selected_log_name = func_args.get('log_name', default_log_name)
120+
should_profile = (func_args.get('prof', False) or comms_logger.prof_all
121+
or selected_log_name in comms_logger.prof_ops)
122+
if should_profile:
123+
# Ops that do not declare a log_name are logged under their own name
124+
func_args['log_name'] = selected_log_name
117125
msg_size = get_msg_size_from_args(func, *args, **kwargs)
118126
log_name = get_debug_log_name(func_args, comms_logger.debug)
119127
timers(log_name).start()
@@ -127,8 +135,7 @@ def log_wrapper(*args, **kwargs):
127135
# If we're using MPI, we can't simply sync the stream
128136
if cdb.using_mpi:
129137
cdb.barrier()
130-
if ('prof' in kwargs and kwargs['prof']) or comms_logger.prof_all or (
131-
'log_name' in kwargs and kwargs['log_name'] in comms_logger.prof_ops):
138+
if should_profile:
132139
log_name = get_debug_log_name(func_args, comms_logger.debug)
133140
raw_name = func.__name__
134141
timers(log_name).stop()
@@ -230,7 +237,13 @@ def broadcast(tensor, src, group=None, async_op=False, prof=False, log_name='bro
230237

231238

232239
@timed_op
233-
def broadcast_object_list(object_list, src, group=None, device=None):
240+
def broadcast_object_list(object_list,
241+
src,
242+
group=None,
243+
device=None,
244+
prof=False,
245+
log_name='broadcast_object_list',
246+
debug=get_caller_func()):
234247
global cdb
235248
return cdb.broadcast_object_list(object_list=object_list, src=src, group=group, device=device)
236249

@@ -364,7 +377,13 @@ def all_to_all_single(output,
364377

365378

366379
@timed_op
367-
def all_to_all(output_tensor_list, input_tensor_list, group=None, async_op=False):
380+
def all_to_all(output_tensor_list,
381+
input_tensor_list,
382+
group=None,
383+
async_op=False,
384+
prof=False,
385+
log_name='all_to_all',
386+
debug=get_caller_func()):
368387
global cdb
369388
return cdb.all_to_all(output_tensor_list, input_tensor_list, group=group, async_op=async_op)
370389

deepspeed/utils/comms_logging.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -37,7 +37,7 @@ def calc_bw_log(comm_op, size, duration):
3737
n = dist.get_world_size()
3838
tput = 0
3939
busbw = 0
40-
if comm_op == "all_to_all_single":
40+
if comm_op == "all_to_all_single" or comm_op == "all_to_all":
4141
tput = (size / duration)
4242
busbw = (size / duration) * ((n - 1) / n)
4343
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):
4747
elif comm_op == "all_reduce" or comm_op == "all_reduce_coalesced" or comm_op == "inference_all_reduce":
4848
tput = (size * 2 / duration)
4949
busbw = (size / duration) * (2 * (n - 1) / n)
50-
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":
50+
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":
5151
tput = (size / duration)
5252
busbw = tput
5353
else:

tests/unit/comm/test_comms_logger.py

Lines changed: 149 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,12 @@
33

44
# DeepSpeed Team
55

6+
from types import SimpleNamespace
7+
from unittest.mock import Mock
8+
9+
import pytest
10+
import torch
11+
612
from deepspeed.utils.comms_logging import CommsLogger
713

814

@@ -49,3 +55,146 @@ def test_trim_mean_does_not_mutate_its_argument():
4955
data = [3.0, 1.0, 2.0]
5056
assert trim_mean(data, 0.1) == 2.0
5157
assert data == [3.0, 1.0, 2.0]
58+
59+
60+
def test_timed_op_falls_back_to_the_op_name_when_log_name_is_missing(monkeypatch):
61+
# timed_op looks up func_args['log_name'], so an op whose signature does not
62+
# declare log_name used to raise KeyError as soon as profiling was turned on.
63+
# Such an op must still be logged, under its own name.
64+
from deepspeed.comm import comm
65+
66+
monkeypatch.setattr(comm, 'comms_logger', CommsLogger())
67+
monkeypatch.setattr(
68+
comm, 'cdb', SimpleNamespace(using_mpi=False, is_initialized=lambda: True,
69+
get_world_size=lambda group=None: 1))
70+
monkeypatch.setattr(comm, 'get_accelerator', lambda: SimpleNamespace(synchronize=lambda: None))
71+
72+
@comm.timed_op
73+
def barrier():
74+
return 'done'
75+
76+
comm.comms_logger.enabled = True
77+
comm.comms_logger.start_profiling_comms()
78+
79+
assert barrier() == 'done'
80+
assert 'barrier' in comm.comms_logger.comms_dict
81+
82+
83+
def test_timed_op_profiles_default_log_name_with_prof_ops(monkeypatch):
84+
from deepspeed.comm import comm
85+
86+
monkeypatch.setattr(comm, 'comms_logger', CommsLogger())
87+
monkeypatch.setattr(
88+
comm, 'cdb', SimpleNamespace(using_mpi=False, is_initialized=lambda: True,
89+
get_world_size=lambda group=None: 1))
90+
monkeypatch.setattr(comm, 'get_accelerator', lambda: SimpleNamespace(synchronize=lambda: None))
91+
92+
@comm.timed_op
93+
def barrier(log_name='barrier'):
94+
return 'done'
95+
96+
comm.comms_logger.enabled = True
97+
comm.comms_logger.prof_ops = ['barrier']
98+
99+
assert barrier() == 'done'
100+
assert 'barrier' in comm.comms_logger.comms_dict
101+
102+
103+
@pytest.mark.parametrize('op_name', ['broadcast_object_list', 'all_to_all'])
104+
@pytest.mark.parametrize('positional', [False, True])
105+
@pytest.mark.parametrize('debug', [False, True])
106+
@pytest.mark.parametrize('profile_mode', ['prof_all', 'prof', 'prof_ops_default', 'prof_ops_custom', 'unselected'])
107+
def test_collective_profiling(monkeypatch, op_name, positional, debug, profile_mode):
108+
from deepspeed.comm import comm
109+
110+
backend_op = Mock(return_value='done')
111+
backend = SimpleNamespace(using_mpi=False, is_initialized=lambda: True, get_world_size=lambda group=None: 2)
112+
setattr(backend, op_name, backend_op)
113+
monkeypatch.setattr(comm, 'cdb', backend)
114+
monkeypatch.setattr(comm, 'get_accelerator', lambda: SimpleNamespace(synchronize=lambda: None))
115+
op_timer = Mock()
116+
op_timer.elapsed.return_value = 1.0
117+
timers = Mock(return_value=op_timer)
118+
monkeypatch.setattr(comm, 'timers', timers)
119+
120+
monkeypatch.setattr(comm, 'comms_logger', CommsLogger())
121+
comm.comms_logger.enabled = True
122+
comm.comms_logger.debug = debug
123+
comm.comms_logger.prof_all = profile_mode == 'prof_all'
124+
log_name = 'custom_collective' if profile_mode in ('prof', 'prof_ops_custom') else op_name
125+
comm.comms_logger.prof_ops = [log_name] if profile_mode.startswith('prof_ops') else []
126+
127+
if op_name == 'broadcast_object_list':
128+
object_list = [{'value': 1}]
129+
args = (object_list, 0, None, None)
130+
expected_size = 0
131+
else:
132+
input_list = [torch.ones(4), torch.ones(4)]
133+
output_list = [torch.empty_like(tensor) for tensor in input_list]
134+
args = (output_list, input_list, None, False)
135+
expected_size = sum(tensor.numel() * tensor.element_size() for tensor in input_list)
136+
137+
kwargs = {}
138+
if profile_mode in ('prof', 'prof_ops_custom'):
139+
prof = profile_mode == 'prof'
140+
if positional:
141+
args += (prof, log_name)
142+
else:
143+
kwargs = {'prof': prof, 'log_name': log_name}
144+
145+
assert getattr(comm, op_name)(*args, **kwargs) == 'done'
146+
if op_name == 'broadcast_object_list':
147+
backend_op.assert_called_once_with(object_list=object_list, src=0, group=None, device=None)
148+
else:
149+
backend_op.assert_called_once_with(output_list, input_list, group=None, async_op=False)
150+
151+
if profile_mode == 'unselected':
152+
assert comm.comms_logger.comms_dict == {}
153+
timers.assert_not_called()
154+
else:
155+
record_name, = comm.comms_logger.comms_dict
156+
if debug:
157+
assert record_name.startswith(log_name + ' | [Caller Func: ')
158+
else:
159+
assert record_name == log_name
160+
record = comm.comms_logger.comms_dict[record_name][expected_size]
161+
assert record[0] == 1
162+
assert record[1] == [1.0]
163+
op_timer.start.assert_called_once_with()
164+
op_timer.stop.assert_called_once_with()
165+
op_timer.elapsed.assert_called_once_with(reset=False)
166+
assert all(call.args == (record_name, ) for call in timers.call_args_list)
167+
168+
169+
def test_timed_op_disabled_does_not_access_profiling_state(monkeypatch):
170+
from deepspeed.comm import comm
171+
172+
@comm.timed_op
173+
def barrier(prof=False, log_name='barrier'):
174+
return 'done'
175+
176+
# No other logger attributes or backend should be needed on the disabled path.
177+
monkeypatch.setattr(comm, 'comms_logger', SimpleNamespace(enabled=False))
178+
monkeypatch.setattr(comm, 'cdb', None)
179+
timers = Mock()
180+
monkeypatch.setattr(comm, 'timers', timers)
181+
182+
assert barrier(True, 'custom_barrier') == 'done'
183+
timers.assert_not_called()
184+
185+
186+
@pytest.mark.parametrize('world_size', [1, 2, 4])
187+
@pytest.mark.parametrize('comm_op', ['broadcast_object_list', 'all_to_all'])
188+
def test_calc_bw_log_supports_object_and_list_collectives(monkeypatch, comm_op, world_size):
189+
import deepspeed.comm as dist
190+
from deepspeed.utils.comms_logging import calc_bw_log
191+
192+
monkeypatch.setattr(dist, 'get_world_size', lambda group=None: world_size)
193+
194+
tput, busbw = calc_bw_log(comm_op, 1024, 2.0)
195+
expected_tput = 1024 / 2.0 * 8 / 1e6
196+
expected_busbw = expected_tput
197+
if comm_op == 'all_to_all':
198+
expected_busbw *= (world_size - 1) / world_size
199+
assert tput == pytest.approx(expected_tput)
200+
assert busbw == pytest.approx(expected_busbw)

0 commit comments

Comments
 (0)