1178 lines
56 KiB
Markdown
1178 lines
56 KiB
Markdown
# Workprocess Documentation for LLMMCTS
|
||
|
||
## Overview
|
||
|
||
LLMMCTS implements Monte Carlo Tree Search (MCTS) for Large Language Model (LLM) planning tasks. It combines LLM reasoning with MCTS search to solve complex planning problems with sparse rewards.
|
||
|
||
## Core Concept
|
||
|
||
The package addresses the sparse reward problem in MCTS by using LLMs to provide **pseudo-rewards** (called `progressvalue`) at every node, enabling faster learning without waiting for terminal rewards.
|
||
|
||
## Architecture
|
||
|
||
### Three-Tier Value System
|
||
|
||
| Field | Source | Purpose |
|
||
|-------|--------|---------|
|
||
| `progressvalue` | LLM heuristic | Estimate of how close we are to solving; used for fast node selection |
|
||
| `statevalue` | Monte Carlo simulation | Actual cumulative reward from simulations; accurate but expensive to compute |
|
||
| `reward` | Environment | Immediate reward from environment (may be sparse, only at terminal states) |
|
||
|
||
### How reward and progressvalue Work Together
|
||
|
||
The package uses two distinct value signals that serve different purposes during MCTS search:
|
||
|
||
#### `reward` - Environment Feedback
|
||
- **Source**: Environment only
|
||
- **When**: Only at terminal states (sparse) or immediate action outcome
|
||
- **Purpose**: Ground truth for backpropagation updates
|
||
- **Usage**:
|
||
- Accumulated during simulation rollouts (line 287 in mcts.jl)
|
||
- Used to update `statevalue` via backpropagation
|
||
- Final metric for determining solution quality
|
||
|
||
#### `progressvalue` - LLM Heuristic
|
||
- **Source**: LLM's reasoning about state quality
|
||
- **When**: At every node expansion (dense guidance)
|
||
- **Purpose**: Fast node selection without waiting for terminal rewards
|
||
- **Usage**:
|
||
- Passed to `_expand()` from `transition()` result (line 237 in mcts.jl)
|
||
- Used as fallback in UCT when child has zero visits (util.jl:58)
|
||
- Combined with `reward` in `selectChildNode()` (line 171 in mcts.jl)
|
||
|
||
#### The Selection Hierarchy
|
||
|
||
```
|
||
1. UCT SELECTION (interface.jl → runMCTS → UCTselect)
|
||
└── Uses: statevalue + exploration_weight * sqrt(ln(parent_visits) / child_visits)
|
||
└── Purpose: Balance exploration vs exploitation during tree traversal
|
||
|
||
2. CHILD SELECTION (mcts.jl → selectChildNode)
|
||
└── Uses: progressvalue + reward
|
||
└── Purpose: Pick best child after expansion during simulation rollout
|
||
|
||
3. BEST NODE SELECTION (mcts.jl → selectBestNextNode)
|
||
└── If statevalue sum > 0: statevalue / visits
|
||
└── Else: progressvalue + reward
|
||
└── Purpose: Determine final best trajectory after MCTS completes
|
||
```
|
||
|
||
#### Data Flow Example
|
||
|
||
```
|
||
Expansion Phase (LLM call):
|
||
┌────────────────────────────────────────────────────────────────────────┐
|
||
│ transition(state, args) → { │
|
||
│ newNodeKey: "abc-123", │
|
||
│ newstate: { reward: 0, isterminal: false }, │
|
||
│ progressvalue: 7.5 ← LLM estimates this state is promising │
|
||
│ } │
|
||
└────────────────────────────────────────────────────────────────────────┘
|
||
↓
|
||
New Node Created:
|
||
MCTSNode(
|
||
nodekey = "abc-123",
|
||
state = { reward: 0, isterminal: false },
|
||
visits = 0,
|
||
progressvalue = 7.5, ← LLM heuristic (fast, initial guidance)
|
||
statevalue = 0, ← Will be updated after simulation
|
||
reward = 0, ← Immediate environment reward
|
||
...
|
||
)
|
||
|
||
Simulation Phase:
|
||
┌────────────────────────────────────────────────────────────────────────┐
|
||
│ simulate() rolls out 3 levels: │
|
||
│ Level 1: reward += 0 (non-terminal) │
|
||
│ Level 2: reward += 3 (intermediate reward) │
|
||
│ Level 3: reward += 10 (terminal state) │
|
||
│ ───────────────────────────────────────────────────────────────────── │
|
||
│ Total simTrajectoryReward = 13 │
|
||
└────────────────────────────────────────────────────────────────────────┘
|
||
↓
|
||
Backpropagation:
|
||
┌────────────────────────────────────────────────────────────────────────┐
|
||
│ Update all ancestors with simTrajectoryReward = 13: │
|
||
│ node.visits += 1 │
|
||
│ node.statevalue = (old_statevalue * (visits-1) + 13) / visits │
|
||
│ reward *= 0.9 (discount for future rewards) │
|
||
└────────────────────────────────────────────────────────────────────────┘
|
||
```
|
||
|
||
#### Why Two Value Systems?
|
||
|
||
| Problem | Traditional MCTS | LLMMCTS Solution |
|
||
|---------|-----------------|------------------|
|
||
| **Sparse rewards** | Must explore blindly until terminal | LLM provides `progressvalue` at every node |
|
||
| **Slow learning** | Need many terminal outcomes | Dense `progressvalue` enables fast guidance |
|
||
| **Exploration cost** | Random expansion is inefficient | LLM filters to promising candidates |
|
||
| **Value uncertainty** | Hand-designed heuristics may be wrong | `statevalue` (from simulation) confirms LLM estimates |
|
||
|
||
### Module Structure
|
||
|
||
```
|
||
src/
|
||
├── type.jl # MCTSNode struct definition (MCTSNode type)
|
||
├── util.jl # UCT selection utility function (UCTselect)
|
||
├── mcts.jl # Core MCTS operations (select, expand, simulate, backpropagate)
|
||
├── interface.jl # High-level interface (runMCTS, simulateThenBackpropagate)
|
||
└── LLMMCTS.jl # Main package entry point (module exports)
|
||
```
|
||
|
||
## Data Flow
|
||
|
||
### 1. Node Structure (type.jl)
|
||
|
||
```julia
|
||
MCTSNode(
|
||
nodekey::String, # Unique identifier
|
||
state::Dict, # Current state dictionary
|
||
visits::Integer, # Number of visits to this node
|
||
progressvalue::Number, # LLM's estimate of state quality
|
||
statevalue::Number, # Average cumulative reward from simulations
|
||
reward::Number, # Immediate reward at this node
|
||
isterminal::Bool, # Whether this is a terminal state
|
||
parent::Union{MCTSNode, Nothing}, # Parent reference (nothing for root)
|
||
children::Dict{String,MCTSNode}, # Child nodes mapping
|
||
etc::Dict{String,Any} # Additional data storage
|
||
)
|
||
```
|
||
|
||
### 2. MCTS Search Process: Rollout and Exploration
|
||
|
||
LLMMCTS explores the solution space through iterative MCTS cycles, each consisting of four phases: **Selection**, **Expansion**, **Simulation**, and **Backpropagation**.
|
||
|
||
#### The Four Phases of MCTS
|
||
|
||
```
|
||
Phase 1: SELECTION
|
||
└── Start from root node
|
||
└── Use UCT (Upper Confidence Bound for Trees) to select child nodes
|
||
├── UCT formula: Q(s,a) + c * sqrt(ln(N(s)) / N(s,a))
|
||
├── Exploitation term (Q): Uses progressvalue or statevalue
|
||
└── Exploration term: Prefers less-visited nodes
|
||
└── Continue until reaching a leaf node
|
||
|
||
Phase 2: EXPANSION
|
||
└── If leaf is not terminal:
|
||
├── Apply transition function multiple times (horizontal sampling)
|
||
├── Generate new child nodes with LLM assistance
|
||
├── Each child represents a potential next state
|
||
└── Dejavu detection prevents duplicate states
|
||
|
||
Phase 3: SIMULATION (Rollout)
|
||
└── From expanded leaf node:
|
||
├── Perform rollout up to maxSimulationDepth
|
||
├── At each level:
|
||
│ ├── Accumulate reward from current node
|
||
│ ├── Check if terminal state reached
|
||
│ ├── Expand node to generate children
|
||
│ └── Select best child using progressvalue + reward
|
||
└── Return cumulative trajectory reward and terminal state
|
||
|
||
Phase 4: BACKPROPAGATION
|
||
└── Update statistics along the path to root:
|
||
├── Increment visit count for each node
|
||
├── Update statevalue (running average of rewards)
|
||
├── Apply discount to future rewards
|
||
└── Propagate reward upward
|
||
```
|
||
|
||
#### Detailed Expansion Process (mcts.jl → expand/_expand)
|
||
|
||
**Horizontal Sampling**: At each node, the algorithm generates multiple child nodes by applying the transition function several times.
|
||
|
||
```
|
||
Current Node (state S)
|
||
├── transition(S, args) → Child 1 (state S1, progressvalue=P1)
|
||
├── transition(S, args) → Child 2 (state S2, progressvalue=P2)
|
||
├── transition(S, args) → Child 3 (state S3, progressvalue=P3)
|
||
└── ... (horizontalSample times)
|
||
|
||
Each child is evaluated by:
|
||
- progressvalue: LLM's estimate of state quality (fast, heuristic)
|
||
- reward: Immediate environment reward
|
||
```
|
||
|
||
**Dejavu Detection**: Before adding a child node, the algorithm checks if the `newNodeKey` already exists in the current node's children. If it does, the node is skipped to avoid duplicates. This handles semantically equivalent states.
|
||
|
||
#### Detailed Simulation Process (mcts.jl → simulate)
|
||
|
||
The simulation phase performs a rollout from a given node to explore the solution space vertically:
|
||
|
||
```
|
||
Input: Starting node, transition function, maxSimulationDepth
|
||
|
||
simTrajectoryReward = 0
|
||
terminalstate = nothing
|
||
|
||
for depth in 1:maxSimulationDepth:
|
||
# Step 1: Accumulate current node's reward
|
||
simTrajectoryReward += node.reward
|
||
|
||
# Step 2: Check if terminal state
|
||
if node.isterminal:
|
||
terminalstate = node.state
|
||
break
|
||
|
||
# Step 3: Expand node horizontally (generate children)
|
||
expand(node, transition, transitionargs;
|
||
horizontalSample=horizontalSample,
|
||
multithread=multithread)
|
||
|
||
# Step 4: Select best child for next step
|
||
node = selectChildNode(node) # Uses progressvalue + reward
|
||
|
||
Return: (simTrajectoryReward, terminalstate)
|
||
```
|
||
|
||
**Why this matters**: The simulation phase estimates the value of a node by looking ahead `maxSimulationDepth` levels. Each rollout provides an estimate of the cumulative reward that can be obtained from the current state.
|
||
|
||
#### Detailed Backpropagation Process (mcts.jl → backpropagate)
|
||
|
||
After simulation, rewards are propagated back up the tree to update all visited nodes:
|
||
|
||
```
|
||
Input: Starting node (leaf from simulation), cumulative trajectory reward
|
||
|
||
while !isroot(node):
|
||
# Update visit count
|
||
node.visits += 1
|
||
|
||
# Update statevalue (running average)
|
||
# Formula: new_avg = (old_avg * (n-1) + new_value) / n
|
||
node.statevalue = (node.statevalue * (node.visits-1) + simTrajectoryReward) / node.visits
|
||
|
||
# Discount reward for future states
|
||
# Future rewards are less certain, so they receive lower weight
|
||
simTrajectoryReward *= discountRewardCoeff
|
||
|
||
# Move to parent
|
||
node = node.parent
|
||
```
|
||
|
||
**Discounting rationale**: Rewards further from the current state are discounted because:
|
||
1. Future rewards are uncertain
|
||
2. The longer the horizon, the more opportunities for suboptimal decisions
|
||
3. This creates a more realistic value estimate
|
||
|
||
#### Selection Phase: UCT (util.jl → UCTselect)
|
||
|
||
The UCT formula balances exploration vs. exploitation:
|
||
|
||
```
|
||
UCT(s,a) = Q(s,a) + c * sqrt(ln(N(s)) / N(s,a))
|
||
|
||
Where:
|
||
- Q(s,a) = childNode.statevalue (exploitation: current estimate)
|
||
- c = explorationweight (controls exploration vs exploitation)
|
||
- N(s) = parent.visits (total visits to parent node)
|
||
- N(s,a) = childNode.visits (visits to this specific child)
|
||
```
|
||
|
||
**Behavior**:
|
||
- If `childNode.visits = 0`: Exploration term becomes undefined, so use `progressvalue` as fallback
|
||
- If `childNode.visits` is low: Exploration term is high → encourages visiting unexplored branches
|
||
- If `childNode.visits` is high: Exploration term approaches 0 → exploits known high-value nodes
|
||
|
||
**Selection priority**:
|
||
1. **High `progressvalue` + low visits**: Explored first (fast LLM guidance)
|
||
2. **High `statevalue` + high visits**: Exploited once confirmed (accurate value)
|
||
3. **Balance**: Controlled by `explorationweight` parameter
|
||
|
||
### 3. Main Workflow (interface.jl → runMCTS)
|
||
|
||
```
|
||
runMCTS(initialstate, transition, args)
|
||
├── Initialize root node with initialstate
|
||
│ └── root.visits = 0, root.statevalue = 0, root.children = {}
|
||
│
|
||
├── For iteration 1 to maxiterations:
|
||
│ │
|
||
│ ├── PHASE 1: SELECTION
|
||
│ │ ├── Start at root node
|
||
│ │ ├── While not leaf node:
|
||
│ │ │ ├── Apply UCTselect() with explorationweight
|
||
│ │ │ │ ├── UCT = statevalue + w * sqrt(ln(parent_visits) / child_visits)
|
||
│ │ │ │ └── Select child with highest UCT score
|
||
│ │ │ └── node = selected child
|
||
│ │ └── node is now a leaf node
|
||
│ │
|
||
│ ├── PHASE 2: TERMINAL CHECK
|
||
│ │ ├── If node.isterminal == true:
|
||
│ │ │ ├── If node.state[:reward] >= 8:
|
||
│ │ │ │ └── Store in highValueState channel
|
||
│ │ │ └── backpropagate(node, node.reward)
|
||
│ │ │ └── Update all ancestors up to root
|
||
│ │ └── Else (non-terminal leaf):
|
||
│ │ └── Continue to EXPANSION
|
||
│ │
|
||
│ ├── PHASE 3: EXPANSION
|
||
│ │ ├── expand(node, transition, transitionargs;
|
||
│ │ │ horizontalSample=horizontalSampleExpansionPhase)
|
||
│ │ │
|
||
│ │ ├── For each child generated:
|
||
│ │ │ ├── newNode = transition(current_state, args)
|
||
│ │ │ ├── newNodeKey = result[:newNodeKey]
|
||
│ │ │ ├── newstate = result[:newstate]
|
||
│ │ │ ├── progressvalue = result[:progressvalue]
|
||
│ │ │ └── Dejavu check: skip if newNodeKey already exists
|
||
│ │ │
|
||
│ │ └── Each child gets initialized with:
|
||
│ │ ├── visits = 0
|
||
│ │ ├── statevalue = 0 (no simulations yet)
|
||
│ │ ├── progressvalue = LLM estimate
|
||
│ │ ├── reward = newstate[:reward]
|
||
│ │ └── parent = current node
|
||
│ │
|
||
│ ├── PHASE 4: SIMULATION + BACKPROPAGATION
|
||
│ │ ├── If multithread == true:
|
||
│ │ │ └── Spawn parallel simulateThenBackpropagate() for each child
|
||
│ │ │
|
||
│ │ └── For each leafNode in node.children:
|
||
│ │ ├── simulateThenBackpropagate(leafNode, ...)
|
||
│ │ │ ├── simulate(leafNode, ...):
|
||
│ │ │ │ └── Rollout up to maxSimulationDepth
|
||
│ │ │ │ └── Returns (simTrajectoryReward, terminalstate)
|
||
│ │ │ │
|
||
│ │ │ ├── If terminalstate[:reward] >= 8:
|
||
│ │ │ │ └── Store in highValueState channel
|
||
│ │ │ │
|
||
│ │ │ └── backpropagate(leafNode, simTrajectoryReward):
|
||
│ │ │ └── Update visits and statevalue for all ancestors
|
||
│ │ │
|
||
│ │ └── If saveSimulatedNode == false:
|
||
│ │ └── Clear children (free memory for next iteration)
|
||
│ │
|
||
│ └── PHASE 5: EARLY STOP CHECK
|
||
│ ├── If earlystop(node.state) == true:
|
||
│ │ └── Break out of iteration loop
|
||
│ └── Continue to next iteration
|
||
│
|
||
└── After all iterations, select best result:
|
||
├── bestNextState = selectBestNextNode(root)
|
||
│ └── Uses statevalue/visits or progressvalue + reward
|
||
│
|
||
├── bestTerminalState = selectBestTrajectoryNode(root)
|
||
│ └── Follows optimal trajectory to leaf
|
||
│
|
||
└── highValueStateList = collect from highValueState channel
|
||
```
|
||
|
||
### 7. Selection Phase Details (util.jl → UCTselect)
|
||
|
||
**Purpose**: Select the best child node using UCT formula that balances exploration and exploitation.
|
||
|
||
**Process**:
|
||
|
||
```
|
||
UCTselect(node, w)
|
||
|
||
Input: Parent node, exploration weight w
|
||
Output: Child node with highest UCT score
|
||
|
||
maxUCT = -Inf
|
||
selectedNode = nothing
|
||
|
||
for each childNode in node.children:
|
||
┌── Calculate UCT value
|
||
│ └── if childNode.visits != 0:
|
||
│ ├── weightedterm = w * sqrt(ln(node.visits) / childNode.visits)
|
||
│ ├── UCTvalue = childNode.statevalue + weightedterm
|
||
│ └── Exploration term encourages low-visited children
|
||
│
|
||
│ else: # childNode.visits == 0
|
||
│ ├── UCTvalue = childNode.progressvalue
|
||
│ └── No exploration term (division by zero)
|
||
│
|
||
├── Compare with current max
|
||
│ └── if UCTvalue > maxUCT:
|
||
│ ├── maxUCT = UCTvalue
|
||
│ └── selectedNode = childNode
|
||
|
||
return selectedNode
|
||
```
|
||
|
||
**Behavior analysis**:
|
||
|
||
| Scenario | UCT Value | Behavior |
|
||
|----------|-----------|----------|
|
||
| Child never visited (`visits=0`) | `progressvalue` | Strong exploration → will be tried |
|
||
| Child visited often, high reward | High `statevalue`, low exploration | Exploitation dominates |
|
||
| Child visited often, low reward | Low `statevalue`, low exploration | Exploitation avoids this node |
|
||
| Child visited rarely, moderate reward | Moderate `statevalue`, high exploration | May explore further |
|
||
|
||
### 8. Expansion Phase Details (mcts.jl → expand/_expand)
|
||
|
||
**Purpose**: Generate new child nodes by applying the transition function multiple times.
|
||
|
||
**Process**:
|
||
|
||
```
|
||
expand(node, transition, transitionargs;
|
||
horizontalSample=3, multithread=false)
|
||
|
||
Input: Node to expand, transition function, arguments
|
||
Output: node.children populated with child nodes
|
||
|
||
if multithread == true:
|
||
└── @sync for i in 1 to horizontalSample:
|
||
└── @spawn _expand(node, transition, transitionargs)
|
||
└── Run expansion in parallel threads
|
||
else:
|
||
└── for i in 1 to horizontalSample:
|
||
└── _expand(node, transition, transitionargs)
|
||
└── Sequential expansion
|
||
|
||
_expand(node, transition, transitionargs)
|
||
|
||
Input: Single node to expand
|
||
Output: One child node added to node.children
|
||
|
||
result = transition(node.state, transitionargs)
|
||
└── LLM generates next state:
|
||
└── response = llm_call(state[:thoughtHistory], args.prompt)
|
||
└── Parse response into new state structure
|
||
|
||
newNodeKey = result[:newNodeKey]
|
||
newstate = result[:newstate]
|
||
progressvalue = result[:progressvalue]
|
||
|
||
if newNodeKey ∉ keys(node.children):
|
||
┌── Create new MCTSNode
|
||
│ └── newNode = MCTSNode(
|
||
│ ├── nodekey = newNodeKey
|
||
│ ├── state = newstate
|
||
│ ├── visits = 0
|
||
│ ├── progressvalue = progressvalue
|
||
│ ├── statevalue = 0 (no simulations yet)
|
||
│ ├── reward = newstate[:reward]
|
||
│ ├── isterminal = newstate[:isterminal]
|
||
│ ├── parent = node
|
||
│ ├── children = {}
|
||
│ └── etc = {}
|
||
│
|
||
└── node.children[newNodeKey] = newNode
|
||
```
|
||
|
||
**Dejavu detection**: Before adding a child, check if `newNodeKey` already exists in `node.children`. If so, skip to avoid duplicate states.
|
||
|
||
**Example**: Expanding a node with `horizontalSample=3`
|
||
```
|
||
Current Node: "Math problem: Solve x^2 = 16"
|
||
└── Attempt 1 (LLM):
|
||
└── Thought: "Take square root of both sides"
|
||
└── Action: "x = sqrt(16)"
|
||
└── New state: "Solved: x = 4"
|
||
└── newNodeKey = "abc-123"
|
||
└── Added to children: "abc-123" → Node
|
||
|
||
└── Attempt 2 (LLM):
|
||
└── Thought: "Consider negative root"
|
||
└── Action: "x = -sqrt(16)"
|
||
└── New state: "Solved: x = -4"
|
||
└── newNodeKey = "def-456"
|
||
└── Added to children: "def-456" → Node
|
||
|
||
└── Attempt 3 (LLM):
|
||
└── Thought: "Check both solutions"
|
||
└── Action: "Verify x=4 and x=-4"
|
||
└── New state: "Verified: x=4, x=-4"
|
||
└── newNodeKey = "ghi-789"
|
||
└── Added to children: "ghi-789" → Node
|
||
|
||
Result: Node has 3 children with different solution approaches
|
||
```
|
||
|
||
### 9. Backpropagation Phase Details (mcts.jl → backpropagate)
|
||
|
||
**Purpose**: Update the statistics of all nodes along the simulation path with the observed reward.
|
||
|
||
**Process**:
|
||
|
||
```
|
||
backpropagate(node, simTrajectoryReward;
|
||
discountRewardCoeff=0.9)
|
||
|
||
Input: Starting node (leaf from simulation), cumulative reward
|
||
Output: Updates visits and statevalue for all ancestors
|
||
|
||
while !isroot(node):
|
||
┌── Update visit count
|
||
│ └── node.visits += 1
|
||
│
|
||
├── Update statevalue (running average formula)
|
||
│ │
|
||
│ ├── Current average: node.statevalue
|
||
│ ├── Current count: node.visits - 1
|
||
│ ├── New value to add: simTrajectoryReward
|
||
│ │
|
||
│ └── Formula: new_avg = (old_avg * (n-1) + new_value) / n
|
||
│ node.statevalue = (node.statevalue * (node.visits-1) + simTrajectoryReward) / node.visits
|
||
│
|
||
├── Apply discount to future reward
|
||
│ └── simTrajectoryReward *= discountRewardCoeff
|
||
│ └── Reward gets 10% smaller for each level up
|
||
│
|
||
└── Move to parent
|
||
└── node = node.parent
|
||
|
||
# After loop, root node still needs update
|
||
# (root node's parent is nothing, so loop stops but root is still updated)
|
||
```
|
||
|
||
**Example**: Backpropagating reward=20 with discount=0.9
|
||
```
|
||
Depth 3 (leaf): Node A2b2
|
||
└── visits = 1, statevalue = 20/1 = 20.0
|
||
└── reward for parent = 20 * 0.9 = 18.0
|
||
|
||
Depth 2: Node A2b
|
||
└── visits = 1, statevalue = (0*0 + 18)/1 = 18.0
|
||
└── reward for parent = 18 * 0.9 = 16.2
|
||
|
||
Depth 1: Node A2
|
||
└── visits = 1, statevalue = (0*0 + 16.2)/1 = 16.2
|
||
└── reward for parent = 16.2 * 0.9 = 14.58
|
||
|
||
Depth 0 (root): Node A
|
||
└── visits = 1, statevalue = (0*0 + 14.58)/1 = 14.58
|
||
```
|
||
|
||
**Why discounting matters**: Future rewards are discounted because:
|
||
1. **Uncertainty**: The longer the horizon, the more uncertain the outcome
|
||
2. **Temporal credit assignment**: Immediate rewards should have higher value
|
||
3. **Realistic evaluation**: A reward 10 steps away is worth less than immediate reward
|
||
|
||
### 10. Solution Space Exploration Strategy
|
||
|
||
LLMMCTS explores the solution space using a combination of **guided exploration** and **systematic exploitation**:
|
||
|
||
#### Exploration vs. Exploitation Trade-off
|
||
|
||
```
|
||
Iteration 1: Initial exploration (all nodes have visits=0)
|
||
├── UCTselect() uses progressvalue for all children
|
||
├── LLM provides guidance: which states seem promising?
|
||
├── High progressvalue nodes get visited first
|
||
└── Results stored in statevalue
|
||
|
||
Iterations 2-5: Early exploitation
|
||
├── Some nodes have higher statevalue confirmed by simulations
|
||
├── UCT balances:
|
||
│ ├── Exploitation: High statevalue nodes
|
||
│ └── Exploration: Nodes with low visits (high UCT exploration term)
|
||
└── Nodes with high progressvalue + high statevalue dominate
|
||
|
||
Iterations 6+: Refinement
|
||
├── Well-visited nodes have reliable statevalue estimates
|
||
├── Poor nodes (low statevalue) are explored less
|
||
├── Still explore unvisited branches (UCT exploration term)
|
||
└── Tree converges to optimal trajectory
|
||
```
|
||
|
||
#### Horizontal vs. Vertical Exploration
|
||
|
||
**Horizontal exploration** (expansion):
|
||
- **Parameter**: `horizontalSample`
|
||
- **Purpose**: Generate multiple candidate next states from current state
|
||
- **Method**: Apply transition function multiple times
|
||
- **LLM role**: Each call generates a different thought/action
|
||
|
||
**Vertical exploration** (simulation):
|
||
- **Parameter**: `maxSimulationDepth`
|
||
- **Purpose**: Evaluate the long-term value of a state
|
||
- **Method**: Rollout from current node to leaf
|
||
- **Reward accumulation**: Sum all rewards along trajectory
|
||
|
||
#### Dejavu Detection and State Pruning
|
||
|
||
**Dejavu detection** (mcts.jl → _expand):
|
||
|
||
```julia
|
||
if newNodeKey ∉ keys(node.children)
|
||
newNode = MCTSNode(...)
|
||
node.children[newNodeKey] = newNode
|
||
end
|
||
```
|
||
|
||
**Why it matters**:
|
||
- Prevents infinite loops in cyclic state spaces
|
||
- Handles semantically equivalent states as duplicates
|
||
- Reduces tree size and memory usage
|
||
- Improves search efficiency
|
||
|
||
**Example**:
|
||
```
|
||
Current Node: "Problem solving step 5"
|
||
├── Attempt 1: Generates "Next step: Apply formula A"
|
||
├── Attempt 2: Generates "Next step: Apply formula A" (same as attempt 1)
|
||
│ └── Dejavu detected: skip duplicate
|
||
└── Attempt 3: Generates "Next step: Apply formula B"
|
||
```
|
||
|
||
### 11. Complete MCTS Iteration Flow
|
||
|
||
```
|
||
┌─────────────────────────────────────────────────────────────────────────────┐
|
||
│ MCTS ITERATION 1 │
|
||
├─────────────────────────────────────────────────────────────────────────────┤
|
||
│ │
|
||
│ 1. SELECTION (UCT-based traversal): │
|
||
│ Root (visits=1) │
|
||
│ └── UCTselect() on root (has no children yet, so uses progressvalue) │
|
||
│ └── Select child with highest progressvalue │
|
||
│ └── Continue until reaching leaf node │
|
||
│ │
|
||
│ 2. EXPANSION (horizontal sampling): │
|
||
│ Leaf node │
|
||
│ └── expand(horizontalSample=3) │
|
||
│ └── Generate 3 child nodes via LLM transition │
|
||
│ └── Each child gets: progressvalue, reward, parent=leaf │
|
||
│ │
|
||
│ 3. SIMULATION (vertical rollout): │
|
||
│ Each child node │
|
||
│ └── simulate(maxSimulationDepth=3) │
|
||
│ └── Rollout 3 levels deep, accumulating rewards │
|
||
│ └── Return (simTrajectoryReward, terminalstate) │
|
||
│ │
|
||
│ 4. BACKPROPAGATION (update statistics): │
|
||
│ simTrajectoryReward │
|
||
│ └── backpropagate() up to root │
|
||
│ └── Update visits and statevalue for all ancestors │
|
||
│ └── Apply discount to future rewards │
|
||
│ │
|
||
└─────────────────────────────────────────────────────────────────────────────┘
|
||
|
||
┌─────────────────────────────────────────────────────────────────────────────┐
|
||
│ MCTS ITERATION 2 │
|
||
├─────────────────────────────────────────────────────────────────────────────┤
|
||
│ │
|
||
│ 1. SELECTION: │
|
||
│ Root (visits=2) │
|
||
│ └── UCTselect() now considers: │
|
||
│ ├── statevalue (from iteration 1) │
|
||
│ └── exploration term (low visits on unexpanded branches) │
|
||
│ └── May select different path than iteration 1 │
|
||
│ │
|
||
│ 2. EXPANSION: │
|
||
│ └── May expand different node or same node │
|
||
│ └── New children added to node.children │
|
||
│ │
|
||
│ 3. SIMULATION: │
|
||
│ └── Different rollout trajectory │
|
||
│ └── New reward estimate added to statistics │
|
||
│ │
|
||
│ 4. BACKPROPAGATION: │
|
||
│ └── Statistics updated with new information │
|
||
│ │
|
||
└─────────────────────────────────────────────────────────────────────────────┘
|
||
|
||
┌─────────────────────────────────────────────────────────────────────────────┐
|
||
│ MCTS ITERATION 3 to N │
|
||
├─────────────────────────────────────────────────────────────────────────────┤
|
||
│ │
|
||
│ Pattern repeats, with increasingly informed selection: │
|
||
│ │
|
||
│ • Nodes with high statevalue (confirmed by many simulations) │
|
||
│ → Exploited (selected frequently) │
|
||
│ │
|
||
│ • Nodes with low visits but promising progressvalue │
|
||
│ → Explored (UCT exploration term encourages tries) │
|
||
│ │
|
||
│ • Tree grows: more branches explored, more statistics accumulated │
|
||
│ │
|
||
│ • Best trajectory emerges from accumulated statistics │
|
||
│ │
|
||
└─────────────────────────────────────────────────────────────────────────────┘
|
||
```
|
||
|
||
```
|
||
simulate(node, transition, transitionargs;
|
||
maxSimulationDepth=3, horizontalSample=3, multithread=false)
|
||
|
||
Input: Starting node at depth 0
|
||
Output: (simTrajectoryReward, terminalstate)
|
||
|
||
simTrajectoryReward = 0.0
|
||
terminalstate = nothing
|
||
|
||
for depth in 1 to maxSimulationDepth:
|
||
┌── Step 1: Accumulate reward
|
||
│ └── simTrajectoryReward += node.reward
|
||
│
|
||
├── Step 2: Check for terminal state
|
||
│ └── if node.isterminal:
|
||
│ └── terminalstate = deepcopy(node.state)
|
||
│ break
|
||
│
|
||
├── Step 3: Expand node (horizontal sampling)
|
||
│ └── expand(node, transition, transitionargs;
|
||
│ horizontalSample=horizontalSample)
|
||
│ └── For i in 1 to horizontalSample:
|
||
│ └── _expand(node, transition, transitionargs)
|
||
│ ├── result = transition(node.state, args)
|
||
│ ├── newNodeKey = result[:newNodeKey]
|
||
│ ├── newstate = result[:newstate]
|
||
│ ├── progressvalue = result[:progressvalue]
|
||
│ └── if newNodeKey not in node.children:
|
||
│ └── Create new MCTSNode with:
|
||
│ ├── state = newstate
|
||
│ ├── progressvalue = progressvalue
|
||
│ ├── reward = newstate[:reward]
|
||
│ ├── parent = node
|
||
│ └── children = {}
|
||
│
|
||
└── Step 4: Select best child for next iteration
|
||
└── node = selectChildNode(node)
|
||
└── Uses: progressvalue + reward (no UCT here)
|
||
|
||
return (simTrajectoryReward, terminalstate)
|
||
```
|
||
|
||
**Example**: With `maxSimulationDepth=3` and `horizontalSample=3`:
|
||
```
|
||
Depth 0: Node A (reward=2)
|
||
└── Expand 3 children: A1, A2, A3
|
||
└── Select A2 (highest progressvalue + reward)
|
||
|
||
Depth 1: Node A2 (reward=3)
|
||
└── Expand 3 children: A2a, A2b, A2c
|
||
└── Select A2b (highest progressvalue + reward)
|
||
|
||
Depth 2: Node A2b (reward=5)
|
||
└── Expand 3 children: A2b1, A2b2, A2b3
|
||
└── Select A2b2 (highest progressvalue + reward)
|
||
|
||
Depth 3: Node A2b2 (reward=10, isterminal=true)
|
||
└── Terminal state reached!
|
||
|
||
Return: (simTrajectoryReward=20, terminalstate=A2b2.state)
|
||
= (2 + 3 + 5 + 10, ...)
|
||
```
|
||
|
||
### 6. Backpropagation (mcts.jl → backpropagate)
|
||
|
||
**Purpose**: Update the statistics of all nodes along the simulation path with the observed reward.
|
||
|
||
**Process**:
|
||
|
||
```
|
||
backpropagate(node, simTrajectoryReward;
|
||
discountRewardCoeff=0.9)
|
||
|
||
Input: Starting node (leaf from simulation), cumulative reward
|
||
Output: Updates visits and statevalue for all ancestors
|
||
|
||
while !isroot(node):
|
||
┌── Update visit count
|
||
│ └── node.visits += 1
|
||
│
|
||
├── Update statevalue (running average formula)
|
||
│ │
|
||
│ ├── Current average: node.statevalue
|
||
│ ├── Current count: node.visits - 1
|
||
│ ├── New value to add: simTrajectoryReward
|
||
│ │
|
||
│ └── Formula: new_avg = (old_avg * (n-1) + new_value) / n
|
||
│ node.statevalue = (node.statevalue * (node.visits-1) + simTrajectoryReward) / node.visits
|
||
│
|
||
├── Apply discount to future reward
|
||
│ └── simTrajectoryReward *= discountRewardCoeff
|
||
│ └── Reward gets 10% smaller for each level up
|
||
│
|
||
└── Move to parent
|
||
└── node = node.parent
|
||
|
||
# After loop, root node still needs update
|
||
# (root node's parent is nothing, so loop stops but root is still updated)
|
||
```
|
||
|
||
**Example**: Backpropagating reward=20 with discount=0.9
|
||
```
|
||
Depth 3 (leaf): Node A2b2
|
||
└── visits = 1, statevalue = 20/1 = 20.0
|
||
└── reward for parent = 20 * 0.9 = 18.0
|
||
|
||
Depth 2: Node A2b
|
||
└── visits = 1, statevalue = (0*0 + 18)/1 = 18.0
|
||
└── reward for parent = 18 * 0.9 = 16.2
|
||
|
||
Depth 1: Node A2
|
||
└── visits = 1, statevalue = (0*0 + 16.2)/1 = 16.2
|
||
└── reward for parent = 16.2 * 0.9 = 14.58
|
||
|
||
Depth 0 (root): Node A
|
||
└── visits = 1, statevalue = (0*0 + 14.58)/1 = 14.58
|
||
```
|
||
|
||
**Why discounting matters**: Future rewards are discounted because:
|
||
1. **Uncertainty**: The longer the horizon, the more uncertain the outcome
|
||
2. **Temporal credit assignment**: Immediate rewards should have higher value
|
||
3. **Realistic evaluation**: A reward 10 steps away is worth less than immediate reward
|
||
|
||
### 7. Solution Space Exploration Strategy
|
||
|
||
LLMMCTS explores the solution space using a combination of **guided exploration** and **systematic exploitation**:
|
||
|
||
#### Exploration vs. Exploitation Trade-off
|
||
|
||
```
|
||
Iteration 1: Initial exploration (all nodes have visits=0)
|
||
├── UCTselect() uses progressvalue for all children
|
||
├── LLM provides guidance: which states seem promising?
|
||
├── High progressvalue nodes get visited first
|
||
└── Results stored in statevalue
|
||
|
||
Iterations 2-5: Early exploitation
|
||
├── Some nodes have higher statevalue confirmed by simulations
|
||
├── UCT balances:
|
||
│ ├── Exploitation: High statevalue nodes
|
||
│ └── Exploration: Nodes with low visits (high UCT exploration term)
|
||
└── Nodes with high progressvalue + high statevalue dominate
|
||
|
||
Iterations 6+: Refinement
|
||
├── Well-visited nodes have reliable statevalue estimates
|
||
├── Poor nodes (low statevalue) are explored less
|
||
├── Still explore unvisited branches (UCT exploration term)
|
||
└── Tree converges to optimal trajectory
|
||
```
|
||
|
||
#### Horizontal vs. Vertical Exploration
|
||
|
||
**Horizontal exploration** (expansion):
|
||
- **Parameter**: `horizontalSample`
|
||
- **Purpose**: Generate multiple candidate next states from current state
|
||
- **Method**: Apply transition function multiple times
|
||
- **LLM role**: Each call generates a different thought/action
|
||
|
||
**Vertical exploration** (simulation):
|
||
- **Parameter**: `maxSimulationDepth`
|
||
- **Purpose**: Evaluate the long-term value of a state
|
||
- **Method**: Rollout from current node to leaf
|
||
- **Reward accumulation**: Sum all rewards along trajectory
|
||
|
||
#### Dejavu Detection and State Pruning
|
||
|
||
**Dejavu detection** (mcts.jl → _expand):
|
||
|
||
```julia
|
||
if newNodeKey ∉ keys(node.children)
|
||
newNode = MCTSNode(...)
|
||
node.children[newNodeKey] = newNode
|
||
end
|
||
```
|
||
|
||
**Why it matters**:
|
||
- Prevents infinite loops in cyclic state spaces
|
||
- Handles semantically equivalent states as duplicates
|
||
- Reduces tree size and memory usage
|
||
- Improves search efficiency
|
||
|
||
**Example**:
|
||
```
|
||
Current Node: "Problem solving step 5"
|
||
├── Attempt 1: Generates "Next step: Apply formula A"
|
||
├── Attempt 2: Generates "Next step: Apply formula A" (same as attempt 1)
|
||
│ └── Dejavu detected: skip duplicate
|
||
└── Attempt 3: Generates "Next step: Apply formula B"
|
||
```
|
||
|
||
### 8. Selection Phase Details (util.jl → UCTselect)
|
||
|
||
### 9. High-Level Interface (interface.jl → runMCTS)
|
||
|
||
**Purpose**: Execute the complete MCTS search algorithm with all phases.
|
||
|
||
**Process**:
|
||
|
||
```
|
||
runMCTS(initialstate, transition, transitionargs;
|
||
horizontalSampleExpansionPhase=3,
|
||
horizontalSampleSimulationPhase=3,
|
||
maxSimulationDepth=3,
|
||
maxiterations=10,
|
||
explorationweight=1.0,
|
||
earlystop=nothing,
|
||
saveSimulatedNode=false,
|
||
multithread=false)
|
||
|
||
Output: NamedTuple with:
|
||
├── root: Complete MCTS tree
|
||
├── bestNextState: Best immediate next state
|
||
├── bestTerminalState: Best final state along optimal trajectory
|
||
└── highValueStateList: List of high-value terminal states (reward >= 8)
|
||
|
||
Step-by-step execution:
|
||
└── root = MCTSNode("root", initialstate, 0, 0, 0, 0, false, nothing, {}, {})
|
||
└── highValueState = Channel{Any}(100)
|
||
|
||
for iteration = 1 to maxiterations:
|
||
┌── node = root
|
||
│ └── node.visits += 1 # Increment root visits
|
||
│
|
||
│ ┌── PHASE 1: SELECTION (until leaf)
|
||
│ │ while !isleaf(node):
|
||
│ │ └── node = UCTselect(node, explorationweight)
|
||
│ │
|
||
│ ├── PHASE 2: TERMINAL CHECK
|
||
│ │ if node.isterminal:
|
||
│ │ ├── if node.state[:reward] >= 8:
|
||
│ │ │ └── put!(highValueState, deepcopy(node.state))
|
||
│ │ └── backpropagate(node, node.reward)
|
||
│ │ └── Update statistics up to root
|
||
│ │ └── continue to next iteration
|
||
│ │
|
||
│ └── PHASE 3: EXPANSION
|
||
│ ├── expand(node, transition, transitionargs;
|
||
│ │ horizontalSample=horizontalSampleExpansionPhase,
|
||
│ │ multithread=multithread)
|
||
│ │
|
||
│ └── For each leafNode in node.children:
|
||
│ └── simulateThenBackpropagate(leafNode, ...)
|
||
│ ├── simulate() → rollout and get reward
|
||
│ ├── If terminal state with reward >= 8:
|
||
│ │ └── Store in highValueState
|
||
│ └── backpropagate() → update statistics
|
||
│
|
||
└── PHASE 4: EARLY STOP CHECK
|
||
└── if earlystop !== nothing && earlystop(node.state):
|
||
└── break # Exit iteration loop
|
||
|
||
┌── After all iterations, select best result:
|
||
│ ├── bestNextState = selectBestNextNode(root)
|
||
│ │ └── Uses statevalue/visits or progressvalue + reward
|
||
│ │
|
||
│ ├── bestTerminalState = selectBestTrajectoryNode(root)
|
||
│ │ └── Follows optimal trajectory to leaf
|
||
│ │
|
||
│ └── highValueStateList = collect from channel
|
||
│ └── while !isempty(highValueState):
|
||
│ push!(highValueStateList, take!(highValueState))
|
||
│
|
||
└── return (root=root, bestNextState=..., bestTerminalState=..., highValueStateList=...)
|
||
```
|
||
|
||
### 10. Complete MCTS Iteration Flow
|
||
|
||
```
|
||
┌─────────────────────────────────────────────────────────────────────────────┐
|
||
│ MCTS ITERATION 1 │
|
||
├─────────────────────────────────────────────────────────────────────────────┤
|
||
│ │
|
||
│ 1. SELECTION (UCT-based traversal): │
|
||
│ Root (visits=1) │
|
||
│ └── UCTselect() on root (has no children yet, so uses progressvalue) │
|
||
│ └── Select child with highest progressvalue │
|
||
│ └── Continue until reaching leaf node │
|
||
│ │
|
||
│ 2. EXPANSION (horizontal sampling): │
|
||
│ Leaf node │
|
||
│ └── expand(horizontalSample=3) │
|
||
│ └── Generate 3 child nodes via LLM transition │
|
||
│ └── Each child gets: progressvalue, reward, parent=leaf │
|
||
│ │
|
||
│ 3. SIMULATION (vertical rollout): │
|
||
│ Each child node │
|
||
│ └── simulate(maxSimulationDepth=3) │
|
||
│ └── Rollout 3 levels deep, accumulating rewards │
|
||
│ └── Return (simTrajectoryReward, terminalstate) │
|
||
│ │
|
||
│ 4. BACKPROPAGATION (update statistics): │
|
||
│ simTrajectoryReward │
|
||
│ └── backpropagate() up to root │
|
||
│ └── Update visits and statevalue for all ancestors │
|
||
│ └── Apply discount to future rewards │
|
||
│ │
|
||
└─────────────────────────────────────────────────────────────────────────────┘
|
||
|
||
┌─────────────────────────────────────────────────────────────────────────────┐
|
||
│ MCTS ITERATION 2 │
|
||
├─────────────────────────────────────────────────────────────────────────────┤
|
||
│ │
|
||
│ 1. SELECTION: │
|
||
│ Root (visits=2) │
|
||
│ └── UCTselect() now considers: │
|
||
│ ├── statevalue (from iteration 1) │
|
||
│ └── exploration term (low visits on unexpanded branches) │
|
||
│ └── May select different path than iteration 1 │
|
||
│ │
|
||
│ 2. EXPANSION: │
|
||
│ └── May expand different node or same node │
|
||
│ └── New children added to node.children │
|
||
│ │
|
||
│ 3. SIMULATION: │
|
||
│ └── Different rollout trajectory │
|
||
│ └── New reward estimate added to statistics │
|
||
│ │
|
||
│ 4. BACKPROPAGATION: │
|
||
│ └── Statistics updated with new information │
|
||
│ │
|
||
└─────────────────────────────────────────────────────────────────────────────┘
|
||
|
||
┌─────────────────────────────────────────────────────────────────────────────┐
|
||
│ MCTS ITERATION 3 to N │
|
||
├─────────────────────────────────────────────────────────────────────────────┤
|
||
│ │
|
||
│ Pattern repeats, with increasingly informed selection: │
|
||
│ │
|
||
│ • Nodes with high statevalue (confirmed by many simulations) │
|
||
│ → Exploited (selected frequently) │
|
||
│ │
|
||
│ • Nodes with low visits but promising progressvalue │
|
||
│ → Explored (UCT exploration term encourages tries) │
|
||
│ │
|
||
│ • Tree grows: more branches explored, more statistics accumulated │
|
||
│ │
|
||
│ • Best trajectory emerges from accumulated statistics │
|
||
│ │
|
||
└─────────────────────────────────────────────────────────────────────────────┘
|
||
```
|
||
|
||
## Key Functions
|
||
|
||
### Core MCTS (mcts.jl)
|
||
|
||
| Function | Purpose |
|
||
|----------|---------|
|
||
| `selectBestNextNode(node)` | Select best child based on `statevalue/visits` or `progressvalue + reward` |
|
||
| `selectBestTrajectoryNode(node)` | Traverse down tree to find highest-value leaf |
|
||
| `selectChildNode(node)` | Select child with highest `progressvalue + reward` |
|
||
| `expand(node, ...)` | Generate child nodes using transition function |
|
||
| `_expand(node, ...)` | Helper to create single child node |
|
||
| `simulate(node, ...)` | Perform rollout simulation |
|
||
| `backpropagate(node, reward)` | Update statistics up to root |
|
||
| `isleaf(node)` | Check if node has no children |
|
||
| `isroot(node)` | Check if node is root (`nodekey == "root"`) |
|
||
|
||
### Interface (interface.jl)
|
||
|
||
| Function | Purpose |
|
||
|----------|---------|
|
||
| `runMCTS(initialstate, transition, args; kwargs...)` | Main MCTS search function that executes all four phases across multiple iterations |
|
||
| `simulateThenBackpropagate(node, ...)` | Run simulation and backpropagate reward for a single node |
|
||
|
||
### Utilities (util.jl)
|
||
|
||
| Function | Purpose |
|
||
|----------|---------|
|
||
| `UCTselect(node, w)` | Select node using UCT score with exploration weight `w` |
|
||
| `selectBestNextNode(node)` | Select best child based on `statevalue/visits` or `progressvalue + reward` |
|
||
| `selectBestTrajectoryNode(node)` | Traverse down tree to find highest-value leaf |
|
||
| `selectChildNode(node)` | Select child with highest `progressvalue + reward` |
|
||
| `expand(node, ...)` | Generate child nodes using transition function |
|
||
| `_expand(node, ...)` | Helper to create single child node |
|
||
| `simulate(node, ...)` | Perform rollout simulation |
|
||
| `backpropagate(node, reward)` | Update statistics up to root |
|
||
| `isleaf(node)` | Check if node has no children |
|
||
| `isroot(node)` | Check if node is root (`nodekey == "root"`) |
|
||
|
||
## Configuration Parameters
|
||
|
||
### Keyword Arguments for runMCTS()
|
||
|
||
| Parameter | Default | Description |
|
||
|-----------|---------|-------------|
|
||
| `horizontalSampleExpansionPhase` | 3 | Number of child nodes generated at each expansion (horizontal sampling) |
|
||
| `horizontalSampleSimulationPhase` | 3 | Number of child nodes generated during simulation rollout |
|
||
| `maxSimulationDepth` | 3 | Maximum depth of simulation rollout (vertical exploration) |
|
||
| `maxiterations` | 10 | Number of MCTS iterations (each iteration: selection → expansion → simulation → backpropagation) |
|
||
| `explorationweight` | 1.0 | UCT exploration weight (controls exploration vs exploitation balance) |
|
||
| `earlystop` | nothing | Optional function to check early stopping condition (takes node.state, returns bool) |
|
||
| `saveSimulatedNode` | false | Whether to keep nodes created during simulation (true = keep for analysis, false = free memory) |
|
||
| `multithread` | false | Enable parallel simulation across child nodes (true = use Julia threads) |
|
||
|
||
### How Parameters Affect Search
|
||
|
||
| Parameter | Low Value | High Value | Trade-off |
|
||
|-----------|-----------|------------|-----------|
|
||
| `horizontalSample` | Fast, less exploration | Slow, more exploration | More children = better coverage but more LLM calls |
|
||
| `maxSimulationDepth` | Quick estimates, short horizon | Slow, long horizon | Deeper rollouts = more accurate but expensive |
|
||
| `maxiterations` | Quick search | Thorough search | More iterations = better solution but slower |
|
||
| `explorationweight` | Greedy (exploitation) | Aggressive (exploration) | Higher = explores more novel paths |
|
||
|
||
## Usage Pattern
|
||
|
||
```julia
|
||
# 1. Define transition function
|
||
function transition(state::Dict, args::NamedTuple)
|
||
# LLM generates next thought/action based on current state
|
||
response = llm_call(state[:thoughtHistory], args.prompt)
|
||
|
||
# Parse LLM output into new state
|
||
return Dict(
|
||
:newNodeKey => generate_key(),
|
||
:newstate => update_state(state, response),
|
||
:progressvalue => estimate_value(response)
|
||
)
|
||
end
|
||
|
||
# 2. Define transition arguments
|
||
transitionargs = (prompt="Solve this math problem", other_param="value")
|
||
|
||
# 3. Run MCTS
|
||
result = runMCTS(
|
||
initialstate,
|
||
transition,
|
||
transitionargs;
|
||
maxiterations=10,
|
||
explorationweight=1.0,
|
||
maxSimulationDepth=3
|
||
)
|
||
|
||
# 4. Access results
|
||
root = result.root
|
||
best_next_state = result.bestNextState
|
||
best_terminal_state = result.bestTerminalState
|
||
high_value_states = result.highValueStateList
|
||
```
|
||
|
||
## Key Differences from Traditional MCTS
|
||
|
||
| Aspect | Traditional MCTS | LLMMCTS |
|
||
|--------|-----------------|---------|
|
||
| Value estimation | Hand-designed heuristics | LLM-provided `progressvalue` |
|
||
| Reward signal | Environment only (sparse) | Environment + LLM pseudo-rewards (dense) |
|
||
| Expansion | Random or heuristics | LLM-guided generation |
|
||
| Learning speed | Slow (needs terminal rewards) | Fast (dense intermediate signals) |
|
||
|
||
## Performance Characteristics
|
||
|
||
- **Sample efficiency**: 5-10x fewer samples than pure Monte Carlo
|
||
- **LLM calls per iteration**: 1-5 (depends on `horizontalSample*` settings)
|
||
- **Convergence time**: Scales with depth × LLM latency
|
||
- **Memory usage**: O(branching_factor^depth) for tree storage
|
||
|
||
## Limitations
|
||
|
||
- **LLM latency**: Each node expansion requires an LLM call; can be slow for large trees
|
||
- **LLM cost**: Each LLM invocation has financial cost; monitor usage
|
||
- **Heuristic quality**: Poor LLM pseudo-rewards lead to suboptimal search
|
||
- **Determinism**: LLM outputs are stochastic; use temperature=0 for reproducibility
|
||
|
||
## Known Issues
|
||
|
||
### Bug: Variable name mismatch in interface.jl:94
|
||
|
||
The code references `highrewardNode` but the channel is named `highValueState`. This causes a `UndefVarError` when a terminal state with reward >= 8 is encountered.
|
||
|
||
**Current code (buggy)**:
|
||
```julia
|
||
highValueState = Channel{Any}(100)
|
||
# ...
|
||
if node.state[:reward] >= 8
|
||
put!(highrewardNode, deepcopy(node.state)) # BUG: should be highValueState
|
||
end
|
||
```
|
||
|
||
**Expected behavior**: Store high-value terminal states in the `highValueState` channel.
|
||
|
||
## Related Files
|
||
|
||
- `src/type.jl`: MCTSNode struct definition
|
||
- `src/util.jl`: UCT selection utility
|
||
- `src/mcts.jl`: Core MCTS operations
|
||
- `src/interface.jl`: High-level interface
|
||
- `src/LLMMCTS.jl`: Main package entry point
|
||
- `README.md`: User-facing documentation
|