diff --git a/.gitignore b/.gitignore index f151518..b33b451 100644 --- a/.gitignore +++ b/.gitignore @@ -17,6 +17,9 @@ src/coolrl_lost_cities/games/classic/deep_cfr/*.c # Rust build output target/ +# Local toolchains +tools/julia/ + # Virtual environments .venv diff --git a/experiments/julia_safe_heuristic/Manifest.toml b/experiments/julia_safe_heuristic/Manifest.toml new file mode 100644 index 0000000..03e83c8 --- /dev/null +++ b/experiments/julia_safe_heuristic/Manifest.toml @@ -0,0 +1,148 @@ +# This file is machine-generated - editing it directly is not advised + +julia_version = "1.11.9" +manifest_format = "2.0" +project_hash = "81d3e26811b735b24d9bccb7176024c1ed320ee8" + +[[deps.Artifacts]] +uuid = "56f22d72-fd6d-98f1-02f0-08ddc0907c33" +version = "1.11.0" + +[[deps.BenchmarkTools]] +deps = ["Compat", "JSON", "Logging", "PrecompileTools", "Printf", "Profile", "Statistics", "UUIDs"] +git-tree-sha1 = "9670d3febc2b6da60a0ae57846ba74670290653f" +uuid = "6e4b80f9-dd63-53aa-95a3-0cdb28fa8baf" +version = "1.8.0" + +[[deps.Compat]] +deps = ["TOML", "UUIDs"] +git-tree-sha1 = "9d8a54ce4b17aa5bdce0ea5c34bc5e7c340d16ad" +uuid = "34da2185-b29b-5c13-b0c7-acf172513d20" +version = "4.18.1" +weakdeps = ["Dates", "LinearAlgebra"] + + [deps.Compat.extensions] + CompatLinearAlgebraExt = "LinearAlgebra" + +[[deps.CompilerSupportLibraries_jll]] +deps = ["Artifacts", "Libdl"] +uuid = "e66e0078-7015-5450-92f7-15fbd957f2ae" +version = "1.1.1+0" + +[[deps.Dates]] +deps = ["Printf"] +uuid = "ade2ca70-3891-5945-98fb-dc099432e06a" +version = "1.11.0" + +[[deps.JSON]] +deps = ["Dates", "Logging", "Parsers", "PrecompileTools", "StructUtils", "UUIDs", "Unicode"] +git-tree-sha1 = "fe23330af47b8ab4e135b2ff65f7398c3a2bfc65" +uuid = "682c06a0-de6a-54ab-a142-c8b1cf79cde6" +version = "1.5.2" + + [deps.JSON.extensions] + JSONArrowExt = ["ArrowTypes"] + + [deps.JSON.weakdeps] + ArrowTypes = "31f734f8-188a-4ce0-8406-c8a06bd891cd" + +[[deps.Libdl]] +uuid = "8f399da3-3557-5675-b5ff-fb832c97cbdb" +version = "1.11.0" + +[[deps.LinearAlgebra]] +deps = ["Libdl", "OpenBLAS_jll", "libblastrampoline_jll"] +uuid = "37e2e46d-f89d-539d-b4ee-838fcccc9c8e" +version = "1.11.0" + +[[deps.Logging]] +uuid = "56ddb016-857b-54e1-b83d-db4d58db5568" +version = "1.11.0" + +[[deps.OpenBLAS_jll]] +deps = ["Artifacts", "CompilerSupportLibraries_jll", "Libdl"] +uuid = "4536629a-c528-5b80-bd46-f80d51c5b363" +version = "0.3.27+1" + +[[deps.Parsers]] +deps = ["Dates", "PrecompileTools", "UUIDs"] +git-tree-sha1 = "5d5e0a78e971354b1c7bff0655d11fdc1b0e12c8" +uuid = "69de0a69-1ddd-5017-9359-2bf0b02dc9f0" +version = "2.8.4" + +[[deps.PrecompileTools]] +deps = ["Preferences"] +git-tree-sha1 = "5aa36f7049a63a1528fe8f7c3f2113413ffd4e1f" +uuid = "aea7be01-6a6a-4083-8856-8a6e6704d82a" +version = "1.2.1" + +[[deps.Preferences]] +deps = ["TOML"] +git-tree-sha1 = "8b770b60760d4451834fe79dd483e318eee709c4" +uuid = "21216c6a-2e73-6563-6e65-726566657250" +version = "1.5.2" + +[[deps.Printf]] +deps = ["Unicode"] +uuid = "de0858da-6303-5e67-8744-51eddeeeb8d7" +version = "1.11.0" + +[[deps.Profile]] +uuid = "9abbd945-dff8-562f-b5e8-e1ebf5ef1b79" +version = "1.11.0" + +[[deps.Random]] +deps = ["SHA"] +uuid = "9a3f8284-a2c9-5f02-9a11-845980a1fd5c" +version = "1.11.0" + +[[deps.SHA]] +uuid = "ea8e919c-243c-51af-8825-aaa63cd721ce" +version = "0.7.0" + +[[deps.Statistics]] +deps = ["LinearAlgebra"] +git-tree-sha1 = "ae3bb1eb3bba077cd276bc5cfc337cc65c3075c0" +uuid = "10745b16-79ce-11e8-11f9-7d13ad32a3b2" +version = "1.11.1" + + [deps.Statistics.extensions] + SparseArraysExt = ["SparseArrays"] + + [deps.Statistics.weakdeps] + SparseArrays = "2f01184e-e22b-5df5-ae63-d93ebab69eaf" + +[[deps.StructUtils]] +deps = ["Dates", "UUIDs"] +git-tree-sha1 = "dd974aefe288ef2898733aecf40858dc86742d74" +uuid = "ec057cc2-7a8d-4b58-b3b3-92acb9f63b42" +version = "2.8.1" + + [deps.StructUtils.extensions] + StructUtilsMeasurementsExt = ["Measurements"] + StructUtilsStaticArraysCoreExt = ["StaticArraysCore"] + StructUtilsTablesExt = ["Tables"] + + [deps.StructUtils.weakdeps] + Measurements = "eff96d63-e80a-5855-80a2-b1b0885c5ab7" + StaticArraysCore = "1e83bf80-4336-4d27-bf5d-d5a4f845583c" + Tables = "bd369af6-aec1-5ad0-b16a-f7cc5008161c" + +[[deps.TOML]] +deps = ["Dates"] +uuid = "fa267f1f-6049-4f14-aa54-33bafae1ed76" +version = "1.0.3" + +[[deps.UUIDs]] +deps = ["Random", "SHA"] +uuid = "cf7118a7-6976-5b1a-9a39-7adc72f591a4" +version = "1.11.0" + +[[deps.Unicode]] +uuid = "4ec0a83e-493e-50e2-b9ac-8f72acf5a8f5" +version = "1.11.0" + +[[deps.libblastrampoline_jll]] +deps = ["Artifacts", "Libdl"] +uuid = "8e850b90-86db-534c-a0d3-1478176c7d93" +version = "5.11.0+0" diff --git a/experiments/julia_safe_heuristic/Project.toml b/experiments/julia_safe_heuristic/Project.toml new file mode 100644 index 0000000..fd6d794 --- /dev/null +++ b/experiments/julia_safe_heuristic/Project.toml @@ -0,0 +1,7 @@ +name = "JuliaSafeHeuristic" +uuid = "e0293fdf-6d99-4c3d-bd99-7fd1bdcf64f1" +version = "0.1.0" + +[deps] +BenchmarkTools = "6e4b80f9-dd63-53aa-95a3-0cdb28fa8baf" +JSON = "682c06a0-de6a-54ab-a142-c8b1cf79cde6" diff --git a/experiments/julia_safe_heuristic/README.md b/experiments/julia_safe_heuristic/README.md new file mode 100644 index 0000000..4bbec46 --- /dev/null +++ b/experiments/julia_safe_heuristic/README.md @@ -0,0 +1,34 @@ +# Julia Safe-Heuristic Experiment + +This experiment keeps Julia out of the production bot registry. Python exports +`GameState.to_snapshot()` parity cases with expected Python actions, then Julia +reads those snapshots and computes actions in-process. + +Generate a corpus: + +```bash +uv run python scripts/export_safe_heuristic_snapshots.py \ + --output runs/tmp/safe_heuristic_snapshots.jsonl \ + --seeds 50 +``` + +Run parity: + +```bash +tools/julia/current/bin/julia --project=experiments/julia_safe_heuristic -e 'using Pkg; Pkg.instantiate()' +tools/julia/current/bin/julia --project=experiments/julia_safe_heuristic \ + experiments/julia_safe_heuristic/test/parity.jl \ + runs/tmp/safe_heuristic_snapshots.jsonl +``` + +Run throughput benchmark: + +```bash +tools/julia/current/bin/julia --project=experiments/julia_safe_heuristic \ + experiments/julia_safe_heuristic/bench/bench_snapshots.jl \ + runs/tmp/safe_heuristic_snapshots.jsonl +``` + +The benchmark measures pure Julia action selection over already-loaded JSON +records. It does not measure Python-to-Julia per-turn FFI, which is deliberately +not part of this experiment. diff --git a/experiments/julia_safe_heuristic/bench/bench_snapshots.jl b/experiments/julia_safe_heuristic/bench/bench_snapshots.jl new file mode 100644 index 0000000..39f9096 --- /dev/null +++ b/experiments/julia_safe_heuristic/bench/bench_snapshots.jl @@ -0,0 +1,27 @@ +using BenchmarkTools +using JSON + +include("../src/SafeHeuristic.jl") +using .SafeHeuristic + +function main() + if length(ARGS) != 1 + println(stderr, "usage: julia --project=experiments/julia_safe_heuristic experiments/julia_safe_heuristic/bench/bench_snapshots.jl ") + exit(2) + end + records = [JSON.parse(line) for line in eachline(ARGS[1]) if !isempty(strip(line))] + println("loaded $(length(records)) snapshots") + for record in records[1:min(end, 100)] + safe_heuristic_action(record) + end + result = @benchmark begin + total = 0 + for record in $records + total += safe_heuristic_action(record) + end + total + end + display(result) +end + +main() diff --git a/experiments/julia_safe_heuristic/bench_threaded.jl b/experiments/julia_safe_heuristic/bench_threaded.jl new file mode 100644 index 0000000..ac65408 --- /dev/null +++ b/experiments/julia_safe_heuristic/bench_threaded.jl @@ -0,0 +1,138 @@ +import Pkg + +const PROJECT_DIR = @__DIR__ +Pkg.activate(PROJECT_DIR; io=devnull) + +using JSON +using Printf +using Statistics +using Base.Threads + +include(joinpath(PROJECT_DIR, "src", "SafeHeuristic.jl")) +using .SafeHeuristic + +const DEFAULT_SNAPSHOT_PATH = joinpath( + dirname(dirname(PROJECT_DIR)), + "runs", + "tmp", + "safe_heuristic_snapshots_3seeds.jsonl", +) +const RUNS_PER_CASE = 5 + +function load_records(path::AbstractString) + return [JSON.parse(line) for line in eachline(path) if !isempty(strip(line))] +end + +function chunk_ranges(length::Int, chunks::Int)::Vector{UnitRange{Int}} + ranges = Vector{UnitRange{Int}}(undef, chunks) + base = length ÷ chunks + extra = length % chunks + start = 1 + for chunk in 1:chunks + width = base + (chunk <= extra ? 1 : 0) + stop = start + width - 1 + ranges[chunk] = start:stop + start = stop + 1 + end + return ranges +end + +function run_sequential(records)::Vector{Int} + actions = Vector{Int}(undef, length(records)) + @inbounds for idx in eachindex(records) + actions[idx] = safe_heuristic_action(records[idx]) + end + return actions +end + +function run_threaded(records, chunks::Int)::Vector{Int} + actions = Vector{Int}(undef, length(records)) + ranges = chunk_ranges(length(records), chunks) + @threads for chunk in eachindex(ranges) + @inbounds for idx in ranges[chunk] + actions[idx] = safe_heuristic_action(records[idx]) + end + end + return actions +end + +function run_case(records, chunks::Int)::Vector{Int} + if chunks == 1 + return run_sequential(records) + end + return run_threaded(records, chunks) +end + +function measure_case(records, chunks::Int, baseline::Vector{Int})::Float64 + warmup_actions = run_case(records, chunks) + if warmup_actions != baseline + error("action parity failed during warmup for $(chunks) thread case") + end + + times = Vector{Float64}(undef, RUNS_PER_CASE) + for run in 1:RUNS_PER_CASE + actions = Vector{Int}() + elapsed = @elapsed begin + actions = run_case(records, chunks) + end + if actions != baseline + error("action parity failed on run $(run) for $(chunks) thread case") + end + times[run] = elapsed * 1000.0 + end + return median(times) +end + +function scaling_label(efficiency::Float64)::String + if efficiency >= 0.80 + return "near-linear" + elseif efficiency >= 0.50 + return "partial scale" + end + return "contention" +end + +function main() + if length(ARGS) > 1 + println(stderr, "usage: julia --threads=8 experiments/julia_safe_heuristic/bench_threaded.jl [snapshots.jsonl]") + exit(2) + end + + path = length(ARGS) == 1 ? ARGS[1] : DEFAULT_SNAPSHOT_PATH + if !isfile(path) + println(stderr, "snapshot file not found: $path") + exit(2) + end + + records = load_records(path) + available = Threads.nthreads() + grid = [threads for threads in (1, 2, 4, 8) if threads <= available && threads <= length(records)] + if isempty(grid) + println(stderr, "no thread counts are available") + exit(2) + end + + baseline = run_sequential(records) + results = Dict{Int,Float64}() + for threads in grid + results[threads] = measure_case(records, threads, baseline) + end + + one_thread_ms = results[1] + println("threads total_ms μs/call speedup vs 1T efficiency") + for threads in grid + total_ms = results[threads] + us_per_call = total_ms * 1000.0 / length(records) + speedup = one_thread_ms / total_ms + efficiency = speedup / threads + @printf("%-7d %8.1f %8.1f %13.2f× %10.0f%%\n", threads, total_ms, us_per_call, speedup, efficiency * 100.0) + end + + if 8 in grid + speedup = one_thread_ms / results[8] + efficiency = speedup / 8 + @printf("8T efficiency %.0f%% — %s\n", efficiency * 100.0, scaling_label(efficiency)) + end +end + +main() diff --git a/experiments/julia_safe_heuristic/src/JuliaSafeHeuristic.jl b/experiments/julia_safe_heuristic/src/JuliaSafeHeuristic.jl new file mode 100644 index 0000000..7c8e6c4 --- /dev/null +++ b/experiments/julia_safe_heuristic/src/JuliaSafeHeuristic.jl @@ -0,0 +1,9 @@ +module JuliaSafeHeuristic + +include("SafeHeuristic.jl") + +using .SafeHeuristic: SafeHeuristicParams, safe_heuristic_action + +export SafeHeuristicParams, safe_heuristic_action + +end diff --git a/experiments/julia_safe_heuristic/src/SafeHeuristic.jl b/experiments/julia_safe_heuristic/src/SafeHeuristic.jl new file mode 100644 index 0000000..23ce623 --- /dev/null +++ b/experiments/julia_safe_heuristic/src/SafeHeuristic.jl @@ -0,0 +1,642 @@ +module SafeHeuristic + +export SafeHeuristicParams, safe_heuristic_action + +const PLAY_OR_DISCARD_ACTIONS_PER_SLOT = 2 +const DRAW_FROM_DECK_ACTION = 0 + +Base.@kwdef struct SafeHeuristicParams + open_target_ratio::Float64 = 0.50 + open_min_card_ratio::Float64 = 0.40 + handshake_target_multiplier::Float64 = 1.15 + handshake_min_card_ratio::Float64 = 0.34 + late_deck_ratio::Float64 = 0.20 + mid_deck_ratio::Float64 = 0.35 + commitment_weight::Float64 = 1.00 + gift_penalty_weight::Float64 = 1.00 + discard_safety_bonus::Float64 = 6.00 + unusable_discard_bonus::Float64 = 20.00 + deck_draw_early_value::Float64 = 2.00 + deck_draw_mid_value::Float64 = 1.00 + deck_draw_late_value::Float64 = -1.00 + deny_opponent_weight::Float64 = 0.40 + winning_deck_bonus::Float64 = 0.75 + losing_deck_penalty::Float64 = 1.25 + losing_visible_draw_bonus::Float64 = 1.50 + speculative_visible_draw_bonus::Float64 = 1.50 + dead_visible_draw_penalty::Float64 = 2.00 + unopened_draw_penalty_three_open::Float64 = 10.00 + unopened_draw_penalty_four_open::Float64 = 20.00 + strong_deny_threshold::Float64 = 10.00 + late_open_block_ratio::Float64 = 0.20 + low_card_sequence_bonus::Float64 = 5.00 + started_expedition_play_bonus::Float64 = 4.00 + started_expedition_followup_bonus::Float64 = 3.00 +end + +struct DerivedHeuristicConfig + middle_rank::Int + max_color_sum::Int + break_even_sum::Int + open_target_sum::Float64 + min_open_cards::Int + min_handshake_numeric_cards::Int + late_deck_threshold::Int + mid_deck_threshold::Int + late_open_block_threshold::Int + bonus_possible::Bool + max_expedition_cards::Int +end + +play_action(slot0::Int)::Int = PLAY_OR_DISCARD_ACTIONS_PER_SLOT * slot0 +discard_action(slot0::Int)::Int = PLAY_OR_DISCARD_ACTIONS_PER_SLOT * slot0 + 1 +draw_from_discard_action(color::Int)::Int = 1 + color + +card_color(card)::Int = Int(card["color"]) +card_rank(card)::Int = Int(card["rank"]) +is_handshake(card)::Bool = card_rank(card) == 0 +num(config, card)::Int = is_handshake(card) ? 0 : Int(config["min_rank"]) + card_rank(card) - 1 +deck_size(config)::Int = Int(config["n_colors"]) * (Int(config["n_ranks"]) + Int(config["n_handshakes"])) +max_rank(config)::Int = Int(config["min_rank"]) + Int(config["n_ranks"]) - 1 +card_action_size(config)::Int = 2 * Int(config["hand_size"]) +draw_action_size(config)::Int = 1 + Int(config["n_colors"]) + +function params_for_variant(name::AbstractString)::SafeHeuristicParams + if name == "loose" + return SafeHeuristicParams( + open_target_ratio=0.42, + open_min_card_ratio=0.30, + handshake_target_multiplier=1.00, + handshake_min_card_ratio=0.25, + late_open_block_ratio=0.12, + ) + elseif name == "strict" + return SafeHeuristicParams( + open_target_ratio=0.62, + open_min_card_ratio=0.50, + handshake_target_multiplier=1.35, + handshake_min_card_ratio=0.45, + late_open_block_ratio=0.30, + ) + end + return SafeHeuristicParams() +end + +function derive(config, params::SafeHeuristicParams)::DerivedHeuristicConfig + n_ranks = Int(config["n_ranks"]) + min_rank = Int(config["min_rank"]) + hand_size = Int(config["hand_size"]) + total_deck_size = deck_size(config) + n_handshakes = Int(config["n_handshakes"]) + bonus_threshold = Int(config["bonus_threshold"]) + max_color_sum = sum(min_rank + rank - 1 for rank in 1:n_ranks) + break_even_sum = -Int(config["expedition_penalty"]) + open_target_sum = min(0.8 * Float64(break_even_sum), params.open_target_ratio * Float64(max_color_sum)) + min_open_cards = max(1, min(hand_size, round(Int, hand_size * params.open_min_card_ratio))) + min_handshake_numeric_cards = max( + 1, + min(hand_size, round(Int, hand_size * params.handshake_min_card_ratio)), + ) + late_deck_threshold = max(1, round(Int, total_deck_size * params.late_deck_ratio)) + mid_deck_threshold = max(late_deck_threshold + 1, round(Int, total_deck_size * params.mid_deck_ratio)) + late_open_block_threshold = max(1, round(Int, total_deck_size * params.late_open_block_ratio)) + max_expedition_cards = n_handshakes + n_ranks + return DerivedHeuristicConfig( + (n_ranks + 1) ÷ 2, + max_color_sum, + break_even_sum, + open_target_sum, + min_open_cards, + min_handshake_numeric_cards, + late_deck_threshold, + mid_deck_threshold, + late_open_block_threshold, + max_expedition_cards >= bonus_threshold, + max_expedition_cards, + ) +end + +function hand(state, player0::Int) + return state["hands"][player0 + 1] +end + +function expeditions(state, player0::Int) + return state["expeditions"][player0 + 1] +end + +function color_expedition(state, player0::Int, color0::Int) + return state["expeditions"][player0 + 1][color0 + 1] +end + +function discard_pile(state, color0::Int) + return state["discards"][color0 + 1] +end + +function last_numeric_rank(state, player0::Int, color0::Int)::Int + last = 0 + for card in color_expedition(state, player0, color0) + rank = card_rank(card) + if rank > last + last = rank + end + end + return last +end + +has_numeric(state, player0::Int, color0::Int)::Bool = last_numeric_rank(state, player0, color0) > 0 + +function can_play_card(state, player0::Int, card)::Bool + rank = card_rank(card) + color = card_color(card) + config = state["config"] + if color < 0 || color >= Int(config["n_colors"]) || rank < 0 || rank > Int(config["n_ranks"]) + return false + end + last_rank = last_numeric_rank(state, player0, color) + if rank == 0 + return last_rank == 0 + end + return rank > last_rank +end + +function score_from_summary(config, len::Int, handshakes::Int, numeric_sum::Int)::Int + if len == 0 + return 0 + end + score = (numeric_sum + Int(config["expedition_penalty"])) * (handshakes + 1) + if len >= Int(config["bonus_threshold"]) + score += Int(config["bonus_amount"]) + end + return score +end + +function total_score(state, player0::Int)::Int + config = state["config"] + total = 0 + for expedition in expeditions(state, player0) + handshakes = 0 + numeric_sum = 0 + for card in expedition + if is_handshake(card) + handshakes += 1 + else + numeric_sum += num(config, card) + end + end + total += score_from_summary(config, length(expedition), handshakes, numeric_sum) + end + return total +end + +score_diff(state, player0::Int)::Int = total_score(state, player0) - total_score(state, 1 - player0) + +function legal_card_mask(state)::Vector{Bool} + h = hand(state, Int(state["current_player"])) + mask = fill(false, card_action_size(state["config"])) + for (idx, card) in enumerate(h) + slot0 = idx - 1 + if can_play_card(state, Int(state["current_player"]), card) + mask[play_action(slot0) + 1] = true + end + mask[discard_action(slot0) + 1] = true + end + return mask +end + +function legal_draw_mask(state)::Vector{Bool} + config = state["config"] + mask = fill(false, draw_action_size(config)) + pending = get(state, "pending_discarded_color", nothing) + if length(state["deck"]) > 0 + mask[DRAW_FROM_DECK_ACTION + 1] = true + end + for color in 0:(Int(config["n_colors"]) - 1) + if pending !== nothing && color == Int(pending) + continue + end + if !isempty(discard_pile(state, color)) + mask[draw_from_discard_action(color) + 1] = true + end + end + return mask +end + +function first_legal(mask)::Int + for (idx, value) in enumerate(mask) + if value + return idx - 1 + end + end + return 0 +end + +function best_pair(candidates) + isempty(candidates) && return nothing + best_value, best_action = candidates[1] + for (value, action) in candidates[2:end] + if value > best_value || (value == best_value && action > best_action) + best_value = value + best_action = action + end + end + return best_action +end + +function best_triple(candidates) + isempty(candidates) && return nothing + best = candidates[1] + for item in candidates[2:end] + if item[1] > best[1] || (item[1] == best[1] && (item[2] > best[2] || (item[2] == best[2] && item[3] > best[3]))) + best = item + end + end + return best[3] +end + +function late_penalty(derived::DerivedHeuristicConfig, deck_left::Int)::Float64 + if deck_left <= derived.late_deck_threshold + return 15.0 + elseif deck_left <= derived.mid_deck_threshold + return 8.0 + end + return 0.0 +end + +function new_color_open_penalty(opened_colors::Int)::Float64 + opened_colors <= 1 && return 0.0 + opened_colors == 2 && return 6.0 + opened_colors == 3 && return 14.0 + return 28.0 +end + +function bonus_potential(state, player0::Int, color0::Int, extra_cards::Int, derived::DerivedHeuristicConfig; committed_cards::Int=0, exclude_slot::Union{Nothing,Int}=nothing)::Float64 + !derived.bonus_possible && return 0.0 + config = state["config"] + expedition_len = length(color_expedition(state, player0, color0)) + committed_cards + need = Int(config["bonus_threshold"]) - expedition_len + need <= 0 && return Float64(config["bonus_amount"]) + playable_count = 0 + for (idx, card) in enumerate(hand(state, player0)) + if exclude_slot !== nothing && idx == exclude_slot + continue + end + if card_color(card) == color0 && can_play_card(state, player0, card) + playable_count += 1 + end + end + if playable_count + extra_cards >= need + return 0.4 * Float64(config["bonus_amount"]) + end + return 0.0 +end + +function opening_plan_value(state, params::SafeHeuristicParams, player0::Int, color0::Int, opening_card, derived::DerivedHeuristicConfig, deck_left::Int)::Float64 + config = state["config"] + numbers = [card for card in hand(state, player0) if card_color(card) == color0 && !is_handshake(card) && card_rank(card) >= card_rank(opening_card)] + handshakes = [card for card in hand(state, player0) if card_color(card) == color0 && is_handshake(card)] + opened_colors = count(expedition -> !isempty(expedition), expeditions(state, player0)) + number_sum = sum(num(config, card) for card in numbers; init=0) + high_cards = [card for card in numbers if card_rank(card) >= derived.middle_rank] + high_count = length(high_cards) + opening_value = num(config, opening_card) + penalty = new_color_open_penalty(opened_colors) + strong_open = length(numbers) >= derived.min_open_cards && number_sum >= derived.open_target_sum && (!isempty(high_cards) || number_sum >= 0.85 * derived.max_color_sum) + speculative_open = opened_colors <= 2 && length(numbers) >= 2 && number_sum >= 0.65 * derived.open_target_sum && !isempty(high_cards) + single_late_open = deck_left <= derived.mid_deck_threshold && length(numbers) >= 1 && opening_value >= 8 + exceptional_open = length(numbers) >= derived.min_open_cards + 1 && number_sum >= max(Float64(derived.break_even_sum), derived.open_target_sum * 1.4) && high_count >= 2 && deck_left > derived.mid_deck_threshold + if opened_colors == 3 + speculative_open = false + end + if opened_colors >= 4 + strong_open = false + speculative_open = false + single_late_open = false + end + strong_open && return 6.0 + 0.25 * number_sum + 0.8 * length(numbers) + 0.5 * length(handshakes) - penalty + speculative_open && return 3.0 + 0.18 * number_sum + 0.7 * length(numbers) + 0.4 * length(handshakes) - penalty + opened_colors == 3 && return 0.0 + exceptional_open && return 10.0 + 0.3 * number_sum + 1.0 * length(numbers) + 0.7 * high_count - penalty + single_late_open && return 1.5 + 0.2 * opening_value - penalty + return 0.0 +end + +function color_commitment(state, params::SafeHeuristicParams, player0::Int, color0::Int, derived::DerivedHeuristicConfig)::Float64 + config = state["config"] + expedition = color_expedition(state, player0, color0) + value = isempty(expedition) ? 0.0 : 5.0 + for card in expedition + value += is_handshake(card) ? 2.0 : 0.25 * num(config, card) + end + playable_cards = [card for card in hand(state, player0) if card_color(card) == color0 && can_play_card(state, player0, card)] + playable_numbers = [card for card in playable_cards if !is_handshake(card)] + playable_handshakes = [card for card in playable_cards if is_handshake(card)] + value += 1.2 * length(playable_numbers) + value += 1.5 * length(playable_handshakes) + value += 0.15 * sum(num(config, card) for card in playable_numbers; init=0) + value += 0.05 * bonus_potential(state, player0, color0, 0, derived) + return value +end + +function public_color_commitment_for_opponent(state, opponent0::Int, color0::Int, derived::DerivedHeuristicConfig)::Float64 + config = state["config"] + expedition = color_expedition(state, opponent0, color0) + discard = discard_pile(state, color0) + value = isempty(expedition) ? 0.0 : 5.0 + handshake_count = count(is_handshake, expedition) + value += 2.0 * handshake_count + numeric_cards = [card for card in expedition if !is_handshake(card)] + value += 0.25 * sum(num(config, card) for card in numeric_cards; init=0) + if !isempty(numeric_cards) + value += 0.4 * num(config, numeric_cards[end]) + end + if !isempty(discard) + top_card = discard[end] + if can_play_card(state, opponent0, top_card) + value += is_handshake(top_card) ? 1.5 : 1.0 + 0.1 * num(config, top_card) + end + end + if derived.bonus_possible + expedition_len = length(expedition) + if expedition_len + 1 >= Int(config["bonus_threshold"]) + value += 0.2 * Float64(config["bonus_amount"]) + end + end + return value +end + +function card_value_for_opponent(state, params::SafeHeuristicParams, opponent0::Int, card, derived::DerivedHeuristicConfig)::Float64 + !can_play_card(state, opponent0, card) && return 0.0 + interest = public_color_commitment_for_opponent(state, opponent0, card_color(card), derived) + is_handshake(card) && return 8.0 + 1.5 * interest + numeric_value = num(state["config"], card) + return numeric_value * (0.4 + 0.25 * interest) +end + +function card_value_for_me(state, params::SafeHeuristicParams, player0::Int, card, derived::DerivedHeuristicConfig)::Float64 + !can_play_card(state, player0, card) && return 0.0 + commitment = color_commitment(state, params, player0, card_color(card), derived) + if is_handshake(card) + return 7.0 + 1.2 * commitment + end + numeric_value = num(state["config"], card) + value = 0.8 * numeric_value + params.commitment_weight * commitment + if !isempty(color_expedition(state, player0, card_color(card))) + value += params.started_expedition_play_bonus + params.started_expedition_followup_bonus + end + if commitment >= 6.0 && card_rank(card) <= derived.middle_rank + value += params.low_card_sequence_bonus + end + return value +end + +function started_expedition_play_value(state, params::SafeHeuristicParams, player0::Int, card, derived::DerivedHeuristicConfig, deck_left::Int)::Float64 + config = state["config"] + color = card_color(card) + expedition = color_expedition(state, player0, color) + numeric_value = num(config, card) + current_sum = sum(num(config, played) for played in expedition if !is_handshake(played); init=0) + followups = [followup for followup in hand(state, player0) if followup !== card && card_color(followup) == color && !is_handshake(followup) && card_rank(followup) > card_rank(card)] + projected_sum = current_sum + numeric_value + sum(num(config, followup) for followup in followups; init=0) + value = params.started_expedition_play_bonus + params.started_expedition_followup_bonus + value += Float64(max_rank(config) + 1 - numeric_value) + if deck_left <= derived.late_deck_threshold + value += 2.0 * numeric_value + elseif deck_left <= derived.mid_deck_threshold + value += 0.8 * numeric_value + end + if projected_sum < derived.open_target_sum + value -= 6.0 + end + handshakes = count(is_handshake, expedition) + value += 3.0 * handshakes + value += bonus_potential(state, player0, color, 0, derived, committed_cards=1) + return value +end + +function visible_number_can_help_open(state, params::SafeHeuristicParams, player0::Int, card, derived::DerivedHeuristicConfig)::Bool + numbers = [other for other in hand(state, player0) if card_color(other) == card_color(card) && !is_handshake(other) && card_rank(other) >= card_rank(card)] + push!(numbers, card) + length(numbers) < derived.min_open_cards && return false + number_sum = sum(num(state["config"], other) for other in numbers; init=0) + number_sum < derived.open_target_sum && return false + return any(card_rank(other) >= derived.middle_rank for other in numbers) +end + +function visible_open_support_value(state, params::SafeHeuristicParams, player0::Int, card, derived::DerivedHeuristicConfig)::Float64 + color = card_color(card) + opened_colors = count(expedition -> !isempty(expedition), expeditions(state, player0)) + same_color_numbers = [other for other in hand(state, player0) if card_color(other) == color && !is_handshake(other)] + same_color_handshakes = [other for other in hand(state, player0) if card_color(other) == color && is_handshake(other)] + future_numbers = [other for other in same_color_numbers if other !== card && card_rank(other) >= card_rank(card)] + value = 0.8 * length(future_numbers) + 1.0 * length(same_color_handshakes) + if card_rank(card) <= derived.middle_rank + value += params.speculative_visible_draw_bonus + end + if visible_number_can_help_open(state, params, player0, card, derived) + value += 4.0 + elseif opened_colors <= 2 && (!isempty(future_numbers) || !isempty(same_color_handshakes)) + value += params.speculative_visible_draw_bonus + end + if opened_colors <= 2 + value += 0.25 * opening_plan_value(state, params, player0, color, card, derived, length(state["deck"])) + elseif opened_colors == 3 + value += 0.1 * max(0.0, opening_plan_value(state, params, player0, color, card, derived, length(state["deck"]))) + end + return value +end + +function visible_draw_value(state, params::SafeHeuristicParams, player0::Int, card, derived::DerivedHeuristicConfig)::Float64 + color = card_color(card) + opponent0 = 1 - player0 + opened_colors = count(expedition -> !isempty(expedition), expeditions(state, player0)) + is_unopened_color = isempty(color_expedition(state, player0, color)) + commitment = color_commitment(state, params, player0, color, derived) + opponent_value = card_value_for_opponent(state, params, opponent0, card, derived) + diff = score_diff(state, player0) + value = params.deny_opponent_weight * opponent_value + diff <= 0 && (value += params.losing_visible_draw_bonus) + exceptional_support = false + if is_unopened_color + opened_colors >= 4 ? (value -= params.unopened_draw_penalty_four_open) : opened_colors >= 3 && (value -= params.unopened_draw_penalty_three_open) + end + if is_handshake(card) + has_numeric(state, player0, color) && return value - params.dead_visible_draw_penalty + if isempty(color_expedition(state, player0, color)) + playable_numbers = [other for other in hand(state, player0) if card_color(other) == color && !is_handshake(other) && can_play_card(state, player0, other)] + number_sum = sum(num(state["config"], other) for other in playable_numbers; init=0) + required_sum = derived.open_target_sum * params.handshake_target_multiplier + if length(playable_numbers) < derived.min_handshake_numeric_cards || number_sum < required_sum + support = visible_open_support_value(state, params, player0, card, derived) + exceptional_support = support >= 6.0 + if is_unopened_color && opened_colors >= 4 && !exceptional_support && opponent_value < params.strong_deny_threshold && diff > -15 + return -8.0 + end + return value + support - 0.5 + end + end + return value + 6.0 + commitment + end + immediate_playable = can_play_card(state, player0, card) + if immediate_playable + value += Float64(num(state["config"], card)) + value += 0.7 * commitment + if !isempty(color_expedition(state, player0, color)) + value += 5.0 + else + support = visible_open_support_value(state, params, player0, card, derived) + exceptional_support = support >= 6.0 + value += support + end + else + value -= params.dead_visible_draw_penalty + if isempty(color_expedition(state, player0, color)) + support = visible_open_support_value(state, params, player0, card, derived) + exceptional_support = support >= 6.0 + value += support + end + end + if is_unopened_color && opened_colors >= 4 && !exceptional_support && opponent_value < params.strong_deny_threshold && diff > -15 + return -8.0 + end + value += bonus_potential(state, player0, color, 1, derived) + return value +end + +function best_handshake_play(state, params::SafeHeuristicParams, player0::Int, legal, derived::DerivedHeuristicConfig, deck_left::Int) + Int(state["config"]["n_handshakes"]) <= 0 && return nothing + candidates = Tuple{Float64,Int}[] + for (idx, card) in enumerate(hand(state, player0)) + slot0 = idx - 1 + action = play_action(slot0) + (!legal[action + 1] || !is_handshake(card)) && continue + color = card_color(card) + expedition = color_expedition(state, player0, color) + any(!is_handshake(played) for played in expedition) && continue + playable_numbers = [other for (other_idx, other) in enumerate(hand(state, player0)) if other_idx != idx && card_color(other) == color && !is_handshake(other) && can_play_card(state, player0, other)] + number_count = length(playable_numbers) + number_sum = sum(num(state["config"], other) for other in playable_numbers; init=0) + number_count < derived.min_handshake_numeric_cards && continue + required_sum = derived.open_target_sum * params.handshake_target_multiplier + number_sum < required_sum && continue + deck_left <= derived.late_open_block_threshold && continue + value = number_sum + 2.0 * number_count + value += bonus_potential(state, player0, color, 0, derived, committed_cards=1, exclude_slot=idx) + value -= late_penalty(derived, deck_left) + push!(candidates, (value, action)) + end + return best_pair(candidates) +end + +function best_number_play(state, params::SafeHeuristicParams, player0::Int, legal, derived::DerivedHeuristicConfig, deck_left::Int) + candidates = Tuple{Float64,Int}[] + for (idx, card) in enumerate(hand(state, player0)) + slot0 = idx - 1 + action = play_action(slot0) + (!legal[action + 1] || is_handshake(card)) && continue + color = card_color(card) + if !isempty(color_expedition(state, player0, color)) + push!(candidates, (started_expedition_play_value(state, params, player0, card, derived, deck_left), action)) + elseif deck_left > derived.late_open_block_threshold && opening_plan_value(state, params, player0, color, card, derived, deck_left) > 0.0 + numbers = [c for c in hand(state, player0) if card_color(c) == color && !is_handshake(c) && card_rank(c) >= card_rank(card)] + handshakes = [c for c in hand(state, player0) if card_color(c) == color && is_handshake(c)] + number_sum = sum(num(state["config"], c) for c in numbers; init=0) + value = number_sum + 2.0 * length(numbers) + 1.5 * length(handshakes) + value += opening_plan_value(state, params, player0, color, card, derived, deck_left) + value += Float64(state["config"]["expedition_penalty"]) + value += max(0.0, Float64(derived.middle_rank - card_rank(card))) + value += bonus_potential(state, player0, color, 0, derived, committed_cards=1, exclude_slot=idx) + value -= late_penalty(derived, deck_left) + push!(candidates, (value, action)) + end + end + return best_pair(candidates) +end + +function best_forced_open(state, params::SafeHeuristicParams, player0::Int, legal, derived::DerivedHeuristicConfig, deck_left::Int) + candidates = Tuple{Float64,Int}[] + for (idx, card) in enumerate(hand(state, player0)) + slot0 = idx - 1 + action = play_action(slot0) + (!legal[action + 1] || is_handshake(card)) && continue + !isempty(color_expedition(state, player0, card_color(card))) && continue + opening_value = opening_plan_value(state, params, player0, card_color(card), card, derived, deck_left) + color_numbers = [other for other in hand(state, player0) if card_color(other) == card_color(card) && !is_handshake(other)] + number_sum = sum(num(state["config"], other) for other in color_numbers; init=0) + if opening_value <= 0.0 && length(color_numbers) < 2 && number_sum < 0.5 * derived.open_target_sum && deck_left > derived.mid_deck_threshold + continue + end + forced_value = opening_value + 0.2 * number_sum + Float64(max_rank(state["config"]) + 1 - num(state["config"], card)) + push!(candidates, (forced_value, action)) + end + return best_pair(candidates) +end + +function best_discard(state, params::SafeHeuristicParams, player0::Int, legal, derived::DerivedHeuristicConfig) + candidates = Tuple{Float64,Int}[] + opponent0 = 1 - player0 + for (idx, card) in enumerate(hand(state, player0)) + slot0 = idx - 1 + action = discard_action(slot0) + !legal[action + 1] && continue + my_value = card_value_for_me(state, params, player0, card, derived) + opponent_value = card_value_for_opponent(state, params, opponent0, card, derived) + score = -my_value - params.gift_penalty_weight * opponent_value + !can_play_card(state, player0, card) && (score += params.unusable_discard_bonus) + !can_play_card(state, opponent0, card) && (score += params.discard_safety_bonus) + is_handshake(card) && can_play_card(state, player0, card) && (score -= 4.0) + push!(candidates, (score, action)) + end + return best_pair(candidates) +end + +function deck_draw_value(state, params::SafeHeuristicParams, derived::DerivedHeuristicConfig)::Float64 + deck_left = length(state["deck"]) + diff = score_diff(state, Int(state["current_player"])) + value = deck_left > derived.mid_deck_threshold ? params.deck_draw_early_value : deck_left > derived.late_deck_threshold ? params.deck_draw_mid_value : params.deck_draw_late_value + value += diff > 0 ? params.winning_deck_bonus : -params.losing_deck_penalty + return value +end + +function act_card(state, params::SafeHeuristicParams, derived::DerivedHeuristicConfig)::Int + player0 = Int(state["current_player"]) + legal = legal_card_mask(state) + deck_left = length(state["deck"]) + action = best_handshake_play(state, params, player0, legal, derived, deck_left) + action !== nothing && return action + action = best_number_play(state, params, player0, legal, derived, deck_left) + action !== nothing && return action + if all(isempty(expedition) for expedition in expeditions(state, player0)) + action = best_forced_open(state, params, player0, legal, derived, deck_left) + action !== nothing && return action + end + action = best_discard(state, params, player0, legal, derived) + action !== nothing && return action + return first_legal(legal) +end + +function act_draw(state, params::SafeHeuristicParams, derived::DerivedHeuristicConfig)::Int + legal = legal_draw_mask(state) + player0 = Int(state["current_player"]) + candidates = Tuple{Float64,Int,Int}[] + legal[DRAW_FROM_DECK_ACTION + 1] && push!(candidates, (deck_draw_value(state, params, derived), 1, DRAW_FROM_DECK_ACTION)) + for color in 0:(Int(state["config"]["n_colors"]) - 1) + action = draw_from_discard_action(color) + pile = discard_pile(state, color) + (!legal[action + 1] || isempty(pile)) && continue + card = pile[end] + value = visible_draw_value(state, params, player0, card, derived) + push!(candidates, (value, 0, action)) + end + action = best_triple(candidates) + action !== nothing && return action + return first_legal(legal) +end + +function safe_heuristic_action(record_or_state; variant::AbstractString="default")::Int + state = haskey(record_or_state, "state") ? record_or_state["state"] : record_or_state + params = params_for_variant(haskey(record_or_state, "variant") ? record_or_state["variant"] : variant) + derived = derive(state["config"], params) + return state["phase"] == "card" ? act_card(state, params, derived) : act_draw(state, params, derived) +end + +end diff --git a/experiments/julia_safe_heuristic/test/parity.jl b/experiments/julia_safe_heuristic/test/parity.jl new file mode 100644 index 0000000..e069c49 --- /dev/null +++ b/experiments/julia_safe_heuristic/test/parity.jl @@ -0,0 +1,27 @@ +using JSON + +include("../src/SafeHeuristic.jl") +using .SafeHeuristic + +function main() + if length(ARGS) != 1 + println(stderr, "usage: julia --project=experiments/julia_safe_heuristic experiments/julia_safe_heuristic/test/parity.jl ") + exit(2) + end + path = ARGS[1] + checked = 0 + for line in eachline(path) + isempty(strip(line)) && continue + record = JSON.parse(line) + actual = safe_heuristic_action(record) + expected = Int(record["expected_action"]) + if actual != expected + println(stderr, "mismatch config=$(record["config_name"]) variant=$(record["variant"]) seed=$(record["seed"]) turn=$(record["turn"]) phase=$(record["phase"]) player=$(record["current_player"]) expected=$expected actual=$actual") + exit(1) + end + checked += 1 + end + println("checked $checked snapshots") +end + +main() diff --git a/scripts/export_safe_heuristic_snapshots.py b/scripts/export_safe_heuristic_snapshots.py new file mode 100644 index 0000000..7a67eaf --- /dev/null +++ b/scripts/export_safe_heuristic_snapshots.py @@ -0,0 +1,105 @@ +from __future__ import annotations + +import argparse +import json +from pathlib import Path +from typing import Any + +from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig + +from coolrl_lost_cities.games.classic.bots.heuristic_py import SafeHeuristicBot +from coolrl_lost_cities.games.classic.bots.registry import ( + LOOSE_SAFE_HEURISTIC_PARAMS, + STRICT_SAFE_HEURISTIC_PARAMS, +) + +VARIANTS = { + "default": None, + "loose": LOOSE_SAFE_HEURISTIC_PARAMS, + "strict": STRICT_SAFE_HEURISTIC_PARAMS, +} + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Export safe-heuristic bot parity snapshots for external implementations." + ) + parser.add_argument("--output", required=True, help="JSONL output path.") + parser.add_argument("--seeds", type=int, default=50, help="Number of seeds per config.") + parser.add_argument("--max-steps", type=int, default=10_000) + return parser.parse_args() + + +def _configs() -> list[tuple[str, LostCitiesConfig]]: + return [ + ("classic", LostCitiesConfig()), + ("small", LostCitiesConfig(n_colors=2, n_ranks=8, hand_size=3)), + ( + "no-handshakes", + LostCitiesConfig(n_colors=3, n_ranks=5, n_handshakes=0, hand_size=5), + ), + ] + + +def _record( + *, + config_name: str, + variant_name: str, + seed: int, + turn: int, + state: GameState, + action: int, +) -> dict[str, Any]: + return { + "config_name": config_name, + "variant": variant_name, + "seed": seed, + "turn": turn, + "phase": state.phase, + "current_player": state.current_player, + "expected_action": action, + "state": state.to_snapshot(), + } + + +def main() -> None: + args = parse_args() + output = Path(args.output) + output.parent.mkdir(parents=True, exist_ok=True) + count = 0 + with output.open("w", encoding="utf-8") as handle: + for config_name, config in _configs(): + for variant_name, params in VARIANTS.items(): + for seed in range(args.seeds): + bot = SafeHeuristicBot(params) + state = GameState.new_game(config, seed=seed) + for turn in range(args.max_steps): + if state.terminal: + break + action = bot.act(state) + handle.write( + json.dumps( + _record( + config_name=config_name, + variant_name=variant_name, + seed=seed, + turn=turn, + state=state, + action=action, + ), + sort_keys=True, + ) + + "\n" + ) + count += 1 + state.apply_action(action) + else: + raise RuntimeError( + f"game did not terminate: config={config_name} " + f"variant={variant_name} seed={seed}" + ) + print(f"Wrote {count} snapshots to {output}") + + +if __name__ == "__main__": + main()