Meta raddoppia l'efficienza di addestramento di GEM al 20-25% MFU con kernel personalizzati e parallelismo 5D
Meta ha presentato il suo modello generativo di raccomandazione pubblicitaria (GEM), che ora opera a scala LLM utilizzando migliaia di GPU. Nell'ultimo anno, ha raggiunto un'utilizzazione dei FLOP del modello (MFU) del 20-25% e ha quadruplicato i suoi FLOP di addestramento. Questo successo deriva da sforzi collaborativi nell'ottimizzazione di kernel, precisione, parallelismo, rete e memoria. Il design di GEM integra trilioni di parametri sparsi insieme a miliardi di parametri densi, utilizzando contenuti pubblicitari e dati di coinvolgimento degli utenti per l'addestramento. È stata creata una libreria di kernel specializzata, con Jagged Flash Attention (JFA) e Generalized Dot-Product Attention (GDPA), insieme a un parallelismo 5D consapevole della topologia. Questi progressi hanno portato a un aumento del 40-140% dei TFLOPS da JFA v4, un raddoppio della velocità da GDPA e un aumento del 30,6% del MFU da BlockAttention.
Fatti principali
- Il modello GEM di Meta ora si addestra a scala LLM su diverse migliaia di GPU di ultima generazione.
- L'efficienza di addestramento end-to-end è raddoppiata al 20-25% MFU.
- I FLOP di addestramento sono aumentati di 4 volte in 12 mesi.
- La libreria di kernel personalizzati include JFA, GDPA e BlockAttention.
- Addestramento misto a precisione ultra-bassa con attenzione MXFP8 e MLP.
- Parallelismo 5D consapevole della topologia con collettive senza SM.
- JFA v4 raggiunge un miglioramento dei TFLOPS del 40-140% rispetto a JFA v2.
- Il kernel GDPA raggiunge un'accelerazione forward di 2 volte e fino a 3,5 volte rispetto a Flash Attention 4.
- BlockAttention migliora il MFU del layer di self-attention di +30,6% rispetto al block attention di Triton.
- Base Batch Shuffling (BBS) riduce lo squilibrio di carico con zero comunicazione cross-rank.
Entità
Istituzioni
- Meta
- PyTorch
- NVIDIA