A Google TPU-kon reprodukálta az Olmo 3 7B tanítását

A Google mérnökei a MaxText JAX/XLA-alapú keretrendszerrel, Google Cloud TPU-kon reprodukálták az Allen Institute for AI Olmo 3 7B modelljének első két tanítási szakaszát. A vállalat szerint a kísérlet a veszteségfüggvény mellett különálló értékeléseken is igazolta a GPU-s PyTorch-referencia és a TPU-s megvalósítás egyezését.
- A Google MaxTextben, Google Cloud TPU-kon reprodukálta az Olmo 3 7B első két tanítási szakaszát.
- A MaxText futása az Ai2 referenciaeredményeit a veszteséggörbén túl különálló értékeléseken is követte.
- Egy adatbetöltési hiba látszólagos javulást okozott, amelyről kiderült, hogy memorizációból származott.
- A Google 44,5 százalékos MFU-t mért Ironwoodon, a második szakaszban v5p-n 57,4 százalékot.
- A hosszú kontextusú szakaszt és az utótréninget a közlés szerint még nem futtatták le.
A teljesen nyílt modell szolgált referenciaként
A Google Developers Blog szeptember 24-i beszámolója szerint az Olmo 3 7B-t azért választották, mert ritkán együtt elérhető tulajdonságokat egyesít. Az Ai2 által fejlesztett modell modern, 7 milliárd paraméteres nyelvi modell, amelyet éles léptékű tanítással készítettek. A szervezet közzétette a modell teljes fejlesztési folyamatát, többek között az adatokat, a kódot, a konfigurációkat, az ellenőrzőpontokat, a naplókat és az értékeléseket.
Ez független PyTorch- és GPU-referenciát biztosított a Google számára a MaxText és a TPU-k teszteléséhez. A cél annak vizsgálata volt, hogy egy GPU-kra készült tanítási recept hűen megismételhető-e JAX-alapú környezetben, TPU-hardveren, és az eredményesség a fontos mérőszámokon is igazolható-e.
Az első két tanítási szakaszt futtatták le
Az Olmo 3 receptje háromlépcsős. Az első szakasz az általános előtanítás, a második a köztes tanítás, vagyis az annealing, a harmadik pedig a hosszú kontextushoz való alkalmazkodás. A Google most az első, körülbelül 5,9 billió tokenes előtanítást és a második szakaszt futtatta végig, majd az eredményeket az Ai2 referenciafutásaihoz hasonlította.
A reprodukció az Ai2 0. lépésnél kiadott PyTorch-ellenőrzőpontjából indult. A súlyokat MaxTexthez Orbax-formátumba alakították át. A konvertált modell első lépésének kimenete a HuggingFace-referenciához képest körülbelül 1,5×10-3 KL-eltérést mutatott, ami a Google szerint a különböző keretrendszerek közötti zajszintnek felel meg. A 8192 tokenes kontextus és bfloat16 használata mellett a két modell az esetek 98,75 százalékában ugyanazt az első helyezett tokent adta.
Az Olmo 3 7B 32 rétegű, 4096 dimenziós, sűrű transzformer. Architektúrájában átrendezett normálási blokk, QK-norm és 3:1 arányú, ablakos, illetve globális figyelem szerepel. Ezeket a MaxTextben is megvalósították, a tanítási adatfolyam pedig az Olmo-core működését követte: a dokumentumokat tokenizálták és összefűzték, 8192 tokenes, nem átfedő részekre vágták, majd rögzített kezdőértékkel összekeverték.
A mérés hibát is feltárt
A MaxText első szakaszának veszteséggörbéje körülbelül 800 ezer lépésig legfeljebb ±0,012 eltéréssel követte az Ai2 közzétett görbéjét. Körülbelül 900 ezer lépéstől a MaxText eredménye kedvezőbbnek látszott, ám a Google későbbi vizsgálata szerint ezt egy adatbetöltési hiba okozta. A javulás memorizációból származott, ezért a különálló, tanítás közben nem használt értékelések fontos szerepet kaptak az ellenőrzésben.
A mérnökök a futtatás megbízhatóságát is vizsgálták. Egy újraindítási tesztben az ellenőrzőpontból folytatott futás minden lépésnél 0,000 eltérést mutatott. Amikor egy hardverhiba megszakította a második szakaszt, a folytatás 127 lépést tanított újra, miközben a naplózott veszteség és a perplexity eltérése szintén 0,000 maradt.
A TPU-kon a teljesítményt is javították
A futás közben a Google a kapacitás háromnegyedét elveszítette, ezért a tanítást az eredeti méret negyedét jelentő szeletén folytatta, a recept módosítása nélkül. A közlés szerint az eszközönkénti áteresztőképesség 1 százalékon belül megmaradt. A második szakaszban a futtatót ugyanazzal az indítóval v5p TPU-kra irányították az Ironwood helyett, és 57,4 százalékos MFU-t értek el.
Ironwoodon 7 milliárd paraméteres modellnél 44,5 százalékos MFU-t mértek. Ehhez a SparseCore kollektív műveletkiszervezését, az újraszámolás hangolását és az optimális felosztást használták. Egy külön vizsgálatban a figyelmi fejek számát 32-ről 16-ra csökkentették, miközben a fejdimenziót 128-ról 256-ra növelték. Azonos paraméterszám és lebegőpontos műveletszám mellett ez 12,4 százalékkal gyorsabbnak bizonyult, a veszteséggörbe pedig 120 milliárd tokenig egyezett az eredeti változatéval. A tényleges reprodukcióban az eredeti architektúrát tartották meg.
A harmadik, 65 ezres szekvenciát és YaRN-t használó szakaszt, valamint az SFT- és GRPO-alapú utótréninget a Google leírta, de a beszámoló szerint ezeket még nem futtatták le. A mostani eredmény a MaxText TPU-s használatát az első két szakaszra támasztja alá, miközben az is kiderült, hogy a különálló értékelések a látszólagos tanítási előnyök ellenőrzéséhez nélkülözhetetlenek.
Google for Developers: Reproducing Olmo 3 7B Pre-training in MaxText: case study of large scale training on TPUs- Google Developers Blog


