Skip to content

Enable ZeRO-3 linear wrapper for existing models - #8189

Open
tohtana wants to merge 2 commits into
deepspeedai:masterfrom
tohtana:tohtana/fix/zero3-linear-direct-init
Open

Enable ZeRO-3 linear wrapper for existing models#8189
tohtana wants to merge 2 commits into
deepspeedai:masterfrom
tohtana:tohtana/fix/zero3-linear-direct-init

Conversation

@tohtana

@tohtana tohtana commented Jul 28, 2026

Copy link
Copy Markdown
Collaborator

Passing an already-constructed model to deepspeed.initialize() with ZeRO-3 and memory_efficient_linear=true does not install the ZeRO-3 Linear wrapper. The wrapper is currently installed only when the model is constructed inside a deepspeed.zero.Init() context.

Without the wrapper, the standard Linear implementation can retain the gathered weight storage until backward completes, significantly increasing memory usage.

This PR activates the existing ZeRO-3 Linear wrapper for the deepspeed.initialize(model=...) path when memory_efficient_linear=true, without requiring a deepspeed.zero.Init() context.

Signed-off-by: Masahiro Tanaka <mtanaka@anyscale.com>
@tohtana
tohtana requested review from loadams and tjruwase as code owners July 28, 2026 20:01

@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: 455574c290

ℹ️ About Codex in GitHub

Your team has set up Codex to 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 👍.

Codex can also answer questions or update the PR. Try commenting "@codex address that feedback".

Comment thread tests/unit/runtime/zero/test_zero_linear_direct_init.py
Comment thread deepspeed/runtime/zero/partition_parameters.py
Comment thread deepspeed/runtime/zero/partition_parameters.py
Signed-off-by: Masahiro Tanaka <mtanaka@anyscale.com>
else:
grad_bias = grad_output.sum(0)
return grad_input, grad_weight, grad_bias
weight_was_partitioned = (hasattr(weight, "ds_status")

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

It occurs to me that this better be a function of weight as well i.e. weight.is_partitioned(), maybe worth a seperate PR.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

How about using is_zero_param() for detection?

@delock

delock commented Jul 30, 2026

Copy link
Copy Markdown
Collaborator

Hi @tohtana I have left my comments. How many memory we may save from this PR? Is a documentation change needed?

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