Skip to content
Open
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
10 changes: 5 additions & 5 deletions deepspeed/runtime/bf16_optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,9 +12,10 @@
from deepspeed.runtime.base_optimizer import ZeROOptimizer
from packaging import version as pkg_version
from deepspeed.git_version_info import version
from deepspeed.runtime.utils import (get_global_norm_of_tensors, clip_tensors_by_global_norm, DummyOptim,
align_dense_tensors, all_gather_dp_groups, is_model_parallel_parameter,
see_memory_usage, graph_process, get_norm_with_moe_layers)
from deepspeed.runtime.utils import (bind_flat_views, get_global_norm_of_tensors, clip_tensors_by_global_norm,
DummyOptim, align_dense_tensors, all_gather_dp_groups,
is_model_parallel_parameter, see_memory_usage, graph_process,
get_norm_with_moe_layers)
from deepspeed.utils import link_hp_params, lazy_init_hp_params_optimizer_state, fragment_address, groups
from deepspeed.moe.utils import is_moe_param, is_moe_param_group
from deepspeed.utils.bwc import bwc_tensor_model_parallel_rank
Expand Down Expand Up @@ -293,8 +294,7 @@ def _split_flat_tensor(self, flat_tensor, num_elem_list):

def _update_storage_to_flattened_tensor(self, tensor_list, flat_tensor):
updated_params = self.unflatten(flat_tensor, tensor_list)
for p, q in zip(tensor_list, updated_params):
p.data = q.data
bind_flat_views(tensor_list, updated_params)

