Skip to content

Reduce memory consumption in batched_forward_pass - #234

Merged
lvwerra merged 2 commits into
huggingface:mainfrom
atsumoto:not-storing-logits
Mar 22, 2023
Merged

lvwerra merged 2 commits into
huggingface:mainfrom
atsumoto:not-storing-logits

Conversation

@atsumoto

Copy link
Copy Markdown
Contributor

This PR reduces memory consumption in batched_forward_pass of PPOTrainer, by avoiding the storage of logits when they are not necessary.

Before this PR, batched_forward_pass stored all of the model's logits all the time like other tensors such as values and logprobs. Here, logits tensors have a much larger size (batch_size * tokens * vocabulary_size) compared to logprobs and values tensors (batch_size × tokens), consuming a significant amount of cuda memory.

I have modified batched_forward_pass to avoid unnecessary storage of logits, which is only required when calculating entropy in the loss method.

@atsumoto atsumoto changed the title Reduce memory consumption by avoiding logits storage in forward_pass Reduce memory consumption in batched_forward_pass Mar 21, 2023
@HuggingFaceDocBuilderDev

HuggingFaceDocBuilderDev commented Mar 21, 2023 •

Copy link
Copy Markdown

The documentation is not available anymore as the PR was closed or merged.

@younesbelkada younesbelkada left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Thanks a lot for fixing and for taking care of the memory consumption
This looks very good to me!
Would love to hear @lvwerra 's thoughts here

Comment thread trl/trainer/ppo_trainer.py
@younesbelkada
younesbelkada requested a review from lvwerra March 21, 2023 12:27

@lvwerra lvwerra left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

Looks great, thanks!

@lvwerra
lvwerra merged commit a6ebdb6 into huggingface:main Mar 22, 2023
yxliu-TAMU pushed a commit to mincheolseong/ECEN743-GRPO-Project-Proposal that referenced this pull request Apr 20, 2025
* Reduce memory consumption by not storing logits in forward_pass

* Add docstring of return_logits
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.

4 participants