diff --git a/deepspeed/runtime/zero/stage_1_and_2.py b/deepspeed/runtime/zero/stage_1_and_2.py index f05a53867c93..028dc0aa7136 100644 --- a/deepspeed/runtime/zero/stage_1_and_2.py +++ b/deepspeed/runtime/zero/stage_1_and_2.py @@ -64,14 +64,16 @@ def input(msg): def split_half_float_double(tensors): - device_type = get_accelerator().device_name() - dtypes = [ - "torch.{}.HalfTensor".format(device_type), "torch.{}.FloatTensor".format(device_type), - "torch.{}.DoubleTensor".format(device_type), "torch.{}.BFloat16Tensor".format(device_type) - ] + # Legacy type strings omit the device prefix on CPU ("torch.FloatTensor"), so + # building them from the accelerator device name matches nothing on CPU and + # silently drops every gradient bucket, skipping the all-reduce entirely. + # Compare dtypes directly so the buckets are device-independent. Only strided + # tensors can be flattened into the dense all-reduce buffer, so every sparse + # layout stays excluded. + dtypes = [torch.half, torch.float, torch.double, torch.bfloat16] buckets = [] for i, dtype in enumerate(dtypes): - bucket = [t for t in tensors if t.type() == dtype] + bucket = [t for t in tensors if t.dtype == dtype and t.layout == torch.strided] if bucket: buckets.append(bucket) return buckets diff --git a/tests/unit/v1/zero/test_zero.py b/tests/unit/v1/zero/test_zero.py index a56d873d8408..131e60ea945e 100644 --- a/tests/unit/v1/zero/test_zero.py +++ b/tests/unit/v1/zero/test_zero.py @@ -24,12 +24,35 @@ from deepspeed.runtime.engine import DeepSpeedEngine from deepspeed.runtime.bf16_optimizer import BF16_Optimizer from deepspeed.runtime.zero.partition_parameters import ZeroParamStatus +from deepspeed.runtime.zero.stage_1_and_2 import split_half_float_double from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint from deepspeed.runtime.zero.utils import ZeRORuntimeException from deepspeed.accelerator import get_accelerator from deepspeed.utils import safe_get_full_fp32_param, safe_get_full_grad +class TestSplitHalfFloatDouble: + + def test_device_independent_buckets_exclude_sparse(self): + # Pins two fixed membership bugs: the legacy accelerator-prefixed type + # strings matched nothing on CPU, silently dropping every bucket, and + # dtype-only matching would admit sparse layouts, which cannot be + # flattened into a dense all-reduce buffer. The CSR sample also pins that + # the exclusion covers layouts where is_sparse is False. + dense_grads = [ + torch.zeros(2, dtype=dtype) for dtype in (torch.half, torch.float, torch.double, torch.bfloat16) + ] + sparse_grad = torch.sparse_coo_tensor(torch.tensor([[0]]), torch.tensor([1.0]), (1, )) + csr_grad = torch.sparse_csr_tensor(torch.tensor([0, 1]), torch.tensor([0]), torch.tensor([1.0]), (1, 1)) + + buckets = split_half_float_double(dense_grads + [sparse_grad, csr_grad]) + + assert len(buckets) == 4 + for bucket, grad in zip(buckets, dense_grads): + assert len(bucket) == 1 + assert bucket[0] is grad + + @pytest.mark.parametrize("zero_stage", [0, 1, 2]) @pytest.mark.parametrize("gradient_allreduce_op,expected_scale", [("mean", 1.0), ("sum", 2.0)]) @pytest.mark.parametrize(