deepspeedai / deepspeedai/DeepSpeed
Support scan over sequence dimension when computing FFN
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 43.1k
- Forks
- 5k
- Avg merge
- 4d 15h
- Merged PRs (30d)
- 112
Description
Is your feature request related to a problem? Please describe.
It's difficult to train large context LLMs without out-of-memory errors, e.g., when training LLMs on repository level code.
Describe the solution you'd like
Blockwise scan a rematerialized FFN over sequence dimension, ie, computing FFN block by block. This can reduce memory cost by 4x over state-of-the-art Memeff / Flashattention as proposed in BPT.
The analysis of memory cost is shown in Section 3.1.
Additional context
- A Jax implementation of BPT is available and the most relevant part. The code allows us to train models on small HBM TPUv3, it would be great to if Megatron-LM supports the feature.
Contributor guide
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Read Section 3.1 of the BPT paper and the linked bpt.py lines 15-36 first. Then identify the corresponding FFN and training entry points in DeepSpeed; done means blockwise scanning over the sequence dimension is supported, lowers memory use, and preserves training behavior.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python, pytorch
- Domain
- machine-learning, performance
- Issue type
- Feature
- Difficulty
- 5/5
- Estimated time
- Over a week
- Activity status
- Stale
- Clarity
- Mostly clear
- Newbie friendliness
- 20/100