Stop Regressing: Training Value Functions via Classification for Scalable Deep RL

Jesse Farebrother, Jordi Orbay, Q. Vuong, Adrien Ali Taiga, Yevgen Chebotar, Ted Xiao, A. Irpan, Sergey Levine, Pablo Samuel Castro, Aleksandra Faust, Aviral Kumar, Rishabh Agarwal

arXiv:2403.03950 · 2026-07-27 공개 · arXiv · PDF

reinforcement-learning classification transformers scalability deep-rl value-functions atari-2600 q-transformers

Abstract

Value functions are a central component of deep reinforcement learning (RL). These functions, parameterized by neural networks, are trained using a mean squared error regression objective to match bootstrapped target values. However, scaling value-based RL methods that use regression to large networks, such as high-capacity Transformers, has proven challenging. This difficulty is in stark contrast to supervised learning: by leveraging a cross-entropy classification loss, supervised methods have scaled reliably to massive networks. Observing this discrepancy, in this paper, we investigate whether the scalability of deep RL can also be improved simply by using classification in place of regression for training value functions. We demonstrate that value functions trained with categorical cross-entropy significantly improves performance and scalability in a variety of domains. These include: single-task RL on Atari 2600 games with SoftMoEs, multi-task RL on Atari with large-scale ResNets, robotic manipulation with Q-transformers, playing Chess without search, and a language-agent Wordle task with high-capacity Transformers, achieving state-of-the-art results on these domains. Through careful analysis, we show that the benefits of categorical cross-entropy primarily stem from its ability to mitigate issues inherent to value-based RL, such as noisy targets and non-stationarity. Overall, we argue that a simple shift to training value functions with categorical cross-entropy can yield substantial improvements in the scalability of deep RL at little-to-no cost.

한국어 요약

한 줄 요약

가치 함수 학습에서 회귀 대신 분류 손실을 사용하면 다양한 도메인에서 성능과 확장성이 크게 향상된다.

핵심 기여도

핵심 아이디어

기존 가치 기반 강화 학습(RL)은 **회귀 손실**(mean squared error, MSE)을 사용해 가치 함수를 학습하지만, 이는 대규모 네트워크(예: Transformer)에서 확장성이 떨어진다. 반면, **감독 학습**(supervised learning)에서는 **분류 손실**(cross-entropy)이 대규모 모델 확장에 효과적임을 고려, 본 연구는 **가치 함수 학습을 분류 문제로 재구성**함으로써 RL의 확장성을 개선할 수 있는지 조사했다.

**가치 함수 학습을 분류 문제로 바꾸는 핵심 아이디어**는 다음과 같다:

이러한 접근은 **노이즈 타겟**과 **비정상성** 문제를 완화하고, 네트워크가 더 **유연하게 학습**할 수 있도록 돕는다. 특히, **Transformer**와 같은 대규모 모델에서 **성능과 확장성**을 동시에 개선하는 데 기여함을 실험적으로 입증했다.

기술적 접근법

주요 결과

의의 및 한계

본 연구는 **가치 함수 학습을 분류 문제로 재구성**함으로써 **확장성과 성능**을 동시에 개선할 수 있음을 입증했다. 특히, **Transformer**와 같은 대규모 모델에서 **기존 회귀 기반 방법 대비 훨씬 높은 성능**을 달성한 점이 주목할 만하다. 이는 **강화 학습 알고리즘 설계**에 있어 **분류 손실의 활용**이 중요한 단서를 제공한다.

하지만, **분류 손실**은 **모든 문제를 완전히 해결하지는 않으며**, **타겟 분포의 정확한 정의**나 **분류 라벨 간 간격**(spacing) 설정 등 **추가적인 연구**가 필요하다. 또한, **분류 손실**이 **모든 RL 도메인**에서 동일한 효과를 보이는지에 대한 **추가 실험**이 필요하다는 한계점도 존재한다.

실용적 활용

본 연구의 접근법은 **대규모 네트워크**(예: Transformer)를 사용하는 **로봇 조작**, **언어 에이전트**, **게임 플레이** 등 다양한 **산업 및 연구 분야**에서 활용 가능하다. 특히, **강화 학습 알고리즘의 확장성**을 향상시키는 데 기여할 수 있으며, **CQL**, **DQN** 등 기존 알고리즘에 **분류 손실을 적용**함으로써 **성능 개선**을 기대할 수 있다.