5 ms·
Bugs in LLM Training – Gradient Accumulation Fix
- danielhanchen 2y agoOh hey! :) TLDR naively gradient accumulation was over-weighting short sequence lengths in LLM finetuning and training runs, and under-weighting long sequence lengths. For eg a text with sequence lengths of [1, 100] would be scaled by 1/(100+1) in full batch training, but grad accum of 2 would weight [1] as 1/1 * 1/2 = 1/2, whilst [100] as 1/100 * 1/2 = 1/200. (1/2 since grad accum needs to divide by the # of grad accum steps)
- ejddhbrbrrnrn 2y agoIs this a general issue rather than unsloth specific. How wide is this problem? Sounds wild if it has been affecting everyones training.
- danielhanchen 2y agoUnfortunately it's not an Unsloth issue but a general issue affecting nearly all trainers which use grad accum. We worked with Huggingface so their trainers should be fixed now though in the main branch
- imjonse 2y agoSame issue described on HF: https://huggingface.co/blog/gradient_accumulation https://huggingface.co/blog/gradient_accumulation It also highlights the main disadvantage of Transformers codebase using the copy-paste method for models, where this fix needs to be applied to every single model separately.
- CraigJPerry 2y ago>> disadvantage of Transformers codebase using the copy-paste method for models, where this fix needs to be applied to every single model separately What are the best tools we have available for tackling this kind of large scale copy-paste change? https://github.com/huggingface/transformers/pull/34191/commits https://github.com/huggingface/transformers/pull/34191/commi... This feels too complex to tackle with PyCharm structural find and replace, even a more powerful structural find and replace like https://comby.dev/ https://comby.dev/ feels underpowered here. Sourcegraph batch changes? That solves broadcasting the change but doesn’t help with capturing the change to make. Open rewrite? The python implementation is early stages, not prod ready as I understand it. Plus this change is too complex to use refaster templates even if we could use orw so you’d be debugging a fairly involved method visitor which in this case is probably orders of magnitude more time consuming than just making the changes manually. What else is there that I don’t know about?
- danielhanchen 2y agoYe a complete change was necessary for now - HF had to isolate the cross entropy loss and make another class for it, and it had to be applied to all model archs.
- danielhanchen 2y agoUnfortunately transformers is a general library for many models, and so there are tonnes of different architectures. Unfortunately copy paste and changing some parts of the arch is the only way feasible in the meantime.
- xcodevn 2y agoLook from a different point of view: this is a feature, not a bug. With this, every example has equal weight, while with the fix, every token has equal weight.
- danielhanchen 2y agoYes you're correct, but in normal full batch training without gradient accumulation, all tokens are weighted equally. Standard grad accum does not, and so the "fix" makes grad accum and full batch training finally mathematically equivalent
- oergiR 2y agoThat makes it sound like it’s a choice, which it isn’t really. The way to look at it is from a probabilistic perspective: with the fix, you maximise the probability of the data. Without the fix, you fairly arbitrarily raise some probabilities to a power greater than one, and some to a power less than one.
- danielhanchen 2y agoYes exactly- mathematically it was incorrect to begin with.
- pama 2y agoAlthough there may be uses for such a modified loss, based on the tone of the writeup it feels like this was an unintended bug in their training code. Training llms with variable max sequence length on different GPU is a recipe for inefficient training anyways, so careful optimizion of MFU at scale, or fixed max sequence length per batch, would have avoided this “bug”.
- danielhanchen 2y agoYe one way to fix it is to use fixed sequence lengths, but it'll still be a tad bit off. Packing say 1000 small sequences to fit a large sequence lengths still will incur the same issue since the denominator will be off by 1000, but yes the problem is much less pronounced.