Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 9 additions & 4 deletions deepspeed/runtime/engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@
from threading import Lock
from weakref import ref

from typing import Callable, Dict, Union, Iterable, Container, List
from typing import Callable, Dict, Union, Iterable, Container, List, Optional

import deepspeed

Expand Down Expand Up @@ -1041,15 +1041,16 @@ def set_custom_curriculum_learning_schedule(self, schedule_func_dict):
if self.training_dataloader is not None and self.curriculum_learning_enabled():
self.training_dataloader.data_sampler.set_custom_curriculum_learning_schedule(schedule_func_dict)

def get_global_grad_norm(self) -> float:
def get_global_grad_norm(self) -> Optional[float]:
"""Return the 2-norm of all gradients. If there is model parallelism,
the norm will be global.
The computed norm will be cached and reused until the next step() pass.
Returns ``None`` when ZeRO Stage 1/2 gradient-norm computation is disabled.
.. note::
In the presence of model parallelism, this is a collective call
and acts as a barrier among ``mpu.get_model_parallel_group()``.
Returns:
float: norm
Optional[float]: norm, or ``None`` when disabled
"""
return self._global_grad_norm

Expand Down Expand Up @@ -1369,6 +1370,9 @@ def zero_multi_rank_bucket_allreduce(self):
def zero_allgather_bucket_size(self):
return self._config.zero_config.allgather_bucket_size

def zero_compute_grad_norm(self):
return self._config.zero_config.compute_grad_norm

def zero_optimization_partition_gradients(self):
return self.zero_optimization_stage() >= ZeroStageEnum.gradients

Expand Down Expand Up @@ -2518,7 +2522,8 @@ def _configure_zero_optimizer(self, optimizer):
gradient_accumulation_dtype=gradient_accumulation_dtype,
communication_data_type=self.communication_data_type,
elastic_checkpoint=self.zero_elastic_checkpoint(),
check_grad_overflow=check_grad_overflow)
check_grad_overflow=check_grad_overflow,
compute_grad_norm=self.zero_compute_grad_norm())

elif zero_stage == ZeroStageEnum.weights:
self._validate_zero3_moe_compatibility()
Expand Down
12 changes: 12 additions & 0 deletions deepspeed/runtime/zero/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -140,6 +140,12 @@ class DeepSpeedZeroConfig(DeepSpeedConfigModel):
Attempts to overlap the reduction of the gradients with backward computation
"""

compute_grad_norm: bool = True
"""
Compute and retain the global gradient norm during ZeRO Stage 1/2 optimizer steps.
Disable only when gradient clipping is off and callers do not use ``get_global_grad_norm()``.
"""

load_from_fp32_weights: bool = True
"""
Boolean indicating whether to initialize fp32 master weights from fp32
Expand Down Expand Up @@ -385,6 +391,12 @@ def overlap_comm_valid(self):
self.overlap_comm = self.stage == ZeroStageEnum.weights
return self

@model_validator(mode="after")
def compute_grad_norm_valid(self):
if not self.compute_grad_norm and self.stage not in (ZeroStageEnum.optimizer_states, ZeroStageEnum.gradients):
raise ValueError("compute_grad_norm=false is supported only with ZeRO Stage 1 or 2")
return self

@model_validator(mode="after")
def offload_ratio_check(self):
offload_config = self.offload_optimizer
Expand Down
35 changes: 28 additions & 7 deletions deepspeed/runtime/zero/stage_1_and_2.py
Original file line number Diff line number Diff line change
Expand Up @@ -180,10 +180,19 @@ def __init__(self,
bf16_master_weights_and_gradients=False,
bf16_optimizer_states=False,
elastic_checkpoint=False,
check_grad_overflow=True):
check_grad_overflow=True,
compute_grad_norm=True):

super().__init__()

if not compute_grad_norm and clip_grad > 0.0:
raise ValueError("zero_optimization.compute_grad_norm=false requires gradient_clipping=0")
if not compute_grad_norm and zenflow_config is not None:
raise ValueError("zero_optimization.compute_grad_norm=false does not support ZenFlow")
if (not compute_grad_norm and offload_optimizer_config is not None
and offload_optimizer_config.device != OffloadDeviceEnum.none):
raise ValueError("zero_optimization.compute_grad_norm=false does not support optimizer offload")

