Stanford CS329A Self-Improving AI Agents | Part 3 | Robust Verification
Watch on YouTube →
Overview
This lecture explores the evolution of verification techniques for Large Language Models (LLMs), starting with OpenAI's 2021 paper on training verifiers for math problems using the GSM8K dataset. It progresses through OpenAI's 2023 'Verify, Step-by-Step' paper introducing process-based reward models (PRMs) over outcome-based reward models (ORMs), Stanford's 2024 'Math-Shepherd' which automates PRM label collection, and finally Stanford's recent 'Weaver' that uses an ensemble of weak verifiers with weak-to-strong supervision for robust verification without extensive human annotation.
Key takeaways
- The GSM8K dataset, introduced in 2021, became a crucial benchmark for evaluating LLM reasoning capabilities in math problems.
- Process-based reward models (PRMs) significantly outperform outcome-based reward models (ORMs) by evaluating each step of an LLM's reasoning, as demonstrated by the PRM800K dataset and Math-Shepherd.
- Automating label collection for PRMs, as done in Math-Shepherd, drastically reduces the need for expensive human annotation, enabling more scalable verification.
- Ensembling multiple weak verifiers and using weak-to-strong supervision, as in Weaver, can create a robust verification system that significantly boosts LLM performance, even with limited labeled data.
- Distilling complex ensemble verifiers like Weaver into smaller models is a viable strategy to achieve high accuracy with significantly reduced computational cost at inference time.
- Verification techniques, including PRMs and ensemble methods, are crucial for improving LLM reasoning, reducing hallucinations, and bridging the generation-verification gap across various domains, not just math.
Chapters
- Lecture 3 focuses on verification to address the generation-verification gap where LLMs can produce plausible but incorrect answers.
- The goal is to automatically select correct answers or guide generation, building on last lecture's inference-time scaling techniques.
- Four papers will be discussed, showing a progression in verification approaches over time.
- Motivation: LLMs hallucinate and confidently present wrong solutions.
- Introduced GSM8K, a reasoning benchmark of 8,500 grade-school math problems requiring multi-step reasoning.
- Trained a verifier model to output the probability of a solution being correct.
- Verifier trained on question-solution pairs with binary correct/incorrect labels.
- Generator samples 100 solutions per problem; labels are created based on known ground truth.
- Test-time usage involves generating multiple answers and using the verifier's score to select the best one.
- Verifier trained with both a binary correctness loss and a standard language modeling (next token prediction) loss.
- Architecture is a language model with a scalar head for binary prediction on a per-token basis.
- Tokens in the question are masked; optimization focuses on solution tokens.
- Generator fine-tuned for two epochs, then 100 completions sampled and labeled.
- Verifier trained for a single epoch on this labeled dataset.
- Ablations explored sentence-level vs. token-level labeling; final solution score derived from the last token's prediction.
- Comparison shows verification outperforms fine-tuning-only approaches, especially with larger verifier training sets.
- For smaller verifier datasets (<1000 samples for a 175B model), verification offers less benefit.
- Larger generators with smaller verifiers performed better than smaller generators with larger verifiers.
- Increasing the number of completions per problem improves accuracy up to around 400 samples.
- Beyond 400 samples, the verifier's precision drops, failing to consistently track the best solution.
- This contrasts with majority voting, which fails to track after ~50 samples.
- Addresses LLM missteps in multi-step reasoning that derail entire answers.
- Introduces outcome-based reward models (ORMs) and process-based reward models (PRMs).
- ORMs reward the entire solution; PRMs reward each step of the reasoning process.
- PRM training involves human annotators scoring each step of a generated solution.
- This creates a dataset of step-level correct/incorrect labels (PRM800K).
- Process supervision provides more precise data and better credit assignment than outcome supervision.
- PRM outperforms ORM and majority voting, especially for rare correct answer occurrences.
- PRM is more data-efficient than ORM; fewer labeled steps are needed for comparable performance.
- PRM generalizes better to new domains and tolerates distribution shifts more than majority voting.
- Addresses the high cost of human annotation for PRMs.
- Automates step-level annotation by measuring the potential of a step to reach a final correct answer.
- Uses 'hard' (any correct path) and 'soft' (frequency of correct paths) estimates, finding hard estimates sufficient with enough samples.
- Samples N candidate solutions, scores them with a PRM, and selects the highest-scoring one.
- PRM can also be used as a reward model for RL fine-tuning the generator.
- Outperforms baselines like self-consistency (majority voting) and ORMs on GSM8K and MATH datasets without human annotation.
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.