Reinforcement Learning for Hierarchical Reasoning Rewards: Minimax-Optimal Rates with Transformers
Organizations: The University of Tokyo, RIKEN AIP
Abstract
Reinforcement learning (RL) has become a standard tool for post-training language models on reasoning tasks, where the policy is updated by reward feedback while exploring the space of responses. Despite its empirical success, theoretical understanding of RL post-training remains limited, in particular of why on-policy exploration combined with a neural reward model is effective. In this paper, we address this question by modeling the reward as a hierarchical function on the response space: the reward consists of infinitely many local components, each of which becomes relevant only after the preceding ones have been resolved. We show that a natural Transformer-based actor--critic algorithm, which alternates between sampling from the current KL-regularized policy, fitting a Transformer critic to the observed rewards, and updating the policy, achieves the minimax optimal rates in the query budget and in the regularization strength up to logarithmic factors, and is minimax optimal for a fixed number of prompts. In contrast, we prove that sampling from the fixed reference distribution, as in offline reward modeling, can limit regret decay to a logarithmic rate. These results show that on-policy exploration progressively zooms in on the region where the reward is concentrated, and quantify its benefit for RL post-training.
Figures & tables
Appendix figures & tables3 assets
Supplementary material from the paper’s appendix.
Appendix
| Symbol | Meaning |
|---|---|
| weight of level of the hierarchy | |
| weight of the levels deeper than | |
| depth at which level is fully resolved | |
| , | active cell of level ; shell |
| reward truncated at level | |
| optimal reward of prompt |
| Symbol | Meaning | Defined in |
|---|---|---|
| accuracy scale attached to depth | ( 14 ) | |
| exponent of in the batch sizes | ( 14 ) | |
| , | score at phase ; partition on which it is constant | Lemma 19 |
| birth-accuracy level | Definition 11 | |
| , | correct depth- cube; the level it has reached | Section B.3 |
| Gibbs mass of the correct cube | Lemma 12 |
| Symbol | Meaning |
|---|---|
| number of front-phase iterations | |
| , | refinement depth of the lookahead; the lookahead itself |
| , | birth-accuracy level; moment order used in the regret transfer |
| , | total and per-phase failure budgets |
| input length of the front-phase critic | |
| , , | batch, query depth and approximation size of the final refinement step |