LongStraw ha annunciato che la sua tecnica di branch-replay può elaborare 2,1 milioni di posizioni di token per il post-training tramite reinforcement learning (RL) utilizzando solo otto GPU H20, riducendo i costi dell'hardware di un ordine di grandezza. Questa affermazione è importante perché l'addestramento di modelli a lungo contesto ha tradizionalmente richiesto decine di GPU di fascia alta, rappresentando una barriera per la maggior parte dei laboratori di ricerca e delle startup.
Perché il RL a lungo contesto è costoso
Il fine-tuning basato su RL dei grandi modelli linguistici esegue tipicamente dei rollout che generano molteplici completamenti alternativi per lo stesso prompt. Ogni rollout deve essere sottoposto a back-propagation, quindi il costo computazionale scala con il numero totale di posizioni di token elaborate. Le attuali pipeline che mirano a un contesto di un milione di token richiedono spesso da 64 a 128 GPU per terminare in un tempo ragionevole. Il costo di tale hardware, unito all'elettricità e al raffreddamento richiesti, limita quanto gli esperti possano spingere la lunghezza del contesto.
Come il branch replay riduce il carico di lavoro
L'approccio di LongStraw si basa su due osservazioni riguardanti la generazione dei transformer:
- Il prompt e la parte iniziale di una risposta sono identici tra i vari rollout.
- Solo la parte finale divergente di ogni risposta necessita effettivamente di nuova computazione.
Il sistema costruisce uno stack di esecuzione consapevole dell'architettura che registra le attivazioni per il prefisso condiviso. Quando viene esplorata una nuova ramificazione, il sistema riproduce il prefisso memorizzato nella cache invece di ricalcolarlo, eseguendo poi il passaggio backward solo sul segmento nuovo. In pratica, ciò significa che il passaggio backward tocca molte meno posizioni di token, garantendo una riduzione della computazione pura da 8 a 16 volte.
Impatto immediato
- 2,1 milioni di posizioni di token elaborate su otto GPU H20, un budget hardware che normalmente supporterebbe solo una frazione di tale carico di lavoro.
- Attacco diretto al collo di bottiglia del RL a lungo contesto, dove i costi di memoria e computazione esplodono all'aumentare del contesto.
- I laboratori potrebbero spostare l'allocazione delle GPU: lo stesso hardware, principalmente un acceleratore di inferenza, può ora essere utilizzato per l'addestramento, sebbene i risultati possano variare su altre schede.
Domande aperte e limiti
L'annuncio omette i dati sulla velocità di addestramento e le curve di convergenza, quindi non sappiamo se il taglio computazionale si traduca in un tempo di esecuzione (wall-clock time) più rapido o semplicemente in una minore occupazione delle GPU. Il metodo è descritto per il campionamento autoregressivo; il suo comportamento con strategie non autoregressive o ibride rimane non testato. Poiché l'H20 è principalmente un acceleratore di inferenza, le prestazioni su schede di addestramento più comuni come H100 o B200 potrebbero variare.
Benchmark indipendenti non hanno ancora verificato i numeri di LongStraw. In assenza di una validazione da parte di terzi, la comunità dovrebbe considerare i risultati come promettenti ma provvisori.
Cosa c'è in gioco
Se l'idea del branch-replay si estendesse ad altri algoritmi di fine-tuning RL come il Direct Preference Optimization (DPO) o il Proximal Policy Optimization (PPO), la barriera dei costi per i modelli a lungo contesto potrebbe venire meno.
Cosa osservare
- Tentativi di replicazione da parte di terzi su una gamma di architetture GPU.
- Aggiornamenti da parte di LongStraw sul throughput di addestramento e sulla qualità finale del modello rispetto alle pipeline di riferimento.
