Files
LLMMCTS/examples/configuration_examples.jl
2026-07-04 13:00:54 +07:00

240 lines
7.0 KiB
Julia

# MCTS Configuration Examples
This file demonstrates different MCTS configuration options and their effects on search behavior.
```julia
using LLMMCTS
# Simple transition function for demonstration
function simple_transition(state::Dict, args::NamedTuple)
current_step = get(state, :step, 0)
newstate = Dict(
:step => current_step + 1,
:reward => (current_step + 1) * 2,
:isterminal => current_step >= args.max_steps - 1
)
progressvalue = (current_step / args.max_steps) * 10
return Dict(
:newNodeKey => "step_$current_step",
:newstate => newstate,
:progressvalue => progressvalue
)
end
initialstate = Dict(
:step => 0,
:reward => 0,
:isterminal => false
)
transitionargs = (max_steps = 5,)
# ============================================================================
# Example 1: Balanced Search (Default)
# ============================================================================
println("Example 1: Balanced Search (Default)")
println("=" ^ 50)
result1 = runMCTS(
initialstate,
simple_transition,
transitionargs;
maxiterations = 10,
explorationweight = 1.0, # Balanced exploration/exploitation
maxSimulationDepth = 3,
horizontalSampleExpansionPhase = 3
)
println("Exploration weight: 1.0 (balanced)")
println("Root visits: ", result1.root.visits)
println("Best terminal step: ", result1.bestTerminalState[:step])
println()
# ============================================================================
# Example 2: Aggressive Exploration
# ============================================================================
println("Example 2: Aggressive Exploration")
println("=" * 50)
result2 = runMCTS(
initialstate,
simple_transition,
transitionargs;
maxiterations = 10,
explorationweight = 2.0, # More exploration
maxSimulationDepth = 3,
horizontalSampleExpansionPhase = 5 # More children per node
)
println("Exploration weight: 2.0 (aggressive exploration)")
println("Root visits: ", result2.root.visits)
println("Children explored: ", length(result2.root.children))
println()
# ============================================================================
# Example 3: Deep Search (Long Horizon)
# ============================================================================
println("Example 3: Deep Search (Long Horizon)")
println("=" * 50)
result3 = runMCTS(
initialstate,
simple_transition,
transitionargs;
maxiterations = 20,
explorationweight = 1.0,
maxSimulationDepth = 5, # Deeper search
horizontalSampleExpansionPhase = 3
)
println("Max simulation depth: 5 (deep search)")
println("Root visits: ", result3.root.visits)
println("Search explores further into the future")
println()
# ============================================================================
# Example 4: Fast Search (Shallow, Many Iterations)
# ============================================================================
println("Example 4: Fast Search (Shallow, Many Iterations)")
println("=" * 50)
result4 = runMCTS(
initialstate,
simple_transition,
transitionargs;
maxiterations = 50, # Many iterations
explorationweight = 1.0,
maxSimulationDepth = 2, # Shallow search
horizontalSampleExpansionPhase = 3
)
println("Many iterations (50), shallow depth (2)")
println("Root visits: ", result4.root.visits)
println("Faster but less thorough search")
println()
# ============================================================================
# Example 5: Parallel Simulation (Multithreading)
# ============================================================================
println("Example 5: Parallel Simulation (Multithreading)")
println("=" * 50)
result5 = runMCTS(
initialstate,
simple_transition,
transitionargs;
maxiterations = 10,
explorationweight = 1.0,
maxSimulationDepth = 3,
horizontalSampleExpansionPhase = 3,
multithread = true # Enable parallel simulation
)
println("Multithreading enabled")
println("Root visits: ", result5.root.visits)
println("Parallel simulation across child nodes")
println()
# ============================================================================
# Example 6: Early Stopping
# ============================================================================
println("Example 6: Early Stopping")
println("=" * 50)
# Define early stopping function
function early_stop(state::Dict)
# Stop when we reach a good enough solution
return get(state, :step, 0) >= 3
end
result6 = runMCTS(
initialstate,
simple_transition,
transitionargs;
maxiterations = 20, # Would run more if not for early stop
explorationweight = 1.0,
maxSimulationDepth = 3,
horizontalSampleExpansionPhase = 3,
earlystop = early_stop
)
println("Early stopping enabled (stops at step >= 3)")
println("Actual iterations: ", result6.root.visits)
println("Early stopping saved unnecessary computation")
println()
# ============================================================================
# Example 7: Save Simulation Nodes (for Analysis)
# ============================================================================
println("Example 7: Save Simulation Nodes")
println("=" * 50)
result7 = runMCTS(
initialstate,
simple_transition,
transitionargs;
maxiterations = 5,
explorationweight = 1.0,
maxSimulationDepth = 3,
horizontalSampleExpansionPhase = 3,
saveSimulatedNode = true # Keep simulation nodes
)
println("saveSimulatedNode = true")
println("Simulation nodes are preserved")
println("Root children: ", length(result7.root.children))
println("Useful for debugging or further analysis")
println()
# ============================================================================
# Example 8: High-Value State Tracking
# ============================================================================
println("Example 8: High-Value State Tracking")
println("=" * 50)
# Transition that can produce high-value states
function high_value_transition(state::Dict, args::NamedTuple)
current_step = get(state, :step, 0)
reward = current_step * 3
# Occasionally produce high-value states
if current_step == 2 || current_step == 4
reward = 9.0 # High value
end
newstate = Dict(
:step => current_step + 1,
:reward => reward,
:isterminal => current_step >= args.max_steps - 1
)
progressvalue = (current_step / args.max_steps) * 10
return Dict(
:newNodeKey => "step_$current_step",
:newstate => newstate,
:progressvalue => progressvalue
)
end
high_value_initial = Dict(
:step => 0,
:reward => 0,
:isterminal => false
)
result8 = runMCTS(
high_value_initial,
high_value_transition,
transitionargs;
maxiterations = 15,
explorationweight = 1.0,
maxSimulationDepth = 3,
horizontalSampleExpansionPhase = 3
)
println("High-value states found: ", length(result8.highValueStateList))
println("States with reward >= 8 were tracked")
for (i, state) in enumerate(result8.highValueStateList)
println(" High-value state $i: step = ", state[:step])
end
```