Skip to content
Merged
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
24 changes: 19 additions & 5 deletions deepspeed/ops/adam/zenflow_torch_adam.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,13 @@ def wrapper(fn):

class ZenFlowSelectiveAdamW(torch.optim.AdamW):

def __setstate__(self, state):
super().__setstate__(state)
for param_state in self.state.values():
step = param_state.get("step")
if torch.is_tensor(step) and step.dtype in (torch.float16, torch.bfloat16):
param_state["step"] = step.float()
Comment thread
vineethsaivs marked this conversation as resolved.

def __init__(self, *args, offload=False, bucket_size=5e8, **kwargs):
if not _ZENFLOW_AVAILABLE:
raise RuntimeError("ZenFlow features are not available with PyTorch < 2.0. "
Expand Down Expand Up @@ -115,7 +122,7 @@ def _step_without_offload(self):

state = self.state.setdefault(param, {})
if len(state) == 0:
state["step"] = torch.zeros((), dtype=param.dtype, device=selected_param.device)
state["step"] = torch.zeros((), dtype=torch.float32, device=selected_param.device)
state["exp_avg"] = torch.zeros_like(selected_param)
state["exp_avg_sq"] = torch.zeros_like(selected_param)
if amsgrad:
Expand Down Expand Up @@ -218,7 +225,7 @@ def group_step(self, group_to_paramlist):

state = self.state.setdefault(param, {})
if len(state) == 0:
state["step"] = torch.zeros((), dtype=param.dtype, device=selected_param.device)
state["step"] = torch.zeros((), dtype=torch.float32, device=selected_param.device)
if amsgrad:
state["max_exp_avg_sq"] = torch.zeros_like(selected_param)
if not self.offload:
Expand Down Expand Up @@ -269,6 +276,13 @@ def group_step(self, group_to_paramlist):

class ZenFlowSelectiveAdamW_stage3(torch.optim.AdamW):

def __setstate__(self, state):
super().__setstate__(state)
for param_state in self.state.values():
step = param_state.get("step")
if torch.is_tensor(step) and step.dtype in (torch.float16, torch.bfloat16):
param_state["step"] = step.float()

def __init__(self, *args, offload=False, bucket_size=5e8, **kwargs):
super(ZenFlowSelectiveAdamW_stage3, self).__init__(*args, **kwargs)
self.offload = offload
Expand Down Expand Up @@ -346,7 +360,7 @@ def _step_without_offload(self):

state = self.state.setdefault(param, {})
if len(state) == 0:
state["step"] = torch.zeros((), dtype=param.dtype, device=selected_param.device)
state["step"] = torch.zeros((), dtype=torch.float32, device=selected_param.device)
state["exp_avg"] = torch.zeros_like(selected_param)
state["exp_avg_sq"] = torch.zeros_like(selected_param)
if amsgrad:
Expand Down Expand Up @@ -463,7 +477,7 @@ def _selective_step_one_param(self, param, group, apply_temp=False):

state = self.state.setdefault(param, {})
if len(state) == 0:
state["step"] = torch.zeros((), dtype=param.dtype, device=compute_device)
state["step"] = torch.zeros((), dtype=torch.float32, device=compute_device)
if amsgrad:
state["max_exp_avg_sq"] = torch.zeros_like(selected_param)
if not self.offload:
Expand Down Expand Up @@ -568,7 +582,7 @@ def group_step(self, paramlist):

state = self.state.setdefault(param, {})
if len(state) == 0:
state["step"] = torch.zeros((), dtype=param.dtype, device=selected_param.device)
state["step"] = torch.zeros((), dtype=torch.float32, device=selected_param.device)
if amsgrad:
state["max_exp_avg_sq"] = torch.zeros_like(selected_param)
if not self.offload:
Expand Down
53 changes: 53 additions & 0 deletions tests/unit/runtime/zenflow/test_zf.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,11 @@
# DeepSpeed Team

import pytest
import torch
import deepspeed.comm as dist
from deepspeed.accelerator import get_accelerator
from deepspeed.runtime.zenflow.zenflow_stage_1_and_2 import _num_selected_columns
from deepspeed.ops.adam.zenflow_torch_adam import ZenFlowSelectiveAdamW, ZenFlowSelectiveAdamW_stage3

from unit.common import DistributedTest
from unit.simple_model import SimpleModel, random_dataloader
Expand All @@ -22,6 +24,57 @@ def test_num_selected_columns_has_nonzero_floor(num_columns, topk_ratio, expecte
assert _num_selected_columns(num_columns, topk_ratio) == expected


@pytest.mark.parametrize("dtype,steps", [(torch.bfloat16, 300), (torch.float16, 2100), (torch.float32, 300)])
def test_selective_adamw_step_counter_keeps_advancing(dtype, steps):
model = torch.nn.Linear(2, 1, bias=False, dtype=dtype)
optimizer = ZenFlowSelectiveAdamW(model.parameters(), lr=0.01)
inputs = torch.ones(1, 2, dtype=dtype)
for _ in range(steps):
optimizer.zero_grad()
model(inputs).float().square().sum().backward()
model.weight.selected_indices = torch.arange(2)
model.weight.selected_grad = model.weight.grad
optimizer.step()

# BF16 cannot represent 257 and FP16 cannot represent 2049. The bias-correction
# clock must keep advancing even when the model uses either of those dtypes.
assert optimizer.state_dict()["state"][0]["step"].item() == steps


@pytest.mark.parametrize("optimizer_cls", [ZenFlowSelectiveAdamW, ZenFlowSelectiveAdamW_stage3])
@pytest.mark.parametrize("counter_dtype", [torch.bfloat16, torch.float16, torch.float32, torch.float64])
def test_selective_adamw_loads_counter_without_losing_precision(optimizer_cls, counter_dtype):
param = torch.nn.Parameter(torch.ones(2))
original = torch.optim.AdamW([param])
param.grad = torch.ones_like(param)
original.step()
checkpoint = original.state_dict()
checkpoint["state"][0]["step"] = torch.tensor(256, dtype=counter_dtype)

optimizer = optimizer_cls([param])
optimizer.load_state_dict(checkpoint)
restored_step = optimizer.state_dict()["state"][0]["step"]
assert restored_step.item() == 256
expected_dtype = torch.float64 if counter_dtype == torch.float64 else torch.float32
assert restored_step.dtype == expected_dtype
assert checkpoint["state"][0]["step"].dtype == counter_dtype


@pytest.mark.parametrize("dtype,boundary", [(torch.bfloat16, 256), (torch.float16, 2048)])
def test_selective_adamw_resumed_counter_advances_past_old_limit(dtype, boundary):
param = torch.nn.Parameter(torch.ones(2, dtype=dtype))
optimizer = ZenFlowSelectiveAdamW([param])
param.selected_grad = torch.ones_like(param)
optimizer.step()
checkpoint = optimizer.state_dict()
checkpoint["state"][0]["step"] = torch.tensor(boundary, dtype=dtype)

optimizer.load_state_dict(checkpoint)
param.selected_grad = torch.ones_like(param)
optimizer.step()
assert optimizer.state_dict()["state"][0]["step"].item() == boundary + 1


class BaseZenFlowTest:
hidden_dim = 10
batch_size = 4
Expand Down
Loading