Google의 Tunix 시스템은 대규모 에이전트 기반 강화 학습(agentic RL)이 TPU를 효율적으로 사용하지 못하게 만들던 병목 지점을 해소합니다. 상호작용 데이터를 생성하는 작업과 정책을 업데이트하는 작업을 분리함으로써, Tunix는 TPU 활용률을 한 자릿수 퍼센트에서 거의 전체 용량에 가깝게 끌어올려 컴퓨팅 낭비를 획기적으로 줄입니다.
에이전트 기반 RL의 병목 현상
에이전트 기반 RL은 우리에게 익숙한 "next-token" 방식의 언어 모델 학습과는 다릅니다. 에이전트는 API 호출을 보내거나, 코드를 실행하거나, 시뮬레이션 환경을 단계별로 진행한 뒤 그 결과에 반응해야 합니다. 따라서 학습 루프는 동기식(synchronous)으로 작동합니다. 모델이 행동을 생성하면 환경이 실행되고, 결과가 반환된 후에야 모델이 그래디언트 업데이트를 받을 수 있습니다. 단일 환경 단계(step)에 몇 초가 소요될 경우, 고가의 TPU 하드웨어는 유휴 상태로 방치되며 보고된 활용률은 10% 미만으로 떨어질 수 있습니다. 이러한 비효율성은 클라우드 비용 상승과 연구 주기 지연으로 직결됩니다.
Tunix의 디커플링된 아키텍처
Tunix는 궤적 생성(trajectory generation)과 정책 최적화(policy optimization)라는 두 단계를 별도의 하드웨어 풀로 분리하여 이 문제를 해결합니다.
- **비동기 액터(Asynchronous actors)**는 저렴한 CPU 또는 GPU에서 실행됩니다. 각 액터는 할당된 환경과 지속적으로 상호작용하며 행동과 관측치를 기록하고, 생성된 궤적을 공유 저장소로 스트리밍합니다.
- **지속적 학습자(Continuous learners)**는 전용 TPU Pod를 점유합니다. 학습자는 중앙 버퍼에서 배치를 가져와 단일 액터의 롤아웃(rollout)이 끝나기를 기다리지 않고 그래디언트 업데이트를 수행합니다.
- **고처리량 버퍼(High-throughput buffer)**는 중간에서 궤적을 위한 스테이징 영역 역할을 합니다. 학습자가 버퍼가 데이터를 공급하는 속도만큼 빠르게 읽을 수 있기 때문에 TPU는 멈추지 않습니다.
결과적으로 TPU가 거의 항상 작동하여 활용률을 100%에 가깝게 밀어붙이는 학습 파이프라인이 구축됩니다.
기술적 난제와 Tunix의 해결 방법
가변 길이 에피소드 및 XLA 재컴파일
JAX의 XLA 컴파일러는 고정된 텐서 형태(tensor shapes)에 최적화되어 있습니다. 그러나 에이전트 작업은 길이가 서로 다른 시퀀스를 생성하며, 이는 보통 비용이 많이 드는 재컴파일을 유발합니다. Tunix는 짧은 시퀀스들을 함께 묶고 유사한 길이의 에피소드를 버킷(bucket)으로 그룹화하여, XLA가 컴파일된 커널을 재사용할 수 있을 만큼 형태를 안정적으로 유지합니다. 그 결과, 성능을 저하시킬 수 있는 컴파일 오버헤드 없이 꾸준한 처리량을 유지합니다.
다수의 TPU 칩에 걸친 대규모 모델 스케일링
700억 개 이상의 파라미터를 가진 에이전트를 학습시키려면 여러 TPU 노드에 가중치와 데이터를 분산해야 합니다. Tunix는 JAX의 ShardMap 프리미티브를 사용하여 모델 파라미터와 활성화 함수(activations)를 모두 샤딩(sharding)함으로써, 학습자가 전체 모델을 메모리에 유지하면서도 데이터를 고속으로 공급할 수 있게 합니다. 이러한 샤딩 전략을 통해 이전에는 단일 TPU Pod로는 불가능했던 모델 학습이 가능해졌습니다.
디커플링된 파이프라인으로 인한 오래된(stale) 그래디언트
액터가 학습자보다 앞서 실행될 경우, 그들이 공급하는 데이터는 현재 정책과 비교했을 때 "오래된(stale)" 상태가 될 수 있습니다. Tunix는 두 가지 메커니즘으로 이 드리프트(drift)를 완화합니다. 첫째, 중요도 샘플링(importance-sampling)을 통해 오래된 샘플에 관련성을 반영하도록 재가중치를 부여합니다. 둘째, 설정 가능한 신선도 임계값(staleness threshold)을 통해 미리 설정된 연령을 초과하는 궤적은 폐기합니다. 이 두 가지 방식은 파이프라인이 비동기적으로 작동하는 동안에도 학습을 안정적으로 유지합니다.
도입 시 주의 사항
- 지연 시간(Latency) 감사 – 디커플링의 이점은 환경의 응답 시간에 달려 있습니다. 팀은 엔드 투 엔드(end-to-end) 지연 시간을 측정하고, 버퍼가 충분히 채워질 수 있도록 액터 풀의 규모를 조정해야 합니다.
- 워커 풀(Worker pool) 설계 – 저렴한 CPU나 GPU에서 많은 액터를 호스팅할 수 있지만, 과도하게 할당하면 네트워크나 스토리지에서 경합(contention)이 발생할 수 있습니다. 버퍼의 수집 속도와 일치하는 균형 잡힌 풀을 구성하는 것이 필수적입니다.
- 버퍼 견고성 – 중앙 저장소는 새로운 병목 지점이 되지 않으면서 높은 쓰기 및 읽기 속도를 처리할 수 있어야 합니다. 낮은 테일 레이턴시(tail latency)와 충분한 대역폭을 갖춘 스토리지 시스템을 선택하는 것은 아키텍처의 필수 요소입니다.
잠재적 단점
분리된 아키텍처는 별도의 하드웨어 플릿(fleet), 지속적인 버퍼, 신선도 제한을 강제하기 위한 조정 로직 등 더 많은 구성 요소를 도입하게 됩니다.
요약
Tunix는 에이전트 기반 RL에서 지배적인 비용이 모델 자체가 아니라 동기식 상호작용 루프로 인해 발생하는 유휴 시간임을 보여줍니다. 롤아웃 작업을 저렴한 하드웨어로 오프로딩하고 고처리량 버퍼를 통해 지속적으로 학습하는 TPU Pod에 데이터를 공급함으로써, Google은 10% 미만의 활용률 문제를 거의 전체 용량을 사용하는 워크플로우로 탈바꿈시켰습니다.
