Fix comms logger KeyError when log_name is omitted - #8267
Conversation
|
Codex usage limits have been reached for code reviews. Please check with the admins of this repo to increase the limits by adding credits. |
7dc2dcc to
657bd1a
Compare
* 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>
ebarkhordar
left a comment
There was a problem hiding this comment.
prof_ops never matches an op that takes its log_name from the signature default. The gate at comm.py:111-113 tests 'log_name' in kwargs, so it fires only when a caller passes the name explicitly, while config-json.md documents "prof_ops": ["all_reduce", "all_gather"] against ordinary calls.
At 1f95164 in a clean container, CPU torch and a stub cdb:
prof_ops = ['all_reduce'] prof_all = False
A. dist.all_reduce(t) -> comms_dict keys: []
B. dist.all_reduce(t, log_name=...) -> comms_dict keys: ['all_reduce']
Your setdefault is one step from covering this. Resolving the name once per op keeps the per-call fast path at two conditions:
def timed_op(func):
default_log_name = get_default_args(func).get('log_name', func.__name__)
def log_wrapper(*args, **kwargs):
if comms_logger.enabled:
selected = kwargs.get('log_name', default_log_name)
if kwargs.get('prof') or comms_logger.prof_all or selected in comms_logger.prof_ops:then func_args['log_name'] = selected in place of the setdefault, and the same condition in the finally gate. With that, A logs and tests/unit/comm/test_comms_logger.py is still 4 passed. It is a separate bug from the KeyError you are fixing, so it may belong in its own PR.
Wow, that's a very insightful observation! I agree that |
|
Your call as the author, but I would take it in this PR. The repair replaces the Either way I am not going to open a competing PR for it. One thing to keep if you do take it: the I re-read |
Summary
Fix a
KeyError: 'log_name'raised by the DeepSpeed communication loggerwhen a wrapped collective is called without an explicit
log_name.This is exposed by multi-rank AutoTP input consistency checks, which call
broadcast_object_listwithout passing profiling metadata. Single-rank TPdoes not exercise this communication path.
Changes
broadcast_object_listandall_to_allfunc.__name__as the defaultlog_nameto cover missing statusValidation
python -m pytest -q tests/unit/comm/test_comms_logger.py