LongStraw는 자사의 브랜치 리플레이(branch-replay) 기술을 통해 단 8개의 H20 GPU만으로 강화 학습(RL) 사후 학습을 위한 210만 개의 토큰 위치를 처리할 수 있으며, 이를 통해 하드웨어 비용을 획기적으로 절감할 수 있다고 발표했습니다. 이 발표가 중요한 이유는 롱 컨텍스트 모델을 학습시키는 데 전통적으로 수십 개의 고성능 GPU가 필요했으며, 이는 대부분의 연구소와 스타트업에 진입 장벽이었기 때문입니다.
롱 컨텍스트 RL이 비용이 많이 드는 이유
대규모 언어 모델의 RL 기반 미세 조정은 일반적으로 동일한 프롬프트에 대해 여러 대안적인 결과물을 생성하는 롤아웃(rollout)을 실행합니다. 각 롤아웃은 역전파(back-propagation) 과정을 거쳐야 하므로, 연산 비용은 처리되는 총 토큰 위치 수에 비례하여 증가합니다. 100만 토큰 컨텍스트를 목표로 하는 현재의 파이프라인은 적절한 시간 내에 작업을 완료하기 위해 종종 64~128개의 GPU를 필요로 합니다. 이러한 하드웨어 비용과 더불어 전력 및 냉각 비용은 실무자들이 컨텍스트 길이를 어디까지 확장할 수 있는지를 제한합니다.
브랜치 리플레이가 작업량을 줄이는 방법
LongStraw의 접근 방식은 트랜스포머 생성에 관한 두 가지 관찰 결과에 기반합니다:
- 프롬프트와 응답의 초기 부분은 모든 롤아웃에서 동일합니다.
- 각 응답의 갈라지는 뒷부분(divergent tail)만이 실제로 새로운 연산을 필요로 합니다.
이 시스템은 공유된 접두사(prefix)에 대한 활성화 값(activations)을 기록하는 아키텍처 인식 실행 스택을 구축합니다. 새로운 브랜치가 탐색될 때, 시스템은 이를 다시 계산하는 대신 캐시된 접두사를 리플레이하고, 새로운 세그먼트에 대해서만 역전파 패스를 실행합니다. 실제로 이는 역전파 패스가 훨씬 적은 토큰 위치를 건드리게 됨을 의미하며, 결과적으로 순수 연산량을 8~16배까지 줄여줍니다.
즉각적인 영향
- 8개의 H20 GPU로 210만 개의 토큰 위치를 처리했습니다. 이는 일반적인 하드웨어 예산으로는 해당 작업량의 극히 일부만 처리할 수 있는 수준입니다.
- 컨텍스트가 커짐에 따라 메모리 및 연산 비용이 폭증하는 롱 컨텍스트 RL의 병목 현상을 직접적으로 겨냥했습니다.
- 연구소들은 GPU 할당을 변경할 수 있습니다. 주로 추론 가속기로 사용되던 동일한 하드웨어를 이제 학습에 사용할 수 있게 되었으나, 다른 카드의 경우 결과가 다를 수 있습니다.
남은 과제 및 한계
이번 발표에는 학습 속도 수치와 수렴 곡선이 생략되어 있어, 연산량 절감이 실제 소요 시간(wall-clock time)의 단축으로 이어지는지, 아니면 단순히 GPU 점유율을 낮추는 것인지는 알 수 없습니다. 또한 이 방법은 자기 회귀(autoregressive) 샘플링을 기준으로 설명되었으며, 비자기 회귀(non-autoregressive) 또는 하이브리드 전략에서의 동작은 아직 테스트되지 않았습니다. H20은 주로 추론 가속기이므로, H100 또는 B200과 같은 더 일반적인 학습용 카드에서의 성능은 다를 수 있습니다.
독립적인 벤치마크를 통해 LongStraw의 수치가 아직 검증되지는 않았습니다. 제3자의 검증이 없으므로, 커뮤니티는 이 결과를 유망하지만 잠정적인 것으로 간주해야 합니다.
중요한 점
만약 브랜치 리플레이 아이디어가 Direct Preference Optimization (DPO) 또는 Proximal Policy Optimization (PPO)와 같은 다른 RL 미세 조정 알고리즘으로 확장된다면, 롱 컨텍스트 모델의 비용 장벽은 사라질 수 있습니다.
주목해야 할 사항
- 다양한 GPU 아키텍처에서의 제3자 재현 시도.
- 베이스라인 파이프라인과 비교한 학습 처리량(throughput) 및 최종 모델 품질에 대한 LongStraw의 업데이트.
