Skip to content

[bug fix] ZeRO2: preserve fp32 grad accumulation precision in partition storage - #8749

Open
stas00 wants to merge 5 commits into
masterfrom
stas00/fp32-gas-reduction
Open

stas00 wants to merge 5 commits into
masterfrom
stas00/fp32-gas-reduction

Conversation

@stas00

@stas00 stas00 commented Oct 5, 2026 •

Copy link
Copy Markdown
Collaborator

When training bf16 models with data_types.grad_accum_dtype: fp32, ZeRO-2 could still store the reduced owned gradient partition in bf16 storage. That defeats the fp32 accumulation setting after communication: values that are representable in fp32 but lie between bf16 values are rounded before APIs such as safe_get_full_grad() read them back.

This matters for request accumulation and similar uneven-gradient cases. If two ranks contribute gradients 1.0 and 1.0078125, fp32 communication produces the exact average 1.00390625. Storing that result in a bf16 partition rounds it back to 1.0, so the training step loses the precision that grad_accum_dtype: fp32 requested.

This PR makes ZeRO-1/2 use fp32 owned gradient partition storage when fp32 gradient accumulation is active. The ZeRO-2 path now writes reduced fp32 communication results into a compact fp32 partition buffer and skips allocating the legacy bf16 grads_in_partition buffer for that path. This keeps the memory behavior aligned with the feature: the old bf16 owned partition is replaced rather than kept alongside the fp32 one, so bf16 + fp32 grad accumulation should use less GPU memory than the previous mixed-storage path.

ZeRO-3 was checked separately. It already allocates grad_partitions_flat_buffer with gradient_accumulation_dtype and casts reduced partitions to that dtype before storing them, so it does not have the same persistent bf16 partition-storage issue. A ZeRO-3 regression case is included to guard that behavior.

Memory saved

Run with a 1B-parameter bf16 toy model on 2 H200 ranks.

  Per rank:

   metric                       before a77aeb676    after 87e384ec3        change
  ━━━━━━━━━━━━━━━━━━━━━━━━━━━  ━━━━━━━━━━━━━━━━━━  ━━━━━━━━━━━━━━━━━  ━━━━━━━━━━━━
   peak allocated                     14,336 MiB         14,336 MiB     unchanged
  ───────────────────────────  ──────────────────  ─────────────────  ────────────
   peak reserved                      16,386 MiB         14,338 MiB    -2,048 MiB
  ───────────────────────────  ──────────────────  ─────────────────  ────────────
   reserved delta after init          12,290 MiB         10,242 MiB    -2,048 MiB 

Config was ZeRO-2, bf16 model, communication_data_type=fp32, data_types.grad_accum_dtype=fp32, GAS=4, numel=1,073,741,824, 2 ranks.

PyTorch peak allocated is unchanged, but peak reserved GPU memory drops by 2 GiB per rank.

Implementation

  • Adds compact fp32 owned gradient partition buffers for ZeRO-1/2 when the model dtype differs from fp32 gradient accumulation dtype.
  • Writes fp32 all-reduce results into that compact partition storage instead of relying on the bf16-owned partition copy.
  • Avoids allocating the legacy bf16 grads_in_partition buffer on the compact fp32 path.
  • Keeps the regression observable through safe_get_full_grad() rather than depending on private storage layout, with the ZeRO-2 internal dtype assertion kept only where it directly checks the fixed path.

@stas00 stas00 changed the title [bug fix] ZeRO1/2: preserve fp32 grad accumulation precision in partition storage [bug fix] ZeRO2: preserve fp32 grad accumulation precision in partition storage Oct 5, 2026
stas00 added 2 commits October 5, 2026 20:18
Signed-off-by: Stas Bekman <stas@stason.org>
Signed-off-by: Stas Bekman <stas@stason.org>
@stas00
stas00 marked this pull request as ready for review October 5, 2026 21:36

@chatgpt-codex-connector chatgpt-codex-connector Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

💡 Codex Review

Here are some automated review suggestions for this pull request.

Reviewed commit: 87e384ec3c

ℹ️ About Codex in GitHub

Codex has been enabled to automatically review pull requests in this repo. Reviews are triggered when you

  • Open a pull request for review
  • Mark a draft as ready
  • Comment "@codex review".

