update
This commit is contained in:
+219
@@ -0,0 +1,219 @@
|
||||
# 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) |
|
||||
|
||||
### Module Structure
|
||||
|
||||
```
|
||||
src/
|
||||
├── type.jl # MCTSNode struct definition
|
||||
├── util.jl # UCT selection utility function
|
||||
├── mcts.jl # Core MCTS operations (select, expand, simulate, backpropagate)
|
||||
├── interface.jl # High-level interface (runMCTS)
|
||||
└── LLMMCTS.jl # Main package entry point
|
||||
```
|
||||
|
||||
## 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. Main Workflow (interface.jl → runMCTS)
|
||||
|
||||
```
|
||||
runMCTS()
|
||||
├── Initialize root node
|
||||
└── For each iteration (1 to maxiterations):
|
||||
├── Traverse tree using UCT select (until leaf)
|
||||
│ └── UCTselect() → balances exploration/exploitation
|
||||
├── Check if leaf is terminal
|
||||
│ ├── Yes → backpropagate reward
|
||||
│ └── No → expand and simulate
|
||||
│ ├── expand() → generate child nodes (horizontal sampling)
|
||||
│ └── simulateThenBackpropagate()
|
||||
│ ├── simulate() → rollout to max depth
|
||||
│ └── backpropagate() → update statistics up to root
|
||||
└── Early stop check (if provided)
|
||||
```
|
||||
|
||||
### 3. Node Selection (util.jl → UCTselect)
|
||||
|
||||
Uses Upper Confidence Bound for Trees formula:
|
||||
|
||||
```
|
||||
UCT(s,a) = Q(s,a) + c * sqrt(ln(N(s)) / N(s,a))
|
||||
```
|
||||
|
||||
Where:
|
||||
- **Exploitation term** (`Q(s,a)`) — Uses `progressvalue` for fast guidance, refined by `statevalue`
|
||||
- **Exploration term** — Encourages visiting less-explored branches
|
||||
|
||||
### 4. Expansion (mcts.jl → expand/_expand)
|
||||
|
||||
Generates child nodes by applying the transition function multiple times (horizontal sampling).
|
||||
|
||||
**Dejavu detection**: Checks for semantically equivalent states to avoid duplicate nodes.
|
||||
|
||||
### 5. Simulation (mcts.jl → simulate)
|
||||
|
||||
Performs rollout from a node up to `maxSimulationDepth`:
|
||||
1. Accumulate current node's reward
|
||||
2. Check for terminal state
|
||||
3. Expand node (generate children)
|
||||
4. Select best child (using `progressvalue + reward`)
|
||||
5. Repeat until max depth reached
|
||||
|
||||
Returns cumulative trajectory reward and terminal state (if any).
|
||||
|
||||
### 6. Backpropagation (mcts.jl → backpropagate)
|
||||
|
||||
Updates statistics along the path from leaf to root:
|
||||
|
||||
```
|
||||
node.visits += 1
|
||||
node.statevalue = (node.statevalue * (node.visits-1) + simTrajectoryReward) / node.visits
|
||||
simTrajectoryReward *= discountRewardCoeff # Discount future rewards
|
||||
move to parent node
|
||||
```
|
||||
|
||||
## 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 |
|
||||
| `simulateThenBackpropagate(node, ...)` | Run simulation and backpropagate reward |
|
||||
|
||||
### Utilities (util.jl)
|
||||
|
||||
| Function | Purpose |
|
||||
|----------|---------|
|
||||
| `UCTselect(node, w)` | Select node using UCT score with exploration weight `w` |
|
||||
|
||||
## Configuration Parameters
|
||||
|
||||
### Keyword Arguments for runMCTS()
|
||||
|
||||
| Parameter | Default | Description |
|
||||
|-----------|---------|-------------|
|
||||
| `horizontalSampleExpansionPhase` | 3 | Children per node during expansion phase |
|
||||
| `horizontalSampleSimulationPhase` | 3 | Children per node during simulation |
|
||||
| `maxSimulationDepth` | 3 | Maximum depth during rollout |
|
||||
| `maxiterations` | 10 | Number of MCTS iterations |
|
||||
| `explorationweight` | 1.0 | UCT exploration weight (higher = more exploration) |
|
||||
| `earlystop` | nothing | Optional function to check early stopping condition |
|
||||
| `saveSimulatedNode` | false | Whether to keep nodes created during simulation |
|
||||
| `multithread` | false | Enable parallel simulation |
|
||||
|
||||
## 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
|
||||
|
||||
## 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
|
||||
Reference in New Issue
Block a user