Skip to content

Fix disabled ZeRO-3 initialization contexts - #4319

Open
1fanwang wants to merge 2 commits into
huggingface:mainfrom
1fanwang:1fannnw/fix-zero3-init-context
Open

1fanwang wants to merge 2 commits into
huggingface:mainfrom
1fanwang:1fannnw/fix-zero3-init-context

Conversation

@1fanwang

@1fanwang 1fanwang commented Sep 23, 2026 •

Copy link
Copy Markdown

What does this PR do?

Disabling ZeRO-3 initialization leaves Transformers' partitioning check enabled. The plugin reports False, but Transformers reports True, so model loading can still enter the ZeRO-3 path. An exception inside the context also leaves the plugin flag changed.

The initialization-only configuration now respects the flag without changing the engine's training configuration. Context exit restores the previous state, including after an exception.

Related to #3170. A two-H100 regression reproduces the reported embedding error and passes with this fix; the original eight-GPU diffusion workflow has not been rerun.

Testing Done

On two H100s connected by NVLink, the same regression was run with the baseline package and the fixed package. The baseline run retained the new test files. The stack was CUDA 12.9, PyTorch 2.11.0+cu129, Transformers 5.15.0 and DeepSpeed 0.18.8; both runs used the same synthetic BERT configuration without downloading a model.

python -m pytest -q -s tests/deepspeed/test_deepspeed_multiple_model.py::DeepSpeedInitializationContextTests::test_disabled_model_initialization_on_two_gpus

Before, both ranks failed with exit 1:

{"rank": 1, "phase": "inside", "plugin": false, "transformers": true}
{"rank": 0, "phase": "inside", "plugin": false, "transformers": true}
{"rank": 1, "embedding_shape": [0], "partitioned": true}
{"rank": 0, "embedding_shape": [0], "partitioned": true}
[rank1]: RuntimeError: 'weight' must be 2-D
[rank0]: RuntimeError: 'weight' must be 2-D

After, both ranks passed with exit 0. The outside-context control confirms that normal ZeRO-3 partitioning still works:

{"rank": 1, "phase": "inside", "plugin": false, "transformers": false}
{"rank": 1, "embedding_shape": [16, 8], "partitioned": false}
{"rank": 0, "phase": "inside", "plugin": false, "transformers": false}
{"rank": 0, "embedding_shape": [16, 8], "partitioned": false}
{"rank": 1, "forward_shape": [1, 3, 8], "forward_matches_cpu": true, "restored_zero3": true, "outside_partitioned": true}
{"rank": 0, "forward_shape": [1, 3, 8], "forward_matches_cpu": true, "restored_zero3": true, "outside_partitioned": true}
Initialization and exception-restoration probe

I ran the same probe on the baseline and this branch with Python 3.12.13, PyTorch 2.14.0, Transformers 5.17.0 and DeepSpeed 0.19.7 on macOS arm64.

from accelerate.utils import DeepSpeedPlugin
from transformers.integrations.deepspeed import is_deepspeed_zero3_enabled


def observe(label: str) -> tuple[bool, bool]:
    state = plugin.is_zero3_init_enabled(), is_deepspeed_zero3_enabled()
    print(f"{label}: plugin={state[0]} transformers={state[1]}")
    return state


plugin = DeepSpeedPlugin(
    hf_ds_config={
        "zero_optimization": {"stage": 3},
        "train_micro_batch_size_per_gpu": 1,
        "gradient_accumulation_steps": 1,
    },
    zero3_init_flag=True,
)
plugin.set_deepspeed_weakref()
assert observe("before") == (True, True)
with plugin.zero3_init_context_manager(enable=False):
    inside = observe("inside")
assert observe("after") == (True, True)
try:
    with plugin.zero3_init_context_manager(enable=False):
        raise RuntimeError("model initialization failed")
except RuntimeError as error:
    print(error)
after_error = observe("after_error")
assert inside == (False, False)
assert after_error == (True, True)

Save the probe as probe_zero3_context.py and run it with the selected checkout installed:

python -u probe_zero3_context.py

Before, exit 1:

before: plugin=True transformers=True
inside: plugin=False transformers=True
after: plugin=True transformers=True
model initialization failed
after_error: plugin=False transformers=True
AssertionError

After, exit 0:

before: plugin=True transformers=True
inside: plugin=False transformers=False
after: plugin=True transformers=True
model initialization failed
after_error: plugin=True transformers=True

This also verifies restoration after a caller raises inside the context.

Before submitting

  • This PR fixes a typo or improves the docs (you can dismiss the other checks if that's the case).
  • Did you read the contributor guideline,
    Pull Request section?
  • Was this discussed/approved via a Github issue or the forum? Please add a link
    to it if that's the case.
  • Did you make sure to update the documentation with your changes? Here are the
    documentation guidelines, and
    here are tips on formatting docstrings.
  • Did you write any new necessary tests?

Who can review?

Anyone in the community is free to review the PR once the tests have passed. Feel free to tag
members/contributors who may be interested in your PR.

Keep Transformers initialization consistent with the plugin flag without changing the training configuration, and restore the state when model loading fails.

Signed-off-by: 1fanwang <1fannnw@gmail.com>
Exercise dense model loading and inference inside the disabled context, then confirm ordinary ZeRO-3 partitioning still resumes afterward.

Signed-off-by: 1fanwang <1fannnw@gmail.com>
@1fanwang
1fanwang marked this pull request as ready for review September 23, 2026 16:05
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.

1 participant