if offload_optimizer_config is not None and offload_optimizer_config.device != OffloadDeviceEnum.none:
self.cpu_offload = True
self.cpu_offload_pin_memory = offload_optimizer_config.pin_memory
Expand All @@ -205,6 +214,7 @@ def __init__(self,

self.elastic_checkpoint = elastic_checkpoint
self.check_grad_overflow = check_grad_overflow
self.compute_grad_norm = compute_grad_norm
self.param_names = param_names
self.mpu = mpu
# differences from apex.fp16_utils:
Expand Down Expand Up @@ -2285,6 +2295,9 @@ def step(self, closure=None):
if self.cpu_offload:
self._offload_accumulated_param_ids = set()

if not self.compute_grad_norm:
self._global_grad_norm = None

see_memory_usage("In step before checking overflow")

# First compute norm for all group so we know if there is overflow
Expand All @@ -2311,11 +2324,12 @@ def step(self, closure=None):
self.timers(timer).stop()
return

# Step 1:- Calculate gradient norm using bit-16 grads
see_memory_usage('Before norm calculation')
scaled_global_grad_norm = self.scaled_global_norm()
self._global_grad_norm = scaled_global_grad_norm / prev_scale
see_memory_usage('After norm before optimizer')
scaled_global_grad_norm = None
if self.compute_grad_norm:
see_memory_usage('Before norm calculation')
scaled_global_grad_norm = self.scaled_global_norm()
self._global_grad_norm = scaled_global_grad_norm / prev_scale
see_memory_usage('After norm before optimizer')

# Step 2:- run optimizer and upscaling simultaneously
for i, group in enumerate(self.bit16_groups):
Expand Down Expand Up @@ -2433,6 +2447,9 @@ def _average_expert_grad_norms(self, norm_groups):
norm_groups[i] = scaled_norm_tensor.to(self.device)

def unscale_and_clip_grads(self, grad_groups_flat, total_norm):
if self.clip_grad == 0.0 and self.loss_scale == 1.0:
return

# compute combined scale factor for this group
combined_scale = self.loss_scale
if self.clip_grad > 0.:
Expand Down Expand Up @@ -2761,7 +2778,11 @@ def _load_global_state(self, sd):
self.loss_scaler = sd.get(LOSS_SCALER, self.loss_scaler)
self.dynamic_loss_scale = sd.get('dynamic_loss_scale', self.dynamic_loss_scale)
self.overflow = sd.get('overflow', self.overflow)
self.clip_grad = sd.get(CLIP_GRAD, self.clip_grad)
checkpoint_clip_grad = sd.get(CLIP_GRAD, self.clip_grad)
if not self.compute_grad_norm and checkpoint_clip_grad > 0.0:
raise ValueError("Cannot load a checkpoint with gradient clipping into "
"zero_optimization.compute_grad_norm=false")
self.clip_grad = checkpoint_clip_grad

ckpt_version = sd.get(DS_VERSION, False)
assert ckpt_version, "Empty ds_version in checkpoint, not clear how to proceed"
Expand Down
7 changes: 7 additions & 0 deletions docs/_pages/config-json.md
Original file line number Diff line number Diff line change
Expand Up @@ -461,6 +461,7 @@ Enabling and configuring ZeRO memory optimizations
"stage": [0|1|2|3],
"allgather_partitions": [true|false],
"allgather_bucket_size": 5e8,
"compute_grad_norm": [true|false],
"overlap_comm": false,
"reduce_scatter": [true|false],
"reduce_bucket_size": 5e8,
Expand Down Expand Up @@ -511,6 +512,12 @@ Enabling and configuring ZeRO memory optimizations
| ------------------------------------------------------------------------------------------------------------ | ------- |
| Number of elements allgathered at a time. Limits the memory required for the allgather for large model sizes | `5e8` |

***compute_grad_norm***: [boolean]

| Description | Default |
| ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | ------- |
| Compute and retain the global gradient norm during ZeRO Stage 1/2 optimizer steps. Set to `false` only with a GPU optimizer, without ZenFlow or gradient clipping, and when callers do not use `get_global_grad_norm()`; finite/overflow checking is unchanged. | `true` |

<i>**overlap_comm**</i>: [boolean]

| Description | Default |
Expand Down
169 changes: 169 additions & 0 deletions tests/unit/runtime/zero/test_zero1_optimizer_fastpath.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,169 @@
# Copyright (c) Microsoft Corporation.
# SPDX-License-Identifier: Apache-2.0
# DeepSpeed Team

import copy
from types import SimpleNamespace

import pytest
import torch

import deepspeed
from deepspeed.checkpoint.constants import CLIP_GRAD, DS_VERSION
from deepspeed.runtime.zero.stage_1_and_2 import DeepSpeedZeroOptimizer
from unit.common import DistributedTest
from unit.simple_model import SimpleModel


def _config(*, compute_grad_norm, gradient_clipping=0.0, stage=1, offload_optimizer=None, zenflow=None):
return {
"train_micro_batch_size_per_gpu": 1,
"bf16": {
"enabled": True,
"check_grad_overflow": True,
},
"optimizer": {
"type": "AdamW",
"params": {
"lr": 1e-3,
},
},
"zero_optimization": {
"stage": stage,
"compute_grad_norm": compute_grad_norm,
"offload_optimizer": offload_optimizer,
"zenflow": zenflow,
},
"gradient_clipping": gradient_clipping,
}


def test_identity_unscale_is_skipped():
optimizer = object.__new__(DeepSpeedZeroOptimizer)
optimizer.clip_grad = 0.0
optimizer.custom_loss_scaler = False
optimizer.loss_scaler = SimpleNamespace(cur_scale=1.0)
gradient = torch.ones(4)
with torch.profiler.profile(activities=[torch.profiler.ProfilerActivity.CPU]) as profiler:
optimizer.unscale_and_clip_grads([gradient], total_norm=None)

assert "aten::mul_" not in {event.key for event in profiler.key_averages()}


def test_non_identity_scale_still_unscales():
optimizer = object.__new__(DeepSpeedZeroOptimizer)
optimizer.clip_grad = 0.0
optimizer.custom_loss_scaler = False
optimizer.loss_scaler = SimpleNamespace(cur_scale=2.0)
gradient = torch.ones(4)

optimizer.unscale_and_clip_grads([gradient], total_norm=None)

torch.testing.assert_close(gradient, torch.full_like(gradient, 0.5))


def test_checkpoint_clipping_rejects_disabled_norm():
optimizer = object.__new__(DeepSpeedZeroOptimizer)
optimizer.compute_grad_norm = False
optimizer.loss_scaler = SimpleNamespace()
optimizer.dynamic_loss_scale = False
optimizer.overflow = False
optimizer.clip_grad = 0.0

with pytest.raises(ValueError, match="checkpoint with gradient clipping"):
optimizer._load_global_state({
CLIP_GRAD: 1.0,
DS_VERSION: "0.18.0",
})


class TestZero1OptimizerFastPath(DistributedTest):
world_size = 1

@pytest.mark.parametrize("stage", [1, 2])
def test_fast_path_matches_default_update(self, stage):
torch.manual_seed(123)
baseline_model = SimpleModel(hidden_dim=4)
fast_model = copy.deepcopy(baseline_model)
baseline_engine, baseline_optimizer, _, _ = deepspeed.initialize(
model=baseline_model,
model_parameters=baseline_model.parameters(),
config=_config(compute_grad_norm=True, stage=stage))
fast_engine, fast_optimizer, _, _ = deepspeed.initialize(
model=fast_model,
model_parameters=fast_model.parameters(),
config=_config(compute_grad_norm=False, stage=stage))
inputs = torch.randn(1, 4, device=baseline_engine.device, dtype=torch.bfloat16)
targets = torch.randn(1, 4, device=baseline_engine.device, dtype=torch.bfloat16)

baseline_loss = baseline_engine(inputs, targets)
fast_loss = fast_engine(inputs, targets)
baseline_engine.backward(baseline_loss)
fast_engine.backward(fast_loss)
baseline_engine.step()
fast_engine.step()

torch.testing.assert_close(fast_loss, baseline_loss)
assert baseline_optimizer._global_grad_norm is not None
assert fast_optimizer._global_grad_norm is None
for baseline_parameter, fast_parameter in zip(baseline_engine.module.parameters(),
fast_engine.module.parameters()):
torch.testing.assert_close(fast_parameter, baseline_parameter)

@pytest.mark.parametrize("stage", [1, 2])
def test_finite_step_skips_norm_but_updates_parameters(self, stage):
model = SimpleModel(hidden_dim=4)
engine, optimizer, _, _ = deepspeed.initialize(model=model,
model_parameters=model.parameters(),
config=_config(compute_grad_norm=False, stage=stage))
inputs = torch.randn(1, 4, device=engine.device, dtype=torch.bfloat16)
targets = torch.randn(1, 4, device=engine.device, dtype=torch.bfloat16)
before = [parameter.detach().clone() for parameter in engine.module.parameters()]

engine.backward(engine(inputs, targets))
engine.step()

assert optimizer.check_grad_overflow
assert optimizer._global_grad_norm is None
assert engine.get_global_grad_norm() is None
assert any(not torch.equal(previous, current) for previous, current in zip(before, engine.module.parameters()))

def test_overflow_check_still_skips_the_step(self):
model = SimpleModel(hidden_dim=4)
engine, optimizer, _, _ = deepspeed.initialize(model=model,
model_parameters=model.parameters(),
config=_config(compute_grad_norm=False))
inputs = torch.randn(1, 4, device=engine.device, dtype=torch.bfloat16)
targets = torch.randn(1, 4, device=engine.device, dtype=torch.bfloat16)
engine.backward(engine(inputs, targets))
gradient = next(gradient for gradients in optimizer.averaged_gradients.values() for gradient in gradients
if gradient is not None)
gradient.view(-1)[0] = float("nan")
before = [parameter.detach().clone() for parameter in engine.module.parameters()]

engine.step()

assert optimizer.overflow
assert optimizer._global_grad_norm is None
assert all(torch.equal(previous, current) for previous, current in zip(before, engine.module.parameters()))

def test_gradient_clipping_rejects_disabled_norm(self):
model = SimpleModel(hidden_dim=4)
with pytest.raises(ValueError, match="requires gradient_clipping=0"):
deepspeed.initialize(model=model,
model_parameters=model.parameters(),
config=_config(compute_grad_norm=False, gradient_clipping=1.0))

def test_optimizer_offload_rejects_disabled_norm(self):
model = SimpleModel(hidden_dim=4)
with pytest.raises(ValueError, match="does not support optimizer offload"):
deepspeed.initialize(model=model,
model_parameters=model.parameters(),
config=_config(compute_grad_norm=False, offload_optimizer={"device": "cpu"}))

def test_zenflow_rejects_disabled_norm(self):
model = SimpleModel(hidden_dim=4)
with pytest.raises(ValueError, match="does not support ZenFlow"):
deepspeed.initialize(model=model,
model_parameters=model.parameters(),
config=_config(compute_grad_norm=False, zenflow={}))
9 changes: 9 additions & 0 deletions tests/unit/runtime/zero/test_zero_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,15 @@ def test_zero_config_overlapcomm():
assert config.overlap_comm == True


def test_zero_config_compute_grad_norm():
assert DeepSpeedZeroConfig(stage=1).compute_grad_norm is True
assert DeepSpeedZeroConfig(stage=1, compute_grad_norm=False).compute_grad_norm is False

for stage in (0, 3):
with pytest.raises(ValueError, match="only with ZeRO Stage 1 or 2"):
DeepSpeedZeroConfig(stage=stage, compute_grad_norm=False)


def test_zero_config_offload_configs():
config = DeepSpeedZeroConfig()
assert config.offload_param is None
Expand Down