If Codex has suggestions, it will comment; otherwise it will react with 👍.

When you sign up for Codex through ChatGPT, Codex can also answer questions or update the PR, like "@codex address that feedback".

Comment thread deepspeed/runtime/zero/stage_1_and_2.py Outdated
Comment thread deepspeed/runtime/zero/stage_1_and_2.py
Comment thread deepspeed/runtime/zero/stage_1_and_2.py
Signed-off-by: Stas Bekman <stas@stason.org>
@sfc-gh-truwase
sfc-gh-truwase added this pull request to the merge queue Oct 5, 2026
@github-merge-queue
github-merge-queue Bot removed this pull request from the merge queue due to failed status checks Oct 6, 2026
@stas00
stas00 added this pull request to the merge queue Oct 6, 2026
@github-merge-queue
github-merge-queue Bot removed this pull request from the merge queue due to failed status checks Oct 6, 2026
@sfc-gh-truwase
sfc-gh-truwase added this pull request to the merge queue Oct 6, 2026
@stas00

stas00 commented Oct 6, 2026 •

Copy link
Copy Markdown
Collaborator Author

@sfc-gh-truwase, I think the modal ci is the problem, it skips in the PR and then runs in the merge queue and fails there, trying to figure out if I can find the actual failures.

Details

@stas00

stas00 commented Oct 6, 2026 •

Copy link
Copy Markdown
Collaborator Author

So modal CI workload is broken, it completes only 24% of tests and then exits:

DS_CI_FAILURE_CLASS=timeout: Sandbox lifetime exhausted after 5400s (run pytest failed with exit code 137)
Error: Process completed with exit code 124.

@github-merge-queue
github-merge-queue Bot removed this pull request from the merge queue due to failed status checks Oct 6, 2026
@stas00
stas00 added this pull request to the merge queue Oct 6, 2026
@github-merge-queue
github-merge-queue Bot removed this pull request from the merge queue due to failed status checks Oct 7, 2026
@stas00
stas00 added this pull request to the merge queue Oct 7, 2026
@github-merge-queue
github-merge-queue Bot removed this pull request from the merge queue due to failed status checks Oct 7, 2026
@sfc-gh-truwase
sfc-gh-truwase added this pull request to the merge queue Oct 9, 2026
@github-merge-queue
github-merge-queue Bot removed this pull request from the merge queue due to failed status checks Oct 9, 2026
@hwchen2017
hwchen2017 added this pull request to the merge queue Oct 9, 2026
@github-merge-queue
github-merge-queue Bot removed this pull request from the merge queue due to failed status checks Oct 9, 2026
@hwchen2017
hwchen2017 requested a balanced review from Copilot October 10, 2026 04:39

Copilot AI left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

🟡 Changes recommended

CPU-offload, non-contiguous, and dynamic state-offload paths do not correctly preserve or manage the new fp32 partition storage.

3 open findings
What changed in this PR

Preserves fp32 gradient accumulation precision in ZeRO-1/2 partition storage.

Changes:

  • Adds compact fp32 owned-gradient buffers.
  • Routes reduced gradients into fp32 storage.
  • Adds precision and compatibility tests.
File Description
deepspeed/​runtime/​zero/​stage_1_and_2.py Implements fp32 partition storage and reduction routing.
tests/​unit/​runtime/​zero/​test_zero_tensor_fragment.py Adds gradient precision regression coverage.
tests/​unit/​v1/​zero/​test_zero2_gradient_safety.py Tests legacy bucket compatibility.

🧠 Review effort: Balanced


💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.

Comment on lines +143 to +145
def __init__(self, grad_buffers, flat_partition):
super().__init__(grad_buffers)
self.flat_partition = flat_partition
Comment on lines +1247 to +1251
if (self.cpu_offload or self.gradient_accumulation_dtype != torch.float32
or self.gradient_accumulation_dtype == self.dtype):
return False
if not self.contiguous_gradients:
return False
class TestZeroBf16Fp32GradAccum(DistributedTest):
world_size = 2

@pytest.mark.parametrize('zero_stage', [2, 3])
@sfc-gh-truwase
sfc-gh-truwase added this pull request to the merge queue Oct 10, 2026
@github-merge-queue
github-merge-queue Bot removed this pull request from the merge queue due to failed status checks Oct 10, 2026

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants