Stanford CS329A Self-Improving AI Agents | Part 6 | Train Time Scaling/Scaling RL
Watch on YouTube →
Overview
This lecture explores train-time scaling techniques to enhance AI model reasoning capabilities, focusing on three papers: STaR, DeepSeekMath, and DAPO. These methods aim to improve model performance on complex reasoning tasks like the AIME benchmark by leveraging self-generated data and reinforcement learning, demonstrating that increased training compute can compensate for fewer model parameters. Key insights include the effectiveness of rationalization, the challenges of RL implementation, and the importance of verifiability in domains like mathematics and coding.
Key takeaways
- Train-time scaling, exemplified by STaR, DeepSeekMath, and DAPO, significantly boosts AI reasoning by leveraging self-generated data and RL, outperforming larger models on benchmarks like AIME.
- STaR iteratively refines reasoning by generating solutions, filtering correct ones, and using hints to create rationales for incorrect attempts, assuming output correctness proxies reasoning quality.
- DeepSeekMath's GRPO algorithm enables scalable RL for mathematical reasoning by reducing memory overhead, achieving 51.7% on the Math benchmark with a 7B model.
- DAPO addresses RL instability in complex reasoning by employing asymmetric clipping for exploration and dynamic sampling to maintain a useful reward signal, improving Qwen-32B on AIME to 50%.
- While these techniques enhance model consistency and reasoning chain coherence, they do not yet fundamentally improve a model's ability to solve entirely new, out-of-domain problems.
Chapters
- Lecture 6 covers train time scaling, focusing on improving models through self-generated outputs.
- Three papers will be discussed: STaR (reasoning with rationales), DeepSeekMath (mathematical reasoning with RL), and DAPO (stabilizing RL for long reasoning chains).
- AIME benchmark performance shows significant gains with train-time scaling (DeepSeekMath 51.7%, DAPO on Qwen-32B 50%) compared to GPT-3.5 (5%).
- Train time scaling involves using model outputs, filtered cleverly, to further fine-tune the model.
- Compute invested in training on self-outputs can substitute for model parameters.
- This contrasts with test-time scaling (inference-based techniques like majority voting) and pre-training.
- Test-time compute is generally cheaper, allowing for multiple inferences to find correct answers.
- Train-time scaling requires careful implementation and sufficient successes in the feedback loop.
- The AIME benchmark shows accuracy increases with both train-time and test-time compute.
- Reasoning enables solving difficult problems by dedicating tokens to step-by-step processes.
- Reasoning involves problem analysis, task decomposition, self-evaluation, and parallel search.
- Domains with verifiability (coding, math) benefit more from reasoning techniques than creative writing.
- STaR addresses the lack of reasoning steps in internet-scale data and the cost of manual annotation.
- It iteratively generates solutions, filters for correct answers, and uses hints to generate rationales for incorrect attempts.
- The core assumption is that final output correctness is a proxy for reasoning quality.
- STaR fine-tunes on successful reasoning paths and uses hints to generate rationales for failed attempts.
- It assumes the initial language model is strong enough to bootstrap from few-shot examples.
- Challenges include evaluating rationale quality and potential issues with rationalizing incorrect reasoning chains.
- The algorithm starts with a small rationale dataset and a large training dataset (questions and answers).
- It generates rationales for training data, fine-tunes on correct solutions, and uses hints for incorrect ones.
- Experiments were conducted on GPT-J (6B parameters) using GSM8K, CommonsenseQA, and synthetic arithmetic problems.
- STaR boosted accuracy on math benchmarks (e.g., 51.7% on GSM8K) using less data than direct fine-tuning.
- It showed reasonable qualitative results on CommonsenseQA rationales.
- Limitations include plateauing performance after iterations and potential issues with simple problems where reasoning is less critical.
- V-STaR adds a verifier to the STaR loop, while Quiet-STaR uses internal latent space reasoning.
- STaR's performance is bounded by the base model's reasoning capability and the quality of rationalization data.
- The ability to make logical leaps into new domains is limited by the training data.
- DeepSeekMath achieved 51.7% accuracy on the Math benchmark with a 7B model, surpassing larger models.
- Key innovations include pre-training on curated Common Crawl data (OpenWebMath) from a code-trained model (DeepSeek Coder).
- It proposed GRPO (Generalized Proximal Policy Optimization) to reduce memory requirements for RL.
- GRPO replaces the critic and value function in PPO with a generalized advantage estimation using a group baseline.
- This reduces model copies from four to three, enabling RL scaling.
- The advantage is calculated as (reward - mean reward) / standard deviation of rewards.
Summary, takeaways, and chapters were generated by AI from the video's transcript and may contain errors. The video belongs to its creator, Stanford Online.