deepspeedai / deepspeedai/DeepSpeed

[QUESTION/HELP] ZERO3 get weight participate in loss

Open
#7,464 1 comment 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
Python
Stars
43.1k
Forks
5k
Avg merge
4d 15h
Merged PRs (30d)
112

Description

Hi, i am trying to engage part of the weights in my loss calculation, say something like l2norm, and i am currently using zero-stage3
However, after initialization, model_engine.parameters seem to just be placeholders until loss.backward
I tried using l2_norm(safe_get_full_fp32_param(params)) but this function seems to only provide copied buffer, thus no gradient derives from this operation.
safe_set_full_grad would work but that would be involving tons of manual calculation and rule-setting, also breaking the code guideline.
What should i do? Would using GatheredParams context work?If so, should i just do it on rank0 or all ranks..?
I am still struggling here so it would really be a lot of help if you could provide some suggestions.
Thanks!

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Begin with the ZeRO stage 3 parameter-handling implementation and the safe_get_full_fp32_param and GatheredParams entry points named in the report. Reproduce the l2-norm use case across ranks, then establish whether the completed approach preserves gradient flow without manual safe_set_full_grad updates.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
distributed-systems, machine-learning
Issue type
Feature
Difficulty
5/5
Estimated time
Over a week
Activity status
Stale
Clarity
Needs clarification
Newbie friendliness
18/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.