Stanford CS329A Self-Improving AI Agents | Part 5 | Planning and Multi-Step Reasoning
Watch on YouTube →
Overview
This lecture explores advanced AI planning and multi-step reasoning through three papers: LATs, SPRINT, and SWiRL. LATs unifies reasoning, acting, and planning in LLMs using Monte Carlo Tree Search (MCTS) and reflection for improved exploration. SPRINT enables LLMs to identify and execute parallelizable reasoning steps, significantly reducing inference time and improving accuracy. SWiRL focuses on training LLMs for multi-step reasoning and tool use via RL, demonstrating generalization across datasets and tools without direct tool execution during training.
Key takeaways
- LATs unifies reasoning, acting, and planning in LLMs using MCTS, scoring actions based on outcomes and incorporating reflection for improved exploration.
- SPRINT enables LLMs to identify and execute parallelizable reasoning steps, significantly reducing inference time and improving accuracy by restructuring thought processes.
- SWiRL trains LLMs for multi-step reasoning and tool use via RL, demonstrating generalization across datasets and tools without direct tool execution during training, and showing RL's advantage over SFT.
- The effectiveness of SWiRL's training is enhanced by using process-filtered synthetic data, where LLM-as-a-Judge validates reasoning steps rather than solely relying on final outcome correctness.
- SPRINT's approach to parallelizing LLM reasoning leads to substantial reductions in sequential token generation (up to 40%) and improved accuracy, particularly for complex tasks requiring extensive thinking.
- SWiRL's RL-based training allows models to learn generalized multi-step thinking and tool invocation capabilities, transferring effectively to new tools and domains beyond the training set.
Chapters
- Lecture covers three papers on multi-step reasoning and planning.
- Focus on tasks requiring reasoning, acting, and planning trajectories.
- Examples include trip planning, involving reasoning, information gathering (acting), and plan refinement (search).
- LATs paper from ICML aims to unify reasoning, acting, and planning in LLMs.
- Addresses the challenge of encouraging diversification and exploration of different plans.
- Integrates techniques from reinforcement learning and multi-step planning, including MCTS.
- Example: Planning a trip to Hawaii involves sampling actions and scoring them.
- Actions like asking friends or reading subreddits are evaluated.
- Scores update the search space and trajectory expansion based on feedback.
- Math-Shepherd uses a verifier to score reasoning trajectories for search guidance.
- LATs scores are based on the outcomes of actions, not just reasoning steps.
- LATs incorporates reflection and environment observations, improving upon ReAct's memorization and interaction.
- Intuition: Chain-of-thought reasoning decomposed into steps, forming a tree via actions.
- Search and planning occur through this tree.
- Feedback from actions is incorporated into the search process.
- Example: Navigating a maze to reach an exit.
- Initial state: dimly lit room with two doors.
- Selection stage uses UCT (Upper Confidence bounds applied to Trees) to choose a node to expand.
- Sampled actions: 'open left door', 'open right door', 'inspect room'.
- Actions are executed, and observations are appended to context.
- Evaluation combines an LLM-as-a-Judge score (0-1) and a self-consistency score (action frequency).
- Simulation: Greedy expansion from the highest-scoring state until an end state is reached.
- Backpropagation: The return of a trajectory (success/failure) updates values of previous actions and states.
- Value function is a weighted average of LLM-as-a-Judge and self-consistency.
- UCT balances exploration (encouraging less visited nodes) and exploitation (choosing high-value nodes).
- Formula: V(s) + C * sqrt(ln(Np) / Ns), where Np is parent visits, Ns is node visits.
- Backpropagation updates state values using the trajectory's return and node visit counts.
- Reflection: Model adds thinking about why a trajectory succeeded or failed.
- Tested on HotPotQA (multi-hop QA requiring retrieval from multiple pages).
- Performance improves with more samples; reflection boosts results.
- Tested on WebShop (e-commerce task requiring multi-step search and selection).
- Achieved high results without fine-tuning, close to human experts.
- LATs unifies reasoning, acting, planning using MCTS, but incurs significant computational cost.
- SPRINT addresses the observation that harder problems require longer thinking, correlating with higher accuracy.
- Identifies independent reasoning steps that can be executed in parallel.
- Framework enables LLMs to plan and execute responses in parallel, accelerating reasoning.
- Uses LLMs (e.g., GPT-4o) to annotate existing reasoning trajectories.
- Annotates planning vs. execution steps and identifies parallelizable subtasks.
- Creates a DAG (Directed Acyclic Graph) of steps to optimize execution order.
- Fine-tunes models (e.g., DeepSeek-R1) on annotated data to generate parallel plans and executions.
- At inference, model outputs tags indicating parallel plans and their executions.
- Enables parallel execution of independent plans, reducing sequential token generation and cost.
- Generates 6k thinking trajectories for MATH dataset, keeping highly parallelizable ones.
- Supervised fine-tuning on models like DeepSeek-R1 and Distill-Qwen-7B.
- SPRINT improves accuracy and efficiency, reducing sequential tokens by ~40% compared to baselines.
- Demonstrates out-of-domain generalization (e.g., MATH training, GPQA testing).
- Parallelism is task-dependent; harder problems with more thinking benefit most.
- Savings are larger for problems requiring more sequential tokens and parallelizable steps.
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.