Learn AI Series (#142) - AI Reasoning and Planning
What will I learn
- The System 1 vs System 2 framing (fast intuition vs slow deliberation) and why a plain transformer is a System 1 machine bolted shut at a fixed amount of thinking;
- chain-of-thought: why simply making a model "show its work" hands it more effective compute, and how that connects to adaptive computation;
- tree-of-thought: exploring several reasoning paths at once with an evaluator to prune the bad branches, which is just beam search wearing a different hat;
- planning with MCTS and A*: classical search algorithms steered by learned neural heuristics, the exact recipe behind AlphaGo;
- process reward models: scoring each reasoning STEP instead of only the final answer, and why that dense signal is such a big deal;
- an honest, no-hype assessment of the reasoning gap -- what current systems genuinely do versus what they only appear to do.
Requirements
- A working modern computer running macOS, Windows or Ubuntu;
- Python 3.10+ with PyTorch installed (
pip install torch) -- everything here is illustrative, nothing needs a GPU or a fine-tuned model to make the point; - You've been through language modeling (#57), the GPT architecture (#58), the AI agents arc (#67-68), and it helps to have just read the robotics episode (#141), since we pick up its final thread. We also lean on beam search (#71) and RL for games (#112).
Difficulty
- Beginner
Curriculum (of the Learn AI Series):
- Learn AI Series (#1) - What Machine Learning Actually Is
- Learn AI Series (#2) - Setting Up Your AI Workbench - Python and NumPy
- Learn AI Series (#3) - Your Data Is Just Numbers - How Machines See the World
- Learn AI Series (#4) - Your First Prediction - No Math, Just Intuition
- Learn AI Series (#5) - Patterns in Data - What "Learning" Actually Looks Like
- Learn AI Series (#6) - From Intuition to Math - Why We Need Formulas
- Learn AI Series (#7) - The Training Loop - See It Work Step by Step
- Learn AI Series (#8) - The Math You Actually Need (Part 1) - Linear Algebra
- Learn AI Series (#9) - The Math You Actually Need (Part 2) - Calculus and Probability
- Learn AI Series (#10) - Your First ML Model - Linear Regression From Scratch
- Learn AI Series (#11) - Making Linear Regression Real
- Learn AI Series (#12) - Classification - Logistic Regression From Scratch
- Learn AI Series (#13) - Evaluation - How to Know If Your Model Actually Works
- Learn AI Series (#14) - Data Preparation - The 80% Nobody Talks About
- Learn AI Series (#15) - Feature Engineering and Selection
- Learn AI Series (#16) - Scikit-Learn - The Standard Library of ML
- Learn AI Series (#17) - Decision Trees - How Machines Make Decisions
- Learn AI Series (#18) - Random Forests - Wisdom of Crowds
- Learn AI Series (#19) - Gradient Boosting - The Kaggle Champion
- Learn AI Series (#20) - Support Vector Machines - Drawing the Perfect Boundary
- Learn AI Series (#21) - Mini Project - Predicting Crypto Market Regimes
- Learn AI Series (#22) - K-Means Clustering - Finding Groups
- Learn AI Series (#23) - Advanced Clustering - Beyond K-Means
- Learn AI Series (#24) - Dimensionality Reduction - PCA
- Learn AI Series (#25) - Advanced Dimensionality Reduction - t-SNE and UMAP
- Learn AI Series (#26) - Anomaly Detection - Finding What Doesn't Belong
- Learn AI Series (#27) - Recommendation Systems - "Users Like You Also Liked..."
- Learn AI Series (#28) - Time Series Fundamentals - When Order Matters
- Learn AI Series (#29) - Time Series Forecasting - Predicting What Comes Next
- Learn AI Series (#30) - Natural Language Processing - Text as Data
- Learn AI Series (#31) - Word Embeddings - Meaning in Numbers
- Learn AI Series (#32) - Bayesian Methods - Thinking in Probabilities
- Learn AI Series (#33) - Ensemble Methods Deep Dive - Stacking and Blending
- Learn AI Series (#34) - ML Engineering - From Notebook to Production
- Learn AI Series (#35) - Data Ethics and Bias in ML
- Learn AI Series (#36) - Mini Project - Complete ML Pipeline
- Learn AI Series (#37) - The Perceptron - Where It All Started
- Learn AI Series (#38) - Neural Networks From Scratch - Forward Pass
- Learn AI Series (#39) - Neural Networks From Scratch - Backpropagation
- Learn AI Series (#40) - Training Neural Networks - Practical Challenges
- Learn AI Series (#41) - Optimization Algorithms - SGD, Momentum, Adam
- Learn AI Series (#42) - PyTorch Fundamentals - Tensors and Autograd
- Learn AI Series (#43) - PyTorch Data and Training
- Learn AI Series (#44) - PyTorch nn.Module - Building Real Networks
- Learn AI Series (#45) - Convolutional Neural Networks - Theory
- Learn AI Series (#46) - CNNs in Practice - Classic to Modern Architectures
- Learn AI Series (#47) - CNN Applications - Detection, Segmentation, Style Transfer
- Learn AI Series (#48) - Recurrent Neural Networks - Sequences
- Learn AI Series (#49) - LSTM and GRU - Solving the Memory Problem
- Learn AI Series (#50) - Sequence-to-Sequence Models
- Learn AI Series (#51) - Attention Mechanisms
- Learn AI Series (#52) - The Transformer Architecture (Part 1)
- Learn AI Series (#53) - The Transformer Architecture (Part 2)
- Learn AI Series (#54) - Vision Transformers
- Learn AI Series (#55) - Generative Adversarial Networks
- Learn AI Series (#56) - Mini Project - Building a Transformer From Scratch
- Learn AI Series (#57) - Language Modeling - Predicting the Next Word
- Learn AI Series (#58) - GPT Architecture - Decoder-Only Transformers
- Learn AI Series (#59) - BERT and Encoder Models
- Learn AI Series (#60) - Training Large Language Models
- Learn AI Series (#61) - Instruction Tuning and Alignment
- Learn AI Series (#62) - Prompt Engineering - Getting the Most from LLMs
- Learn AI Series (#63) - Embeddings and Vector Search
- Learn AI Series (#64) - Retrieval-Augmented Generation (RAG) - Basics
- Learn AI Series (#65) - RAG - Advanced Techniques
- Learn AI Series (#66) - Working with LLM APIs
- Learn AI Series (#67) - Building AI Agents (Part 1) - Foundations
- Learn AI Series (#68) - Building AI Agents (Part 2) - Advanced Patterns
- Learn AI Series (#69) - Fine-Tuning Language Models
- Learn AI Series (#70) - Running Local Models
- Learn AI Series (#71) - Text Generation Techniques
- Learn AI Series (#72) - Tokenization Deep Dive
- Learn AI Series (#73) - LLM Evaluation
- Learn AI Series (#74) - The Hugging Face Ecosystem
- Learn AI Series (#75) - Multimodal Models - Text Meets Vision
- Learn AI Series (#76) - Mini Project - Your Own AI Assistant
- Learn AI Series (#77) - Image Processing Fundamentals
- Learn AI Series (#78) - Object Detection (Part 1) - Foundations
- Learn AI Series (#79) - Object Detection (Part 2) - Modern Approaches
- Learn AI Series (#80) - Image Segmentation
- Learn AI Series (#81) - Pose Estimation and Tracking
- Learn AI Series (#82) - Optical Character Recognition
- Learn AI Series (#83) - Video Understanding
- Learn AI Series (#84) - Generative Images - Diffusion Models (Part 1)
- Learn AI Series (#85) - Generative Images - Diffusion Models (Part 2)
- Learn AI Series (#86) - Image-to-Image and Editing
- Learn AI Series (#87) - 3D Vision
- Learn AI Series (#88) - Face Analysis
- Learn AI Series (#89) - Medical and Scientific Imaging
- Learn AI Series (#90) - Self-Supervised Learning for Vision
- Learn AI Series (#91) - Mini Project - Building a Visual AI System
- Learn AI Series (#92) - Audio Fundamentals for AI
- Learn AI Series (#93) - Speech Recognition
- Learn AI Series (#94) - Text-to-Speech (TTS)
- Learn AI Series (#95) - Audio Classification
- Learn AI Series (#96) - Music Generation
- Learn AI Series (#97) - Speaker Recognition and Diarization
- Learn AI Series (#98) - Natural Language Understanding for Voice
- Learn AI Series (#99) - Audio Enhancement
- Learn AI Series (#100) - Multimodal Audio-Visual Models
- Learn AI Series (#101) - Mini Project: Voice-Controlled AI Assistant
- Learn AI Series (#102) - What Is Reinforcement Learning?
- Learn AI Series (#103) - Multi-Armed Bandits
- Learn AI Series (#104) - Dynamic Programming
- Learn AI Series (#105) - Monte Carlo Methods
- Learn AI Series (#106) - Temporal Difference Learning
- Learn AI Series (#107) - Deep Q-Networks (DQN)
- Learn AI Series (#108) - Policy Gradient Methods
- Learn AI Series (#109) - Advanced Policy Optimization
- Learn AI Series (#110) - Model-Based Reinforcement Learning
- Learn AI Series (#111) - Multi-Agent Reinforcement Learning
- Learn AI Series (#112) - RL for Games
- Learn AI Series (#113) - RL for Real-World Applications
- Learn AI Series (#114) - Inverse Reinforcement Learning
- Learn AI Series (#115) - Offline Reinforcement Learning
- Learn AI Series (#116) - Mini Project: Training a Game-Playing AI
- Learn AI Series (#117) - ML System Design
- Learn AI Series (#118) - Data Engineering for AI
- Learn AI Series (#119) - Experiment Tracking and Reproducibility
- Learn AI Series (#120) - Model Optimization: Making Models Fast
- Learn AI Series (#121) - Model Serving Architecture
- Learn AI Series (#122) - Edge AI: Running Models on Devices
- Learn AI Series (#123) - Monitoring ML in Production
- Learn AI Series (#124) - CI/CD for Machine Learning
- Learn AI Series (#125) - GPU Programming Basics
- Learn AI Series (#126) - Distributed Training
- Learn AI Series (#127) - AI Security
- Learn AI Series (#128) - Privacy-Preserving AI
- Learn AI Series (#129) - AutoML and Neural Architecture Search
- Learn AI Series (#130) - Causal Inference and ML
- Learn AI Series (#131) - Graph Neural Networks
- Learn AI Series (#132) - AI for Structured Data
- Learn AI Series (#133) - Synthetic Data Generation
- Learn AI Series (#134) - AI Infrastructure Economics
- Learn AI Series (#135) - Building AI Teams and Processes
- Learn AI Series (#136) - Mini Project: Production AI Platform
- Learn AI Series (#137) - Foundation Models
- Learn AI Series (#138) - Multimodal AI
- Learn AI Series (#139) - AI for Code
- Learn AI Series (#140) - Scientific AI
- Learn AI Series (#141) - Robotics and Embodied AI
- Learn AI Series (#142) - AI Reasoning and Planning (this post)
Learn AI Series (#142) - AI Reasoning and Planning
I left you last time with a deliberate loose thread. All of #141's robotics lived on reflex -- see a state, emit an action, see the next state, react again -- and I asked what it takes to build a model that does not just REACT, but looks several steps ahead, weighs options it has never tried, and commits to a PLAN before it moves. That question is not really about robots at all. It is about the deepest weakness in everything we've built across 141 episodes: our models are magnificent at fast pattern matching and surprisingly bad at slow, deliberate thinking.
Ask a language model "what's the capital of France?" and it fires back "Paris" instantly. Ask it "if I have 3 boxes, each holding 2 red balls and 1 blue ball, and I draw one ball from each box, what's the probability that exactly two are red?" and it stumbles -- unless you first tell it to think step by step, at which point it often gets it right. That gap, between the instant answer and the one that needs working-out, is the whole subject of today. Let's dive right in ;-)
Solutions to episode #141's exercises
House rules, same as ever: we clear last time's homework before we touch anything new. Episode #141 was robotics, and all three tasks were about feeling a failure mode in your own hands rather than reading about it.
Exercise 1 -- Feel the compounding error. Train a BehaviorCloningPolicy on a trivial 1D "track the target" env where the expert always starts near 0, then evaluate from a starting position FAR from 0 and watch it wander. Here is the whole thing, expert, training and the damning evaluation.
import torch
import torch.nn as nn
class BehaviorCloningPolicy(nn.Module):
"""Map a 1D position to a 1D move (see #141)."""
def __init__(self, hidden=64):
super().__init__()
self.net = nn.Sequential(
nn.Linear(1, hidden), nn.ReLU(),
nn.Linear(hidden, 1), nn.Tanh(),
)
def forward(self, obs):
return self.net(obs)
def expert_action(pos):
# the expert always nudges toward 0, capped at [-1, 1]
return torch.clamp(-pos, -1.0, 1.0)
# expert demonstrations: ALWAYS start near 0
obs, act = [], []
for _ in range(2000):
p = torch.randn(1) * 0.3 # starts hug the origin
obs.append(p); act.append(expert_action(p))
obs = torch.stack(obs); act = torch.stack(act)
policy = BehaviorCloningPolicy()
opt = torch.optim.Adam(policy.parameters(), lr=3e-3)
for epoch in range(300):
loss = nn.functional.mse_loss(policy(obs), act)
opt.zero_grad(); loss.backward(); opt.step()
# now roll out from FAR away and watch the drift
def rollout(start, steps=15):
pos = torch.tensor([start])
traj = [start]
for _ in range(steps):
with torch.no_grad():
pos = pos + policy(pos.unsqueeze(0)).squeeze(0)
traj.append(round(pos.item(), 3))
return traj
print("from 0.2 :", rollout(0.2))
print("from 5.0 :", rollout(5.0))
Started near the origin, the policy walks calmly to 0. Started at 5.0 -- a position it NEVER saw in training -- it makes a slightly wrong first move, lands somewhere even stranger, makes an even worse move, and the errors snowball. The one-sentence explanation: the further you start from the training distribution, the more off-distribution each state the policy visits becomes, so its errors compound instead of correcting -- and the fix, named in #141, is DAgger, which folds the states the policy actually visits back into the training set with expert labels.
Exercise 2 -- Randomize until it transfers. Extend the DomainRandomizer with a per-episode sampler, draw 1000 configs, and bucket-count the friction values to confirm they cover the full range.
import numpy as np
class DomainRandomizer:
def __init__(self):
self.ranges = {"friction": (0.3, 1.5), "mass_scale": (0.7, 1.3)}
def sample(self):
return {k: np.random.uniform(lo, hi) for k, (lo, hi) in self.ranges.items()}
def randomize_episode(self, env_params):
"""Return a FRESH config each episode, merged onto the base params."""
cfg = dict(env_params)
cfg.update(self.sample())
return cfg
rng = DomainRandomizer()
base = {"gravity": -9.81}
frictions = [rng.randomize_episode(base)["friction"] for _ in range(1000)]
# bucket-count into 4 bins across the declared range
bins = np.linspace(0.3, 1.5, 5)
counts = np.histogram(frictions, bins=bins)[0]
for i, c in enumerate(counts):
print(f"[{bins[i]:.2f}, {bins[i+1]:.2f}): {c:4d} " + "#" * (c // 20))
Roughly 250 samples land in each of the four buckets -- a flat spread across the whole (0.3, 1.5) band, which is exactly what "cover the range" means. And the argument the exercise asked for: if you randomize friction over (0.9, 1.0) but the real floor has friction 0.4, then reality sits OUTSIDE the training distribution entirely, and every guarantee evaporates. Domain randomization only converts sim-to-real into a generalization problem when reality falls inside your ranges. Randomize too narrowly and you are right back to a brittle policy that has never met the world it will actually stand on.
Exercise 3 -- Give the VLA a memory. Modify the VisionLanguageActionPolicy to swallow a short HISTORY of vision frames instead of a single one, keep the instruction as a single vector, and confirm the output is still one action.
class HistoryVLAPolicy(nn.Module):
"""Vision history + one instruction -> one action vector."""
def __init__(self, vision_dim=512, language_dim=512, action_dim=7, k=4):
super().__init__()
self.k = k
self.vision_encoder = nn.Sequential(nn.Linear(vision_dim * k, 256), nn.ReLU())
self.language_encoder = nn.Sequential(nn.Linear(language_dim, 256), nn.ReLU())
self.policy_head = nn.Sequential(
nn.Linear(512, 256), nn.ReLU(),
nn.Linear(256, action_dim), nn.Tanh(),
)
def forward(self, vision_window, language):
# vision_window: (batch, k, vision_dim) -> flatten the time axis
v = self.vision_encoder(vision_window.reshape(vision_window.shape[0], -1))
l = self.language_encoder(language)
return self.policy_head(torch.cat([v, l], dim=-1))
vla = HistoryVLAPolicy()
frames = torch.randn(1, 4, 512) # last 4 encoded frames
instr = torch.randn(1, 512) # "catch the rolling can"
print("motor command ->", vla(frames, instr).shape) # (1, 7)
Still (1, 7) -- one action, regardless of how many frames feed in. And a manipulation task that is IMPOSSIBLE from a single frame but solvable with a 4-frame window: catching a can rolling across the table. One photograph cannot tell you which way the can is moving or how fast; four frames reveal the velocity, and velocity is the whole game. Right -- homework settled. Now we teach the machine to think slowly ;-)
System 1 vs System 2: the framing that explains everything
Daniel Kahneman's dual-process theory splits human thinking into two modes, and it is the single most useful lens I know for what follows.
System 1 is fast, automatic, intuitive. You recognise a face. You read a word. You catch a ball. No deliberation -- just an immediate response bubbling up.
System 2 is slow, effortful, deliberate. You multiply 17 by 24. You plan a route through an unfamiliar city. You check whether an argument is actually valid or merely sounds nice.
Here is the uncomfortable truth: every standard neural network we've built, transformers included (#52-53), is a System 1 machine. A forward pass through a transformer is a FIXED-depth computation. The model spends exactly the same amount of compute on every input, no matter how hard it is. "What's 2+2?" and "prove that the square root of 2 is irrational" get the identical number of floating-point operations. That is a deep architectural limitation, not a bug you can patch with more data. Human reasoning allocates VARIABLE effort to variable difficulty -- you slow down on the hard part -- and a fixed forward pass simply cannot.
Every technique in this episode is, in its own way, an attempt to bolt something like System 2 onto a System 1 substrate. Keep that frame in your head and the rest falls into place.
Chain-of-thought: thinking out loud
The simplest and most jaw-droppingly effective reasoning trick is almost embarrassing: just make the model show its work.
Chain-of-thought (CoT) prompting was crystallised by Wei et al. in 2022. Instead of asking for a direct answer, you ask the model to reason step by step, and the intermediate tokens act as a scratchpad -- each reasoning step conditioning the next. Why on earth does that help? Because autoregressive generation (#58) is fundamentally NOT a single forward pass. Every token the model generates gets to attend to all the tokens before it. A ten-step chain of thought therefore gives the model roughly ten times more effective compute on the problem than a blurted one-token answer. The intermediate tokens are not decoration -- they are doing computation. This is the crux: CoT lets a fixed-depth network fake variable depth by unrolling its thinking across the sequence dimension in stead of the layer dimension.
We can build a tiny sketch of that "variable depth" idea directly. Instead of unrolling in text, we unroll a hidden state through a repeated reasoning block and -- crucially -- let the model decide WHEN to stop.
import torch
import torch.nn as nn
class ReasoningStep(nn.Module):
"""One reasoning step: a residual refinement of a hidden state."""
def __init__(self, hidden_dim=256):
super().__init__()
self.reason = nn.Sequential(
nn.Linear(hidden_dim, hidden_dim), nn.GELU(),
nn.Linear(hidden_dim, hidden_dim),
)
self.norm = nn.LayerNorm(hidden_dim)
def forward(self, state):
return self.norm(state + self.reason(state))
class IterativeReasoner(nn.Module):
"""Apply a VARIABLE number of reasoning steps, decided at runtime."""
def __init__(self, input_dim=64, hidden_dim=256, n_classes=10, max_steps=8):
super().__init__()
self.encoder = nn.Linear(input_dim, hidden_dim)
self.step = ReasoningStep(hidden_dim)
self.classifier = nn.Linear(hidden_dim, n_classes)
self.halt_predictor = nn.Linear(hidden_dim, 1)
self.max_steps = max_steps
def forward(self, x):
state = self.encoder(x)
cumulative_halt = torch.zeros(x.size(0), 1)
output = torch.zeros(x.size(0), self.classifier.out_features)
for t in range(self.max_steps):
state = self.step(state)
halt = torch.sigmoid(self.halt_predictor(state)) # "am I done?"
output = output + halt * self.classifier(state) # weighted vote
cumulative_halt = cumulative_halt + halt
if (cumulative_halt > 0.99).all():
break # stop thinking early
return output
model = IterativeReasoner()
print("output ->", model(torch.randn(3, 64)).shape) # (3, 10)
That halt_predictor is the whole idea of adaptive computation (Graves, 2016): simple inputs halt after a step or two, hard inputs keep grinding. It is a crude mechanical cousin of "spend longer on the harder question", and it captures why CoT works -- you are letting the model buy more thinking when the problem is worth it.
Tree-of-thought: exploring more than one path
Chain-of-thought is a SINGLE path through reasoning space. But what if the first step wanders off in the wrong direction? You're now stuck on a bad road with no way to turn around. Humans do not do this -- we consider a couple of approaches, sense which one smells promising, and abandon the dead ends.
Tree-of-thought (ToT) gives the model that ability: explore several reasoning paths at once and keep the best. If that sounds familiar, it should -- it is beam search (#71) applied to REASONING steps rather than to individual tokens.
class TreeOfThought:
"""Explore multiple reasoning paths, keep the most promising."""
def __init__(self, propose_fn, evaluate_fn, n_branches=3, max_depth=5):
self.propose = propose_fn # generate candidate next steps
self.evaluate = evaluate_fn # score a partial reasoning path
self.n_branches = n_branches
self.max_depth = max_depth
def search(self, problem, beam_width=3):
frontier = [(problem, 0.0)] # (reasoning_so_far, score)
for _ in range(self.max_depth):
candidates = []
for state, _ in frontier:
for step in self.propose(state, self.n_branches):
new_state = state + "\n" + step
candidates.append((new_state, self.evaluate(new_state)))
candidates.sort(key=lambda x: x[1], reverse=True)
frontier = candidates[:beam_width] # prune to the top-k
return frontier[0] # best complete path
The evaluate_fn is the beating heart of this. It scores partial reasoning so you can kill bad branches early, and it can be a separate "verifier" model, the same model prompted to critique itself, or even a dumb heuristic. Better evaluator, more efficient search -- and this is exactly why the next section on rewarding reasoning steps matters so much. Having said that, ToT is not free: exploring a tree multiplies your inference cost, so you spend compute to buy reliability. There is no such thing as a free lunch, only lunches you pay for in GPU-seconds.
Planning with MCTS and A*
For STRUCTURED problems -- board games, route finding, task scheduling -- we don't have to invent reasoning from nothing. We have decades of classical planning algorithms, and the modern move is to steer them with learned neural heuristics.
Monte Carlo Tree Search (MCTS) is how AlphaGo beat the world champion (#112). It grows a search tree by repeating four steps: SELECT a promising node, EXPAND it, SIMULATE (or evaluate) the result, and BACKPROPAGATE that value up to the ancestors. The neural network supplies the prior policy (which moves are worth trying first) and a value estimate (how good is this position), so the search never wastes time on obviously silly branches.
import math
class MCTSNode:
"""A node in the Monte Carlo search tree."""
def __init__(self, state, parent=None, prior=0.0):
self.state = state
self.parent = parent
self.prior = prior
self.children = {}
self.visit_count = 0
self.value_sum = 0.0
@property
def value(self):
return 0.0 if self.visit_count == 0 else self.value_sum / self.visit_count
def ucb_score(self, c_puct=1.4):
"""Balance exploitation (value) against exploration (unvisited priors)."""
if self.visit_count == 0:
return float("inf")
explore = c_puct * self.prior * math.sqrt(self.parent.visit_count) / (1 + self.visit_count)
return self.value + explore
def mcts_search(root, policy_value_fn, apply_action, n_simulations=100):
for _ in range(n_simulations):
node = root
# 1. SELECT: descend by UCB until we reach a leaf
while node.children:
node = max(node.children.values(), key=lambda n: n.ucb_score())
# 2. EXPAND + EVALUATE: ask the network for priors and a value
policy, value = policy_value_fn(node.state)
for action, prob in policy.items():
node.children[action] = MCTSNode(apply_action(node.state, action), node, prob)
# 3. BACKPROPAGATE: push the value up every ancestor
while node is not None:
node.visit_count += 1
node.value_sum += value
node = node.parent
# robust choice: the MOST-VISITED child, not the highest-value one
return max(root.children.items(), key=lambda kv: kv[1].visit_count)
Notice the last line -- we return the most-VISITED child, not the highest-scoring one. That is deliberate: visit count is a far more robust signal than a raw value estimate, because a node only accrues visits if the search kept coming back to it. AlphaGo learned that lesson so we don't have to.
A* is the classic optimal shortest-path algorithm, and it plugs a heuristic into a priority queue: always expand the node with the smallest "cost so far plus estimated cost to go". The neural twist is to LEARN that estimated-cost-to-go instead of hand-crafting it. Here is A* with a heuristic you could swap for a network:
import heapq
def a_star(start, goal, neighbors, heuristic):
"""Classic A*; `heuristic(node, goal)` can be a learned network."""
frontier = [(0.0, start)] # (priority, node)
cost_so_far = {start: 0.0}
came_from = {start: None}
while frontier:
_, current = heapq.heappop(frontier)
if current == goal:
break
for nxt, step_cost in neighbors(current):
new_cost = cost_so_far[current] + step_cost
if nxt not in cost_so_far or new_cost < cost_so_far[nxt]:
cost_so_far[nxt] = new_cost
priority = new_cost + heuristic(nxt, goal) # f = g + h
heapq.heappush(frontier, (priority, nxt))
came_from[nxt] = current
return came_from, cost_so_far
The beauty here is that A* is provably OPTIMAL as long as the heuristic never overestimates (it is "admissible"). Marry that guarantee to a learned heuristic and you get the best of both worlds: classical rigour plus learned intuition about which direction the goal probably lies. That marriage -- hard guarantees from search, soft intuition from the network -- is the whole spirit of modern planning.
Process reward models: grading the working, not just the answer
Now the piece I find genuinely exciting, and the most important recent shift in how we train reasoning: process reward models (PRMs). Instead of rewarding only the final answer (outcome-based), you reward EACH reasoning step (process-based).
Why does that matter so much? Because outcome-only reward is horribly SPARSE -- the model learns "right" or "wrong" once, at the very end of a long chain, and has to guess which of its ten steps deserved the blame or the credit. That is a brutal credit-assignment problem (we first met it back in RL, #106). Process reward is DENSE -- every intermediate step gets its own feedback -- so the model learns precisely which moves lead somewhere good.
class ProcessRewardModel(nn.Module):
"""Score each reasoning STEP, not just the final answer."""
def __init__(self, hidden_dim=256):
super().__init__()
self.step_encoder = nn.TransformerEncoder(
nn.TransformerEncoderLayer(d_model=hidden_dim, nhead=8, batch_first=True),
num_layers=4,
)
self.step_scorer = nn.Linear(hidden_dim, 1)
def forward(self, step_embeddings):
# step_embeddings: (batch, n_steps, hidden_dim)
encoded = self.step_encoder(step_embeddings) # each step sees the earlier ones
return torch.sigmoid(self.step_scorer(encoded).squeeze(-1)) # (batch, n_steps)
def training_loss(self, step_embeddings, step_labels):
"""step_labels: 1.0 for a correct step, 0.0 for a wrong one."""
return nn.functional.binary_cross_entropy(self.forward(step_embeddings), step_labels)
prm = ProcessRewardModel()
steps = torch.randn(1, 6, 256) # a 6-step chain of thought
print("per-step scores ->", prm(steps).shape) # (1, 6)
OpenAI's work on PRMs (the "Let's Verify Step by Step" line of research) showed that models trained with process supervision clearly OUTPERFORM outcome-supervised ones on hard math -- and, tellingly, the gap WIDENS as the problems get harder. The practical catch is nasty: you need step-level labels, which are expensive to collect from humans. So a lot of current research is about generating those labels automatically -- using one model to judge another model's steps, or inferring step quality from whether the paths that pass through a step tend to reach correct answers. Nota bene: this is also exactly the evaluate_fn that made tree-of-thought work earlier. A good process reward model IS a good reasoning-path evaluator. The pieces of this episode are more connected than they first appear.
The reasoning gap: an honest assessment
Let me be blunt, because the hype around this topic is thick enough to spread on toast. Current AI systems -- the very best LLMs included -- do NOT reason the way you do. They are extraordinary pattern matchers that can simulate reasoning over short chains, especially on problems that rhyme with their training data. But they still fumble genuinely novel logic puzzles, struggle to plan many steps ahead under hard constraints, and cannot reliably tell a valid argument from a plausible-sounding invalid one.
Chain-of-thought helps, but it does not conjure reasoning out of thin air -- it creates a FORMAT in which the model's pattern matching operates more effectively. Tree-of-thought helps more, at the cost of multiplying compute, and still guarantees nothing about correctness. Process reward models are, to my eye, the most promising direction of the three, precisely because they create a training signal for the reasoning PROCESS itself rather than only its output.
The honest answer -- and I would distrust anyone who tells you otherwise with a straight face -- is that we do not yet know whether scaling these approaches produces genuine reasoning or an ever-more-convincing imitation of it. That is one of the most important open questions in the whole field, and it is very much unsettled in 2026. I would rather hand you that honest uncertainty than a comforting story.
Exercises
Get your hands dirty before next time. Three tasks, climbing in difficulty:
Watch adaptive computation actually adapt. Take the
IterativeReasonerand log how many steps it runs beforecumulative_haltcrosses 0.99 for each input in a batch (add a step counter to the loop). Feed it two batches -- one oftorch.zeros-ish "easy" inputs and one of large-magnitude "hard" inputs -- and report the average step counts. Then, in a sentence, explain why an UNTRAINED halt predictor gives you a noisy, meaningless answer, and what you would need to train it against.Give tree-of-thought a real evaluator. Instantiate
TreeOfThoughtwith a toypropose_fnthat appends random digits to a string and anevaluate_fnthat scores a path by how close the digits sum to a target you pick. Run the search and confirm the winning path's digits sum near your target. Then explain, in one sentence, how the search quality changes if you replace your hand-written evaluator with a random one.Prove A stays optimal -- then break it.* Build a small 5-by-5 grid with a few blocked cells, run
a_starwith the Manhattan-distance heuristic, and confirm it finds the shortest path. Now multiply the heuristic by 5 (making it INADMISSIBLE, i.e. it overestimates), rerun, and show it can return a longer path. Write one sentence connecting "admissible heuristic" to "optimality guarantee".
We'll open next episode with full solutions, as always.
Quick recap
- System 1 vs System 2: a plain transformer is a fixed-depth System 1 machine that spends equal compute on every input; reasoning techniques are all attempts to graft on System 2 deliberation;
- chain-of-thought turns intermediate tokens into a scratchpad, letting a fixed-depth network buy variable effective depth along the sequence -- with adaptive computation as the mechanical version of "think longer on the hard bit";
- tree-of-thought is beam search over reasoning steps, keeping several paths alive and pruning with an evaluator, at the price of multiplied compute;
- MCTS (select, expand, evaluate, backpropagate) and A* (learned heuristic in a priority queue) fuse classical search guarantees with learned neural intuition -- the AlphaGo recipe;
- process reward models score each reasoning step instead of the final answer only, turning a sparse credit-assignment nightmare into dense feedback -- and doubling as the evaluator ToT needs;
- the reasoning gap is real: today's systems simulate reasoning through pattern matching and do not yet exhibit robust, generalizable logic -- and whether scaling closes that gap is genuinely unknown.
And here is the thread I want to leave dangling this time. Everything today assumed a model that is already TRAINED and now merely thinks harder at inference time. But a human who reasons also keeps LEARNING -- you pick up a new fact today and it does not erase what you knew yesterday. Our models don't do that. Freeze the weights and they are frozen; retrain them on something new and they tend to forget the old. What would it take for a model to keep learning after deployment, without quietly overwriting everything it already knew? That is where we head next ;-)