def _flatten_dense_tensors_aligned(self, tensor_list, alignment):
return self.flatten(align_dense_tensors(tensor_list, alignment))
Expand Down
14 changes: 9 additions & 5 deletions deepspeed/runtime/fp16/fused_optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,8 @@
import torch
from torch._utils import _flatten_dense_tensors, _unflatten_dense_tensors
from deepspeed.runtime.base_optimizer import DeepSpeedOptimizer
from deepspeed.runtime.utils import get_global_norm, get_flattened_grad_norm, CheckOverflow, get_weight_norm, get_norm_with_moe_layers, is_model_parallel_parameter
from deepspeed.runtime.utils import (bind_flat_views, get_global_norm, get_flattened_grad_norm, CheckOverflow,
get_weight_norm, get_norm_with_moe_layers, is_model_parallel_parameter)
from deepspeed.runtime.fp16.loss_scaler import LossScaleConfig, LossScaleProfile
from deepspeed.utils import logger, log_dist
from deepspeed.utils.torch import required_torch_version
Expand Down Expand Up @@ -92,8 +93,7 @@ def __init__(self,
self.fp16_groups_flat.append(_flatten_dense_tensors([p.clone().detach() for p in self.fp16_groups[i]]))
# set model fp16 weight to slices of flattened buffer
updated_params = _unflatten_dense_tensors(self.fp16_groups_flat[i], self.fp16_groups[i])
for p, q in zip(self.fp16_groups[i], updated_params):
p.data = q.data
bind_flat_views(self.fp16_groups[i], updated_params)
# init master weight, flattened
self.fp32_groups_flat.append(self.fp16_groups_flat[i].clone().float().detach())
# modify optimizer of have flat master weight
Expand Down Expand Up @@ -187,8 +187,7 @@ def step_fused_adam(self, closure=None):
# TODO: we probably don't need this? just to be safe
for i in range(len(norm_groups)):
updated_params = _unflatten_dense_tensors(self.fp16_groups_flat[i], self.fp16_groups[i])
for p, q in zip(self.fp16_groups[i], updated_params):
p.data = q.data
bind_flat_views(self.fp16_groups[i], updated_params)
return self.overflow

def set_lr(self, lr):
Expand Down Expand Up @@ -354,6 +353,11 @@ def step(self, closure=None):
for i in range(len(self.fp16_groups)):
updated_params = _unflatten_dense_tensors(self.fp32_groups_flat[i], self.fp16_groups[i])
for p, q in zip(self.fp16_groups[i], updated_params):
if p.numel() == 0:
# See bind_flat_views: `q` is a 1-D zeros({0}) here, not a view of
# `p`'s shape, so this copy would raise on the shape mismatch. There
# are no elements to copy either way.
continue
p.data.copy_(q.data)
self.has_executed_step = True
if self.timers:
Expand Down
15 changes: 15 additions & 0 deletions deepspeed/runtime/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -823,6 +823,21 @@ def empty_cache():
get_accelerator().reset_peak_memory_stats()


def bind_flat_views(tensors, views):
"""Point each tensor at its view of a flat buffer, skipping zero-element ones.

torch's ``unflatten_dense_tensors`` special-cases ``numel == 0`` and returns a
freshly allocated 1-D ``zeros({0})`` rather than a view of the requested shape,
so assigning it would replace e.g. a ``(0, 8)`` parameter with a ``(0,)`` one and
break the owning module's own forward. There is no slice of the flat buffer for
such a tensor to point at either, so nothing is left unbound by skipping it.
"""
for tensor, view in zip(tensors, views):
if tensor.numel() == 0:
continue
tensor.data = view.data


def see_memory_usage(message, force=False):
if not force:
return
Expand Down
9 changes: 4 additions & 5 deletions deepspeed/runtime/zero/stage_1_and_2.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,9 +21,9 @@
from deepspeed.runtime.base_optimizer import ZeROOptimizer
from deepspeed.runtime.fp16.loss_scaler import CreateLossScaler
from deepspeed.runtime.torch_autocast import get_autocast_dtype, get_all_comm_dtypes, is_autocast_initialized, sort_dtypes
from deepspeed.runtime.utils import (empty_cache, see_memory_usage, has_inf_or_nan, inf, is_model_parallel_parameter,
align_dense_tensors, all_gather_dp_groups, mask_nan_or_inf_with_val_inplace,
count_used_parameters_in_backward)
from deepspeed.runtime.utils import (bind_flat_views, empty_cache, see_memory_usage, has_inf_or_nan, inf,
is_model_parallel_parameter, align_dense_tensors, all_gather_dp_groups,
mask_nan_or_inf_with_val_inplace, count_used_parameters_in_backward)
from deepspeed.runtime.zero.config import ZeroStageEnum
from deepspeed.runtime.zero.utils import get_norm_dtype
from deepspeed.runtime.zero.offload_config import OffloadDeviceEnum, OffloadStateTypeEnum
Expand Down Expand Up @@ -804,8 +804,7 @@ def _configure_moe_settings(self):

def _update_model_bit16_weights(self, group_index):
updated_params = self.unflatten(self.bit16_groups_flat[group_index], self.round_robin_bit16_meta[group_index])
for p, q in zip(self.round_robin_bit16_groups[group_index], updated_params):
p.data = q.data
bind_flat_views(self.round_robin_bit16_groups[group_index], updated_params)

# set model fp16 weight to slices of reordered flattened buffer
for param_index, param in enumerate(self.bit16_groups[group_index]):
Expand Down
143 changes: 143 additions & 0 deletions tests/unit/runtime/zero/test_zero_numel_param_shape.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,143 @@
# Copyright (c) Microsoft Corporation.
# SPDX-License-Identifier: Apache-2.0

# DeepSpeed Team
"""A zero-element parameter must keep its shape when it is bound to a flat buffer.

Every optimizer wrapper that flattens the parameters repoints each one at its slice of
the flat buffer. torch's `unflatten_dense_tensors` special-cases a zero-element tensor
and returns a freshly allocated 1-D `zeros({0})` rather than a view of the requested
shape, so a `(0, 8)` parameter came back as `(0,)` and the module's own forward then
dispatched `F.linear` to `addmv`:

RuntimeError: size mismatch, got input (1), mat (1x8), vec (0)

The parameters are rebuilt on every `step()` as well as at init, so the shape did not
survive one iteration either.
"""

import pytest
import torch

from unit.common import DistributedTest

import deepspeed

HIDDEN = 8

# One case per wrapper that binds parameters to a flat buffer.
CONFIGS = {
"fp16_stage0": ({
"fp16": {
"enabled": True,
"loss_scale": 1.0
},
"zero_optimization": {
"stage": 0
}
}, torch.float16),
"bf16_stage0": ({
"bf16": {
"enabled": True
},
"zero_optimization": {
"stage": 0
}
}, torch.bfloat16),
"bf16_stage1_fp32_accum": ({
"bf16": {
"enabled": True
},
"zero_optimization": {
"stage": 1
},
"data_types": {
"grad_accum_dtype": "fp32"
}
}, torch.bfloat16),
"zero1": ({
"bf16": {
"enabled": True
},
"zero_optimization": {
"stage": 1
}
}, torch.bfloat16),
"zero2": ({
"bf16": {
"enabled": True
},
"zero_optimization": {
"stage": 2
}
}, torch.bfloat16),
}


class EmptyTailModel(torch.nn.Module):
"""A trainable parameter with no elements, kept in the autograd graph by the loss."""

def __init__(self, hidden=HIDDEN):
super().__init__()
self.dense = torch.nn.Linear(hidden, hidden, bias=False)
self.empty = torch.nn.Linear(hidden, 0, bias=False)

def forward(self, x):
hidden = self.dense(x)
# `empty(hidden)` is (batch, 0); summing it keeps the parameter in the graph.
return hidden.sum() + self.empty(hidden).sum()


def _engine(case):
extra, _ = CONFIGS[case]
config = {
"train_micro_batch_size_per_gpu": 1,
"optimizer": {
"type": "Adam",
"params": {
"lr": 1e-3
}
},
**extra,
}
model = EmptyTailModel()
engine, *_ = deepspeed.initialize(model=model, model_parameters=model.parameters(), config=config)
return engine


def _step(engine, case):
_, dtype = CONFIGS[case]
loss = engine(torch.randn(1, HIDDEN, device=engine.device, dtype=dtype))
engine.backward(loss)
engine.step()


@pytest.mark.parametrize("case", list(CONFIGS))
class TestZeroNumelParameterShape(DistributedTest):
world_size = 1

def test_shape_survives_initialize(self, case):
engine = _engine(case)

assert engine.module.empty.weight.shape == torch.Size([0, HIDDEN])
# The sized parameter shares the flat buffer, which is what makes the
# zero-element one the special case rather than the rule.
assert engine.module.dense.weight.shape == torch.Size([HIDDEN, HIDDEN])

def test_shape_survives_a_step(self, case):
engine = _engine(case)

_step(engine, case)

assert engine.global_steps == 1
assert engine.module.empty.weight.shape == torch.Size([0, HIDDEN])

def test_a_second_step_still_runs(self, case):
# step() rebuilds the parameters from the flat buffer, so a shape lost there
# only shows up on the forward of the iteration after it.
engine = _engine(case)

_step(engine, case)
_step(engine, case)

assert engine.global_steps == 2
Loading