diff --git a/workprocess.md b/workprocess.md new file mode 100644 index 0000000..41fed5d --- /dev/null +++ b/workprocess.md @@ -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