System Tunix od Google eliminuje wąskie gardło, które uniemożliwiało efektywne wykorzystanie jednostek TPU w wielkoskalowym agentowym uczeniu ze wzmocnieniem (RL). Poprzez oddzielenie generowania danych z interakcji od aktualizacji polityki, Tunix zwiększa utylizację TPU z pojedynczych procent do niemal pełnej wydajności, drastycznie ograniczając marnotrawstwo mocy obliczeniowej.
Wąskie gardło w agentowym RL
Agentowe RL różni się od bardziej znanego trenowania modeli językowych metodą „next-token”. Agent musi wysyłać wywołania API, wykonywać kod lub przechodzić przez symulowane środowisko, a następnie reagować na wynik. Pętla treningowa jest zatem synchroniczna: model generuje akcję, środowisko działa, wynik zostaje zwrócony i dopiero wtedy model otrzymuje aktualizację gradientu. Gdy jeden krok w środowisku zajmuje kilka sekund, kosztowne jednostki TPU pozostają bezczynne, a raportowana utylizacja może spaść poniżej 10%. Ta nieefektywność przekłada się bezpośrednio na wyższe rachunki za chmurę i wolniejsze cykle badawcze.
Rozdzielona architektura Tunix
Tunix rozwiązuje ten problem, przenosząc dwa etapy — generowanie trajektorii i optymalizację polityki — na oddzielne zasoby sprzętowe.
- Asynchroniczni aktorzy działają na tanich procesorach CPU lub GPU. Każdy aktor nieustannie wchodzi w interakcję z przypisanym mu środowiskiem, rejestruje akcje i obserwacje, a następnie przesyła powstałe trajektorie do współdzielonego magazynu.
- Ciągłe jednostki uczące (Continuous learners) zajmują dedykowane jednostki TPU Pods. Jednostka ucząca pobiera partie danych z centralnego bufora i wykonuje aktualizacje gradientu bez czekania na zakończenie procesu rollout przez któregokolwiek z aktorów.
- Bufor o wysokiej przepustowości znajduje się pośrodku, pełniąc rolę obszaru przygotowawczego dla trajektorii. Ponieważ jednostka ucząca może czytać dane tak szybko, jak bufor jest w stanie je dostarczać, TPU nigdy nie przestaje pracować.
Efektem netto jest potok treningowy, w którym jednostki TPU pozostają zajęte prawie przez cały czas, zbliżając utylizację do 100%.
Przeszkody techniczne i sposób, w jaki Tunix im przeciwdziała
Epizody o zmiennej długości i rekompilacja XLA
Kompilator XLA w JAX optymalizuje pod kątem stałych kształtów tensorów. Zadania agentowe generują jednak sekwencje o różnej długości, co normalnie wywoływałoby kosztowne rekompilacje. Tunix pakuje krótsze sekwencje razem i grupuje epizody o podobnej długości w koszyki (buckets), utrzymując stabilne kształty wystarczająco długo, aby XLA mógł ponownie wykorzystać skompilowane jądra (kernels). Wynikiem jest stała przepustowość bez narzutu kompilatora, który w przeciwnym razie sparaliżowałby wydajność.
Skalowanie ogromnych modeli na wielu chipach TPU
Trenowanie agentów posiadających ponad 70 miliardów parametrów wymaga rozproszenia wag i danych na wiele węzłów TPU. Tunix wykorzystuje prymityw ShardMap z biblioteki JAX do shardingu zarówno parametrów modelu, jak i aktywacji, co pozwala jednostce uczącej na utrzymanie całego modelu w pamięci przy jednoczesnym szybkim dostarczaniu danych. Ta strategia shardingu umożliwia trenowanie modeli, które wcześniej były poza zasięgiem pojedynczego TPU pod.
Nieaktualne gradienty w rozdzielonych potokach
Gdy aktorzy wyprzedzają jednostkę uczącą, dostarczane przez nich dane mogą stać się „nieaktualne” (stale) względem obecnej polityki. Tunix łagodzi to przesunięcie za pomocą dwóch mechanizmów: próbkowanie ważone (importance sampling) nadaje starszym próbkom nowe wagi, aby odzwierciedlić ich istotność, a konfigurowalny próg nieaktualności odrzuca trajektorie, których wiek przekracza zadaną wartość. Razem zapewniają one stabilność uczenia, nawet gdy potok działa asynchronicznie.
Na co muszą zwrócić uwagę użytkownicy
- Audyt opóźnień – Korzyść z rozdzielenia komponentów zależy od czasu odpowiedzi środowiska. Zespoły powinny mierzyć opóźnienia end-to-end i upewnić się, że pule aktorów są odpowiednio skalowane, aby bufor był stale wypełniony.
- Projektowanie puli workerów – Tanie procesory CPU lub GPU mogą obsługiwać wielu aktorów, ale nadmierne ich zagęszczenie może spowodować rywalizację o zasoby sieci lub pamięci masowej. Niezbędna jest zrównoważona pula, która dopasuje tempo pobierania danych do tempa zapisu w buforze.
- Solidność bufora – Centralny magazyn musi obsługiwać wysokie tempo zapisu i odczytu, nie stając się nowym wąskim gardłem. Wybór systemu przechowywania danych o niskim opóźnieniu ogonowym (tail latency) i wystarczającej przepustowości jest kluczowym elementem architektury.
Potencjalne wady
Rozdzielona architektura wprowadza więcej ruchomych elementów: oddzielne floty sprzętowe, trwały bufor oraz logikę koordynacji w celu wymuszenia limitów nieaktualności danych.
Podsumowanie
Tunix pokazuje, że dominującym kosztem w agentowym RL nie jest sam model, lecz czas bezczynności spowodowany synchronicznymi pętlami interakcji. Przenosząc pracę nad generowaniem ról (rollout) na tani sprzęt i zasilając stale uczącą się jednostkę TPU pod z bufora o wysokiej przepustowości, Google przekształciło problem utylizacji poniżej 10% w przepływ pracy o niemal pełnej wydajności.
