Compare commits

..
Author SHA1 Message Date
admin e49af0868f bad 2025-09-28 08:29:45 -07:00
admin 08bbbdd651 let's actually look at the code 2025-09-28 08:28:49 -07:00
admin 16eb808a73 what a mess 2025-09-28 08:28:49 -07:00
admin d2a1cc77e9 slow progress 2025-09-28 08:28:49 -07:00
admin 4a974b265a don't stop simulation on a player flip 2025-09-28 08:28:49 -07:00
admin f181c9241d set the current playerId correctly 2025-09-28 08:28:49 -07:00
admin 6cad64d348 investigate playerflips issue 2025-09-28 08:28:48 -07:00
admin 83780ca7bb hide info 2025-09-28 08:28:48 -07:00
admin 1a9b307922 logging fixes and use correct settings 2025-09-28 08:28:48 -07:00
admin c887b868fd player flips 2025-09-28 08:28:48 -07:00
admin 05bacede4f maybe kinda working 2025-09-28 08:28:48 -07:00
admin d3a1088b4c END_TURN not marked as terminal 2025-09-28 08:28:47 -07:00
admin 0a33ad4433 still a little drunk but END_TURN is scoring correctly 2025-09-28 08:28:47 -07:00
admin 3f6a8340d3 MCTSAI as a separate target 2025-09-28 08:27:49 -07:00
admin 811fcec712 MCTS integration complete 2025-09-28 08:25:46 -07:00
admin 3c2f5eeab1 store the decision tree 2025-09-28 08:25:45 -07:00
514 changed files with 15876 additions and 40403 deletions
+2 -5
View File
@@ -19,9 +19,9 @@ common --worker_sandboxing
common --local_test_jobs=64
common --jobs=64
common --cxxopt="--std=c++23"
common --cxxopt="--std=c++20"
common --cxxopt="-Wno-deprecated-non-prototype"
common --host_cxxopt="--std=c++23"
common --host_cxxopt="--std=c++20"
common --javacopt="-Xlint:-options"
@@ -29,9 +29,6 @@ common --javacopt="-Xlint:-options"
common --linkopt=-Wl
common:macos --linkopt=-Wl,-no_warn_duplicate_libraries
# Fix Xcode version caching issue - avoids need for `bazel clean --expunge` after Xcode updates
common:macos --repo_env=DEVELOPER_DIR=/Applications/Xcode.app/Contents/Developer
common --java_language_version=17
common --java_runtime_version=remotejdk_17
common --tool_java_language_version=17
-3
View File
@@ -1,3 +0,0 @@
CompileFlags:
Add:
- "-std=c++23"
+1 -45
View File
@@ -34,54 +34,10 @@ jobs:
with:
lfs: false
- name: Run tests
id: test
continue-on-error: true
run: bazel test --build_event_json_file=test.json //src/test/... //src/main/go/...
- name: Collect failed test logs
if: always()
run: |
# Remove any existing failed_test_logs directory and create fresh
rm -rf failed_test_logs
mkdir -p failed_test_logs
# Extract failed test targets from test.json and copy their logs
# The test.json is in JSONL format - one JSON object per line
# We look for lines with testResult that have a status other than PASSED
if [ -f test.json ]; then
grep '"testResult"' test.json | \
grep '"status"' | \
grep -v '"status":"PASSED"' | \
grep -o '"label":"[^"]*"' | \
cut -d'"' -f4 | \
sort -u | \
while read target; do
# Convert target like //src/test/cpp/...:test_name to path
log_path=$(echo "$target" | sed 's|^//||' | sed 's|:|/|')
if [ -f "bazel-testlogs/$log_path/test.log" ]; then
log_name=$(echo "$log_path" | tr '/' '_')
if cp "bazel-testlogs/$log_path/test.log" "failed_test_logs/${log_name}.log"; then
echo "Collected log for failed test: $target"
else
echo "Error: Failed to copy log for $target"
fi
fi
done
fi
# List what we collected
echo "Collected logs:"
ls -lh failed_test_logs/ 2>/dev/null || echo "No logs collected"
- name: Archive test results
if: always()
if: success() || failure()
uses: actions/upload-artifact@v4
with:
name: test.json
path: test.json
- name: Archive failed test logs
if: always()
uses: actions/upload-artifact@v4
with:
name: failed-test-logs
path: failed_test_logs/
if-no-files-found: ignore
- name: Fail if tests failed
if: steps.test.outcome == 'failure'
run: exit 1
-1
View File
@@ -37,4 +37,3 @@ scripts/refresh_name_layers/refresh_name_layers.zip
.metals
api_keys.txt
src/main/csharp/net/eagle0/clients/unity/eagle0/ProjectSettings/Packages/com.unity.dedicated-server/
+1 -2
View File
@@ -32,9 +32,8 @@ repos:
- id: gazelle
name: gazelle
language: system
entry: ./scripts/pre-commit-gazelle.sh
entry: bazel run //:gazelle
files: '(\.go|\.proto|BUILD\.bazel|BUILD|WORKSPACE|WORKSPACE\.bazel|\.bzl)$'
pass_filenames: false
- repo: local
hooks:
- id: update-action-result-types
+7 -93
View File
@@ -4,32 +4,26 @@ This file provides guidance to Claude Code (claude.ai/code) when working with co
## Project Overview
Eagle0 is a multi-language gaming system combining strategic turn-based gameplay (Eagle) with tactical hex-based
combat (Shardok). The system integrates LLM-based narrative generation and supports both human and AI players.
Eagle0 is a multi-language gaming system combining strategic turn-based gameplay (Eagle) with tactical hex-based combat (Shardok). The system integrates LLM-based narrative generation and supports both human and AI players.
## Architecture
**Three-Tier Game System:**
- **Unity Client (C#)**: Real-time strategy game client with integrated tactical combat UI
- **Eagle (Scala)**: Strategic layer managing turn-based gameplay, diplomacy, hero progression, and province control
- **Shardok (C++)**: Tactical layer handling real-time hex-based combat simulation with performance-critical battle
resolution
- **Shardok (C++)**: Tactical layer handling real-time hex-based combat simulation with performance-critical battle resolution
**Communication Flow:**
```
Unity Client ↔ Eagle (gRPC streaming) ↔ Shardok (internal gRPC)
```
**Key Entry Points:**
- `/src/main/csharp/net/eagle0/clients/unity/eagle0/` - Unity C# game client
- `/src/main/scala/net/eagle0/eagle/Main.scala` - Eagle strategic game server
- `/src/main/cpp/net/eagle0/shardok/shardok_server_main.cpp` - Shardok tactical server
**Protocol Buffer Architecture:**
- Extensive use of protobuf for type-safe communication
- Separate packages: `api/` (client-facing), `internal/` (server state), `views/` (client projections)
- Event sourcing pattern with immutable action history
@@ -37,7 +31,6 @@ Unity Client ↔ Eagle (gRPC streaming) ↔ Shardok (internal gRPC)
## Essential Commands
### Building
```bash
# Build Eagle server (Scala strategic layer)
bazel build //src/main/scala/net/eagle0/eagle:eagle_server_deploy.jar
@@ -56,7 +49,6 @@ bazel build //src/main/cpp/net/eagle0/shardok:shardok-server
```
### Running Services
```bash
# Eagle server (port 40032)
bazel run //src/main/scala/net/eagle0/eagle:eagle_server -- --eagle-grpc-port 40032
@@ -68,7 +60,6 @@ bazel run //src/main/cpp/net/eagle0/shardok:shardok-server --compilation_mode=op
```
### Testing
```bash
# Run all tests
bazel test //src/test/... //src/main/go/...
@@ -79,24 +70,12 @@ bazel test //src/test/cpp/... # C++ Shardok tests
```
### Code Generation
```bash
bazel run gazelle # Update Go build files
./scripts/updateActionResultTypes.sh # Update protocol buffer mappings
```
### Pre-Commit Checklist
**MANDATORY: Before running `git commit`, verify:**
1. **If you modified any BUILD.bazel file:** Run `bazel run gazelle` and stage any changes it makes
2. **If you modified C++ or C# files:** Run `clang-format -i` on the modified files
3. **If you modified Scala files:** scalafmt will run automatically via pre-commit hook
The pre-commit hook runs gazelle but only checks if it succeeds - it does NOT verify the BUILD files are in canonical format. The `gazelle_test` will fail if deps are not alphabetically sorted. **Always run gazelle manually after BUILD file changes.**
### Code Formatting
```bash
# ALWAYS run clang-format after making any C++ or C# code changes
clang-format -i <modified_files>
@@ -109,14 +88,13 @@ find . -name "*.cs" | xargs clang-format -i
```
### Static Analysis
```bash
# Run clang-tidy static analysis on C++ files
# Note: This may show some header include errors but will still analyze the main file
bazel run @llvm_toolchain//:clang-tidy -- --checks='readability-*,bugprone-*,clang-analyzer-*' <file_path> -- -I/Users/dancrosby/CodingProjects/github/eagle0 -std=c++23
bazel run @llvm_toolchain//:clang-tidy -- --checks='readability-*,bugprone-*,clang-analyzer-*' <file_path> -- -I/Users/dancrosby/CodingProjects/github/eagle0 -std=c++20
# Example for AI files:
bazel run @llvm_toolchain//:clang-tidy -- --checks='readability-*,bugprone-*,clang-analyzer-*' /Users/dancrosby/CodingProjects/github/eagle0/src/main/cpp/net/eagle0/shardok/ai/AIUnitScoreCalculator.cpp -- -I/Users/dancrosby/CodingProjects/github/eagle0 -std=c++23
bazel run @llvm_toolchain//:clang-tidy -- --checks='readability-*,bugprone-*,clang-analyzer-*' /Users/dancrosby/CodingProjects/github/eagle0/src/main/cpp/net/eagle0/shardok/ai/AIUnitScoreCalculator.cpp -- -I/Users/dancrosby/CodingProjects/github/eagle0 -std=c++20
```
## AI Algorithm Selection
@@ -124,17 +102,13 @@ bazel run @llvm_toolchain//:clang-tidy -- --checks='readability-*,bugprone-*,cla
Eagle0 supports two AI algorithms for tactical combat decision-making:
### Iterative Deepening AI (Default)
The original minimax-based AI with sophisticated randomness handling:
- **Advantages**: Proven, sophisticated randomness evaluation, comprehensive lookahead
- **Use cases**: Production builds, scenarios requiring precise evaluation
- **Performance**: Single-threaded, thorough evaluation
### Monte Carlo Tree Search AI (MCTS)
Modern MCTS-based AI with multithreading support:
- **Advantages**: Multithreaded, better performance on modern CPUs, anytime algorithm
- **Use cases**: Performance testing, scenarios requiring fast decisions
- **Performance**: Multithreaded, adaptive depth based on time budget
@@ -167,36 +141,30 @@ bazel test //src/test/cpp/net/eagle0/shardok/ai:ai_mcts_test # If available
Both implementations are compatible with all existing interfaces and produce the same `SearchResult` structure.
**Note**: Both implementations are documented in `src/main/cpp/net/eagle0/shardok/ai/AI_SCORING_SYSTEM.md`, including
recommendations for improving MCTS randomness handling.
**Note**: Both implementations are documented in `src/main/cpp/net/eagle0/shardok/ai/AI_SCORING_SYSTEM.md`, including recommendations for improving MCTS randomness handling.
The AI algorithm selection is made at runtime when creating ShardokAIClient instances, allowing different AI strategies
to be used for different players or game situations within the same server process.
The AI algorithm selection is made at runtime when creating ShardokAIClient instances, allowing different AI strategies to be used for different players or game situations within the same server process.
## Language-Specific Patterns
**Scala (Strategic Layer):**
- Use `EngineImpl.scala` for core game logic modifications
- Follow event sourcing pattern - all changes through immutable actions
- gRPC streaming for real-time client updates via `EagleServiceImpl.scala`
- LLM integration in `/common/llm_integration/` for narrative generation
**C++ (Tactical Layer):**
- Performance-critical combat in `ShardokEngine.hpp/.cpp`
- FlatBuffers for efficient serialization in `/flatbuffer/` directory
- AI systems in `/ai/` subdirectory with pluggable strategy selectors
- Extensive unit testing with Google Test framework
**Protocol Buffers:**
- Three-layer structure: `api/` (client), `internal/` (server), `views/` (projections)
- Use `shardok_internal_interface.proto` for Eagle-Shardok communication
- Maintain backward compatibility when modifying existing messages
**C# (Unity Client):**
- Located in `/src/main/csharp/net/eagle0/clients/unity/eagle0/`
- Uses Unity 6 (6000.0.32f1) with comprehensive protobuf integration (100+ .proto files)
- Key components: `EagleConnection.cs` (gRPC client), `EagleGameController.cs` (main game logic)
@@ -205,7 +173,6 @@ to be used for different players or game situations within the same server proce
- Seamless transition between strategic gameplay and hex-based tactical combat
**Go (Build Tools):**
- Build automation and code generation utilities
- AWS S3 integration for deployment artifacts
@@ -216,31 +183,6 @@ to be used for different players or game situations within the same server proce
- Map validation tests ensure game content integrity
- Use `GameSettings_test_utils.cpp` and `ShardokEngineBasedTestData.cpp` for C++ test helpers
### Scala Testing Patterns
**Use `inside()` instead of `asInstanceOf` for type matching in tests:**
Never use `asInstanceOf` in tests. Instead, use ScalaTest's `inside()` pattern for safe type matching:
```scala
// BAD - don't do this
val changedHero = result.changedHeroes.head.asInstanceOf[ChangedHeroC]
changedHero.heroId shouldBe 19
// GOOD - use inside() pattern
import org.scalatest.Inside.inside
inside(result.changedHeroes.head) { case changedHero: ChangedHeroC =>
changedHero.heroId shouldBe 19
changedHero.vigorChange shouldBe StatDelta(17.2)
}
```
The `inside()` pattern:
- Provides better error messages when the type doesn't match
- Is idiomatic ScalaTest
- Works with pattern matching for more complex assertions
## Performance Testing
When making performance-related changes to the AI or engine:
@@ -272,38 +214,10 @@ done
```
**Important notes:**
- Run tests multiple times (3-5) to account for performance variance
- Focus on commands evaluated at each depth rather than total commands
- Commands at different depths aren't directly comparable (depth 3 is more valuable than depth 2)
- **Always test performance changes** - what seems like an optimization may sometimes have unexpected overhead or
behavior changes.
## Troubleshooting Scala Build Errors
### MissingType Errors
When you see errors like:
```
dotty.tools.dotc.core.MissingType: Cannot resolve reference to type net.eagle0.eagle.internal.game_state.type.GameState
```
**This is NOT a Scala compiler crash.** This is a missing dependency in BUILD.bazel.
**How to fix:**
1. Identify the missing type from the error message (e.g., `game_state.GameState`)
2. Find the Bazel target that provides this type (e.g., `//src/main/protobuf/net/eagle0/eagle/internal:game_state_scala_proto`)
3. Add it to the `deps` of the failing target
4. If the type appears in a public method signature, also add it to `exports` so downstream targets can see it
**Common pattern:** When adding a method to a class that takes or returns a proto type, the proto dependency often needs to be added to both `deps` AND `exports`.
### Bazel Clean
**NEVER run `bazel clean` without asking first.** It rarely fixes actual issues and wastes significant rebuild time. The issues that seem like they need `bazel clean` are usually:
- Missing imports in Scala code
- Missing dependencies in BUILD.bazel
- Missing exports for types used in public signatures
- **Always test performance changes** - what seems like an optimization may sometimes have unexpected overhead or behavior changes.
## Game Content
+177
View File
@@ -0,0 +1,177 @@
# MCTS with Player Flips - Clean Design
## Core Principles
1. **Single Perspective**: All scoring is from our AI's perspective (positive = good for us, negative = bad for us)
2. **Player Flips**: Continue expansion/simulation until N player turn changes occur
3. **Minimax Integration**: Our turns maximize our score, opponent turns minimize our score
4. **Simplicity**: No special casing - just track when currentPlayer changes
## Node Structure
```cpp
struct MCTSNode {
// Command that led to this state
size_t commandIndex;
CommandType commandType;
// Game state after executing the command
GameStateW resultingGameState;
PlayerId currentPlayer; // Whose turn it is in this state
// Tree position
int playerFlipsFromRoot; // Number of player changes from root
bool isOurTurn; // currentPlayer == our AI's playerId
// MCTS statistics (always from our perspective)
int visitCount = 0;
double totalScore = 0.0;
double averageScore = 0.0;
// Tree structure
std::vector<std::unique_ptr<MCTSNode>> children;
std::vector<size_t> untriedCommands;
};
```
## Expansion Algorithm
```cpp
MCTSNode* Expand(MCTSNode* node) {
// Check if we've reached max player flips
if (node->playerFlipsFromRoot >= maxPlayerFlips) {
return node; // Don't expand further
}
// Pick untried command
size_t cmdIndex = PickUntriedCommand(node);
// Execute command to create child state
auto childEngine = std::make_shared<ShardokEngine>(parentEngine);
PlayerId playerBefore = childEngine->GetCurrentPlayerId();
childEngine->PostCommand(playerBefore, cmdIndex, averageGenerator);
PlayerId playerAfter = childEngine->GetCurrentPlayerId();
// Create child node
auto child = std::make_unique<MCTSNode>();
child->commandIndex = cmdIndex;
child->resultingGameState = childEngine->GetCurrentGameState();
child->currentPlayer = playerAfter;
child->isOurTurn = (playerAfter == ourPlayerId);
// Track player flips
child->playerFlipsFromRoot = node->playerFlipsFromRoot;
if (playerBefore != playerAfter) {
child->playerFlipsFromRoot++;
}
// Score from our perspective
child->immediateScore = ScoreFromOurPerspective(child->resultingGameState);
return child;
}
```
## Simulation Algorithm
```cpp
double Simulate(const GameStateW& startState, PlayerId startPlayer, int startFlips) {
auto simEngine = CreateEngine(startState);
int currentFlips = startFlips;
while (currentFlips < maxPlayerFlips) {
PlayerId currentPlayer = simEngine->GetCurrentPlayerId();
bool isOurTurn = (currentPlayer == ourPlayerId);
// Get available commands
auto commands = simEngine->GetAvailableCommandsForAIPlayer(currentPlayer);
if (commands->empty()) break;
// Pick best command based on whose turn it is
size_t bestCmd = 0;
double bestScore = isOurTurn ? -infinity : +infinity;
for (size_t i = 0; i < commands->size(); ++i) {
auto testEngine = CreateEngine(simEngine);
testEngine->PostCommand(currentPlayer, i, averageGenerator);
// Always score from our perspective
double score = ScoreFromOurPerspective(testEngine->GetCurrentGameState());
// Our turn: maximize our score, Opponent turn: minimize our score
bool shouldSelect = isOurTurn ? (score > bestScore) : (score < bestScore);
if (shouldSelect) {
bestScore = score;
bestCmd = i;
}
}
// Execute chosen command
PlayerId playerBefore = simEngine->GetCurrentPlayerId();
simEngine->PostCommand(currentPlayer, bestCmd, averageGenerator);
PlayerId playerAfter = simEngine->GetCurrentPlayerId();
// Track player flips
if (playerBefore != playerAfter) {
currentFlips++;
}
}
return ScoreFromOurPerspective(simEngine->GetCurrentGameState());
}
```
## Selection Algorithm
```cpp
MCTSNode* SelectChild(MCTSNode* node) {
MCTSNode* bestChild = nullptr;
double bestUCB1 = node->isOurTurn ? -infinity : +infinity;
for (auto& child : node->children) {
double ucb1 = child->averageScore + explorationTerm;
// Our turn: pick highest UCB1, Opponent turn: pick lowest UCB1
bool shouldSelect = node->isOurTurn ? (ucb1 > bestUCB1) : (ucb1 < bestUCB1);
if (shouldSelect) {
bestUCB1 = ucb1;
bestChild = child.get();
}
}
return bestChild;
}
```
## Backpropagation Algorithm
```cpp
void Backpropagate(MCTSNode* node, double score) {
while (node != nullptr) {
node->visitCount++;
node->totalScore += score; // Always from our perspective
node->averageScore = node->totalScore / node->visitCount;
node = node->parent;
}
}
```
## Key Simplifications
1. **No END_TURN special cases** - just check if currentPlayer changed after any command
2. **Consistent scoring** - always from our AI's perspective throughout
3. **Clear minimax** - our nodes maximize, opponent nodes minimize
4. **Simple state tracking** - just count player flips, no complex inheritance
5. **Unified command handling** - all commands handled the same way
## Implementation Plan
1. **Refactor MCTSNode structure** - simplify to core fields needed
2. **Rewrite MCTSExpansion** - remove END_TURN special cases, just track player changes
3. **Fix MCTSSimulation** - ensure consistent perspective and proper minimax
4. **Simplify MCTSSelection** - clean minimax logic
5. **Clean up MCTSBackpropagation** - single perspective throughout
6. **Add comprehensive testing** - verify player flips are tracked correctly
This design eliminates the confusion around perspectives and special cases, making the algorithm much easier to understand and debug.
-1
View File
@@ -76,7 +76,6 @@ use_repo(
"com_github_aws_aws_sdk_go_v2_config",
"com_github_aws_aws_sdk_go_v2_credentials",
"com_github_aws_aws_sdk_go_v2_service_s3",
"org_golang_google_grpc",
"org_golang_google_protobuf",
)
+2 -1
View File
@@ -1 +1,2 @@
UNITY_VERSION='6000.3.0f1'
UNITY_VERSION='6000.1.11f1'
-280
View File
@@ -1,280 +0,0 @@
# CommandProto Usage Analysis in shardok/ai
This document analyzes all remaining usages of `CommandProto` (protocol buffer representation) in the AI code and identifies opportunities to eliminate proto conversion by using `ShardokCommand` directly.
## Summary
**Total CommandProto usages found:** 42 locations across 9 files
**Eliminated:** 6 usages (14%) - ✅ **Phase 1 Complete**
**Can be eliminated:** ~14 usages (33%)
**Must keep (for now):** ~22 usages (53%)
---
## Files with CommandProto Usage
### 1. AICommandFilter.cpp (6 usages) - ✅ **COMPLETED** (PR #4505)
**Location:** Lines 146, 189, 252, 356, 387, 428
**Original usage:**
```cpp
const auto cmdProto = cmd.GetCommandProto();
if (!cmdProto.has_target()) { ... }
const auto& targetCoords = cmdProto.target();
if (!cmdProto.has_actor()) { ... }
const auto unitId = cmdProto.actor().value();
```
**Replaced with:**
```cpp
const int targetRow = cmd.GetTargetRow();
const int targetCol = cmd.GetTargetColumn();
if (targetRow < 0 || targetCol < 0) {
throw ShardokInternalErrorException("Command missing required target");
}
const Coords targetCoords(targetRow, targetCol);
const int actorId = cmd.GetActorUnitId();
if (actorId < 0) {
throw ShardokInternalErrorException("Command missing required actor");
}
```
**Status:****ELIMINATED** - Replaced with direct accessors + exception handling
**Impact:** Eliminated 6 proto conversions in hot path (command filtering)
**Completed:** Phase 1, PR #4505
---
### 2. ShardokAIClient.cpp (8 usages)
**Location:** Lines 83, 86, 87, 102, 105, 237, 261, 311, 356
**Usage breakdown:**
#### a) Command validation (lines 83-87)
```cpp
void CheckCommand(const CommandProto &realDescriptor, const CommandProto &guessedDescriptor) {
differencer.IgnoreField(CommandProto::descriptor()->FindFieldByNumber(
CommandProto::kFollowUpCommandTypesFieldNumber));
```
**Status:****MUST KEEP** - Uses protobuf reflection for comparison
**Reason:** Comparing proto messages for correctness checking requires proto API
#### b) GetAvailableCommandProtos calls (lines 105, 356)
```cpp
const auto guessedCommands = guessedEngine.GetAvailableCommandProtos(playerId, false);
if (const auto &availableCommands = engine.GetAvailableCommandProtos(playerId, false);
```
**Status:****CAN REPLACE** - Should use `GetAvailableCommandsForAIPlayer()` instead
**Impact:** This is a major conversion point - converts entire command list to protos
**Priority:** HIGH (converts all commands to proto unnecessarily)
#### c) Strategy selector methods (lines 102, 237, 261, 311)
```cpp
const vector<CommandProto> &realAvailableCommands) const -> CommandChoiceResults
```
**Status:****CAN REPLACE** - Depends on fixing strategy selector signatures
**Priority:** MEDIUM (depends on other refactors)
---
### 3. IterativeDeepeningAI.cpp/hpp (4 usages)
**Location:** Lines 41, 272 (cpp), 73, 96 (hpp)
**Current usage:**
```cpp
const std::vector<CommandProto>& commands,
```
**Status:****CAN REPLACE** - These methods should accept `CommandListSPtr` instead
**Impact:** Major - this is the main AI search algorithm
**Priority:** HIGH (core AI algorithm)
**Note:** IterativeDeepeningAI already receives commands as proto vectors. The conversion happens upstream at the entry point. Need to trace back to find where `GetAvailableCommandProtos` is called.
---
### 4. AIFleeDecisionCalculator.cpp/hpp (6 usages)
**Location:** Lines 17, 38, 39, 62, 63 (hpp), 18, 19, 137, 138 (cpp)
**Current usage:**
```cpp
const vector<CommandProto>& availableCommands,
const vector<CommandProto>::const_iterator& fleeCommand,
```
**Status:****CAN REPLACE** - Should use `CommandListSPtr` and indices instead
**Impact:** Flee decision logic could avoid proto conversion
**Priority:** MEDIUM
---
### 5. AIAttackerStrategySelector.cpp/hpp (2 usages)
**Location:** Line 30 in both files
**Current usage:**
```cpp
const vector<CommandProto>& availableCommands) -> AIStrategy
```
**Status:** ⚠️ **PARTIALLY REPLACEABLE** - Currently doesn't use the commands parameter
**Current implementation:**
```cpp
const vector<CommandProto>& /*availableCommands*/) -> AIStrategy {
// Parameter is commented out - not used!
return AIStrategy::DEFAULT;
}
```
**Priority:** LOW (parameter unused, but signature should be consistent)
---
### 6. AICommandEvaluator.hpp (1 usage)
**Location:** Line 27
**Current usage:**
```cpp
using CommandProto = net::eagle0::shardok::api::CommandDescriptor;
```
**Status:** ⚠️ **CHECK USAGE** - Type alias, need to check if used
**Priority:** LOW (just a type alias)
---
### 7. AIScoreCalculator.hpp (1 usage)
**Location:** Line 24
**Current usage:**
```cpp
using CommandProto = net::eagle0::shardok::api::CommandDescriptor;
```
**Status:** ⚠️ **CHECK USAGE** - Type alias, need to check if used
**Priority:** LOW (just a type alias)
---
### 8. AIWaterCrossingCommandChooser.hpp (1 usage)
**Location:** Line 20
**Current usage:**
```cpp
using CommandProto = net::eagle0::shardok::api::CommandDescriptor;
```
**Status:** ⚠️ **CHECK USAGE** - Type alias, need to check if used
**Priority:** LOW (just a type alias)
---
## Key Conversion Points (Entry Points)
### ShardokEngine::GetAvailableCommandProtos()
This method converts the entire command list from `CommandListSPtr` to `vector<CommandProto>`.
**Current flow:**
```
ShardokEngine::GetAvailableCommandsForAIPlayer() → CommandListSPtr
↓ (conversion)
ShardokEngine::GetAvailableCommandProtos() → vector<CommandProto>
AI algorithms (IterativeDeepeningAI, etc.)
```
**Desired flow:**
```
ShardokEngine::GetAvailableCommandsForAIPlayer() → CommandListSPtr
↓ (no conversion!)
AI algorithms use CommandSPtr directly
```
---
## Recommendations by Priority
### HIGH Priority (Performance-critical hot paths)
1. **AICommandFilter.cpp (6 usages)**
- Replace `cmd.GetCommandProto()` with direct accessor methods
- Use `GetActorUnitId()`, `GetTargetRow()`, `GetTargetColumn()`
- Impact: Eliminates 6 proto conversions per filtered command
2. **ShardokAIClient.cpp - GetAvailableCommandProtos calls**
- Replace calls to `GetAvailableCommandProtos()` with `GetAvailableCommandsForAIPlayer()`
- Impact: Eliminates conversion of entire command list
3. **IterativeDeepeningAI**
- Change signature from `vector<CommandProto>` to `CommandListSPtr`
- Impact: Main AI search algorithm avoids proto conversion
### MEDIUM Priority
4. **AIFleeDecisionCalculator**
- Change to use `CommandListSPtr` and indices
- Impact: Flee decision logic avoids proto
5. **ShardokAIClient strategy methods**
- Update signatures to use `CommandListSPtr`
- Cascades to strategy selectors
### LOW Priority
6. **Type aliases**
- Remove unused `using CommandProto` declarations
- Clean up imports
---
## Migration Strategy
### Phase 1: Low-hanging fruit (AICommandFilter) - ✅ **COMPLETED** (PR #4505)
- ✅ Replaced 6 proto conversions with direct accessor calls
- ✅ Added exception handling for missing actor/target data
- ✅ No signature changes needed
- ✅ Immediate performance benefit
- **PR:** #4505
### Phase 2: Entry point (ShardokAIClient)
- Replace `GetAvailableCommandProtos()` calls with `GetAvailableCommandsForAIPlayer()`
- Update method signatures in ShardokAIClient
### Phase 3: Core AI (IterativeDeepeningAI)
- Change IterativeDeepeningAI to accept `CommandListSPtr`
- This is the biggest change but has highest impact
### Phase 4: Supporting systems
- Update AIFleeDecisionCalculator
- Update strategy selectors
- Clean up type aliases
### Phase 5: Validation code
- Keep proto-based validation as-is (uses reflection)
- Consider if validation is still needed in production
---
## Notes
- **MCTS already converted**: The MCTS code path already uses `CommandListSPtr` directly
- **Proto still needed**: For serialization/network communication (not in AI hot path)
- **Validation**: Proto comparison in CheckCommand() should remain (uses proto reflection)
---
## Estimated Impact
**Proto conversions eliminated:** ~20-25 per command choice
**Performance gain:** Eliminates hundreds of allocations per AI decision
**Code simplification:** Removes proto conversion layer from AI
**Before:**
```
Command → Proto → AI Decision
```
**After:**
```
Command → AI Decision (direct)
```
File diff suppressed because it is too large Load Diff
-297
View File
@@ -1,297 +0,0 @@
# Deproto Migration Plan
## Vision
**Protocol buffers should only be used at the edges** — for network serialization (gRPC) and disk persistence. Inside the Eagle game engine, all logic should operate on native Scala models.
```
┌─────────────────────────────────────────────────────────────────────┐
│ GRPC BOUNDARY │
│ EagleServiceImpl.scala ←→ Proto Messages ←→ Unity Client │
└─────────────────────────────────────────────────────────────────────┘
GameStateConverter
┌─────────────────────────────────────────────────────────────────────┐
│ SCALA ENGINE │
│ │
│ GameStateC ───→ Actions ───→ ActionResultT ───→ New GameStateC │
│ ↑ │ │
│ │ (Pure Scala models) │ │
│ └───────────────────────────────────────────────────┘ │
│ │
│ HeroC, FactionC, ProvinceC, BattalionC, ArmyC, etc. │
└─────────────────────────────────────────────────────────────────────┘
GameStateConverter
┌─────────────────────────────────────────────────────────────────────┐
│ PERSISTENCE BOUNDARY │
│ GameHistory.scala ←→ Proto Messages ←→ File/Database │
└─────────────────────────────────────────────────────────────────────┘
```
---
## Current State
### Completed Phases
| Phase | Status | Summary |
|-------|--------|---------|
| Phase 1: GameStateC | **Complete** | Scala `GameState` model with 22 fields |
| Phase 2: EngineImpl | **Complete** | Holds Scala `GameState` internally |
| Phase 3: GameHistory | **Complete** | `stateAfter` returns Scala GameState |
| Phase 4: ActionResultT | **Complete** | All 59 actions return `ActionResultT` |
| Phase 5: Action Base Classes | **Complete** | All `RandomSequentialResultsAction` and `DeterministicSingleResultAction` converted to T-type base classes |
| Phase 5b: Base Class Cleanup | **Complete** | `RandomSequentialResultsAction` and `DeterministicSingleResultAction` deleted |
| Phase 5c: RoundPhaseAdvancer Actions | **Complete** | All actions called by RoundPhaseAdvancer accept Scala GameState |
| Phase 5d: RoundPhaseAdvancer Itself | **Complete** | RoundPhaseAdvancer.checkForPhaseAdvancement takes Scala GameState |
### Phase 5c/5d Progress (Complete)
`RoundPhaseAdvancer.checkForPhaseAdvancement` now accepts Scala `GameState` and `ActionResultApplier` directly (PR #4677).
| Action | PR | Status |
|--------|-----|--------|
| `PrisonerExchangeAction` | #4670 | ✅ Merged |
| `PerformForcedTurnBackAction` | #4671 | ✅ Merged |
| `PerformHeroDeparturesAction` | #4672 | ✅ Merged |
| `RequestFreeForAllBattlesAction` | #4673 | ✅ Merged |
| `EndPlayerCommandsPhaseAction` | #4674 | ✅ Merged |
| `EndDiplomacyResolutionPhaseAction` | #4675 | ✅ Merged |
| `RoundPhaseAdvancer` itself | #4677 | ✅ Merged |
### EngineImpl Progress
| Change | PR | Status |
|--------|-----|--------|
| `recursiveTransform` deleted | #4677 | ✅ Merged |
| `recursiveTransformT` uses `RandomStateTSequencer` | #4677 | ✅ Merged |
### Current Architecture
**ActionResultT Production (100% Complete):**
- All actions produce `ActionResultT`
- Conversion to `ActionResultProto` happens via `ActionResultProtoConverter.toProto()`
- No direct `ActionResultProto` construction outside the converter
**ActionResultProto Consumption (Next Target):**
- `ActionResultProtoApplierImpl` - applies proto results to proto GameState
- `RoundPhaseAdvancer` - calls converter, passes protos to applier
- `InMemoryHistory` / `PersistedHistory` - stores proto results
- Service layer (`GameController`, `GamesManager`, etc.) - uses proto for client communication
---
## Phase 6: Migrate to ActionResultT Consumers
### Objective
Eliminate internal consumption of `ActionResultProto`. Everything inside the engine should work with `ActionResultT`.
### Current Flow (Proto-Heavy)
```
Action.execute()
→ ActionResultT
→ ActionResultProtoConverter.toProto()
→ ActionResultProto
→ ActionResultProtoApplierImpl.applyActionResults()
→ GameStateProto
→ GameStateConverter.fromProto()
→ GameStateC
```
### Target Flow (T-Types Throughout)
```
Action.execute()
→ ActionResultT
→ ActionResultTApplier.applyActionResults()
→ GameStateC
(Proto conversion only at boundaries)
```
### Key Files to Convert
**Tier 1 - Core Applier:**
```
src/main/scala/net/eagle0/eagle/library/actions/applier/ActionResultProtoApplierImpl.scala
```
Create `ActionResultApplier` that applies `ActionResultT` directly to Scala `GameState`.
**Tier 2 - RoundPhaseAdvancer:****Complete**
```
src/main/scala/net/eagle0/eagle/library/RoundPhaseAdvancer.scala
```
Now accepts Scala `GameState` and `ActionResultApplier`. Only converts to proto lazily for `AvailableCommandsFactory` calls.
**Tier 3 - Sequencers:**
```
src/main/scala/net/eagle0/eagle/library/actions/impl/common/RandomStateTSequencer.scala
src/main/scala/net/eagle0/eagle/library/actions/impl/common/RandomStateProtoSequencer.scala
```
Modify `RandomStateTSequencer` to thread Scala `GameState` throughout (currently converts to proto internally). Then evaluate whether `RandomStateProtoSequencer` is still needed at all.
**Current State**: `RandomStateTSequencer` accepts Scala `GameState` via its `apply()` method but internally converts to proto. All callback methods (`withRandomActionResult`, `withActionResults`, etc.) pass `GameStateProto` to callers, forcing actions that use the sequencer to work with proto types internally.
**Target State**: Create a fully protoless sequencer where:
1. `lastState` returns Scala `GameState` (not `lastStateProto`)
2. All callback methods pass Scala `GameState` to callers
3. Actions using the sequencer can be fully protoless
**Migration Path**:
1. Add `lastState: GameState` method alongside `lastStateProto` (non-breaking)
2. Add parallel callback methods that pass Scala GameState (e.g., `withScalaActionResult`)
3. Migrate actions one by one to use the new Scala-based callbacks
4. Once all actions migrated, deprecate/remove proto-based callbacks
5. Remove `lastStateProto` once no longer used
**RandomStateSequencer Migration Progress** (PR #4679 introduced protoless `RandomStateSequencer`):
| Action | Status |
|--------|--------|
| `TruceTurnBackPhaseAction` | ✅ Migrated (PR #4680) |
| `EndHandleRiotsPhaseAction` | ✅ Migrated (PR #4684) |
| `PerformVassalCommandsPhaseAction` | ✅ Migrated |
| `PerformVassalDefenseDecisionsAction` | ✅ Migrated |
| `EndVassalCommandsPhaseAction` | ✅ Migrated |
| `PerformReconResolutionAction` | ✅ Migrated |
| `NewRoundAction` | ✅ Migrated (PR #4698) |
| `EndBattleAftermathPhaseAction` | ✅ Migrated (PR #4699) |
| `EndDiplomacyResolutionPhaseAction` | ✅ Migrated |
| `PerformUnaffiliatedHeroesAction` | ✅ Migrated |
| `EngineImpl.recursiveTransformT` | ✅ Migrated (PR #4704) |
| `ProtolessSequentialResultsActionWrapper` | ✅ Migrated (PR #4705) |
| `LegacyRandomStateTSequencer` | ✅ **Deleted** (PR #4705) |
**TCommandFactory Extraction** (PR #4684):
To enable lightweight mocking of command creation in tests, `TCommandFactory` trait was extracted from `CommandFactory`. This allows tests to mock just the `makeTCommand` method without pulling in all 40+ command dependencies that `CommandFactory` requires.
- `TCommandFactory` - lightweight trait with just `makeTCommand`
- `CommandFactory extends TCommandFactory` - maintains backward compatibility
- Actions accepting command factories now use `TCommandFactory` type for better testability
**Tier 4 - History APIs:**
```
src/main/scala/net/eagle0/eagle/service/InMemoryHistory.scala
src/main/scala/net/eagle0/eagle/service/PersistedHistory.scala
```
Change APIs to vend Scala `GameState` and `ActionResultT` instead of proto versions. `PersistedHistory` converts to proto internally for disk persistence; `InMemoryHistory` doesn't need proto at all.
### ActionResultProto Consumer Inventory
| File | Usage | Target |
|------|-------|--------|
| `ActionResultTApplierImpl.scala` | Converts T→Proto, delegates to proto applier | Replace with `ActionResultApplier` |
| `RoundPhaseAdvancer.scala` | ~~20 converter calls~~ | ✅ **Complete** - uses Scala GameState |
| `RandomStateTSequencer.scala` | Converts T→Proto internally | Thread Scala GameState throughout |
| `RandomStateProtoSequencer.scala` | Returns `Vector[ActionResultProto]` | Evaluate if still needed |
| `VigorXPApplier.scala` | Wraps proto results | Convert to work with T |
| `ResolveBattleAction.scala` | 2 converter calls | Convert after dependencies |
| `PerformForcedTurnBackAction.scala` | 1 converter call | Convert after dependencies |
| `EndFreeForAllDecisionPhaseAction.scala` | ~~fromProtoState~~ | **Complete** - now takes Scala GameState |
| `EndBattleRequestPhaseAction.scala` | ~~fromProtoState~~ | **Complete** - now takes Scala GameState |
| `EndDefenseDecisionPhaseAction.scala` | ~~fromProtoState~~ | **Complete** - now takes Scala GameState |
| `EndPleaseRecruitMePhaseAction.scala` | ~~fromProtoState~~ | **Complete** - now takes Scala GameState |
| `InMemoryHistory.scala` | Stores proto results | Vend Scala types, remove proto entirely |
| `PersistedHistory.scala` | Stores proto results | Vend Scala types, convert internally for disk |
| `GameController.scala` | Uses proto for client communication | Keep proto (gRPC boundary) |
### Estimated Effort
| Component | Lines | Complexity |
|-----------|-------|------------|
| `ActionResultApplier` | ~900 | High (port of proto applier) |
| `RandomStateTSequencer` refactor | ~150 | Medium |
| `RoundPhaseAdvancer` updates | ~100 | Medium |
| History API updates | ~100 | Low |
| Action/utility updates | ~200 | Low |
| **Total** | **~1450** | |
### Validation
- [ ] `ActionResultApplier` created and tested
- [ ] `RandomStateTSequencer` threads Scala GameState throughout
- [ ] `RoundPhaseAdvancer` uses T-types internally
- [ ] History APIs vend Scala types
- [ ] No `ActionResultProtoConverter.toProto()` calls except at persistence/gRPC boundaries
- [ ] All tests pass
---
## Phase 7: Clean Up Legacy Utilities
### Objective
Remove remaining direct proto imports from utility classes.
### Files to Modify
| File | Status |
|------|--------|
| `CommandChoiceHelpers.scala` | Accepts proto `GameState`; blocks full deproto of `PerformVassalCommandsPhaseAction` and `PerformVassalDefenseDecisionsAction` |
| `LegacyProvinceUtils.scala` | Replace with `ProvinceUtils.scala` - `hasImminentRiot` added (PR #4683) |
| `LegacyFactionUtils.scala` | Replace proto imports with `FactionT` |
| `LegacyUnaffiliatedHeroUtils.scala` | Replace proto imports with Scala models |
| `BattalionTypeLoader.scala` | Keep proto for file loading, convert immediately after |
| `BeastUtils.scala` | **Complete** - now uses Scala `BeastInfo` only |
### View Filters (Blocking Full Deproto)
The `ProvinceViewFilter` utility currently works entirely with proto types, blocking full deproto of actions that generate province views:
| File | Issue | Needed |
|------|-------|--------|
| `ProvinceViewFilter.scala` | Takes proto `Province`/`GameState`, returns proto `ProvinceView` | Scala `ProvinceViewT` model |
| `GameStateViewFilter.scala` | Uses proto types throughout | Depends on `ProvinceViewT` |
| `GameStateViewDiffer.scala` | Works with view protos | Depends on `ProvinceViewT` |
**Blocked Actions**:
- `EndBattleAftermathPhaseAction` - uses `ProvinceViewFilter` for `revelationChange`, requires lazy proto conversion
- `PerformReconResolutionAction` - uses `ProvinceViewFilter` for reconned provinces
- `GameStateFactionExtensions` - uses `ProvinceViewFilter` for `updatedReconnedProvinces`
**Solution**: Create Scala `ProvinceViewT` (and possibly `ProvinceViewC`) that mirrors the proto `ProvinceView`. Then create a protoless `ProvinceViewFilter` that operates on Scala types. The proto version can delegate to the Scala version + convert, or we maintain both during transition.
---
## Phase 8: Verify Boundaries
### Objective
Confirm protos are used correctly at boundaries — and ONLY there.
### Expected Proto Usage (Keep)
- `EagleServiceImpl.scala` - gRPC boundary
- `InMemoryHistory.scala` / `PersistedHistory.scala` - Persistence boundary
- `*Converter.scala` - Explicit conversion utilities
- `*Loader.scala` - File loading utilities
### Expected No Proto Usage (Verify)
- `/library/actions/impl/` - Pure Scala models
- `/library/util/` - Pure Scala models (except loaders)
- `/model/state/` - Pure Scala models
---
## Open Questions
1. **Persistence Format**: Currently game state is persisted as proto. Should we keep proto for persistence (good for schema evolution) or switch to a different format?
2. **Shardok Integration**: `ResolveBattleAction` communicates with Shardok. Should the Shardok interface use protos (external service) or Scala models?
3. **View Generation**: `GameStateViewDiffer` works with view protos for client updates. Views need Scala models (`ProvinceViewT`, etc.) to allow actions like `EndBattleAftermathPhaseAction` to be fully protoless. The Scala views would be converted to proto only at the gRPC boundary when sending updates to clients.
---
## Success Criteria
### Code Quality
- [ ] Zero proto imports in `/library/actions/` (except boundaries)
- [ ] Zero proto imports in `/library/` utilities (except loaders)
- [ ] `GameStateT` used throughout engine internals
- [ ] Proto usage limited to: `EagleServiceImpl`, loaders, converters, persistence
### Architecture
- [ ] Clear separation: Scala models (internal) vs Proto (boundaries)
- [ ] Converters as the only bridge between domains
- [ ] No "proto creep" into business logic
Binary file not shown.
-1
View File
@@ -9,7 +9,6 @@ require (
github.com/aws/aws-sdk-go-v2/config v1.28.10
github.com/aws/aws-sdk-go-v2/credentials v1.17.51
github.com/aws/aws-sdk-go-v2/service/s3 v1.72.2
google.golang.org/grpc v1.68.0
google.golang.org/protobuf v1.36.3
)
-2
View File
@@ -40,8 +40,6 @@ github.com/google/go-cmp v0.5.5/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/
golang.org/x/text v0.25.0 h1:qVyWApTSYLk/drJRO5mDlNYskwQznZmkpV2c8q9zls4=
golang.org/x/text v0.25.0/go.mod h1:WEdwpYrmk1qmdHvhkSTNPm3app7v4rsT8F2UD6+VHIA=
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
google.golang.org/grpc v1.68.0 h1:aHQeeJbo8zAkAa3pRzrVjZlbz6uSfeOXlJNQM0RAbz0=
google.golang.org/grpc v1.68.0/go.mod h1:fmSPC5AsjSBCK54MyHRx48kpOti1/jRfOlwEWywNjWA=
google.golang.org/protobuf v1.26.0-rc.1 h1:7QnIQpGRHE5RnLKnESfDoxm2dTapTZua5a0kS0A+VXQ=
google.golang.org/protobuf v1.26.0-rc.1/go.mod h1:jlhhOSvTdKEhbULTjvd4ARK9grFBp09yW+WbY/TyQbw=
google.golang.org/protobuf v1.36.3 h1:82DV7MYdb8anAVi3qge1wSnMDrnKK7ebr+I0hHRN1BU=
+2 -3
View File
@@ -3,8 +3,7 @@
set -euxo pipefail
/bin/echo "building darwin bundle"
bazel build --config=mactools @net_eagle0_unity_godice//darwin/framework:DarwinGodiceBundle
ZIP_LOCATION=$(bazel cquery --config=mactools --output=files @net_eagle0_unity_godice//darwin/framework:DarwinGodiceBundle 2>/dev/null)
/usr/bin/unzip -o $ZIP_LOCATION -d src/main/csharp/net/eagle0/clients/unity/eagle0/Assets/Plugins/
bazel build --noincompatible_enable_cc_toolchain_resolution @net_eagle0_unity_godice//darwin/framework:DarwinGodiceBundle
/usr/bin/unzip -o bazel-bin/external/net_eagle0_unity_godice/darwin/framework/DarwinGodiceBundle.zip -d src/main/csharp/net/eagle0/clients/unity/eagle0/Assets/Plugins/
/usr/bin/plutil -convert xml1 src/main/csharp/net/eagle0/clients/unity/eagle0/Assets/Plugins/DarwinGodiceBundle.bundle/Contents/Info.plist
+2 -3
View File
@@ -5,9 +5,8 @@ set -euxo pipefail
/bin/echo "build plugins"
/bin/echo "building darwin bundle"
bazel build --config=mactools @net_eagle0_unity_godice//darwin/framework:DarwinGodiceBundle
ZIP_LOCATION=$(bazel cquery --config=mactools --output=files @net_eagle0_unity_godice//darwin/framework:DarwinGodiceBundle 2>/dev/null)
/usr/bin/unzip -o $ZIP_LOCATION -d src/main/csharp/net/eagle0/clients/unity/eagle0/Assets/Plugins/
bazel build --noincompatible_enable_cc_toolchain_resolution @net_eagle0_unity_godice//darwin/framework:DarwinGodiceBundle
/usr/bin/unzip -o bazel-bin/external/net_eagle0_unity_godice/darwin/framework/DarwinGodiceBundle.zip -d src/main/csharp/net/eagle0/clients/unity/eagle0/Assets/Plugins/
/usr/bin/plutil -convert xml1 src/main/csharp/net/eagle0/clients/unity/eagle0/Assets/Plugins/DarwinGodiceBundle.bundle/Contents/Info.plist
+2 -4
View File
@@ -1,10 +1,8 @@
#!/usr/bin/env bash
curl -L "https://docs.google.com/spreadsheets/d/1pv-WMXReccddPwev_YG9IXEGznuGHrYjNNEZ0Rb-ZhM/export?gid=0&format=tsv" | tr -d '\r' > src/main/resources/net/eagle0/shardok/settings.tsv
curl -L "https://docs.google.com/spreadsheets/d/1p6I5nUMcoAPHIcqikVgbBCFVnqN9dpOEVClbS_wOI7M/export?gid=0&format=tsv" | tr -d '\r' > src/main/resources/net/eagle0/eagle/settings.tsv
curl -L "https://docs.google.com/spreadsheets/d/1pv-WMXReccddPwev_YG9IXEGznuGHrYjNNEZ0Rb-ZhM/export?gid=0&format=tsv" > src/main/resources/net/eagle0/shardok/settings.tsv
curl -L "https://docs.google.com/spreadsheets/d/1p6I5nUMcoAPHIcqikVgbBCFVnqN9dpOEVClbS_wOI7M/export?gid=0&format=tsv" > src/main/resources/net/eagle0/eagle/settings.tsv
bazel run //src/main/go/net/eagle0/build/settings_generator:settings_generator -- \
${PWD}/src/main/resources/net/eagle0/eagle/settings.tsv \
${PWD}/src/main/scala/net/eagle0/eagle/library/settings/
bazel run gazelle
+4 -4
View File
@@ -1,11 +1,11 @@
#!/usr/bin/env bash
curl -L "https://docs.google.com/spreadsheets/d/1DHEsiv4cY4gE6AX3sVH82K__mpBD1aznIYCQwQxA_F0/export?gid=0&format=tsv" | tr -d '\r' > /tmp/names.tsv
curl -L "https://docs.google.com/spreadsheets/d/1DHEsiv4cY4gE6AX3sVH82K__mpBD1aznIYCQwQxA_F0/export?gid=0&format=tsv" > /tmp/names.tsv
bazel run //src/main/scala/net/eagle0/util:name_list_checker -- /tmp/names.tsv > src/main/resources/net/eagle0/names.tsv
bazel run //src/main/scala/net/eagle0/util:name_list_json_maker -- /tmp/names.tsv > src/main/resources/net/eagle0/names.json
curl -L "https://docs.google.com/spreadsheets/d/1NhvG73HKyVE36yGpkV2oJiSIXoNqQOYTr5ArLnucYL0/export?gid=0&format=tsv" | tr -d '\r' > src/main/resources/net/eagle0/shardok/battalionTypes.tsv
curl -L "https://docs.google.com/spreadsheets/d/1pNWiyxIks2wJ1v7jRLFD24zrKHG2AfhC-nkWmQKQGN4/export?gid=0&format=tsv" | tr -d '\r' > src/main/resources/net/eagle0/eagle/heroes.tsv
curl -L "https://docs.google.com/spreadsheets/d/1RUguq5eAQprsZwOOqiCc-1dg4Urc_6iJ6awZsFU4MeI/export?gid=0&format=tsv" | tr -d '\r' > src/main/resources/net/eagle0/eagle/beasts.tsv
curl -L "https://docs.google.com/spreadsheets/d/1NhvG73HKyVE36yGpkV2oJiSIXoNqQOYTr5ArLnucYL0/export?gid=0&format=tsv" > src/main/resources/net/eagle0/shardok/battalionTypes.tsv
curl -L "https://docs.google.com/spreadsheets/d/1pNWiyxIks2wJ1v7jRLFD24zrKHG2AfhC-nkWmQKQGN4/export?gid=0&format=tsv" > src/main/resources/net/eagle0/eagle/heroes.tsv
curl -L "https://docs.google.com/spreadsheets/d/1RUguq5eAQprsZwOOqiCc-1dg4Urc_6iJ6awZsFU4MeI/export?gid=0&format=tsv" > src/main/resources/net/eagle0/eagle/beasts.tsv
#curl -L "https://docs.google.com/spreadsheets/d/1Z-60cJ_N1IasvqpVb5awKEkIYznEeR2IZSdli47oW88/export?gid=0&format=tsv" > src/main/resources/net/eagle0/eagle/province_map.tsv
${PWD}/scripts/dlSettings.sh
-19
View File
@@ -1,19 +0,0 @@
#!/bin/bash
# Pre-commit hook wrapper for gazelle that fails if files are modified.
# This ensures BUILD files are in canonical format before committing.
set -e
# Run gazelle
bazel run //:gazelle 2>/dev/null
# Check if any BUILD files were modified
if ! git diff --quiet -- '*.bazel' '**/BUILD' 'WORKSPACE*'; then
echo ""
echo "ERROR: gazelle modified BUILD files. Please stage the changes and retry:"
echo ""
git diff --name-only -- '*.bazel' '**/BUILD' 'WORKSPACE*'
echo ""
echo "Run: git add -u && git commit"
exit 1
fi
+2 -22
View File
@@ -18,31 +18,11 @@ static inline auto MixIn(uint64_t& hash, const uint8_t byte) {
}
// Hash an entire buffer using FNV-1a
// Fast word-at-a-time implementation - processes 8 bytes at once for better performance
// while maintaining good distribution properties for hash table use
static inline auto HashBuffer(const uint8_t* data, size_t size) -> uint64_t {
if (data == nullptr) { return FNV_OFFSET_BASIS; }
uint64_t hash = FNV_OFFSET_BASIS;
const uint8_t* end = data + size;
// Process 8 bytes at a time
while (data + 8 <= end) {
uint64_t word;
// Use memcpy to avoid alignment issues and let compiler optimize
__builtin_memcpy(&word, data, sizeof(word));
hash ^= word;
hash *= FNV_PRIME;
data += 8;
if (data != nullptr) {
for (size_t i = 0; i < size; ++i) { MixIn(hash, data[i]); }
}
// Process remaining bytes
while (data < end) {
hash ^= static_cast<uint64_t>(*data);
hash *= FNV_PRIME;
data++;
}
return hash;
}
@@ -30,7 +30,7 @@ auto rloc(const string& execPath) -> string {
const std::unique_ptr<Runfiles> runfiles(Runfiles::Create(execPath, &error));
if (runfiles == nullptr) {
fprintf(stderr, "Error! %s\n", error.c_str());
printf("Error! %s\n", error.c_str());
abort();
// error handling
}
@@ -67,9 +67,9 @@ auto FilesystemUtils::MapFilesDirectory() -> string {
void FilesystemUtils::MakeDirectoryIfNecessary(const string& directoryPath) {
if (fs::create_directories(directoryPath))
fprintf(stderr, "Directory %s created\n", directoryPath.c_str());
printf("Directory %s created\n", directoryPath.c_str());
else
fprintf(stderr, "No new directory created for %s\n", directoryPath.c_str());
printf("No new directory created for %s\n", directoryPath.c_str());
}
auto FilesystemUtils::SaveFilesDirectory() -> string {
@@ -129,11 +129,11 @@ auto FilesystemUtils::AtomicallySaveToPath(const string& path, const byte_vector
if (ostr.good()) {
const int err = rename(tempPath.c_str(), path.c_str());
if (err == -1) {
fprintf(stderr, "Failed to move file to %s! Errno %d\n", path.c_str(), errno);
printf("Failed to move file to %s! Errno %d\n", path.c_str(), errno);
return false;
}
} else {
fprintf(stderr, "Failed writing to %s!\n", tempPath.c_str());
printf("Failed writing to %s!\n", tempPath.c_str());
return false;
}
@@ -14,14 +14,6 @@
#include "src/main/cpp/net/eagle0/common/RandomGenerator.hpp"
// A deterministic random generator that returns values from a fixed sequence.
// Used for testing and MCTS simulation where we want specific, predictable outcomes.
//
// Values in the sequence are treated as [0, 1] probabilities that are returned
// by DoubleZeroToOne(). The normal percentile methods (including open-ended
// variants) work as usual, so callers must provide appropriate sequences.
// For example, to get an open-ended low result of -50, provide [0.02, 0.52]
// which produces: initial=2 (triggers open-ended), accumulated=52, final=2-52=-50
class SequenceRandomGenerator : public ::RandomGenerator {
private:
const std::vector<double> sequence;
@@ -1,139 +0,0 @@
# MCTS (Monte Carlo Tree Search) Framework
This directory contains a game-agnostic Monte Carlo Tree Search implementation that can be used with any turn-based game. The framework separates the MCTS algorithm from game-specific logic through abstract interfaces.
## Core Abstract Classes
### `MCTSAction` (abstract/MCTSAction.hpp)
Abstract interface for representing game actions/moves.
**Key Methods:**
- `getIndex()` - Returns the action's unique identifier
- `getDescription()` - Human-readable description for debugging/logging
- `clone()` - Creates a deep copy of the action
- `equals()` - Compares actions for equality
### `MCTSGameState` (abstract/MCTSGameState.hpp)
Abstract interface for representing game states.
**Key Methods:**
- `hash()` - Returns a hash for transposition table lookups
- `score(playerId)` - Evaluates the state's value for a given player
- `currentPlayerId()` - Returns whose turn it is
- `isTerminal()` - Checks if the game has ended
- `getWinner()` - Returns the winning player (if terminal)
- `clone()` - Creates a deep copy of the state
- `equals()` - Compares states for equality
### `MCTSGameEngine` (abstract/MCTSGameEngine.hpp)
Abstract interface for game rule enforcement and state transitions. Many methods have efficient default implementations.
**Must Override (Pure Virtual):**
- `applyAction(state, action)` - Applies an action to create a new state
- `getLegalActions(state)` - Returns all valid moves from a state
- `isTerminal(state)` - Checks if a state is game-ending
- `evaluateState(state, playerId)` - Scores a state for a player
**Optional Overrides (Have Default Implementations):**
- `applyActionMutable(state, action)` - Apply action in-place for efficiency (default: calls applyAction)
- `filterActions(actions, state)` - Applies heuristic filtering (default: no filtering)
- `simulateRandomPlayout(state, playerId, maxDepth, policy)` - Runs simulation (default: efficient mutable implementation)
- `getActionScore(state, action, playerId)` - Scores an action (default: apply and evaluate)
- `shouldStopSearch(state, iterations, startTime)` - Early termination (default: no early stop)
**Performance Features:**
- The default `simulateRandomPlayout` clones the state once and mutates it throughout simulation for efficiency
- Games can override `applyActionMutable` to provide even more efficient in-place updates
- Games can override `simulateRandomPlayout` for custom optimizations (e.g., using internal engine state)
## MCTS Algorithm Implementation
### `AbstractMCTSAI` (abstract/AbstractMCTSAI.hpp)
The main MCTS algorithm implementation that works with any game implementing the abstract interfaces.
**Key Features:**
- **Selection**: Uses UCB1 (Upper Confidence Bound) for node selection
- **Expansion**: Adds new nodes to the search tree
- **Simulation**: Runs random playouts to estimate node values
- **Backpropagation**: Updates node statistics with simulation results
- **Multithreading**: Supports parallel MCTS with configurable thread count
- **Path Compression**: Optimizes move sequences for better performance
**Configuration Options:**
- `explorationConstant` - UCB1 exploration parameter (default: √2)
- `maxSimulationDepth` - Maximum depth for random playouts
- `maxTreeDepth` - Maximum tree depth to prevent stack overflow
- `useMultithreading` - Enable parallel search
- `numThreads` - Number of worker threads
- `simulationPolicy` - Strategy for action selection during simulation
### `MCTSNode` (abstract/MCTSNode.hpp)
Represents nodes in the MCTS search tree.
**Core Data:**
- `action` - The action that led to this node
- `actionIndex` - Index in the original actions array
- `gameState` - The game state at this node
- `visitCount` - Number of times this node was visited
- `totalReward` - Sum of simulation rewards
- `averageReward` - Average reward (totalReward / visitCount)
- `children` - Child nodes in the search tree
- `parent` - Parent node reference
**Key Methods:**
- `CanExpand()` - Checks if node has untried actions
- `GetBestChild(explorationConstant)` - UCB1-based child selection
- `GetBestFinalChild()` - Most-visited child (for final move selection)
- `CalculateUCB1(explorationConstant)` - Computes UCB1 value
## Simulation Policies
The framework supports multiple strategies for action selection during random playouts:
- **RANDOM** - Uniform random selection
- **FILTERED_RANDOM** - Random selection from filtered action set
- **BEST_IMMEDIATE** - Always choose the highest-scoring immediate action
- **WEIGHTED_BEST_IMMEDIATE** - Weighted random selection based on action scores
## Type Definitions
### `MCTSTypes` (abstract/MCTSTypes.hpp)
- `MCTSPlayerId` - Player identifier type (int)
- `MCTSSimulationPolicy` - Enumeration of simulation strategies
- `MCTSConfig` - Configuration structure for MCTS parameters
## Usage Pattern
To use this framework with your game:
1. **Implement the abstract interfaces** for your game:
```cpp
class MyGameAction : public MCTSAction { /* ... */ };
class MyGameState : public MCTSGameState { /* ... */ };
class MyGameEngine : public MCTSGameEngine { /* ... */ };
```
2. **Create and configure the AI**:
```cpp
MCTSConfig config;
config.explorationConstant = 1.414;
config.maxSimulationDepth = 100;
AbstractMCTSAI ai(playerId, config);
```
3. **Run the search**:
```cpp
auto actions = engine.getLegalActions(currentState);
auto result = ai.Search(engine, currentState, actions, timeLimit);
auto bestAction = actions[result.bestActionIndex];
```
## Testing
The framework includes comprehensive tests using a Tic-Tac-Toe implementation:
- `MockTicTacToe.hpp` - Example implementation of all abstract interfaces
- `AbstractMCTSAI_test.cpp` - Unit tests for the core algorithm
- `MCTSIntegration_test.cpp` - Integration tests with complete games
- `MCTSNode_test.cpp` - Tests for the node data structure
This demonstrates how to implement the interfaces and validates that the MCTS algorithm works correctly with any turn-based game.
File diff suppressed because it is too large Load Diff
@@ -1,106 +0,0 @@
//
// Abstract MCTS AI implementation - game agnostic
//
#ifndef EAGLE0_ABSTRACT_MCTSAI_HPP
#define EAGLE0_ABSTRACT_MCTSAI_HPP
#include <chrono>
#include <memory>
#include <unordered_map>
#include <vector>
#include "MCTSAction.hpp"
#include "MCTSGameEngine.hpp"
#include "MCTSGameState.hpp"
#include "MCTSNode.hpp"
#include "MCTSTypes.hpp"
namespace shardok {
namespace mcts {
class AbstractMCTSAI {
public:
// Search result structure
struct SearchResult {
size_t bestActionIndex = 0;
double bestScore = 0.0;
int searchDepth = 0;
int nodesEvaluated = 0;
std::chrono::milliseconds searchTime{0};
bool foundWinningMove = false;
};
explicit AbstractMCTSAI(MCTSPlayerId playerId, MCTSConfig config = MCTSConfig{});
// Main search interface
[[nodiscard]] auto Search(
const MCTSGameEngine& engine,
const MCTSGameState& initialState,
std::chrono::milliseconds timeLimit) const -> SearchResult;
// Configuration
[[nodiscard]] auto GetConfig() const -> const MCTSConfig& { return config_; }
void SetConfig(const MCTSConfig& newConfig) { config_ = newConfig; }
[[nodiscard]] auto FindNodeAtDepthWithHash(
const MCTSNode* root,
int maxDepth,
uint64_t targetHash) -> const MCTSNode*;
private:
MCTSPlayerId playerId_;
MCTSConfig config_;
// Transposition table: maps state hash -> minimum depth at which state was reached
// Used to detect and penalize longer paths to the same game state
// Cleared at the start of each Search() call
mutable std::unordered_map<uint64_t, int> transpositionTable_;
// Core MCTS algorithm
[[nodiscard]] auto BuildMCTSTree(
const MCTSGameEngine& engine,
const MCTSGameState& initialState,
std::chrono::steady_clock::time_point deadline) const -> std::unique_ptr<MCTSNode>;
// MCTS phases
[[nodiscard]] auto MCTSSelection(MCTSNode* root) const -> MCTSNode*;
[[nodiscard]] auto MCTSExpansion(MCTSNode* node, const MCTSGameEngine& engine) const
-> MCTSNode*;
[[nodiscard]] auto MCTSSimulation(
const MCTSGameEngine& engine,
const MCTSGameState& state,
MCTSPlayerId startingPlayer,
int startingPlayerFlips = 0) const -> double;
auto MCTSBackpropagation(MCTSNode* node, double reward, MCTSBackpropagationPolicy policy) const
-> void;
// Helper functions
[[nodiscard]] auto SelectSimulationAction(
const MCTSGameEngine& engine,
const MCTSGameState& state,
const std::vector<std::unique_ptr<MCTSAction>>& actions,
bool isMaximizing) const -> size_t;
// Logging
static auto LogSearchResults(
const MCTSNode* rootNode,
const MCTSNode* bestChild,
const SearchResult& result) -> void;
// Debug tree dumping
static auto DumpTreeToFile(const MCTSNode* root, const std::string& filepath) -> void;
private:
static auto
DumpNodeRecursive(const MCTSNode* node, std::ostream& out, int indentLevel, bool isLastChild)
-> void;
};
} // namespace mcts
} // namespace shardok
#endif // EAGLE0_ABSTRACT_MCTSAI_HPP
@@ -1,94 +0,0 @@
load("//tools:copts.bzl", "COPTS")
cc_library(
name = "mcts_types",
hdrs = ["MCTSTypes.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai/mcts:__subpackages__",
"//src/test/cpp/net/eagle0/common/mcts:__subpackages__",
"//src/test/cpp/net/eagle0/shardok/ai/mcts:__subpackages__",
],
)
cc_library(
name = "mcts_action",
hdrs = ["MCTSAction.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai/mcts:__subpackages__",
"//src/test/cpp/net/eagle0/common/mcts:__subpackages__",
"//src/test/cpp/net/eagle0/shardok/ai/mcts:__subpackages__",
],
)
cc_library(
name = "mcts_game_state",
hdrs = ["MCTSGameState.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai/mcts:__subpackages__",
"//src/test/cpp/net/eagle0/common/mcts:__subpackages__",
"//src/test/cpp/net/eagle0/shardok/ai/mcts:__subpackages__",
],
deps = [
":mcts_types",
],
)
cc_library(
name = "mcts_game_engine",
srcs = ["MCTSGameEngine.cpp"],
hdrs = ["MCTSGameEngine.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai/mcts:__subpackages__",
"//src/test/cpp/net/eagle0/common/mcts:__subpackages__",
"//src/test/cpp/net/eagle0/shardok/ai/mcts:__subpackages__",
],
deps = [
":mcts_action",
":mcts_game_state",
":mcts_types",
],
)
cc_library(
name = "mcts_node",
hdrs = ["MCTSNode.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai:__pkg__",
"//src/main/cpp/net/eagle0/shardok/ai/mcts:__subpackages__",
"//src/test/cpp/net/eagle0/common/mcts:__subpackages__",
"//src/test/cpp/net/eagle0/shardok/ai/mcts:__subpackages__",
],
deps = [
":mcts_action",
":mcts_game_state",
":mcts_types",
],
)
cc_library(
name = "abstract_mcts_ai",
srcs = ["AbstractMCTSAI.cpp"],
hdrs = ["AbstractMCTSAI.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai/mcts:__subpackages__",
"//src/test/cpp/net/eagle0/common/mcts:__subpackages__",
"//src/test/cpp/net/eagle0/shardok/ai/mcts:__subpackages__",
],
deps = [
":mcts_action",
":mcts_game_engine",
":mcts_game_state",
":mcts_node",
":mcts_types",
"//src/main/cpp/net/eagle0/common/mcts/util:tree_indent_util",
],
)
# Individual targets are exposed above - no need for a catch-all target
# Each component should be imported explicitly by its consumers
@@ -1,40 +0,0 @@
//
// Abstract action interface for MCTS
//
#ifndef EAGLE0_MCTS_ACTION_HPP
#define EAGLE0_MCTS_ACTION_HPP
#include <memory>
#include <string>
namespace shardok {
namespace mcts {
// Abstract interface for game actions
class MCTSAction {
public:
virtual ~MCTSAction() = default;
// Get a unique index for this action (used for command indexing)
[[nodiscard]] virtual size_t getIndex() const = 0;
// Get a human-readable description for debugging/logging
[[nodiscard]] virtual std::string getDescription() const = 0;
// Create a deep copy of this action
[[nodiscard]] virtual std::unique_ptr<MCTSAction> clone() const = 0;
// Check if two actions are equivalent
[[nodiscard]] virtual bool equals(const MCTSAction& other) const = 0;
// Check if this action requires a chance node (binary success/failure outcome)
// Examples: START_FIRE, RAISE_DEAD, EXTINGUISH_FIRE
// If true, the game engine should provide outcome probabilities
[[nodiscard]] virtual bool requiresChanceNode() const = 0;
};
} // namespace mcts
} // namespace shardok
#endif // EAGLE0_MCTS_ACTION_HPP
@@ -1,148 +0,0 @@
//
// Default implementations for MCTSGameEngine
//
#include "MCTSGameEngine.hpp"
#include <algorithm>
#include <limits>
#include <random>
#include <vector>
#include "MCTSTypes.hpp" // For MCTSInternalError
namespace shardok {
namespace mcts {
double MCTSGameEngine::simulateRandomPlayout(
const MCTSGameState& state,
MCTSPlayerId playerId,
int maxDepth,
MCTSSimulationPolicy policy) const {
// Clone state once and mutate it throughout simulation for efficiency
auto currentState = state.clone();
int depth = 0;
// Use thread-local random generator for thread safety
static thread_local std::mt19937 gen(std::random_device{}());
// Simulate until terminal or max depth
while (!currentState->isTerminal() && depth < maxDepth) {
auto actions = getLegalActions(*currentState, playerId, 0, 0);
if (actions.empty()) { break; }
size_t selectedIndex = 0;
// Select action based on policy
switch (policy) {
case MCTSSimulationPolicy::RANDOM: {
std::uniform_int_distribution<> dis(0, actions.size() - 1);
selectedIndex = dis(gen);
break;
}
case MCTSSimulationPolicy::FILTERED_RANDOM: {
auto filteredIndices = filterActions(actions, *currentState);
if (!filteredIndices.empty()) {
std::uniform_int_distribution<> dis(0, filteredIndices.size() - 1);
selectedIndex = filteredIndices[dis(gen)];
} else {
// Fall back to random if no actions pass filter
std::uniform_int_distribution<> dis(0, actions.size() - 1);
selectedIndex = dis(gen);
}
break;
}
case MCTSSimulationPolicy::BEST_IMMEDIATE: {
double bestScore = -std::numeric_limits<double>::infinity();
for (size_t i = 0; i < actions.size(); ++i) {
double score = getActionScore(
*currentState,
*actions[i],
currentState->currentPlayerId());
if (score > bestScore) {
bestScore = score;
selectedIndex = i;
}
}
break;
}
case MCTSSimulationPolicy::WEIGHTED_BEST_IMMEDIATE: {
// Score all actions and weight by ranking
std::vector<std::pair<size_t, double>> scores;
scores.reserve(actions.size());
for (size_t i = 0; i < actions.size(); ++i) {
double score = getActionScore(
*currentState,
*actions[i],
currentState->currentPlayerId());
scores.emplace_back(i, score);
}
// Sort by score (descending)
std::sort(scores.begin(), scores.end(), [](const auto& a, const auto& b) {
return a.second > b.second;
});
// Create weights based on ranking (1/rank)
std::vector<double> weights;
weights.reserve(scores.size());
for (size_t i = 0; i < scores.size(); ++i) { weights.push_back(1.0 / (i + 1.0)); }
// Select based on weights
std::discrete_distribution<> dis(weights.begin(), weights.end());
selectedIndex = scores[dis(gen)].first;
break;
}
case MCTSSimulationPolicy::WEIGHTED_HEURISTIC: {
// Get heuristic weights (fast O(1) per action)
const auto weights = getActionWeights(actions, *currentState);
// Filter out zero-weight actions
std::vector<size_t> validIndices;
std::vector<double> validWeights;
validIndices.reserve(actions.size());
validWeights.reserve(actions.size());
for (size_t i = 0; i < weights.size() && i < actions.size(); ++i) {
if (weights[i] > 0.0) {
validIndices.push_back(i);
validWeights.push_back(weights[i]);
}
}
// If all actions filtered out, this is a bug in the weighting logic
if (validWeights.empty()) {
throw MCTSInternalError(
"MCTS simulation (playout): All actions have zero weight in "
"WEIGHTED_HEURISTIC policy (action count: " +
std::to_string(actions.size()) +
") - this indicates incorrect weighting");
}
// Select based on heuristic weights
std::discrete_distribution<> dis(validWeights.begin(), validWeights.end());
selectedIndex = validIndices[dis(gen)];
break;
}
}
// Apply selected action using mutable version for efficiency
applyActionMutable(currentState, *actions[selectedIndex]);
if (!currentState) {
break; // Failed to apply action
}
depth++;
}
// Return evaluation from original player's perspective
return evaluateState(*currentState, playerId);
}
} // namespace mcts
} // namespace shardok
@@ -1,161 +0,0 @@
//
// Abstract game engine interface for MCTS
//
#ifndef EAGLE0_MCTS_GAME_ENGINE_HPP
#define EAGLE0_MCTS_GAME_ENGINE_HPP
#include <chrono>
#include <memory>
#include <vector>
#include "MCTSAction.hpp"
#include "MCTSGameState.hpp"
#include "MCTSTypes.hpp"
namespace shardok {
namespace mcts {
// Information about chance outcomes (supports both binary and multi-outcome)
struct ChanceOutcomeInfo {
std::vector<double> probabilities; // Probability of each outcome (must sum to 1.0)
std::vector<double> rolls; // Roll values for each outcome
// Factory for binary success/failure outcomes (e.g., START_FIRE)
[[nodiscard]] static ChanceOutcomeInfo binary(double successProbability) {
// -100: triggers open-ended low sequence, succeeds against any threshold
// 150: triggers open-ended high sequence, fails against any threshold
return {{successProbability, 1.0 - successProbability}, {-100.0, 150.0}};
}
// Factory for multi-outcome with fixed seeds (e.g., END_TURN)
// Uses uniformly distributed roll values to sample different random outcomes
[[nodiscard]] static ChanceOutcomeInfo multiOutcome(int numOutcomes) {
std::vector<double> probs(numOutcomes, 1.0 / numOutcomes);
std::vector<double> rollValues;
rollValues.reserve(numOutcomes);
// Spread rolls across the percentile range: 10, 30, 50, 70, 90 for 5 outcomes
for (int i = 0; i < numOutcomes; ++i) {
rollValues.push_back(10.0 + (80.0 * i) / (numOutcomes - 1));
}
return {probs, rollValues};
}
[[nodiscard]] const std::vector<double>& getRepresentativeRolls() const { return rolls; }
[[nodiscard]] const std::vector<double>& getProbabilities() const { return probabilities; }
};
// Backward compatibility alias
using BinaryOutcomeInfo = ChanceOutcomeInfo;
// Abstract interface for game engines
class MCTSGameEngine {
public:
virtual ~MCTSGameEngine() = default;
// Apply an action to a state and return the resulting state
// If deterministicRoll is provided (0.0-100.0), use that for any random outcomes
[[nodiscard]] virtual std::unique_ptr<MCTSGameState> applyAction(
const MCTSGameState& state,
const MCTSAction& action,
double deterministicRoll = -1.0) const = 0;
// Apply an action to a mutable state in-place (for efficient simulation)
// Default: clone, apply, and move the result back
// Override this for better performance
virtual void applyActionMutable(std::unique_ptr<MCTSGameState>& state, const MCTSAction& action)
const {
state = applyAction(*state, action);
}
// Get all legal actions for the current state with player flip tracking
// Default implementation ignores flip tracking and calls base version
[[nodiscard]] virtual std::vector<std::unique_ptr<MCTSAction>> getLegalActions(
const MCTSGameState& state,
MCTSPlayerId /*rootPlayerId*/,
int /*currentPlayerFlips*/,
int /*maxPlayerFlips*/) const = 0;
// Check if a state is terminal
[[nodiscard]] virtual bool isTerminal(const MCTSGameState& state) const = 0;
// Evaluate a state from the perspective of a player
[[nodiscard]] virtual double evaluateState(const MCTSGameState& state, MCTSPlayerId playerId)
const = 0;
// Filter actions based on game-specific heuristics
// Returns indices of actions to keep
// Default: no filtering (return all indices)
[[nodiscard]] virtual std::vector<size_t> filterActions(
const std::vector<std::unique_ptr<MCTSAction>>& actions,
const MCTSGameState& /*state*/) const {
std::vector<size_t> indices;
indices.reserve(actions.size());
for (size_t i = 0; i < actions.size(); ++i) { indices.push_back(i); }
return indices;
}
// Get heuristic weights for actions (used by WEIGHTED_HEURISTIC simulation policy)
// Returns weights corresponding to each action (same size as actions vector)
// Weight of 0.0 = never select, higher = more likely to select
// Default: uniform weights (all actions equally likely)
[[nodiscard]] virtual std::vector<double> getActionWeights(
const std::vector<std::unique_ptr<MCTSAction>>& actions,
const MCTSGameState& /*state*/) const {
// Default: uniform weights
return std::vector<double>(actions.size(), 1.0);
}
// Simulate a random playout from the given state
// Default implementation uses policy to select actions
[[nodiscard]] virtual double simulateRandomPlayout(
const MCTSGameState& state,
MCTSPlayerId playerId,
int maxDepth,
MCTSSimulationPolicy policy) const;
// Get the immediate score of applying an action
// Default: apply the action and evaluate the resulting state
[[nodiscard]] virtual double getActionScore(
const MCTSGameState& state,
const MCTSAction& action,
MCTSPlayerId playerId) const {
auto newState = applyAction(state, action);
if (!newState) { return 0.0; }
return evaluateState(*newState, playerId);
}
// Check if we should stop searching (e.g., time limit, found winning move)
[[nodiscard]] virtual bool shouldStopSearch(
const MCTSGameState& /*state*/,
int /*iterations*/,
std::chrono::steady_clock::time_point /*startTime*/) const {
// Default: no early stopping
return false;
}
// Map a filtered action index back to the original unfiltered index
// This is needed when getLegalActions() applies filtering - the returned actions
// may be a subset of all available actions, and this maps back to the original index.
// Default implementation: no filtering, so filtered index = original index
[[nodiscard]] virtual size_t mapFilteredIndexToOriginal(
size_t filteredIndex,
const MCTSGameState& state) const {
// Default: no filtering, index stays the same
(void)state; // Suppress unused parameter warning
return filteredIndex;
}
// Get binary outcome information for an action that requires a chance node
// Only called for actions where action.requiresChanceNode() returns true
// Returns success probability for binary success/failure actions
[[nodiscard]] virtual BinaryOutcomeInfo getBinaryOutcomeInfo(
const MCTSGameState& state,
const MCTSAction& action) const = 0;
};
} // namespace mcts
} // namespace shardok
#endif // EAGLE0_MCTS_GAME_ENGINE_HPP
@@ -1,50 +0,0 @@
//
// Abstract game state interface for MCTS
//
#ifndef EAGLE0_MCTS_GAME_STATE_HPP
#define EAGLE0_MCTS_GAME_STATE_HPP
#include <cstdint>
#include <memory>
#include <string>
#include "MCTSTypes.hpp"
namespace shardok {
namespace mcts {
// Abstract interface for game states
class MCTSGameState {
public:
virtual ~MCTSGameState() = default;
// Compute hash for transposition table
[[nodiscard]] virtual uint64_t hash() const = 0;
// Evaluate the state from the perspective of the given player
[[nodiscard]] virtual double score(MCTSPlayerId playerId) const = 0;
// Get the player whose turn it is
[[nodiscard]] virtual MCTSPlayerId currentPlayerId() const = 0;
// Check if the game has ended
[[nodiscard]] virtual bool isTerminal() const = 0;
// Create a deep copy of the state
[[nodiscard]] virtual std::unique_ptr<MCTSGameState> clone() const = 0;
// Check if two states are equivalent
[[nodiscard]] virtual bool equals(const MCTSGameState& other) const = 0;
// Get winner if terminal, or -1 if not terminal or draw
[[nodiscard]] virtual MCTSPlayerId getWinner() const = 0;
// Optional: Get a string representation for debugging
[[nodiscard]] virtual std::string toString() const { return "MCTSGameState"; }
};
} // namespace mcts
} // namespace shardok
#endif // EAGLE0_MCTS_GAME_STATE_HPP
@@ -1,277 +0,0 @@
//
// Abstract MCTS Node structure for game-agnostic implementation
//
#ifndef EAGLE0_ABSTRACT_MCTSNODE_HPP
#define EAGLE0_ABSTRACT_MCTSNODE_HPP
#include <cmath>
#include <limits>
#include <memory>
#include <vector>
#include "MCTSAction.hpp"
#include "MCTSGameState.hpp"
#include "MCTSTypes.hpp"
namespace shardok {
namespace mcts {
// Node type for MCTS tree
enum class NodeType {
DECISION, // Player chooses an action (standard MCTS node)
CHANCE // Nature determines outcome (for probabilistic actions)
};
// Abstract MCTS Node structure
struct MCTSNode {
// Node type
NodeType nodeType = NodeType::DECISION;
// Action information
std::unique_ptr<MCTSAction> action; // The action that led to this node (null for root)
size_t actionIndex = SIZE_MAX; // Index in the original actions array (SIZE_MAX for root)
// Score information
double immediateScore = 0.0;
double lookaheadScore = 0.0;
// Game state after this action
std::unique_ptr<MCTSGameState> gameState;
// MCTS statistics
int visitCount = 0;
double totalReward = 0.0;
double averageReward = 0.0;
mutable double ucb1Value = 0.0;
double actionWeight = 1.0; // Prior probability/weight for this action (from heuristics)
// Tree structure
std::vector<std::unique_ptr<MCTSNode>> children;
size_t nextUntriedActionIndex = 0; // Next action to expand
size_t totalActions = 0; // Total number of available actions
MCTSNode* parent = nullptr;
// Chance node specific fields (only used when nodeType == CHANCE)
std::vector<double> outcomeProbabilities; // Probability of each outcome
std::vector<double> outcomeRolls; // Representative roll for each outcome
// Game context
MCTSPlayerId playerId;
int depth = 0;
bool isTerminal = false;
int playerFlips = 0; // Number of times the active player has changed from root player
bool isMaximizingPlayer = true; // True if this node is maximizing for root player
// Transposition detection
uint64_t stateHash = 0;
bool isRedundant = false; // True if this node represents a duplicate state
// Constructor for root node
MCTSNode(std::unique_ptr<MCTSGameState> state, MCTSPlayerId pid, int d)
: gameState(std::move(state)),
playerId(pid),
depth(d),
playerFlips(0),
isMaximizingPlayer(true) {
if (gameState) {
stateHash = gameState->hash();
isTerminal = gameState->isTerminal();
}
}
// Constructor for child node
MCTSNode(
std::unique_ptr<MCTSAction> act,
std::unique_ptr<MCTSGameState> state,
MCTSPlayerId pid,
int d,
size_t actIdx = SIZE_MAX,
int flips = 0,
bool isMaximizing = true,
double weight = 1.0)
: action(std::move(act)),
actionIndex(actIdx),
gameState(std::move(state)),
actionWeight(weight),
playerId(pid),
depth(d),
playerFlips(flips),
isMaximizingPlayer(isMaximizing) {
if (gameState) {
stateHash = gameState->hash();
isTerminal = gameState->isTerminal();
}
}
// Iterative destructor to avoid stack overflow with deep trees
~MCTSNode() {
std::vector<std::unique_ptr<MCTSNode>> nodesToDestroy;
nodesToDestroy.swap(children);
while (!nodesToDestroy.empty()) {
std::vector<std::unique_ptr<MCTSNode>> currentBatch;
currentBatch.swap(nodesToDestroy);
for (const auto& node : currentBatch) {
if (node && !node->children.empty()) {
for (auto& child : node->children) {
nodesToDestroy.push_back(std::move(child));
}
node->children.clear();
}
}
}
}
// Calculate UCB1 value for this node from parent's perspective
// Uses prior-weighted formula similar to AlphaGo:
// UCB = Q + c * P * sqrt(N_parent) / (1 + N_child)
// Where P is the action weight (prior probability from heuristics)
[[nodiscard]] double CalculateUCB1(
const double explorationConstant,
const int parentVisitCount,
const bool parentIsMaximizing) const {
// Exploitation: use lookahead score (minimax value)
// For minimizing nodes, negate the score to prefer low child values
const double exploitationValue = parentIsMaximizing ? lookaheadScore : -lookaheadScore;
// Exploration: prior-weighted formula (AlphaGo-style)
// Actions with weight 0.0 (like FLEE_COMMAND) get no exploration bonus
// Unvisited nodes get: c * weight * sqrt(N_parent)
// This prevents bad actions from dominating exploration due to infinite UCB
const double explorationValue = explorationConstant * actionWeight *
std::sqrt(parentVisitCount) / (1.0 + visitCount);
return exploitationValue + explorationValue;
}
// Check if this node can be expanded
[[nodiscard]] bool CanExpand() const { return nextUntriedActionIndex < totalActions; }
// Check if this is a chance node
[[nodiscard]] bool IsChanceNode() const { return nodeType == NodeType::CHANCE; }
// Check if this is a decision node
[[nodiscard]] bool IsDecisionNode() const { return nodeType == NodeType::DECISION; }
// Get best child from chance node (probability-weighted selection)
// For chance nodes, we want to explore outcomes proportionally to their probability
[[nodiscard]] MCTSNode* GetBestChanceChild() const {
if (children.empty() || !IsChanceNode()) return nullptr;
// Find the outcome that is most under-explored relative to its probability
// Expected visits for outcome i: total_visits * probability[i]
// Actual visits: child[i]->visitCount
// Deficit: expected - actual
size_t bestIndex = 0;
double bestDeficit = -std::numeric_limits<double>::max();
for (size_t i = 0; i < children.size(); i++) {
if (!children[i] || children[i]->isRedundant) continue;
const double expectedVisits = visitCount * outcomeProbabilities[i];
const double actualVisits = static_cast<double>(children[i]->visitCount);
const double deficit = expectedVisits - actualVisits;
if (deficit > bestDeficit) {
bestDeficit = deficit;
bestIndex = i;
}
}
return children[bestIndex].get();
}
// Get best child based on UCB1
[[nodiscard]] MCTSNode* GetBestChild(const double explorationConstant) const {
if (children.empty()) return nullptr;
MCTSNode* bestChild = nullptr;
double bestValue = -std::numeric_limits<double>::max();
for (auto& child : children) {
// Skip redundant nodes
if (child->isRedundant) continue;
// Calculate UCB1 value using the helper function
const double value =
child->CalculateUCB1(explorationConstant, visitCount, isMaximizingPlayer);
// Debug logging for UCB selection
static bool enableUCBDebug = false;
if (enableUCBDebug && child->visitCount > 0) {
const double exploitationValue =
isMaximizingPlayer ? child->lookaheadScore : -child->lookaheadScore;
const double explorationValue =
explorationConstant * std::sqrt(std::log(visitCount) / child->visitCount);
printf(" UCB: %s lookahead=%.2f expl=%.2f (+%.2f) = %.2f [%s]\n",
isMaximizingPlayer ? "MAX" : "MIN",
child->lookaheadScore,
exploitationValue,
explorationValue,
value,
child->action ? child->action->getDescription().c_str() : "root");
}
if (value > bestValue) {
bestValue = value;
bestChild = child.get();
}
}
return bestChild;
}
// Get best child based on visit count (for final selection)
[[nodiscard]] MCTSNode* GetBestFinalChild() const {
if (children.empty()) return nullptr;
MCTSNode* bestChild = nullptr;
int bestVisits = 0;
double bestScore = isMaximizingPlayer ? -std::numeric_limits<double>::max()
: std::numeric_limits<double>::max();
for (const auto& child : children) {
// Skip redundant nodes
if (child->isRedundant) continue;
// Prefer most-visited node (robust child selection)
if (child->visitCount > bestVisits) {
bestVisits = child->visitCount;
bestScore = child->lookaheadScore;
bestChild = child.get();
} else if (child->visitCount == bestVisits) {
// Tie-break on lookahead score (minimax value, not poisoned average)
// Maximizing: prefer higher score (better for root player)
// Minimizing: prefer lower score (worse for root player)
const bool shouldReplace = isMaximizingPlayer ? (child->lookaheadScore > bestScore)
: (child->lookaheadScore < bestScore);
if (shouldReplace) {
bestScore = child->lookaheadScore;
bestChild = child.get();
}
}
}
// If no child was visited, fall back to lookahead score
if (!bestChild && !children.empty()) {
for (const auto& child : children) {
if (child->isRedundant) continue;
const bool shouldReplace = isMaximizingPlayer ? (child->lookaheadScore > bestScore)
: (child->lookaheadScore < bestScore);
if (shouldReplace) {
bestScore = child->lookaheadScore;
bestChild = child.get();
}
}
}
return bestChild;
}
};
} // namespace mcts
} // namespace shardok
#endif // EAGLE0_ABSTRACT_MCTSNODE_HPP
@@ -1,60 +0,0 @@
//
// Core types for abstract MCTS implementation
//
#ifndef EAGLE0_MCTS_TYPES_HPP
#define EAGLE0_MCTS_TYPES_HPP
#include <stdexcept>
#include <string>
namespace shardok {
namespace mcts {
// Exception thrown when MCTS encounters an internal error that indicates a bug
class MCTSInternalError : public std::logic_error {
public:
explicit MCTSInternalError(const std::string& message) : std::logic_error(message) {}
};
// Abstract player identifier type
using MCTSPlayerId = int;
// Simulation policy for MCTS rollouts
enum class MCTSSimulationPolicy {
RANDOM, // Pure random selection
FILTERED_RANDOM, // Random from filtered actions
BEST_IMMEDIATE, // Choose best immediate score
WEIGHTED_BEST_IMMEDIATE, // Random weighted by score ranking
WEIGHTED_HEURISTIC // Random weighted by fast heuristics (no score evaluation)
};
// Backpropagation policy for MCTS tree updates
enum class MCTSBackpropagationPolicy {
AVERAGING, // Traditional MCTS averaging (for stochastic/single-player games)
MINIMAX // Minimax backup (for deterministic adversarial games)
};
// Configuration for MCTS algorithm
struct MCTSConfig {
double explorationConstant = 1.414; // UCB1 constant (sqrt(2) by default)
int maxSimulationDepth = 1000; // Maximum depth for rollout
int maxTreeDepth = 2000; // Maximum tree depth to prevent stack overflow
bool useMultithreading = true; // Enable parallel MCTS
int numThreads = 16; // Number of threads for parallel MCTS
MCTSSimulationPolicy simulationPolicy = MCTSSimulationPolicy::BEST_IMMEDIATE;
MCTSBackpropagationPolicy backpropagationPolicy = MCTSBackpropagationPolicy::AVERAGING;
int maxPlayerFlips = 0; // Maximum number of player changes for tree expansion
// (0 = expand through current player's turn only,
// 1 = expand through opponent's first response, etc.)
int maxSimulationFlips = 0; // Maximum player flips for leaf evaluation
// When evaluating a leaf at playerFlips < maxSimulationFlips,
// simulate forward to this phase for fair comparison
// (default 0 = evaluate leaves as-is, backward compatible)
std::string debugDumpPath = ""; // If non-empty, dump MCTS tree to this file path
};
} // namespace mcts
} // namespace shardok
#endif // EAGLE0_MCTS_TYPES_HPP
@@ -1,8 +0,0 @@
load("@rules_cc//cc:defs.bzl", "cc_library")
cc_library(
name = "tree_indent_util",
srcs = ["TreeIndentUtil.cpp"],
hdrs = ["TreeIndentUtil.hpp"],
visibility = ["//visibility:public"],
)
@@ -1,53 +0,0 @@
//
// Utility functions for processing tree indentation with UTF-8 box drawing characters
//
#include "TreeIndentUtil.hpp"
namespace mcts::util {
namespace {
// Box drawing characters for tree visualization
constexpr const char* kBranch = "\xE2\x94\x9C"; // ├
constexpr const char* kCorner = "\xE2\x94\x94"; // └
constexpr const char* kVertical = "\xE2\x94\x82"; // │
constexpr const char* kHorizontal = "\xE2\x94\x80"; // ─
} // namespace
std::string BuildTreeIndent(int indentLevel, bool isLastChild) {
std::string indent;
for (int i = 0; i < indentLevel; ++i) {
if (i == indentLevel - 1) {
indent += isLastChild ? kCorner : kBranch;
indent += kHorizontal;
indent += " ";
} else {
indent += " ";
}
}
return indent;
}
std::string ConvertBranchToContinuation(const std::string& indent) {
std::string result = indent;
const std::string replacement = std::string(kVertical) + " ";
// Replace ├ and └ with │
size_t pos = 0;
while ((pos = result.find(kBranch, pos)) != std::string::npos) {
result.replace(pos, 3, replacement); // UTF-8 chars are 3 bytes
pos += replacement.size();
}
pos = 0;
while ((pos = result.find(kCorner, pos)) != std::string::npos) {
result.replace(pos, 3, replacement);
pos += replacement.size();
}
return result;
}
} // namespace mcts::util
@@ -1,22 +0,0 @@
//
// Utility functions for processing tree indentation with UTF-8 box drawing characters
//
#ifndef EAGLE0_TREE_INDENT_UTIL_HPP
#define EAGLE0_TREE_INDENT_UTIL_HPP
#include <string>
namespace mcts::util {
// Builds tree indentation string for a node at a given depth
// Returns string like " ├─ " or " └─ " with proper spacing
std::string BuildTreeIndent(int indentLevel, bool isLastChild);
// Converts tree branch characters (├ and └) to continuation lines (│) for sub-content
// This preserves the tree structure when displaying additional info below a node
std::string ConvertBranchToContinuation(const std::string& indent);
} // namespace mcts::util
#endif // EAGLE0_TREE_INDENT_UTIL_HPP
@@ -9,6 +9,7 @@
#include <ranges>
#include <unordered_map>
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/util/HexMapUtils.hpp"
namespace shardok {
@@ -75,28 +76,30 @@ auto MinDistanceIncludingBraving(
auto EffectiveDistance(
const Unit* unit,
const HexMap* map,
const AttackLocations& attackLocations,
const MapId& mapId,
const APDCache& apdCache,
const BattalionTypeGetter& battalionTypeGetter,
ActionPoints braveWaterCost) -> DIST_T {
const AttackLocations& attackLocations,
const SettingsGetter& settings,
const int braveWaterCost) -> DIST_T {
return EffectiveDistance(
unit,
map,
attackLocations.LocationsWithEnemyInRange(unit),
mapId,
apdCache,
battalionTypeGetter,
attackLocations.LocationsWithEnemyInRange(unit),
settings,
braveWaterCost);
}
auto EffectiveDistance(
const Unit* unit,
const HexMap* map,
const CoordsSet& locations,
const MapId& mapId,
const APDCache& apdCache,
const BattalionTypeGetter& battalionTypeGetter,
ActionPoints braveWaterCost) -> DIST_T {
const auto mapId = ActionPointDistancesCache::GetMapId(map);
const auto& battType = battalionTypeGetter(unit->battalion().type());
const CoordsSet& locations,
const SettingsGetter& settings,
const int braveWaterCost) -> DIST_T {
const auto& battType = settings.GetBattalionType(unit->battalion().type());
const auto* notBravingApd = apdCache->GetRaw(map, mapId, battType, false);
const ActionPointDistances* bravingApd = nullptr;
if (battType->allowsBraveWater) {
@@ -129,12 +132,12 @@ auto GenerateTargetPriorities(
const vector<const Unit*>& remainingUnits,
const APDCache& apdCache,
const ALCache& alCache,
const BattalionTypeGetter& battalionTypeGetter,
ActionPoints braveWaterCost,
const MapId& mapId,
const SettingsGetter& settings,
const bool isLateGame) -> vector<TargetPriorityList> {
auto cc = map->column_count();
const auto mapId = ActionPointDistancesCache::GetMapId(map);
const auto braveWaterCost = settings.Backing().brave_water_action_point_cost();
vector<TargetPriorityList> allTargetsUnitsAndDistances{};
allTargetsUnitsAndDistances.reserve(remainingUnits.size());
@@ -158,7 +161,7 @@ auto GenerateTargetPriorities(
vector<TargetAndDistance> targetsWithDistance;
// Get APDs directly from cache (now with built-in thread-local optimization)
const auto& battType = battalionTypeGetter(unit->battalion().type());
const auto& battType = settings.GetBattalionType(unit->battalion().type());
const auto* notBravingApd = apdCache->GetRaw(map, mapId, battType, false);
const ActionPointDistances* bravingApd = nullptr;
if (battType->allowsBraveWater) {
@@ -5,12 +5,12 @@
#ifndef EAGLE0_AIATTACKGROUPS_HPP
#define EAGLE0_AIATTACKGROUPS_HPP
#include <functional>
#include <vector>
#include "src/main/cpp/net/eagle0/shardok/ai/AIAttackLocations.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistancesCache.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/map/CoordsSet.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
#include "src/main/flatbuffer/net/eagle0/shardok/storage/game_state.hpp"
#include "src/main/flatbuffer/net/eagle0/shardok/storage/player_info.hpp"
#include "src/main/flatbuffer/net/eagle0/shardok/storage/unit.hpp"
@@ -22,8 +22,6 @@ using Unit = net::eagle0::shardok::storage::fb::Unit;
using net::eagle0::shardok::storage::fb::PlayerInfo;
using std::vector;
using BattalionTypeGetter = std::function<BattalionTypeSPtr(BattalionTypeId)>;
struct TargetAndAttackLocations {
Coords target;
CoordsSet attackLocations;
@@ -43,18 +41,20 @@ struct TargetPriorityList {
auto EffectiveDistance(
const Unit* unit,
const HexMap* map,
const AttackLocations& attackLocations,
const MapId& mapId,
const APDCache& apdCache,
const BattalionTypeGetter& battalionTypeGetter,
ActionPoints braveWaterCost) -> DIST_T;
const AttackLocations& attackLocations,
const SettingsGetter& settings,
int braveWaterCost) -> DIST_T;
auto EffectiveDistance(
const Unit* unit,
const HexMap* map,
const CoordsSet& locations,
const MapId& mapId,
const APDCache& apdCache,
const BattalionTypeGetter& battalionTypeGetter,
ActionPoints braveWaterCost) -> DIST_T;
const CoordsSet& locations,
const SettingsGetter& settings,
int braveWaterCost) -> DIST_T;
auto EffectiveDistance(
const Unit* unit,
@@ -71,8 +71,8 @@ auto GenerateTargetPriorities(
const vector<const Unit*>& remainingUnits,
const APDCache& apdCache,
const ALCache& alCache,
const BattalionTypeGetter& battalionTypeGetter,
ActionPoints braveWaterCost,
const MapId& mapId,
const SettingsGetter& settings,
bool isLateGame = false) -> vector<TargetPriorityList>;
} // namespace shardok
@@ -4,7 +4,6 @@
#include "AIAttackerStrategySelector.hpp"
#include "AIAttackGroups.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIFleeDecisionCalculator.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIScoreUtilities.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/util/HexMapUtils.hpp"
@@ -21,13 +20,11 @@ auto AIAttackerStrategySelector::BestAttackerStrategy(
const PlayerId attackerPid,
const GameStateW& gameState,
const CoordsSet& criticalTileCoords,
int maxRounds,
const APDCache& apdCache,
const ALCache& alCache,
const BattalionTypeGetter& battalionTypeGetter,
ActionPoints braveWaterCost,
const SettingsGetter& settings,
const AIWaterCrossingCommandChooser& waterCrossingCommandChooser,
const CommandListSPtr& /*availableCommands*/) -> AIStrategy {
const vector<CommandProto>& /*availableCommands*/) -> AIStrategy {
uint32_t attackerUnitCount = 0;
int defenderOccupiedCriticalTileCount = 0;
bool canFlee = false;
@@ -66,14 +63,12 @@ auto AIAttackerStrategySelector::BestAttackerStrategy(
if (canFlee && AIFleeDecisionCalculator::ShouldConsiderFleeing(
attackerPid,
gameState,
maxRounds,
settings,
FLEE_CONSIDERATION_THRESHOLD)) {
chosenStrategy = FleeStrategy;
} else if (const CoordsSet startCrossingLocations =
waterCrossingCommandChooser.StartCrossingFrom(
battalionTypeGetter,
gameState,
criticalTileCoords);
waterCrossingCommandChooser
.StartCrossingFrom(settings, gameState, criticalTileCoords);
!startCrossingLocations.empty()) {
chosenStrategy = CrossRiversStrategy(startCrossingLocations);
} else if (attackerUnitCount < criticalTileCoords.size()) {
@@ -88,8 +83,8 @@ auto AIAttackerStrategySelector::BestAttackerStrategy(
attackerUnits,
apdCache,
alCache,
battalionTypeGetter,
braveWaterCost));
ActionPointDistancesCache::GetMapId(gameState->hex_map()),
settings));
}
// If any critical tile is occupied by the defender, attack the castles.
// Otherwise, try to hold the castles.
@@ -105,8 +100,8 @@ auto AIAttackerStrategySelector::BestAttackerStrategy(
attackerUnits,
apdCache,
alCache,
battalionTypeGetter,
braveWaterCost));
ActionPointDistancesCache::GetMapId(gameState->hex_map()),
settings));
} else {
chosenStrategy = HoldCastlesStrategy;
}
@@ -6,12 +6,9 @@
#define EAGLE0_AIATTACKERSTRATEGYSELECTOR_HPP
#include "src/main/cpp/net/eagle0/shardok/ai/AIAttackLocations.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AICommonTypes.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIStrategy.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIWaterCrossingCommandChooser.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/GameStateW.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokCommand.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistancesCache.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/map/CoordsSet.hpp"
namespace shardok {
@@ -22,13 +19,11 @@ public:
PlayerId attackerPid,
const GameStateW& gameState,
const CoordsSet& criticalTileCoords,
int maxRounds,
const APDCache& apdCache,
const ALCache& alCache,
const BattalionTypeGetter& battalionTypeGetter,
ActionPoints braveWaterCost,
const SettingsGetter& settings,
const AIWaterCrossingCommandChooser& waterCrossingCommandChooser,
const CommandListSPtr& availableCommands) -> AIStrategy;
const vector<CommandProto>& availableCommands) -> AIStrategy;
};
} // namespace shardok
@@ -1,560 +0,0 @@
//
// Command evaluator for AI lookahead search.
// Extracted from AIScoreCalculator to separate concerns.
//
#include "AICommandEvaluator.hpp"
#include <chrono>
#include <cmath>
#include <future>
#include <limits>
#include "AICommandFilter.hpp"
#include "TranspositionTable.hpp"
#include "src/main/cpp/net/eagle0/common/SequenceRandomGenerator.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/score/AIScoreCalculator.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokEngine.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/util/HexCubeUtils.hpp"
namespace shardok {
// No need to forward declare internal functions - use the public interface instead
// Helper constants and static variables
static const std::vector<double> _averageSequence = {0.5};
static const auto _averageGenerator = std::make_shared<SequenceRandomGenerator>(_averageSequence);
#define MULTITHREAD true
#define LOGGING_ 0
// Helper function to determine if a command type is deterministic
static auto IsDeterministic(const CommandType type) -> bool {
switch (type) {
case net::eagle0::shardok::common::MOVE_COMMAND:
case net::eagle0::shardok::common::CONTROL_COMMAND:
case net::eagle0::shardok::common::METEOR_START_COMMAND:
case net::eagle0::shardok::common::METEOR_TARGET_COMMAND:
case net::eagle0::shardok::common::METEOR_CANCEL_COMMAND:
case net::eagle0::shardok::common::END_TURN_COMMAND:
case net::eagle0::shardok::common::PLACE_UNIT_COMMAND:
case net::eagle0::shardok::common::PLACE_HIDDEN_UNIT_COMMAND:
case net::eagle0::shardok::common::UNIT_STOP_COMMAND:
case net::eagle0::shardok::common::UNIT_REST_COMMAND:
case net::eagle0::shardok::common::FLEE_COMMAND:
case net::eagle0::shardok::common::REINFORCE_COMMAND:
case net::eagle0::shardok::common::RETREAT_COMMAND:
case net::eagle0::shardok::common::END_PLAYER_SETUP_COMMAND:
case net::eagle0::shardok::common::HIDE_COMMAND:
case net::eagle0::shardok::common::FORTIFY_COMMAND:
case net::eagle0::shardok::common::BECOME_OUTLAW_COMMAND:
case net::eagle0::shardok::common::HOLY_WAVE_COMMAND:
case net::eagle0::shardok::common::REPAIR_COMMAND: return true;
default: return false;
}
}
// Helper function to sort commands by score
static auto CommandSorter(
const AICommandEvaluator::IndexAndScore& l,
const AICommandEvaluator::IndexAndScore& r) -> bool {
if (l.lookaheadScore < r.lookaheadScore) return true;
if (l.lookaheadScore > r.lookaheadScore) return false;
// At this point the scores are tied
if (l.immediateScore < r.immediateScore) return true;
if (l.immediateScore > r.immediateScore) return false;
return false;
}
AICommandEvaluator::AICommandEvaluator(
const AIScoreCalculator& scorer,
const APDCache& apdCache,
BattalionTypeGetter battalionTypeGetter)
: scorer_(scorer),
apdCache_(apdCache),
battalionTypeGetter_(std::move(battalionTypeGetter)) {} // Move the function object
auto AICommandEvaluator::PerformLookahead(
const PlayerId pid,
const bool isDefender,
const int remainingLookahead,
const int maxRepeatCount,
const std::shared_ptr<ShardokEngine>& innerEngine,
const ScoreValue currentUtility,
const AIStrategy& attackerStrategy,
const CoordsSet& allCastleCoords,
std::chrono::steady_clock::time_point deadline) const -> std::future<ScoreValue> {
// Check transposition table before expensive computation
auto cachedScore =
g_transpositionTable.probe(innerEngine->GetCurrentGameState(), remainingLookahead, pid);
if (cachedScore.has_value()) {
// Return cached result immediately
std::promise<ScoreValue> p;
p.set_value(*cachedScore);
return p.get_future();
}
const auto nextUtility = currentUtility;
// Check if we've reached the depth limit before making recursive calls
if (remainingLookahead <= 0) {
// Store the current utility in the transposition table and return it
// Note: Store with depth 1 since depth 0 indicates an empty entry in the transposition
// table
g_transpositionTable.store(innerEngine->GetCurrentGameState(), 1, pid, nextUtility);
std::promise<ScoreValue> p;
p.set_value(nextUtility);
return p.get_future();
}
if (const CommandListSPtr nextCommands = innerEngine->GetAvailableCommandsForAIPlayer(pid);
nextCommands && !nextCommands->empty()) {
// Get the future from FindBestCommand without calling .get()
auto bestCommandFuture = FindBestCommand(
pid,
isDefender,
remainingLookahead - 1,
maxRepeatCount,
*innerEngine,
attackerStrategy,
nextUtility,
allCastleCoords,
deadline);
// Return a future that chains the best command evaluation
return std::async(
std::launch::deferred,
[bestCommandFuture = std::move(bestCommandFuture),
innerEngine,
pid,
nextUtility,
remainingLookahead]() mutable -> ScoreValue {
const auto [index, type, lookaheadScore, immediateScore] =
bestCommandFuture.get();
ScoreValue resultScore;
if (auto& nextCommand =
innerEngine->GetAvailableCommandsForAIPlayer(pid)->at(index);
nextCommand->GetCommandType() !=
net::eagle0::shardok::common::END_TURN_COMMAND) {
resultScore = immediateScore;
} else {
resultScore = nextUtility;
}
// Store in transposition table before returning
g_transpositionTable.store(
innerEngine->GetCurrentGameState(),
remainingLookahead,
pid,
resultScore);
return resultScore;
});
}
// No commands available, store and return the current utility as a future
g_transpositionTable
.store(innerEngine->GetCurrentGameState(), remainingLookahead, pid, nextUtility);
std::promise<ScoreValue> p;
p.set_value(nextUtility);
return p.get_future();
}
auto AICommandEvaluator::EvaluateWithRandomness(
const PlayerId pid,
const bool isDefender,
const uint32_t commandIndex,
const int remainingLookahead,
const int maxRepeatCount,
const std::shared_ptr<RandomGenerator>& randomGenerator,
const ShardokEngine& guessedEngine,
const AIStrategy& attackerStrategy,
const CoordsSet& allCastleCoords,
std::chrono::steady_clock::time_point deadline) const -> ImmediateAndLookaheadScore {
ImmediateAndLookaheadScore returnValue{};
// Check if we've exceeded the deadline
if (std::chrono::steady_clock::now() > deadline) {
// Return with a default score and an empty future that resolves immediately
std::promise<ScoreValue> p;
p.set_value(0.0); // Default timeout score
returnValue.immediateScore = 0.0;
returnValue.lookaheadScore = p.get_future();
return returnValue;
}
auto innerEngine = std::make_shared<ShardokEngine>(guessedEngine, false);
innerEngine->PostCommand(pid, commandIndex, randomGenerator);
auto innerUtility = scorer_.GuessedStateScore(
isDefender,
innerEngine->GetCurrentGameState(),
attackerStrategy,
allCastleCoords);
returnValue.immediateScore = innerUtility;
if (remainingLookahead <= 0) {
std::promise<ScoreValue> p;
returnValue.lookaheadScore = p.get_future();
p.set_value(innerUtility);
} else {
auto lookaheadLambda = [this,
pid,
isDefender,
remainingLookahead,
maxRepeatCount,
innerEngine,
attackerStrategy,
innerUtility,
&allCastleCoords,
deadline]() -> ScoreValue {
auto lookaheadFuture = PerformLookahead(
pid,
isDefender,
remainingLookahead,
maxRepeatCount,
innerEngine,
innerUtility,
attackerStrategy,
allCastleCoords,
deadline);
return lookaheadFuture.get();
};
#if MULTITHREAD
auto launchPolicy = remainingLookahead == 1 ? std::launch::async : std::launch::deferred;
returnValue.lookaheadScore = std::async(launchPolicy, lookaheadLambda);
#else
std::promise<ScoreValue> p;
returnValue.lookaheadScore = p.get_future();
auto lambdaResult = lookaheadLambda();
p.set_value(lambdaResult);
#endif
}
return returnValue;
}
auto AICommandEvaluator::FindBestCommand(
const PlayerId pid,
const bool isDefender,
const int remainingLookahead,
const int maxRepeatCount,
const ShardokEngine& guessedEngine,
const AIStrategy& attackerStrategy,
const ScoreValue currentUtility,
const CoordsSet& allCastleCoords,
std::chrono::steady_clock::time_point deadline) const -> std::future<IndexAndScore> {
const CommandListSPtr guessedDescriptors = guessedEngine.GetAvailableCommandsForAIPlayer(pid);
// Filter out obviously bad commands to reduce search space
const std::vector<size_t> filteredIndices = AICommandFilter::FilterCommands(
guessedDescriptors,
pid,
isDefender,
guessedEngine.GetCurrentGameState(),
apdCache_,
battalionTypeGetter_);
const auto& gameState = guessedEngine.GetCurrentGameState();
// Calculate minimum hex distance to enemies for this player
double minDistToEnemies = std::numeric_limits<double>::max();
const auto* units = gameState->units();
for (size_t i = 0; i < units->size(); ++i) {
if (const auto* playerUnit = units->Get(static_cast<unsigned int>(i));
playerUnit->player_id() == pid) {
const auto& playerCoords = playerUnit->location();
for (size_t j = 0; j < units->size(); ++j) {
if (const auto* enemyUnit = units->Get(static_cast<unsigned int>(j));
enemyUnit->player_id() != pid) {
const auto& enemyCoords = enemyUnit->location();
// Proper hex distance calculation using cube coordinates
const Cube playerCube = OffsetToCube(playerCoords);
const Cube enemyCube = OffsetToCube(enemyCoords);
const int hexDistance = CubeDistance(playerCube, enemyCube);
minDistToEnemies = std::min(minDistToEnemies, static_cast<double>(hexDistance));
}
}
}
}
if (minDistToEnemies == std::numeric_limits<double>::max()) {
minDistToEnemies = 0.0; // No enemies found
}
#if LOGGING_
// Log command count and distance metrics for performance analysis
const auto allCommandCount = guessedDescriptors->size();
const auto filteredCommandCount = filteredIndices.size();
const int currentRound = gameState->current_round();
printf("AI_COMMAND_COUNT: Round %d, Player %d, Defender %d, MinDist %.1f, Commands %zu -> %zu "
"(%.1f%% filtered)\n",
currentRound,
static_cast<int>(pid),
isDefender ? 1 : 0,
minDistToEnemies,
allCommandCount,
filteredCommandCount,
100.0 * (allCommandCount - filteredCommandCount) / allCommandCount);
#endif
const auto commandCount = filteredIndices.size();
// Structure to hold all command evaluation data
struct CommandEvaluation {
size_t index;
CommandType type;
ScoreValue immediateScore;
std::vector<std::future<ScoreValue>> lookaheadFutures;
};
std::vector<CommandEvaluation> commandEvaluations(commandCount);
for (uint32_t index = 0; index < commandCount; index++) {
const auto originalIndex = filteredIndices[index];
const auto& guessedDescriptor = guessedDescriptors->at(originalIndex);
const auto guessedCommandType = guessedDescriptor->GetCommandType();
commandEvaluations[index].index = originalIndex;
commandEvaluations[index].type = guessedCommandType;
if (guessedCommandType == net::eagle0::shardok::common::END_TURN_COMMAND) {
std::promise<ScoreValue> p;
commandEvaluations[index].lookaheadFutures.push_back(p.get_future());
p.set_value(currentUtility);
commandEvaluations[index].immediateScore = currentUtility;
} else if (IsDeterministic(guessedCommandType)) {
auto [immediateScore, lookaheadScore] = EvaluateWithRandomness(
pid,
isDefender,
originalIndex,
remainingLookahead,
maxRepeatCount,
_averageGenerator,
guessedEngine,
attackerStrategy,
allCastleCoords,
deadline);
commandEvaluations[index].immediateScore = immediateScore;
commandEvaluations[index].lookaheadFutures.push_back(std::move(lookaheadScore));
} else if (guessedDescriptor->HasOdds()) {
const auto successChancePercentile = guessedDescriptor->GetOddsPercentile();
const double successChance = static_cast<double>(successChancePercentile) / 100.0;
// Success attempt uses 1.0 - (successChance / 2) as the roll
auto [successImmediateScore, successLookaheadScore] = EvaluateWithRandomness(
pid,
isDefender,
originalIndex,
remainingLookahead,
maxRepeatCount,
std::make_shared<SequenceRandomGenerator>(
std::vector{1.0 - successChance / 2.0}),
guessedEngine,
attackerStrategy,
allCastleCoords,
deadline);
// Failure attempt uses the average of (1 - successChance) and 0 as the roll
auto [failureImmediateScore, failureLookaheadScore] = EvaluateWithRandomness(
pid,
isDefender,
originalIndex,
remainingLookahead,
maxRepeatCount,
std::make_shared<SequenceRandomGenerator>(
std::vector{(1.0 - successChance) / 2.0}),
guessedEngine,
attackerStrategy,
allCastleCoords,
deadline);
commandEvaluations[index].immediateScore =
std::lerp(failureImmediateScore, successImmediateScore, successChance);
auto successSF = successLookaheadScore.share();
auto failureSF = failureLookaheadScore.share();
commandEvaluations[index].lookaheadFutures.push_back(std::async(
std::launch::deferred,
[successSF, failureSF, successChance]() -> double {
return std::lerp(failureSF.get(), successSF.get(), successChance);
}));
} else {
ScoreValue sum = 0.0;
for (int repeatIteration = 0; repeatIteration < maxRepeatCount; repeatIteration++) {
// In each iteration, use a double from [0, 1] as the random roll
auto sequence = std::vector{
static_cast<double>(repeatIteration) /
static_cast<double>(maxRepeatCount - 1)};
auto [immediateScore, lookaheadScore] = EvaluateWithRandomness(
pid,
isDefender,
originalIndex,
remainingLookahead,
maxRepeatCount,
std::make_shared<SequenceRandomGenerator>(sequence),
guessedEngine,
attackerStrategy,
allCastleCoords,
deadline);
sum += immediateScore;
commandEvaluations[index].lookaheadFutures.push_back(std::move(lookaheadScore));
}
commandEvaluations[index].immediateScore = sum / maxRepeatCount;
}
}
// Return a future that will wait for all evaluations and find the best one
return std::async(
std::launch::deferred,
[evals = std::move(commandEvaluations)]() mutable -> IndexAndScore {
std::vector<IndexAndScore> allResults;
allResults.reserve(evals.size());
// Wait for all futures and compute final scores
for (auto& eval : evals) {
ScoreValue totalLookaheadScore = 0.0;
for (auto& future : eval.lookaheadFutures) {
totalLookaheadScore += future.get();
}
ScoreValue avgLookaheadScore =
eval.lookaheadFutures.empty()
? eval.immediateScore
: totalLookaheadScore / eval.lookaheadFutures.size();
allResults.push_back(IndexAndScore{
.index = eval.index,
.type = eval.type,
.lookaheadScore = avgLookaheadScore,
.immediateScore = eval.immediateScore});
}
// Find the best command using the existing sorter
auto bestIt = std::ranges::max_element(allResults, CommandSorter);
return *bestIt;
});
}
auto AICommandEvaluator::EvaluateCommand(
const PlayerId pid,
const bool isDefender,
const int remainingLookahead,
const int maxRepeatCount,
const ShardokEngine& guessedEngine,
const AIStrategy& attackerStrategy,
const ScoreValue currentUtility,
const CoordsSet& allCastleCoords,
const size_t commandIndex,
std::chrono::steady_clock::time_point deadline) const -> std::future<ScoreValue> {
const CommandListSPtr guessedDescriptors = guessedEngine.GetAvailableCommandsForAIPlayer(pid);
if (commandIndex >= guessedDescriptors->size()) {
std::promise<ScoreValue> p;
p.set_value(currentUtility);
return p.get_future();
}
const auto& guessedDescriptor = guessedDescriptors->at(commandIndex);
if (const auto guessedCommandType = guessedDescriptor->GetCommandType();
guessedCommandType == net::eagle0::shardok::common::END_TURN_COMMAND) {
std::promise<ScoreValue> p;
p.set_value(currentUtility);
return p.get_future();
} else if (IsDeterministic(guessedCommandType)) {
auto [immediateScore, lookaheadScore] = EvaluateWithRandomness(
pid,
isDefender,
commandIndex,
remainingLookahead,
maxRepeatCount,
_averageGenerator,
guessedEngine,
attackerStrategy,
allCastleCoords,
deadline);
return std::move(lookaheadScore);
} else if (guessedDescriptor->HasOdds()) {
const auto successChancePercentile = guessedDescriptor->GetOddsPercentile();
const double successChance = static_cast<double>(successChancePercentile) / 100.0;
// Success attempt
auto [successImmediateScore, successLookaheadScore] = EvaluateWithRandomness(
pid,
isDefender,
commandIndex,
remainingLookahead,
maxRepeatCount,
std::make_shared<SequenceRandomGenerator>(std::vector{1.0 - successChance / 2.0}),
guessedEngine,
attackerStrategy,
allCastleCoords,
deadline);
// Failure attempt
auto [failureImmediateScore, failureLookaheadScore] = EvaluateWithRandomness(
pid,
isDefender,
commandIndex,
remainingLookahead,
maxRepeatCount,
std::make_shared<SequenceRandomGenerator>(std::vector{(1.0 - successChance) / 2.0}),
guessedEngine,
attackerStrategy,
allCastleCoords,
deadline);
// Return weighted average of success and failure
auto successSF = successLookaheadScore.share();
auto failureSF = failureLookaheadScore.share();
return std::async(std::launch::deferred, [successSF, failureSF, successChance]() -> double {
return std::lerp(failureSF.get(), successSF.get(), successChance);
});
} else {
// For non-deterministic commands without odds, use multiple attempts
std::vector<std::future<ScoreValue>> lookaheadFutures;
lookaheadFutures.reserve(maxRepeatCount);
for (int repeatIteration = 0; repeatIteration < maxRepeatCount; repeatIteration++) {
auto sequence = std::vector{
static_cast<double>(repeatIteration) / static_cast<double>(maxRepeatCount - 1)};
auto [immediateScore, lookaheadScore] = EvaluateWithRandomness(
pid,
isDefender,
commandIndex,
remainingLookahead,
maxRepeatCount,
std::make_shared<SequenceRandomGenerator>(sequence),
guessedEngine,
attackerStrategy,
allCastleCoords,
deadline);
lookaheadFutures.push_back(std::move(lookaheadScore));
}
// Return a future that computes the average when needed
return std::async(
std::launch::deferred,
[lookaheadFutures = std::move(lookaheadFutures),
maxRepeatCount]() mutable -> double {
ScoreValue total = 0.0;
for (auto& future : lookaheadFutures) { total += future.get(); }
return total / maxRepeatCount;
});
}
}
} // namespace shardok
@@ -1,110 +0,0 @@
//
// Command evaluator for AI lookahead search.
// Separated from AIScoreCalculator to isolate pure state scoring from lookahead logic.
//
#ifndef EAGLE0_AICOMMANDEVALUATOR_HPP
#define EAGLE0_AICOMMANDEVALUATOR_HPP
#include <chrono>
#include <future>
#include "src/main/cpp/net/eagle0/shardok/ai/AIStrategy.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/GameStateW.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokCTypes.h"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistancesCache.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/map/CoordsSet.hpp"
#include "src/main/protobuf/net/eagle0/shardok/common/command_type.pb.h"
namespace shardok {
// Forward declarations
class AIScoreCalculator;
class ShardokEngine;
using ScoreValue = double;
using CommandType = net::eagle0::shardok::common::CommandType;
using BattalionTypeGetter = std::function<BattalionTypeSPtr(BattalionTypeId)>;
/// Evaluates commands with lookahead using minimax-style search.
/// Uses AIScoreCalculator for pure state evaluation, adds recursive lookahead logic.
class AICommandEvaluator {
public:
/// Construct evaluator with a scorer for state evaluation and dependencies for command
/// filtering
AICommandEvaluator(
const AIScoreCalculator& scorer,
const APDCache& apdCache,
BattalionTypeGetter battalionTypeGetter); // Pass by value
/// Evaluates the score for a particular command index with lookahead.
[[nodiscard]] auto EvaluateCommand(
PlayerId pid,
bool isDefender,
int remainingLookahead,
int maxRepeatCount,
const ShardokEngine& guessedEngine,
const AIStrategy& attackerStrategy,
ScoreValue currentUtility,
const CoordsSet& allCastleCoords,
size_t commandIndex,
std::chrono::steady_clock::time_point deadline) const -> std::future<ScoreValue>;
/// Find the best command among all available commands at the given depth.
struct IndexAndScore {
size_t index;
CommandType type;
ScoreValue lookaheadScore;
ScoreValue immediateScore;
};
[[nodiscard]] auto FindBestCommand(
PlayerId pid,
bool isDefender,
int remainingLookahead,
int maxRepeatCount,
const ShardokEngine& guessedEngine,
const AIStrategy& attackerStrategy,
ScoreValue currentUtility,
const CoordsSet& allCastleCoords,
std::chrono::steady_clock::time_point deadline) const -> std::future<IndexAndScore>;
private:
const AIScoreCalculator& scorer_;
const APDCache& apdCache_;
BattalionTypeGetter battalionTypeGetter_; // Store by value
struct ImmediateAndLookaheadScore {
ScoreValue immediateScore;
std::future<ScoreValue> lookaheadScore;
};
/// Recursive lookahead calculator
[[nodiscard]] auto PerformLookahead(
PlayerId pid,
bool isDefender,
int remainingLookahead,
int maxRepeatCount,
const std::shared_ptr<ShardokEngine>& innerEngine,
ScoreValue currentUtility,
const AIStrategy& attackerStrategy,
const CoordsSet& allCastleCoords,
std::chrono::steady_clock::time_point deadline) const -> std::future<ScoreValue>;
/// Evaluate single command execution with randomness handling
[[nodiscard]] auto EvaluateWithRandomness(
PlayerId pid,
bool isDefender,
uint32_t commandIndex,
int remainingLookahead,
int maxRepeatCount,
const std::shared_ptr<class RandomGenerator>& randomGenerator,
const ShardokEngine& guessedEngine,
const AIStrategy& attackerStrategy,
const CoordsSet& allCastleCoords,
std::chrono::steady_clock::time_point deadline) const -> ImmediateAndLookaheadScore;
};
} // namespace shardok
#endif // EAGLE0_AICOMMANDEVALUATOR_HPP
@@ -7,7 +7,6 @@
#include <algorithm>
#include "src/main/cpp/net/eagle0/shardok/library/BattalionType.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokException.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/util/HexCubeUtils.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/util/HexMapUtils.hpp"
#include "src/main/protobuf/net/eagle0/shardok/common/command_type.pb.h"
@@ -37,8 +36,8 @@ std::vector<size_t> AICommandFilter::FilterCommands(
PlayerId pid,
bool isDefender,
const GameStateW& gameState,
const APDCache& apdCache,
const BattalionTypeGetter& battalionTypeGetter) {
const SettingsGetter& settings,
const APDCache& apdCache) {
std::vector<size_t> filteredIndices;
filteredIndices.reserve(commands->size());
@@ -67,8 +66,8 @@ std::vector<size_t> AICommandFilter::FilterCommands(
pid,
isDefender,
gameState,
settings,
apdCache,
battalionTypeGetter,
enemyLocations,
castleLocations,
minDistToEnemies)) {
@@ -81,22 +80,16 @@ std::vector<size_t> AICommandFilter::FilterCommands(
pid,
isDefender,
gameState,
settings,
apdCache,
battalionTypeGetter,
enemyLocations,
minDistToEnemies)) {
shouldFilter = true;
}
// Check strategic blunders
if (!shouldFilter && IsStrategicBlunder(
*cmd,
pid,
isDefender,
gameState,
apdCache,
battalionTypeGetter,
minDistToEnemies)) {
if (!shouldFilter &&
IsStrategicBlunder(*cmd, pid, isDefender, gameState, settings, minDistToEnemies)) {
shouldFilter = true;
}
@@ -111,8 +104,8 @@ bool AICommandFilter::IsWastefulAction(
PlayerId pid,
bool isDefender,
const GameStateW& gameState,
const SettingsGetter& settings,
const APDCache& apdCache,
const BattalionTypeGetter& battalionTypeGetter,
const CoordsSet& enemyLocations,
const CoordsSet& castleLocations,
double minDistToEnemies) {
@@ -144,16 +137,15 @@ bool AICommandFilter::IsWastefulAction(
if (!isDefender) {
// Attackers: Only allow fire if the target location is on or adjacent to an enemy
const int targetRow = cmd.GetTargetRow();
const int targetCol = cmd.GetTargetColumn();
if (targetRow < 0 || targetCol < 0) {
throw ShardokInternalErrorException(
"START_FIRE_COMMAND missing required target information");
const auto cmdProto = cmd.GetCommandProto();
if (!cmdProto.has_target()) {
return true; // Can't analyze without target info
}
const Coords fireLocation(
static_cast<int8_t>(targetRow),
static_cast<int8_t>(targetCol));
const auto& targetCoords = cmdProto.target();
const Coords fireLocation{
static_cast<int8_t>(targetCoords.row()),
static_cast<int8_t>(targetCoords.column())};
// Check if any enemy is on the fire location or adjacent to it
bool enemyNearFireLocation = false;
@@ -188,12 +180,13 @@ bool AICommandFilter::IsWastefulAction(
if (!isDefender) {
// Attackers: Only allow fortify if within 3 hexes of enemies or castles
const int unitId = cmd.GetActorUnitId();
if (unitId < 0) {
throw ShardokInternalErrorException(
"FORTIFY_COMMAND missing required actor information");
const auto cmdProto = cmd.GetCommandProto();
if (!cmdProto.has_actor()) {
return true; // Can't analyze without actor info
}
const auto unitId = cmdProto.actor().value();
// Get the acting unit directly by ID
const Unit* actingUnit = gameState->units()->Get(unitId);
// verify the unit is still active
@@ -250,18 +243,16 @@ bool AICommandFilter::IsWastefulAction(
// These actions can fail, so we need high confidence of benefit (8+ action points
// saved)
const int unitId = cmd.GetActorUnitId();
const int targetRow = cmd.GetTargetRow();
const int targetCol = cmd.GetTargetColumn();
if (unitId < 0 || targetRow < 0 || targetCol < 0) {
throw ShardokInternalErrorException(
"BUILD_BRIDGE/FREEZE_WATER_COMMAND missing required actor or target "
"information");
const auto cmdProto = cmd.GetCommandProto();
if (!cmdProto.has_actor() || !cmdProto.has_target()) {
return true; // Can't analyze without full command info
}
const Coords waterLocation(
static_cast<int8_t>(targetRow),
static_cast<int8_t>(targetCol));
const auto unitId = cmdProto.actor().value();
const auto& targetCoords = cmdProto.target();
const Coords waterLocation{
static_cast<int8_t>(targetCoords.row()),
static_cast<int8_t>(targetCoords.column())};
// Get the acting unit directly by ID
const Unit* actingUnit = gameState->units()->Get(unitId);
@@ -278,7 +269,7 @@ bool AICommandFilter::IsWastefulAction(
}
// Get action point distances for this unit's battalion type
const auto& battType = battalionTypeGetter(actingUnit->battalion().type());
const auto& battType = settings.GetBattalionType(actingUnit->battalion().type());
const auto* apd = apdCache->GetRaw(
gameState->hex_map(),
ActionPointDistancesCache::GetMapId(gameState->hex_map()),
@@ -356,16 +347,15 @@ bool AICommandFilter::IsWastefulAction(
case CommandType::REPAIR_COMMAND: {
// Repair filtering - filter repairs with high integrity targets
// Note: RepairCommandFactory already filters enemy-occupied targets
const int targetRow = cmd.GetTargetRow();
const int targetCol = cmd.GetTargetColumn();
if (targetRow < 0 || targetCol < 0) {
throw ShardokInternalErrorException(
"REPAIR_COMMAND missing required target information");
const auto cmdProto = cmd.GetCommandProto();
if (!cmdProto.has_target()) {
return true; // Can't analyze without target info
}
const Coords repairLocation(
static_cast<int8_t>(targetRow),
static_cast<int8_t>(targetCol));
const auto& targetCoords = cmdProto.target();
const Coords repairLocation{
static_cast<int8_t>(targetCoords.row()),
static_cast<int8_t>(targetCoords.column())};
// Check terrain modifiers at target location
const auto* terrain = GetTerrain(gameState->hex_map(), repairLocation);
@@ -388,16 +378,15 @@ bool AICommandFilter::IsWastefulAction(
case CommandType::EXTINGUISH_FIRE_COMMAND: {
// Extinguish fire filtering - don't extinguish fires on enemy-occupied tiles
const int targetRow = cmd.GetTargetRow();
const int targetCol = cmd.GetTargetColumn();
if (targetRow < 0 || targetCol < 0) {
throw ShardokInternalErrorException(
"EXTINGUISH_FIRE_COMMAND missing required target information");
const auto cmdProto = cmd.GetCommandProto();
if (!cmdProto.has_target()) {
return true; // Can't analyze without target info
}
const Coords fireLocation(
static_cast<int8_t>(targetRow),
static_cast<int8_t>(targetCol));
const auto& targetCoords = cmdProto.target();
const Coords fireLocation{
static_cast<int8_t>(targetCoords.row()),
static_cast<int8_t>(targetCoords.column())};
// Check if any enemy occupies the fire location - let them burn!
std::vector<PlayerId> allyPids; // Empty for now - assume 2-player game
@@ -418,8 +407,8 @@ bool AICommandFilter::IsWastefulMovement(
PlayerId pid,
bool isDefender,
const GameStateW& gameState,
const SettingsGetter& settings,
const APDCache& apdCache,
const BattalionTypeGetter& battalionTypeGetter,
const CoordsSet& enemyLocations,
double minDistToEnemies) {
if (cmd.GetCommandType() != CommandType::MOVE_COMMAND) { return false; }
@@ -429,17 +418,17 @@ bool AICommandFilter::IsWastefulMovement(
return false; // Don't filter defender movement or when close to enemies
}
// Get unit and target information directly from command
const int unitId = cmd.GetActorUnitId();
const int targetRow = cmd.GetTargetRow();
const int targetCol = cmd.GetTargetColumn();
// Get the command proto to access unit and target information
const auto cmdProto = cmd.GetCommandProto();
// Check if we have the required information
if (unitId < 0 || targetRow < 0 || targetCol < 0) {
throw ShardokInternalErrorException(
"MOVE_COMMAND missing required actor or target information");
if (!cmdProto.has_actor() || !cmdProto.has_target()) {
return false; // Can't analyze without unit and target info
}
const auto unitId = cmdProto.actor().value();
const auto& targetCoords = cmdProto.target();
// Get the acting unit directly by ID
const Unit* actingUnit = gameState->units()->Get(unitId);
// Verify the unit is still active
@@ -455,10 +444,12 @@ bool AICommandFilter::IsWastefulMovement(
}
const auto& currentCoords = actingUnit->location();
const Coords targetCoordsFlat(static_cast<int8_t>(targetRow), static_cast<int8_t>(targetCol));
const Coords targetCoordsFlat{
static_cast<int8_t>(targetCoords.row()),
static_cast<int8_t>(targetCoords.column())};
// Get action point distances for this unit's battalion type
const auto& battType = battalionTypeGetter(actingUnit->battalion().type());
const auto& battType = settings.GetBattalionType(actingUnit->battalion().type());
const auto* apd = apdCache->GetRaw(
gameState->hex_map(),
ActionPointDistancesCache::GetMapId(gameState->hex_map()),
@@ -499,8 +490,7 @@ bool AICommandFilter::IsStrategicBlunder(
PlayerId /*pid*/,
bool /*isDefender*/,
const GameStateW& /*gameState*/,
const APDCache& /*apdCache*/,
const BattalionTypeGetter& /*battalionTypeGetter*/,
const SettingsGetter& /*settings*/,
double /*minDistToEnemies*/) {
// Simplified strategic blunder detection for now
// TODO: Implement proper castle abandonment detection
@@ -8,12 +8,12 @@
#include <memory>
#include <vector>
#include "src/main/cpp/net/eagle0/shardok/ai/AICommonTypes.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokEngine.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistancesCache.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/map/CoordsSet.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
#include "src/main/flatbuffer/net/eagle0/shardok/storage/game_state.hpp"
#include "src/main/protobuf/net/eagle0/shardok/api/command_descriptor.pb.h"
namespace shardok {
@@ -32,8 +32,8 @@ public:
* @param pid Player ID making the move
* @param isDefender True if this player is the defender
* @param gameState Current game state
* @param settings Game settings for parameter lookup
* @param apdCache Action point distance cache for distance calculations
* @param battalionTypeLookup Function to look up battalion types by ID
* @return Filtered list of commands worth evaluating
*/
static std::vector<size_t> FilterCommands(
@@ -41,8 +41,8 @@ public:
PlayerId pid,
bool isDefender,
const GameStateW& gameState,
const APDCache& apdCache,
const BattalionTypeGetter& battalionTypeLookup);
const SettingsGetter& settings,
const APDCache& apdCache);
private:
// Helper to build enemy locations once for efficiency
@@ -54,8 +54,8 @@ private:
PlayerId pid,
bool isDefender,
const GameStateW& gameState,
const SettingsGetter& settings,
const APDCache& apdCache,
const BattalionTypeGetter& battalionTypeLookup,
const CoordsSet& enemyLocations,
const CoordsSet& castleLocations,
double minDistToEnemies);
@@ -66,8 +66,8 @@ private:
PlayerId pid,
bool isDefender,
const GameStateW& gameState,
const SettingsGetter& settings,
const APDCache& apdCache,
const BattalionTypeGetter& battalionTypeLookup,
const CoordsSet& enemyLocations,
double minDistToEnemies);
@@ -77,8 +77,7 @@ private:
PlayerId pid,
bool isDefender,
const GameStateW& gameState,
const APDCache& apdCache,
const BattalionTypeGetter& battalionTypeLookup,
const SettingsGetter& settings,
double minDistToEnemies);
// Helper functions for distance and position analysis
@@ -1,23 +0,0 @@
//
// AICommonTypes.hpp
// Common type definitions used across AI utility functions
//
#ifndef EAGLE0_AICOMMONTYPES_HPP
#define EAGLE0_AICOMMONTYPES_HPP
#include <functional>
#include "src/main/cpp/net/eagle0/shardok/library/BattalionType.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokCTypes.h"
namespace shardok {
// Function type for looking up battalion types by ID
// Used across AI utilities to get battalion type information without
// needing to pass the entire scorer object
using BattalionTypeGetter = std::function<BattalionTypeSPtr(BattalionTypeId)>;
} // namespace shardok
#endif // EAGLE0_AICOMMONTYPES_HPP
@@ -10,14 +10,8 @@ namespace shardok {
// Enum for AI algorithm selection
enum class AIAlgorithmType {
ITERATIVE_DEEPENING, // Default: Minimax with sophisticated randomness
MCTS // Monte Carlo Tree Search with multithreading
};
// Enum for scoring calculator selection
enum class ScoringCalculatorType {
STANDARD, // Default: Unbounded raw scores
NORMALIZED, // Normalized scores in [0, 1] range for ML training
MCTS_OPTIMIZED // Bounded linear scores tuned for MCTS
MCTS, // Monte Carlo Tree Search with multithreading (original)
MCTS_CLEAN // Clean MCTS implementation with fixed simulation perspective
};
} // namespace shardok
@@ -7,10 +7,8 @@
#include <algorithm>
#include <ranges>
#include "src/main/cpp/net/eagle0/shardok/ai/AIAttackGroups.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIScoreUtilities.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIWaterCrossingCalculator.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/score/AIScoreCalculator.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/util/HexMapUtils.hpp"
namespace shardok {
@@ -21,9 +19,8 @@ constexpr double MINIMUM_RATIO_FOR_DEFENDER_TO_HOLD = 0.60;
auto AIDefenderStrategySelector::BestDefenderStrategy(
const GameStateW& gameState,
const CoordsSet& criticalTileCoords,
int maxRounds,
const APDCache& apdCache,
const BattalionTypeGetter& battalionTypeGetter) -> AIStrategy {
const SettingsGetter& settings) -> AIStrategy {
uint32_t attackerNonUndeadUnitCount = 0;
uint32_t attackerNonUndeadUnitNotRequiringWaterCrossingCount = 0;
int attackerTroops = 0;
@@ -39,7 +36,7 @@ auto AIDefenderStrategySelector::BestDefenderStrategy(
player->player_id(),
criticalTileCoords,
apdCache,
battalionTypeGetter);
settings);
attackerUnitIdsRequiringWaterCrossing.insert(
attackerUnitIdsRequiringWaterCrossing.end(),
unitIdsRequiringWaterCrossing.begin(),
@@ -74,7 +71,7 @@ auto AIDefenderStrategySelector::BestDefenderStrategy(
}
}
const int roundsRemaining = maxRounds - gameState->current_round();
const int roundsRemaining = 32 - gameState->current_round();
AIStrategy chosenStrategy;
// Defender will flee if
@@ -6,23 +6,19 @@
#define EAGLE0_AIDEFENDERSTRATEGYSELECTOR_HPP
#include "src/main/cpp/net/eagle0/shardok/ai/AIAttackLocations.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AICommonTypes.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIStrategy.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/GameStateW.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistancesCache.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/map/CoordsSet.hpp"
#include "src/main/flatbuffer/net/eagle0/shardok/storage/game_state.hpp"
namespace shardok {
class AIDefenderStrategySelector {
public:
static auto BestDefenderStrategy(
const GameStateW& gameState,
const CoordsSet& criticalTileCoords,
int maxRounds,
const APDCache& apdCache,
const BattalionTypeGetter& battalionTypeGetter) -> AIStrategy;
const SettingsGetter& settings) -> AIStrategy;
};
} // namespace shardok
@@ -49,8 +49,8 @@ auto DefenderDistanceBuf(
const vector<const Unit *> &attackerUnits,
const APDCache &apdCache,
const ALCache &alCache,
const BattalionTypeGetter &battalionTypeGetter,
ActionPoints braveWaterCost,
const SettingsGetter &settings,
const int braveWaterActionPointCost,
const bool lateGame,
const bool includeUndead) -> double {
const auto &locationsToAttackMe = alCache->CachedLocations(defenderLocation, lateGame);
@@ -73,14 +73,14 @@ auto DefenderDistanceBuf(
notBravingDistances[typeInt] = apdCache->GetRaw(
hexMap,
mapId,
battalionTypeGetter(attacker->battalion().type()),
settings.GetBattalionType(attacker->battalion().type()),
false);
bravingDistances[typeInt] = apdCache->GetRaw(
hexMap,
mapId,
battalionTypeGetter(attacker->battalion().type()),
settings.GetBattalionType(attacker->battalion().type()),
true,
braveWaterCost);
braveWaterActionPointCost);
}
}
@@ -6,10 +6,10 @@
#define EAGLE0_AIDISTANCEDEBUF_HPP
#include "AIAttackLocations.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AICommonTypes.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistances.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistancesCache.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/map/CoordsSet.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
namespace shardok {
@@ -23,8 +23,8 @@ auto DefenderDistanceBuf(
const vector<const Unit *> &attackerUnits,
const APDCache &apdCache,
const ALCache &alCache,
const BattalionTypeGetter &battalionTypeGetter,
ActionPoints braveWaterCost,
const SettingsGetter &settings,
int braveWaterActionPointCost,
bool lateGame,
bool includeUndead) -> double;
@@ -15,15 +15,15 @@
namespace shardok {
auto AIFleeDecisionCalculator::GetFleeCommandIndex(
const CommandList::const_iterator& fleeCommand,
const CommandListSPtr& availableCommands) -> size_t {
return static_cast<size_t>(std::distance(availableCommands->begin(), fleeCommand));
const vector<CommandProto>::const_iterator& fleeCommand,
const vector<CommandProto>& availableCommands) -> size_t {
return static_cast<size_t>(std::distance(availableCommands.begin(), fleeCommand));
}
auto AIFleeDecisionCalculator::EstimateCombatSuccess(
PlayerId attackerPlayerId,
const GameStateW& gameState,
int maxRounds) -> double {
const SettingsGetter& settings) -> double {
if (gameState->status() == nullptr ||
gameState->status()->state() !=
net::eagle0::shardok::storage::fb::GameStatus_::State_GAME_RUNNING) {
@@ -68,7 +68,7 @@ auto AIFleeDecisionCalculator::EstimateCombatSuccess(
}
}
const int roundsRemaining = maxRounds - gameState->current_round();
const int roundsRemaining = settings.Backing().max_rounds() - gameState->current_round();
// Special case: Attacker has no heroes - automatic loss
if (attackerHeroes == 0) {
@@ -133,15 +133,17 @@ auto AIFleeDecisionCalculator::EstimateCombatSuccess(
auto AIFleeDecisionCalculator::EvaluateFleeVsFight(
PlayerId playerId,
const SettingsGetter& settingsGetter,
const GameStateW& guessedState,
const CommandListSPtr& availableCommands,
const CommandList::const_iterator& fleeCommand,
int maxRounds,
int minimumFleeOddsThreshold,
int desperateFleeThreshold,
const vector<CommandProto>& availableCommands,
const vector<CommandProto>::const_iterator& fleeCommand,
bool enableDebugLogging) -> FleeDecision {
// Get flee success odds
const int fleeSuccessChance = (*fleeCommand)->GetOddsPercentile();
const int fleeSuccessChance = fleeCommand->odds().success_chance();
// Get thresholds from settings
const int minimumFleeOddsThreshold = settingsGetter.Backing().ai_minimum_flee_odds_threshold();
const int desperateFleeThreshold = settingsGetter.Backing().ai_desperate_flee_threshold();
if (enableDebugLogging) {
printf("AI FinalRound: Evaluating flee (odds=%d%%)...\n", fleeSuccessChance);
@@ -161,7 +163,7 @@ auto AIFleeDecisionCalculator::EvaluateFleeVsFight(
}
// Low flee odds - evaluate if fighting might be better
const double combatWinChance = EstimateCombatSuccess(playerId, guessedState, maxRounds);
const double combatWinChance = EstimateCombatSuccess(playerId, guessedState, settingsGetter);
// If combat situation is hopeless, even bad flee odds are better than certain death
if (combatWinChance <= 0.05 && fleeSuccessChance >= desperateFleeThreshold) {
@@ -213,11 +215,11 @@ auto AIFleeDecisionCalculator::EvaluateFleeVsFight(
auto AIFleeDecisionCalculator::ShouldConsiderFleeing(
PlayerId attackerPlayerId,
const GameStateW& guessedState,
int maxRounds,
const SettingsGetter& settings,
double fleeConsiderationThreshold) -> bool {
// Get combat success probability
const double combatSuccessChance =
EstimateCombatSuccess(attackerPlayerId, guessedState, maxRounds);
EstimateCombatSuccess(attackerPlayerId, guessedState, settings);
// Consider fleeing if combat success chance is below threshold
return combatSuccessChance < fleeConsiderationThreshold;
@@ -9,11 +9,14 @@
#ifndef AIFleeDecisionCalculator_hpp
#define AIFleeDecisionCalculator_hpp
#include "src/main/cpp/net/eagle0/shardok/library/ShardokCommand.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokEngine.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
#include "src/main/protobuf/net/eagle0/shardok/api/command_descriptor.pb.h"
namespace shardok {
using CommandProto = net::eagle0::shardok::api::CommandDescriptor;
class AIFleeDecisionCalculator {
public:
// Configuration for flee decision thresholds
@@ -32,33 +35,31 @@ public:
// Evaluate whether to flee or fight in the final round
[[nodiscard]] static auto EvaluateFleeVsFight(
PlayerId playerId,
const SettingsGetter& settings,
const GameStateW& guessedState,
const CommandListSPtr& availableCommands,
const CommandList::const_iterator& fleeCommand,
int maxRounds,
int minimumFleeOddsThreshold,
int desperateFleeThreshold,
const vector<CommandProto>& availableCommands,
const vector<CommandProto>::const_iterator& fleeCommand,
bool enableDebugLogging = false) -> FleeDecision;
// Estimate probability of combat success for the attacker
[[nodiscard]] static auto EstimateCombatSuccess(
PlayerId attackerPlayerId,
const GameStateW& guessedState,
int maxRounds) -> double;
const SettingsGetter& settings) -> double;
// Determine if the attacker should consider fleeing based on combat odds
// Returns true if fleeing should be considered as an option
[[nodiscard]] static auto ShouldConsiderFleeing(
PlayerId attackerPlayerId,
const GameStateW& guessedState,
int maxRounds,
const SettingsGetter& settings,
double fleeConsiderationThreshold = 0.5) -> bool;
private:
// Helper to get flee command index
[[nodiscard]] static auto GetFleeCommandIndex(
const CommandList::const_iterator& fleeCommand,
const CommandListSPtr& availableCommands) -> size_t;
const vector<CommandProto>::const_iterator& fleeCommand,
const vector<CommandProto>& availableCommands) -> size_t;
};
} // namespace shardok
@@ -1,251 +0,0 @@
//
// Fast heuristic weighting implementation with context-aware logic
//
#include "AIHeuristicWeighting.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokException.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistances.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/util/HexMapUtils.hpp"
namespace shardok {
using CommandType = net::eagle0::shardok::common::CommandType;
using Coords = net::eagle0::shardok::storage::fb::Coords;
using ProtoCoords = net::eagle0::shardok::common::Coords;
double AIHeuristicWeighting::GetCommandWeight(
const CommandType commandType,
const UnitId actorUnitId,
const PlayerId actorPlayerId,
const Coords& targetCoords,
const GameStateW& state,
const CoordsSet& castleCoords,
const APDCache* apdCache,
bool isDefender,
std::function<BattalionTypeSPtr(BattalionTypeId)> getBattalionType) {
// Fast O(1) heuristic weights based on command type and game context
// Higher weight = more likely to select in simulation
// 0.0 = never select (filtered out)
const auto* hexMap = state->hex_map();
const auto* units = state->units();
const bool hasTarget = (targetCoords.row() >= 0 && targetCoords.column() >= 0);
switch (commandType) {
// === HIGH VALUE OFFENSIVE (10.0) ===
// Ranged attacks - very valuable, typically available when in range
case CommandType::ARCHERY_COMMAND: return 20.0;
case CommandType::LIGHTNING_BOLT_COMMAND: return 10.0;
case CommandType::FEAR_COMMAND: return 10.0;
// Area/tactical spells - high impact
case CommandType::METEOR_START_COMMAND: {
// METEOR_START doesn't have a target - it's based on actor location
if (hasTarget) {
throw ShardokInternalErrorException(
"METEOR_START_COMMAND should not have target coordinates");
}
// Get actor's location
const auto* actorUnit = units->Get(actorUnitId);
if (!actorUnit) {
throw ShardokInternalErrorException(
"METEOR_START_COMMAND actor unit not found in game state");
}
const Coords& actorLocation = actorUnit->location();
int enemyCount = 0;
// Count enemies within meteor range (3 hexes) of actor location
constexpr int METEOR_RANGE = 3;
const auto tilesInRange = TilesWithinDistance(hexMap, actorLocation, METEOR_RANGE);
for (const auto& tileCoords : tilesInRange) {
if (const auto* unit = Occupant(units, tileCoords)) {
if (unit->player_id() != actorPlayerId) { enemyCount++; }
}
}
return 1.0 + (enemyCount * 15.0); // Base 1 + 15 per enemy in range
}
case CommandType::METEOR_TARGET_COMMAND: {
// High weight per enemy unit at or adjacent to target
if (!hasTarget) {
throw ShardokInternalErrorException(
"METEOR_TARGET_COMMAND requires target coordinates for heuristic "
"weighting");
}
int enemyCount = 0;
// Count enemies at target
if (const auto* targetUnit = Occupant(units, targetCoords)) {
if (targetUnit->player_id() != actorPlayerId) { enemyCount++; }
}
// Count enemies adjacent to target
for (const auto& neighbor : HexMapUtils::GetAdjacentTiles(hexMap, targetCoords)) {
if (const auto* unit = Occupant(units, neighbor.coords)) {
if (unit->player_id() != actorPlayerId) { enemyCount++; }
}
}
return 1.0 + (enemyCount * 15.0); // Base 1 + 15 per enemy in range
}
case CommandType::RAISE_DEAD_COMMAND: return 10.0;
case CommandType::HOLY_WAVE_COMMAND: return 8.0;
// Fire on enemy (context-dependent)
case CommandType::START_FIRE_COMMAND: {
// High if enemy at target, low otherwise
if (!hasTarget) {
throw ShardokInternalErrorException(
"START_FIRE_COMMAND requires target coordinates for heuristic weighting");
}
if (const auto* targetUnit = Occupant(units, targetCoords)) {
if (targetUnit->player_id() != actorPlayerId) {
return 10.0; // Enemy at target - high value
}
}
return 1.0; // No enemy - low value but still valid
}
// === MEDIUM-HIGH OFFENSIVE (5.0-7.0) ===
// Direct damage melee
case CommandType::MELEE_COMMAND: return 7.0;
case CommandType::CHARGE_COMMAND: return 7.0; // Damage + movement
case CommandType::CHALLENGE_DUEL_COMMAND: return 5.0;
// Control and tactical magic
case CommandType::CONTROL_COMMAND: return 6.0;
case CommandType::METEOR_CAST_COMMAND: return 6.0; // Finish meteor
case CommandType::REDUCE_COMMAND: {
// High if enemy at target, zero otherwise
if (!hasTarget) return 0.0;
if (const auto* targetUnit = Occupant(units, targetCoords)) {
if (targetUnit->player_id() != actorPlayerId) {
return 10.0; // Enemy at target - very high value
}
}
return 0.0; // No enemy - don't use
}
// === MOVEMENT - Context-dependent ===
case CommandType::MOVE_COMMAND: {
if (isDefender) {
return 0.0; // Defenders don't move
}
// Attackers: weight based on distance improvement towards castle
if (!hasTarget) {
throw ShardokInternalErrorException(
"MOVE_COMMAND requires target coordinates for heuristic weighting");
}
// Get actor unit to determine battalion type and start position
const auto* actorUnit = units->Get(actorUnitId);
if (!actorUnit) return 4.0; // Default if can't find actor
// Get battalion type for distance calculation
const auto battalionTypeId = actorUnit->battalion().type();
const auto battalionTypePtr = getBattalionType(battalionTypeId);
if (!battalionTypePtr) return 4.0; // Default if can't get battalion type
// Get ActionPointDistances for this battalion type
const auto mapId = ActionPointDistancesCache::GetMapId(hexMap);
const auto* apd = (*apdCache)->GetRaw(hexMap, mapId, battalionTypePtr, false, -1);
if (!apd) return 4.0; // Default if can't get distances
// Calculate minimum distance from start to any castle
const Coords startCoords = actorUnit->location();
auto minStartDistance = ActionPointDistances::IMPOSSIBLE;
for (const auto& castleCoord : castleCoords) {
const auto dist = apd->Distance(startCoords, castleCoord);
if (dist < minStartDistance) { minStartDistance = dist; }
}
// Calculate minimum distance from end to any castle
const Coords& endCoords = targetCoords;
auto minEndDistance = ActionPointDistances::IMPOSSIBLE;
for (const auto& castleCoord : castleCoords) {
const auto dist = apd->Distance(endCoords, castleCoord);
if (dist < minEndDistance) { minEndDistance = dist; }
}
// Return weight based on distance improvement
// Higher weight if we're moving closer to castle
if (minStartDistance == ActionPointDistances::IMPOSSIBLE ||
minEndDistance == ActionPointDistances::IMPOSSIBLE) {
return 4.0; // Default if distances are impossible
}
const auto improvement = static_cast<double>(minStartDistance - minEndDistance);
return std::max(0.0, improvement);
}
case CommandType::BRAVE_WATER_COMMAND: return 3.0; // Tactical movement
case CommandType::SCOUT_COMMAND:
return 2.0; // Information gathering
// Terrain manipulation
case CommandType::FREEZE_WATER_COMMAND: return 3.0;
case CommandType::BUILD_BRIDGE_COMMAND: return 3.0;
// === LOW VALUE DEFENSIVE/UTILITY (1.0-2.0) ===
case CommandType::EXTINGUISH_FIRE_COMMAND: {
// High if friendly at target, low otherwise
if (!hasTarget) {
throw ShardokInternalErrorException(
"EXTINGUISH_FIRE_COMMAND requires target coordinates for heuristic "
"weighting");
}
if (const auto* targetUnit = Occupant(units, targetCoords)) {
if (targetUnit->player_id() == actorPlayerId) {
return 8.0; // Friendly at target - high value
}
}
return 1.0; // No friendly - low value but still valid
}
case CommandType::UNIT_REST_COMMAND: return 1.5;
case CommandType::FORTIFY_COMMAND: return 2.0;
// Zero weight - don't use in simulation
case CommandType::REPAIR_COMMAND: return 0.0;
case CommandType::HIDE_COMMAND: return 0.0;
case CommandType::RELEASE_UNIT_COMMAND: return 0.0;
case CommandType::REINFORCE_COMMAND: return 10.0;
case CommandType::MANAGE_PRISONER: return 1.0;
// === ZERO WEIGHT - NEVER SELECT (0.0) ===
// Explicitly bad actions
case CommandType::FLEE_COMMAND: return 0.0; // Never flee in simulation
case CommandType::RETREAT_COMMAND: return 0.0;
case CommandType::BECOME_OUTLAW_COMMAND: return 0.0; // Never become outlaw
case CommandType::DISMISS_UNIT_COMMAND:
return 0.0; // Never dismiss in combat
// Actions that are fine as a fallback
case CommandType::END_TURN_COMMAND: return 1.0;
case CommandType::UNIT_STOP_COMMAND: return 1.0;
case CommandType::METEOR_CANCEL_COMMAND: return 1.0;
// Setup commands (shouldn't appear in combat, but filter anyway)
case CommandType::PLACE_UNIT_COMMAND: return 10.0;
case CommandType::PLACE_HIDDEN_UNIT_COMMAND: return 1.0;
case CommandType::END_PLAYER_SETUP_COMMAND: return 1.0;
// Unknown/unhandled
case CommandType::UNKNOWN_COMMAND:
default: return 0.0; // Don't select unknown commands
}
}
} // namespace shardok
@@ -1,40 +0,0 @@
//
// Fast heuristic weighting for MCTS simulations
// Provides O(1) weights based on command type and context
//
#ifndef EAGLE0_AI_HEURISTIC_WEIGHTING_HPP
#define EAGLE0_AI_HEURISTIC_WEIGHTING_HPP
#pragma clang diagnostic push
#pragma clang diagnostic ignored "-Wdeprecated-redundant-constexpr-static-def"
#include "src/main/protobuf/net/eagle0/shardok/common/command_type.pb.h"
#pragma clang diagnostic pop
#include "src/main/cpp/net/eagle0/shardok/library/GameStateW.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistancesCache.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/map/CoordsSet.hpp"
namespace shardok {
// Fast heuristic-based command weighting for MCTS simulation policy
// Avoids expensive score calculation while maintaining intelligent bias
class AIHeuristicWeighting {
public:
// Get weight for a command using fast heuristics with game context
// Returns weight >= 0.0, where 0.0 means "never select" and higher is more likely
static double GetCommandWeight(
net::eagle0::shardok::common::CommandType commandType,
UnitId actorUnitId,
PlayerId actorPlayerId,
const Coords& targetCoords,
const GameStateW& state,
const CoordsSet& castleCoords,
const APDCache* apdCache,
bool isDefender,
std::function<BattalionTypeSPtr(BattalionTypeId)> getBattalionType);
};
} // namespace shardok
#endif // EAGLE0_AI_HEURISTIC_WEIGHTING_HPP
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,64 @@
//
// Created by dancrosby on 3/4/20.
//
#ifndef EAGLE0_AISCORECALCULATOR_HPP
#define EAGLE0_AISCORECALCULATOR_HPP
#include <chrono>
#include <future>
#include "src/main/cpp/net/eagle0/shardok/ai/AIAttackLocations.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIStrategy.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokCTypes.h"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokEngine.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistancesCache.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
#include "src/main/flatbuffer/net/eagle0/shardok/storage/game_state.hpp"
#include "src/main/protobuf/net/eagle0/shardok/api/command_descriptor.pb.h"
#include "src/main/protobuf/net/eagle0/shardok/api/game_state_view.pb.h"
namespace shardok {
using net::eagle0::shardok::api::GameStateView;
using GameState = fb::GameState;
using shardok::PlayerId;
using std::future;
using std::vector;
using ScoreValue = double;
using CommandProto = net::eagle0::shardok::api::CommandDescriptor;
class AIScoreCalculator {
public:
// Evaluate the score of a guessed game state based on the current AI strategy. DOES NOT perform
// or evaluate any commands.
[[nodiscard]] static auto GuessedStateScore(
bool isDefender,
const GameStateW &state,
const AIStrategy &aiStrategy,
const CoordsSet &allCastleCoords,
const SettingsGetter &settingsGetter,
const APDCache &apdCache,
const ALCache &alCache) -> ScoreValue;
// Evaluates the score for a particular command index for the given player, using lookahead.
[[nodiscard]] static auto CommandScore(
PlayerId pid,
bool isDefender,
int remainingLookahead,
int maxRepeatCount,
const ShardokEngine &guessedEngine,
const AIStrategy &attackerStrategy,
ScoreValue currentUtility,
const SettingsGetter &settingsGetter,
const CoordsSet &allCastleCoords,
const APDCache &apdCache,
const ALCache &alCache,
size_t commandIndex,
std::chrono::steady_clock::time_point deadline) -> std::future<ScoreValue>;
};
} // namespace shardok
#endif // EAGLE0_AISCORECALCULATOR_HPP
@@ -3,7 +3,6 @@
//
#include "src/main/cpp/net/eagle0/shardok/ai/AIStrategy.hpp"
namespace shardok {
AIStrategy FleeStrategy = AIStrategy{AIStrategy::STRATEGY_FLEE};
AIStrategy HoldCastlesStrategy = AIStrategy{AIStrategy::STRATEGY_HOLD_CASTLES};
@@ -24,40 +24,10 @@ int AIEvaluationCounter::GetCurrentCount() { return activeCount.load(); }
auto CalculateTimeBudget(
const PlayerId playerId,
const GameSettingsSPtr &settings,
const GameStateW &state,
const size_t numCommands) -> AITimeBudget {
const GameStateW &state) -> AITimeBudget {
const auto settingsGetter = settings->GetGetter();
const auto castleCoords = AllCastleCoords(state->hex_map());
// Check if we're in setup phase
const bool isSetupPhase = state->status()->state() ==
net::eagle0::shardok::storage::fb::GameStatus_::State_SET_UP;
// Get maximum budget cap from settings (in seconds)
const double maxBudgetSeconds =
settingsGetter.Backing().lookahead_time_budget_maximum_seconds();
const double maxBudgetMs = maxBudgetSeconds * 1000.0;
// During setup, use the setup-specific time budget
if (isSetupPhase) {
// Dynamic budget: msPerCommand × numCommands
const double msPerCommand =
settingsGetter.Backing().lookahead_time_budget_per_command_setup_ms();
const double budgetMs = msPerCommand * static_cast<double>(numCommands);
// Clamp to reasonable bounds: 200ms minimum, maxBudgetMs maximum
const auto clampedBudgetMs = std::clamp(budgetMs, 200.0, maxBudgetMs);
const auto remainingBudget =
std::chrono::milliseconds(static_cast<int64_t>(clampedBudgetMs));
const size_t minDepth = settingsGetter.Backing().min_lookahead_turns();
return AITimeBudget{
.remainingBudget = remainingBudget,
.minDepthRequired = minDepth,
.isCloseToEnemy = false}; // Not relevant during setup
}
// Determine proximity (≤4 hex distance) - applies to both attackers and defenders
bool isClose = false;
const auto *units = state->units();
@@ -102,25 +72,12 @@ auto CalculateTimeBudget(
}
}
// Get time budget from settings - dynamic based on number of commands
// Dynamic budget: msPerCommand × numCommands
const double msPerCommand =
isClose ? settingsGetter.Backing().lookahead_time_budget_per_command_close_ms()
: settingsGetter.Backing().lookahead_time_budget_per_command_far_ms();
const double budgetMs = msPerCommand * static_cast<double>(numCommands);
// Get time budget from settings
const auto budget = std::chrono::duration<double>(
isClose ? settingsGetter.Backing().lookahead_time_budget_close_in_seconds()
: settingsGetter.Backing().lookahead_time_budget_far_in_seconds());
// Clamp to reasonable bounds: 200ms minimum, maxBudgetMs maximum
const auto clampedBudgetMs = std::clamp(budgetMs, 200.0, maxBudgetMs);
const auto remainingBudget = std::chrono::milliseconds(static_cast<int64_t>(clampedBudgetMs));
// TEMPORARY DEBUG OUTPUT
printf("[DEBUG CalculateTimeBudget] numCommands=%zu, msPerCommand=%.2f, budgetMs=%.2f, "
"clampedBudgetMs=%.2f, isClose=%d\n",
numCommands,
msPerCommand,
budgetMs,
clampedBudgetMs,
isClose);
const auto remainingBudget = std::chrono::duration_cast<std::chrono::milliseconds>(budget);
// Get minimum depth requirement
const size_t minDepth = settingsGetter.Backing().min_lookahead_turns();
@@ -36,13 +36,10 @@ struct AITimeBudget {
};
// Calculate time budget based on proximity to enemies and castles
// Time budget is calculated dynamically based on number of available commands:
// budget = msPerCommand × numCommands (clamped to 200-5000ms)
auto CalculateTimeBudget(
PlayerId playerId,
const GameSettingsSPtr &settings,
const GameStateW &state,
size_t numCommands) -> AITimeBudget;
const GameStateW &state) -> AITimeBudget;
} // namespace shardok
@@ -17,10 +17,9 @@ using std::end;
using std::shared_ptr;
constexpr double kProfessionValue = 200;
constexpr double kVigorScoreMultiplier = 5.0;
constexpr double kCastleMultiplierBonus = 1.0;
constexpr double kOnFireMultiplier = 0.25;
constexpr double kAdjacentFireMultiplier = 0.80;
constexpr double kAdjacentFireMultiplier = 0.99;
constexpr double kOnIceMultiplier = 0.25;
constexpr double kMeteorStartInRangeValue = 50;
constexpr double kMeteorDirectTargetingEnemy = 2;
@@ -64,8 +63,7 @@ auto ContextFreeUnitValue(const Unit *unit) -> ScoreValue {
4.0;
}
const double vigorValue =
unit->has_attached_hero() ? unit->attached_hero().vigor() * kVigorScoreMultiplier : 0.0;
const double vigorValue = unit->has_attached_hero() ? unit->attached_hero().vigor() : 0.0;
double battalionTypeMultiplier = 1.0;
switch (unit->battalion().type()) {
@@ -337,8 +335,7 @@ auto UnitValue(
const AttackLocations &locationsThisSideCanAttackFrom,
const CoordsSet &locationsInDangerFromEnemy,
const ActionPointDistances *distances,
int meteorRange,
double meteorCastVigorCost) -> ScoreValue {
const SettingsGetter &settings) -> ScoreValue {
const auto &location = unit->location();
if (location.row() < 0) return 0; // unplaced unit
@@ -357,7 +354,9 @@ auto UnitValue(
kCastleMultiplierBonus * (terrain->modifier().castle().integrity() + 25) / 100.0;
}
double onFireMultiplier = 1.0;
if (terrain->modifier().fire().present()) { onFireMultiplier *= kOnFireMultiplier; }
if (terrain->modifier().fire().present() && (isAttacker || attackerWantsCastles)) {
onFireMultiplier *= kOnFireMultiplier;
}
{
for (const auto adjacentCoords = HexMapUtils::GetAdjacentCoords(map, location);
const auto &c : adjacentCoords) {
@@ -381,8 +380,8 @@ auto UnitValue(
roundsRemaining,
attackerUnits,
defenderUnits,
meteorRange,
meteorCastVigorCost);
settings.Backing().meteor_range(),
settings.Backing().meteor_cast_vigor_cost());
// scouting values
// attack range
@@ -46,8 +46,7 @@ auto UnitValue(
const AttackLocations &locationsThisSideCanAttackFrom,
const CoordsSet &locationsInDangerFromEnemy,
const ActionPointDistances *distances,
int meteorRange,
double meteorCastVigorCost) -> ScoreValue;
const SettingsGetter &settings) -> ScoreValue;
} // namespace shardok
@@ -7,9 +7,9 @@
#include <algorithm>
#include <ranges>
#include "AIAttackLocations.hpp"
#include "AIDistanceDebuf.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIAttackGroups.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIAttackLocations.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIDistanceDebuf.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/util/HexMapUtils.hpp"
#include "src/main/flatbuffer/net/eagle0/shardok/storage/victory_condition.hpp"
@@ -43,8 +43,8 @@ auto AttackerDebufForOnFireCriticalTile(
const vector<const Unit*>& extinguishingUnits,
const APDCache& apdCache,
const ALCache& alCache,
const BattalionTypeGetter& battalionTypeGetter,
ActionPoints braveWaterCost,
const SettingsGetter& settings,
const int braveWaterActionPointCost,
const bool lateGame) -> double {
double minDebuf = 99999.9;
@@ -60,8 +60,8 @@ auto AttackerDebufForOnFireCriticalTile(
extinguishingUnits,
apdCache,
alCache,
battalionTypeGetter,
braveWaterCost,
settings,
braveWaterActionPointCost,
lateGame,
/* includeUndead = */ false);
if (newDebuf < minDebuf) minDebuf = newDebuf;
@@ -77,8 +77,8 @@ auto AttackerDebufForUnoccupiedCriticalTile(
const vector<const Unit*>& claimableUnits,
const APDCache& apdCache,
const ALCache& alCache,
const BattalionTypeGetter& battalionTypeGetter,
ActionPoints braveWaterCost,
const SettingsGetter& settings,
const int braveWaterActionPointCost,
const bool lateGame) -> double {
return UNHELD_VALUE * DefenderDistanceBuf(
criticalTileLocation,
@@ -87,8 +87,8 @@ auto AttackerDebufForUnoccupiedCriticalTile(
claimableUnits,
apdCache,
alCache,
battalionTypeGetter,
braveWaterCost,
settings,
braveWaterActionPointCost,
lateGame,
/* includeUndead = */ false);
}
@@ -100,8 +100,8 @@ auto AttackerDebufForDefenderOccupiedCriticalTile(
const vector<const Unit*>& attackerUnits,
const APDCache& apdCache,
const ALCache& alCache,
const BattalionTypeGetter& battalionTypeGetter,
ActionPoints braveWaterCost,
const SettingsGetter& settings,
const int braveWaterActionPointCost,
const bool lateGame) {
const double baseUnitValue =
defenderUnit->battalion().size() +
@@ -117,8 +117,8 @@ auto AttackerDebufForDefenderOccupiedCriticalTile(
attackerUnits,
apdCache,
alCache,
battalionTypeGetter,
braveWaterCost,
settings,
braveWaterActionPointCost,
lateGame,
/* includeUndead = */ false);
}
@@ -126,7 +126,10 @@ auto AttackerDebufForDefenderOccupiedCriticalTile(
auto DefenderHoldsCriticalTilesVictoryScore(
const GameStateW& gameState,
const CoordsSet& criticalTileLocations,
const PlayerInfo* player) -> ScoreValue {
const PlayerInfo* player,
const APDCache& /*apdCache*/,
const ALCache& /*alCache*/,
const SettingsGetter& /*settings*/) -> ScoreValue {
ScoreValue total = 0.0;
const auto rc = gameState->hex_map()->row_count();
@@ -156,8 +159,7 @@ auto AttackerHoldsCriticalTilesVictoryScore(
const PlayerInfo* player,
const APDCache& apdCache,
const ALCache& alCache,
const BattalionTypeGetter& battalionTypeGetter,
ActionPoints braveWaterCost) -> ScoreValue {
const SettingsGetter& settings) -> ScoreValue {
vector<const Unit*> playerUnits{};
vector<const Unit*> claimablePlayerUnits{};
for (const Unit* unit : *gameState->units()) {
@@ -173,6 +175,7 @@ auto AttackerHoldsCriticalTilesVictoryScore(
return criticalTileLocations.size() * MAX_DEFENDER_HELD_VALUE;
}
const int braveWaterActionPointCost = settings.Backing().brave_water_action_point_cost();
const MapId mapId = apdCache->GetMapId(gameState->hex_map());
ScoreValue total = 0.0;
@@ -199,8 +202,8 @@ auto AttackerHoldsCriticalTilesVictoryScore(
claimablePlayerUnits,
apdCache,
alCache,
battalionTypeGetter,
braveWaterCost,
settings,
braveWaterActionPointCost,
IsLateGame(gameState));
total += BADLY_HELD_VALUE;
}
@@ -212,8 +215,8 @@ auto AttackerHoldsCriticalTilesVictoryScore(
playerUnits,
apdCache,
alCache,
battalionTypeGetter,
braveWaterCost,
settings,
braveWaterActionPointCost,
IsLateGame(gameState));
}
} else if (terrain->modifier().fire().present()) {
@@ -224,8 +227,8 @@ auto AttackerHoldsCriticalTilesVictoryScore(
claimablePlayerUnits,
apdCache,
alCache,
battalionTypeGetter,
braveWaterCost,
settings,
braveWaterActionPointCost,
IsLateGame(gameState));
} else {
total -= AttackerDebufForUnoccupiedCriticalTile(
@@ -235,8 +238,8 @@ auto AttackerHoldsCriticalTilesVictoryScore(
claimablePlayerUnits,
apdCache,
alCache,
battalionTypeGetter,
braveWaterCost,
settings,
braveWaterActionPointCost,
IsLateGame(gameState));
}
}
@@ -249,8 +252,7 @@ auto LastPlayerStandingVictoryScore(
const PlayerInfo* player,
const APDCache& apdCache,
const ALCache& alCache,
const BattalionTypeGetter& battalionTypeGetter,
ActionPoints braveWaterCost) -> ScoreValue {
const SettingsGetter& settings) -> ScoreValue {
if (!std::ranges::contains(
*player->victory_conditions(),
net::eagle0::shardok::storage::fb::
@@ -283,8 +285,8 @@ auto LastPlayerStandingVictoryScore(
playerUnits,
apdCache,
alCache,
battalionTypeGetter,
braveWaterCost,
settings,
5,
IsLateGame(gameState),
/* includeUndead = */ true);
}
@@ -9,7 +9,6 @@
#include "src/main/cpp/net/eagle0/shardok/ai/AIAttackGroups.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIAttackLocations.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AICommonTypes.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/GameStateW.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistancesCache.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/map/CoordsSet.hpp"
@@ -30,21 +29,22 @@ auto AttackerHoldsCriticalTilesVictoryScore(
const PlayerInfo* player,
const APDCache& apdCache,
const ALCache& alCache,
const BattalionTypeGetter& battalionTypeGetter,
ActionPoints braveWaterCost) -> ScoreValue;
const SettingsGetter& settings) -> ScoreValue;
auto DefenderHoldsCriticalTilesVictoryScore(
const GameStateW& gameState,
const CoordsSet& criticalTileLocations,
const PlayerInfo* player) -> ScoreValue;
const PlayerInfo* player,
const APDCache& apdCache,
const ALCache& alCache,
const SettingsGetter& settings) -> ScoreValue;
auto LastPlayerStandingVictoryScore(
const GameStateW& gameState,
const PlayerInfo* player,
const APDCache& apdCache,
const ALCache& alCache,
const BattalionTypeGetter& battalionTypeGetter,
ActionPoints braveWaterCost) -> ScoreValue;
const SettingsGetter& settings) -> ScoreValue;
} // namespace shardok
@@ -15,7 +15,7 @@ auto UnitIdsRequiringWaterCrossing(
const PlayerId pid,
const CoordsSet &destinations,
const APDCache &apdCache,
const BattalionTypeGetter &battalionTypeGetter) -> vector<UnitId> {
const SettingsGetter &settings) -> vector<UnitId> {
// Put out all the fires, except on bridges
fb::HexMapW mapCopy = fb::CopyHexMap(gameState->hex_map());
for (uint32_t index = 0; index < mapCopy->terrain()->size(); index++) {
@@ -36,7 +36,7 @@ auto UnitIdsRequiringWaterCrossing(
for (const auto *unit : *gameState->units()) {
if (unit->player_id() != pid) continue;
const auto &battType = battalionTypeGetter(unit->battalion().type());
const auto &battType = settings.GetBattalionType(unit->battalion().type());
if (unit->status() == net::eagle0::shardok::storage::fb::UnitStatus_NORMAL_UNIT) {
for (const Coords &destination : destinations) {
@@ -76,7 +76,8 @@ auto UnitIdsRequiringWaterCrossing(
auto UnitIdsToCreateWaterCrossing(
const GameStateW &gameState,
const PlayerId pid,
const BattalionTypeGetter &battalionTypeGetter) -> vector<UnitId> {
const APDCache & /*apdCache*/,
const SettingsGetter &settings) -> vector<UnitId> {
vector<UnitId> unitIds{};
for (const auto *unit : *gameState->units()) {
@@ -87,7 +88,7 @@ auto UnitIdsToCreateWaterCrossing(
if (!unit->has_attached_hero()) continue;
const auto profession = unit->attached_hero().profession_info().profession();
const auto &battalionType = battalionTypeGetter(unit->battalion().type());
const auto &battalionType = settings.GetBattalionType(unit->battalion().type());
if (profession == net::eagle0::shardok::storage::fb::Profession_ENGINEER ||
(profession == net::eagle0::shardok::storage::fb::Profession_MAGE &&
@@ -198,14 +199,14 @@ auto IntendedCrossingStarts(
const GameStateW &gameState,
const vector<UnitId> &unitIdsCreatingCrossing,
const CoordsSet &tilesToStartCrossingFrom,
const MapId &mapId,
const APDCache &apdCache,
const BattalionTypeGetter &battalionTypeGetter) -> CoordsSet {
const SettingsGetter &settings) -> CoordsSet {
CoordsSet intendedCrossingStarts(gameState->hex_map());
const MapId mapId = apdCache->GetMapId(gameState->hex_map());
for (const UnitId uid : unitIdsCreatingCrossing) {
const Unit *unit = gameState->units()->Get(uid);
const Coords &location = unit->location();
const auto &battalionType = battalionTypeGetter(unit->battalion().type());
const auto &battalionType = settings.GetBattalionType(unit->battalion().type());
const auto *apd = apdCache->GetRaw(gameState->hex_map(), mapId, battalionType, false);
if (location.row() >= 0) {
@@ -218,111 +219,4 @@ auto IntendedCrossingStarts(
return intendedCrossingStarts;
}
using Unit = net::eagle0::shardok::storage::fb::Unit;
constexpr double kNoRequiredCrossingScore = std::numeric_limits<double>::max();
constexpr double kNoCrossingCreatorsScore = std::numeric_limits<double>::min();
auto WaterCrossingScore(
const PlayerId playerId,
const BattalionTypeGetter &battalionTypeGetter,
const GameStateW &gameState,
const CoordsSet &castleCoords,
const CoordsSet &startCrossingFrom,
const APDCache &apdCache) -> double {
uint32_t castleClaimCount = 0;
for (const auto *unit : *gameState->units()) {
if (unit->player_id() != playerId) continue;
const auto status = unit->status();
if (status != net::eagle0::shardok::storage::fb::UnitStatus_NORMAL_UNIT &&
status != net::eagle0::shardok::storage::fb::UnitStatus_RESERVE_UNIT)
continue;
if (!unit->has_attached_hero()) continue;
++castleClaimCount;
}
CoordsSet destinations = castleCoords;
if (castleClaimCount < castleCoords.size()) {
destinations = CoordsSet(gameState->hex_map());
for (const auto *enemyUnit : *gameState->units()) {
if (enemyUnit->player_id() == playerId) continue;
const auto status = enemyUnit->status();
if (status != net::eagle0::shardok::storage::fb::UnitStatus_NORMAL_UNIT) continue;
AssertValid(enemyUnit->location(), gameState->hex_map());
destinations.Add(enemyUnit->location());
}
}
const auto unitIdsRequiringCrossing = UnitIdsRequiringWaterCrossing(
gameState,
playerId,
castleCoords,
apdCache,
battalionTypeGetter);
if (unitIdsRequiringCrossing.empty()) return kNoRequiredCrossingScore;
const auto unitIdsCreatingCrossing =
UnitIdsToCreateWaterCrossing(gameState, playerId, battalionTypeGetter);
if (unitIdsCreatingCrossing.empty()) return kNoCrossingCreatorsScore;
double totalScore = 0;
const auto mapId = ActionPointDistancesCache::GetMapId(gameState->hex_map());
// First put a big penalty on the distance for units that can create a crossing
for (const UnitId uid : unitIdsCreatingCrossing) {
const Unit *unit = gameState->units()->Get(uid);
const auto &battalionType = battalionTypeGetter(unit->battalion().type());
Coords location = unit->location();
int thisDistance;
if (location.row() < 0) thisDistance = 1000;
else {
const auto *apd = apdCache->GetRaw(gameState->hex_map(), mapId, battalionType, false);
thisDistance = MinimumDistance(apd, location, startCrossingFrom);
}
totalScore -= thisDistance * 100.0;
}
// Now a smaller penalty for distance for units that need to cross, except if they block -- then
// a large penalty
for (const UnitId uid : unitIdsRequiringCrossing) {
// If this unit ID can also create a crossing, we already handled it
if (std::ranges::contains(unitIdsCreatingCrossing, uid)) continue;
const Unit *unit = gameState->units()->Get(uid);
const auto &battalionType = battalionTypeGetter(unit->battalion().type());
Coords location = unit->location();
const auto *apd = apdCache->GetRaw(gameState->hex_map(), mapId, battalionType, false);
int thisDistance;
if (location.row() < 0) thisDistance = 1000;
else { thisDistance = MinimumDistance(apd, location, startCrossingFrom); }
bool targetBlocks = false;
// If we're not capable of creating a crossing, don't get in the way of somebody that is.
for (const UnitId crossingUid : unitIdsCreatingCrossing) {
const auto *crossingCapableUnit = gameState->units()->Get(crossingUid);
// Don't check for units that aren't yet placed
if (crossingCapableUnit->location().row() < 0) continue;
AssertValid(crossingCapableUnit->location(), gameState->hex_map());
if (thisDistance <
MinimumDistance(apd, crossingCapableUnit->location(), startCrossingFrom)) {
targetBlocks = true;
break;
}
}
if (targetBlocks) continue;
totalScore -= thisDistance;
}
return totalScore;
}
} // namespace shardok
@@ -5,7 +5,6 @@
#ifndef EAGLE0_AIWATERCROSSINGCALCULATOR_HPP
#define EAGLE0_AIWATERCROSSINGCALCULATOR_HPP
#include "src/main/cpp/net/eagle0/shardok/ai/AICommonTypes.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/GameStateW.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistancesCache.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/map/CoordsSet.hpp"
@@ -35,13 +34,14 @@ auto UnitIdsRequiringWaterCrossing(
PlayerId pid,
const CoordsSet& destinations,
const APDCache& apdCache,
const BattalionTypeGetter& battalionTypeGetter) -> vector<UnitId>;
const SettingsGetter& settings) -> vector<UnitId>;
// Units belonging to the player that are capable of creating water crossings
auto UnitIdsToCreateWaterCrossing(
const GameStateW& gameState,
PlayerId pid,
const BattalionTypeGetter& battalionTypeGetter) -> vector<UnitId>;
const APDCache& apdCache,
const SettingsGetter& settings) -> vector<UnitId>;
// Whether a unit of the given type can reach destination from origin, given the current state
// of the map
@@ -71,17 +71,9 @@ auto IntendedCrossingStarts(
const GameStateW& gameState,
const vector<UnitId>& unitIdsCreatingCrossing,
const CoordsSet& tilesToStartCrossingFrom,
const MapId& mapId,
const APDCache& apdCache,
const BattalionTypeGetter& battalionTypeGetter) -> CoordsSet;
// Calculate score based on water crossing strategy
auto WaterCrossingScore(
PlayerId playerId,
const BattalionTypeGetter& battalionTypeGetter,
const GameStateW& gameState,
const CoordsSet& castleCoords,
const CoordsSet& startCrossingFrom,
const APDCache& apdCache) -> double;
const SettingsGetter& settings) -> CoordsSet;
} // namespace shardok
@@ -18,7 +18,7 @@ constexpr ScoreValue kNoRequiredCrossingScore = std::numeric_limits<ScoreValue>:
constexpr ScoreValue kNoCrossingCreatorsScore = std::numeric_limits<ScoreValue>::min();
[[nodiscard]] auto AIWaterCrossingCommandChooser::WaterCrossingScore(
const BattalionTypeGetter &battalionTypeGetter,
const SettingsGetter &settingsGetter,
const GameStateW &gameState,
const CoordsSet &castleCoords,
const CoordsSet &startCrossingFrom) const -> ScoreValue {
@@ -51,13 +51,15 @@ constexpr ScoreValue kNoCrossingCreatorsScore = std::numeric_limits<ScoreValue>:
playerId,
castleCoords,
apdCache,
battalionTypeGetter);
settingsGetter);
if (unitIdsRequiringCrossing.empty()) return kNoRequiredCrossingScore;
const auto unitIdsCreatingCrossing =
UnitIdsToCreateWaterCrossing(gameState, playerId, battalionTypeGetter);
UnitIdsToCreateWaterCrossing(gameState, playerId, apdCache, settingsGetter);
if (unitIdsCreatingCrossing.empty()) return kNoCrossingCreatorsScore;
fprintf(stderr, "%lu units require a water crossing\n", unitIdsRequiringCrossing.size());
ScoreValue totalScore = 0;
const auto mapId = ActionPointDistancesCache::GetMapId(gameState->hex_map());
@@ -65,7 +67,7 @@ constexpr ScoreValue kNoCrossingCreatorsScore = std::numeric_limits<ScoreValue>:
// First put a big penalty on the distance for units that can create a crossing
for (const UnitId uid : unitIdsCreatingCrossing) {
const Unit *unit = gameState->units()->Get(uid);
const auto &battalionType = battalionTypeGetter(unit->battalion().type());
const auto &battalionType = settingsGetter.GetBattalionType(unit->battalion().type());
Coords location = unit->location();
int thisDistance;
@@ -86,7 +88,7 @@ constexpr ScoreValue kNoCrossingCreatorsScore = std::numeric_limits<ScoreValue>:
if (std::ranges::contains(unitIdsCreatingCrossing, uid)) continue;
const Unit *unit = gameState->units()->Get(uid);
const auto &battalionType = battalionTypeGetter(unit->battalion().type());
const auto &battalionType = settingsGetter.GetBattalionType(unit->battalion().type());
Coords location = unit->location();
const auto *apd = apdCache->GetRaw(gameState->hex_map(), mapId, battalionType, false);
@@ -118,7 +120,7 @@ constexpr ScoreValue kNoCrossingCreatorsScore = std::numeric_limits<ScoreValue>:
}
auto AIWaterCrossingCommandChooser::StartCrossingFrom(
const BattalionTypeGetter &battalionTypeGetter,
const SettingsGetter &settingsGetter,
const GameStateW &gameState,
const CoordsSet &castleCoords) const -> CoordsSet {
CoordsSet startCrossingFrom(gameState->hex_map());
@@ -152,16 +154,16 @@ auto AIWaterCrossingCommandChooser::StartCrossingFrom(
playerId,
castleCoords,
apdCache,
battalionTypeGetter);
settingsGetter);
if (unitIdsRequiringCrossing.empty()) return startCrossingFrom;
const auto unitIdsCreatingCrossing =
UnitIdsToCreateWaterCrossing(gameState, playerId, battalionTypeGetter);
UnitIdsToCreateWaterCrossing(gameState, playerId, apdCache, settingsGetter);
if (unitIdsCreatingCrossing.empty()) return startCrossingFrom;
for (const UnitId uid : unitIdsRequiringCrossing) {
const Unit *unit = gameState->units()->Get(uid);
const auto &battalionType = battalionTypeGetter(unit->battalion().type());
const auto &battalionType = settingsGetter.GetBattalionType(unit->battalion().type());
Coords origin = unit->location();
// FIXME: this is just grabbing the first starting position, ideally we'd try them all
@@ -6,16 +6,19 @@
#define EAGLE0_AIWATERCROSSINGCOMMANDCHOOSER_HPP
#include <utility>
#include <vector>
#include "src/main/cpp/net/eagle0/shardok/ai/AICommonTypes.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/GameStateW.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistancesCache.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/fb_helpers/FlatbufferWrapper.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/map/CoordsSet.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
#include "src/main/flatbuffer/net/eagle0/shardok/storage/game_state.hpp"
#include "src/main/protobuf/net/eagle0/shardok/api/command_descriptor.pb.h"
namespace shardok {
using CommandProto = net::eagle0::shardok::api::CommandDescriptor;
using GameState = net::eagle0::shardok::storage::fb::GameState;
using Unit = net::eagle0::shardok::storage::fb::Unit;
using ScoreValue = double;
@@ -30,13 +33,13 @@ public:
: playerId(pid),
apdCache(std::move(apdCache)) {}
[[nodiscard]] auto StartCrossingFrom(
const BattalionTypeGetter &battalionTypeGetter,
auto StartCrossingFrom(
const SettingsGetter &settingsGetter,
const GameStateW &gameState,
const CoordsSet &castleCoords) const -> CoordsSet;
[[nodiscard]] auto WaterCrossingScore(
const BattalionTypeGetter &battalionTypeGetter,
const SettingsGetter &settingsGetter,
const GameStateW &gameState,
const CoordsSet &castleCoords,
const CoordsSet &startCrossingFrom) const -> ScoreValue;
@@ -619,7 +619,6 @@ auto result = mctsAI.Search(settings, state, commands, budget);
```
#### Algorithm Comparison
| Feature | Iterative Deepening | MCTS |
|---------|-------------------|------|
| **Randomness Handling** | Sophisticated (chance nodes, multi-sample) | Simplified (average rolls) |
@@ -636,6 +635,7 @@ The MCTS implementation provides a solid foundation. Known limitations:
Note: The APD cache is fully thread-safe using thread-local storage and mutex-protected shared cache.
<<<<<<< HEAD
Adding chance node handling and ensuring thread safety would make it a superior replacement for the iterative deepening approach while maintaining the sophisticated randomness evaluation that makes the current system effective.
## MCTS Configuration Options
@@ -762,4 +762,4 @@ MCTSAI ai(playerId, isDefender, strategy, castleCoords, apdCache, alCache, confi
- Higher `immediateScoreTieBreakThreshold` to emphasize direct paths
- `BEST_IMMEDIATE` simulation for most predictable behavior
The configuration system allows fine-tuning MCTS behavior for different scenarios while maintaining compatibility with the existing AI infrastructure.
The configuration system allows fine-tuning MCTS behavior for different scenarios while maintaining compatibility with the existing AI infrastructure.
+91 -80
View File
@@ -1,16 +1,5 @@
load("//tools:copts.bzl", "COPTS")
cc_library(
name = "ai_common_types",
hdrs = ["AICommonTypes.hpp"],
copts = COPTS,
visibility = ["//visibility:public"],
deps = [
"//src/main/cpp/net/eagle0/shardok/library:battalion_type",
"//src/main/cpp/net/eagle0/shardok/library:shardok_c_types",
],
)
cc_library(
name = "ai_attacker_strategy_selector",
srcs = ["AIAttackerStrategySelector.cpp"],
@@ -39,15 +28,14 @@ cc_library(
hdrs = ["AIAttackGroups.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai:__subpackages__",
"//src/main/cpp/net/eagle0/shardok/ai_performance_runner:__pkg__",
"//src/test/cpp/net/eagle0/shardok/ai:__subpackages__",
"//src/test/cpp/net/eagle0/shardok/ai:__pkg__",
],
deps = [
":ai_attack_locations",
":ai_common_types",
":ai_score_utilities",
"//src/main/cpp/net/eagle0/shardok/library/action_point_distances:action_point_distances_cache",
"//src/main/cpp/net/eagle0/shardok/library/settings:game_settings",
"//src/main/cpp/net/eagle0/shardok/library/util:hex_map_utils",
"//src/main/flatbuffer/net/eagle0/shardok/storage:hex_map_cc_fbs",
"//src/main/flatbuffer/net/eagle0/shardok/storage:unit_cc_fbs",
@@ -59,10 +47,6 @@ cc_library(
srcs = ["AIAttackLocations.cpp"],
hdrs = ["AIAttackLocations.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai:__subpackages__",
"//src/test/cpp/net/eagle0/shardok/ai:__subpackages__",
],
deps = [
":ai_score_utilities",
"//src/main/cpp/net/eagle0/shardok/library/map:terrain",
@@ -101,14 +85,11 @@ cc_library(
hdrs = ["AIDistanceDebuf.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai:__subpackages__",
"//src/main/cpp/net/eagle0/shardok/ai_performance_runner:__pkg__",
"//src/test/cpp/net/eagle0/shardok/ai:__pkg__",
"//src/test/cpp/net/eagle0/shardok/ai:__subpackages__",
],
deps = [
":ai_attack_locations",
":ai_common_types",
":ai_score_utilities",
"//src/main/cpp/net/eagle0/shardok/library/action_point_distances",
"//src/main/cpp/net/eagle0/shardok/library/action_point_distances:action_point_distances_cache",
@@ -137,10 +118,8 @@ cc_library(
hdrs = ["AIScoreUtilities.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai:__subpackages__",
"//src/main/cpp/net/eagle0/shardok/ai_performance_runner:__pkg__",
"//src/test/cpp/net/eagle0/shardok/ai:__pkg__",
"//src/test/cpp/net/eagle0/shardok/ai:__subpackages__",
],
deps = [
"//src/main/cpp/net/eagle0/shardok/library:game_state_w",
@@ -163,48 +142,9 @@ cc_library(
":ai_score_utilities",
":ai_unit_score_calculator",
"//src/main/cpp/net/eagle0/shardok/library:engine",
"//src/main/cpp/net/eagle0/shardok/library/settings:game_settings",
"//src/main/cpp/net/eagle0/shardok/library/util:hex_map_utils",
],
)
cc_library(
name = "ai_heuristic_weighting",
srcs = ["AIHeuristicWeighting.cpp"],
hdrs = ["AIHeuristicWeighting.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai/mcts/adapters:__pkg__",
"//src/test/cpp/net/eagle0/shardok/ai:__pkg__",
],
deps = [
"//src/main/cpp/net/eagle0/shardok/library:game_state_w",
"//src/main/cpp/net/eagle0/shardok/library/action_point_distances:action_point_distances_cache",
"//src/main/cpp/net/eagle0/shardok/library/map:coords_set",
"//src/main/cpp/net/eagle0/shardok/library/util:hex_map_utils",
"//src/main/protobuf/net/eagle0/shardok/common:command_type_cc_proto",
],
)
cc_library(
name = "ai_command_evaluator",
srcs = ["AICommandEvaluator.cpp"],
hdrs = ["AICommandEvaluator.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai_performance_runner:__pkg__",
"//src/test/cpp/net/eagle0/shardok/ai:__pkg__",
],
deps = [
":ai_command_filter",
":ai_strategy",
":transposition_table",
"//src/main/cpp/net/eagle0/common:sequence_random_generator",
"//src/main/cpp/net/eagle0/shardok/ai/score:ai_score_calculator_interface",
"//src/main/cpp/net/eagle0/shardok/library:engine",
"//src/main/cpp/net/eagle0/shardok/library:game_state_w",
"//src/main/cpp/net/eagle0/shardok/library/map:coords_set",
"//src/main/cpp/net/eagle0/shardok/library/util:hex_cube_utils",
"//src/main/protobuf/net/eagle0/shardok/common:command_type_cc_proto",
"//src/main/protobuf/net/eagle0/shardok/api:command_descriptor_cc_proto",
],
)
@@ -215,17 +155,16 @@ cc_library(
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai/mcts:__pkg__",
"//src/main/cpp/net/eagle0/shardok/ai/mcts/adapters:__pkg__",
"//src/main/cpp/net/eagle0/shardok/ai_performance_runner:__pkg__",
"//src/test/cpp/net/eagle0/shardok/ai:__pkg__",
],
deps = [
":ai_common_types",
"//src/main/cpp/net/eagle0/shardok/library:engine",
"//src/main/cpp/net/eagle0/shardok/library/action_point_distances:action_point_distances_cache",
"//src/main/cpp/net/eagle0/shardok/library/settings:game_settings",
"//src/main/cpp/net/eagle0/shardok/library/util:hex_map_utils",
"//src/main/flatbuffer/net/eagle0/shardok/storage:game_state_cc_fbs",
"//src/main/protobuf/net/eagle0/shardok/api:command_descriptor_cc_proto",
"//src/main/protobuf/net/eagle0/shardok/common:command_type_cc_proto",
],
)
@@ -244,19 +183,40 @@ cc_library(
],
)
cc_library(
name = "ai_score_calculator",
srcs = ["AIScoreCalculator.cpp"],
hdrs = ["AIScoreCalculator.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai/mcts:__pkg__",
"//src/main/cpp/net/eagle0/shardok/ai_performance_runner:__pkg__",
"//src/test/cpp/net/eagle0/shardok/ai:__pkg__",
],
deps = [
":ai_attacker_strategy_selector",
":ai_command_filter",
":ai_unit_score_calculator",
":ai_victory_condition_score_calculator",
":transposition_table",
"//src/main/cpp/net/eagle0/common:sequence_random_generator",
"//src/main/cpp/net/eagle0/shardok/library:engine",
"//src/main/cpp/net/eagle0/shardok/library/view_filters:game_state_guesser",
],
)
cc_library(
name = "ai_strategy",
srcs = ["AIStrategy.cpp"],
hdrs = ["AIStrategy.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai:__subpackages__",
"//src/main/cpp/net/eagle0/shardok/ai/mcts:__pkg__",
"//src/main/cpp/net/eagle0/shardok/ai_performance_runner:__pkg__",
"//src/test/cpp/net/eagle0/shardok/ai:__subpackages__",
],
deps = [
":ai_attack_groups",
"//src/main/cpp/net/eagle0/shardok/library/map:coords_set",
],
)
@@ -266,7 +226,6 @@ cc_library(
hdrs = ["AIUnitScoreCalculator.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai:__subpackages__",
"//src/main/cpp/net/eagle0/shardok/ai_performance_runner:__pkg__",
"//src/test/cpp/net/eagle0/shardok/ai:__subpackages__",
],
@@ -277,6 +236,27 @@ cc_library(
],
)
cc_library(
name = "ai_victory_condition_score_calculator",
srcs = ["AIVictoryConditionScoreCalculator.cpp"],
hdrs = ["AIVictoryConditionScoreCalculator.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai_performance_runner:__pkg__",
"//src/test/cpp/net/eagle0/shardok/ai:__subpackages__",
],
deps = [
":ai_attack_groups",
":ai_attack_locations",
":ai_distance_debuf",
":ai_score_utilities",
"//src/main/cpp/net/eagle0/shardok/library:game_state_w",
"//src/main/cpp/net/eagle0/shardok/library/action_point_distances",
"//src/main/cpp/net/eagle0/shardok/library/action_point_distances:action_point_distances_cache",
"//src/main/flatbuffer/net/eagle0/shardok/storage:game_state_cc_fbs",
],
)
cc_library(
name = "ai_water_crossing_calculator",
srcs = ["AIWaterCrossingCalculator.cpp"],
@@ -284,13 +264,10 @@ cc_library(
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok:__pkg__",
"//src/main/cpp/net/eagle0/shardok/ai:__subpackages__",
"//src/test/cpp/net/eagle0/shardok/ai:__subpackages__",
],
deps = [
":ai_common_types",
":ai_minimum_distance_and_target",
"//src/main/cpp/net/eagle0/shardok/ai/score:ai_score_calculator_interface",
"//src/main/cpp/net/eagle0/shardok/library:game_state_w",
"//src/main/cpp/net/eagle0/shardok/library/action_point_distances",
"//src/main/cpp/net/eagle0/shardok/library/action_point_distances:action_point_distances_cache",
@@ -312,6 +289,7 @@ cc_library(
"//src/main/cpp/net/eagle0/shardok/library:game_state_w",
"//src/main/cpp/net/eagle0/shardok/library/action_point_distances",
"//src/main/cpp/net/eagle0/shardok/library/action_point_distances:action_point_distances_cache",
"//src/main/protobuf/net/eagle0/shardok/api:command_descriptor_cc_proto",
],
)
@@ -346,14 +324,51 @@ cc_library(
],
deps = [
":ai_attacker_strategy_selector",
":ai_command_evaluator",
":ai_defender_strategy_selector",
":ai_score_calculator",
":ai_time_budget",
":ai_water_crossing_command_chooser",
"//src/main/cpp/net/eagle0/common:time_utils",
"//src/main/cpp/net/eagle0/shardok/ai/score:ai_score_calculator_interface",
"//src/main/cpp/net/eagle0/shardok/library:engine",
"//src/main/cpp/net/eagle0/shardok/library/util:hex_map_utils",
"//src/main/protobuf/net/eagle0/shardok/api:command_descriptor_cc_proto",
],
)
cc_library(
name = "ai_config",
hdrs = ["AIConfig.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai_performance_runner:__pkg__",
"//src/test/cpp/net/eagle0/shardok/ai:__pkg__",
],
)
cc_library(
name = "ai_mcts_clean",
srcs = ["MCTSCleanAI.cpp"],
hdrs = ["MCTSCleanAI.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai_performance_runner:__pkg__",
"//src/test/cpp/net/eagle0/shardok/ai:__pkg__",
],
deps = [
":ai_attack_locations",
":ai_config",
":ai_iterative_deepening", # For SearchResult compatibility
":ai_score_calculator",
":ai_strategy",
":ai_time_budget",
"//src/main/cpp/net/eagle0/common:random_generator",
"//src/main/cpp/net/eagle0/common:sequence_random_generator",
"//src/main/cpp/net/eagle0/shardok/library:engine",
"//src/main/cpp/net/eagle0/shardok/library:game_state_w",
"//src/main/cpp/net/eagle0/shardok/library/action_point_distances:action_point_distances_cache",
"//src/main/cpp/net/eagle0/shardok/library/settings:game_settings",
"//src/main/protobuf/net/eagle0/shardok/api:command_descriptor_cc_proto",
"//src/main/protobuf/net/eagle0/shardok/common:command_type_cc_proto",
],
)
@@ -379,17 +394,13 @@ cc_library(
":ai_defender_strategy_selector",
":ai_flee_decision_calculator",
":ai_iterative_deepening", # Direct dependency for runtime selection
":ai_score_calculator",
":ai_time_budget",
":ai_water_crossing_command_chooser",
"//src/main/cpp/net/eagle0/common:time_utils",
"//src/main/cpp/net/eagle0/shardok/ai/mcts:shardok_mcts_ai", # MCTS with abstraction layer
"//src/main/cpp/net/eagle0/shardok/ai/score:ai_score_calculator_interface",
"//src/main/cpp/net/eagle0/shardok/ai/score:mcts_optimized_ai_score_calculator", # Bounded linear scorer for MCTS
"//src/main/cpp/net/eagle0/shardok/ai/score:normalized_ai_score_calculator", # Normalized [0,1] scorer for ML training
"//src/main/cpp/net/eagle0/shardok/ai/score:standard_ai_score_calculator", # Standard unbounded scorer (default)
"//src/main/cpp/net/eagle0/shardok/ai/mcts:mcts_ai", # Direct dependency for runtime selection
"//src/main/cpp/net/eagle0/shardok/library/action_point_distances",
"//src/main/cpp/net/eagle0/shardok/library/action_point_distances:action_point_distances_cache",
"//src/main/cpp/net/eagle0/shardok/library/util:game_state_dumper",
"@com_google_protobuf//:protobuf",
],
)
@@ -10,9 +10,8 @@
#include <utility>
#include "AIAttackerStrategySelector.hpp"
#include "AICommandEvaluator.hpp"
#include "AIScoreCalculator.hpp"
#include "TranspositionTable.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/score/AIScoreCalculator.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokEngine.hpp"
namespace shardok {
@@ -24,21 +23,19 @@ IterativeDeepeningAI::IterativeDeepeningAI(
const bool isDefender,
AIStrategy strategy,
const CoordsSet& castleCoords,
const AIScoreCalculator& scorer,
const APDCache& apdCache,
BattalionTypeGetter battalionTypeGetter)
const ALCache& alCache)
: playerId(playerId),
isDefender(isDefender),
strategy(std::move(strategy)),
castleCoords(castleCoords),
scorer(scorer),
apdCache(apdCache),
battalionTypeGetter(std::move(battalionTypeGetter)) {} // Move the function object
alCache(alCache) {}
auto IterativeDeepeningAI::IterativeSearch(
const GameSettingsSPtr& settings,
const GameStateW& state,
const CommandListSPtr& commands,
const std::vector<CommandProto>& commands,
const AITimeBudget& initialBudget) const -> SearchResult {
// Make a mutable copy of the time budget to track remaining time
AITimeBudget timeBudget = initialBudget;
@@ -51,7 +48,7 @@ auto IterativeDeepeningAI::IterativeSearch(
// DEBUG: Clear TT to see if that's causing the suspicious depth reaching
// g_transpositionTable.clear(); // Uncomment to test without cross-search caching
if (commands->empty()) {
if (commands.empty()) {
#if DEBUG_ITERATIVE_DEEPENING_TIMINGS
printf("ID AI: Commands are empty, returning early\n");
#endif
@@ -70,14 +67,20 @@ auto IterativeDeepeningAI::IterativeSearch(
const auto& settingsGetter = settings->GetGetter();
const auto guessedEngine = ShardokEngine(settings, state);
const auto maxRepeatCount = settingsGetter.Backing().ai_utility_repeat_count();
const ScoreValue currentUtility =
scorer.GuessedStateScore(isDefender, state, strategy, castleCoords);
const ScoreValue currentUtility = AIScoreCalculator::GuessedStateScore(
isDefender,
state,
strategy,
castleCoords,
settingsGetter,
apdCache,
alCache);
// Initialize data structures for tracking scores at each depth
scoresByDepth.clear();
scoresByDepth.resize(commands->size());
scoresByDepth.resize(commands.size());
highestDepthCompleted.clear();
highestDepthCompleted.resize(commands->size(), 0);
highestDepthCompleted.resize(commands.size(), 0);
size_t currentDepth = 1;
size_t previousBestCommand = 0; // Track best command from previous depth
@@ -108,7 +111,7 @@ auto IterativeDeepeningAI::IterativeSearch(
auto future = SearchCommandAtDepthWithEngine(
guessedEngine,
scorer,
settingsGetter,
maxRepeatCount,
commands,
cmdIndex,
@@ -132,8 +135,7 @@ auto IterativeDeepeningAI::IterativeSearch(
evaluatedCount++;
// Check if this command is not END_TURN_COMMAND
if ((*commands)[cmdIndex]->GetCommandType() !=
net::eagle0::shardok::common::END_TURN_COMMAND) {
if (commands[cmdIndex].type() != net::eagle0::shardok::common::END_TURN_COMMAND) {
allEndTurnCommands = false;
}
}
@@ -144,7 +146,7 @@ auto IterativeDeepeningAI::IterativeSearch(
size_t currentBestCommand = 0;
ScoreValue currentBestScore = -std::numeric_limits<ScoreValue>::infinity();
for (size_t i = 0; i < commands->size(); ++i) {
for (size_t i = 0; i < commands.size(); ++i) {
if (highestDepthCompleted[i] >= currentDepth) {
if (scoresByDepth[i][currentDepth] > currentBestScore) {
currentBestScore = scoresByDepth[i][currentDepth];
@@ -157,20 +159,16 @@ auto IterativeDeepeningAI::IterativeSearch(
if (currentDepth > 1 && currentBestCommand != previousBestCommand) {
#if DEBUG_ITERATIVE_DEEPENING_TIMINGS
printf("ID AI: Best command changed at depth %lu:\n", currentDepth);
printf(" Depth %lu best: command %zu (score %.2f) - type: %s\n",
printf(" Depth %lu best: command %zu (score %.2f) - %s\n",
currentDepth - 1,
previousBestCommand,
scoresByDepth[previousBestCommand][currentDepth - 1],
net::eagle0::shardok::common::CommandType_Name(
(*commands)[previousBestCommand]->GetCommandType())
.c_str());
printf(" Depth %lu best: command %zu (score %.2f) - type: %s\n",
commands[previousBestCommand].DebugString().c_str());
printf(" Depth %lu best: command %zu (score %.2f) - %s\n",
currentDepth,
currentBestCommand,
currentBestScore,
net::eagle0::shardok::common::CommandType_Name(
(*commands)[currentBestCommand]->GetCommandType())
.c_str());
commands[currentBestCommand].DebugString().c_str());
#endif
}
@@ -249,7 +247,7 @@ auto IterativeDeepeningAI::IterativeSearch(
result.searchCompleted = result.minimumDepthCompleted;
result.timeUsed = std::chrono::duration_cast<std::chrono::milliseconds>(
std::chrono::steady_clock::now() - startTime);
result.availableCommandCount = commands->size();
result.availableCommandCount = commands.size();
result.commandCountEvaluated = evaluatedCountAtHighestDepth;
result.completionReason = completionReason;
@@ -272,9 +270,9 @@ bool IterativeDeepeningAI::IsTimeExpired(const AITimeBudget& budget) {
auto IterativeDeepeningAI::SearchCommandAtDepthWithEngine(
const ShardokEngine& guessedEngine,
const AIScoreCalculator& scorer,
const GameSettings::Getter& settingsGetter,
const int maxRepeatCount,
const CommandListSPtr& commands,
const std::vector<CommandProto>& commands,
const size_t commandIndex,
const int desiredDepth,
const ScoreValue currentUtility,
@@ -284,57 +282,65 @@ auto IterativeDeepeningAI::SearchCommandAtDepthWithEngine(
result.depthAchieved = desiredDepth;
result.searchCompleted = true;
result.minimumDepthCompleted = true;
result.availableCommandCount = commands->size();
result.availableCommandCount = commands.size();
result.commandCountEvaluated = 1; // We're evaluating just this command
if (commandIndex >= commands->size()) {
if (commandIndex >= commands.size()) {
result.bestScore = 0.0;
std::promise<SearchResult> p;
p.set_value(result);
return p.get_future();
}
// Track concurrent evaluations and adjust time accounting
AIEvaluationCounter counter;
const auto startTime = std::chrono::steady_clock::now();
try {
// Track concurrent evaluations and adjust time accounting
AIEvaluationCounter counter;
const auto startTime = std::chrono::steady_clock::now();
// Calculate deadline from remaining time budget
const auto deadline = startTime + timeBudget.remainingBudget;
// Calculate deadline from remaining time budget
const auto deadline = startTime + timeBudget.remainingBudget;
// Create command evaluator for lookahead search
AICommandEvaluator evaluator(scorer, apdCache, battalionTypeGetter);
// Get the future from CommandScore - don't wait yet
// Note: CommandScore expects remainingLookahead, not desiredDepth
// desiredDepth 1 = evaluate immediate (remainingLookahead 0)
// desiredDepth 2 = look 1 move ahead (remainingLookahead 1)
// desiredDepth N = look N-1 moves ahead (remainingLookahead N-1)
auto commandScoreFuture = AIScoreCalculator::CommandScore(
playerId,
isDefender,
desiredDepth - 1, // Convert desiredDepth to remainingLookahead
maxRepeatCount,
guessedEngine,
strategy,
currentUtility,
settingsGetter,
castleCoords,
apdCache,
alCache,
commandIndex,
deadline);
// Get the future from EvaluateCommand - don't wait yet
// Note: EvaluateCommand expects remainingLookahead, not desiredDepth
// desiredDepth 1 = evaluate immediate (remainingLookahead 0)
// desiredDepth 2 = look 1 move ahead (remainingLookahead 1)
// desiredDepth N = look N-1 moves ahead (remainingLookahead N-1)
auto commandScoreFuture = evaluator.EvaluateCommand(
playerId,
isDefender,
desiredDepth - 1, // Convert desiredDepth to remainingLookahead
maxRepeatCount,
guessedEngine,
strategy,
currentUtility,
castleCoords,
commandIndex,
deadline);
// Calculate time and adjust budget before waiting
// This is needed because we need to update timeBudget synchronously
const auto commandScore = commandScoreFuture.get();
// Calculate time and adjust budget before waiting
// This is needed because we need to update timeBudget synchronously
const auto commandScore = commandScoreFuture.get();
const auto elapsed = std::chrono::steady_clock::now() - startTime;
const int concurrentCount = AIEvaluationCounter::GetCurrentCount();
const auto adjustedElapsed = elapsed / std::max(1, concurrentCount);
const auto adjustedElapsedMs =
std::chrono::duration_cast<std::chrono::milliseconds>(adjustedElapsed);
const auto elapsed = std::chrono::steady_clock::now() - startTime;
const int concurrentCount = AIEvaluationCounter::GetCurrentCount();
const auto adjustedElapsed = elapsed / std::max(1, concurrentCount);
const auto adjustedElapsedMs =
std::chrono::duration_cast<std::chrono::milliseconds>(adjustedElapsed);
// Deduct adjusted time from remaining budget
timeBudget.remainingBudget -= adjustedElapsedMs;
// Deduct adjusted time from remaining budget
timeBudget.remainingBudget -= adjustedElapsedMs;
result.bestScore = commandScore;
result.bestScore = commandScore;
} catch (const std::exception& e) {
// If evaluation fails, return a neutral score rather than crashing
#if DEBUG_ITERATIVE_DEEPENING_TIMINGS
printf("SearchCommandAtDepthWithEngine: evaluation failed with exception: %s\n", e.what());
#endif
result.bestScore = 0.0;
}
std::promise<SearchResult> p;
p.set_value(result);
@@ -12,18 +12,17 @@
#include "AIStrategy.hpp"
#include "AITimeBudget.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIAttackLocations.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/score/AIScoreCalculator.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokCTypes.h"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokCommand.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistancesCache.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/util/HexMapUtils.hpp"
#include "src/main/protobuf/net/eagle0/shardok/api/command_descriptor.pb.h"
namespace shardok {
// Forward declarations
class ShardokEngine;
using ScoreValue = double;
using BattalionTypeGetter = std::function<BattalionTypeSPtr(BattalionTypeId)>;
using CommandProto = net::eagle0::shardok::api::CommandDescriptor;
/// Reason why AI evaluation completed at the achieved depth.
enum class EvaluationCompletionReason {
@@ -62,14 +61,13 @@ public:
bool isDefender,
AIStrategy strategy,
const CoordsSet& castleCoords,
const AIScoreCalculator& scorer,
const APDCache& apdCache,
BattalionTypeGetter battalionTypeGetter); // Pass by value
const ALCache& alCache);
[[nodiscard]] SearchResult IterativeSearch(
const GameSettingsSPtr& settings,
const GameStateW& state,
const CommandListSPtr& commands,
const std::vector<CommandProto>& commands,
const AITimeBudget& initialBudget) const;
private:
@@ -77,9 +75,8 @@ private:
bool isDefender;
AIStrategy strategy;
CoordsSet castleCoords;
const AIScoreCalculator& scorer;
const APDCache& apdCache;
BattalionTypeGetter battalionTypeGetter; // Store by value, not reference!
const ALCache& alCache;
// Reusable vectors to reduce memory allocations
mutable std::vector<std::vector<ScoreValue>> scoresByDepth;
@@ -90,9 +87,9 @@ private:
[[nodiscard]] std::future<SearchResult> SearchCommandAtDepthWithEngine(
const ShardokEngine& guessedEngine,
const AIScoreCalculator& scorer,
const GameSettings::Getter& settingsGetter,
int maxRepeatCount,
const CommandListSPtr& commands,
const std::vector<CommandProto>& commands,
size_t commandIndex,
int desiredDepth,
ScoreValue currentUtility,
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,105 @@
//
// MCTS-based AI system for Shardok
// Alternative to IterativeDeepeningAI using Monte Carlo Tree Search
//
#ifndef EAGLE0_MCTSAI_HPP
#define EAGLE0_MCTSAI_HPP
#include <chrono>
#include <future>
#include <memory>
#include <vector>
#include "AIStrategy.hpp"
#include "AITimeBudget.hpp"
#include "IterativeDeepeningAI.hpp" // For SearchResult compatibility
#include "src/main/cpp/net/eagle0/shardok/library/GameStateW.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokCTypes.h"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistancesCache.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
#include "src/main/protobuf/net/eagle0/shardok/api/command_descriptor.pb.h"
#include "src/main/protobuf/net/eagle0/shardok/common/command_type.pb.h"
namespace shardok {
// Forward declarations
class ShardokEngine;
struct MCTSNode;
// Configuration for MCTS algorithm
struct MCTSConfig {
double explorationConstant = 1.414; // UCB1 constant (sqrt(2) by default)
int maxPlayerFlips = 1; // Number of player turn changes to evaluate
bool useMultithreading = true; // Enable parallel MCTS
int numThreads = 16; // Number of threads for parallel MCTS (when enabled)
};
class MCTSAI {
public:
using CommandProto = net::eagle0::shardok::api::CommandDescriptor;
using SearchResult = IterativeDeepeningAI::SearchResult;
MCTSAI(PlayerId playerId,
bool isDefender,
AIStrategy strategy,
const CoordsSet& castleCoords,
const APDCache& apdCache,
const ALCache& alCache,
MCTSConfig config = MCTSConfig{});
// Main search interface - compatible with IterativeDeepeningAI
[[nodiscard]] auto Search(
const GameSettingsSPtr& settings,
const GameStateW& state,
const std::vector<CommandProto>& commands,
const AITimeBudget& budget) const -> SearchResult;
// Get/set configuration
[[nodiscard]] auto GetConfig() const -> const MCTSConfig& { return config; }
void SetConfig(const MCTSConfig& newConfig) { config = newConfig; }
private:
PlayerId playerId;
bool isDefender;
AIStrategy strategy;
const CoordsSet& castleCoords;
const APDCache& apdCache;
const ALCache& alCache;
MCTSConfig config;
// Internal MCTS tree building
[[nodiscard]] auto BuildMCTSTree(
const ShardokEngine& engine,
const SettingsGetter& settingsGetter,
std::chrono::steady_clock::time_point deadline) const -> std::unique_ptr<MCTSNode>;
// MCTS algorithm phases
auto MCTSSelection(MCTSNode* root) const -> MCTSNode*;
auto MCTSExpansion(
MCTSNode* node,
const ShardokEngine& engine,
const SettingsGetter& settingsGetter) const -> MCTSNode*;
auto MCTSSimulation(
const ShardokEngine& engineState,
PlayerId currentPlayer,
const SettingsGetter& settingsGetter) const -> double;
auto MCTSBackpropagation(MCTSNode* node, double reward) const -> void;
// Helper functions
[[nodiscard]] auto IsTerminalForPlayer(
const GameStateW& gameState,
PlayerId currentPlayer,
const SettingsGetter& settingsGetter) const -> bool;
// Get coordinate information for logging
[[nodiscard]] std::string GetCommandCoordinateInfo(
net::eagle0::shardok::common::CommandType commandType,
size_t commandIndex,
const GameStateW& gameState,
const GameSettingsSPtr& settings) const;
};
} // namespace shardok
#endif // EAGLE0_MCTSAI_HPP
@@ -0,0 +1,292 @@
#include "MCTSCleanAI.hpp"
#include <algorithm>
#include <chrono>
#include <cmath>
#include <limits>
#include "AIScoreCalculator.hpp"
#include "src/main/cpp/net/eagle0/common/SequenceRandomGenerator.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokEngine.hpp"
#include "src/main/protobuf/net/eagle0/shardok/common/command_type.pb.h"
namespace shardok {
MCTSCleanAI::MCTSCleanAI(
PlayerId playerId,
bool isDefender,
const AIStrategy& strategy,
const CoordsSet& castleCoords,
const APDCache& apdCache,
const ALCache& alCache)
: ourPlayerId(playerId),
isDefender(isDefender),
strategy(strategy),
castleCoords(castleCoords),
apdCache(apdCache),
alCache(alCache),
rng(std::chrono::steady_clock::now().time_since_epoch().count()) {}
IterativeDeepeningAI::SearchResult MCTSCleanAI::Search(
const GameSettingsSPtr& settings,
const GameStateW& gameState,
const std::vector<CommandProto>& availableCommands,
const AITimeBudget& timeBudget) {
const auto startTime = std::chrono::steady_clock::now();
const auto maxTime = std::chrono::duration_cast<std::chrono::milliseconds>(
timeBudget.remainingBudget * 0.9); // 90% of budget
// Create root node
auto root = std::make_unique<MCTSNode>();
root->resultingGameState = gameState;
root->currentPlayer = ourPlayerId;
root->isOurTurn = true;
root->playerFlipsFromRoot = 0;
// Initialize untried commands for root
for (size_t i = 0; i < availableCommands.size(); ++i) { root->untriedCommands.push_back(i); }
int iterationCount = 0;
// Main MCTS loop
while (std::chrono::steady_clock::now() - startTime < maxTime) {
// 1. Selection - find leaf node to expand
MCTSNode* selected = Selection(root.get());
// 2. Expansion - add new child if possible
MCTSNode* expanded = Expansion(selected, settings, availableCommands);
// 3. Simulation - run random playout
double reward = Simulation(
expanded->resultingGameState,
expanded->currentPlayer,
expanded->playerFlipsFromRoot,
settings);
// 4. Backpropagation - update statistics
Backpropagation(expanded, reward);
iterationCount++;
// Early exit if all commands tried
if (root->untriedCommands.empty() && root->children.size() == availableCommands.size()) {
bool allChildrenFullyExplored = true;
for (const auto& child : root->children) {
if (child->visitCount < 10) { // Minimum visits per child
allChildrenFullyExplored = false;
break;
}
}
if (allChildrenFullyExplored) { break; }
}
}
// Select best command based on visit counts
size_t bestCommandIndex = 0;
int maxVisits = 0;
for (size_t i = 0; i < root->children.size(); ++i) {
if (root->children[i]->visitCount > maxVisits) {
maxVisits = root->children[i]->visitCount;
bestCommandIndex = root->children[i]->commandIndex;
}
}
// Create search result
IterativeDeepeningAI::SearchResult result;
result.bestCommandIndex = bestCommandIndex;
result.availableCommandCount = availableCommands.size();
result.depthAchieved = maxPlayerFlips; // Our max search depth
result.commandCountEvaluated = iterationCount;
result.completionReason = EvaluationCompletionReason::RAN_OUT_OF_TIME;
printf("MCTS Clean: %d iterations, best command %zu with %d visits\n",
iterationCount,
bestCommandIndex,
maxVisits);
return result;
}
// Selection - navigate to leaf using UCB1
MCTSCleanAI::MCTSNode* MCTSCleanAI::Selection(MCTSNode* root) {
MCTSNode* current = root;
while (!current->children.empty() && current->untriedCommands.empty()) {
current = SelectChild(current);
}
return current;
}
// Expansion - add new child node
MCTSCleanAI::MCTSNode* MCTSCleanAI::Expansion(
MCTSNode* node,
const GameSettingsSPtr& settings,
const std::vector<CommandProto>& availableCommands) {
// If we've reached max player flips, don't expand
if (node->playerFlipsFromRoot >= maxPlayerFlips) { return node; }
// If no untried commands, return current node
if (node->untriedCommands.empty()) { return node; }
// Pick an untried command
size_t cmdIndex = PickUntriedCommand(node);
// Create child engine and execute command
auto childEngine = CreateEngine(node->resultingGameState, settings);
PlayerId playerBefore = childEngine->GetCurrentPlayerId();
// Use sequence random generator for deterministic results
auto averageGenerator = std::make_shared<SequenceRandomGenerator>(std::vector<double>{0.5});
childEngine->PostCommand(playerBefore, cmdIndex, averageGenerator);
PlayerId playerAfter = childEngine->GetCurrentPlayerId();
// Create child node
bool isOurTurnAfter = (playerAfter == ourPlayerId);
auto child = std::make_unique<MCTSNode>(
cmdIndex,
availableCommands[cmdIndex].type(),
playerAfter,
isOurTurnAfter);
child->resultingGameState = childEngine->GetCurrentGameState();
child->parent = node;
// Track player flips
child->playerFlipsFromRoot = node->playerFlipsFromRoot;
if (playerBefore != playerAfter) { child->playerFlipsFromRoot++; }
// Initialize child's untried commands if we haven't hit max flips
if (child->playerFlipsFromRoot < maxPlayerFlips) {
auto childCommands = childEngine->GetAvailableCommandsForAIPlayer(playerAfter);
for (size_t i = 0; i < childCommands->size(); ++i) { child->untriedCommands.push_back(i); }
}
MCTSNode* childPtr = child.get();
node->children.push_back(std::move(child));
return childPtr;
}
// Simulation - run random playout from current state
double MCTSCleanAI::Simulation(
const GameStateW& startState,
PlayerId /* startPlayer */,
int startFlips,
const GameSettingsSPtr& settings) {
auto simEngine = CreateEngine(startState, settings);
int currentFlips = startFlips;
auto averageGenerator = std::make_shared<SequenceRandomGenerator>(std::vector<double>{0.5});
while (currentFlips < maxPlayerFlips) {
PlayerId currentPlayer = simEngine->GetCurrentPlayerId();
bool isOurTurn = (currentPlayer == ourPlayerId);
// Get available commands
auto commands = simEngine->GetAvailableCommandsForAIPlayer(currentPlayer);
if (commands->empty()) break;
// Pick best command based on whose turn it is
size_t bestCmd = 0;
double bestScore = isOurTurn ? -std::numeric_limits<double>::infinity()
: std::numeric_limits<double>::infinity();
for (size_t i = 0; i < commands->size(); ++i) {
auto testEngine = CreateEngine(simEngine->GetCurrentGameState(), settings);
testEngine->PostCommand(currentPlayer, i, averageGenerator);
// Always score from our perspective
double score = ScoreFromOurPerspective(
testEngine->GetCurrentGameState(),
settings->GetGetter());
// Our turn: maximize our score, Opponent turn: minimize our score
bool shouldSelect = isOurTurn ? (score > bestScore) : (score < bestScore);
if (shouldSelect) {
bestScore = score;
bestCmd = i;
}
}
// Execute chosen command
PlayerId playerBefore = simEngine->GetCurrentPlayerId();
simEngine->PostCommand(currentPlayer, bestCmd, averageGenerator);
PlayerId playerAfter = simEngine->GetCurrentPlayerId();
// Track player flips
if (playerBefore != playerAfter) { currentFlips++; }
}
return ScoreFromOurPerspective(simEngine->GetCurrentGameState(), settings->GetGetter());
}
// Backpropagation - update node statistics
void MCTSCleanAI::Backpropagation(MCTSNode* node, double score) {
while (node != nullptr) {
node->visitCount++;
node->totalScore += score; // Always from our perspective
node->averageScore = node->totalScore / node->visitCount;
node = node->parent;
}
}
// Helper functions
double MCTSCleanAI::ScoreFromOurPerspective(
const GameStateW& gameState,
const SettingsGetter& settingsGetter) {
return AIScoreCalculator::GuessedStateScore(
isDefender,
gameState,
strategy,
castleCoords,
settingsGetter,
apdCache,
alCache);
}
MCTSCleanAI::MCTSNode* MCTSCleanAI::SelectChild(MCTSNode* node) {
MCTSNode* bestChild = nullptr;
double bestUCB1 = node->isOurTurn ? -std::numeric_limits<double>::infinity()
: std::numeric_limits<double>::infinity();
for (const auto& child : node->children) {
double exploitation = child->averageScore;
double exploration =
explorationConstant * std::sqrt(std::log(node->visitCount) / child->visitCount);
double ucb1 = exploitation + exploration;
// Our turn: pick highest UCB1, Opponent turn: pick lowest UCB1
bool shouldSelect = node->isOurTurn ? (ucb1 > bestUCB1) : (ucb1 < bestUCB1);
if (shouldSelect) {
bestUCB1 = ucb1;
bestChild = child.get();
}
}
return bestChild;
}
size_t MCTSCleanAI::PickUntriedCommand(MCTSNode* node) {
if (node->untriedCommands.empty()) {
return 0; // Should not happen
}
// Pick random untried command
std::uniform_int_distribution<size_t> indexDist(0, node->untriedCommands.size() - 1);
size_t randomIndex = indexDist(rng);
size_t cmdIndex = node->untriedCommands[randomIndex];
node->untriedCommands.erase(node->untriedCommands.begin() + randomIndex);
return cmdIndex;
}
std::shared_ptr<ShardokEngine> MCTSCleanAI::CreateEngine(
const GameStateW& gameState,
const GameSettingsSPtr& settings) {
return std::make_shared<ShardokEngine>(settings, gameState);
}
} // namespace shardok
@@ -0,0 +1,133 @@
#pragma once
#include <memory>
#include <random>
#include <vector>
#include "AIAttackLocations.hpp"
#include "AIConfig.hpp"
#include "AIStrategy.hpp"
#include "AITimeBudget.hpp"
#include "IterativeDeepeningAI.hpp" // For SearchResult
#include "src/main/cpp/net/eagle0/common/RandomGenerator.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/GameStateW.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistancesCache.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/map/CoordsSet.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
#include "src/main/protobuf/net/eagle0/shardok/api/command_descriptor.pb.h"
#include "src/main/protobuf/net/eagle0/shardok/common/command_type.pb.h"
namespace shardok {
// Forward declare cache type aliases
using APDCache = std::shared_ptr<ActionPointDistancesCache>;
using ALCache = std::unique_ptr<AttackLocationsCache>;
} // namespace shardok
namespace shardok {
// Clean MCTS implementation following the design document
class MCTSCleanAI {
public:
// Clean node structure with single perspective scoring
struct MCTSNode {
// Command that led to this state
size_t commandIndex;
net::eagle0::shardok::common::CommandType commandType;
// Game state after executing the command
GameStateW resultingGameState;
PlayerId currentPlayer; // Whose turn it is in this state
// Tree position
int playerFlipsFromRoot; // Number of player changes from root
bool isOurTurn; // currentPlayer == our AI's playerId
// MCTS statistics (always from our perspective)
int visitCount = 0;
double totalScore = 0.0;
double averageScore = 0.0;
// Tree structure
std::vector<std::unique_ptr<MCTSNode>> children;
std::vector<size_t> untriedCommands;
MCTSNode* parent = nullptr;
// Constructor
MCTSNode(
size_t cmdIndex,
net::eagle0::shardok::common::CommandType cmdType,
PlayerId currentPlayerId,
bool ourTurn)
: commandIndex(cmdIndex),
commandType(cmdType),
currentPlayer(currentPlayerId),
playerFlipsFromRoot(0),
isOurTurn(ourTurn) {}
// Root constructor
MCTSNode()
: commandIndex(0),
commandType(net::eagle0::shardok::common::UNKNOWN_COMMAND),
currentPlayer(0),
playerFlipsFromRoot(0),
isOurTurn(true) {}
};
MCTSCleanAI(
PlayerId playerId,
bool isDefender,
const AIStrategy& strategy,
const CoordsSet& castleCoords,
const APDCache& apdCache,
const ALCache& alCache);
// Main search function compatible with existing interface
IterativeDeepeningAI::SearchResult Search(
const GameSettingsSPtr& settings,
const GameStateW& gameState,
const std::vector<CommandProto>& availableCommands,
const AITimeBudget& timeBudget);
private:
PlayerId ourPlayerId;
bool isDefender;
AIStrategy strategy;
CoordsSet castleCoords;
const APDCache& apdCache;
const ALCache& alCache;
// MCTS parameters
static constexpr int maxPlayerFlips = 5; // Stop after N player changes
static constexpr double explorationConstant = 1.414; // sqrt(2)
// Random number generation
std::mt19937 rng;
std::uniform_real_distribution<double> uniformDist{0.0, 1.0};
// Core MCTS algorithms
MCTSNode* Selection(MCTSNode* root);
MCTSNode* Expansion(
MCTSNode* node,
const GameSettingsSPtr& settings,
const std::vector<CommandProto>& availableCommands);
double Simulation(
const GameStateW& startState,
PlayerId startPlayer,
int startFlips,
const GameSettingsSPtr& settings);
void Backpropagation(MCTSNode* node, double score);
// Helper functions
double ScoreFromOurPerspective(
const GameStateW& gameState,
const SettingsGetter& settingsGetter);
MCTSNode* SelectChild(MCTSNode* node);
size_t PickUntriedCommand(MCTSNode* node);
std::shared_ptr<ShardokEngine> CreateEngine(
const GameStateW& gameState,
const GameSettingsSPtr& settings);
};
} // namespace shardok
@@ -0,0 +1,147 @@
# MCTS Player-Flip Depth System Design
## Overview
Replace command-count-based depth limits with player-turn-aware depth that naturally aligns with game structure. All nodes evaluate to the same player-flip depth to ensure comparable scores.
## Core Concepts
### Player-Flip Depth
- **Depth 1**: Evaluate until first player flip (complete our turn)
- **Depth 2**: Continue through opponent's full turn
- **Depth 3**: Continue through our next full turn
- **Depth N**: N complete player turn changes
### Minimax Selection
- **Always evaluate from original AI's perspective**
- **Our turn**: Select moves that maximize our score
- **Opponent's turn**: Select moves that minimize our score
- Same backpropagation value regardless of whose turn
## Implementation Plan
### 1. Configuration Changes
```cpp
struct MCTSConfig {
double explorationConstant = 1.414;
int maxPlayerFlips = 2; // How many player changes to evaluate
bool useMultithreading = true;
int numThreads = 16;
// REMOVED: maxTreeDepth - all nodes go to same player-flip depth
// REMOVED: maxSimulationDepth - replaced by maxPlayerFlips
};
```
### 2. Node Structure Updates
```cpp
struct MCTSNode {
// Existing fields...
// New fields for player-flip tracking
int playerFlipsFromRoot = 0;
bool isMaximizingPlayer = true; // true = our turn, false = opponent's
// Modified selection for minimax
MCTSNode* GetBestChild(double explorationConstant) {
if (isMaximizingPlayer) {
return GetChildWithHighestUCB1(explorationConstant);
} else {
return GetChildWithLowestUCB1(explorationConstant);
}
}
};
```
### 3. Expansion Rules
- **Continue expanding** until reaching `maxPlayerFlips` player changes
- **No arbitrary depth limit** - accept theoretical stack overflow risk
- **Mark as terminal** only when:
- Game is over
- Reached `maxPlayerFlips` player changes
- No available commands
### 4. Key Implementation Details
#### Terminal Detection
```cpp
bool IsTerminalForExpansion() {
return playerFlipsFromRoot >= maxPlayerFlips ||
gameIsOver ||
noCommandsAvailable;
}
```
#### Player Tracking
```cpp
// When expanding END_TURN or END_PLAYER_SETUP
child->playerFlipsFromRoot = parent->playerFlipsFromRoot + 1;
child->isMaximizingPlayer = !parent->isMaximizingPlayer;
```
#### UCB1 for Minimax
- Maximizing player: Choose highest UCB1
- Minimizing player: Choose lowest UCB1
- Unvisited nodes: Extreme values to force exploration
### 5. Simulation Strategy
Simulations run until `maxPlayerFlips` is reached:
- Random moves for both players
- Stop at player-flip boundaries
- Always evaluate from original AI perspective
## Benefits
### Consistent Evaluation
- All leaf nodes at same player-flip depth
- Scores are directly comparable
- No apples-to-oranges comparison issues
### Natural Game Structure
- Respects turn boundaries
- Complete tactical sequences evaluated
- Opponent responses properly modeled
### Strategic Depth Control
- **Setup/Early**: `maxPlayerFlips = 1` (fast, local tactics)
- **Mid-game**: `maxPlayerFlips = 2` (balanced)
- **Critical**: `maxPlayerFlips = 3+` (deep strategy)
## Implementation Order
1. **Phase 1**: Update MCTSConfig
- Remove `maxTreeDepth` and `maxSimulationDepth`
- Add `maxPlayerFlips`
2. **Phase 2**: Modify MCTSNode
- Add `playerFlipsFromRoot` and `isMaximizingPlayer`
- Update child selection for minimax
3. **Phase 3**: Update Expansion
- Track player flips
- Remove depth-based termination
- Only terminate at player-flip boundaries
4. **Phase 4**: Fix Selection
- Implement minimax selection based on `isMaximizingPlayer`
- Modify UCB1 interpretation
5. **Phase 5**: Update Simulation
- Run until `maxPlayerFlips` reached
- Handle both player perspectives
## Testing Strategy
1. Verify all evaluations reach same player-flip depth
2. Confirm opponent chooses minimizing moves
3. Test with different `maxPlayerFlips` settings
4. Validate score consistency across tree
## Notes
- Stack overflow risk accepted for evaluation consistency
- Each unit can make multiple moves per turn (move + scout + attack)
- Maximum ~10 units per player limits practical depth
- END_TURN and END_PLAYER_SETUP both count as player flips
@@ -10,27 +10,17 @@
#define DEBUG_FLEE_DECISIONS
// Enable to dump game state and debug tree to /tmp for debugging
// #define ENABLE_MCTS_DEBUG_DUMP
#ifdef ENABLE_MCTS_DEBUG_DUMP
#include <chrono>
#include <fstream>
#include <iomanip>
#include <sstream>
#endif
#include <google/protobuf/util/message_differencer.h>
#include "AIAttackerStrategySelector.hpp"
#include "AIConfig.hpp"
#include "AIConfig.hpp" // Must come before other AI includes
#include "AIDefenderStrategySelector.hpp"
#include "AIFleeDecisionCalculator.hpp"
#include "AIScoreUtilities.hpp"
#include "AITimeBudget.hpp"
#include "IterativeDeepeningAI.hpp"
#include "mcts/ShardokMCTSAI.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/score/MCTSOptimizedAIScoreCalculator.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/score/NormalizedAIScoreCalculator.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/score/StandardAIScoreCalculator.hpp"
#include "mcts/MCTSAI.hpp"
#include "src/main/cpp/net/eagle0/common/TimeUtils.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/util/HexMapUtils.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/view_filters/GameStateGuesser.hpp"
#include "src/main/protobuf/net/eagle0/shardok/api/action_result_view.pb.h"
@@ -55,16 +45,12 @@ ShardokAIClient::ShardokAIClient(
const bool isDefender,
const HexMap *hexMap,
const SettingsGetter &settings,
const AIAlgorithmType aiAlgorithmType,
const ScoringCalculatorType scoringCalculatorType,
const mcts::MCTSConfig &mctsConfig)
const AIAlgorithmType aiAlgorithmType)
: playerId(playerId),
isDefender(isDefender),
aiAlgorithmType(aiAlgorithmType),
scoringCalculatorType(scoringCalculatorType),
alCache(std::make_unique<AttackLocationsCache>(hexMap, settings)),
waterCrossingCommandChooser(playerId, apdCache),
mctsConfig(mctsConfig) {
waterCrossingCommandChooser(playerId, apdCache) {
// Pre-generate the most common cache entries for better performance
const auto mapId = ActionPointDistancesCache::GetMapId(hexMap);
@@ -88,110 +74,38 @@ ShardokAIClient::ShardokAIClient(
apdCache->ConsolidateThreadLocalCache_Racy();
}
void CheckCommand(const CommandSPtr &realCommand, const CommandSPtr &guessedCommand) {
// Verify that the AI's guessed state produces the same available commands as the real state.
// We only compare fields that uniquely identify a command - metadata fields like action_points,
// will_unhide, next_round_target_info are not part of command identity.
void CheckCommand(const CommandProto &realDescriptor, const CommandProto &guessedDescriptor) {
string diff;
auto differencer = google::protobuf::util::MessageDifferencer();
differencer.IgnoreField(CommandProto::descriptor()->FindFieldByNumber(
CommandProto::kFollowUpCommandTypesFieldNumber));
differencer.ReportDifferencesToString(&diff);
if (!differencer.Compare(realDescriptor, guessedDescriptor)) {
printf("diff: %s\n\n", diff.c_str());
if (realCommand->GetCommandType() != guessedCommand->GetCommandType()) {
throw ShardokInternalErrorException("Command type mismatch between real and guessed state");
}
if (realCommand->GetPlayerId() != guessedCommand->GetPlayerId()) {
throw ShardokInternalErrorException("Player ID mismatch between real and guessed state");
}
if (realCommand->GetActorUnitId() != guessedCommand->GetActorUnitId()) {
throw ShardokInternalErrorException("Actor unit mismatch between real and guessed state");
}
if (realCommand->GetTargetRow() != guessedCommand->GetTargetRow() ||
realCommand->GetTargetColumn() != guessedCommand->GetTargetColumn()) {
throw ShardokInternalErrorException(
"Target coordinates mismatch between real and guessed state");
}
// For commands with odds (like FLEE), verify the odds match
if (realCommand->HasOdds() != guessedCommand->HasOdds()) {
throw ShardokInternalErrorException(
"Odds presence mismatch between real and guessed state");
}
if (realCommand->HasOdds() && guessedCommand->HasOdds()) {
if (realCommand->GetOddsPercentile() != guessedCommand->GetOddsPercentile()) {
throw ShardokInternalErrorException(
"Odds percentile mismatch between real and guessed state");
}
printf("Selected command descriptor\n%s\ndoes not match guessed\n%s\n\n",
realDescriptor.DebugString().c_str(),
guessedDescriptor.DebugString().c_str());
throw ShardokInternalErrorException("Illegal state for AI client");
}
}
auto ShardokAIClient::StandardChooseCommandIndex(
const GameSettingsSPtr &settings,
const GameStateW &guessedState,
const CommandListSPtr &realAvailableCommands) const -> CommandChoiceResults {
const vector<CommandProto> &realAvailableCommands) const -> CommandChoiceResults {
const auto settingsGetter = settings->GetGetter();
const auto guessedEngine = ShardokEngine(settings, guessedState);
const auto guessedCommands = guessedEngine.GetAvailableCommandsForAIPlayer(playerId);
const auto commandCount = guessedCommands->size();
// Calculate time budget based on game situation using new dynamic per-command settings
const auto timeBudget = CalculateTimeBudget(playerId, settings, guessedState, commandCount);
// Calculate time budget based on game situation using new settings
const auto timeBudget = CalculateTimeBudget(playerId, settings, guessedState);
// Configure MCTS based on proximity to enemy
// When far from enemy: use AVERAGING with maxPlayerFlips=0 (single-player lookahead)
// - AVERAGING naturally penalizes longer paths through variance
// - No opponent nodes, so no one-bad-child problem
// When close to enemy: use MINIMAX with maxPlayerFlips=1 (adversarial lookahead)
// - MINIMAX correctly models opponent choosing best response
// - Explores through one opponent turn for tactical accuracy
auto adjustedMCTSConfig = mctsConfig;
const auto guessedCommands = guessedEngine.GetAvailableCommandProtos(playerId, false);
const auto commandCount = guessedCommands.size();
// For fair evaluation: simulate leaves to opponent's turn start (maxSimulationFlips=1)
// This ensures all leaves are scored at the same game phase:
// - Leaves at playerFlips=0 (still my turn): simulate through END_TURN to playerFlips=1
// - Leaves at playerFlips=1 (opponent's turn): evaluate immediately
// Result: consistent comparison of "what happens after I end my turn"
// adjustedMCTSConfig.maxSimulatfixionFlips = 1;
// adjustedMCTSConfig.maxPlayerFlips = 0;
// if (timeBudget.isCloseToEnemy) {
// adjustedMCTSConfig.maxPlayerFlips = 1;
// adjustedMCTSConfig.backpropagationPolicy = mcts::MCTSBackpropagationPolicy::MINIMAX;
// if constexpr (kPerformanceLogging) {
// printf("MCTS Config: Close to enemy - using maxPlayerFlips=1, MINIMAX backprop\n");
// }
// } else {
// adjustedMCTSConfig.maxPlayerFlips = 0;
// adjustedMCTSConfig.backpropagationPolicy = mcts::MCTSBackpropagationPolicy::AVERAGING;
// if constexpr (kPerformanceLogging) {
// printf("MCTS Config: Far from enemy - using maxPlayerFlips=0, AVERAGING backprop\n");
// }
// }
assert(commandCount == realAvailableCommands->size());
// Verify that the AI's guessed state produces the same available commands as reality
assert(commandCount == realAvailableCommands.size());
for (size_t i = 0; i < commandCount; i++) {
CheckCommand((*realAvailableCommands)[i], (*guessedCommands)[i]);
}
// Extract values directly from settings for strategy selection
const auto maxRounds = settingsGetter.Backing().max_rounds();
const auto braveWaterCost = settingsGetter.Backing().brave_water_action_point_cost();
const auto battalionTypeGetter = [&settingsGetter](BattalionTypeId typeId) {
return settingsGetter.GetBattalionType(typeId);
};
// Create scorer for actual scoring during search - type selected at construction
std::unique_ptr<AIScoreCalculator> scorer;
switch (scoringCalculatorType) {
case ScoringCalculatorType::NORMALIZED:
scorer = MakeNormalizedAIScoreCalculator(settingsGetter, apdCache, alCache);
break;
case ScoringCalculatorType::MCTS_OPTIMIZED:
scorer = MakeMCTSOptimizedAIScoreCalculator(settingsGetter, apdCache, alCache);
break;
case ScoringCalculatorType::STANDARD:
default: scorer = MakeStandardAIScoreCalculator(settingsGetter, apdCache, alCache); break;
CheckCommand(realAvailableCommands[i], guessedCommands[i]);
}
// Determine strategy once for consistent scoring throughout iterative deepening
@@ -199,18 +113,15 @@ auto ShardokAIClient::StandardChooseCommandIndex(
const AIStrategy strategy = isDefender ? AIDefenderStrategySelector::BestDefenderStrategy(
guessedState,
castleCoords,
maxRounds,
apdCache,
battalionTypeGetter)
settingsGetter)
: AIAttackerStrategySelector::BestAttackerStrategy(
playerId,
guessedState,
castleCoords,
maxRounds,
apdCache,
alCache,
battalionTypeGetter,
braveWaterCost,
settingsGetter,
waterCrossingCommandChooser,
realAvailableCommands);
@@ -218,58 +129,12 @@ auto ShardokAIClient::StandardChooseCommandIndex(
IterativeDeepeningAI::SearchResult search_result;
if (aiAlgorithmType == AIAlgorithmType::MCTS) {
#ifdef ENABLE_MCTS_DEBUG_DUMP
// Set unique debug dump path for each action using timestamp
const auto now = std::chrono::system_clock::now();
const auto nowTime = std::chrono::system_clock::to_time_t(now);
const auto nowMs =
std::chrono::duration_cast<std::chrono::milliseconds>(now.time_since_epoch()) %
1000;
std::ostringstream pathStream;
pathStream << "/tmp/shardok_debug_"
<< std::put_time(std::localtime(&nowTime), "%Y%m%d_%H%M%S") << "_"
<< std::setfill('0') << std::setw(3) << nowMs.count() << "_p"
<< static_cast<int>(playerId) << ".txt";
adjustedMCTSConfig.debugDumpPath = pathStream.str();
// Also dump the game state to a file for reproduction
std::ostringstream statePathStream;
statePathStream << "/tmp/shardok_state_"
<< std::put_time(std::localtime(&nowTime), "%Y%m%d_%H%M%S") << "_"
<< std::setfill('0') << std::setw(3) << nowMs.count() << "_p"
<< static_cast<int>(playerId) << ".bin";
const std::string statePath = statePathStream.str();
// Write the flatbuffer game state to file using SaveTo method
if (guessedState.SaveTo(statePath)) {
printf("Game state dumped to: %s\n", statePath.c_str());
} else {
printf("Failed to dump game state to: %s\n", statePath.c_str());
}
#endif // ENABLE_MCTS_DEBUG_DUMP
// Using Monte Carlo Tree Search AI (with abstraction layer)
ShardokMCTSAI ai(
playerId,
isDefender,
strategy,
castleCoords,
*scorer,
apdCache,
alCache,
adjustedMCTSConfig);
search_result = ai.Search(settings, guessedState, timeBudget);
// Using Monte Carlo Tree Search AI
MCTSAI ai(playerId, isDefender, strategy, castleCoords, apdCache, alCache);
search_result = ai.Search(settings, guessedState, realAvailableCommands, timeBudget);
} else {
// Using Iterative Deepening AI (default)
IterativeDeepeningAI ai(
playerId,
isDefender,
strategy,
castleCoords,
*scorer,
apdCache,
battalionTypeGetter);
IterativeDeepeningAI ai(playerId, isDefender, strategy, castleCoords, apdCache, alCache);
search_result =
ai.IterativeSearch(settings, guessedState, realAvailableCommands, timeBudget);
}
@@ -288,12 +153,9 @@ auto ShardokAIClient::StandardChooseCommandIndex(
result.commandCountEvaluated,
result.availableCommandCount);
}
const auto chosenCommandType =
(*realAvailableCommands)[result.chosenIndex]->GetCommandType();
printf("ID AI: Search complete - achieved depth %d for best command %zu (%s)\n",
printf("ID AI: Search complete - achieved depth %d for best command %zu\n",
result.depthAchieved,
result.chosenIndex,
net::eagle0::shardok::common::CommandType_Name(chosenCommandType).c_str());
result.chosenIndex);
fflush(stdout);
}
@@ -304,20 +166,19 @@ auto ShardokAIClient::StandardChooseCommandIndex(
auto ShardokAIClient::LateRoundAttackerChooseCommandIndex(
const GameSettingsSPtr &settings,
const GameStateW &guessedState,
const CommandListSPtr &realAvailableCommands) const -> CommandChoiceResults {
const vector<CommandProto> &realAvailableCommands) const -> CommandChoiceResults {
if (const auto dismissCommand = std::ranges::find_if(
*realAvailableCommands,
[](const CommandSPtr &cmd) {
return cmd->GetCommandType() ==
net::eagle0::shardok::common::DISMISS_UNIT_COMMAND;
realAvailableCommands,
[](const net::eagle0::shardok::api::CommandDescriptor &cmd) {
return cmd.type() == net::eagle0::shardok::common::DISMISS_UNIT_COMMAND;
});
dismissCommand == realAvailableCommands->end()) {
dismissCommand == realAvailableCommands.end()) {
return StandardChooseCommandIndex(settings, guessedState, realAvailableCommands);
} else {
CommandChoiceResults results{};
results.chosenIndex =
static_cast<size_t>(std::distance(realAvailableCommands->begin(), dismissCommand));
results.availableCommandCount = realAvailableCommands->size();
static_cast<size_t>(std::distance(realAvailableCommands.begin(), dismissCommand));
results.availableCommandCount = realAvailableCommands.size();
results.depthAchieved = 1; // Simple heuristic choice
results.commandCountEvaluated = 1; // Only evaluated one command type
results.completionReason =
@@ -329,31 +190,24 @@ auto ShardokAIClient::LateRoundAttackerChooseCommandIndex(
auto ShardokAIClient::FinalRoundAttackerChooseCommandIndex(
const GameSettingsSPtr &settings,
const GameStateW &guessedState,
const CommandListSPtr &realAvailableCommands) const -> CommandChoiceResults {
const auto fleeCommand =
std::ranges::find_if(*realAvailableCommands, [](const CommandSPtr &cmd) {
return cmd->GetCommandType() == net::eagle0::shardok::common::FLEE_COMMAND;
const vector<CommandProto> &realAvailableCommands) const -> CommandChoiceResults {
const auto fleeCommand = std::ranges::find_if(
realAvailableCommands,
[](const net::eagle0::shardok::api::CommandDescriptor &cmd) {
return cmd.type() == net::eagle0::shardok::common::FLEE_COMMAND;
});
if (fleeCommand == realAvailableCommands->end()) {
if (fleeCommand == realAvailableCommands.end()) {
return LateRoundAttackerChooseCommandIndex(settings, guessedState, realAvailableCommands);
}
// Extract values directly from settings for flee decision evaluation
const auto settingsGetter = settings->GetGetter();
const auto maxRounds = settingsGetter.Backing().max_rounds();
const auto minimumFleeOddsThreshold = settingsGetter.Backing().ai_minimum_flee_odds_threshold();
const auto desperateFleeThreshold = settingsGetter.Backing().ai_desperate_flee_threshold();
// Use the flee decision calculator
const auto fleeDecision = AIFleeDecisionCalculator::EvaluateFleeVsFight(
playerId,
settings->GetGetter(),
guessedState,
realAvailableCommands,
fleeCommand,
maxRounds,
minimumFleeOddsThreshold,
desperateFleeThreshold,
#ifdef DEBUG_FLEE_DECISIONS
true // Enable debug logging
#else
@@ -364,7 +218,7 @@ auto ShardokAIClient::FinalRoundAttackerChooseCommandIndex(
if (fleeDecision.shouldFlee) {
CommandChoiceResults results{};
results.chosenIndex = fleeDecision.commandIndex;
results.availableCommandCount = realAvailableCommands->size();
results.availableCommandCount = realAvailableCommands.size();
results.depthAchieved = 1; // Heuristic choice
results.commandCountEvaluated = 1; // Only evaluated one command type
results.completionReason = EvaluationCompletionReason::RAN_OUT_OF_COMMANDS;
@@ -378,7 +232,7 @@ auto ShardokAIClient::FinalRoundAttackerChooseCommandIndex(
auto ShardokAIClient::ChooseCommandIndex(
const GameSettingsSPtr &settings,
const GameStateView &gsv,
const CommandListSPtr &realAvailableCommands) const -> CommandChoiceResults {
const vector<CommandProto> &realAvailableCommands) const -> CommandChoiceResults {
static int typeChosenCount[net::eagle0::shardok::common::CommandType_MAX + 1];
static int totalChoices = 0;
@@ -397,7 +251,7 @@ auto ShardokAIClient::ChooseCommandIndex(
results = StandardChooseCommandIndex(settings, guessedState, realAvailableCommands);
}
const auto chosenType = (*realAvailableCommands)[results.chosenIndex]->GetCommandType();
const auto chosenType = realAvailableCommands[results.chosenIndex].type();
typeChosenCount[static_cast<int>(chosenType)]++;
totalChoices++;
@@ -423,8 +277,8 @@ auto ShardokAIClient::ChooseCommandIndex(
auto ShardokAIClient::ChooseCommandIndex(const ShardokEngine &engine) const
-> CommandChoiceResults {
if (const auto &availableCommands = engine.GetAvailableCommandsForAIPlayer(playerId);
availableCommands->empty()) {
if (const auto &availableCommands = engine.GetAvailableCommandProtos(playerId, false);
availableCommands.empty()) {
printf("no commands for player %d\n", playerId);
throw ShardokInternalErrorException(
"Asked to choose a command, but there are none available");
@@ -12,13 +12,11 @@
#include <vector>
#include "src/main/cpp/net/eagle0/common/RandomGenerator.hpp"
#include "src/main/cpp/net/eagle0/common/mcts/abstract/MCTSTypes.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIConfig.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIScoreCalculator.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AITimeBudget.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIWaterCrossingCommandChooser.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/IterativeDeepeningAI.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/score/AIScoreCalculator.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokCommand.hpp"
#include "src/main/protobuf/net/eagle0/shardok/api/game_state_view.pb.h"
namespace shardok {
@@ -42,33 +40,29 @@ private:
const PlayerId playerId;
const bool isDefender;
const AIAlgorithmType aiAlgorithmType;
const ScoringCalculatorType scoringCalculatorType;
APDCache apdCache = std::make_shared<ActionPointDistancesCache>();
ALCache alCache;
const AIWaterCrossingCommandChooser waterCrossingCommandChooser;
// MCTS configuration (only used when aiAlgorithmType == MCTS)
mcts::MCTSConfig mctsConfig;
[[nodiscard]] auto StandardChooseCommandIndex(
const GameSettingsSPtr& settings,
const GameStateW& guessedState,
const CommandListSPtr& realAvailableCommands) const -> CommandChoiceResults;
const vector<CommandProto>& realAvailableCommands) const -> CommandChoiceResults;
[[nodiscard]] auto LateRoundAttackerChooseCommandIndex(
const GameSettingsSPtr& settings,
const GameStateW& guessedState,
const CommandListSPtr& realAvailableCommands) const -> CommandChoiceResults;
const vector<CommandProto>& realAvailableCommands) const -> CommandChoiceResults;
[[nodiscard]] auto FinalRoundAttackerChooseCommandIndex(
const GameSettingsSPtr& settings,
const GameStateW& guessedState,
const CommandListSPtr& realAvailableCommands) const -> CommandChoiceResults;
const vector<CommandProto>& realAvailableCommands) const -> CommandChoiceResults;
[[nodiscard]] auto ChooseCommandIndex(
const GameSettingsSPtr& settings,
const net::eagle0::shardok::api::GameStateView& gsv,
const CommandListSPtr& realAvailableCommands) const -> CommandChoiceResults;
const vector<CommandProto>& realAvailableCommands) const -> CommandChoiceResults;
public:
explicit ShardokAIClient(
@@ -76,19 +70,13 @@ public:
bool isDefender,
const HexMap* hexMap,
const SettingsGetter& settings,
AIAlgorithmType aiAlgorithmType,
ScoringCalculatorType scoringCalculatorType,
const mcts::MCTSConfig& mctsConfig);
AIAlgorithmType aiAlgorithmType = AIAlgorithmType::ITERATIVE_DEEPENING);
~ShardokAIClient() = default;
[[nodiscard]] auto GetPlayerId() const -> PlayerId { return playerId; }
[[nodiscard]] auto ChooseCommandIndex(const ShardokEngine& engine) const
-> CommandChoiceResults;
// MCTS configuration methods (only relevant when using MCTS algorithm)
[[nodiscard]] auto GetMCTSConfig() const -> const mcts::MCTSConfig& { return mctsConfig; }
void SetMCTSConfig(const mcts::MCTSConfig& config) { mctsConfig = config; }
};
} // namespace shardok
@@ -1,24 +1,29 @@
load("//tools:copts.bzl", "COPTS")
cc_library(
name = "shardok_mcts_ai",
srcs = ["ShardokMCTSAI.cpp"],
hdrs = ["ShardokMCTSAI.hpp"],
name = "mcts_ai",
srcs = ["MCTSAI.cpp"],
hdrs = ["MCTSAI.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai:__pkg__",
"//src/main/cpp/net/eagle0/shardok/ai_performance_runner:__pkg__",
"//src/test/cpp/net/eagle0/shardok/ai/mcts:__pkg__",
"//src/test/cpp/net/eagle0/shardok/ai:__pkg__",
],
deps = [
"//src/main/cpp/net/eagle0/common/mcts/abstract:abstract_mcts_ai",
"//src/main/cpp/net/eagle0/common:random_generator",
"//src/main/cpp/net/eagle0/common:sequence_random_generator",
"//src/main/cpp/net/eagle0/shardok/ai:ai_command_filter",
"//src/main/cpp/net/eagle0/shardok/ai:ai_iterative_deepening", # For SearchResult compatibility
"//src/main/cpp/net/eagle0/shardok/ai:ai_score_calculator",
"//src/main/cpp/net/eagle0/shardok/ai:ai_strategy",
"//src/main/cpp/net/eagle0/shardok/ai:ai_time_budget",
"//src/main/cpp/net/eagle0/shardok/ai/mcts/adapters:shardok_mcts_factory",
"//src/main/cpp/net/eagle0/shardok/ai/mcts/internal:mcts_node",
"//src/main/cpp/net/eagle0/shardok/library:engine",
"//src/main/cpp/net/eagle0/shardok/library:game_state_w",
"//src/main/cpp/net/eagle0/shardok/library:shardok_c_types",
"//src/main/cpp/net/eagle0/shardok/library/action_point_distances:action_point_distances_cache",
"//src/main/cpp/net/eagle0/shardok/library/settings:game_settings",
"//src/main/protobuf/net/eagle0/shardok/api:command_descriptor_cc_proto",
"//src/main/protobuf/net/eagle0/shardok/common:command_type_cc_proto",
],
)
@@ -1,813 +0,0 @@
# Chance Nodes in MCTS for Shardok
## Problem Statement
### Current Behavior
The current MCTS implementation uses a fixed roll (50th percentile) for all probabilistic outcomes during simulation. This creates several issues:
1. **Binary success actions overvalued**: A START_FIRE command with 51% success is treated as always succeeding, making it appear better than it actually is.
2. **Discontinuity at 50%**: Actions with 49% vs 51% success have dramatically different evaluations, when they should be similar.
3. **Variable-outcome actions simplified**: Melee/archery attacks with damage ranges are evaluated at a single point rather than their full distribution.
### Example Issue
```
START_FIRE with 51% success:
- Current MCTS: Assumes always succeeds (roll = 50)
- Reality: Succeeds 51% of time, fails 49% of time
- Result: AI overvalues this action
```
### How Iterative Deepening Solves This
The iterative deepening AI (see `AICommandEvaluator.cpp:352-393`) handles randomness correctly:
```cpp
// For actions with odds (binary success/fail):
// 1. Evaluate success outcome with representative roll
auto [successScore, successLookahead] = EvaluateWithRandomness(
...,
std::make_shared<SequenceRandomGenerator>(std::vector{1.0 - successChance / 2.0})
);
// 2. Evaluate failure outcome with representative roll
auto [failureScore, failureLookahead] = EvaluateWithRandomness(
...,
std::make_shared<SequenceRandomGenerator>(std::vector{(1.0 - successChance) / 2.0})
);
// 3. Compute weighted average (expected value)
immediateScore = std::lerp(failureScore, successScore, successChance);
lookaheadScore = std::lerp(failureLookahead.get(), successLookahead.get(), successChance);
```
This is essentially an implicit form of chance nodes - evaluating both outcomes and weighting by probability.
## Chance Nodes Concept
### Classic MCTS with Chance Nodes
In games with randomness (e.g., backgammon), MCTS uses two types of nodes:
1. **Decision Nodes**: Player chooses an action
- Selection uses UCB formula (exploration/exploitation tradeoff)
- One child per legal action
2. **Chance Nodes**: Nature determines outcome
- Selection uses expectation (weighted by probability)
- One child per possible outcome
```
Decision Node (Player to move)
├─ Action A
│ └─ Chance Node
│ ├─ Outcome 1 (prob 0.3) → Game State
│ ├─ Outcome 2 (prob 0.5) → Game State
│ └─ Outcome 3 (prob 0.2) → Game State
└─ Action B
└─ Deterministic → Game State
```
### Example: START_FIRE in Shardok
**Current approach:**
```
State S
└─ START_FIRE (roll=50)
└─ State S' (fire always starts)
```
**With chance nodes:**
```
State S
└─ START_FIRE action
└─ Chance Node
├─ Success (51%) → State S_success (fire started)
└─ Failure (49%) → State S_failure (no fire, vigor spent)
```
### Value Propagation
**Decision nodes:** Maximize/minimize over children (depending on player)
**Chance nodes:** Expected value over children (weighted by probability)
```cpp
// Decision node value (max for current player)
value = max(child.value for child in children)
// Chance node value (expectation)
value = sum(prob[i] * child[i].value for i in outcomes)
```
## Implementation Approaches
### Option 1: Explicit Chance Nodes (Full Implementation)
Modify the MCTS tree structure to explicitly represent chance nodes.
**Pros:**
- Theoretically sound
- Handles arbitrary outcome distributions
- Clear separation of decision vs chance
**Cons:**
- Significant code changes
- Larger tree (more memory)
- More complex tree traversal
**Tree Structure:**
```cpp
enum class NodeType { DECISION, CHANCE };
struct MCTSNode {
NodeType type;
// For decision nodes
MCTSPlayerId player;
std::vector<std::unique_ptr<MCTSAction>> actions;
std::vector<std::unique_ptr<MCTSNode>> children; // One per action
// For chance nodes
std::vector<double> probabilities; // One per outcome
std::vector<std::unique_ptr<MCTSNode>> outcomes; // One per outcome
double visits;
double totalReward;
};
```
**Selection Phase:**
```cpp
MCTSNode* select(MCTSNode* node) {
while (!node->isLeaf()) {
if (node->type == DECISION) {
// Use UCB to select action
node = selectChildUCB(node);
} else { // CHANCE node
// Use probability-weighted selection
node = selectOutcomeByProbability(node);
}
}
return node;
}
```
**Backpropagation:**
```cpp
void backpropagate(MCTSNode* node, double reward) {
while (node != nullptr) {
node->visits++;
if (node->type == DECISION) {
node->totalReward += reward; // Sum for averaging
} else { // CHANCE node
node->totalReward += reward; // Still sum, but averaged differently
}
node = node->parent;
}
}
```
### Option 2: Implicit Chance Nodes (Hybrid Approach)
Keep the current tree structure but sample outcomes during expansion/simulation.
**Pros:**
- Smaller code changes
- More memory efficient
- Easier to implement incrementally
**Cons:**
- Less theoretically pure
- May need more visits to converge
- Sampling introduces variance
**Approach:**
```cpp
// During expansion
std::unique_ptr<MCTSGameState> expand(
const MCTSGameState& state,
const MCTSAction& action
) {
if (action.isDeterministic()) {
return applyActionDeterministic(state, action);
} else {
// Sample an outcome based on probabilities
auto outcome = sampleOutcome(action);
return applyActionWithOutcome(state, action, outcome);
}
}
```
**For binary actions (e.g., START_FIRE):**
```cpp
// Expand creates one of two children based on sampling
if (random() < successProbability) {
return applySuccess(state, action);
} else {
return applyFailure(state, action);
}
// Over many visits, visit ratio will approach probability ratio
// E.g., 51% success action will have ~51% success children, 49% failure children
```
### Option 3: Determinized Sampling (Simplest)
Pre-sample all random outcomes at the start of each simulation rollout.
**Pros:**
- Minimal code changes
- Easy to understand
- Works with existing tree structure
**Cons:**
- May converge slowly
- Doesn't explicitly represent probability
- Can waste simulations on unlikely outcomes
**Approach:**
```cpp
// At start of each simulation
std::vector<double> rollSequence = generateRollSequence(maxDepth);
// Use sequence during simulation
auto state = rootState;
for (int depth = 0; depth < maxDepth; depth++) {
auto action = selectAction(state);
state = applyAction(state, action, rollSequence[depth]);
}
```
## Recommended Approach: Progressive Enhancement
Implement in phases to manage complexity:
### Phase 1: Binary Chance Nodes (Explicit)
Start with actions that have clear success/failure outcomes (e.g., START_FIRE, EXTINGUISH_FIRE, RAISE_DEAD):
1. Identify binary actions (commands with `HasOdds()`)
2. Add chance node support for these actions only
3. Modify tree expansion to create chance nodes
4. Update selection/backpropagation for chance nodes
**Implementation:**
```cpp
// In ShardokGameEngine::getLegalActions()
// Mark which actions require chance nodes
struct ActionMetadata {
std::unique_ptr<MCTSAction> action;
bool requiresChanceNode;
double successProbability; // If requiresChanceNode = true
};
```
```cpp
// In tree expansion
if (action.requiresChanceNode) {
// Create chance node with two children
auto chanceNode = std::make_unique<MCTSNode>(CHANCE);
chanceNode->probabilities = {successProb, 1.0 - successProb};
// Expand both outcomes
chanceNode->outcomes.push_back(applySuccess(state, action));
chanceNode->outcomes.push_back(applyFailure(state, action));
return chanceNode;
} else {
// Normal deterministic expansion
return applyAction(state, action);
}
```
### Phase 2: Multi-Outcome Actions
Extend to actions with multiple outcomes (e.g., melee damage ranges):
1. Discretize continuous distributions into buckets
2. For melee/archery, use 3-5 representative damage values (min, low, avg, high, max)
3. Compute probabilities for each bucket
4. Create chance nodes with multiple children
**Example: Melee Attack**
```cpp
// Instead of sampling full damage distribution,
// use representative values
struct DamageBucket {
int damageValue; // Representative damage
double probability; // Probability of this range
};
// For a melee attack that can deal 10-20 damage
std::vector<DamageBucket> buckets = {
{10, 0.1}, // Min damage (unlucky)
{13, 0.2}, // Low damage
{15, 0.4}, // Average damage
{17, 0.2}, // High damage
{20, 0.1} // Max damage (lucky)
};
```
### Phase 3: Optimization
Once chance nodes work correctly:
1. Add transposition table support for chance nodes
2. Optimize memory layout
3. Consider progressive widening (start with 2 outcomes, expand to more if visited often)
4. Profile and tune
## Design Decisions
### How to Represent Outcomes?
**Option A: Explicit state copies**
```cpp
struct ChanceNode {
std::vector<std::unique_ptr<MCTSGameState>> outcomeStates;
std::vector<double> probabilities;
};
```
**Option B: Lazy evaluation**
```cpp
struct ChanceNode {
MCTSGameState baseState;
MCTSAction action;
std::vector<int> outcomeRolls; // Roll values for each outcome
std::vector<double> probabilities;
// Compute state on-demand
MCTSGameState getOutcome(size_t index) {
return applyActionWithRoll(baseState, action, outcomeRolls[index]);
}
};
```
**Recommendation:** Option B - lazy evaluation. Only materialize states when visited.
### How Many Outcomes per Action?
**Binary actions (START_FIRE, etc.):**
- Exactly 2 outcomes (success/fail)
- Use exact probabilities from `GetOddsPercentile()`
**Damage actions (MELEE, ARCHERY):**
- Start with 3 outcomes (low/med/high)
- Can expand to 5 if needed for accuracy
- Use representative rolls: 10th, 50th, 90th percentile
**Complex actions (METEOR):**
- Consider 2-3 outcomes initially
- Can model as "hits N enemies" for N in {0, 1, 2, 3+}
### How to Handle Transposition Table?
**Challenge:** Same state can be reached via different chance outcomes
**Solution:**
- Hash based on game state only (not the path taken)
- When looking up, return cached evaluation if state matches
- This is already how transposition tables work!
```cpp
// Current approach works fine:
auto hash = computeHash(gameState); // Doesn't include how we got here
if (auto cached = transpositionTable.lookup(hash)) {
return cached->value;
}
```
### Selection at Chance Nodes
**During tree traversal:**
```cpp
size_t selectOutcome(const ChanceNode& node) {
// Option 1: Sample by probability (introduces variance)
double r = random();
double cumulative = 0.0;
for (size_t i = 0; i < node.probabilities.size(); i++) {
cumulative += node.probabilities[i];
if (r < cumulative) return i;
}
// Option 2: Round-robin weighted by visit count vs probability
// (Explore under-visited outcomes more)
size_t leastVisited = findMostUnderExploredOutcome(node);
return leastVisited;
}
```
**Recommendation:** Use Option 2 to ensure all outcomes get explored proportionally.
## Integration Points
### Modified Functions
1. **`ShardokGameEngine::getLegalActions()`**
- Add metadata about which actions need chance nodes
- Return action + probability information
2. **`ShardokGameEngine::applyAction()`**
- For binary actions, return both possible outcomes
- Or: take an explicit outcome index parameter
3. **`AbstractMCTSAI::selection()`**
- Handle chance nodes differently from decision nodes
- Use probability-weighted selection instead of UCB
4. **`AbstractMCTSAI::expand()`**
- Create chance node children for probabilistic actions
- May create multiple child nodes per action
5. **`AbstractMCTSAI::backpropagate()`**
- Update all nodes in path (both decision and chance)
- Value calculation already handles this correctly (just averages)
### New Functions Needed
```cpp
// In ShardokGameEngine
struct ChanceOutcome {
int roll; // The dice roll that produces this outcome
double probability; // Probability of this outcome
};
std::vector<ChanceOutcome> getChanceOutcomes(const MCTSAction& action) const;
```
```cpp
// In MCTSNode
bool isChanceNode() const;
const std::vector<double>& getOutcomeProbabilities() const;
```
## Testing Strategy
### Unit Tests
1. **Binary action correctness**
```cpp
TEST(ChanceNodes, BinaryActionExpectedValue) {
// START_FIRE with 60% success
// Run MCTS with chance nodes
// Verify: visits to success ~= 60%, visits to failure ~= 40%
// Verify: expected value matches manual calculation
}
```
2. **Comparison with iterative deepening**
```cpp
TEST(ChanceNodes, MatchesIterativeDeepening) {
// Same position, both AIs
// Should choose same action
// Scores should be similar (within variance)
}
```
3. **Transposition table with chance**
```cpp
TEST(ChanceNodes, TranspositionConsistency) {
// Two paths to same state via different chance outcomes
// Should reuse cached evaluation
}
```
### Integration Tests
1. Compare MCTS with/without chance nodes on test positions
2. Verify that chance nodes reduce overvaluation of marginal actions
3. Performance test: measure slowdown (expect 1.5-2x for binary actions)
### Real-World Validation
Run the problematic START_FIRE scenario:
- With current MCTS: Should overvalue START_FIRE
- With chance nodes: Should correctly weight success/failure
- Expected: END_TURN should get significantly more visits
## Performance Considerations
### Memory Overhead
**Per chance node:**
- Probability vector: `N * sizeof(double)` (N = number of outcomes)
- Outcome children: `N * sizeof(unique_ptr)`
- For binary: ~32 bytes per chance node
**Estimate:**
- Current tree: ~100K nodes per search
- With chance nodes: ~150K nodes (50% actions are probabilistic)
- Extra memory: ~50K * 32 bytes = ~1.6 MB
- **Acceptable overhead**
### Computational Overhead
**Per simulation:**
- Current: 1 path through tree
- With chance nodes: Still 1 path, but more nodes
- Overhead: ~20-30% (more node visits)
**Mitigation:**
- Transposition table helps (same states via different paths)
- Progressive widening (start with 2 outcomes, expand if visited often)
- Lazy state evaluation (don't materialize until needed)
### Convergence Speed
Chance nodes may require more visits to converge because:
- More children per action (branching factor increases)
- Outcomes need proportional exploration
**Mitigation:**
- Use visit count thresholds before expanding chance nodes
- Consider progressive widening (UCT-ProgressiveWidening)
## Migration Path
### Step 1: Infrastructure (1-2 days)
- Add `NodeType` enum and metadata to MCTSNode
- Implement chance node creation (without using them yet)
- Add unit tests for chance node structure
### Step 2: Binary Actions (2-3 days)
- Identify all binary success/fail actions
- Modify expansion to create chance nodes for these
- Update selection/backpropagation
- Test on START_FIRE scenario
### Step 3: Integration Testing (1 day)
- Run full MCTS tests with chance nodes enabled
- Compare with iterative deepening on test positions
- Validate that it fixes the START_FIRE overvaluation
### Step 4: Multi-Outcome Actions (2-3 days)
- Implement damage bucketing for MELEE/ARCHERY
- Create chance nodes with 3-5 outcomes
- Test on combat scenarios
### Step 5: Optimization (1-2 days)
- Profile performance
- Add progressive widening if needed
- Tune outcome granularity
### Step 6: Documentation & Cleanup (1 day)
- Document the new approach
- Clean up code
- Add comprehensive tests
## Alternative: Simpler Hybrid Approach
If full chance nodes are too complex, consider a hybrid:
1. **Keep current tree structure** (no explicit chance nodes)
2. **During expansion:** Sample outcome and create one child
3. **Over many simulations:** Statistics converge to correct probabilities
4. **Add outcome tracking:** Store "which outcome" in edge/node metadata
**Example:**
```cpp
// Expansion samples an outcome
auto expand(state, action) {
if (action.hasBinaryOutcome()) {
// Sample once
bool success = (random() < successProb);
// Store which outcome this edge represents
edge.metadata.outcome = success ? OUTCOME_SUCCESS : OUTCOME_FAILURE;
return applyWithOutcome(state, action, success);
}
}
// Selection prioritizes under-explored outcomes
auto selectChild(node) {
// Find action where outcome distribution is unbalanced
// E.g., 60% success action should have ~60% success children
// If we have 80% success children, prefer exploring failure
}
```
This is simpler but less theoretically sound. It's a reasonable starting point if full chance nodes prove too complex.
## Comparison: Chance Nodes vs Open-Loop MCTS
### What is Open-Loop MCTS?
**Open-loop MCTS** (also called "determinization MCTS" or "information set MCTS") is an alternative approach to handling randomness:
1. At the **start of each simulation**, sample all random outcomes needed for that simulation
2. Play out the entire simulation using those fixed random values
3. Different simulations use different random seeds
4. The tree structure doesn't explicitly model randomness - it's all in the rollouts
**Example implementation:**
```cpp
// At start of simulation
std::vector<double> rollSequence = sampleRolls(maxDepth); // Pre-sample all rolls
// During simulation
MCTSNode* node = root;
for (int depth = 0; depth < maxDepth; depth++) {
Action action = selectAction(node);
node = applyAction(node, action, rollSequence[depth]); // Use pre-sampled roll
}
```
### Open-Loop MCTS for Shardok
**How it would work:**
```cpp
// Each simulation samples a "possible world"
void simulate(MCTSNode* root) {
// Sample random rolls for this simulation
auto rolls = generateRollSequence(); // e.g., {0.45, 0.78, 0.23, ...}
// Play out simulation using these fixed rolls
auto state = root->state;
for (int depth = 0; depth < maxDepth; depth++) {
auto action = selectAction(state);
state = applyAction(state, action, rolls[depth]);
}
double reward = evaluate(state);
backpropagate(root, reward);
}
```
**Would this fix the START_FIRE issue?**
**Yes** - partially. Different simulations would see different outcomes:
- Some simulations: START_FIRE succeeds (roll < 0.51)
- Some simulations: START_FIRE fails (roll >= 0.51)
- Over many simulations, the action's value would approach the expected value
**However**, it's less efficient than chance nodes because:
- Needs MORE simulations to converge
- Wastes effort exploring unlikely scenarios equally with likely ones
- Doesn't explicitly guide exploration based on probability
### Detailed Comparison
| Aspect | Chance Nodes (Closed-Loop) | Open-Loop MCTS | Current (Fixed Roll) |
|--------|---------------------------|----------------|----------------------|
| **Randomness Handling** | Explicit in tree structure | Implicit in simulation sampling | Fixed roll=50 |
| **Convergence Speed** | Fast - probabilities guide search | Slower - needs more samples | N/A (wrong answer) |
| **Memory Usage** | Higher (more nodes) | Lower (no extra nodes) | Lowest |
| **Implementation Complexity** | High (tree structure changes) | Medium (sampling layer) | Low (current) |
| **Theoretical Soundness** | Highest (models true game tree) | Medium (approximation via sampling) | Low (assumes fixed outcome) |
| **START_FIRE Fix** | ✅ Yes, accurately | ✅ Yes, eventually | ❌ No |
| **Efficiency** | Most efficient per simulation | Less efficient (wasted samples) | Efficient but wrong |
| **Handles Hidden Information** | Poor | Excellent | N/A |
### When to Prefer Each Approach
**Prefer Chance Nodes when:**
- Randomness outcomes are discrete and enumerable (e.g., binary success/fail)
- Probabilities are known precisely
- You want fastest convergence to correct answer
- Game tree is the primary concern (no hidden information)
- **This is Shardok's situation**
**Prefer Open-Loop when:**
- Randomness is continuous and high-dimensional
- Hidden information or imperfect information is present
- Simplicity is paramount
- You can afford many simulations
- Used in games like poker, bridge, Skat
### Why Chance Nodes are Better for Shardok
1. **Discrete outcomes**: Most Shardok randomness is binary (success/fail) or small discrete sets (damage ranges)
- START_FIRE: 2 outcomes (success/fail)
- MELEE: Can bucket into 3-5 damage ranges
- Not continuous - perfect fit for chance nodes
2. **Known probabilities**: We have exact probabilities from `GetOddsPercentile()`
- Chance nodes can use exact probabilities
- Open-loop just samples blindly
3. **No hidden information**: Shardok is perfect information (all units visible to AI)
- Chance nodes' main weakness doesn't apply
- Open-loop's main strength doesn't help
4. **Convergence matters**: Limited simulation budget
- Need to converge quickly
- Chance nodes achieve this better
5. **Existing infrastructure**: We already have deterministic state transitions
- Adding chance nodes builds on what we have
- Open-loop would need different rollout structure
### Performance Analysis
**Chance Nodes:**
```
Time per simulation: 1.3x current
Simulations needed: 10,000 to converge
Total time: 13,000x units
Memory: 1.5x current (extra chance nodes)
```
**Open-Loop:**
```
Time per simulation: 1.0x current (same as now)
Simulations needed: 30,000 to converge (more variance)
Total time: 30,000x units
Memory: 1.0x current (no extra nodes)
```
**Result:** Chance nodes are **2.3x faster overall** despite being slower per simulation, because they converge with fewer simulations.
### Hybrid Approach: Best of Both Worlds?
Could we combine them?
**Idea:** Use chance nodes for high-probability branches, open-loop for rare events
```cpp
if (probability > 0.1 && outcomeCount <= 5) {
// Use explicit chance node
createChanceNode(outcomes, probabilities);
} else {
// Use open-loop sampling
sampleOutcome();
}
```
**Verdict:** Probably not worth the complexity. Shardok's randomness is simple enough that chance nodes handle everything well.
### Recommendation for Shardok
**Use Chance Nodes**, specifically:
1. **Phase 1:** Binary actions (START_FIRE, RAISE_DEAD, etc.)
- 2 outcomes, exact probabilities
- Biggest bang for buck
2. **Phase 2:** Damage ranges (MELEE, ARCHERY)
- 3-5 buckets
- Still manageable
3. **If needed:** Could fall back to open-loop for complex actions
- E.g., METEOR with many possible outcomes
- But likely unnecessary
### Why Not Open-Loop?
While open-loop would eventually fix the START_FIRE issue, it has significant downsides for Shardok:
1. **Slower convergence**: Needs 2-3x more simulations
2. **Doesn't leverage known probabilities**: We have exact odds, why ignore them?
3. **Less interpretable**: Harder to debug why AI chose an action
4. **Doesn't align with iterative deepening**: We want MCTS to match the proven algorithm
The only advantage of open-loop (simplicity) is outweighed by chance nodes' efficiency and correctness.
### Could We Use Current Approach + Better Sampling?
**Idea:** Keep fixed rolls but use different rolls per simulation?
```cpp
// Instead of always roll=50
double roll = random(); // Different each simulation
```
**Problem:** This is essentially open-loop without the tree!
- Even slower to converge
- Tree doesn't learn the outcome probabilities
- Worst of both worlds
**Verdict:** No, this doesn't help. If we're going to sample, do it properly (open-loop). Otherwise, use chance nodes.
### Final Verdict
**For Shardok, chance nodes are clearly superior:**
- ✅ Faster convergence (2-3x vs open-loop)
- ✅ Leverages exact probabilities
- ✅ Perfect fit for discrete outcomes
- ✅ Aligns with iterative deepening approach
- ✅ Better debuggability and interpretability
- ❌ More complex implementation (but manageable)
Open-loop would be a fallback if chance nodes prove too difficult, but given the benefits and the bounded complexity (only binary and small discrete outcomes), chance nodes are the right choice.
## Conclusion
Implementing chance nodes will fix the overvaluation of marginal probabilistic actions like START_FIRE with 51% success. The recommended approach is:
1. Start with **explicit chance nodes for binary actions**
2. Use **lazy state evaluation** to minimize memory
3. **Progressive enhancement** - binary first, then multi-outcome
4. Compare with iterative deepening to validate correctness
Expected benefits:
- More accurate action evaluation
- Better handling of probabilistic outcomes
- Closer alignment with theoretical MCTS
- Fixes the START_FIRE issue without tuning heuristics
Expected costs:
- ~20-30% slower per simulation (more nodes)
- ~1-2MB extra memory
- ~1-2 weeks development time
The benefits significantly outweigh the costs for a more theoretically sound and accurate AI.
@@ -0,0 +1,869 @@
//
// MCTS-based AI implementation for Shardok
//
#include "MCTSAI.hpp"
#include <algorithm>
#include <cmath>
#include <future>
#include <limits>
#include <mutex>
#include <random>
#include <thread>
#include "internal/MCTSNode.hpp"
#include "src/main/cpp/net/eagle0/common/SequenceRandomGenerator.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AICommandFilter.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIScoreCalculator.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokEngine.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/util/HexMapUtils.hpp"
#include "src/main/protobuf/net/eagle0/shardok/common/command_type.pb.h"
namespace shardok {
// Type alias for internal MCTSNode
using MCTSNode = internal::MCTSNode;
// Static helper for average random generator
static const std::vector _averageSequence = {0.5};
static const auto _averageGenerator = std::make_shared<SequenceRandomGenerator>(_averageSequence);
MCTSAI::MCTSAI(
const PlayerId playerId,
const bool isDefender,
AIStrategy strategy,
const CoordsSet& castleCoords,
const APDCache& apdCache,
const ALCache& alCache,
MCTSConfig config)
: playerId(playerId),
isDefender(isDefender),
strategy(std::move(strategy)),
castleCoords(castleCoords),
apdCache(apdCache),
alCache(alCache),
config(std::move(config)) {}
auto MCTSAI::Search(
const GameSettingsSPtr& settings,
const GameStateW& state,
const std::vector<CommandProto>& commands,
const AITimeBudget& budget) const -> SearchResult {
const auto startTime = std::chrono::steady_clock::now();
SearchResult result;
if (commands.empty()) {
result.searchCompleted = true;
result.completionReason = EvaluationCompletionReason::RAN_OUT_OF_COMMANDS;
return result;
}
if (commands.size() == 1) {
result.searchCompleted = true;
result.bestCommandIndex = 0;
result.bestScore = 0;
result.availableCommandCount = 1;
result.depthAchieved = 1;
result.commandCountEvaluated = 1;
result.completionReason = EvaluationCompletionReason::RAN_OUT_OF_COMMANDS;
return result;
}
const auto& settingsGetter = settings->GetGetter();
// Compute critical tiles once to avoid 8.5% runtime overhead in ShardokEngine construction
const auto criticalTiles = GetCriticalTileLocations(state->hex_map());
const auto guessedEngine = ShardokEngine(settings, state, criticalTiles);
const auto deadline = startTime + budget.remainingBudget;
// Build MCTS tree
auto rootNode = BuildMCTSTree(guessedEngine, settingsGetter, criticalTiles, deadline);
if (rootNode) {
// Get best command from tree
const MCTSNode* bestChild = rootNode->GetBestFinalChild();
if (bestChild) {
result.searchCompleted = true;
result.bestCommandIndex = bestChild->commandIndex;
result.bestScore = bestChild->lookaheadScore;
result.completionReason = EvaluationCompletionReason::RAN_OUT_OF_COMMANDS;
// Calculate max depth reached in the tree
std::function<int(const MCTSNode*)> getMaxDepth = [&](const MCTSNode* node) -> int {
int maxChildDepth = node->depth;
for (const auto& child : node->children) {
maxChildDepth = std::max(maxChildDepth, getMaxDepth(child.get()));
}
return maxChildDepth;
};
result.depthAchieved = getMaxDepth(rootNode.get());
// Count total nodes visited
std::function<size_t(const MCTSNode*)> countVisited =
[&](const MCTSNode* node) -> size_t {
size_t count = (node->visitCount > 0) ? 1 : 0;
for (const auto& child : node->children) { count += countVisited(child.get()); }
return count;
};
result.commandCountEvaluated = countVisited(rootNode.get());
result.availableCommandCount = commands.size();
// MCTS-specific logging
printf("MCTS: Selected command %zu (visit:%d, reward:%.2f, lookahead:%.2f) from %zu "
"options\n",
bestChild->commandIndex,
bestChild->visitCount,
bestChild->averageReward,
bestChild->lookaheadScore,
rootNode->children.size());
// Log top 3 commands for debugging with their best sequences
std::vector<MCTSNode*> sortedChildren;
for (const auto& child : rootNode->children) { sortedChildren.push_back(child.get()); }
std::ranges::sort(sortedChildren, [](const MCTSNode* a, const MCTSNode* b) {
return a->visitCount > b->visitCount;
});
printf("MCTS: Top commands by visits:\n");
for (size_t i = 0; i < std::min(static_cast<size_t>(3), sortedChildren.size()); ++i) {
auto* child = sortedChildren[i];
printf(" [%zu] cmd:%zu visits:%d immediate:%.2f backprop:%.2f type:%s",
i,
child->commandIndex,
child->visitCount,
child->immediateScore,
child->averageReward,
CommandType_Name(child->commandType).c_str());
// Show unit and target info for commands that have them
if (child->actorUnitId >= 0) { printf(" unit:%d", child->actorUnitId); }
if (child->targetRow >= 0 && child->targetCol >= 0) {
printf(" target:(%d,%d)", child->targetRow, child->targetCol);
}
// Show sequence preview for this command
if (!child->children.empty()) {
// Find best child by visits
MCTSNode* bestNext = nullptr;
int maxVisits = 0;
for (const auto& grandchild : child->children) {
if (grandchild->visitCount > maxVisits) {
maxVisits = grandchild->visitCount;
bestNext = grandchild.get();
}
}
if (bestNext) {
printf(" -> %s", CommandType_Name(bestNext->commandType).c_str());
}
}
printf("\n");
}
printf("MCTS: Tree stats - max depth:%zu, total nodes:%zu, root visits:%d\n",
result.depthAchieved,
result.commandCountEvaluated,
rootNode->visitCount);
// Log the best command sequence from the chosen command
struct SequenceNode {
size_t commandIndex;
std::string commandType;
int actorUnitId;
int targetRow;
int targetCol;
double immediateScore;
double backpropScore;
};
std::vector<SequenceNode> bestSequence;
bestSequence.reserve(5);
auto current = const_cast<MCTSNode*>(bestChild);
double sequenceScore = bestChild->averageReward;
// Trace the best path from chosen command (most visited child at each level)
while (current) {
bestSequence.push_back(
{current->commandIndex,
CommandType_Name(current->commandType),
current->actorUnitId,
current->targetRow,
current->targetCol,
current->immediateScore,
current->averageReward});
// If no children, we've reached the end of the sequence
if (current->children.empty()) { break; }
// Find most visited child
MCTSNode* bestChildNode = nullptr;
int maxVisits = 0;
for (const auto& child : current->children) {
if (child->visitCount > maxVisits) {
maxVisits = child->visitCount;
bestChildNode = child.get();
}
}
current = bestChildNode;
if (current) {
sequenceScore = current->averageReward; // Update to final score
}
}
if (!bestSequence.empty()) {
printf("MCTS: Best sequence from chosen command (final: %.2f):\n", sequenceScore);
for (size_t i = 0; i < bestSequence.size(); ++i) {
const auto& [commandIndex, commandType, actorUnitId, targetRow, targetCol, immediateScore, backpropScore] =
bestSequence[i];
printf(" %zu. cmd:%zu %s", i + 1, commandIndex, commandType.c_str());
// Add unit and target info if present
if (actorUnitId >= 0) { printf(" unit:%d", actorUnitId); }
if (targetRow >= 0 && targetCol >= 0) {
printf(" target:(%d,%d)", targetRow, targetCol);
}
printf(" (immediate:%.2f, backprop:%.2f)\n", immediateScore, backpropScore);
}
}
}
}
const auto endTime = std::chrono::steady_clock::now();
result.timeUsed = std::chrono::duration_cast<std::chrono::milliseconds>(endTime - startTime);
return result;
}
auto MCTSAI::BuildMCTSTree(
const ShardokEngine& engine,
const SettingsGetter& settingsGetter,
const CoordsSet& criticalTileCoords,
const std::chrono::steady_clock::time_point deadline) const -> std::unique_ptr<MCTSNode> {
// Clear transposition registry for this search
if (config.enableTranspositionDetection) { stateRegistry.clear(); }
// Create root node
auto root = std::make_unique<MCTSNode>(
0,
net::eagle0::shardok::common::END_TURN_COMMAND,
playerId,
0,
isDefender);
root->resultingGameState = engine.GetCurrentGameState();
// Register root node in transposition table if enabled
if (config.enableTranspositionDetection) {
root->stateHash = root->resultingGameState.ComputeFNV1aHash();
stateRegistry[root->stateHash] = root.get();
}
// Initialize root with available commands
const CommandListSPtr rootCommands = engine.GetAvailableCommandsForAIPlayer(playerId);
if (!rootCommands || rootCommands->empty()) { return root; }
// Filter commands for better performance
const std::vector<size_t> filteredIndices = AICommandFilter::FilterCommands(
rootCommands,
playerId,
isDefender,
engine.GetCurrentGameState(),
settingsGetter,
apdCache);
root->untriedCommands = filteredIndices;
root->fullyExpanded = filteredIndices.empty();
// Main MCTS loop
int iterations = 0;
printf("MCTS: Starting search with %zu filtered commands (budget: %.0fms)\n",
filteredIndices.size(),
std::chrono::duration<double, std::milli>(deadline - std::chrono::steady_clock::now())
.count());
if (config.useMultithreading) {
// Parallel MCTS: Run iterations until time budget expires
const int numThreads =
std::min(config.numThreads, static_cast<int>(std::thread::hardware_concurrency()));
std::vector<std::future<void>> futures;
std::mutex treeMutex; // Protect tree updates
std::atomic totalIterations{0}; // Track iterations across threads
for (int t = 0; t < numThreads; ++t) {
futures.push_back(std::async(std::launch::async, [&, this] {
while (std::chrono::steady_clock::now() < deadline) {
++totalIterations;
// Selection and expansion need locking
MCTSNode* selected;
{
std::lock_guard lock(treeMutex);
selected = MCTSSelection(root.get());
if (selected && !selected->isTerminal && selected->CanExpand()) {
selected = MCTSExpansion(
selected,
engine,
settingsGetter,
criticalTileCoords);
}
}
// Skip simulation if selection failed (all children redundant)
if (!selected) continue;
// Simulation can run in parallel from the selected node's state
// Note: Creating engine from node state is correct for MCTS simulation
// Backpropagation needs locking
{
ShardokEngine nodeEngine(
engine.GetGameSettings(),
selected->resultingGameState,
criticalTileCoords);
const double reward =
MCTSSimulation(nodeEngine, selected->playerId, settingsGetter);
std::lock_guard lock(treeMutex);
MCTSBackpropagation(selected, reward);
}
}
}));
}
// Wait for all threads to complete
for (auto& future : futures) { future.get(); }
iterations = totalIterations.load(); // Get total from all threads
} else {
// Sequential MCTS - run until time budget expires
while (std::chrono::steady_clock::now() < deadline) {
// Selection
MCTSNode* selected = MCTSSelection(root.get());
// Skip if selection failed (all children redundant)
if (!selected) continue;
// Expansion
if (!selected->isTerminal && selected->CanExpand()) {
selected = MCTSExpansion(selected, engine, settingsGetter, criticalTileCoords);
}
// Simulation (from selected node's state)
ShardokEngine nodeEngine(
engine.GetGameSettings(),
selected->resultingGameState,
criticalTileCoords);
double reward = MCTSSimulation(
nodeEngine,
selected->playerId,
settingsGetter); // Use node's player, not original player
// Backpropagation
MCTSBackpropagation(selected, reward);
iterations++;
}
}
printf("MCTS: Completed %d iterations, root has %zu children\n",
iterations,
root->children.size());
return root;
}
auto MCTSAI::MCTSSelection(MCTSNode* root) const -> MCTSNode* {
MCTSNode* current = root;
while (!current->isTerminal && !current->isRedundant) {
if (current->CanExpand()) {
return current; // Node has untried commands
} else if (!current->children.empty()) {
current = current->GetBestChild(config.explorationConstant);
if (!current) break;
} else {
break; // Leaf node
}
}
return current;
}
auto MCTSAI::MCTSExpansion(
MCTSNode* node,
const ShardokEngine& engine,
const SettingsGetter& settingsGetter,
const CoordsSet& criticalTileCoords) const -> MCTSNode* {
static int expansionCallCount = 0;
if (expansionCallCount < 3) {
printf("MCTSExpansion called %d: node depth:%d untried:%zu\n",
expansionCallCount++,
node->depth,
node->untriedCommands.size());
}
if (node->untriedCommands.empty()) return node;
// Don't expand beyond maximum depth to prevent unbounded tree growth
if (node->depth >= config.maxTreeDepth) {
node->fullyExpanded = true;
node->untriedCommands.clear();
return node;
}
// Pick a random untried command
std::random_device rd;
std::mt19937 gen(rd());
std::uniform_int_distribution dis(0, static_cast<int>(node->untriedCommands.size() - 1));
const size_t randomIndex = dis(gen);
const auto commandIndex = node->untriedCommands[randomIndex];
node->untriedCommands.erase(node->untriedCommands.begin() + randomIndex);
// Create engine from the node's current state (not root state!)
const auto nodeEngine = std::make_shared<ShardokEngine>(
engine.GetGameSettings(),
node->resultingGameState,
criticalTileCoords);
// Get command descriptor from the node's state
const CommandListSPtr commands = nodeEngine->GetAvailableCommandsForAIPlayer(node->playerId);
if (!commands || commandIndex >= commands->size()) return node;
const auto& command = commands->at(commandIndex);
const auto commandType = command->GetCommandType();
const auto descriptor = command->GetCommandProto();
// Create child node
auto child = std::make_unique<MCTSNode>(
commandIndex,
commandType,
node->playerId,
node->depth + 1,
node->isDefender);
child->parent = node;
// Extract actor unit ID if present
if (descriptor.has_actor()) { child->actorUnitId = descriptor.actor().value(); }
// Extract target coordinates if present
// Note: In protobuf3, target is always present but may have default values
// We'll always capture the coordinates - commands without targets will have (-1,-1) by default
const auto& target = descriptor.target();
child->targetRow = target.row();
child->targetCol = target.column();
// Declare variables that will be used later
// Handle randomness appropriately based on command type
if (command->HasOdds()) {
// For commands with odds, use average roll for expansion
// For expansion, use average roll regardless of success chance
const auto generator = std::make_shared<SequenceRandomGenerator>(std::vector{0.5});
nodeEngine->PostCommand(node->playerId, commandIndex, generator);
} else {
// Use average generator for deterministic evaluation
nodeEngine->PostCommand(node->playerId, commandIndex, _averageGenerator);
}
child->resultingGameState = nodeEngine->GetCurrentGameState();
// Check whose turn it is after executing the command
PlayerId currentPlayer = nodeEngine->GetCurrentPlayerId();
bool isOurTurn = (currentPlayer == playerId);
// Update child's player ID to reflect whose turn it actually is
child->playerId = currentPlayer;
// Calculate immediate score (always from our perspective)
child->immediateScore = AIScoreCalculator::GuessedStateScore(
isDefender, // Use our original role, not node's
child->resultingGameState,
strategy,
castleCoords,
settingsGetter,
apdCache,
alCache);
// Initially, lookahead score equals immediate score
child->lookaheadScore = child->immediateScore;
// Debug: Log first few expansions to see what's happening
static int expansionCount = 0;
if (expansionCount < 5) {
printf("MCTS Expansion %d: cmd:%zu type:%s immediate_score:%.2f\n",
expansionCount++,
commandIndex,
CommandType_Name(commandType).c_str(),
child->immediateScore);
}
// Check if terminal
child->isTerminal = IsTerminalForPlayer(child->resultingGameState, playerId, settingsGetter);
// Transposition detection
if (config.enableTranspositionDetection) {
child->stateHash = child->resultingGameState.ComputeFNV1aHash();
auto existingIt = stateRegistry.find(child->stateHash);
if (existingIt != stateRegistry.end()) {
MCTSNode* existingNode = existingIt->second;
// Apply tie-breaking rules to determine which node to keep
bool shouldPruneChild = false;
if (child->depth > existingNode->depth) {
// Rule 1: Prune deeper node (current child is deeper)
shouldPruneChild = true;
} else if (child->depth == existingNode->depth) {
// Rule 2: At same depth, prune node with higher command index
if (child->commandIndex > existingNode->commandIndex) {
shouldPruneChild = true;
} else {
// Current child wins - mark existing node as redundant
existingNode->isRedundant = true;
stateRegistry[child->stateHash] = child.get(); // Update registry
}
} else {
// Child is shallower - mark existing node as redundant
existingNode->isRedundant = true;
stateRegistry[child->stateHash] = child.get(); // Update registry
}
if (shouldPruneChild) {
child->isRedundant = true;
// Don't expand redundant nodes
}
} else {
// New state - register it
stateRegistry[child->stateHash] = child.get();
}
}
// Get available commands for child - only if it's still our turn and not redundant
if (!child->isTerminal && !child->isRedundant && isOurTurn) {
const CommandListSPtr childCommands =
nodeEngine->GetAvailableCommandsForAIPlayer(currentPlayer);
if (childCommands) {
const std::vector<size_t> childFiltered = AICommandFilter::FilterCommands(
childCommands,
currentPlayer,
isDefender, // Use our original role
child->resultingGameState,
settingsGetter,
apdCache);
child->untriedCommands = childFiltered;
child->fullyExpanded = childFiltered.empty();
}
} else if (!isOurTurn) {
// Mark as terminal if it's not our turn - we can't expand opponent moves
child->isTerminal = true;
child->fullyExpanded = true;
}
// Update parent's expansion status
if (node->untriedCommands.empty()) { node->fullyExpanded = true; }
MCTSNode* childPtr = child.get();
node->children.push_back(std::move(child));
return childPtr;
}
auto MCTSAI::MCTSSimulation(
const ShardokEngine& engineState,
PlayerId startingPlayer,
const SettingsGetter& settingsGetter) const -> double {
// Create copy for simulation
auto simEngine = std::make_shared<ShardokEngine>(engineState, false);
// Always evaluate from our AI's perspective (not the startingPlayer's perspective)
const double initialScore = AIScoreCalculator::GuessedStateScore(
isDefender,
simEngine->GetCurrentGameState(),
strategy,
castleCoords,
settingsGetter,
apdCache,
alCache);
// Debug: Note that we're always scoring from our AI's perspective regardless of startingPlayer
(void)startingPlayer; // Acknowledge parameter to avoid warning
// Fast rollout with random/heuristic moves until terminal
int simulationSteps = 0;
for (int step = 0; step < config.maxSimulationDepth; ++step) {
const GameStateW& currentState = simEngine->GetCurrentGameState();
// Get whose turn it is
const PlayerId currentPlayer = simEngine->GetCurrentPlayerId();
// Check if terminal
if (IsTerminalForPlayer(currentState, playerId, settingsGetter)) { break; }
// Continue as long as it's still our turn (don't stop after each individual command)
// In this game, a player can move multiple units before turn switches
if (currentPlayer != playerId) {
// Turn switched to opponent - stop simulation immediately
break;
}
simulationSteps++;
// Get available commands for current player
const CommandListSPtr commands = simEngine->GetAvailableCommandsForAIPlayer(currentPlayer);
if (!commands || commands->empty()) break;
// Select command based on simulation policy
const auto commandIndex = static_cast<int>(
SelectSimulationCommand(commands, currentPlayer, simEngine, settingsGetter));
// Execute command
simEngine->PostCommand(currentPlayer, commandIndex, _averageGenerator);
// Debug: Only log if turn changed unexpectedly
const PlayerId newPlayer = simEngine->GetCurrentPlayerId();
static int debugCount = 0;
if (newPlayer != currentPlayer && debugCount < 5) {
const auto commandType = commands->at(commandIndex)->GetCommandType();
printf("MCTS Sim step %d: cmd_type:%s player_before:%d player_after:%d\n",
simulationSteps,
CommandType_Name(commandType).c_str(),
currentPlayer,
newPlayer);
printf(" WARNING: Turn changed after command!\n");
debugCount++;
}
}
// Evaluate final position
const double finalScore = AIScoreCalculator::GuessedStateScore(
isDefender,
simEngine->GetCurrentGameState(),
strategy,
castleCoords,
settingsGetter,
apdCache,
alCache);
// Debug: Log first few simulations to see depth and score change
static int simCount = 0;
if (simCount < 3) {
printf("MCTS Simulation %d: steps:%d initial:%.2f final:%.2f delta:%.2f\n",
simCount++,
simulationSteps,
initialScore,
finalScore,
finalScore - initialScore);
}
return finalScore;
}
auto MCTSAI::MCTSBackpropagation(MCTSNode* node, double reward) -> void {
while (node) {
node->visitCount++;
node->totalReward += reward;
node->averageReward = node->totalReward / node->visitCount;
// Update lookahead score as weighted average
if (node->visitCount == 1) {
node->lookaheadScore = reward;
} else {
node->lookaheadScore =
(node->lookaheadScore * (node->visitCount - 1) + reward) / node->visitCount;
}
node = node->parent;
}
}
auto MCTSAI::IsTerminalForPlayer(
const GameStateW& gameState,
PlayerId /*currentPlayer*/,
const SettingsGetter& settingsGetter) -> bool {
// Check if game is over
if (gameState->status()->state() ==
net::eagle0::shardok::storage::fb::GameStatus_::State_VICTORY ||
gameState->status()->state() ==
net::eagle0::shardok::storage::fb::GameStatus_::State_DRAW) {
return true;
}
// Check max rounds
if (gameState->current_round() >= settingsGetter.Backing().max_rounds()) { return true; }
return false;
}
auto MCTSAI::SelectSimulationCommand(
const CommandListSPtr& commands,
PlayerId currentPlayer,
const std::shared_ptr<ShardokEngine>& simEngine,
const SettingsGetter& settingsGetter) const -> size_t {
if (commands->size() == 1) {
return 0; // Only one choice
}
std::random_device rd;
std::mt19937 gen(rd());
switch (config.simulationPolicy) {
case MCTSSimulationPolicy::RANDOM: {
// Pure random selection
std::uniform_int_distribution<> dis(0, commands->size() - 1);
return dis(gen);
}
case MCTSSimulationPolicy::FILTERED_RANDOM: {
// Filter commands first, then random selection
const auto filteredIndices = AICommandFilter::FilterCommands(
commands,
currentPlayer,
isDefender,
simEngine->GetCurrentGameState(),
settingsGetter,
apdCache);
if (filteredIndices.empty()) {
// Fallback to random if no commands pass filter
std::uniform_int_distribution<> dis(0, commands->size() - 1);
return dis(gen);
}
std::uniform_int_distribution<> dis(0, filteredIndices.size() - 1);
return filteredIndices[dis(gen)];
}
case MCTSSimulationPolicy::BEST_IMMEDIATE: {
// Evaluate all commands and pick the best
double bestScore = -std::numeric_limits<double>::max();
size_t bestIndex = 0;
for (size_t i = 0; i < commands->size(); ++i) {
// Create a temporary engine to evaluate this command
const auto testEngine = std::make_shared<ShardokEngine>(*simEngine, false);
// Get score BEFORE executing command (for potential player flip comparison)
const double preScore = AIScoreCalculator::GuessedStateScore(
isDefender,
testEngine->GetCurrentGameState(),
strategy,
castleCoords,
settingsGetter,
apdCache,
alCache);
const PlayerId playerBefore = testEngine->GetCurrentPlayerId();
testEngine->PostCommand(currentPlayer, i, _averageGenerator);
const PlayerId playerAfter = testEngine->GetCurrentPlayerId();
double score;
if (playerAfter != playerBefore) {
// Player flipped - use pre-execution score to avoid opponent turn effects
score = preScore;
} else {
// Normal command - use post-execution score
score = AIScoreCalculator::GuessedStateScore(
isDefender,
testEngine->GetCurrentGameState(),
strategy,
castleCoords,
settingsGetter,
apdCache,
alCache);
}
if (score > bestScore) {
bestScore = score;
bestIndex = i;
}
}
return bestIndex;
}
case MCTSSimulationPolicy::WEIGHTED_BEST_IMMEDIATE: {
// Evaluate all commands and weight by ranking
struct CommandScore {
size_t index;
double score;
};
std::vector<CommandScore> commandScores;
commandScores.reserve(commands->size());
for (size_t i = 0; i < commands->size(); ++i) {
// Create a temporary engine to evaluate this command
const auto testEngine = std::make_shared<ShardokEngine>(*simEngine, false);
// Get score BEFORE executing command (for potential player flip comparison)
const double preScore = AIScoreCalculator::GuessedStateScore(
isDefender,
testEngine->GetCurrentGameState(),
strategy,
castleCoords,
settingsGetter,
apdCache,
alCache);
const PlayerId playerBefore = testEngine->GetCurrentPlayerId();
testEngine->PostCommand(currentPlayer, i, _averageGenerator);
const PlayerId playerAfter = testEngine->GetCurrentPlayerId();
double score;
if (playerAfter != playerBefore) {
// Player flipped - use pre-execution score to avoid opponent turn effects
score = preScore;
} else {
// Normal command - use post-execution score
score = AIScoreCalculator::GuessedStateScore(
isDefender,
testEngine->GetCurrentGameState(),
strategy,
castleCoords,
settingsGetter,
apdCache,
alCache);
}
commandScores.push_back({i, score});
}
// Sort by score (best first)
std::sort(
commandScores.begin(),
commandScores.end(),
[](const CommandScore& a, const CommandScore& b) { return a.score > b.score; });
// Assign weights: 1.0 for best, 0.5 for second, 0.33 for third, etc.
std::vector<double> weights;
weights.reserve(commandScores.size());
double totalWeight = 0.0;
for (size_t i = 0; i < commandScores.size(); ++i) {
double weight = 1.0 / (i + 1); // 1/1, 1/2, 1/3, ...
weights.push_back(weight);
totalWeight += weight;
}
// Random selection based on weights
std::uniform_real_distribution dis(0.0, totalWeight);
const double target = dis(gen);
double cumulative = 0.0;
for (size_t i = 0; i < weights.size(); ++i) {
cumulative += weights[i];
if (cumulative >= target) { return commandScores[i].index; }
}
// Fallback (shouldn't happen)
return commandScores[0].index;
}
}
// Fallback to random (shouldn't reach here)
std::uniform_int_distribution dis(0, static_cast<int>(commands->size() - 1));
return dis(gen);
}
} // namespace shardok
@@ -0,0 +1,128 @@
//
// MCTS-based AI system for Shardok
// Alternative to IterativeDeepeningAI using Monte Carlo Tree Search
//
#ifndef EAGLE0_MCTSAI_HPP
#define EAGLE0_MCTSAI_HPP
#include <chrono>
#include <future>
#include <memory>
#include <unordered_map>
#include <vector>
#include "src/main/cpp/net/eagle0/shardok/ai/AIStrategy.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AITimeBudget.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/IterativeDeepeningAI.hpp" // For SearchResult compatibility
#include "src/main/cpp/net/eagle0/shardok/library/GameStateW.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokCTypes.h"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokCommand.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistancesCache.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
#include "src/main/protobuf/net/eagle0/shardok/api/command_descriptor.pb.h"
namespace shardok {
// Forward declarations
class ShardokEngine;
// MCTSNode is defined in internal/MCTSNode.hpp
namespace internal {
struct MCTSNode;
}
// Simulation policy for MCTS rollouts
enum class MCTSSimulationPolicy {
RANDOM, // Pure random selection
FILTERED_RANDOM, // Random from filtered commands
BEST_IMMEDIATE, // Choose best immediate score
WEIGHTED_BEST_IMMEDIATE // Random weighted by score ranking
};
// Configuration for MCTS algorithm
struct MCTSConfig {
double explorationConstant = 1.414; // UCB1 constant (sqrt(2) by default)
int maxSimulationDepth = 1000; // Maximum depth for rollout
int maxTreeDepth = 2000; // Maximum tree depth to prevent stack overflow
bool useMultithreading = true; // Enable parallel MCTS
int numThreads = 16; // Number of threads for parallel MCTS (when enabled)
MCTSSimulationPolicy simulationPolicy = MCTSSimulationPolicy::BEST_IMMEDIATE;
bool enableTranspositionDetection = true; // Enable pruning of duplicate states
};
class MCTSAI {
public:
using CommandProto = net::eagle0::shardok::api::CommandDescriptor;
using SearchResult = IterativeDeepeningAI::SearchResult;
MCTSAI(PlayerId playerId,
bool isDefender,
AIStrategy strategy,
const CoordsSet& castleCoords,
const APDCache& apdCache,
const ALCache& alCache,
MCTSConfig config = MCTSConfig{});
// Main search interface - compatible with IterativeDeepeningAI
[[nodiscard]] auto Search(
const GameSettingsSPtr& settings,
const GameStateW& state,
const std::vector<CommandProto>& commands,
const AITimeBudget& budget) const -> SearchResult;
// Get/set configuration
[[nodiscard]] auto GetConfig() const -> const MCTSConfig& { return config; }
void SetConfig(const MCTSConfig& newConfig) { config = newConfig; }
private:
PlayerId playerId;
bool isDefender;
AIStrategy strategy;
const CoordsSet& castleCoords;
const APDCache& apdCache;
const ALCache& alCache;
MCTSConfig config;
// Transposition detection infrastructure
mutable std::unordered_map<uint64_t, internal::MCTSNode*>
stateRegistry; // Hash -> first node mapping
// Internal MCTS tree building
[[nodiscard]] auto BuildMCTSTree(
const ShardokEngine& engine,
const SettingsGetter& settingsGetter,
const CoordsSet& criticalTileCoords,
std::chrono::steady_clock::time_point deadline) const
-> std::unique_ptr<internal::MCTSNode>;
// MCTS algorithm phases
auto MCTSSelection(internal::MCTSNode* root) const -> internal::MCTSNode*;
auto MCTSExpansion(
internal::MCTSNode* node,
const ShardokEngine& engine,
const SettingsGetter& settingsGetter,
const CoordsSet& criticalTileCoords) const -> internal::MCTSNode*;
auto MCTSSimulation(
const ShardokEngine& engineState,
PlayerId currentPlayer,
const SettingsGetter& settingsGetter) const -> double;
static auto MCTSBackpropagation(internal::MCTSNode* node, double reward) -> void;
// Helper functions
[[nodiscard]] static auto IsTerminalForPlayer(
const GameStateW& gameState,
PlayerId currentPlayer,
const SettingsGetter& settingsGetter) -> bool;
// Simulation command selection based on policy
[[nodiscard]] auto SelectSimulationCommand(
const CommandListSPtr& commands,
PlayerId currentPlayer,
const std::shared_ptr<ShardokEngine>& simEngine,
const SettingsGetter& settingsGetter) const -> size_t;
};
} // namespace shardok
#endif // EAGLE0_MCTSAI_HPP
@@ -1,111 +0,0 @@
//
// Shardok-specific MCTS AI implementation using abstract interfaces
//
#include "ShardokMCTSAI.hpp"
#include "adapters/ShardokGameEngine.hpp"
#include "adapters/ShardokGameState.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AICommandFilter.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/score/AIScoreCalculator.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokEngine.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/util/HexMapUtils.hpp"
namespace shardok {
ShardokMCTSAI::ShardokMCTSAI(
PlayerId playerId,
bool isDefender,
AIStrategy strategy,
const CoordsSet& castleCoords,
const AIScoreCalculator& scoreCalculator,
const APDCache& apdCache,
const ALCache& alCache,
MCTSConfig config)
: abstractAI_(std::make_unique<mcts::AbstractMCTSAI>(
static_cast<mcts::MCTSPlayerId>(
playerId), // Use actual player ID for correct scoring
config)),
isDefender_(isDefender),
strategy_(strategy),
castleCoords_(castleCoords),
scoreCalculator_(scoreCalculator),
apdCache_(apdCache),
alCache_(alCache) {}
auto ShardokMCTSAI::Search(
const GameSettingsSPtr& settings,
const GameStateW& state,
const AITimeBudget& budget) const -> SearchResult {
// Compute critical tiles once to avoid 8.5% runtime overhead in ShardokEngine construction
const auto criticalTiles = GetCriticalTileLocations(state->hex_map());
// Create Shardok engine for simulation
ShardokEngine engine(settings, state, criticalTiles, 0, false);
// Create game state adapter
auto gameState = mcts::ShardokMCTSFactory::createGameState(
state,
&scoreCalculator_, // Pass the score calculator
settings, // Pass shared_ptr directly
isDefender_,
strategy_,
castleCoords_,
apdCache_,
alCache_,
criticalTiles);
// Create game engine adapter (passing critical tiles to avoid recomputation)
auto gameEngine = mcts::ShardokMCTSFactory::createGameEngine(
engine,
&scoreCalculator_, // Pass the score calculator
settings,
apdCache_,
alCache_,
isDefender_,
strategy_,
castleCoords_,
criticalTiles);
// Perform abstract search
const auto timeLimit = budget.remainingBudget;
const auto abstractResult = abstractAI_->Search(*gameEngine, *gameState, timeLimit);
// Report cache statistics for performance analysis
if (auto* shardokEngine = dynamic_cast<mcts::ShardokGameEngine*>(gameEngine.get())) {
shardokEngine->reportCacheStatistics();
}
// Get unfiltered command count for consistent reporting with IterativeDeepeningAI
// (MCTS uses filtered commands internally, but we report unfiltered count for metrics)
const auto unfilteredCommands = engine.GetAvailableCommandsForAIPlayer(
static_cast<PlayerId>(gameState->currentPlayerId()));
const size_t unfilteredCount = unfilteredCommands ? unfilteredCommands->size() : 0;
// Convert result back to Shardok format
SearchResult result;
// Map filtered index back to original unfiltered index
result.bestCommandIndex =
gameEngine->mapFilteredIndexToOriginal(abstractResult.bestActionIndex, *gameState);
result.bestScore = abstractResult.bestScore;
result.depthAchieved = static_cast<size_t>(abstractResult.searchDepth);
result.commandCountEvaluated = static_cast<size_t>(abstractResult.nodesEvaluated);
result.timeUsed = abstractResult.searchTime;
result.availableCommandCount = unfilteredCount;
result.minimumDepthCompleted =
(abstractResult.searchDepth >= static_cast<int>(budget.minDepthRequired));
result.searchCompleted = true; // MCTS is anytime - always returns a valid result
// Determine completion reason based on what actually happened
if (abstractResult.foundWinningMove || unfilteredCount == 0) {
// Found a terminal winning state or no commands available
result.completionReason = EvaluationCompletionReason::RAN_OUT_OF_COMMANDS;
} else {
// Normal case - time budget exhausted while exploring
result.completionReason = EvaluationCompletionReason::RAN_OUT_OF_TIME;
}
return result;
}
} // namespace shardok
@@ -1,71 +0,0 @@
//
// Shardok-specific MCTS AI that wraps the abstract implementation
//
#ifndef EAGLE0_SHARDOK_MCTSAI_HPP
#define EAGLE0_SHARDOK_MCTSAI_HPP
#include <memory>
#include <vector>
#include "adapters/ShardokMCTSFactory.hpp"
#include "src/main/cpp/net/eagle0/common/mcts/abstract/AbstractMCTSAI.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIStrategy.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AITimeBudget.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/IterativeDeepeningAI.hpp" // For SearchResult compatibility
#include "src/main/cpp/net/eagle0/shardok/library/GameStateW.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokCTypes.h"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistancesCache.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
#pragma clang diagnostic push
#pragma clang diagnostic ignored "-Wdeprecated-redundant-constexpr-static-def"
#pragma clang diagnostic pop
namespace shardok {
// Forward declarations
class ShardokEngine;
class AICommandFilter;
class AIScoreCalculator;
class ShardokMCTSAI {
public:
using SearchResult = IterativeDeepeningAI::SearchResult;
using MCTSConfig = mcts::MCTSConfig;
ShardokMCTSAI(
PlayerId playerId,
bool isDefender,
AIStrategy strategy,
const CoordsSet& castleCoords,
const AIScoreCalculator& scoreCalculator,
const APDCache& apdCache,
const ALCache& alCache,
MCTSConfig config = MCTSConfig{});
// Main search interface - compatible with IterativeDeepeningAI
[[nodiscard]] auto Search(
const GameSettingsSPtr& settings,
const GameStateW& state,
const AITimeBudget& budget) const -> SearchResult;
// Configuration
[[nodiscard]] auto GetConfig() const -> const MCTSConfig& { return abstractAI_->GetConfig(); }
void SetConfig(const MCTSConfig& newConfig) { abstractAI_->SetConfig(newConfig); }
private:
std::unique_ptr<mcts::AbstractMCTSAI> abstractAI_;
// Shardok-specific context
bool isDefender_;
AIStrategy strategy_;
const CoordsSet& castleCoords_;
const AIScoreCalculator& scoreCalculator_;
const APDCache& apdCache_;
const ALCache& alCache_;
};
} // namespace shardok
#endif // EAGLE0_SHARDOK_MCTSAI_HPP
@@ -1,82 +0,0 @@
load("//tools:copts.bzl", "COPTS")
cc_library(
name = "shardok_action",
srcs = ["ShardokAction.cpp"],
hdrs = ["ShardokAction.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai/mcts:__pkg__",
"//src/test/cpp/net/eagle0/common/mcts:__subpackages__",
],
deps = [
"//src/main/cpp/net/eagle0/common/mcts/abstract:mcts_action",
"//src/main/cpp/net/eagle0/shardok/library:shardok_c_types",
"//src/main/protobuf/net/eagle0/shardok/common:command_type_cc_proto",
],
)
cc_library(
name = "shardok_game_state",
srcs = ["ShardokGameState.cpp"],
hdrs = ["ShardokGameState.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai/mcts:__pkg__",
"//src/test/cpp/net/eagle0/common/mcts:__subpackages__",
],
deps = [
"//src/main/cpp/net/eagle0/common/mcts/abstract:mcts_game_state",
"//src/main/cpp/net/eagle0/shardok/ai:ai_strategy",
"//src/main/cpp/net/eagle0/shardok/ai/score:ai_score_calculator_interface",
"//src/main/cpp/net/eagle0/shardok/library:engine",
"//src/main/cpp/net/eagle0/shardok/library:game_state_w",
"//src/main/cpp/net/eagle0/shardok/library/settings:game_settings",
"//src/main/protobuf/net/eagle0/shardok/common:command_type_cc_proto",
],
)
cc_library(
name = "shardok_game_engine",
srcs = ["ShardokGameEngine.cpp"],
hdrs = ["ShardokGameEngine.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai/mcts:__pkg__",
"//src/test/cpp/net/eagle0/common/mcts:__subpackages__",
],
deps = [
":shardok_action",
":shardok_game_state",
"//src/main/cpp/net/eagle0/common:sequence_random_generator",
"//src/main/cpp/net/eagle0/common/mcts/abstract:mcts_game_engine",
"//src/main/cpp/net/eagle0/shardok/ai:ai_command_filter",
"//src/main/cpp/net/eagle0/shardok/ai:ai_heuristic_weighting",
"//src/main/cpp/net/eagle0/shardok/ai/score:ai_score_calculator_interface",
"//src/main/cpp/net/eagle0/shardok/library:engine",
"//src/main/cpp/net/eagle0/shardok/library:shardok_c_types",
"//src/main/cpp/net/eagle0/shardok/library/settings:game_settings",
],
)
cc_library(
name = "shardok_mcts_factory",
srcs = ["ShardokMCTSFactory.cpp"],
hdrs = ["ShardokMCTSFactory.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai:__pkg__",
"//src/main/cpp/net/eagle0/shardok/ai/mcts:__pkg__",
"//src/test/cpp/net/eagle0/common/mcts:__subpackages__",
],
deps = [
":shardok_action",
":shardok_game_engine",
":shardok_game_state",
"//src/main/cpp/net/eagle0/shardok/ai:ai_command_filter",
"//src/main/cpp/net/eagle0/shardok/ai/score:ai_score_calculator_interface",
"//src/main/cpp/net/eagle0/shardok/library:engine",
"//src/main/cpp/net/eagle0/shardok/library:game_state_w",
"//src/main/cpp/net/eagle0/shardok/library/settings:game_settings",
],
)
@@ -1,88 +0,0 @@
//
// Shardok-specific action adapter implementation
//
#include "ShardokAction.hpp"
#include <sstream>
#pragma clang diagnostic push
#pragma clang diagnostic ignored "-Wdeprecated-redundant-constexpr-static-def"
#include "src/main/protobuf/net/eagle0/shardok/common/command_type.pb.h"
#pragma clang diagnostic pop
namespace shardok::mcts {
// Constructor: extract and store just the essential fields
ShardokAction::ShardokAction(
size_t index,
CommandType type,
PlayerId player,
int actorId,
int targetRow,
int targetCol,
bool hasOdds)
: commandIndex_(index),
type_(type),
player_(player),
actorId_(actorId),
targetRow_(targetRow),
targetCol_(targetCol),
hasOdds_(hasOdds) {}
std::string ShardokAction::getDescription() const {
std::stringstream ss;
// Show player
ss << "P" << static_cast<int>(player_) << " ";
ss << net::eagle0::shardok::common::CommandType_Name(type_);
if (actorId_ >= 0) { ss << " Unit:" << actorId_; }
if (targetRow_ >= 0 && targetCol_ >= 0) {
ss << " @(" << targetRow_ << "," << targetCol_ << ")";
}
return ss.str();
}
std::unique_ptr<MCTSAction> ShardokAction::clone() const {
return std::make_unique<ShardokAction>(
commandIndex_,
type_,
player_,
actorId_,
targetRow_,
targetCol_,
hasOdds_);
}
bool ShardokAction::equals(const MCTSAction& other) const {
const auto* shardokOther = dynamic_cast<const ShardokAction*>(&other);
if (!shardokOther) { return false; }
// Compare by index only - actions from same command list are uniquely identified by index
return commandIndex_ == shardokOther->commandIndex_;
}
bool ShardokAction::requiresChanceNode() const {
// Actions with probabilistic outcomes require chance nodes:
// 1. Binary success/failure actions (hasOdds_): START_FIRE, FEAR, etc.
// 2. END_TURN: random effects (fire spread, weather changes)
// 3. Combat actions: roll affects damage dealt (MELEE, ARCHERY, CHARGE, DUEL)
if (hasOdds_) { return true; }
using namespace net::eagle0::shardok::common;
switch (type_) {
case END_TURN_COMMAND:
case MELEE_COMMAND:
case ARCHERY_COMMAND:
case CHARGE_COMMAND:
case CHALLENGE_DUEL_COMMAND:
case REDUCE_COMMAND: return true;
default: return false;
}
}
} // namespace shardok::mcts
@@ -1,61 +0,0 @@
//
// Shardok-specific action adapter for MCTS
//
#ifndef EAGLE0_SHARDOK_ACTION_HPP
#define EAGLE0_SHARDOK_ACTION_HPP
#include <memory>
#include <string>
#include "src/main/cpp/net/eagle0/common/mcts/abstract/MCTSAction.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokCTypes.h"
#pragma clang diagnostic push
#pragma clang diagnostic ignored "-Wdeprecated-redundant-constexpr-static-def"
#include "src/main/protobuf/net/eagle0/shardok/common/command_type.pb.h"
#pragma clang diagnostic pop
namespace shardok::mcts {
class ShardokAction : public MCTSAction {
public:
using CommandType = net::eagle0::shardok::common::CommandType;
// Constructor: store just the essential fields (no proto, no pointer)
ShardokAction(
size_t index,
CommandType type,
PlayerId player,
int actorId,
int targetRow,
int targetCol,
bool hasOdds);
// MCTSAction interface implementation
[[nodiscard]] size_t getIndex() const override { return commandIndex_; }
[[nodiscard]] std::string getDescription() const override;
[[nodiscard]] std::unique_ptr<MCTSAction> clone() const override;
[[nodiscard]] bool equals(const MCTSAction& other) const override;
[[nodiscard]] bool requiresChanceNode() const override;
// Shardok-specific accessors (O(1), no allocations)
[[nodiscard]] int getType() const { return static_cast<int>(type_); }
[[nodiscard]] PlayerId getPlayer() const { return player_; }
[[nodiscard]] int getActorId() const { return actorId_; }
[[nodiscard]] std::pair<int, int> getTarget() const { return {targetRow_, targetCol_}; }
private:
// Store only essential fields (~25 bytes, all POD, cache-friendly)
size_t commandIndex_;
CommandType type_;
PlayerId player_;
int actorId_; // -1 if no actor
int targetRow_; // -1 if no target
int targetCol_; // -1 if no target
bool hasOdds_; // true if command has probabilistic outcome
};
} // namespace shardok::mcts
#endif // EAGLE0_SHARDOK_ACTION_HPP
@@ -1,628 +0,0 @@
//
// Shardok-specific game engine adapter implementation
//
#include "ShardokGameEngine.hpp"
#include <algorithm>
#include <chrono>
#include <numeric>
#include "ShardokAction.hpp"
#include "ShardokGameState.hpp"
#include "src/main/cpp/net/eagle0/common/SequenceRandomGenerator.hpp"
#include "src/main/cpp/net/eagle0/common/mcts/abstract/MCTSTypes.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AICommandFilter.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIHeuristicWeighting.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/score/AIScoreCalculator.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokEngine.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokException.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
namespace shardok::mcts {
// Shared cache for legal actions (uses lock-free parallel hash map for thread safety)
// Using 8 submaps to reduce contention with 16 MCTS threads
gtl::parallel_flat_hash_map<
uint64_t,
ShardokGameEngine::LegalActionsCache,
std::hash<uint64_t>,
std::equal_to<uint64_t>,
std::allocator<std::pair<const uint64_t, ShardokGameEngine::LegalActionsCache>>,
8,
std::mutex>
ShardokGameEngine::legalActionsCache_;
std::atomic<uint64_t> ShardokGameEngine::cacheHits_{0};
std::atomic<uint64_t> ShardokGameEngine::cacheMisses_{0};
std::atomic<uint64_t> ShardokGameEngine::timeInHashComputation_{0};
std::atomic<uint64_t> ShardokGameEngine::timeInLegalActionsComputation_{0};
ShardokGameEngine::ShardokGameEngine(
[[maybe_unused]] const ShardokEngine* engine,
const AIScoreCalculator* scoreCalculator,
const GameSettingsSPtr& gameSettings,
const APDCache* apdCache,
const ALCache* alCache,
bool isDefender,
const AIStrategy& strategy,
const CoordsSet& castleCoords,
const CoordsSet& criticalTileCoords)
: scoreCalculator_(scoreCalculator),
gameSettings_(gameSettings),
apdCache_(apdCache),
alCache_(alCache),
isDefender_(isDefender),
strategy_(strategy),
castleCoords_(castleCoords),
criticalTileCoords_(criticalTileCoords) {
// Thread-local cache is automatically initialized per thread
// Reserve space to reduce rehashing (based on profiling: ~30-50K unique states per search)
legalActionsCache_.reserve(100000);
}
std::unique_ptr<MCTSGameState> ShardokGameEngine::applyAction(
const MCTSGameState& state,
const MCTSAction& action,
double deterministicRoll) const {
const auto* shardokState = dynamic_cast<const ShardokGameState*>(&state);
const auto* shardokAction = dynamic_cast<const ShardokAction*>(&action);
if (!shardokState || !shardokAction) { return nullptr; }
const auto currentPlayer = static_cast<PlayerId>(state.currentPlayerId());
// Use cached engine if available (avoids recomputing GetAvailableCommands for same state)
std::shared_ptr<ShardokEngine> engine;
if (auto cachedEngine = shardokState->getCachedEngine()) {
// Clone the cached engine to preserve command cache
engine = std::make_shared<ShardokEngine>(*cachedEngine);
} else {
// Create fresh engine and populate command cache
engine = std::make_shared<ShardokEngine>(
gameSettings_,
shardokState->getShardokState(),
criticalTileCoords_,
0,
false);
// Populate command cache (result intentionally unused, just populating cache)
[[maybe_unused]] const auto commands =
engine->GetAvailableCommandsForAIPlayer(currentPlayer);
// Cache the engine for future use with this state
shardokState->setCachedEngine(engine);
// Clone it for applying the action (don't mutate the cached engine)
engine = std::make_shared<ShardokEngine>(*engine);
}
// Create deterministic random generator if a specific roll is requested
// deterministicRoll of -1.0 (default) means use random generator
// Any other value (including negative) creates a deterministic generator
// For open-ended percentile commands, we compute a sequence of values that will
// produce the desired final result through the normal open-ended mechanics
std::shared_ptr<::RandomGenerator> randomGen = nullptr;
constexpr double kNoRollSentinel = -1.0;
if (deterministicRoll != kNoRollSentinel) {
std::vector<double> sequence;
if (deterministicRoll >= 5.0 && deterministicRoll <= 95.0) {
// Normal range: single value works directly
sequence = {deterministicRoll / 100.0};
} else if (deterministicRoll < 5.0) {
// Need open-ended LOW result (e.g., -100 for guaranteed success)
// OpenEndedPercentile: if initial < 5, returns initial - OpenEndedHighImpl(0, 4)
// We want: initial - accumulated = deterministicRoll
// Use initial = 2 (clearly < 5), so accumulated = 2 - deterministicRoll
constexpr double kInitialLow = 2.0;
sequence = {kInitialLow / 100.0};
// OpenEndedHighImpl accumulates rolls until one < 95
// Split accumulated into rolls: 96 (continues) + remaining (stops)
double remaining = kInitialLow - deterministicRoll;
while (remaining > 95.0) {
sequence.push_back(0.96); // 96 > 95, continues accumulation
remaining -= 96.0;
}
sequence.push_back(remaining / 100.0); // Final roll < 95, stops
} else {
// Need open-ended HIGH result (e.g., 150 for guaranteed failure)
// OpenEndedPercentile: if initial > 95, returns OpenEndedHighImpl(initial, 4)
// OpenEndedHighImpl accumulates rolls until one < 95
constexpr double kInitialHigh = 96.0;
sequence = {kInitialHigh / 100.0};
double remaining = deterministicRoll - kInitialHigh;
while (remaining > 95.0) {
sequence.push_back(0.96);
remaining -= 96.0;
}
sequence.push_back(remaining / 100.0);
}
randomGen = std::make_shared<::SequenceRandomGenerator>(sequence);
}
engine->PostCommand(currentPlayer, shardokAction->getIndex(), randomGen);
// Create and return the new state
auto newState = std::make_unique<ShardokGameState>(
engine->GetCurrentGameState(),
scoreCalculator_,
gameSettings_.get(),
isDefender_,
strategy_,
castleCoords_,
*apdCache_,
*alCache_,
criticalTileCoords_);
// Cache the engine on the new state so score() can use it for END_TURN normalization
// The engine's command list may be stale after the action was applied, but that's OK -
// we'll refresh it when we call GetAvailableCommandsForAIPlayer() in score()
newState->setCachedEngine(engine);
// Don't pre-compute hash - let it be computed lazily on first use
// Many states (especially in simulation) never need their hash computed
return newState;
}
void ShardokGameEngine::applyActionMutable(
std::unique_ptr<MCTSGameState>& state,
const MCTSAction& action) const {
auto* shardokState = dynamic_cast<ShardokGameState*>(state.get());
const auto* shardokAction = dynamic_cast<const ShardokAction*>(&action);
if (!shardokState || !shardokAction) {
// Fallback to default implementation
state = applyAction(*state, action);
return;
}
const auto currentPlayer = static_cast<PlayerId>(state->currentPlayerId());
// Use cached engine if available
std::shared_ptr<ShardokEngine> engine;
if (auto cachedEngine = shardokState->getCachedEngine()) {
engine = std::make_shared<ShardokEngine>(*cachedEngine);
} else {
engine = std::make_shared<ShardokEngine>(
gameSettings_,
shardokState->getShardokState(),
criticalTileCoords_,
0,
false);
// Populate command cache (result intentionally unused, just populating cache)
[[maybe_unused]] const auto commands =
engine->GetAvailableCommandsForAIPlayer(currentPlayer);
shardokState->setCachedEngine(engine);
engine = std::make_shared<ShardokEngine>(*engine);
}
engine->PostCommand(currentPlayer, shardokAction->getIndex(), nullptr);
shardokState->getMutableShardokState() = engine->GetCurrentGameState();
// Clear the cached engine and hash since the state has been mutated
shardokState->setCachedEngine(nullptr);
shardokState->invalidateHashCache();
}
std::vector<std::unique_ptr<MCTSAction>> ShardokGameEngine::getLegalActions(
const MCTSGameState& state,
MCTSPlayerId /*rootPlayerId*/,
int currentPlayerFlips,
int maxPlayerFlips) const {
const auto* shardokState = dynamic_cast<const ShardokGameState*>(&state);
if (!shardokState) { return {}; }
const auto currentPlayer = static_cast<PlayerId>(state.currentPlayerId());
// Check if we've exceeded the maximum allowed player flips
// currentPlayerFlips is the number of times the player has changed since root
// maxPlayerFlips is the maximum number of changes we allow
// If maxPlayerFlips is 0, only explore root player's moves (stop when player first changes)
// If maxPlayerFlips is 1, explore through opponent's response (stop after opponent's moves)
if (currentPlayerFlips > maxPlayerFlips) {
return {}; // Stop exploration - we've exceeded the flip limit
}
// Time hash computation
const auto hashStart = std::chrono::high_resolution_clock::now();
const uint64_t stateHash = shardokState->hash();
const auto hashEnd = std::chrono::high_resolution_clock::now();
timeInHashComputation_.fetch_add(
std::chrono::duration_cast<std::chrono::microseconds>(hashEnd - hashStart).count(),
std::memory_order_relaxed);
// Check transposition table for cached legal actions
if (auto it = legalActionsCache_.find(stateHash); it != legalActionsCache_.end()) {
cacheHits_.fetch_add(1, std::memory_order_relaxed);
// Use cached engine
shardokState->setCachedEngine(it->second.engine);
// Get commands from the cached engine (Engine already caches these internally)
const CommandListSPtr commands =
it->second.engine->GetAvailableCommandsForAIPlayer(currentPlayer);
if (!commands || commands->empty()) { return {}; }
// Convert to MCTSActions using stored filtered indices
std::vector<std::unique_ptr<MCTSAction>> actions;
actions.reserve(it->second.filteredIndices.size());
for (const size_t origIdx : it->second.filteredIndices) {
if (origIdx < commands->size()) {
const auto& cmd = commands->at(origIdx);
// Extract essential fields directly from command (no proto conversion!)
actions.push_back(std::make_unique<ShardokAction>(
origIdx,
cmd->GetCommandType(),
cmd->GetPlayerId(),
cmd->GetActorUnitId(),
cmd->GetTargetRow(),
cmd->GetTargetColumn(),
cmd->HasOdds()));
}
}
// Sort actions by weight (descending) to ensure MCTS explores high-value actions first
const std::vector<double> weights = getActionWeights(actions, state);
std::vector<size_t> sortedIndices(actions.size());
std::iota(sortedIndices.begin(), sortedIndices.end(), 0);
std::sort(sortedIndices.begin(), sortedIndices.end(), [&weights](size_t a, size_t b) {
return weights[a] > weights[b];
});
std::vector<std::unique_ptr<MCTSAction>> sortedActions;
sortedActions.reserve(actions.size());
for (size_t idx : sortedIndices) { sortedActions.push_back(std::move(actions[idx])); }
return sortedActions;
}
cacheMisses_.fetch_add(1, std::memory_order_relaxed);
// Time legal actions computation
const auto actionsStart = std::chrono::high_resolution_clock::now();
// Use cached engine if available, otherwise create and cache it
std::shared_ptr<ShardokEngine> engine;
if (auto cachedEngine = shardokState->getCachedEngine()) {
engine = cachedEngine;
} else {
engine = std::make_shared<ShardokEngine>(
gameSettings_,
shardokState->getShardokState(),
criticalTileCoords_);
shardokState->setCachedEngine(engine);
}
const CommandListSPtr commands = engine->GetAvailableCommandsForAIPlayer(currentPlayer);
if (!commands || commands->empty()) { return {}; }
// Filter commands using AICommandFilter (matching original MCTSAI behavior)
// Use gameSettings for battalion type lookups
const std::vector<size_t> filteredIndices = AICommandFilter::FilterCommands(
commands,
currentPlayer,
isDefender_,
shardokState->getShardokState(),
*apdCache_,
[this](BattalionTypeId typeId) {
return gameSettings_->GetGetter().GetBattalionType(typeId);
});
// Convert only filtered commands to MCTSActions
std::vector<std::unique_ptr<MCTSAction>> actions;
actions.reserve(filteredIndices.size());
for (const size_t idx : filteredIndices) {
if (idx < commands->size()) {
const auto& cmd = commands->at(idx);
// Extract essential fields directly from command (no proto conversion!)
actions.push_back(std::make_unique<ShardokAction>(
idx,
cmd->GetCommandType(),
cmd->GetPlayerId(),
cmd->GetActorUnitId(),
cmd->GetTargetRow(),
cmd->GetTargetColumn(),
cmd->HasOdds()));
}
}
// Sort actions by weight (descending) to ensure MCTS explores high-value actions first
// This is critical when maxPlayerFlips is low (e.g., 1), as only the first few actions
// get explored deeply. Original indices are preserved in ShardokAction::getIndex()
const std::vector<double> weights = getActionWeights(actions, state);
// Create index vector for sorting
std::vector<size_t> sortedIndices(actions.size());
std::iota(sortedIndices.begin(), sortedIndices.end(), 0);
// Sort indices by weight (descending)
std::sort(sortedIndices.begin(), sortedIndices.end(), [&weights](size_t a, size_t b) {
return weights[a] > weights[b];
});
// Reorder actions according to sorted indices
std::vector<std::unique_ptr<MCTSAction>> sortedActions;
sortedActions.reserve(actions.size());
for (size_t idx : sortedIndices) { sortedActions.push_back(std::move(actions[idx])); }
actions = std::move(sortedActions);
const auto actionsEnd = std::chrono::high_resolution_clock::now();
timeInLegalActionsComputation_.fetch_add(
std::chrono::duration_cast<std::chrono::microseconds>(actionsEnd - actionsStart)
.count(),
std::memory_order_relaxed);
// Store in transposition table for future lookups
// Note: We only store filtered indices and the engine (which caches commands internally)
// This avoids duplicating heavy protocol buffer objects
// Use lazy_emplace_l to ensure thread-safe insertion (locks the bucket during construction)
legalActionsCache_.lazy_emplace_l(
stateHash,
[&](typename decltype(legalActionsCache_)::value_type& v) {
// Update existing entry
v.second.filteredIndices = filteredIndices;
v.second.engine = engine;
},
[&](const typename decltype(legalActionsCache_)::constructor& ctor) {
// Create new entry
ctor(stateHash, LegalActionsCache{filteredIndices, engine});
});
return actions;
}
bool ShardokGameEngine::isTerminal(const MCTSGameState& state) const { return state.isTerminal(); }
double ShardokGameEngine::evaluateState(const MCTSGameState& state, MCTSPlayerId playerId) const {
return state.score(playerId);
}
std::vector<size_t> ShardokGameEngine::filterActions(
const std::vector<std::unique_ptr<MCTSAction>>& actions,
const MCTSGameState& /*state*/) const {
// All filtering is already done in getLegalActions() using AICommandFilter
// This method is used by simulation policies and doesn't need additional filtering
std::vector<size_t> indices;
indices.reserve(actions.size());
for (size_t i = 0; i < actions.size(); ++i) { indices.push_back(i); }
return indices;
}
std::vector<double> ShardokGameEngine::getActionWeights(
const std::vector<std::unique_ptr<MCTSAction>>& actions,
const MCTSGameState& state) const {
// Cast to ShardokGameState to access Shardok-specific methods
const auto* shardokState = dynamic_cast<const ShardokGameState*>(&state);
if (!shardokState) {
throw MCTSInternalError(
"ShardokGameEngine::getActionWeights called with non-Shardok state - this "
"indicates a type mismatch in the MCTS adapter layer");
}
// Get cached engine and command list for looking up command protos
auto cachedEngine = shardokState->getCachedEngine();
if (!cachedEngine) {
throw MCTSInternalError(
"ShardokGameEngine::getActionWeights called with state that has no cached engine");
}
const auto currentPlayer = static_cast<PlayerId>(state.currentPlayerId());
const CommandListSPtr commands = cachedEngine->GetAvailableCommandsForAIPlayer(currentPlayer);
// Determine if current player is defender (not root player!)
// During simulation we need to use the correct perspective for action weighting
bool currentPlayerIsDefender = false;
const auto& gameState = shardokState->getShardokState();
for (const auto* pi : *gameState->player_infos()) {
if (pi->player_id() == currentPlayer) {
currentPlayerIsDefender = pi->is_defender();
break;
}
}
// Use AIHeuristicWeighting for fast O(1) context-aware command weighting
std::vector<double> weights;
weights.reserve(actions.size());
for (const auto& action : actions) {
const auto* shardokAction = dynamic_cast<const ShardokAction*>(action.get());
if (!shardokAction) {
throw MCTSInternalError(
"ShardokGameEngine::getActionWeights encountered non-Shardok action - this "
"indicates a type mismatch in the MCTS adapter layer");
}
// Look up command proto from cached engine using action's index
const size_t cmdIndex = shardokAction->getIndex();
if (cmdIndex >= commands->size()) {
throw MCTSInternalError(
"ShardokGameEngine::getActionWeights: action index out of bounds");
}
const auto& cmd = commands->at(cmdIndex);
weights.push_back(AIHeuristicWeighting::GetCommandWeight(
cmd->GetCommandType(),
cmd->GetActorUnitId(),
cmd->GetPlayerId(),
Coords{cmd->GetTargetRow(), cmd->GetTargetColumn()},
gameState,
castleCoords_,
apdCache_,
currentPlayerIsDefender, // Use current player's role, not root player's!
[this](BattalionTypeId typeId) {
return gameSettings_->GetGetter().GetBattalionType(typeId);
}));
}
return weights;
}
double ShardokGameEngine::getActionScore(
const MCTSGameState& state,
const MCTSAction& action,
MCTSPlayerId playerId) const {
auto newState = applyAction(state, action);
if (!newState) { return 0.0; }
return newState->score(playerId);
}
bool ShardokGameEngine::shouldStopSearch(
const MCTSGameState& /*state*/,
int /*iterations*/,
std::chrono::steady_clock::time_point /*startTime*/) const {
// Could add early termination logic here
return false;
}
size_t ShardokGameEngine::mapFilteredIndexToOriginal(
size_t filteredIndex,
const MCTSGameState& state) const {
// Get the filtered actions (uses cached engine)
auto actions = getLegalActions(state, state.currentPlayerId(), 0, 0);
// Check bounds
if (filteredIndex >= actions.size()) { return filteredIndex; }
// Extract the original index from the ShardokAction
const auto* shardokAction = dynamic_cast<const ShardokAction*>(actions[filteredIndex].get());
if (!shardokAction) { return filteredIndex; }
// ShardokAction stores the original unfiltered index
return shardokAction->getIndex();
}
void ShardokGameEngine::reportCacheStatistics() const {
const uint64_t hits = cacheHits_.load(std::memory_order_relaxed);
const uint64_t misses = cacheMisses_.load(std::memory_order_relaxed);
const uint64_t hashTime = timeInHashComputation_.load(std::memory_order_relaxed);
const uint64_t actionsTime = timeInLegalActionsComputation_.load(std::memory_order_relaxed);
const uint64_t totalLookups = hits + misses;
if (totalLookups > 0) {
const double hitRate = static_cast<double>(hits) / static_cast<double>(totalLookups);
const double avgHashTimeUs =
static_cast<double>(hashTime) / static_cast<double>(totalLookups);
const double avgActionsTimeUs =
misses > 0 ? static_cast<double>(actionsTime) / static_cast<double>(misses) : 0.0;
printf("Legal Actions Cache Stats:\n");
printf(" Lookups: %llu hits, %llu misses, %.1f%% hit rate, %zu entries\n",
static_cast<unsigned long long>(hits),
static_cast<unsigned long long>(misses),
hitRate * 100.0,
legalActionsCache_.size());
printf(" Timing: %.2f us avg hash, %.2f us avg actions (on miss)\n",
avgHashTimeUs,
avgActionsTimeUs);
printf(" Total time: %.2f ms in hash, %.2f ms in actions\n",
hashTime / 1000.0,
actionsTime / 1000.0);
// Calculate if transposition table is worth it
const double timeWithCache = hashTime + actionsTime;
const double timeWithoutCache =
avgActionsTimeUs * static_cast<double>(totalLookups); // All lookups recompute
const double savings = (timeWithoutCache - timeWithCache) / timeWithoutCache * 100.0;
printf(" Cache savings: %.1f%% vs. no cache (%.2f ms saved)\n",
savings,
(timeWithoutCache - timeWithCache) / 1000.0);
}
}
void ShardokGameEngine::resetCacheStatistics() {
cacheHits_.store(0, std::memory_order_relaxed);
cacheMisses_.store(0, std::memory_order_relaxed);
timeInHashComputation_.store(0, std::memory_order_relaxed);
timeInLegalActionsComputation_.store(0, std::memory_order_relaxed);
}
ChanceOutcomeInfo ShardokGameEngine::getBinaryOutcomeInfo(
const MCTSGameState& state,
const MCTSAction& action) const {
const auto* shardokState = dynamic_cast<const ShardokGameState*>(&state);
const auto* shardokAction = dynamic_cast<const ShardokAction*>(&action);
if (!shardokState || !shardokAction) {
throw ShardokInternalErrorException("Invalid state or action type in getBinaryOutcomeInfo");
}
// Check for multi-outcome commands (roll affects outcome quality, not just success/failure)
// These use multiOutcome() with fixed seeds to sample the range of possible results
using namespace net::eagle0::shardok::common;
const auto commandType = static_cast<CommandType>(shardokAction->getType());
switch (commandType) {
case END_TURN_COMMAND:
// END_TURN has random effects (fire spread, weather changes)
return ChanceOutcomeInfo::multiOutcome(5);
case MELEE_COMMAND:
case ARCHERY_COMMAND:
case CHARGE_COMMAND:
case REDUCE_COMMAND:
// Combat/siege commands: OpenEndedPercentile roll affects damage dealt
// Use 5 outcomes to sample the roll distribution
return ChanceOutcomeInfo::multiOutcome(5);
case CHALLENGE_DUEL_COMMAND:
// Duels have multiple combat rounds with rolls, so outcomes vary significantly
return ChanceOutcomeInfo::multiOutcome(5);
default:
// Continue to binary outcome handling below
break;
}
const auto currentPlayer = static_cast<PlayerId>(state.currentPlayerId());
// Get or create the engine for this state
std::shared_ptr<ShardokEngine> engine;
if (auto cachedEngine = shardokState->getCachedEngine()) {
engine = cachedEngine;
} else {
engine = std::make_shared<ShardokEngine>(
gameSettings_,
shardokState->getShardokState(),
criticalTileCoords_,
0,
false);
// Populate command cache
[[maybe_unused]] const auto commands =
engine->GetAvailableCommandsForAIPlayer(currentPlayer);
shardokState->setCachedEngine(engine);
}
// Get command descriptors
const auto descriptors = engine->GetAvailableCommandsForAIPlayer(currentPlayer);
const size_t actionIndex = shardokAction->getIndex();
if (actionIndex >= descriptors->size()) {
throw ShardokInternalErrorException("Action index out of range in getBinaryOutcomeInfo");
}
const auto& descriptor = descriptors->at(actionIndex);
// Get success probability for binary outcome actions
if (!descriptor->HasOdds()) {
throw ShardokInternalErrorException("Action does not have odds in getBinaryOutcomeInfo");
}
const auto successChancePercentile = descriptor->GetOddsPercentile();
const double successProbability = static_cast<double>(successChancePercentile) / 100.0;
return ChanceOutcomeInfo::binary(successProbability);
}
void ShardokGameEngine::clearLegalActionsCache() { legalActionsCache_.clear(); }
// Extern-linkage function for testing
void clearLegalActionsCache_ForTesting() { ShardokGameEngine::clearLegalActionsCache(); }
} // namespace shardok::mcts
@@ -1,144 +0,0 @@
//
// Shardok-specific game engine adapter for MCTS
//
#ifndef EAGLE0_SHARDOK_GAME_ENGINE_HPP
#define EAGLE0_SHARDOK_GAME_ENGINE_HPP
#include <atomic>
#include <functional>
#include <gtl/phmap.hpp>
#include <memory>
#include <vector>
#include "src/main/cpp/net/eagle0/common/mcts/abstract/MCTSGameEngine.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIAttackLocations.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIStrategy.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokEngine.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistancesCache.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/map/CoordsSet.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
namespace shardok {
// Forward declarations
class AICommandFilter;
class AIScoreCalculator;
class RandomGenerator;
// Use existing type definitions from the Shardok codebase
// GameSettingsSPtr and SettingsGetter are defined in GameSettings.hpp
namespace mcts {
class ShardokGameEngine : public MCTSGameEngine {
public:
ShardokGameEngine(
const ShardokEngine* engine,
const AIScoreCalculator* scoreCalculator,
const GameSettingsSPtr& gameSettings,
const APDCache* apdCache,
const ALCache* alCache,
bool isDefender,
const AIStrategy& strategy,
const CoordsSet& castleCoords,
const CoordsSet& criticalTileCoords);
// MCTSGameEngine interface implementation
[[nodiscard]] std::unique_ptr<MCTSGameState> applyAction(
const MCTSGameState& state,
const MCTSAction& action,
double deterministicRoll = -1.0) const override;
void applyActionMutable(std::unique_ptr<MCTSGameState>& state, const MCTSAction& action)
const override;
[[nodiscard]] std::vector<std::unique_ptr<MCTSAction>> getLegalActions(
const MCTSGameState& state,
MCTSPlayerId rootPlayerId,
int currentPlayerFlips,
int maxPlayerFlips) const override;
[[nodiscard]] bool isTerminal(const MCTSGameState& state) const override;
[[nodiscard]] double evaluateState(const MCTSGameState& state, MCTSPlayerId playerId)
const override;
[[nodiscard]] std::vector<size_t> filterActions(
const std::vector<std::unique_ptr<MCTSAction>>& actions,
const MCTSGameState& state) const override;
[[nodiscard]] std::vector<double> getActionWeights(
const std::vector<std::unique_ptr<MCTSAction>>& actions,
const MCTSGameState& state) const override;
[[nodiscard]] double getActionScore(
const MCTSGameState& state,
const MCTSAction& action,
MCTSPlayerId playerId) const override;
[[nodiscard]] bool shouldStopSearch(
const MCTSGameState& state,
int iterations,
std::chrono::steady_clock::time_point startTime) const override;
[[nodiscard]] size_t mapFilteredIndexToOriginal(
size_t filteredIndex,
const MCTSGameState& state) const override;
[[nodiscard]] BinaryOutcomeInfo getBinaryOutcomeInfo(
const MCTSGameState& state,
const MCTSAction& action) const override;
// Report transposition table statistics
void reportCacheStatistics() const;
// Reset cache statistics
void resetCacheStatistics();
private:
// Transposition table entry for caching legal actions
// Note: We don't store command protos since the Engine already caches them
struct LegalActionsCache {
std::vector<size_t> filteredIndices;
std::shared_ptr<ShardokEngine> engine; // Engine with populated command cache
};
const AIScoreCalculator* scoreCalculator_;
const GameSettingsSPtr gameSettings_;
const APDCache* apdCache_;
const ALCache* alCache_;
bool isDefender_;
AIStrategy strategy_;
const CoordsSet castleCoords_; // Own the data to avoid dangling references
// Computed once to avoid 8.5% overhead per engine construction
const CoordsSet criticalTileCoords_; // Own the data to avoid dangling references
// Transposition table for legal actions (shared across threads with lock-free hash map)
// parallel_flat_hash_map provides thread-safe concurrent access without explicit locking
// Using 8 submaps (N=8) to reduce contention with default 16 MCTS threads
static gtl::parallel_flat_hash_map<
uint64_t,
LegalActionsCache,
std::hash<uint64_t>,
std::equal_to<uint64_t>,
std::allocator<std::pair<const uint64_t, LegalActionsCache>>,
8,
std::mutex>
legalActionsCache_;
static std::atomic<uint64_t> cacheHits_;
static std::atomic<uint64_t> cacheMisses_;
// Performance timing (in microseconds)
static std::atomic<uint64_t> timeInHashComputation_;
static std::atomic<uint64_t> timeInLegalActionsComputation_;
public:
// Clear the static legal actions cache (useful for tests)
static void clearLegalActionsCache();
};
} // namespace mcts
} // namespace shardok
#endif // EAGLE0_SHARDOK_GAME_ENGINE_HPP
@@ -1,125 +0,0 @@
//
// Shardok-specific game state adapter implementation
//
#include "ShardokGameState.hpp"
#include <sstream>
#include <string>
#include <utility>
#include "src/main/cpp/net/eagle0/shardok/ai/score/AIScoreCalculator.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
namespace shardok::mcts {
ShardokGameState::ShardokGameState(
GameStateW state,
const AIScoreCalculator* calculator,
const GameSettings* settings,
const bool isDefender,
AIStrategy strategy,
const CoordsSet& castleCoords,
const APDCache& apdCache,
const ALCache& alCache,
const CoordsSet& criticalTileCoords)
: state_(std::move(state)),
scoreCalculator_(calculator),
settings_(settings),
isDefender_(isDefender),
strategy_(std::move(strategy)),
castleCoords_(castleCoords),
apdCache_(apdCache),
alCache_(alCache),
criticalTileCoords_(criticalTileCoords) {}
uint64_t ShardokGameState::hash() const {
if (!hashCached_) {
cachedHash_ = state_.ComputeFNV1aHash();
hashCached_ = true;
}
return cachedHash_;
}
double ShardokGameState::score(MCTSPlayerId playerId) const {
// Honor the interface contract: score() should return evaluation from playerId's perspective.
// Map the requested playerId to defender/attacker role to determine scoring perspective.
// Look up which player ID is the defender from game state
bool foundDefender = false;
bool requestedPlayerIsDefender = false;
if (state_->player_infos()) {
for (const auto* pi : *state_->player_infos()) {
if (pi && pi->is_defender()) {
foundDefender = true;
requestedPlayerIsDefender = (static_cast<PlayerId>(playerId) == pi->player_id());
break;
}
}
}
// Fallback: if we can't determine from game state, use isDefender_ which represents
// the root player's role (and playerId is always the root player in practice)
const bool scoreFromDefenderPerspective =
foundDefender ? requestedPlayerIsDefender : isDefender_;
// Score the current state directly
return scoreCalculator_
->GuessedStateScore(scoreFromDefenderPerspective, state_, strategy_, castleCoords_);
}
MCTSPlayerId ShardokGameState::currentPlayerId() const { return state_->current_player(); }
bool ShardokGameState::isTerminal() const {
// Check if game status indicates the game is over
if (state_->status()) {
const auto gameStatus = state_->status()->state();
if (gameStatus == net::eagle0::shardok::storage::fb::GameStatus_::State_VICTORY ||
gameStatus == net::eagle0::shardok::storage::fb::GameStatus_::State_DRAW) {
return true;
}
}
// Check max rounds
if (state_->current_round() >= settings_->GetGetter().Backing().max_rounds()) { return true; }
return false;
}
std::unique_ptr<MCTSGameState> ShardokGameState::clone() const {
auto cloned = std::make_unique<ShardokGameState>(
state_,
scoreCalculator_,
settings_,
isDefender_,
strategy_,
castleCoords_,
apdCache_,
alCache_,
criticalTileCoords_);
// Don't copy the cached engine - each state needs its own
return cloned;
}
bool ShardokGameState::equals(const MCTSGameState& other) const {
const auto* shardokOther = dynamic_cast<const ShardokGameState*>(&other);
if (!shardokOther) { return false; }
return hash() == shardokOther->hash();
}
MCTSPlayerId ShardokGameState::getWinner() const {
// Note: FlatBuffer doesn't have a winner field
// In practice, this would need to determine winner from victory conditions
return -1; // No winner
}
std::string ShardokGameState::toString() const {
std::stringstream ss;
ss << "ShardokGameState[Round:" << static_cast<int>(state_->current_round())
<< " Player:" << currentPlayerId() << " Hash:" << hash() << "]";
return ss.str();
}
} // namespace shardok::mcts
@@ -1,83 +0,0 @@
//
// Shardok-specific game state adapter for MCTS
//
#ifndef EAGLE0_SHARDOK_GAME_STATE_HPP
#define EAGLE0_SHARDOK_GAME_STATE_HPP
#include <memory>
#include "src/main/cpp/net/eagle0/common/mcts/abstract/MCTSGameState.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIStrategy.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/GameStateW.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokEngine.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistancesCache.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/map/CoordsSet.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
namespace shardok {
// Forward declarations
class AIScoreCalculator;
namespace mcts {
class ShardokGameState : public MCTSGameState {
public:
ShardokGameState(
GameStateW state,
const AIScoreCalculator* calculator,
const GameSettings* settings,
bool isDefender,
AIStrategy strategy,
const CoordsSet& castleCoords,
const APDCache& apdCache,
const ALCache& alCache,
const CoordsSet& criticalTileCoords);
// MCTSGameState interface implementation
[[nodiscard]] uint64_t hash() const override;
[[nodiscard]] double score(MCTSPlayerId playerId) const override;
[[nodiscard]] MCTSPlayerId currentPlayerId() const override;
[[nodiscard]] bool isTerminal() const override;
[[nodiscard]] std::unique_ptr<MCTSGameState> clone() const override;
[[nodiscard]] bool equals(const MCTSGameState& other) const override;
[[nodiscard]] MCTSPlayerId getWinner() const override;
[[nodiscard]] std::string toString() const override;
// Shardok-specific accessors
[[nodiscard]] const GameStateW& getShardokState() const { return state_; }
[[nodiscard]] GameStateW& getMutableShardokState() { return state_; }
[[nodiscard]] bool isDefender() const { return isDefender_; }
[[nodiscard]] const GameSettings* getSettings() const { return settings_; }
[[nodiscard]] const CoordsSet& getCriticalTileCoords() const { return criticalTileCoords_; }
// Engine caching for performance (avoids recomputing available commands)
void setCachedEngine(std::shared_ptr<ShardokEngine> engine) const { cachedEngine_ = engine; }
[[nodiscard]] std::shared_ptr<ShardokEngine> getCachedEngine() const { return cachedEngine_; }
// Invalidate hash cache when state is mutated
void invalidateHashCache() const {
hashCached_ = false;
cachedHash_ = 0;
}
private:
GameStateW state_;
const AIScoreCalculator* scoreCalculator_;
const GameSettings* settings_;
bool isDefender_;
AIStrategy strategy_;
const CoordsSet castleCoords_; // Own the data to avoid dangling references
const APDCache& apdCache_;
const ALCache& alCache_;
mutable uint64_t cachedHash_ = 0;
mutable bool hashCached_ = false;
const CoordsSet criticalTileCoords_; // Own the data to avoid dangling references
mutable std::shared_ptr<ShardokEngine> cachedEngine_; // Engine with cached available commands
};
} // namespace mcts
} // namespace shardok
#endif // EAGLE0_SHARDOK_GAME_STATE_HPP
@@ -1,82 +0,0 @@
//
// Factory implementation for creating Shardok-specific MCTS components
//
#include "ShardokMCTSFactory.hpp"
#include "ShardokAction.hpp"
#include "ShardokGameEngine.hpp"
#include "ShardokGameState.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokEngine.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
namespace shardok::mcts {
std::unique_ptr<MCTSGameEngine> ShardokMCTSFactory::createGameEngine(
const ShardokEngine& engine,
const AIScoreCalculator* scoreCalculator,
const GameSettingsSPtr& gameSettings,
const APDCache& apdCache,
const ALCache& alCache,
bool isDefender,
const AIStrategy& strategy,
const CoordsSet& castleCoords,
const CoordsSet& criticalTileCoords) {
return std::make_unique<ShardokGameEngine>(
&engine,
scoreCalculator,
gameSettings,
&apdCache,
&alCache,
isDefender,
strategy,
castleCoords,
criticalTileCoords);
}
std::unique_ptr<MCTSGameState> ShardokMCTSFactory::createGameState(
const GameStateW& state,
const AIScoreCalculator* scoreCalculator,
const GameSettingsSPtr& settings,
bool isDefender,
const AIStrategy& strategy,
const CoordsSet& castleCoords,
const APDCache& apdCache,
const ALCache& alCache,
const CoordsSet& criticalTileCoords) {
return std::make_unique<ShardokGameState>(
state,
scoreCalculator,
settings.get(), // Get raw pointer from shared_ptr
isDefender,
strategy,
castleCoords,
apdCache,
alCache,
criticalTileCoords);
}
std::vector<std::unique_ptr<MCTSAction>> ShardokMCTSFactory::createActionsFromCommandList(
const CommandListSPtr& commands) {
std::vector<std::unique_ptr<MCTSAction>> actions;
if (!commands) { return actions; }
actions.reserve(commands->size());
for (size_t i = 0; i < commands->size(); ++i) {
const auto& cmd = (*commands)[i];
// Extract essential fields directly from command (no proto conversion!)
actions.push_back(std::make_unique<ShardokAction>(
i,
cmd->GetCommandType(),
cmd->GetPlayerId(),
cmd->GetActorUnitId(),
cmd->GetTargetRow(),
cmd->GetTargetColumn(),
cmd->HasOdds()));
}
return actions;
}
} // namespace shardok::mcts
@@ -1,70 +0,0 @@
//
// Factory for creating Shardok-specific MCTS components
//
#ifndef EAGLE0_SHARDOK_MCTS_FACTORY_HPP
#define EAGLE0_SHARDOK_MCTS_FACTORY_HPP
#include <functional>
#include <memory>
#include <vector>
#include "src/main/cpp/net/eagle0/shardok/ai/AIAttackLocations.hpp"
#include "src/main/cpp/net/eagle0/shardok/ai/AIStrategy.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokCTypes.h"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokCommand.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistancesCache.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/map/CoordsSet.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
namespace shardok {
// Forward declarations
class ShardokEngine;
class AICommandFilter;
class AIScoreCalculator;
class GameStateW;
class GameSettings;
namespace mcts {
// Forward declarations
class MCTSGameEngine;
class MCTSGameState;
class MCTSAction;
class ShardokMCTSFactory {
public:
// Create a Shardok game engine adapter
[[nodiscard]] static std::unique_ptr<MCTSGameEngine> createGameEngine(
const ShardokEngine& engine,
const AIScoreCalculator* scoreCalculator,
const GameSettingsSPtr& gameSettings,
const APDCache& apdCache,
const ALCache& alCache,
bool isDefender,
const AIStrategy& strategy,
const CoordsSet& castleCoords,
const CoordsSet& criticalTileCoords);
// Create a Shardok game state adapter
[[nodiscard]] static std::unique_ptr<MCTSGameState> createGameState(
const GameStateW& state,
const AIScoreCalculator* scoreCalculator,
const GameSettingsSPtr& settings,
bool isDefender,
const AIStrategy& strategy,
const CoordsSet& castleCoords,
const APDCache& apdCache,
const ALCache& alCache,
const CoordsSet& criticalTileCoords);
// Convert from command list to MCTS actions
[[nodiscard]] static std::vector<std::unique_ptr<MCTSAction>> createActionsFromCommandList(
const CommandListSPtr& commands);
};
} // namespace mcts
} // namespace shardok
#endif // EAGLE0_SHARDOK_MCTS_FACTORY_HPP
@@ -0,0 +1,16 @@
load("//tools:copts.bzl", "COPTS")
cc_library(
name = "mcts_node",
hdrs = ["MCTSNode.hpp"],
copts = COPTS,
visibility = [
"//src/main/cpp/net/eagle0/shardok/ai/mcts:__pkg__",
"//src/test/cpp/net/eagle0/shardok/ai/mcts:__pkg__",
],
deps = [
"//src/main/cpp/net/eagle0/shardok/library:game_state_w",
"//src/main/cpp/net/eagle0/shardok/library:shardok_c_types",
"//src/main/protobuf/net/eagle0/shardok/common:command_type_cc_proto",
],
)
@@ -0,0 +1,210 @@
//
// Internal MCTS Node structure for Shardok AI
// This is an implementation detail and should not be used by external code
//
#ifndef EAGLE0_INTERNAL_MCTSNODE_HPP
#define EAGLE0_INTERNAL_MCTSNODE_HPP
#include <cmath>
#include <cstdio>
#include <limits>
#include <memory>
#include <vector>
#include "src/main/cpp/net/eagle0/shardok/library/GameStateW.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokCTypes.h"
// Suppress the protobuf deprecation warning temporarily
#pragma GCC diagnostic push
#pragma GCC diagnostic ignored "-Wdeprecated-redundant-constexpr-static-def"
#include "src/main/protobuf/net/eagle0/shardok/common/command_type.pb.h"
#pragma GCC diagnostic pop
namespace shardok {
namespace internal {
// Import CommandType for use within the internal namespace
using CommandType = net::eagle0::shardok::common::CommandType;
// MCTS Node structure
struct MCTSNode {
// Command information
size_t commandIndex;
CommandType commandType;
int actorUnitId = -1; // Unit performing the command (-1 if not applicable)
int targetRow = -1; // Target coordinate row (-1 if not applicable)
int targetCol = -1; // Target coordinate column (-1 if not applicable)
// Score information
double immediateScore;
double lookaheadScore;
// Game state after this command
GameStateW resultingGameState;
// MCTS statistics
int visitCount = 0;
double totalReward = 0.0;
double averageReward = 0.0;
double ucb1Value = 0.0;
// Tree structure
std::vector<std::unique_ptr<MCTSNode>> children;
std::vector<size_t> untriedCommands;
bool fullyExpanded = false;
MCTSNode* parent = nullptr;
// Game context
PlayerId playerId;
int depth = 0;
bool isDefender = false;
bool isTerminal = false;
// Transposition detection
uint64_t stateHash = 0;
bool isRedundant = false; // True if this node represents a duplicate state
MCTSNode(
const size_t cmdIndex,
const CommandType cmdType,
const PlayerId pid,
const int d,
const bool defender)
: commandIndex(cmdIndex),
commandType(cmdType),
immediateScore(0.0),
lookaheadScore(0.0),
playerId(pid),
depth(d),
isDefender(defender) {}
// Iterative destructor to avoid stack overflow with deep trees
~MCTSNode() {
// Use iterative approach to destroy children
std::vector<std::unique_ptr<MCTSNode>> nodesToDestroy;
nodesToDestroy.swap(children);
while (!nodesToDestroy.empty()) {
// Take ownership of all children from the current batch
std::vector<std::unique_ptr<MCTSNode>> currentBatch;
currentBatch.swap(nodesToDestroy);
// Collect grandchildren for next iteration
for (const auto& node : currentBatch) {
if (node && !node->children.empty()) {
for (auto& child : node->children) {
nodesToDestroy.push_back(std::move(child));
}
node->children.clear();
}
}
// currentBatch goes out of scope here, destroying nodes with no children
}
}
// Calculate UCB1 value for this node
void CalculateUCB1(const double explorationConstant) {
if (visitCount == 0) {
ucb1Value = std::numeric_limits<double>::max();
} else if (parent && parent->visitCount > 0) {
ucb1Value = averageReward +
explorationConstant * std::sqrt(std::log(parent->visitCount) / visitCount);
} else {
ucb1Value = averageReward;
}
}
// Check if this node can be expanded
[[nodiscard]] bool CanExpand() const { return !fullyExpanded && !untriedCommands.empty(); }
// Get best child based on UCB1
[[nodiscard]] MCTSNode* GetBestChild(const double explorationConstant) const {
if (children.empty()) return nullptr;
MCTSNode* bestChild = nullptr;
double bestValue = -std::numeric_limits<double>::max();
static int selectionCallCount = 0;
const bool shouldDebug = selectionCallCount < 5;
for (auto& child : children) {
// Skip redundant nodes
if (child->isRedundant) continue;
child->CalculateUCB1(explorationConstant);
if (child->ucb1Value > bestValue) {
bestValue = child->ucb1Value;
bestChild = child.get();
}
if (shouldDebug && child->visitCount > 0) {
printf("UCB1 Debug: cmd:%zu visits:%d reward:%.2f ucb1:%.2f%s\n",
child->commandIndex,
child->visitCount,
child->averageReward,
child->ucb1Value,
child->isRedundant ? " [REDUNDANT]" : "");
}
}
if (shouldDebug) {
if (bestChild) {
printf("UCB1 Selected: cmd:%zu ucb1:%.2f\n",
bestChild->commandIndex,
bestChild->ucb1Value);
} else {
printf("UCB1 Selected: nullptr (all children redundant)\n");
}
selectionCallCount++;
}
return bestChild;
}
// Get best child based on average reward (for final selection)
[[nodiscard]] MCTSNode* GetBestFinalChild() const {
if (children.empty()) return nullptr;
MCTSNode* bestChild = nullptr;
double bestScore = -std::numeric_limits<double>::max();
int bestVisits = 0;
for (const auto& child : children) {
// Skip redundant nodes
if (child->isRedundant) continue;
// For final selection, prefer most-visited node (robust child selection)
// Only consider nodes that have been visited
if (child->visitCount > bestVisits) {
bestVisits = child->visitCount;
bestScore = child->averageReward;
bestChild = child.get();
} else if (child->visitCount == bestVisits && child->averageReward > bestScore) {
// Tie-break on average reward
bestScore = child->averageReward;
bestChild = child.get();
}
}
// If no child was visited (shouldn't happen), fall back to lookahead score
if (!bestChild && !children.empty()) {
for (const auto& child : children) {
// Skip redundant nodes
if (child->isRedundant) continue;
if (child->lookaheadScore > bestScore) {
bestScore = child->lookaheadScore;
bestChild = child.get();
}
}
}
return bestChild;
}
};
} // namespace internal
} // namespace shardok
#endif // EAGLE0_INTERNAL_MCTSNODE_HPP
@@ -1,56 +0,0 @@
//
// Created by dancrosby on 3/4/20.
//
#ifndef EAGLE0_AISCORECALCULATOR_HPP
#define EAGLE0_AISCORECALCULATOR_HPP
#include <future>
#include "src/main/cpp/net/eagle0/shardok/ai/AIAttackLocations.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/GameStateW.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/ShardokCTypes.h"
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/ActionPointDistancesCache.hpp"
#include "src/main/cpp/net/eagle0/shardok/library/map/CoordsSet.hpp"
namespace shardok {
using shardok::PlayerId;
using std::future;
using std::vector;
using ScoreValue = double;
// Forward declarations
class ShardokEngine;
struct AIStrategy;
/// Abstract base class for AI scoring algorithms.
/// Allows testing different scoring strategies by implementing different scorers.
class AIScoreCalculator {
public:
virtual ~AIScoreCalculator() = default;
// Rule of five: explicitly default or delete copy/move operations
AIScoreCalculator(const AIScoreCalculator &) = default;
AIScoreCalculator &operator=(const AIScoreCalculator &) = default;
AIScoreCalculator(AIScoreCalculator &&) = default;
AIScoreCalculator &operator=(AIScoreCalculator &&) = default;
protected:
AIScoreCalculator() = default;
public:
/// Evaluate the score of a guessed game state based on the current AI strategy.
/// DOES NOT perform lookahead - this is pure state evaluation.
/// For lookahead search, use AICommandEvaluator which depends on this interface.
[[nodiscard]] virtual auto GuessedStateScore(
bool isDefender,
const GameStateW &state,
const AIStrategy &aiStrategy,
const CoordsSet &allCastleCoords) const -> ScoreValue = 0;
};
} // namespace shardok
#endif // EAGLE0_AISCORECALCULATOR_HPP

Some files were not shown because too many files have changed in this diff Show More