Skip to content
Merged
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
18 changes: 14 additions & 4 deletions tests/unit/v1/zero/test_offload_states.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,12 @@
from deepspeed.utils import safe_get_local_fp32_param, safe_get_local_optimizer_state
from deepspeed.runtime.zero.offload_states import get_state_devices

# The strict allocated-memory deltas asserted in this file assume memory_allocated()
# is allocator bookkeeping (cuda); on cpu it reports process RSS, which does not
# shrink when tensors are freed.
accelerator_device_mod = torch.get_device_module(get_accelerator().device_name())

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Add the required Signed-off-by trailer

This is a one-parent, non-merge commit, but its commit message has no Signed-off-by trailer. Add the required signoff before merging so the commit satisfies the repository's mandatory commit policy.

AGENTS.md reference: AGENTS.md:L8-L8

Useful? React with 👍 / 👎.

allocator_backed_memory_stats = hasattr(accelerator_device_mod, 'memory_allocated')
Comment on lines +21 to +22

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Avoid requiring torch.get_device_module during collection

On supported PyTorch versions that predate torch.get_device_module (including PyTorch 2.0, which remains allowed by requirements/requirements.txt:10), importing this test module now raises AttributeError during pytest collection, so none of its tests can run. Probe the method with getattr or determine the capability through the DeepSpeed accelerator abstraction instead.

Useful? React with 👍 / 👎.


# ==============================================================================
# ZeRO-1 and ZeRO-2 TESTS
# ==============================================================================
Expand Down Expand Up @@ -80,7 +86,8 @@ def run_model_zero12(model, param_groups, config_dict, hidden_dim, dtype, offloa
optimizer_device = offload_torch_device if is_offload_optimizer_enabled(config_dict) else accelerator_device
offload_only_optimizer_states = is_only_offload_optimizer_states(
offloaded_states, [OffloadStateTypeEnum.optim_states, OffloadStateTypeEnum.hp_params])
expect_memory_change = not (is_offload_optimizer_enabled(config_dict) and offload_only_optimizer_states)
expect_memory_change = allocator_backed_memory_stats and not (is_offload_optimizer_enabled(config_dict)
and offload_only_optimizer_states)

model, _, _, _ = deepspeed.initialize(model=model, model_parameters=param_groups, config=config_dict)

Expand Down Expand Up @@ -116,14 +123,16 @@ def run_model_zero12(model, param_groups, config_dict, hidden_dim, dtype, offloa
alloc_after_offload = get_accelerator().memory_allocated()

if grad_numel > 0:
assert alloc_after_offload < alloc_before_offload, f"FAIL: Allocated memory for grads should decrease after offload {alloc_after_offload=} < {alloc_before_offload=}"
if allocator_backed_memory_stats:
assert alloc_after_offload < alloc_before_offload, f"FAIL: Allocated memory for grads should decrease after offload {alloc_after_offload=} < {alloc_before_offload=}"
validate_grad_device(model, offload_torch_device)

model.reload_states()
alloc_after_reload = get_accelerator().memory_allocated()

if grad_numel > 0:
assert alloc_after_reload > alloc_after_offload, f"FAIL: Allocated memory for grads should increase after reload {alloc_after_reload=} > {alloc_after_offload=}"
if allocator_backed_memory_stats:
assert alloc_after_reload > alloc_after_offload, f"FAIL: Allocated memory for grads should increase after reload {alloc_after_reload=} > {alloc_after_offload=}"
validate_grad_device(model, accelerator_device)

reloaded_grads = [
Expand Down Expand Up @@ -279,7 +288,8 @@ def run_model_zero3(model, param_groups, config_dict, hidden_dim, dtype, offload
offload_only_optimizer_states = is_only_offload_optimizer_states(
offloaded_states,
[OffloadStateTypeEnum.optim_states, OffloadStateTypeEnum.hp_params, OffloadStateTypeEnum.lp_grads])
expect_memory_change = not (is_offload_optimizer_enabled(config_dict) and offload_only_optimizer_states)
expect_memory_change = allocator_backed_memory_stats and not (is_offload_optimizer_enabled(config_dict)
and offload_only_optimizer_states)
Comment on lines +291 to +292

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Preserve the ZeRO-3 offload placement check

Under the CPU accelerator—the environment this change is intended to unblock—this makes expect_memory_change false, but the existing validate_device(model, offload_state_device, offloaded_states) call is inside the same conditional at lines 336–338. Consequently, all ZeRO-3 CPU cases stop verifying that the requested states actually moved to CPU; the later equality checks after reload can pass even if offloading was a no-op. Gate only the memory-delta assertion and keep the observable device-placement validation unconditional.

AGENTS.md reference: AGENTS.md:L30-L32

Useful? React with 👍 / 👎.


offload_state_device: dict[OffloadStateTypeEnum, torch.device] = {
OffloadStateTypeEnum.hp_params: offload_torch_device,
Expand Down
Loading