From f2fde43ecdb16d8f2986b565913823c968904145 Mon Sep 17 00:00:00 2001 From: Masahiro Tanaka Date: Mon, 24 Aug 2026 11:12:58 -0700 Subject: [PATCH] Forward barrier device_ids to communication backends Signed-off-by: Masahiro Tanaka --- deepspeed/comm/ccl.py | 2 +- deepspeed/comm/comm.py | 2 +- tests/unit/comm/test_barrier.py | 37 +++++++++++++++++++++++++++++++++ 3 files changed, 39 insertions(+), 2 deletions(-) create mode 100644 tests/unit/comm/test_barrier.py diff --git a/deepspeed/comm/ccl.py b/deepspeed/comm/ccl.py index e95e40a03087..3e3a42855f44 100644 --- a/deepspeed/comm/ccl.py +++ b/deepspeed/comm/ccl.py @@ -134,7 +134,7 @@ def gather(self, tensor, gather_list, dst, group=None, async_op=False): def scatter(self, tensor, gather_list, dst, group=None, async_op=False): return self.run_collective(name="scatter", tensor=tensor, gather_list=gather_list, dst=dst, group=group) - def barrier(self, group=None, async_op=False): + def barrier(self, group=None, async_op=False, device_ids=None): return self.run_collective(name="barrier", group=group, async_op=async_op) def monitored_barrier(self, group=None, timeout=None, wait_all_ranks=False): diff --git a/deepspeed/comm/comm.py b/deepspeed/comm/comm.py index 468166e35bba..635f64fef8b8 100755 --- a/deepspeed/comm/comm.py +++ b/deepspeed/comm/comm.py @@ -422,7 +422,7 @@ def scatter(tensor, @timed_op def barrier(group=None, async_op=False, device_ids=None, prof=False, log_name='barrier', debug=get_caller_func()): global cdb - return cdb.barrier(group=group, async_op=async_op) + return cdb.barrier(group=group, async_op=async_op, device_ids=device_ids) @timed_op diff --git a/tests/unit/comm/test_barrier.py b/tests/unit/comm/test_barrier.py new file mode 100644 index 000000000000..2db466fbc7eb --- /dev/null +++ b/tests/unit/comm/test_barrier.py @@ -0,0 +1,37 @@ +# Copyright (c) DeepSpeed Team. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +from unittest.mock import MagicMock, sentinel + +import pytest + +import deepspeed.comm.comm as comm +from deepspeed.comm.ccl import CCLBackend + + +@pytest.mark.parametrize("call_kwargs", [{}, {"device_ids": [0]}]) +def test_public_barrier_forwards_device_ids(monkeypatch, call_kwargs): + backend = MagicMock() + backend.barrier.return_value = sentinel.work + monkeypatch.setattr(comm, "cdb", backend) + monkeypatch.setattr(comm.comms_logger, "enabled", False) + expected_device_ids = call_kwargs.get("device_ids") + + result = comm.barrier(**call_kwargs) + + assert result is sentinel.work + backend.barrier.assert_called_once_with(group=None, async_op=False, device_ids=expected_device_ids) + assert backend.barrier.call_args.kwargs["device_ids"] is expected_device_ids + + +@pytest.mark.parametrize("call_kwargs", [{}, {"device_ids": [0]}]) +def test_ccl_barrier_accepts_device_ids(call_kwargs): + backend = CCLBackend.__new__(CCLBackend) + backend.run_collective = MagicMock(return_value=sentinel.work) + + result = backend.barrier(**call_kwargs) + + assert result is sentinel.work + backend.run_collective.assert_called_once_with(name="barrier", group=None, async_op=False)