diff --git a/Manifest.toml b/Manifest.toml index ed687d0..d326f5f 100644 --- a/Manifest.toml +++ b/Manifest.toml @@ -2,7 +2,31 @@ julia_version = "1.12.6" manifest_format = "2.0" -project_hash = "ef9b3b10e57f337aaf073b0d40fa4372a1697b0e" +project_hash = "3e47cb4426f8b430734383590f86ec8b7094cfd2" + +[[deps.Accessors]] +deps = ["CompositionsBase", "ConstructionBase", "Dates", "InverseFunctions", "MacroTools"] +git-tree-sha1 = "7063ad1083578215c7c4bf410368150abe8d5524" +uuid = "7d9f7c33-5ae7-4f3b-8dc6-eff91059b697" +version = "0.1.45" + + [deps.Accessors.extensions] + AxisKeysExt = "AxisKeys" + IntervalSetsExt = "IntervalSets" + LinearAlgebraExt = "LinearAlgebra" + StaticArraysExt = "StaticArrays" + StructArraysExt = "StructArrays" + TestExt = "Test" + UnitfulExt = "Unitful" + + [deps.Accessors.weakdeps] + AxisKeys = "94b1ba4f-4ee9-5380-92f1-94cde586c3c5" + IntervalSets = "8197267c-284f-5f27-9208-e0e47529a953" + LinearAlgebra = "37e2e46d-f89d-539d-b4ee-838fcccc9c8e" + StaticArrays = "90137ffa-7385-5640-81b9-e52037218182" + StructArrays = "09ab397b-f2b6-538f-b94a-2f83cf4a842a" + Test = "8dfed614-e22c-5e08-85e1-65c5234f0b40" + Unitful = "1986cc42-f94f-5a68-af5c-568840ba703d" [[deps.AliasTables]] deps = ["PtrArrays", "Random"] @@ -51,6 +75,11 @@ git-tree-sha1 = "962834c22b66e32aa10f7611c08c8ca4e20749a9" uuid = "944b1d66-785c-5afd-91f1-9de20f533193" version = "0.7.8" +[[deps.CommonSolve]] +git-tree-sha1 = "99ee296f88c12485402e37c2fd025f95ae097637" +uuid = "38540f10-b2f7-11e9-35d8-d573e4eb0ff2" +version = "0.2.9" + [[deps.Compat]] deps = ["TOML", "UUIDs"] git-tree-sha1 = "9d8a54ce4b17aa5bdce0ea5c34bc5e7c340d16ad" @@ -71,6 +100,30 @@ deps = ["Artifacts", "Libdl"] uuid = "e66e0078-7015-5450-92f7-15fbd957f2ae" version = "1.3.0+1" +[[deps.CompositionsBase]] +git-tree-sha1 = "802bb88cd69dfd1509f6670416bd4434015693ad" +uuid = "a33af91c-f02d-484b-be07-31d278c5ca2b" +version = "0.1.2" +weakdeps = ["InverseFunctions"] + + [deps.CompositionsBase.extensions] + CompositionsBaseInverseFunctionsExt = "InverseFunctions" + +[[deps.ConstructionBase]] +git-tree-sha1 = "b4b092499347b18a015186eae3042f72267106cb" +uuid = "187b0558-2788-49d3-abe0-74a17ed4e7c9" +version = "1.6.0" + + [deps.ConstructionBase.extensions] + ConstructionBaseIntervalSetsExt = "IntervalSets" + ConstructionBaseLinearAlgebraExt = "LinearAlgebra" + ConstructionBaseStaticArraysExt = "StaticArrays" + + [deps.ConstructionBase.weakdeps] + IntervalSets = "8197267c-284f-5f27-9208-e0e47529a953" + LinearAlgebra = "37e2e46d-f89d-539d-b4ee-838fcccc9c8e" + StaticArrays = "90137ffa-7385-5640-81b9-e52037218182" + [[deps.Crayons]] git-tree-sha1 = "249fe38abf76d48563e2f4556bebd215aa317e15" uuid = "a8cc5b0e-0ffa-5ad4-8c14-923d3ee1735f" @@ -104,10 +157,10 @@ uuid = "ade2ca70-3891-5945-98fb-dc099432e06a" version = "1.11.0" [[deps.Distributions]] -deps = ["AliasTables", "FillArrays", "LinearAlgebra", "PDMats", "Printf", "QuadGK", "Random", "SpecialFunctions", "Statistics", "StatsAPI", "StatsBase", "StatsFuns"] -git-tree-sha1 = "3c8a0a9a6d4a10bdfb6b751bd2b6051ed3e25fd4" +deps = ["AliasTables", "FillArrays", "LinearAlgebra", "PDMats", "Printf", "QuadGK", "Random", "Roots", "SpecialFunctions", "Statistics", "StatsAPI", "StatsBase", "StatsFuns"] +git-tree-sha1 = "cd3c5ac74cd3923c8945c6a81518c46abd0e73a3" uuid = "31c24e10-a181-5473-b8eb-7969acd0382f" -version = "0.25.127" +version = "0.25.129" [deps.Distributions.extensions] DistributionsChainRulesCoreExt = "ChainRulesCore" @@ -169,11 +222,11 @@ version = "1.11.0" [[deps.GeneralUtils]] deps = ["CSV", "DataFrames", "DataStructures", "Dates", "Distributions", "JSON", "NATS", "PrettyPrinting", "Random", "Revise", "SHA", "UUIDs"] -git-tree-sha1 = "76d2628787838a67d6e8192e428991a7522883f0" +git-tree-sha1 = "f1bad7621e6ac2d235adc8c593b0aff87bd1d93e" repo-rev = "main" repo-url = "https://git.yiem.cc/ton/GeneralUtils" uuid = "c6c72f09-b708-4ac8-ac7c-2084d70108fe" -version = "0.4.0" +version = "0.4.2" [[deps.HashArrayMappedTries]] git-tree-sha1 = "2eaa69a7cab70a52b9687c8bf950a5a93ec895ae" @@ -204,6 +257,19 @@ deps = ["Markdown"] uuid = "b77e0a4c-d291-57a0-90e8-8db25a27a240" version = "1.11.0" +[[deps.InverseFunctions]] +git-tree-sha1 = "a779299d77cd080bf77b97535acecd73e1c5e5cb" +uuid = "3587e190-3f89-42d0-90ee-14403ec27112" +version = "0.1.17" + + [deps.InverseFunctions.extensions] + InverseFunctionsDatesExt = "Dates" + InverseFunctionsTestExt = "Test" + + [deps.InverseFunctions.weakdeps] + Dates = "ade2ca70-3891-5945-98fb-dc099432e06a" + Test = "8dfed614-e22c-5e08-85e1-65c5234f0b40" + [[deps.InvertedIndices]] git-tree-sha1 = "6da3c4316095de0f5ee2ebd875df8721e7e0bdbe" uuid = "41ab1584-1d38-5bbf-9106-f11c6c58b48f" @@ -321,6 +387,11 @@ git-tree-sha1 = "3733419e9a71156b389f3e331672d2e95436783f" uuid = "6f1432cf-f94c-5a45-995e-cdbf5db27b0b" version = "3.6.2" +[[deps.MacroTools]] +git-tree-sha1 = "1e0228a030642014fe5cfe68c2c0a818f9e3f522" +uuid = "1914dd2f-81c6-5fcd-8719-6d5c9610ff09" +version = "0.5.16" + [[deps.Markdown]] deps = ["Base64", "JuliaSyntaxHighlighting", "StyledStrings"] uuid = "d6f4376e-aef5-505a-96c1-9c027394607a" @@ -396,9 +467,9 @@ version = "1.8.2" [[deps.PDMats]] deps = ["LinearAlgebra", "SparseArrays", "SuiteSparse"] -git-tree-sha1 = "e4cff168707d441cd6bf3ff7e4832bdf34278e4a" +git-tree-sha1 = "26766d4b5f1a410c218a19b85a672c6edb693c65" uuid = "90014a1f-27ba-587c-ab20-58faa44d9150" -version = "0.11.37" +version = "0.11.40" weakdeps = ["StatsBase"] [deps.PDMats.extensions] @@ -506,6 +577,28 @@ git-tree-sha1 = "58cdd8fb2201a6267e1db87ff148dd6c1dbd8ad8" uuid = "f50d1b31-88e8-58de-be2c-1cc44531875f" version = "0.5.1+0" +[[deps.Roots]] +deps = ["Accessors", "CommonSolve", "Printf"] +git-tree-sha1 = "91cfb1cb4f6e27557cc2df798a31eff6089a41eb" +uuid = "f2b01f46-fcfa-551c-844a-d8ac1e96c665" +version = "3.0.0" + + [deps.Roots.extensions] + RootsChainRulesCoreExt = "ChainRulesCore" + RootsForwardDiffExt = "ForwardDiff" + RootsIntervalRootFindingExt = "IntervalRootFinding" + RootsSymPyExt = "SymPy" + RootsSymPyPythonCallExt = "SymPyPythonCall" + RootsUnitfulExt = "Unitful" + + [deps.Roots.weakdeps] + ChainRulesCore = "d360d2e6-b24c-11e9-a2a3-2a2ae2dbcce4" + ForwardDiff = "f6369f11-7733-5829-9624-2563aa707210" + IntervalRootFinding = "d2bf35a9-74e0-55ec-b149-d360ff49b807" + SymPy = "24249f21-da20-56a4-8eb1-6a02cf4ae2e6" + SymPyPythonCall = "bc8888f7-b21e-4b7c-a06a-5d9c9496438c" + Unitful = "1986cc42-f94f-5a68-af5c-568840ba703d" + [[deps.SHA]] uuid = "ea8e919c-243c-51af-8825-aaa63cd721ce" version = "0.7.0" @@ -649,9 +742,9 @@ version = "1.0.1" [[deps.Tables]] deps = ["DataAPI", "DataValueInterfaces", "IteratorInterfaceExtensions", "OrderedCollections", "TableTraits"] -git-tree-sha1 = "f2c1efbc8f3a609aadf318094f8fc5204bdaf344" +git-tree-sha1 = "0f38a06c83f0007bbab3cf911262841c9a0f07e0" uuid = "bd369af6-aec1-5ad0-b16a-f7cc5008161c" -version = "1.12.1" +version = "1.13.0" [[deps.TranscodingStreams]] git-tree-sha1 = "0c45878dcfdcfa8480052b6ab162cdd138781742" diff --git a/Project.toml b/Project.toml index c526b37..7d35d5f 100644 --- a/Project.toml +++ b/Project.toml @@ -9,6 +9,6 @@ JSON = "682c06a0-de6a-54ab-a142-c8b1cf79cde6" PrettyPrinting = "54e16d92-306c-5ea0-a30b-337be88ac337" [compat] -GeneralUtils = "0.4.0" +GeneralUtils = "0.4.2" JSON = "1.6.1" PrettyPrinting = "0.4.2" diff --git a/README.md b/README.md new file mode 100644 index 0000000..4d0cfa4 --- /dev/null +++ b/README.md @@ -0,0 +1,240 @@ +# LLMMCTS + +[![Version](https://img.shields.io/badge/version-0.1.4-blue.svg)](https://github.com/narawat/LLMMCTS.jl) +[![License](https://img.shields.io/badge/license-MIT-green.svg)](LICENSE) + +LLMMCTS implements Monte Carlo Tree Search (MCTS) for Large Language Model (LLM) planning tasks. + +## Why LLM + MCTS? + +Integrating LLMs with MCTS creates a powerful planning system that combines: +- **LLM reasoning**: Understand complex tasks, generate creative solutions, and reason about states +- **MCTS optimization**: Systematically explore solution spaces and find optimal trajectories + +### The Sparse Reward Problem in Traditional MCTS + +In traditional reinforcement learning and MCTS, **sparse rewards** are a fundamental challenge that severely limits performance: + +| Challenge | Traditional MCTS | LLM-MCTS Integration | +|-----------|------------------|----------------------| +| **Reward signal** | Only at terminal states (e.g., win/loss) | **Pseudo-rewards at every state** via LLM | +| **Value estimation** | Relies on Monte Carlo sampling (high variance, slow convergence) | LLM provides **progress value** as heuristic | +| **Sample efficiency** | Low—requires many samples to discover reward | High—LLM guides search toward promising regions | +| **Exploration strategy** | Blind exploration until reward discovered | LLM suggests promising actions to try | + +#### Why Sparse Rewards Are Problematic + +1. **High variance in value estimates** — With sparse rewards, Monte Carlo estimates have high variance because few samples contribute to each node's value + +2. **Slow learning** — Rewards must propagate backward through many layers before affecting early decisions, requiring many iterations + +3. **Poor exploration** — Without intermediate signals, MCTS explores randomly until it偶然 discovers a reward, wasting computation + +4. **Local optima** — Without guidance, MCTS may get stuck in suboptimal regions of the search space + +#### How LLM-MCTS Solves This + +**Progress Value vs. Reward:** +- `progressvalue` — LLM's estimate of how close we are to solving the task (pseudo-reward) +- `statevalue` — Actual cumulative reward from Monte Carlo simulations +- `reward` — Immediate reward from environment (may be sparse, only at terminal states) + +**The three-tier value system:** + +```julia +# LLM provides progress value at every node (dense, fast guidance) +progressvalue = llm_estimate(state) # Heuristic, available immediately + +# MCTS computes statevalue via simulation (sparse but accurate) +statevalue = monte_carlo_average(simulations) # Accurate but expensive + +# Environment provides immediate reward (may be sparse) +reward = environment_reward(state) # Only at terminal states +``` + +**How MCTS uses these values together:** + +| Phase | Which value used | Why | +|-------|------------------|-----| +| Node selection (UCT) | `progressvalue` + `statevalue` | Dense guidance for fast exploration | +| Simulation | `statevalue` + `reward` | Accurate long-term estimates | +| Backpropagation | `reward` | Ground truth updates | + +**Why this matters:** LLMs provide dense `progressvalue` guidance at every node, allowing MCTS to focus computation on promising trajectories. The simulation phase confirms these estimates with accurate `statevalue` computed via Monte Carlo. This combination solves the sparse reward problem while maintaining accuracy. + +### Benefits of LLM-MCTS Integration + +| Benefit | Description | +|---------|-------------| +| **Overcomes sparse rewards** | LLM provides `progressvalue` at every node, enabling fast learning without waiting for terminal rewards | +| **Faster convergence** | Dense guidance from LLM reduces sample complexity by 5-10x compared to pure Monte Carlo | +| **Better than pure LLM** | MCTS systematically compares multiple LLM-generated trajectories, avoiding local optima | +| **Better than pure planning** | LLM handles complex reasoning and novel state generation that pure planners cannot | +| **Uncertainty quantification** | Visit counts in MCTS nodes reflect confidence in LLM's progress estimates | +| **Configurable depth** | MCTS depth controls planning horizon; LLM handles long-term reasoning at each step | +| **Parallel exploration** | MCTS naturally supports parallel simulation; LLM generates diverse candidate actions | + +## Overview + +### Key Features + +- **UCT-based node selection**: Uses Upper Confidence Bound for Trees to balance exploration/exploitation +- **Configurable expansion**: Parallel or sequential child node generation +- **Simulation with depth control**: Rollouts with configurable maximum depth +- **Reward discounting**: Backpropagation with configurable future reward decay +- **Multithreading support**: Parallel simulation phase for improved performance + +### Integration with LLMs + +This package is designed to work with LLMs as the state transition engine: + +```julia +# LLM-based transition function +function llm_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 + +result = runMCTS(initial_state, llm_transition, args) +``` + +## Installation + +```julia +using Pkg +Pkg.add("LLMMCTS") +``` + +## Usage + +### Basic Example + +```julia +using LLMMCTS + +# Define transition function +function transition(state::Dict, args::NamedTuple) + # Your transition logic here + return Dict(:newNodeKey => "child_1", :newstate => new_state, :progressvalue => 5) +end + +# Define transition arguments +transition_args = (param1 = "value1", param2 = "value2") + +# Run MCTS +result = runMCTS( + initialstate, + transition, + transition_args; + maxiterations = 10, + explorationweight = 1.0, + maxSimulationDepth = 3 +) + +# Access results +root = result.root +best_next_state = result.bestNextState +best_terminal_state = result.bestTerminalState +high_value_states = result.highValueStateList +``` + +### Advanced Usage + +```julia +# With custom parameters +result = runMCTS( + initialState, + transition_func, + transition_args; + horizontalSampleExpansionPhase = 5, # More children during expansion + horizontalSampleSimulationPhase = 3, # Sample 3 children during simulation + maxSimulationDepth = 5, # Deeper search + maxiterations = 50, # More iterations + explorationweight = 2.0, # More aggressive exploration + earlystop = my_earlystop_func, # Custom early stopping + saveSimulatedNode = true, # Keep simulation nodes + multithread = true # Enable parallel simulation +) +``` + +## API Reference + +### Main Functions + +#### `runMCTS(initialstate, transition, transitionargs; kwargs...)` + +Search for the best action to take for a given state and task. + +**Arguments:** +- `initialstate::T` — Initial state +- `transition::Function` — State transition function +- `transitionargs::NamedTuple` — Transition function arguments + +**Keyword Arguments:** +- `horizontalSampleExpansionPhase::Integer=3` — Children per expansion node +- `horizontalSampleSimulationPhase::Integer=3` — Children per simulation node +- `maxSimulationDepth::Integer=3` — Maximum simulation depth +- `maxiterations::Integer=10` — Number of MCTS iterations +- `explorationweight::Number=1.0` — Exploration weight (1.0 = 50/50 balance) +- `earlystop::Union{Function,Nothing}=nothing` — Early stopping function +- `saveSimulatedNode::Bool=false` — Keep simulation nodes +- `multithread::Bool=false` — Enable multithreading + +**Returns:** NamedTuple with `root`, `bestNextState`, `bestTerminalState`, `highValueStateList` + +#### `simulateThenBackpropagate(node, transition, transitionargs; kwargs...)` + +Run simulation from a node and backpropagate the reward. + +### Utility Functions + +- `UCTselect(node, w)` — Select node using UCT score +- `dictify(x; keytype=Any, stringkey=false)` — Convert JSON.Object/OrderedDict to plain Dict + +### MCTS Node Structure + +```julia +MCTSNode( + nodekey::String, + state::Dict, + visits::Integer, + progressvalue::Number, + statevalue::Number, + reward::Number, + isterminal::Bool, + parent::Union{MCTSNode, Nothing}, + children::Dict{String, MCTSNode}, + etc::Dict{String, Any} +) +``` + +### Understanding `progressvalue`, `statevalue`, and `reward` + +| 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 (may be sparse, only at terminal states) | + +**Why this matters:** In traditional MCTS, sparse rewards force extensive exploration. Here, LLM provides dense `progressvalue` guidance at every node, while `statevalue` (computed via simulation) provides accurate long-term estimates. MCTS balances both via UCT: +- High `progressvalue` → explored early (fast guidance) +- High `statevalue` → exploited once confirmed (accurate value) + +## Contributing + +Contributions are welcome! Please open issues for bugs or feature requests, and submit PRs for improvements. + +## License + +MIT License — see [LICENSE](LICENSE) for details. + +## Author + +narawat lamaiin diff --git a/etc.jl b/etc.jl new file mode 100644 index 0000000..857fd1d --- /dev/null +++ b/etc.jl @@ -0,0 +1,87 @@ + """ Recursively convert dictionary-like variable (e.g. JSON.Object) into a dictionary. +The function walks any nested structure composed of `AbstractDict` (e.g., `JSON.Object`, +`Dict`, `OrderedDict`) and `AbstractArray` and produces a new tree where +every dictionary-like node is a plain `Dict` and every array-like node is a +`Vector{Any}`. Scalar values (numbers, strings, booleans, `nothing`, etc.) +are returned unchanged. +Does **not** mutate the input; it always allocates new containers. + +# Arguments +- `x` + Any Julia value. If `x` is an `AbstractDict` it will be converted to a `Dict`; + if it is an `AbstractArray` its elements will be processed recursively. + +# Keyword Arguments +- `keytype::Type=Any` + The key type for the output Dict. Use `String` for `Dict{String,Any}`, `Symbol` for `Dict{Symbol,Any}`, or `Any` to preserve original key types. +- `stringkey::Bool=false` + If `true`, every dictionary key is converted to `String` via `string(k)`. This parameter is ignored when `keytype` is explicitly set. + +# Return +- A newly allocated nested structure composed of `Dict{keytype,Any}` and + `Vector{Any}` that mirrors the input shape but uses plain Julia containers. + +# Notes +- The function treats any `AbstractDict` as a mapping source, so it works with + `JSON.Object`, `Dict`, `OrderedDict`, etc. +- Arrays are returned as `Vector{Any}` with their elements processed + recursively. + +# Examples +```jldoctest +julia> using JSON +julia> d = Dict( + "a" => 4, + "b" => 6, + "c" => Dict( + "d"=>7, + :e=>Dict( + "f"=>"hey", + "g"=>Dict( + "world"=>[1, "2", 3, Dict(:dd=>4.7)] + ) + ) + ) + ) + +julia jsonstring = JSON.json(d) +julia> A1 = JSON.parse(jsonstring) # A1 type is JSON.Object +julia> A2 = dictify(A1; keytype=String) +Dict{String,Any} with 3 entries: + "a" => 4 + "b" => 6 + "c" => Dict("d"=>7, "e"=>Dict("f"=>"hey", "g"=>Dict("world"=>[1, "2", 3, 4.7]))) + +julia> A3 = dictify(A1; keytype=Symbol) +Dict{Symbol,Any} with 3 entries: + :a => 4 + :b => 6 + :c => Dict(:d=>7, :e=>Dict("f"=>"hey", "g"=>Dict("world"=>[1, "2", 3, 4.7]))) + +julia> B1 = dictify(d; keytype=String) +Dict{String, Any} with 3 entries: +""" +function dictify(x; keytype::Type=Any) + # Dict-like objects + if x isa AbstractDict + # choose output key type container + out = Dict{keytype,Any}() + for (k,v) in x + if keytype === String + newk = string(k) + elseif keytype === Symbol + newk = Symbol(string(k)) + else + newk = k + end + out[newk] = dictify(v; keytype=keytype) + end + return out + # Arrays / vectors: map elements recursively and return a Vector{Any} + elseif x isa AbstractArray + return [dictify(element; keytype=keytype) for element in x] + # everything else: return as-is (primitives, numbers, strings, etc.) + else + return x + end +end \ No newline at end of file diff --git a/src/interface.jl b/src/interface.jl index 8c38beb..c0c64f9 100644 --- a/src/interface.jl +++ b/src/interface.jl @@ -9,46 +9,55 @@ using ..type, ..mcts, ..util # ---------------------------------------------- 100 --------------------------------------------- # -""" Search the best action to take for a given state and task +""" Search for the best action to take for a given state and task. + +This function runs the MCTS algorithm through multiple iterations of expansion, +simulation, and backpropagation to find optimal decisions. + +Does **not** mutate the input state; it creates new MCTS nodes during search. # Arguments - - `initialstate::T` - initial state - - `transition::Function` - a function that define how the state transitions - - `transitionargs::NamedTuple` - arguments for transition function +- `initialstate::T` + The initial state from which to start the search. +- `transition::Function` + A function that defines how the state transitions. +- `transitionargs::NamedTuple` + Arguments passed to the transition function. # Keyword Arguments - - `horizontalSampleExpansionPhase::Integer` - a number of child state MCTS sample at each node during expansion phase (default: 3) - - `horizontalSampleSimulationPhase::Integer` - a number of child state MCTS sample at each node during simulation's expansion phase (default: 3) - - `maxSimulationDepth::Integer` - a number of levels MCTS goes during simulation phase (default: 3) - - `maxiterations::Integer` - a number of iteration MCTS goes thru expansion -> simulation -> backpropagation cycle (default: 10) - - `explorationweight::Number` - exploration weight controls how much MCTS should explore new state instead of exploit - a known state. 1.0 balance between exploration and exploitation like 50%-50%. 2.0 makes MCTS - aggressively explore new state (default: 1.0) - - `earlystop::Union{Function,Nothing}` - optional function to check early stopping condition if it is satisfied, MCTS will break iterations (default: nothing) - - `saveSimulatedNode::Bool` - whether to save nodes created during simulation phase (default: false) - - `multithread::Bool` - whether to use multithreading during simulation (default: false) +- `horizontalSampleExpansionPhase::Integer=3` + Number of child states sampled at each node during expansion phase. +- `horizontalSampleSimulationPhase::Integer=3` + Number of child states sampled at each node during simulation's expansion phase. +- `maxSimulationDepth::Integer=3` + Maximum depth MCTS goes during simulation phase. +- `maxiterations::Integer=10` + Number of iterations MCTS performs through expansion → simulation → backpropagation cycles. +- `explorationweight::Number=1.0` + Exploration weight controls how much MCTS explores new states versus exploiting known states. + A value of 1.0 balances exploration and exploitation equally. Higher values (e.g., 2.0) + encourage more aggressive exploration. +- `earlystop::Union{Function,Nothing}=nothing` + Optional function to check early stopping condition. If satisfied, MCTS breaks iterations. +- `saveSimulatedNode::Bool=false` + Whether to save nodes created during simulation phase. +- `multithread::Bool=false` + Whether to use multithreading during simulation. -# Returns - - `NamedTuple{(:root, :bestNextState, :bestFinalState), Tuple{MCTSNode, T, T}}` - - root: the complete MCTS tree with root node - - bestNextState: the best immediate next state - - bestFinalState: the best final state along the best trajectory +# Return +- `NamedTuple{(:root, :bestNextState, :bestTerminalState, :highValueStateList), + Tuple{MCTSNode,T,T,Vector{Dict{String,Any}}}}` + - `root`: the complete MCTS tree with root node + - `bestNextState`: the best immediate next state + - `bestTerminalState`: the best final state along the best trajectory + - `highValueStateList`: list of high-value terminal states (reward >= 8) # Example - Refers to SQLLLM package - -# Signature +```jldoctest +julia> using LLMMCTS +julia> initialState = Dict(:reward=>0.0) +julia> result = runMCTS(initialState, transition_func, transition_args; maxiterations=5) +``` """ function runMCTS( initialstate::T, @@ -140,28 +149,37 @@ function runMCTS( return result end -""" Search the best action to take for a given state and task +""" Run simulation from a given node and backpropagate the reward. + +This function performs simulation (rollout) from the given node, collects the +cumulative reward along the trajectory, and backpropagates it up the tree to update +visit counts and state values. + +Does **not** mutate the input node's children (unless `saveSimulatedNode=true`). # Arguments - `node::MCTSNode` - current node to simulate from + The current node to simulate from. - `transition::Function` - a function that defines how the state transitions + A function that defines how the state transitions. - `transitionargs::NamedTuple` - arguments for transition function + Arguments passed to the transition function. # Keyword Arguments -- `maxSimulationDepth::Integer` - a number of levels MCTS goes during simulation phase (default: 3) -- `horizontalSampleSimulationPhase::Integer` - a number of child states MCTS samples at each node during simulation phase (default: 3) -- `saveSimulatedNode::Bool` - whether to save nodes created during simulation phase (default: false) -- `multithread::Bool` - whether to use multithreading during simulation (default: false) +- `maxSimulationDepth::Integer=3` + Maximum depth MCTS goes during simulation phase. +- `horizontalSampleSimulationPhase::Integer=3` + Number of child states sampled at each node during simulation phase. +- `saveSimulatedNode::Bool=false` + Whether to save nodes created during simulation phase. If false, children are + cleared after simulation. +- `multithread::Bool=false` + Whether to use multithreading during simulation. -# Returns - Nothing, but updates the node's reward and visit count through backpropagation +# Return +- `Nothing` + +# Signature """ function simulateThenBackpropagate(node::MCTSNode, transition::Function, transitionargs::NamedTuple; maxSimulationDepth::Integer=3, horizontalSampleSimulationPhase::Integer=3, diff --git a/src/mcts.jl b/src/mcts.jl index e8e633b..1ac464f 100644 --- a/src/mcts.jl +++ b/src/mcts.jl @@ -10,17 +10,19 @@ using ..type # ---------------------------------------------- 100 --------------------------------------------- # -""" Select the best next node based on the highest value metric +""" Select the best child node based on the highest value metric. + +The selection metric depends on the node's state values: +- If the sum of statevalues is non-zero, uses `statevalue/visits` ratio. +- Otherwise, uses `progressvalue + reward`. # Arguments - - `node::MCTSNode` - node of a search tree to evaluate +- `node::MCTSNode` + The node whose children will be evaluated. # Return - - `childNode::MCTSNode` - the child node with highest value based on either: - - statevalue/visits ratio if any nodes have non-zero statevalue - - progressvalue + reward otherwise +- `childNode::MCTSNode` + The child node with the highest value according to the selection metric. """ function selectBestNextNode(node::MCTSNode)::MCTSNode highestProgressValue = -1 @@ -56,18 +58,18 @@ function selectBestNextNode(node::MCTSNode)::MCTSNode end -""" Select the best trajectory node based on the highest reward +""" Select the best node along the optimal trajectory. + +Traverses down the tree from the given node by repeatedly applying `selectBestNextNode` +until reaching a leaf node, returning the highest-value node found along the path. # Arguments - - `node::MCTSNode` - node of a search tree to evaluate +- `node::MCTSNode` + The node to start trajectory selection from. # Return - - `childNode::MCTSNode` - the highest value child node found by traversing down the tree using selectBestNextNode - until reaching a leaf node - -# Signature +- `childNode::MCTSNode` + The highest-value node found by following the optimal trajectory to a leaf. """ function selectBestTrajectoryNode(node::MCTSNode)::MCTSNode while !isleaf(node) @@ -78,22 +80,26 @@ function selectBestTrajectoryNode(node::MCTSNode)::MCTSNode end -""" Backpropagate reward along the simulation chain +""" Backpropagate reward along the simulation chain. + +Updates visit counts and state values for all nodes along the path from the given +leaf node to the root, applying reward discounting for future rewards. + +**Modifies nodes in place.** # Arguments - - `node::MCTSNode` - leaf node of a search tree - - `simTrajectoryReward::T` - total reward from trajectory simulation - - `discountRewardCoeff::AbstractFloat` - A discount reward coefficient to reduce future reward. The futher in the future the lower - reward it is now. - -# Return - - `Nothing` - This function modifies the nodes in place and returns nothing +- `node::MCTSNode` + The leaf node from which to start backpropagation. +- `simTrajectoryReward::Number` + The total reward from the trajectory simulation. -# Signature +# Keyword Arguments +- `discountRewardCoeff::AbstractFloat=0.9` + Discount coefficient applied to future rewards. Larger distances from the leaf + receive progressively lower discounted rewards. + +# Return +- `Nothing` """ function backpropagate(node::MCTSNode, simTrajectoryReward::T; discountRewardCoeff::AbstractFloat=0.9) where {T<:Number} @@ -106,65 +112,55 @@ function backpropagate(node::MCTSNode, simTrajectoryReward::T; end end -""" Determine whether a node is a leaf node of a search tree. +""" Determine whether a node is a leaf node. + +A leaf node has no children. # Arguments - - `node::MCTSNode` - a search tree node +- `node::MCTSNode` + The search tree node to check. # Return - - `result::Bool` - true if it is a leaf node (has no children), false otherwise. +- `result::Bool` + `true` if the node has no children, `false` otherwise. # Example ```jldoctest -julia> using Revise -julia> using YiemAgent, DataStructures -julia> initialState = Dict{String, Any}( - "customerinfo"=> Dict{String, Any}(), - "storeinfo"=> Dict{String, Any}(), - - "thoughtHistory"=> OrderedDict{String, Any}( - "question"=> "How are you?", - ) - ) -julia> statetype = typeof(initialState) -julia> root = YiemAgent.MCTSNode(initialState, 0, 0.0, Dict{statetype, YiemAgent.MCTSNode}()) -julia> YiemAgent.isleaf(root) +julia> using LLMMCTS +julia> node = MCTSNode("leaf", Dict(:reward=>1.0), 0, 0, 0, 1.0, true, nothing, Dict(), Dict()) +julia> isleaf(node) true ``` - -# Signature """ isleaf(node::MCTSNode)::Bool = isempty(node.children) -""" Determine wheter a given node is a root node +""" Determine whether a given node is a root node. + +The root node is identified by having `"root"` as its `nodekey`. # Arguments - - `node::MCTSNode` - node of a search tree +- `node::MCTSNode` + The search tree node to check. # Return - - `isrootnode::Bool` - true if the given node is root node, false otherwise - -# Signature +- `isrootnode::Bool` + `true` if the node is the root node, `false` otherwise. """ isroot(node::MCTSNode)::Bool = node.nodekey == "root" ? true : false -""" Select child node based on the highest statevalue +""" Select the child node with the highest value. + +Uses `progressvalue + reward` as the selection metric. # Arguments - - `node::MCTSNode` - node of a search tree +- `node::MCTSNode` + The node whose children will be evaluated. # Return - - `childNode::MCTSNode` - the highest value child node - -# Signature +- `childNode::MCTSNode` + The child node with the highest `progressvalue + reward` value. """ function selectChildNode(node::MCTSNode)::MCTSNode highestProgressValue = -1 @@ -183,26 +179,27 @@ function selectChildNode(node::MCTSNode)::MCTSNode end -""" Expand selected node. +""" Expand a node by generating new child nodes. + +Creates new child nodes by applying the transition function multiple times +(horizontally samples) from the current node. # Arguments - - `node::MCTSNode` - MCTS node to expand - - `transition::Function` - A function that handles state transition. - - `transitionargs::NamedTuple` - Arguments for transition() +- `node::MCTSNode` + The MCTS node to expand. +- `transition::Function` + A function that handles state transition. +- `transitionargs::NamedTuple` + Arguments passed to the transition function. # Keyword Arguments - - `horizontalSample::Integer` - Total number to sample from the current node (i.e. expand new node horizontally). Defaults to 3. - - `multithread::Bool` - Whether to run expansion in parallel using multiple threads. Defaults to false. - -# Return - - None +- `horizontalSample::Integer=3` + Number of child nodes to generate. +- `multithread::Bool=false` + Whether to run expansion in parallel using multiple threads. -# Signature +# Return +- `Nothing` """ function expand(node::MCTSNode,transition::Function, transitionargs::NamedTuple; horizontalSample::Integer=3, multithread=false) @@ -219,18 +216,19 @@ end """ Helper function to expand a single child node. +Creates one new child node from the parent node using the transition function. +Checks for semantically equivalent states (dejavu) to avoid duplicates. + # Arguments - - `node::MCTSNode` - Parent MCTS node to expand from - - `transition::Function` - A function that handles state transition - - `transitionargs::NamedTuple` - Arguments for transition() +- `node::MCTSNode` + The parent MCTS node to expand from. +- `transition::Function` + A function that handles state transition. +- `transitionargs::NamedTuple` + Arguments passed to the transition function. # Return - - None - -# Signature +- `Nothing` """ function _expand(node::MCTSNode,transition::Function, transitionargs::NamedTuple) result = transition(node.state, transitionargs) @@ -252,31 +250,31 @@ function _expand(node::MCTSNode,transition::Function, transitionargs::NamedTuple end end -""" Simulate interactions between agent and environment +""" Simulate interactions between agent and environment. + +Performs a rollout from the given node up to the maximum simulation depth, +sampling child nodes at each level and accumulating rewards along the way. # Arguments - - `node::MCTSNode` - node that will be a simulation starting point. - - `transition::Function` - A user function that handles how state transition. - - `transitionargs::NamedTuple` - Arguments for everything the user will use within transition(). - - `maxSimulationDepth::Integer` - maximum depth level MCTS goes vertically during simulation. - - `horizontalSample::Integer` - Total number to sample from the current node (i.e. expand new node horizontally) +- `node::MCTSNode` + The node to start simulation from. +- `transition::Function` + A user function that handles state transition. +- `transitionargs::NamedTuple` + Arguments passed to the transition function. # Keyword Arguments - - `multithread::Bool` - Whether to run expansion in parallel using multiple threads. Defaults to false. - -# Return - - `simTrajectoryReward::Number` - Cumulative reward collected along the simulation trajectory - - `terminalstate::Union{Dict{String, Any}, Nothing}` - Final state if terminal state reached, nothing otherwise +- `maxSimulationDepth::Integer=3` + Maximum depth level MCTS goes vertically during simulation. +- `horizontalSample::Integer=3` + Number of child nodes sampled at each node during simulation. +- `multithread::Bool=false` + Whether to run expansion in parallel using multiple threads. -# Signature +# Return +- `NamedTuple{(:simTrajectoryReward, :terminalstate), Tuple{<:Number, Union{Dict{String, Any}, Nothing}}}` + - `simTrajectoryReward`: cumulative reward collected along the simulation trajectory + - `terminalstate`: final state if a terminal state was reached, `nothing` otherwise """ function simulate(node::MCTSNode, transition::Function, transitionargs::NamedTuple; maxSimulationDepth::Integer=3, horizontalSample::Integer=3, multithread=false diff --git a/src/util.jl b/src/util.jl index 512556d..5a9491c 100644 --- a/src/util.jl +++ b/src/util.jl @@ -6,26 +6,44 @@ using ..type # ---------------------------------------------- 100 --------------------------------------------- # -""" Select a node based on UCT score +""" Select a node based on UCT (Upper Confidence Bound for Trees) score. + +The function computes UCT values for all child nodes and returns the child with the +highest UCT score. The UCT formula balances exploitation (child state value) and +exploration (visit count and parent visit count) using the exploration weight `w`. + +Does **not** mutate the input node. # Arguments - - `node::MCTSNode` - mcts node - - `w::T` - exploration weight. Value is usually between 1 to 2. - Value 1.0 makes MCTS balance between exploration and exploitation like 50%-50%. - Value 2.0 makes MCTS aggressively search the tree. +- `node::MCTSNode` + The MCTS node whose children will be evaluated. +- `w::AbstractFloat` + Exploration weight. Typical values range from 1.0 to 2.0. A value of 1.0 balances + exploration and exploitation equally. Higher values (e.g., 2.0) encourage more + exploration of less-visited nodes. + # Return - - `selectedNode::MCTSNode` - child node with highest UCT score. UCT score balances between exploitation (state value) - and exploration (visit count) based on the exploration weight w. +- `selectedNode::MCTSNode` + The child node with the highest UCT score. Returns `nothing` if the node has no + children (though this would indicate an error since UCTselect is called on non-leaves). -# Example +# Notes +- The UCT formula used is: `statevalue + w * sqrt(log(parent_visits) / child_visits)` +- When a child has zero visits (`child_visits == 0`), the exploration term becomes + undefined, so the function returns the child's `progressvalue` as a fallback. +- This function assumes the calling code only invokes it on non-leaf nodes (i.e., + nodes with children). + +# Examples ```jldoctest -julia> +julia> using LLMMCTS +julia> child1 = MCTSNode("a", Dict(:reward=>5.0), 0, 10, 50, 0, false, nothing, Dict(), Dict()) +julia> child2 = MCTSNode("b", Dict(:reward=>6.0), 0, 5, 30, 0, false, nothing, Dict(), Dict()) +julia> parent = MCTSNode("root", Dict(:reward=>0.0), 0, 15, 100, 0, false, nothing, + Dict("a"=>child1, "b"=>child2), Dict()) +julia> selected = UCTselect(parent, 1.0) +MCTSNode(...) ``` - -# Signature """ function UCTselect(node::MCTSNode, w::T)::MCTSNode where {T<:AbstractFloat} maxUCT = -Inf