mirror of
https://github.com/nolen777/eagle0.git
synced 2026-07-29 03:05:43 +00:00
Compare commits
219
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5331bdedb1 | ||
|
|
364b0978d8 | ||
|
|
8e2575be50 | ||
|
|
913d927902 | ||
|
|
066381e24e | ||
|
|
ced52b0195 | ||
|
|
262ba36436 | ||
|
|
d16aa63c00 | ||
|
|
636bf8f9f3 | ||
|
|
b84df05953 | ||
|
|
dcf0261ac3 | ||
|
|
7afe4e788a | ||
|
|
2e4fc0d230 | ||
|
|
a5b608d18a | ||
|
|
e6519fef20 | ||
|
|
94d49e61d7 | ||
|
|
f3e2873f34 | ||
|
|
c74ddb8983 | ||
|
|
c05e5f7f37 | ||
|
|
9d3967c58d | ||
|
|
07f27ea0ff | ||
|
|
b73d834fab | ||
|
|
deecd5a9ca | ||
|
|
314ff83d24 | ||
|
|
9adfd84498 | ||
|
|
63901e24e5 | ||
|
|
a4a128fe34 | ||
|
|
9cee497886 | ||
|
|
42294da2f6 | ||
|
|
75a129fb4c | ||
|
|
71a1858168 | ||
|
|
52f0cbe180 | ||
|
|
31dc53cc9e | ||
|
|
2a0654f884 | ||
|
|
1dd6eabc15 | ||
|
|
fcdc7d80b8 | ||
|
|
54db688c4e | ||
|
|
792c4f2b53 | ||
|
|
5aad32f5d9 | ||
|
|
78a833c086 | ||
|
|
f7c382446e | ||
|
|
a380eca47e | ||
|
|
90f239c696 | ||
|
|
be393a4cdd | ||
|
|
9e97f71bb9 | ||
|
|
113d54b936 | ||
|
|
5d6c2fef90 | ||
|
|
5e2e7a454c | ||
|
|
914141aed1 | ||
|
|
2ae235e933 | ||
|
|
14d83def79 | ||
|
|
170e998324 | ||
|
|
0dcdac1719 | ||
|
|
164933dbdd | ||
|
|
bf4db493ab | ||
|
|
89cabe9d17 | ||
|
|
dde7a58b44 | ||
|
|
486e99a02d | ||
|
|
2982927200 | ||
|
|
0551453536 | ||
|
|
ce357c612e | ||
|
|
f9e69b6f75 | ||
|
|
f43e914720 | ||
|
|
6c50c0da24 | ||
|
|
d265b76607 | ||
|
|
09a51e4280 | ||
|
|
5593effe69 | ||
|
|
44c268de93 | ||
|
|
0a40acb84d | ||
|
|
9603b497d2 | ||
|
|
0551dd0f13 | ||
|
|
45c4cf783d | ||
|
|
72c52e0b0d | ||
|
|
dffd569ed7 | ||
|
|
a1ffae91a8 | ||
|
|
45a32af435 | ||
|
|
1df8ec68e8 | ||
|
|
53d6e6f63d | ||
|
|
bc84cf6871 | ||
|
|
865a34d00a | ||
|
|
1a751cea6d | ||
|
|
5a8a343bcc | ||
|
|
5979dc7372 | ||
|
|
87474888f9 | ||
|
|
1b3697a40c | ||
|
|
a45b5dadd8 | ||
|
|
9f910bf849 | ||
|
|
c5466e38a8 | ||
|
|
958104b238 | ||
|
|
95e1d80e78 | ||
|
|
0dce9f47b0 | ||
|
|
6e788f4388 | ||
|
|
90d0918233 | ||
|
|
e6038927f1 | ||
|
|
acad796662 | ||
|
|
83c4ac7d38 | ||
|
|
7db07dc371 | ||
|
|
f1b843873a | ||
|
|
e8aefbb6ee | ||
|
|
1c51cc080f | ||
|
|
3946f2eb2d | ||
|
|
49bbdb1d2c | ||
|
|
bbdc30a4af | ||
|
|
19f5cf9e89 | ||
|
|
9033571110 | ||
|
|
e9e557f8f6 | ||
|
|
b32d252df3 | ||
|
|
d4723db2d1 | ||
|
|
7ce3cca731 | ||
|
|
24d21d402d | ||
|
|
e9fb1c5a87 | ||
|
|
b8b7d3a980 | ||
|
|
e503a8af9d | ||
|
|
9c4f46b6ca | ||
|
|
da453bb353 | ||
|
|
618cd18f44 | ||
|
|
429725c4e1 | ||
|
|
54bdefd75c | ||
|
|
98baf7ec66 | ||
|
|
fe4332c107 | ||
|
|
dfa18cef70 | ||
|
|
7845a54b5e | ||
|
|
3106fd9a40 | ||
|
|
19a14174c5 | ||
|
|
1f460a2777 | ||
|
|
12244fb1d4 | ||
|
|
b87910dcf5 | ||
|
|
214790c5e8 | ||
|
|
ffd4ff29d3 | ||
|
|
63e0334ef8 | ||
|
|
4aae50d72c | ||
|
|
0bd6e5b5d2 | ||
|
|
0560f15d1c | ||
|
|
428b91f337 | ||
|
|
d489857692 | ||
|
|
93e6771ded | ||
|
|
430b16bc86 | ||
|
|
8fb518ccad | ||
|
|
bfd4fcebbf | ||
|
|
7c312eb2ef | ||
|
|
209fab050b | ||
|
|
0bc0cbc738 | ||
|
|
48b561a999 | ||
|
|
535fe76620 | ||
|
|
7d21bbe72d | ||
|
|
d38619acb5 | ||
|
|
09f08fc35f | ||
|
|
41caa802df | ||
|
|
289071e0d0 | ||
|
|
cd27c9d084 | ||
|
|
f4f83ce5b5 | ||
|
|
b09bb8332b | ||
|
|
88904c8d50 | ||
|
|
88a5a62a24 | ||
|
|
167ee625a1 | ||
|
|
8b575f8845 | ||
|
|
ef0811183a | ||
|
|
1ae61b4f15 | ||
|
|
f938e0dfd9 | ||
|
|
8f2406b5bd | ||
|
|
00b072cc9a | ||
|
|
e76e040a07 | ||
|
|
e6fddbac45 | ||
|
|
74dfab9c34 | ||
|
|
83094de34e | ||
|
|
cdb56cb060 | ||
|
|
07a88e8de7 | ||
|
|
100051081d | ||
|
|
d04f004d91 | ||
|
|
30c7b3fab3 | ||
|
|
c954ec7084 | ||
|
|
3b6b2e235d | ||
|
|
48b9c6eccf | ||
|
|
30d6068af2 | ||
|
|
d35ac6f40c | ||
|
|
04f9656e67 | ||
|
|
3d7d4a6f70 | ||
|
|
3f573d82d7 | ||
|
|
4596ec8942 | ||
|
|
82ffa57721 | ||
|
|
b311b69e8e | ||
|
|
890d6ecef6 | ||
|
|
7bdcc511f5 | ||
|
|
9b0322e8a3 | ||
|
|
230b3ed891 | ||
|
|
92591ac26f | ||
|
|
6ffdfc87c6 | ||
|
|
b368c093b8 | ||
|
|
65ee957770 | ||
|
|
2e4e001cf5 | ||
|
|
1a63fd3859 | ||
|
|
0e3febad79 | ||
|
|
301b3fff57 | ||
|
|
cb6cb0b17f | ||
|
|
1f335a0ebc | ||
|
|
c8a70728bb | ||
|
|
fbefed617f | ||
|
|
217333e924 | ||
|
|
aeb52042d4 | ||
|
|
7065288cf2 | ||
|
|
60b4c4fcea | ||
|
|
ce532b4b9a | ||
|
|
9acf324ba1 | ||
|
|
0382d08ed5 | ||
|
|
2159f87dc9 | ||
|
|
ca6770b237 | ||
|
|
d80e5e413c | ||
|
|
3e35e678b3 | ||
|
|
0640ea7542 | ||
|
|
bce577758f | ||
|
|
ad8e34ec3d | ||
|
|
215ebbee24 | ||
|
|
1b7b2a2332 | ||
|
|
d5eb0e95c1 | ||
|
|
f96780ac83 | ||
|
|
837825eb90 | ||
|
|
5ea2d7e4d7 | ||
|
|
ff9dd51418 | ||
|
|
e609fcac17 |
@@ -29,6 +29,9 @@ 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
|
||||
|
||||
@@ -34,10 +34,54 @@ 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: success() || failure()
|
||||
if: always()
|
||||
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
|
||||
|
||||
@@ -37,3 +37,4 @@ 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/
|
||||
|
||||
@@ -32,8 +32,9 @@ repos:
|
||||
- id: gazelle
|
||||
name: gazelle
|
||||
language: system
|
||||
entry: bazel run //:gazelle
|
||||
entry: ./scripts/pre-commit-gazelle.sh
|
||||
files: '(\.go|\.proto|BUILD\.bazel|BUILD|WORKSPACE|WORKSPACE\.bazel|\.bzl)$'
|
||||
pass_filenames: false
|
||||
- repo: local
|
||||
hooks:
|
||||
- id: update-action-result-types
|
||||
|
||||
@@ -85,6 +85,16 @@ 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
|
||||
@@ -206,6 +216,31 @@ 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:
|
||||
@@ -244,6 +279,32 @@ done
|
||||
- **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
|
||||
|
||||
## Game Content
|
||||
|
||||
**Maps:** `.e0mj` files in `/src/main/resources/net/eagle0/shardok/maps/`
|
||||
|
||||
@@ -76,6 +76,7 @@ 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",
|
||||
)
|
||||
|
||||
@@ -160,6 +161,10 @@ maven.install(
|
||||
# Other
|
||||
"org.reactivestreams:reactive-streams:1.0.4",
|
||||
"javax.xml.bind:jaxb-api:2.3.1",
|
||||
|
||||
# OkHttp (for SSE with read timeout support)
|
||||
"com.squareup.okhttp3:okhttp:4.12.0",
|
||||
"com.squareup.okhttp3:okhttp-sse:4.12.0",
|
||||
],
|
||||
duplicate_version_warning = "error",
|
||||
fail_if_repin_required = True,
|
||||
|
||||
+1
-2
@@ -1,2 +1 @@
|
||||
|
||||
UNITY_VERSION='6000.2.7f2'
|
||||
UNITY_VERSION='6000.3.0f1'
|
||||
|
||||
@@ -0,0 +1,280 @@
|
||||
# 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
@@ -0,0 +1,309 @@
|
||||
# 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
|
||||
→ ActionResultApplier.applyActionResults()
|
||||
→ GameStateC
|
||||
|
||||
(Proto conversion only at boundaries)
|
||||
```
|
||||
|
||||
### Key Files to Convert
|
||||
|
||||
**Tier 1 - Core Applier:** ✅ **Complete**
|
||||
```
|
||||
src/main/scala/net/eagle0/eagle/library/actions/applier/ActionResultApplierImpl.scala
|
||||
```
|
||||
`ActionResultApplier` applies `ActionResultT` directly to Scala `GameState`. The legacy `ActionResultTApplierImpl` wraps it and converts to/from proto for callers that still need proto types.
|
||||
|
||||
**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 | Status |
|
||||
|------|-------|--------|
|
||||
| `ActionResultApplierImpl.scala` | Applies ActionResultT to Scala GameState | ✅ **Complete** |
|
||||
| `ActionResultTApplierImpl.scala` | Legacy wrapper - converts to/from proto | Keep until all callers migrated |
|
||||
| `RoundPhaseAdvancer.scala` | Uses Scala GameState | ✅ **Complete** |
|
||||
| `RandomStateSequencer.scala` | Threads Scala GameState | ✅ **Complete** |
|
||||
| `VigorXPApplier.scala` | Has both proto and Scala methods | Scala method exists, delete proto method when unused |
|
||||
| `PerformForcedTurnBackAction.scala` | Fully protoless | ✅ **Complete** |
|
||||
| `ResolveBattleAction.scala` | Heavy proto usage | Blocked by proto dependencies |
|
||||
| `InMemoryHistory.scala` | Stores proto results | Pending - vend Scala types |
|
||||
| `PersistedHistory.scala` | Stores proto results | Pending - vend Scala types, convert for disk |
|
||||
| `GameController.scala` | Uses proto for client communication | Keep proto (gRPC boundary) |
|
||||
|
||||
### Remaining Proto Usage in Actions
|
||||
|
||||
The following actions still have proto usage, blocked by utility dependencies:
|
||||
|
||||
| Action | Proto Usage | Blocker |
|
||||
|--------|-------------|---------|
|
||||
| `EndBattleAftermathPhaseAction` | 1 `toProto` call | `ProvinceViewFilter` needs Scala types |
|
||||
| `NewRoundAction` | 1 `fromProto` call | `ChronicleEventGenerator` returns proto |
|
||||
| `EndHandleRiotsPhaseAction` | 1 `toProto` call | `CommandChoiceHelpers` takes proto GameState |
|
||||
| `PerformVassalCommandsPhaseAction` | 1 `toProto` call | `CommandChoiceHelpers` takes proto GameState |
|
||||
| `PerformVassalDefenseDecisionsAction` | 1 `toProto` call | `CommandChoiceHelpers` takes proto GameState |
|
||||
| `EndVassalCommandsPhaseAction` | 1 `toProto` call | `CommandChoiceHelpers` takes proto GameState |
|
||||
| `PerformReconResolutionAction` | Uses `ProvinceViewFilter` | `ProvinceViewFilter` needs Scala types |
|
||||
| `ResolveBattleAction` | Heavy proto usage | Large refactor needed |
|
||||
|
||||
### Estimated Effort (Remaining)
|
||||
|
||||
| Component | Lines | Complexity |
|
||||
|-----------|-------|------------|
|
||||
| `ProvinceViewFilter` to Scala | ~150 | Medium |
|
||||
| `CommandChoiceHelpers` to Scala | ~2000 | High |
|
||||
| `ChronicleEventGenerator` to Scala | ~400 | Medium |
|
||||
| History API updates | ~100 | Low |
|
||||
| **Total Remaining** | **~2650** | |
|
||||
|
||||
### Validation
|
||||
- [x] `ActionResultApplier` created and tested
|
||||
- [x] `RandomStateSequencer` threads Scala GameState throughout
|
||||
- [x] `RoundPhaseAdvancer` uses T-types internally
|
||||
- [ ] `ProvinceViewFilter` uses Scala types
|
||||
- [ ] `CommandChoiceHelpers` uses Scala types
|
||||
- [ ] 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.
@@ -9,6 +9,7 @@ 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
|
||||
)
|
||||
|
||||
|
||||
@@ -40,6 +40,8 @@ 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=
|
||||
|
||||
+102
-12
@@ -1,9 +1,10 @@
|
||||
{
|
||||
"__AUTOGENERATED_FILE_DO_NOT_MODIFY_THIS_FILE_MANUALLY": "THERE_IS_NO_DATA_ONLY_ZUUL",
|
||||
"__INPUT_ARTIFACTS_HASH": 571423113,
|
||||
"__RESOLVED_ARTIFACTS_HASH": 438039003,
|
||||
"__INPUT_ARTIFACTS_HASH": 289080209,
|
||||
"__RESOLVED_ARTIFACTS_HASH": -131178107,
|
||||
"conflict_resolution": {
|
||||
"com.google.guava:failureaccess:1.0.1": "com.google.guava:failureaccess:1.0.2",
|
||||
"com.squareup.okio:okio:2.10.0": "com.squareup.okio:okio:3.6.0",
|
||||
"io.netty:netty-buffer:4.1.110.Final": "io.netty:netty-buffer:4.1.112.Final",
|
||||
"io.netty:netty-codec-http2:4.1.110.Final": "io.netty:netty-codec-http2:4.1.112.Final",
|
||||
"io.netty:netty-codec-http:4.1.110.Final": "io.netty:netty-codec-http:4.1.112.Final",
|
||||
@@ -155,6 +156,18 @@
|
||||
},
|
||||
"version": "1.4.2"
|
||||
},
|
||||
"com.squareup.okhttp3:okhttp": {
|
||||
"shasums": {
|
||||
"jar": "b1050081b14bb7a3a7e55a4d3ef01b5dcfabc453b4573a4fc019767191d5f4e0"
|
||||
},
|
||||
"version": "4.12.0"
|
||||
},
|
||||
"com.squareup.okhttp3:okhttp-sse": {
|
||||
"shasums": {
|
||||
"jar": "bff4fbcaef7aac2d910d4ff46dafaa4e6d15da127df6bac97216da46943a7d4c"
|
||||
},
|
||||
"version": "4.12.0"
|
||||
},
|
||||
"com.squareup.okhttp:okhttp": {
|
||||
"shasums": {
|
||||
"jar": "88ac9fd1bb51f82bcc664cc1eb9c225c90dc4389d660231b4cc737bebfe7d0aa"
|
||||
@@ -163,9 +176,15 @@
|
||||
},
|
||||
"com.squareup.okio:okio": {
|
||||
"shasums": {
|
||||
"jar": "a27f091d34aa452e37227e2cfa85809f29012a8ef2501a9b5a125a978e4fcbc1"
|
||||
"jar": "8e63292e5c53bb93c4a6b0c213e79f15990fed250c1340f1c343880e1c9c39b5"
|
||||
},
|
||||
"version": "2.10.0"
|
||||
"version": "3.6.0"
|
||||
},
|
||||
"com.squareup.okio:okio-jvm": {
|
||||
"shasums": {
|
||||
"jar": "67543f0736fc422ae927ed0e504b98bc5e269fda0d3500579337cb713da28412"
|
||||
},
|
||||
"version": "3.6.0"
|
||||
},
|
||||
"com.thesamet.scalapb:compilerplugin_3": {
|
||||
"shasums": {
|
||||
@@ -444,15 +463,27 @@
|
||||
},
|
||||
"org.jetbrains.kotlin:kotlin-stdlib": {
|
||||
"shasums": {
|
||||
"jar": "b8ab1da5cdc89cb084d41e1f28f20a42bd431538642a5741c52bbfae3fa3e656"
|
||||
"jar": "55e989c512b80907799f854309f3bc7782c5b3d13932442d0379d5c472711504"
|
||||
},
|
||||
"version": "1.4.20"
|
||||
"version": "1.9.10"
|
||||
},
|
||||
"org.jetbrains.kotlin:kotlin-stdlib-common": {
|
||||
"shasums": {
|
||||
"jar": "a7112c9b3cefee418286c9c9372f7af992bd1e6e030691d52f60cb36dbec8320"
|
||||
"jar": "cde3341ba18a2ba262b0b7cf6c55b20c90e8d434e42c9a13e6a3f770db965a88"
|
||||
},
|
||||
"version": "1.4.20"
|
||||
"version": "1.9.10"
|
||||
},
|
||||
"org.jetbrains.kotlin:kotlin-stdlib-jdk7": {
|
||||
"shasums": {
|
||||
"jar": "ac6361bf9ad1ed382c2103d9712c47cdec166232b4903ed596e8876b0681c9b7"
|
||||
},
|
||||
"version": "1.9.10"
|
||||
},
|
||||
"org.jetbrains.kotlin:kotlin-stdlib-jdk8": {
|
||||
"shasums": {
|
||||
"jar": "a4c74d94d64ce1abe53760fe0389dd941f6fc558d0dab35e47c085a11ec80f28"
|
||||
},
|
||||
"version": "1.9.10"
|
||||
},
|
||||
"org.jetbrains:annotations": {
|
||||
"shasums": {
|
||||
@@ -779,12 +810,23 @@
|
||||
"org.checkerframework:checker-qual",
|
||||
"org.ow2.asm:asm"
|
||||
],
|
||||
"com.squareup.okhttp3:okhttp": [
|
||||
"com.squareup.okio:okio",
|
||||
"org.jetbrains.kotlin:kotlin-stdlib-jdk8"
|
||||
],
|
||||
"com.squareup.okhttp3:okhttp-sse": [
|
||||
"com.squareup.okhttp3:okhttp",
|
||||
"org.jetbrains.kotlin:kotlin-stdlib-jdk8"
|
||||
],
|
||||
"com.squareup.okhttp:okhttp": [
|
||||
"com.squareup.okio:okio"
|
||||
],
|
||||
"com.squareup.okio:okio": [
|
||||
"org.jetbrains.kotlin:kotlin-stdlib",
|
||||
"org.jetbrains.kotlin:kotlin-stdlib-common"
|
||||
"com.squareup.okio:okio-jvm"
|
||||
],
|
||||
"com.squareup.okio:okio-jvm": [
|
||||
"org.jetbrains.kotlin:kotlin-stdlib-common",
|
||||
"org.jetbrains.kotlin:kotlin-stdlib-jdk8"
|
||||
],
|
||||
"com.thesamet.scalapb:compilerplugin_3": [
|
||||
"com.google.protobuf:protobuf-java",
|
||||
@@ -992,6 +1034,13 @@
|
||||
"org.jetbrains.kotlin:kotlin-stdlib-common",
|
||||
"org.jetbrains:annotations"
|
||||
],
|
||||
"org.jetbrains.kotlin:kotlin-stdlib-jdk7": [
|
||||
"org.jetbrains.kotlin:kotlin-stdlib"
|
||||
],
|
||||
"org.jetbrains.kotlin:kotlin-stdlib-jdk8": [
|
||||
"org.jetbrains.kotlin:kotlin-stdlib",
|
||||
"org.jetbrains.kotlin:kotlin-stdlib-jdk7"
|
||||
],
|
||||
"org.json4s:json4s-ast_3": [
|
||||
"org.scala-lang:scala3-library_3"
|
||||
],
|
||||
@@ -1451,6 +1500,29 @@
|
||||
"com.google.truth:truth": [
|
||||
"com.google.common.truth"
|
||||
],
|
||||
"com.squareup.okhttp3:okhttp": [
|
||||
"okhttp3",
|
||||
"okhttp3.internal",
|
||||
"okhttp3.internal.authenticator",
|
||||
"okhttp3.internal.cache",
|
||||
"okhttp3.internal.cache2",
|
||||
"okhttp3.internal.concurrent",
|
||||
"okhttp3.internal.connection",
|
||||
"okhttp3.internal.http",
|
||||
"okhttp3.internal.http1",
|
||||
"okhttp3.internal.http2",
|
||||
"okhttp3.internal.io",
|
||||
"okhttp3.internal.platform",
|
||||
"okhttp3.internal.platform.android",
|
||||
"okhttp3.internal.proxy",
|
||||
"okhttp3.internal.publicsuffix",
|
||||
"okhttp3.internal.tls",
|
||||
"okhttp3.internal.ws"
|
||||
],
|
||||
"com.squareup.okhttp3:okhttp-sse": [
|
||||
"okhttp3.internal.sse",
|
||||
"okhttp3.sse"
|
||||
],
|
||||
"com.squareup.okhttp:okhttp": [
|
||||
"com.squareup.okhttp",
|
||||
"com.squareup.okhttp.internal",
|
||||
@@ -1459,7 +1531,7 @@
|
||||
"com.squareup.okhttp.internal.io",
|
||||
"com.squareup.okhttp.internal.tls"
|
||||
],
|
||||
"com.squareup.okio:okio": [
|
||||
"com.squareup.okio:okio-jvm": [
|
||||
"okio",
|
||||
"okio.internal"
|
||||
],
|
||||
@@ -1814,6 +1886,7 @@
|
||||
"kotlin.annotation",
|
||||
"kotlin.collections",
|
||||
"kotlin.collections.builders",
|
||||
"kotlin.collections.jdk8",
|
||||
"kotlin.collections.unsigned",
|
||||
"kotlin.comparisons",
|
||||
"kotlin.concurrent",
|
||||
@@ -1822,24 +1895,36 @@
|
||||
"kotlin.coroutines.cancellation",
|
||||
"kotlin.coroutines.intrinsics",
|
||||
"kotlin.coroutines.jvm.internal",
|
||||
"kotlin.enums",
|
||||
"kotlin.experimental",
|
||||
"kotlin.internal",
|
||||
"kotlin.internal.jdk7",
|
||||
"kotlin.internal.jdk8",
|
||||
"kotlin.io",
|
||||
"kotlin.io.encoding",
|
||||
"kotlin.io.path",
|
||||
"kotlin.jdk7",
|
||||
"kotlin.js",
|
||||
"kotlin.jvm",
|
||||
"kotlin.jvm.functions",
|
||||
"kotlin.jvm.internal",
|
||||
"kotlin.jvm.internal.markers",
|
||||
"kotlin.jvm.internal.unsafe",
|
||||
"kotlin.jvm.jdk8",
|
||||
"kotlin.jvm.optionals",
|
||||
"kotlin.math",
|
||||
"kotlin.properties",
|
||||
"kotlin.random",
|
||||
"kotlin.random.jdk8",
|
||||
"kotlin.ranges",
|
||||
"kotlin.reflect",
|
||||
"kotlin.sequences",
|
||||
"kotlin.streams.jdk8",
|
||||
"kotlin.system",
|
||||
"kotlin.text",
|
||||
"kotlin.time"
|
||||
"kotlin.text.jdk8",
|
||||
"kotlin.time",
|
||||
"kotlin.time.jdk8"
|
||||
],
|
||||
"org.jetbrains:annotations": [
|
||||
"org.intellij.lang.annotations",
|
||||
@@ -2270,8 +2355,11 @@
|
||||
"com.google.protobuf:protobuf-java",
|
||||
"com.google.re2j:re2j",
|
||||
"com.google.truth:truth",
|
||||
"com.squareup.okhttp3:okhttp",
|
||||
"com.squareup.okhttp3:okhttp-sse",
|
||||
"com.squareup.okhttp:okhttp",
|
||||
"com.squareup.okio:okio",
|
||||
"com.squareup.okio:okio-jvm",
|
||||
"com.thesamet.scalapb:compilerplugin_3",
|
||||
"com.thesamet.scalapb:lenses_3",
|
||||
"com.thesamet.scalapb:protoc-bridge_2.13",
|
||||
@@ -2324,6 +2412,8 @@
|
||||
"org.hamcrest:hamcrest-core",
|
||||
"org.jetbrains.kotlin:kotlin-stdlib",
|
||||
"org.jetbrains.kotlin:kotlin-stdlib-common",
|
||||
"org.jetbrains.kotlin:kotlin-stdlib-jdk7",
|
||||
"org.jetbrains.kotlin:kotlin-stdlib-jdk8",
|
||||
"org.jetbrains:annotations",
|
||||
"org.json4s:json4s-ast_3",
|
||||
"org.json4s:json4s-core_3",
|
||||
|
||||
@@ -3,7 +3,8 @@
|
||||
set -euxo pipefail
|
||||
|
||||
/bin/echo "building darwin bundle"
|
||||
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/
|
||||
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/
|
||||
|
||||
/usr/bin/plutil -convert xml1 src/main/csharp/net/eagle0/clients/unity/eagle0/Assets/Plugins/DarwinGodiceBundle.bundle/Contents/Info.plist
|
||||
|
||||
@@ -5,8 +5,9 @@ set -euxo pipefail
|
||||
/bin/echo "build plugins"
|
||||
|
||||
/bin/echo "building darwin bundle"
|
||||
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/
|
||||
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/
|
||||
|
||||
/usr/bin/plutil -convert xml1 src/main/csharp/net/eagle0/clients/unity/eagle0/Assets/Plugins/DarwinGodiceBundle.bundle/Contents/Info.plist
|
||||
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
#!/usr/bin/env bash
|
||||
|
||||
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
|
||||
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
|
||||
|
||||
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
|
||||
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
#!/usr/bin/env bash
|
||||
|
||||
curl -L "https://docs.google.com/spreadsheets/d/1DHEsiv4cY4gE6AX3sVH82K__mpBD1aznIYCQwQxA_F0/export?gid=0&format=tsv" > /tmp/names.tsv
|
||||
curl -L "https://docs.google.com/spreadsheets/d/1DHEsiv4cY4gE6AX3sVH82K__mpBD1aznIYCQwQxA_F0/export?gid=0&format=tsv" | tr -d '\r' > /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" > 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/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/1Z-60cJ_N1IasvqpVb5awKEkIYznEeR2IZSdli47oW88/export?gid=0&format=tsv" > src/main/resources/net/eagle0/eagle/province_map.tsv
|
||||
|
||||
${PWD}/scripts/dlSettings.sh
|
||||
|
||||
Executable
+19
@@ -0,0 +1,19 @@
|
||||
#!/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
|
||||
@@ -18,11 +18,31 @@ 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;
|
||||
if (data != nullptr) {
|
||||
for (size_t i = 0; i < size; ++i) { MixIn(hash, data[i]); }
|
||||
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;
|
||||
}
|
||||
|
||||
// Process remaining bytes
|
||||
while (data < end) {
|
||||
hash ^= static_cast<uint64_t>(*data);
|
||||
hash *= FNV_PRIME;
|
||||
data++;
|
||||
}
|
||||
|
||||
return hash;
|
||||
}
|
||||
|
||||
|
||||
@@ -14,6 +14,14 @@
|
||||
|
||||
#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;
|
||||
|
||||
@@ -5,12 +5,18 @@
|
||||
#include "AbstractMCTSAI.hpp"
|
||||
|
||||
#include <algorithm>
|
||||
#include <atomic>
|
||||
#include <chrono>
|
||||
#include <fstream>
|
||||
#include <future>
|
||||
#include <iomanip>
|
||||
#include <limits>
|
||||
#include <mutex>
|
||||
#include <random>
|
||||
#include <stdexcept>
|
||||
#include <thread>
|
||||
|
||||
#include "src/main/cpp/net/eagle0/common/mcts/util/TreeIndentUtil.hpp"
|
||||
|
||||
namespace shardok::mcts {
|
||||
|
||||
AbstractMCTSAI::AbstractMCTSAI(MCTSPlayerId playerId, MCTSConfig config)
|
||||
@@ -31,11 +37,32 @@ auto AbstractMCTSAI::Search(
|
||||
result.searchTime = std::chrono::duration_cast<std::chrono::milliseconds>(
|
||||
std::chrono::steady_clock::now() - startTime);
|
||||
|
||||
if (!rootNode || rootNode->children.empty()) {
|
||||
// Fallback to first action if no tree was built
|
||||
result.bestActionIndex = 0;
|
||||
result.bestScore = 0.0;
|
||||
return result;
|
||||
if (!rootNode) {
|
||||
throw MCTSInternalError("MCTS search: BuildMCTSTree returned null root node");
|
||||
}
|
||||
|
||||
if (rootNode->children.empty()) {
|
||||
// This can happen legitimately when:
|
||||
// 1. No legal actions available (terminal state) - return default
|
||||
// 2. Only one action and we early-exited without exploring - return index 0
|
||||
// 3. Multiple actions but none expanded - this is a bug
|
||||
if (rootNode->totalActions == 0) {
|
||||
// Terminal state - no actions available, return default result
|
||||
result.bestActionIndex = 0;
|
||||
result.bestScore = 0.0;
|
||||
result.nodesEvaluated = 0;
|
||||
return result;
|
||||
}
|
||||
// Single action case - should have been expanded in BuildMCTSTree
|
||||
if (rootNode->totalActions == 1) {
|
||||
result.bestActionIndex = 0;
|
||||
result.bestScore = 0.0;
|
||||
return result;
|
||||
}
|
||||
// Multiple actions but no children expanded - this shouldn't happen
|
||||
throw MCTSInternalError(
|
||||
"MCTS search: Root has " + std::to_string(rootNode->totalActions) +
|
||||
" actions but no children expanded - this indicates a bug in BuildMCTSTree");
|
||||
}
|
||||
|
||||
// Find best child
|
||||
@@ -44,7 +71,7 @@ auto AbstractMCTSAI::Search(
|
||||
// Use actionIndex which is the index into the filtered actions from
|
||||
// engine.getLegalActions()
|
||||
result.bestActionIndex = bestChild->actionIndex;
|
||||
result.bestScore = bestChild->averageReward;
|
||||
result.bestScore = bestChild->lookaheadScore; // Use minimax value, not poisoned average
|
||||
result.searchDepth = bestChild->depth;
|
||||
result.nodesEvaluated = rootNode->visitCount;
|
||||
|
||||
@@ -58,6 +85,9 @@ auto AbstractMCTSAI::Search(
|
||||
LogSearchResults(rootNode.get(), bestChild, result);
|
||||
}
|
||||
|
||||
// Dump tree if explicitly requested via config
|
||||
if (!config_.debugDumpPath.empty()) { DumpTreeToFile(rootNode.get(), config_.debugDumpPath); }
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
@@ -65,18 +95,73 @@ auto AbstractMCTSAI::BuildMCTSTree(
|
||||
const MCTSGameEngine& engine,
|
||||
const MCTSGameState& initialState,
|
||||
const std::chrono::steady_clock::time_point deadline) const -> std::unique_ptr<MCTSNode> {
|
||||
// Clear transposition table for this search
|
||||
// Maps state hash -> minimum depth, used to detect redundant longer paths
|
||||
transpositionTable_.clear();
|
||||
|
||||
// Create root node
|
||||
auto root = std::make_unique<MCTSNode>(initialState.clone(), playerId_, 0);
|
||||
// IMPORTANT: Use the initial state's current player, not playerId_
|
||||
// node->playerId represents "whose turn it is", not "who we're searching for"
|
||||
// This is critical for correct player flip tracking
|
||||
auto root = std::make_unique<MCTSNode>(initialState.clone(), initialState.currentPlayerId(), 0);
|
||||
|
||||
// Record root state in transposition table
|
||||
transpositionTable_[root->stateHash] = root->depth;
|
||||
|
||||
// Set whether root is maximizing based on whether current player matches who we're searching
|
||||
// for
|
||||
root->isMaximizingPlayer = (initialState.currentPlayerId() == playerId_);
|
||||
|
||||
// Get legal actions from engine for the root state
|
||||
const auto rootActions = engine.getLegalActions(initialState);
|
||||
// Root has 0 player flips
|
||||
const auto rootActions =
|
||||
engine.getLegalActions(initialState, playerId_, 0, config_.maxPlayerFlips);
|
||||
|
||||
// Initialize untried actions from the root actions
|
||||
root->untriedActionIndices.reserve(rootActions.size());
|
||||
for (size_t i = 0; i < rootActions.size(); ++i) { root->untriedActionIndices.push_back(i); }
|
||||
// Early exit if only one action available - no need to search
|
||||
if (rootActions.size() <= 1) {
|
||||
// Expand the single action so Search() can return it
|
||||
if (!rootActions.empty()) {
|
||||
root->totalActions = 1;
|
||||
[[maybe_unused]] auto* expanded = MCTSExpansion(root.get(), engine);
|
||||
}
|
||||
return root;
|
||||
}
|
||||
|
||||
// Initialize action counter
|
||||
root->totalActions = rootActions.size();
|
||||
|
||||
// CRITICAL: Do at least one expansion before entering the time-bounded loop.
|
||||
// This ensures we always have at least one child to return, even if the deadline
|
||||
// has already passed (e.g., due to debugger pause, system load, etc.)
|
||||
{
|
||||
auto* selected = MCTSSelection(root.get());
|
||||
const bool selectedIsRoot = (selected == root.get());
|
||||
const size_t childrenBeforeExpansion = root->children.size();
|
||||
|
||||
if (selected) {
|
||||
auto* expanded = MCTSExpansion(selected, engine);
|
||||
const double reward =
|
||||
MCTSSimulation(engine, *expanded->gameState, playerId_, expanded->playerFlips);
|
||||
MCTSBackpropagation(expanded, reward, config_.backpropagationPolicy);
|
||||
}
|
||||
|
||||
// Verify we actually have at least one child after the initial expansion
|
||||
if (root->children.empty()) {
|
||||
throw MCTSInternalError(
|
||||
"MCTS BuildMCTSTree: Initial expansion failed to produce any children. "
|
||||
"totalActions=" +
|
||||
std::to_string(root->totalActions) +
|
||||
", selected=" + (selected ? "non-null" : "null") +
|
||||
", selectedIsRoot=" + (selectedIsRoot ? "true" : "false") +
|
||||
", childrenBefore=" + std::to_string(childrenBeforeExpansion) +
|
||||
", childrenAfter=" + std::to_string(root->children.size()) +
|
||||
", root->CanExpand()=" + (root->CanExpand() ? "true" : "false") +
|
||||
", root->nextUntriedActionIndex=" +
|
||||
std::to_string(root->nextUntriedActionIndex));
|
||||
}
|
||||
}
|
||||
|
||||
std::atomic<int> iterations{0};
|
||||
constexpr int maxIterations = 100000;
|
||||
|
||||
if (config_.useMultithreading && config_.numThreads > 1) {
|
||||
// Multithreaded MCTS
|
||||
@@ -86,8 +171,7 @@ auto AbstractMCTSAI::BuildMCTSTree(
|
||||
futures.reserve(config_.numThreads);
|
||||
for (int threadId = 0; threadId < config_.numThreads; ++threadId) {
|
||||
futures.push_back(std::async(std::launch::async, [&] {
|
||||
while (std::chrono::steady_clock::now() < deadline &&
|
||||
iterations.load() < maxIterations) {
|
||||
while (std::chrono::steady_clock::now() < deadline) {
|
||||
// Selection and Expansion (with lock - tree modification must be serialized)
|
||||
MCTSNode* expanded;
|
||||
{
|
||||
@@ -103,10 +187,13 @@ auto AbstractMCTSAI::BuildMCTSTree(
|
||||
|
||||
// Backpropagation (with lock - modifies node statistics)
|
||||
{
|
||||
const double reward =
|
||||
MCTSSimulation(engine, *expanded->gameState, playerId_);
|
||||
const double reward = MCTSSimulation(
|
||||
engine,
|
||||
*expanded->gameState,
|
||||
playerId_,
|
||||
expanded->playerFlips);
|
||||
std::lock_guard lock(treeMutex);
|
||||
MCTSBackpropagation(expanded, reward);
|
||||
MCTSBackpropagation(expanded, reward, config_.backpropagationPolicy);
|
||||
iterations.fetch_add(1);
|
||||
}
|
||||
}
|
||||
@@ -117,19 +204,20 @@ auto AbstractMCTSAI::BuildMCTSTree(
|
||||
for (auto& future : futures) { future.wait(); }
|
||||
} else {
|
||||
// Single-threaded MCTS
|
||||
while (std::chrono::steady_clock::now() < deadline && iterations < maxIterations) {
|
||||
while (std::chrono::steady_clock::now() < deadline) {
|
||||
// Selection
|
||||
auto* selected = MCTSSelection(root.get());
|
||||
if (!selected) break;
|
||||
if (!selected) { break; }
|
||||
|
||||
// Expansion
|
||||
auto* expanded = MCTSExpansion(selected, engine);
|
||||
|
||||
// Simulation
|
||||
const double reward = MCTSSimulation(engine, *expanded->gameState, playerId_);
|
||||
const double reward =
|
||||
MCTSSimulation(engine, *expanded->gameState, playerId_, expanded->playerFlips);
|
||||
|
||||
// Backpropagation
|
||||
MCTSBackpropagation(expanded, reward);
|
||||
MCTSBackpropagation(expanded, reward, config_.backpropagationPolicy);
|
||||
|
||||
++iterations;
|
||||
|
||||
@@ -149,11 +237,27 @@ auto AbstractMCTSAI::BuildMCTSTree(
|
||||
auto AbstractMCTSAI::MCTSSelection(MCTSNode* root) const -> MCTSNode* {
|
||||
MCTSNode* current = root;
|
||||
|
||||
while (!current->isTerminal && current->depth < config_.maxTreeDepth) {
|
||||
while (current->depth < config_.maxTreeDepth) {
|
||||
// Check expansion FIRST - allows expanding "terminal" nodes that still have
|
||||
// untried actions (e.g., final round where we need to pick an action)
|
||||
if (current->CanExpand()) {
|
||||
return current; // Node has untried actions
|
||||
} else if (!current->children.empty()) {
|
||||
current = current->GetBestChild(config_.explorationConstant);
|
||||
return current; // Node has untried actions/outcomes
|
||||
}
|
||||
|
||||
// Only after expansion check: stop if terminal and fully expanded
|
||||
if (current->isTerminal) {
|
||||
break; // Terminal and no more actions to try
|
||||
}
|
||||
|
||||
if (!current->children.empty()) {
|
||||
// Choose child based on node type
|
||||
if (current->IsChanceNode()) {
|
||||
// Chance nodes: select outcome proportional to probability
|
||||
current = current->GetBestChanceChild();
|
||||
} else {
|
||||
// Decision nodes: select using UCB1
|
||||
current = current->GetBestChild(config_.explorationConstant);
|
||||
}
|
||||
if (!current) break;
|
||||
} else {
|
||||
break; // Leaf node
|
||||
@@ -165,56 +269,250 @@ auto AbstractMCTSAI::MCTSSelection(MCTSNode* root) const -> MCTSNode* {
|
||||
|
||||
auto AbstractMCTSAI::MCTSExpansion(MCTSNode* node, const MCTSGameEngine& engine) const
|
||||
-> MCTSNode* {
|
||||
if (node->untriedActionIndices.empty() || node->isTerminal) {
|
||||
// Only skip if we truly can't expand. Allow expansion even if "terminal" as long as
|
||||
// there are untried actions (e.g., final round where we need to pick an action).
|
||||
if (!node->CanExpand()) {
|
||||
return node; // Nothing to expand
|
||||
}
|
||||
|
||||
// Select a random untried action
|
||||
thread_local std::mt19937 gen(std::random_device{}());
|
||||
std::uniform_int_distribution<size_t> dis(0, node->untriedActionIndices.size() - 1);
|
||||
const size_t randomIndex = dis(gen);
|
||||
const size_t actionIndex = node->untriedActionIndices[randomIndex];
|
||||
// Handle chance node expansion (expanding outcomes)
|
||||
if (node->IsChanceNode()) {
|
||||
// Chance nodes expand their outcome children
|
||||
// This should have been set up when the chance node was created
|
||||
if (node->outcomeProbabilities.empty()) {
|
||||
throw MCTSInternalError(
|
||||
"Chance node has no outcome probabilities - this indicates a bug");
|
||||
}
|
||||
|
||||
// Remove from untried list
|
||||
node->untriedActionIndices.erase(std::next(
|
||||
node->untriedActionIndices.begin(),
|
||||
static_cast<std::vector<size_t>::difference_type>(randomIndex)));
|
||||
const size_t outcomeIndex = node->nextUntriedActionIndex++;
|
||||
if (outcomeIndex >= node->outcomeProbabilities.size()) {
|
||||
throw MCTSInternalError(
|
||||
"Chance node outcomeIndex >= outcomeProbabilities.size() - bug in expansion");
|
||||
}
|
||||
|
||||
if (node->untriedActionIndices.empty()) { node->fullyExpanded = true; }
|
||||
// The chance node's action should be the binary action
|
||||
if (!node->action) {
|
||||
throw MCTSInternalError("Chance node has no action - this indicates a bug");
|
||||
}
|
||||
|
||||
// Apply the action with the representative roll for this outcome
|
||||
// Outcome 0 = success, Outcome 1 = failure
|
||||
// Use the representative roll for this specific outcome
|
||||
const double representativeRoll = node->outcomeRolls[outcomeIndex];
|
||||
auto newState = engine.applyAction(*node->gameState, *node->action, representativeRoll);
|
||||
if (!newState) {
|
||||
throw MCTSInternalError(
|
||||
"MCTS expansion: engine.applyAction() returned nullptr for chance node "
|
||||
"outcome - this indicates a game engine error");
|
||||
}
|
||||
|
||||
// Determine if player changed
|
||||
const MCTSPlayerId newPlayerId = newState->currentPlayerId();
|
||||
const bool playerChanged = (newPlayerId != node->playerId);
|
||||
|
||||
// Calculate player flips and maximizing status
|
||||
const int newPlayerFlips = node->playerFlips + (playerChanged ? 1 : 0);
|
||||
const bool newIsMaximizing = (newPlayerId == playerId_);
|
||||
|
||||
// Create outcome child (decision node)
|
||||
auto outcomeChild = std::make_unique<MCTSNode>(
|
||||
node->action->clone(),
|
||||
std::move(newState),
|
||||
newPlayerId,
|
||||
node->depth + 1,
|
||||
outcomeIndex,
|
||||
newPlayerFlips,
|
||||
newIsMaximizing,
|
||||
node->actionWeight); // Inherit action weight from chance node
|
||||
|
||||
// Set up outcome child's actions if not terminal
|
||||
const bool shouldExpand =
|
||||
!outcomeChild->isTerminal && node->playerFlips <= config_.maxPlayerFlips;
|
||||
if (shouldExpand) {
|
||||
const auto childActions = engine.getLegalActions(
|
||||
*outcomeChild->gameState,
|
||||
playerId_,
|
||||
newPlayerFlips,
|
||||
config_.maxPlayerFlips);
|
||||
outcomeChild->totalActions = childActions.size();
|
||||
}
|
||||
|
||||
// Calculate scores
|
||||
outcomeChild->immediateScore = engine.evaluateState(*outcomeChild->gameState, playerId_);
|
||||
outcomeChild->lookaheadScore = outcomeChild->immediateScore;
|
||||
|
||||
// Set parent and add to children
|
||||
outcomeChild->parent = node;
|
||||
node->children.push_back(std::move(outcomeChild));
|
||||
|
||||
// Update chance node's immediate score to expected value of expanded outcomes
|
||||
// This corrects the initial value (which incorrectly used parent state) and ensures
|
||||
// fair UCB comparison with non-chance actions like END_TURN
|
||||
{
|
||||
double expectedImmediate = 0.0;
|
||||
double totalProbability = 0.0;
|
||||
for (size_t i = 0; i < node->children.size(); i++) {
|
||||
const double prob = node->outcomeProbabilities[i];
|
||||
const double childImmediate = node->children[i]->immediateScore;
|
||||
expectedImmediate += prob * childImmediate;
|
||||
totalProbability += prob;
|
||||
}
|
||||
// Normalize by total probability of expanded outcomes
|
||||
if (totalProbability > 0.0) {
|
||||
node->immediateScore = expectedImmediate / totalProbability;
|
||||
// CRITICAL: Always update lookaheadScore to the expected value.
|
||||
// Without this, chance nodes keep their initial lookaheadScore from the parent
|
||||
// state (before the action), while regular actions use the child state (after).
|
||||
// This gives chance nodes an unfair initial UCB advantage.
|
||||
node->lookaheadScore = node->immediateScore;
|
||||
}
|
||||
}
|
||||
|
||||
return node->children.back().get();
|
||||
}
|
||||
|
||||
// Handle decision node expansion (expanding actions)
|
||||
// Get next action to expand (sequential order)
|
||||
const size_t actionIndex = node->nextUntriedActionIndex++;
|
||||
|
||||
// Get legal actions from engine (uses cached engine for performance)
|
||||
const auto nodeActions = engine.getLegalActions(*node->gameState);
|
||||
// Use the parameterized version to respect player flips
|
||||
const auto nodeActions = engine.getLegalActions(
|
||||
*node->gameState,
|
||||
playerId_,
|
||||
node->playerFlips,
|
||||
config_.maxPlayerFlips);
|
||||
|
||||
// Create new child node
|
||||
if (actionIndex >= nodeActions.size()) {
|
||||
return node; // Invalid action index
|
||||
throw MCTSInternalError(
|
||||
"MCTS expansion: actionIndex (" + std::to_string(actionIndex) +
|
||||
") >= nodeActions.size() (" + std::to_string(nodeActions.size()) +
|
||||
") - this indicates a bug in action indexing");
|
||||
}
|
||||
|
||||
// Get action weights from engine (for prior-weighted UCB)
|
||||
const auto actionWeights = engine.getActionWeights(nodeActions, *node->gameState);
|
||||
|
||||
const auto& action = nodeActions[actionIndex];
|
||||
|
||||
const double actionWeight =
|
||||
actionIndex < actionWeights.size() ? actionWeights[actionIndex] : 1.0;
|
||||
|
||||
// Check if this action requires a chance node
|
||||
if (action->requiresChanceNode()) {
|
||||
// Create intermediate chance node
|
||||
auto chanceNode = std::make_unique<MCTSNode>(
|
||||
action->clone(),
|
||||
node->gameState->clone(), // Chance node has same state as parent
|
||||
node->playerId,
|
||||
node->depth + 1,
|
||||
actionIndex,
|
||||
node->playerFlips,
|
||||
node->isMaximizingPlayer,
|
||||
actionWeight);
|
||||
|
||||
chanceNode->nodeType = NodeType::CHANCE;
|
||||
|
||||
// Get outcome information from engine
|
||||
const auto outcomeInfo = engine.getBinaryOutcomeInfo(*node->gameState, *action);
|
||||
|
||||
// Set up outcome metadata (2 outcomes for binary actions)
|
||||
chanceNode->outcomeProbabilities = outcomeInfo.getProbabilities();
|
||||
chanceNode->outcomeRolls = outcomeInfo.getRepresentativeRolls();
|
||||
chanceNode->totalActions = 2; // Binary: success and failure
|
||||
|
||||
// Chance node immediate score will be computed as expected value during backpropagation
|
||||
// For now, initialize to parent's score as a reasonable default
|
||||
chanceNode->immediateScore = engine.evaluateState(*node->gameState, playerId_);
|
||||
chanceNode->lookaheadScore = chanceNode->immediateScore;
|
||||
|
||||
// Set parent and add to children
|
||||
chanceNode->parent = node;
|
||||
node->children.push_back(std::move(chanceNode));
|
||||
|
||||
// CRITICAL: Immediately expand the first outcome and return that instead.
|
||||
// If we returned the chance node itself, MCTSSimulation would run on the parent state
|
||||
// (since chance nodes have parent's gameState), which is wrong. We need to simulate
|
||||
// from an actual outcome state.
|
||||
//
|
||||
// Note: This recursion is bounded because outcome children are decision nodes,
|
||||
// not chance nodes, so the recursion goes exactly one level deep.
|
||||
return MCTSExpansion(node->children.back().get(), engine);
|
||||
}
|
||||
|
||||
// Regular (non-chance) action: create decision node directly
|
||||
auto newState = engine.applyAction(*node->gameState, *action);
|
||||
if (!newState) {
|
||||
return node; // Failed to apply action
|
||||
throw MCTSInternalError(
|
||||
"MCTS expansion: engine.applyAction() returned nullptr for action " +
|
||||
action->getDescription() + " - this indicates a game engine error");
|
||||
}
|
||||
|
||||
// Determine if player changed
|
||||
const MCTSPlayerId newPlayerId = newState->currentPlayerId();
|
||||
const bool playerChanged = (newPlayerId != node->playerId);
|
||||
|
||||
// Calculate player flips and maximizing status
|
||||
const int newPlayerFlips = node->playerFlips + (playerChanged ? 1 : 0);
|
||||
// Node is maximizing if current player is the root player (playerId_)
|
||||
const bool newIsMaximizing = (newPlayerId == playerId_);
|
||||
|
||||
auto child = std::make_unique<MCTSNode>(
|
||||
action->clone(),
|
||||
std::move(newState),
|
||||
node->gameState->currentPlayerId(),
|
||||
newPlayerId,
|
||||
node->depth + 1,
|
||||
actionIndex);
|
||||
actionIndex,
|
||||
newPlayerFlips,
|
||||
newIsMaximizing,
|
||||
actionWeight); // Pass the action weight for prior-weighted UCB
|
||||
|
||||
// Set up child's untried actions if not terminal
|
||||
if (!child->isTerminal) {
|
||||
const auto childActions = engine.getLegalActions(*child->gameState);
|
||||
child->untriedActionIndices.reserve(childActions.size());
|
||||
for (size_t i = 0; i < childActions.size(); ++i) {
|
||||
child->untriedActionIndices.push_back(i);
|
||||
// Check transposition table: mark as redundant if we've reached this state at a shallower depth
|
||||
// This prevents MCTS from exploring longer paths to the same game state
|
||||
// Works best with MINIMAX backpropagation (penalty propagates as min/max)
|
||||
// Also provides benefit with AVERAGING (penalty pulls average down significantly)
|
||||
const uint64_t childHash = child->stateHash;
|
||||
auto it = transpositionTable_.find(childHash);
|
||||
if (it != transpositionTable_.end()) {
|
||||
const int previousDepth = it->second;
|
||||
if (child->depth > previousDepth) {
|
||||
// Longer path to same state - mark as redundant and heavily penalize
|
||||
// Use -infinity to be unambiguously worse than any legitimate score
|
||||
child->isRedundant = true;
|
||||
child->immediateScore = -std::numeric_limits<double>::infinity();
|
||||
child->lookaheadScore = -std::numeric_limits<double>::infinity();
|
||||
} else {
|
||||
// Found shorter or equal path - update table
|
||||
transpositionTable_[childHash] = child->depth;
|
||||
}
|
||||
} else {
|
||||
// First time seeing this state - record it
|
||||
transpositionTable_[childHash] = child->depth;
|
||||
}
|
||||
|
||||
// Calculate immediate and lookahead scores
|
||||
child->immediateScore = engine.evaluateState(*child->gameState, playerId_);
|
||||
child->lookaheadScore = child->immediateScore;
|
||||
// Set up child's untried actions if not terminal and parent hasn't exceeded player flips
|
||||
// playerFlips counts how many times the player has CHANGED from root
|
||||
// We expand children of nodes that are within the maxPlayerFlips limit
|
||||
// maxPlayerFlips=0: same player can take multiple sequential actions
|
||||
// maxPlayerFlips=1: can explore opponent's immediate responses
|
||||
const bool shouldExpand = !child->isTerminal && node->playerFlips <= config_.maxPlayerFlips;
|
||||
|
||||
if (shouldExpand) {
|
||||
const auto childActions = engine.getLegalActions(
|
||||
*child->gameState,
|
||||
playerId_,
|
||||
newPlayerFlips,
|
||||
config_.maxPlayerFlips);
|
||||
child->totalActions = childActions.size();
|
||||
}
|
||||
|
||||
// Calculate immediate and lookahead scores from root player's perspective
|
||||
// Skip for redundant nodes (already have penalty scores)
|
||||
if (!child->isRedundant) {
|
||||
child->immediateScore = engine.evaluateState(*child->gameState, playerId_);
|
||||
child->lookaheadScore = child->immediateScore;
|
||||
}
|
||||
|
||||
// Set parent and add to children
|
||||
child->parent = node;
|
||||
@@ -226,20 +524,49 @@ auto AbstractMCTSAI::MCTSExpansion(MCTSNode* node, const MCTSGameEngine& engine)
|
||||
auto AbstractMCTSAI::MCTSSimulation(
|
||||
const MCTSGameEngine& engine,
|
||||
const MCTSGameState& state,
|
||||
const MCTSPlayerId startingPlayer) const -> double {
|
||||
const MCTSPlayerId startingPlayer,
|
||||
const int startingPlayerFlips) const -> double {
|
||||
if (state.isTerminal()) { return state.score(startingPlayer); }
|
||||
|
||||
// If we've already exceeded the simulation horizon, don't simulate - just return immediate
|
||||
// score This ensures fair comparison: all leaves are evaluated at the same game phase Example:
|
||||
// maxSimulationFlips=1 means simulate THROUGH opponent's first response (i.e., allow one action
|
||||
// at playerFlips=1, then stop)
|
||||
if (startingPlayerFlips > config_.maxSimulationFlips) { return state.score(startingPlayer); }
|
||||
|
||||
// Create a mutable copy for simulation
|
||||
auto currentState = state.clone();
|
||||
int depth = 0;
|
||||
int playerFlips = startingPlayerFlips; // Start from the expanded node's flip count
|
||||
MCTSPlayerId previousPlayer = currentState->currentPlayerId();
|
||||
|
||||
// Simulate until terminal or max depth
|
||||
while (!currentState->isTerminal() && depth < config_.maxSimulationDepth) {
|
||||
const auto actions = engine.getLegalActions(*currentState);
|
||||
// Simulate until we exceed the horizon, hit terminal state, or max depth
|
||||
// Note: We allow one action AT maxSimulationFlips before stopping
|
||||
while (!currentState->isTerminal() && depth < config_.maxSimulationDepth &&
|
||||
playerFlips <= config_.maxSimulationFlips) {
|
||||
// Track player changes
|
||||
const MCTSPlayerId currentPlayer = currentState->currentPlayerId();
|
||||
if (currentPlayer != previousPlayer) {
|
||||
playerFlips++;
|
||||
previousPlayer = currentPlayer;
|
||||
}
|
||||
|
||||
// Get legal actions with player flip tracking
|
||||
const auto actions = engine.getLegalActions(
|
||||
*currentState,
|
||||
playerId_,
|
||||
playerFlips,
|
||||
config_.maxSimulationFlips);
|
||||
if (actions.empty()) { break; }
|
||||
|
||||
// Determine if current player is maximizing or minimizing
|
||||
// Maximizing: current player is root player (trying to maximize root player's score)
|
||||
// Minimizing: current player is opponent (trying to minimize root player's score)
|
||||
const bool isMaximizing = (currentPlayer == playerId_);
|
||||
|
||||
// Select action based on simulation policy
|
||||
const size_t selectedIndex = SelectSimulationAction(engine, *currentState, actions);
|
||||
const size_t selectedIndex =
|
||||
SelectSimulationAction(engine, *currentState, actions, isMaximizing);
|
||||
if (selectedIndex >= actions.size()) { break; }
|
||||
|
||||
// Apply action
|
||||
@@ -253,18 +580,114 @@ auto AbstractMCTSAI::MCTSSimulation(
|
||||
return currentState->score(startingPlayer);
|
||||
}
|
||||
|
||||
auto AbstractMCTSAI::MCTSBackpropagation(MCTSNode* node, const double reward) -> void {
|
||||
auto AbstractMCTSAI::MCTSBackpropagation(
|
||||
MCTSNode* node,
|
||||
const double reward,
|
||||
const MCTSBackpropagationPolicy policy) const -> void {
|
||||
// Backpropagation strategy is configured via MCTSConfig:
|
||||
// - AVERAGING: Traditional MCTS averaging (for stochastic/single-player games)
|
||||
// - MINIMAX: Minimax backup (for deterministic adversarial games)
|
||||
|
||||
const bool useMinimaxBackup = (policy == MCTSBackpropagationPolicy::MINIMAX);
|
||||
|
||||
while (node) {
|
||||
node->visitCount++;
|
||||
|
||||
// Always track average for UCB
|
||||
node->totalReward += reward;
|
||||
node->averageReward = node->totalReward / node->visitCount;
|
||||
|
||||
// Update lookahead score as weighted average
|
||||
if (node->visitCount == 1) {
|
||||
node->lookaheadScore = reward;
|
||||
// Update lookahead score based on node type and strategy
|
||||
if (node->IsChanceNode() && !node->children.empty()) {
|
||||
// Chance nodes: compute expected value (weighted average of outcomes)
|
||||
// lookaheadScore = sum(probability[i] * childValue[i])
|
||||
double expectedValue = 0.0;
|
||||
double totalProbability = 0.0;
|
||||
int visitedChildCount = 0;
|
||||
|
||||
for (size_t i = 0; i < node->children.size(); i++) {
|
||||
const auto& child = node->children[i];
|
||||
if (child->visitCount == 0) continue; // Unvisited outcomes don't contribute
|
||||
|
||||
const double probability = node->outcomeProbabilities[i];
|
||||
const double childValue = child->lookaheadScore;
|
||||
expectedValue += probability * childValue;
|
||||
totalProbability += probability;
|
||||
visitedChildCount++;
|
||||
}
|
||||
|
||||
// Use expected value if we have visited outcomes, else use average
|
||||
if (visitedChildCount > 0) {
|
||||
// CRITICAL: Normalize by total probability to get correct expected value
|
||||
// when not all outcomes have been visited yet
|
||||
if (totalProbability > 0.0 && totalProbability < 1.0) {
|
||||
// Normalize to account for unvisited outcomes
|
||||
// This gives the correct expected value among visited outcomes
|
||||
expectedValue /= totalProbability;
|
||||
}
|
||||
node->lookaheadScore = expectedValue;
|
||||
} else {
|
||||
// No outcomes visited yet, fall back to average
|
||||
if (node->visitCount == 1) {
|
||||
node->lookaheadScore = reward;
|
||||
} else {
|
||||
const double alpha = 1.0 / node->visitCount;
|
||||
node->lookaheadScore = (1.0 - alpha) * node->lookaheadScore + alpha * reward;
|
||||
}
|
||||
}
|
||||
} else if (useMinimaxBackup && !node->children.empty()) {
|
||||
// Minimax backup: use best/worst child value for adversarial games
|
||||
// The operation (MAX or MIN) depends on whose turn it is at THIS node
|
||||
// - If this node is root player's turn: root chooses MAX (best for root)
|
||||
// - If this node is opponent's turn: opponent chooses MIN (best for opponent = worst
|
||||
// for root)
|
||||
//
|
||||
// Note: In setup phase, children can have different isMaximizingPlayer values:
|
||||
// - PLACE_UNIT keeps same player's turn
|
||||
// - END_PLAYER_SETUP flips to opponent's turn
|
||||
// So we must use the PARENT node's isMaximizingPlayer, not the child's.
|
||||
|
||||
const bool thisNodeIsRootPlayer = node->isMaximizingPlayer;
|
||||
|
||||
double minmaxValue = thisNodeIsRootPlayer ? -std::numeric_limits<double>::max()
|
||||
: std::numeric_limits<double>::max();
|
||||
|
||||
int visitedChildCount = 0;
|
||||
for (const auto& child : node->children) {
|
||||
if (child->visitCount == 0) continue; // Unvisited children don't contribute
|
||||
|
||||
const double childValue = child->lookaheadScore;
|
||||
visitedChildCount++;
|
||||
|
||||
if (thisNodeIsRootPlayer) {
|
||||
// Root player chooses: take MAX (best for root)
|
||||
minmaxValue = std::max(minmaxValue, childValue);
|
||||
} else {
|
||||
// Opponent chooses: take MIN (best for opponent = worst for root)
|
||||
minmaxValue = std::min(minmaxValue, childValue);
|
||||
}
|
||||
}
|
||||
|
||||
// Use minimax value if we found any visited children, else use average
|
||||
if (visitedChildCount > 0) {
|
||||
node->lookaheadScore = minmaxValue;
|
||||
} else {
|
||||
// No children visited yet, fall back to average
|
||||
if (node->visitCount == 1) {
|
||||
node->lookaheadScore = reward;
|
||||
} else {
|
||||
const double alpha = 1.0 / node->visitCount;
|
||||
node->lookaheadScore = (1.0 - alpha) * node->lookaheadScore + alpha * reward;
|
||||
}
|
||||
}
|
||||
} else {
|
||||
const double alpha = 1.0 / node->visitCount;
|
||||
node->lookaheadScore = (1.0 - alpha) * node->lookaheadScore + alpha * reward;
|
||||
// Standard MCTS averaging (for maxPlayerFlips=0 or leaf nodes)
|
||||
if (node->visitCount == 1) {
|
||||
node->lookaheadScore = reward;
|
||||
} else {
|
||||
const double alpha = 1.0 / node->visitCount;
|
||||
node->lookaheadScore = (1.0 - alpha) * node->lookaheadScore + alpha * reward;
|
||||
}
|
||||
}
|
||||
|
||||
node = node->parent;
|
||||
@@ -274,8 +697,13 @@ auto AbstractMCTSAI::MCTSBackpropagation(MCTSNode* node, const double reward) ->
|
||||
auto AbstractMCTSAI::SelectSimulationAction(
|
||||
const MCTSGameEngine& engine,
|
||||
const MCTSGameState& state,
|
||||
const std::vector<std::unique_ptr<MCTSAction>>& actions) const -> size_t {
|
||||
if (actions.empty()) { return 0; }
|
||||
const std::vector<std::unique_ptr<MCTSAction>>& actions,
|
||||
const bool isMaximizing) const -> size_t {
|
||||
if (actions.empty()) {
|
||||
throw MCTSInternalError(
|
||||
"MCTSSimulation called with empty actions list - this indicates a bug in the "
|
||||
"MCTS tree building or game state");
|
||||
}
|
||||
|
||||
thread_local std::mt19937 gen(std::random_device{}());
|
||||
|
||||
@@ -297,13 +725,31 @@ auto AbstractMCTSAI::SelectSimulationAction(
|
||||
}
|
||||
|
||||
case MCTSSimulationPolicy::BEST_IMMEDIATE: {
|
||||
double bestScore = -std::numeric_limits<double>::max();
|
||||
size_t bestIndex = 0;
|
||||
// For adversarial search:
|
||||
// - Maximizing nodes select action with HIGHEST score (best for root player)
|
||||
// - Minimizing nodes select action with LOWEST score (worst for root player)
|
||||
|
||||
for (size_t i = 0; i < actions.size(); ++i) {
|
||||
const double score =
|
||||
engine.getActionScore(state, *actions[i], state.currentPlayerId());
|
||||
if (score > bestScore) {
|
||||
// First, filter out obviously bad moves (e.g., BECOME_OUTLAW)
|
||||
const auto filteredIndices = engine.filterActions(actions, state);
|
||||
if (filteredIndices.empty()) {
|
||||
// If all actions filtered out, fall back to first action
|
||||
return 0;
|
||||
}
|
||||
|
||||
double bestScore = isMaximizing ? -std::numeric_limits<double>::max()
|
||||
: std::numeric_limits<double>::max();
|
||||
size_t bestIndex = filteredIndices[0];
|
||||
|
||||
for (const size_t i : filteredIndices) {
|
||||
// CRITICAL: Always get score from ROOT player's perspective for adversarial search
|
||||
// If we use currentPlayerId, opponent actions would be scored from their
|
||||
// perspective, causing them to select moves that help themselves instead of hurt
|
||||
// us!
|
||||
const double score = engine.getActionScore(state, *actions[i], playerId_);
|
||||
|
||||
const bool shouldSelect = isMaximizing ? (score > bestScore) : (score < bestScore);
|
||||
|
||||
if (shouldSelect) {
|
||||
bestScore = score;
|
||||
bestIndex = i;
|
||||
}
|
||||
@@ -317,15 +763,23 @@ auto AbstractMCTSAI::SelectSimulationAction(
|
||||
scores.reserve(actions.size());
|
||||
|
||||
for (size_t i = 0; i < actions.size(); ++i) {
|
||||
const double score =
|
||||
engine.getActionScore(state, *actions[i], state.currentPlayerId());
|
||||
// CRITICAL: Always get score from ROOT player's perspective for adversarial search
|
||||
const double score = engine.getActionScore(state, *actions[i], playerId_);
|
||||
scores.emplace_back(i, score);
|
||||
}
|
||||
|
||||
// Sort by score
|
||||
std::ranges::sort(scores, [](const auto& a, const auto& b) {
|
||||
return a.second > b.second;
|
||||
});
|
||||
// - Maximizing: highest scores first (prefer actions that maximize root player's score)
|
||||
// - Minimizing: lowest scores first (prefer actions that minimize root player's score)
|
||||
if (isMaximizing) {
|
||||
std::ranges::sort(scores, [](const auto& a, const auto& b) {
|
||||
return a.second > b.second; // Descending
|
||||
});
|
||||
} else {
|
||||
std::ranges::sort(scores, [](const auto& a, const auto& b) {
|
||||
return a.second < b.second; // Ascending
|
||||
});
|
||||
}
|
||||
|
||||
// Create weights based on ranking
|
||||
std::vector<double> weights;
|
||||
@@ -338,6 +792,36 @@ auto AbstractMCTSAI::SelectSimulationAction(
|
||||
std::discrete_distribution<> dis(weights.begin(), weights.end());
|
||||
return scores[dis(gen)].first;
|
||||
}
|
||||
|
||||
case MCTSSimulationPolicy::WEIGHTED_HEURISTIC: {
|
||||
// Get heuristic weights from engine (fast O(1) per action)
|
||||
const auto weights = engine.getActionWeights(actions, state);
|
||||
|
||||
// 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: 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 using discrete_distribution
|
||||
std::discrete_distribution<> dis(validWeights.begin(), validWeights.end());
|
||||
return validIndices[dis(gen)];
|
||||
}
|
||||
}
|
||||
|
||||
// Default to random
|
||||
@@ -393,14 +877,19 @@ auto AbstractMCTSAI::LogSearchResults(
|
||||
return a->visitCount > b->visitCount;
|
||||
});
|
||||
|
||||
printf("MCTS: Top actions by visits:\n");
|
||||
for (size_t i = 0; i < std::min(static_cast<size_t>(3), sortedChildren.size()); ++i) {
|
||||
// Show all actions if there are <= 10, otherwise top 5
|
||||
const size_t numToShow = sortedChildren.size() <= 10 ? sortedChildren.size() : 5;
|
||||
printf("MCTS: Top %zu actions by visits (out of %zu total):\n",
|
||||
numToShow,
|
||||
sortedChildren.size());
|
||||
for (size_t i = 0; i < numToShow; ++i) {
|
||||
const auto* child = sortedChildren[i];
|
||||
printf(" [%zu] visits:%d immediate:%.2f backprop:%.2f",
|
||||
printf(" [%zu] visits:%d avgReward:%.2f immediate:%.2f lookahead:%.2f",
|
||||
i,
|
||||
child->visitCount,
|
||||
child->averageReward,
|
||||
child->immediateScore,
|
||||
child->averageReward);
|
||||
child->lookaheadScore);
|
||||
|
||||
// Show the action's own description
|
||||
if (child->action) { printf(" %s", child->action->getDescription().c_str()); }
|
||||
@@ -452,14 +941,165 @@ auto AbstractMCTSAI::LogSearchResults(
|
||||
|
||||
if (!bestSequence.empty()) {
|
||||
printf("MCTS: Best sequence from chosen action (final: %.2f):\n", sequenceScore);
|
||||
int displayedStep = 0;
|
||||
for (size_t i = 0; i < bestSequence.size(); ++i) {
|
||||
const auto* node = bestSequence[i];
|
||||
printf(" %zu.", i + 1);
|
||||
|
||||
// Skip outcome nodes (children of chance nodes) - they're displayed with their parent
|
||||
if (i > 0 && node->parent && node->parent->IsChanceNode()) { continue; }
|
||||
|
||||
displayedStep++;
|
||||
printf(" %d.", displayedStep);
|
||||
if (node->action) { printf(" %s", node->action->getDescription().c_str()); }
|
||||
printf(" (visits:%d, immediate:%.2f, backprop:%.2f)\n",
|
||||
|
||||
// If this is a chance node, display outcome probabilities and scores
|
||||
if (node->IsChanceNode() && !node->outcomeProbabilities.empty()) {
|
||||
printf(" [");
|
||||
for (size_t j = 0; j < node->outcomeProbabilities.size(); ++j) {
|
||||
if (j > 0) printf(", ");
|
||||
const double prob = node->outcomeProbabilities[j] * 100;
|
||||
// Show lookahead score for each outcome if child exists
|
||||
if (j < node->children.size() && node->children[j]->visitCount > 0) {
|
||||
printf("%.0f%%->%.1f", prob, node->children[j]->lookaheadScore);
|
||||
} else {
|
||||
printf("%.0f%%->?", prob);
|
||||
}
|
||||
}
|
||||
printf("]");
|
||||
}
|
||||
|
||||
printf(" (visits:%d, immediate:%.2f, lookahead:%.2f)\n",
|
||||
node->visitCount,
|
||||
node->immediateScore,
|
||||
node->averageReward);
|
||||
node->lookaheadScore);
|
||||
|
||||
// For non-root nodes in the sequence, show what the top alternatives were
|
||||
// Skip showing alternatives for chance nodes (they have outcome children, not action
|
||||
// alternatives)
|
||||
if (i > 0 && node->parent && !node->parent->children.empty() &&
|
||||
!node->parent->IsChanceNode()) {
|
||||
// Collect all siblings (including this node) and sort by visit count
|
||||
std::vector<const MCTSNode*> siblings;
|
||||
siblings.reserve(node->parent->children.size());
|
||||
for (const auto& child : node->parent->children) {
|
||||
if (!child->isRedundant) { siblings.push_back(child.get()); }
|
||||
}
|
||||
|
||||
// Sort by visit count (descending)
|
||||
std::ranges::sort(siblings, [](const MCTSNode* a, const MCTSNode* b) {
|
||||
return a->visitCount > b->visitCount;
|
||||
});
|
||||
|
||||
// Show top 3 alternatives at this decision point
|
||||
printf(" Alternatives at this node (%zu total):\n", siblings.size());
|
||||
const size_t topN = std::min(siblings.size(), size_t(3));
|
||||
for (size_t j = 0; j < topN; ++j) {
|
||||
const auto* alt = siblings[j];
|
||||
printf(" [%zu] visits:%d immediate:%.2f lookahead:%.2f",
|
||||
j,
|
||||
alt->visitCount,
|
||||
alt->immediateScore,
|
||||
alt->lookaheadScore);
|
||||
if (alt->action) { printf(" %s", alt->action->getDescription().c_str()); }
|
||||
printf("\n");
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
auto AbstractMCTSAI::DumpTreeToFile(const MCTSNode* root, const std::string& filepath) -> void {
|
||||
if (!root) return;
|
||||
|
||||
std::ofstream out(filepath);
|
||||
if (!out) {
|
||||
fprintf(stderr, "Failed to open dump file: %s\n", filepath.c_str());
|
||||
return;
|
||||
}
|
||||
|
||||
out << "MCTS Tree Dump\n";
|
||||
out << "==============\n\n";
|
||||
out << "Root Node:\n";
|
||||
out << " Visits: " << root->visitCount << "\n";
|
||||
out << " Immediate Score: " << root->immediateScore << "\n";
|
||||
out << " Lookahead Score: " << root->lookaheadScore << "\n";
|
||||
out << " Average Reward: " << root->averageReward << "\n";
|
||||
out << " Player ID: " << root->playerId << "\n";
|
||||
out << " Depth: " << root->depth << "\n";
|
||||
out << " Is Maximizing: " << (root->isMaximizingPlayer ? "true" : "false") << "\n";
|
||||
out << " State Hash: " << std::hex << root->stateHash << std::dec << "\n";
|
||||
out << "\n";
|
||||
|
||||
if (!root->children.empty()) {
|
||||
out << "Children:\n";
|
||||
for (size_t i = 0; i < root->children.size(); ++i) {
|
||||
const auto& child = root->children[i];
|
||||
const bool isLast = (i == root->children.size() - 1);
|
||||
DumpNodeRecursive(child.get(), out, 1, isLast);
|
||||
}
|
||||
}
|
||||
|
||||
out << "\n=== End of Tree Dump ===\n";
|
||||
out.close();
|
||||
|
||||
printf("MCTS: Tree dumped to %s\n", filepath.c_str());
|
||||
}
|
||||
|
||||
auto AbstractMCTSAI::DumpNodeRecursive(
|
||||
const MCTSNode* node,
|
||||
std::ostream& out,
|
||||
const int indentLevel,
|
||||
const bool isLastChild) -> void {
|
||||
if (!node) return;
|
||||
|
||||
// Create indent string
|
||||
const std::string indent = ::mcts::util::BuildTreeIndent(indentLevel, isLastChild);
|
||||
|
||||
// Write node information
|
||||
out << indent;
|
||||
|
||||
// Show node type for chance nodes
|
||||
if (node->IsChanceNode()) { out << "[CHANCE] "; }
|
||||
|
||||
if (node->action) {
|
||||
out << node->action->getDescription();
|
||||
} else {
|
||||
out << "[ROOT]";
|
||||
}
|
||||
out << " (visits:" << node->visitCount;
|
||||
out << ", immediate:" << std::fixed << std::setprecision(2) << node->immediateScore;
|
||||
out << ", lookahead:" << node->lookaheadScore;
|
||||
out << ", avgReward:" << node->averageReward;
|
||||
out << ", weight:" << node->actionWeight;
|
||||
out << ", depth:" << node->depth;
|
||||
out << ", flips:" << node->playerFlips;
|
||||
out << ", player:" << node->playerId;
|
||||
out << ", max:" << (node->isMaximizingPlayer ? "T" : "F");
|
||||
if (node->isRedundant) { out << ", REDUNDANT"; }
|
||||
if (node->isTerminal) { out << ", TERMINAL"; }
|
||||
out << ")\n";
|
||||
|
||||
// Show outcome probabilities and rolls for chance nodes
|
||||
if (node->IsChanceNode() && !node->outcomeProbabilities.empty()) {
|
||||
const std::string outcomeIndent = ::mcts::util::ConvertBranchToContinuation(indent);
|
||||
out << outcomeIndent << " Outcomes: ";
|
||||
for (size_t i = 0; i < node->outcomeProbabilities.size(); ++i) {
|
||||
if (i > 0) out << ", ";
|
||||
out << "[" << i << "] p=" << std::fixed << std::setprecision(3)
|
||||
<< node->outcomeProbabilities[i];
|
||||
if (i < node->outcomeRolls.size()) {
|
||||
out << " roll=" << std::fixed << std::setprecision(1) << node->outcomeRolls[i];
|
||||
}
|
||||
}
|
||||
out << "\n";
|
||||
}
|
||||
|
||||
// Recursively dump children
|
||||
if (!node->children.empty()) {
|
||||
for (size_t i = 0; i < node->children.size(); ++i) {
|
||||
const auto& child = node->children[i];
|
||||
const bool isLast = (i == node->children.size() - 1);
|
||||
DumpNodeRecursive(child.get(), out, indentLevel + 1, isLast);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,6 +7,7 @@
|
||||
|
||||
#include <chrono>
|
||||
#include <memory>
|
||||
#include <unordered_map>
|
||||
#include <vector>
|
||||
|
||||
#include "MCTSAction.hpp"
|
||||
@@ -51,6 +52,11 @@ 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,
|
||||
@@ -66,21 +72,32 @@ private:
|
||||
[[nodiscard]] auto MCTSSimulation(
|
||||
const MCTSGameEngine& engine,
|
||||
const MCTSGameState& state,
|
||||
MCTSPlayerId startingPlayer) const -> double;
|
||||
MCTSPlayerId startingPlayer,
|
||||
int startingPlayerFlips = 0) const -> double;
|
||||
|
||||
static auto MCTSBackpropagation(MCTSNode* node, double reward) -> void;
|
||||
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) const -> size_t;
|
||||
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
|
||||
|
||||
@@ -86,6 +86,7 @@ cc_library(
|
||||
":mcts_game_state",
|
||||
":mcts_node",
|
||||
":mcts_types",
|
||||
"//src/main/cpp/net/eagle0/common/mcts/util:tree_indent_util",
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
@@ -27,6 +27,11 @@ public:
|
||||
|
||||
// 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
|
||||
|
||||
@@ -9,6 +9,8 @@
|
||||
#include <random>
|
||||
#include <vector>
|
||||
|
||||
#include "MCTSTypes.hpp" // For MCTSInternalError
|
||||
|
||||
namespace shardok {
|
||||
namespace mcts {
|
||||
|
||||
@@ -26,7 +28,7 @@ double MCTSGameEngine::simulateRandomPlayout(
|
||||
|
||||
// Simulate until terminal or max depth
|
||||
while (!currentState->isTerminal() && depth < maxDepth) {
|
||||
auto actions = getLegalActions(*currentState);
|
||||
auto actions = getLegalActions(*currentState, playerId, 0, 0);
|
||||
if (actions.empty()) { break; }
|
||||
|
||||
size_t selectedIndex = 0;
|
||||
@@ -95,6 +97,38 @@ double MCTSGameEngine::simulateRandomPlayout(
|
||||
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
|
||||
|
||||
@@ -16,15 +16,50 @@
|
||||
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) const = 0;
|
||||
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
|
||||
@@ -34,9 +69,13 @@ public:
|
||||
state = applyAction(*state, action);
|
||||
}
|
||||
|
||||
// Get all legal actions for the current state
|
||||
// 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) const = 0;
|
||||
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;
|
||||
@@ -57,6 +96,17 @@ public:
|
||||
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(
|
||||
@@ -96,6 +146,13 @@ public:
|
||||
(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
|
||||
|
||||
@@ -17,8 +17,16 @@
|
||||
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)
|
||||
@@ -35,17 +43,24 @@ struct MCTSNode {
|
||||
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;
|
||||
std::vector<size_t> untriedActionIndices;
|
||||
bool fullyExpanded = false;
|
||||
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;
|
||||
@@ -55,7 +70,9 @@ struct MCTSNode {
|
||||
MCTSNode(std::unique_ptr<MCTSGameState> state, MCTSPlayerId pid, int d)
|
||||
: gameState(std::move(state)),
|
||||
playerId(pid),
|
||||
depth(d) {
|
||||
depth(d),
|
||||
playerFlips(0),
|
||||
isMaximizingPlayer(true) {
|
||||
if (gameState) {
|
||||
stateHash = gameState->hash();
|
||||
isTerminal = gameState->isTerminal();
|
||||
@@ -68,12 +85,18 @@ struct MCTSNode {
|
||||
std::unique_ptr<MCTSGameState> state,
|
||||
MCTSPlayerId pid,
|
||||
int d,
|
||||
size_t actIdx = SIZE_MAX)
|
||||
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) {
|
||||
depth(d),
|
||||
playerFlips(flips),
|
||||
isMaximizingPlayer(isMaximizing) {
|
||||
if (gameState) {
|
||||
stateHash = gameState->hash();
|
||||
isTerminal = gameState->isTerminal();
|
||||
@@ -100,20 +123,64 @@ struct MCTSNode {
|
||||
}
|
||||
}
|
||||
|
||||
// Calculate UCB1 value for this node
|
||||
void CalculateUCB1(const double explorationConstant) const {
|
||||
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;
|
||||
}
|
||||
// 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 !fullyExpanded && !untriedActionIndices.empty(); }
|
||||
[[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 {
|
||||
@@ -126,10 +193,28 @@ struct MCTSNode {
|
||||
// Skip redundant nodes
|
||||
if (child->isRedundant) continue;
|
||||
|
||||
child->CalculateUCB1(explorationConstant);
|
||||
// Calculate UCB1 value using the helper function
|
||||
const double value =
|
||||
child->CalculateUCB1(explorationConstant, visitCount, isMaximizingPlayer);
|
||||
|
||||
if (child->ucb1Value > bestValue) {
|
||||
bestValue = child->ucb1Value;
|
||||
// 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();
|
||||
}
|
||||
}
|
||||
@@ -143,7 +228,8 @@ struct MCTSNode {
|
||||
|
||||
MCTSNode* bestChild = nullptr;
|
||||
int bestVisits = 0;
|
||||
double bestScore = -std::numeric_limits<double>::max();
|
||||
double bestScore = isMaximizingPlayer ? -std::numeric_limits<double>::max()
|
||||
: std::numeric_limits<double>::max();
|
||||
|
||||
for (const auto& child : children) {
|
||||
// Skip redundant nodes
|
||||
@@ -152,12 +238,18 @@ struct MCTSNode {
|
||||
// Prefer most-visited node (robust child selection)
|
||||
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;
|
||||
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();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -166,7 +258,9 @@ struct MCTSNode {
|
||||
for (const auto& child : children) {
|
||||
if (child->isRedundant) continue;
|
||||
|
||||
if (child->lookaheadScore > bestScore) {
|
||||
const bool shouldReplace = isMaximizingPlayer ? (child->lookaheadScore > bestScore)
|
||||
: (child->lookaheadScore < bestScore);
|
||||
if (shouldReplace) {
|
||||
bestScore = child->lookaheadScore;
|
||||
bestChild = child.get();
|
||||
}
|
||||
|
||||
@@ -5,18 +5,34 @@
|
||||
#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
|
||||
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
|
||||
@@ -27,6 +43,15 @@ struct MCTSConfig {
|
||||
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
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
load("@rules_cc//cc:defs.bzl", "cc_library")
|
||||
|
||||
cc_library(
|
||||
name = "tree_indent_util",
|
||||
srcs = ["TreeIndentUtil.cpp"],
|
||||
hdrs = ["TreeIndentUtil.hpp"],
|
||||
visibility = ["//visibility:public"],
|
||||
)
|
||||
@@ -0,0 +1,53 @@
|
||||
//
|
||||
// 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
|
||||
@@ -0,0 +1,22 @@
|
||||
//
|
||||
// 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
|
||||
@@ -27,7 +27,7 @@ auto AIAttackerStrategySelector::BestAttackerStrategy(
|
||||
const BattalionTypeGetter& battalionTypeGetter,
|
||||
ActionPoints braveWaterCost,
|
||||
const AIWaterCrossingCommandChooser& waterCrossingCommandChooser,
|
||||
const vector<CommandProto>& /*availableCommands*/) -> AIStrategy {
|
||||
const CommandListSPtr& /*availableCommands*/) -> AIStrategy {
|
||||
uint32_t attackerUnitCount = 0;
|
||||
int defenderOccupiedCriticalTileCount = 0;
|
||||
bool canFlee = false;
|
||||
|
||||
@@ -10,6 +10,7 @@
|
||||
#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"
|
||||
|
||||
@@ -27,7 +28,7 @@ public:
|
||||
const BattalionTypeGetter& battalionTypeGetter,
|
||||
ActionPoints braveWaterCost,
|
||||
const AIWaterCrossingCommandChooser& waterCrossingCommandChooser,
|
||||
const vector<CommandProto>& availableCommands) -> AIStrategy;
|
||||
const CommandListSPtr& availableCommands) -> AIStrategy;
|
||||
};
|
||||
|
||||
} // namespace shardok
|
||||
|
||||
@@ -11,9 +11,9 @@
|
||||
#include <limits>
|
||||
|
||||
#include "AICommandFilter.hpp"
|
||||
#include "AIScoreCalculator.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"
|
||||
|
||||
|
||||
@@ -14,7 +14,6 @@
|
||||
#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/api/command_descriptor.pb.h"
|
||||
#include "src/main/protobuf/net/eagle0/shardok/common/command_type.pb.h"
|
||||
|
||||
namespace shardok {
|
||||
@@ -24,7 +23,6 @@ class AIScoreCalculator;
|
||||
class ShardokEngine;
|
||||
|
||||
using ScoreValue = double;
|
||||
using CommandProto = net::eagle0::shardok::api::CommandDescriptor;
|
||||
using CommandType = net::eagle0::shardok::common::CommandType;
|
||||
using BattalionTypeGetter = std::function<BattalionTypeSPtr(BattalionTypeId)>;
|
||||
|
||||
|
||||
@@ -7,6 +7,7 @@
|
||||
#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"
|
||||
@@ -143,15 +144,16 @@ bool AICommandFilter::IsWastefulAction(
|
||||
|
||||
if (!isDefender) {
|
||||
// Attackers: Only allow fire if the target location is on or adjacent to an enemy
|
||||
const auto cmdProto = cmd.GetCommandProto();
|
||||
if (!cmdProto.has_target()) {
|
||||
return true; // Can't analyze without target info
|
||||
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& targetCoords = cmdProto.target();
|
||||
const Coords fireLocation{
|
||||
static_cast<int8_t>(targetCoords.row()),
|
||||
static_cast<int8_t>(targetCoords.column())};
|
||||
const Coords fireLocation(
|
||||
static_cast<int8_t>(targetRow),
|
||||
static_cast<int8_t>(targetCol));
|
||||
|
||||
// Check if any enemy is on the fire location or adjacent to it
|
||||
bool enemyNearFireLocation = false;
|
||||
@@ -186,13 +188,12 @@ bool AICommandFilter::IsWastefulAction(
|
||||
|
||||
if (!isDefender) {
|
||||
// Attackers: Only allow fortify if within 3 hexes of enemies or castles
|
||||
const auto cmdProto = cmd.GetCommandProto();
|
||||
if (!cmdProto.has_actor()) {
|
||||
return true; // Can't analyze without actor info
|
||||
const int unitId = cmd.GetActorUnitId();
|
||||
if (unitId < 0) {
|
||||
throw ShardokInternalErrorException(
|
||||
"FORTIFY_COMMAND missing required actor information");
|
||||
}
|
||||
|
||||
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
|
||||
@@ -249,16 +250,18 @@ bool AICommandFilter::IsWastefulAction(
|
||||
// These actions can fail, so we need high confidence of benefit (8+ action points
|
||||
// saved)
|
||||
|
||||
const auto cmdProto = cmd.GetCommandProto();
|
||||
if (!cmdProto.has_actor() || !cmdProto.has_target()) {
|
||||
return true; // Can't analyze without full command info
|
||||
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 unitId = cmdProto.actor().value();
|
||||
const auto& targetCoords = cmdProto.target();
|
||||
const Coords waterLocation{
|
||||
static_cast<int8_t>(targetCoords.row()),
|
||||
static_cast<int8_t>(targetCoords.column())};
|
||||
const Coords waterLocation(
|
||||
static_cast<int8_t>(targetRow),
|
||||
static_cast<int8_t>(targetCol));
|
||||
|
||||
// Get the acting unit directly by ID
|
||||
const Unit* actingUnit = gameState->units()->Get(unitId);
|
||||
@@ -353,15 +356,16 @@ bool AICommandFilter::IsWastefulAction(
|
||||
case CommandType::REPAIR_COMMAND: {
|
||||
// Repair filtering - filter repairs with high integrity targets
|
||||
// Note: RepairCommandFactory already filters enemy-occupied targets
|
||||
const auto cmdProto = cmd.GetCommandProto();
|
||||
if (!cmdProto.has_target()) {
|
||||
return true; // Can't analyze without target info
|
||||
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& targetCoords = cmdProto.target();
|
||||
const Coords repairLocation{
|
||||
static_cast<int8_t>(targetCoords.row()),
|
||||
static_cast<int8_t>(targetCoords.column())};
|
||||
const Coords repairLocation(
|
||||
static_cast<int8_t>(targetRow),
|
||||
static_cast<int8_t>(targetCol));
|
||||
|
||||
// Check terrain modifiers at target location
|
||||
const auto* terrain = GetTerrain(gameState->hex_map(), repairLocation);
|
||||
@@ -384,15 +388,16 @@ bool AICommandFilter::IsWastefulAction(
|
||||
|
||||
case CommandType::EXTINGUISH_FIRE_COMMAND: {
|
||||
// Extinguish fire filtering - don't extinguish fires on enemy-occupied tiles
|
||||
const auto cmdProto = cmd.GetCommandProto();
|
||||
if (!cmdProto.has_target()) {
|
||||
return true; // Can't analyze without target info
|
||||
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& targetCoords = cmdProto.target();
|
||||
const Coords fireLocation{
|
||||
static_cast<int8_t>(targetCoords.row()),
|
||||
static_cast<int8_t>(targetCoords.column())};
|
||||
const Coords fireLocation(
|
||||
static_cast<int8_t>(targetRow),
|
||||
static_cast<int8_t>(targetCol));
|
||||
|
||||
// Check if any enemy occupies the fire location - let them burn!
|
||||
std::vector<PlayerId> allyPids; // Empty for now - assume 2-player game
|
||||
@@ -424,17 +429,17 @@ bool AICommandFilter::IsWastefulMovement(
|
||||
return false; // Don't filter defender movement or when close to enemies
|
||||
}
|
||||
|
||||
// Get the command proto to access unit and target information
|
||||
const auto cmdProto = cmd.GetCommandProto();
|
||||
// Get unit and target information directly from command
|
||||
const int unitId = cmd.GetActorUnitId();
|
||||
const int targetRow = cmd.GetTargetRow();
|
||||
const int targetCol = cmd.GetTargetColumn();
|
||||
|
||||
// Check if we have the required information
|
||||
if (!cmdProto.has_actor() || !cmdProto.has_target()) {
|
||||
return false; // Can't analyze without unit and target info
|
||||
if (unitId < 0 || targetRow < 0 || targetCol < 0) {
|
||||
throw ShardokInternalErrorException(
|
||||
"MOVE_COMMAND missing required actor or target information");
|
||||
}
|
||||
|
||||
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
|
||||
@@ -450,9 +455,7 @@ bool AICommandFilter::IsWastefulMovement(
|
||||
}
|
||||
|
||||
const auto& currentCoords = actingUnit->location();
|
||||
const Coords targetCoordsFlat{
|
||||
static_cast<int8_t>(targetCoords.row()),
|
||||
static_cast<int8_t>(targetCoords.column())};
|
||||
const Coords targetCoordsFlat(static_cast<int8_t>(targetRow), static_cast<int8_t>(targetCol));
|
||||
|
||||
// Get action point distances for this unit's battalion type
|
||||
const auto& battType = battalionTypeGetter(actingUnit->battalion().type());
|
||||
|
||||
@@ -14,7 +14,6 @@
|
||||
#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 {
|
||||
|
||||
|
||||
@@ -13,6 +13,13 @@ enum class AIAlgorithmType {
|
||||
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
|
||||
};
|
||||
|
||||
} // namespace shardok
|
||||
|
||||
#endif // EAGLE0_AI_CONFIG_HPP
|
||||
@@ -8,9 +8,9 @@
|
||||
#include <ranges>
|
||||
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/AIAttackGroups.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/AIScoreCalculator.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 {
|
||||
|
||||
@@ -15,9 +15,9 @@
|
||||
namespace shardok {
|
||||
|
||||
auto AIFleeDecisionCalculator::GetFleeCommandIndex(
|
||||
const vector<CommandProto>::const_iterator& fleeCommand,
|
||||
const vector<CommandProto>& availableCommands) -> size_t {
|
||||
return static_cast<size_t>(std::distance(availableCommands.begin(), fleeCommand));
|
||||
const CommandList::const_iterator& fleeCommand,
|
||||
const CommandListSPtr& availableCommands) -> size_t {
|
||||
return static_cast<size_t>(std::distance(availableCommands->begin(), fleeCommand));
|
||||
}
|
||||
|
||||
auto AIFleeDecisionCalculator::EstimateCombatSuccess(
|
||||
@@ -134,14 +134,14 @@ auto AIFleeDecisionCalculator::EstimateCombatSuccess(
|
||||
auto AIFleeDecisionCalculator::EvaluateFleeVsFight(
|
||||
PlayerId playerId,
|
||||
const GameStateW& guessedState,
|
||||
const vector<CommandProto>& availableCommands,
|
||||
const vector<CommandProto>::const_iterator& fleeCommand,
|
||||
const CommandListSPtr& availableCommands,
|
||||
const CommandList::const_iterator& fleeCommand,
|
||||
int maxRounds,
|
||||
int minimumFleeOddsThreshold,
|
||||
int desperateFleeThreshold,
|
||||
bool enableDebugLogging) -> FleeDecision {
|
||||
// Get flee success odds
|
||||
const int fleeSuccessChance = fleeCommand->odds().success_chance();
|
||||
const int fleeSuccessChance = (*fleeCommand)->GetOddsPercentile();
|
||||
|
||||
if (enableDebugLogging) {
|
||||
printf("AI FinalRound: Evaluating flee (odds=%d%%)...\n", fleeSuccessChance);
|
||||
|
||||
@@ -9,13 +9,11 @@
|
||||
#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/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
|
||||
@@ -35,8 +33,8 @@ public:
|
||||
[[nodiscard]] static auto EvaluateFleeVsFight(
|
||||
PlayerId playerId,
|
||||
const GameStateW& guessedState,
|
||||
const vector<CommandProto>& availableCommands,
|
||||
const vector<CommandProto>::const_iterator& fleeCommand,
|
||||
const CommandListSPtr& availableCommands,
|
||||
const CommandList::const_iterator& fleeCommand,
|
||||
int maxRounds,
|
||||
int minimumFleeOddsThreshold,
|
||||
int desperateFleeThreshold,
|
||||
@@ -59,8 +57,8 @@ public:
|
||||
private:
|
||||
// Helper to get flee command index
|
||||
[[nodiscard]] static auto GetFleeCommandIndex(
|
||||
const vector<CommandProto>::const_iterator& fleeCommand,
|
||||
const vector<CommandProto>& availableCommands) -> size_t;
|
||||
const CommandList::const_iterator& fleeCommand,
|
||||
const CommandListSPtr& availableCommands) -> size_t;
|
||||
};
|
||||
|
||||
} // namespace shardok
|
||||
|
||||
@@ -0,0 +1,251 @@
|
||||
//
|
||||
// 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
|
||||
@@ -0,0 +1,40 @@
|
||||
//
|
||||
// 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
|
||||
@@ -24,10 +24,40 @@ int AIEvaluationCounter::GetCurrentCount() { return activeCount.load(); }
|
||||
auto CalculateTimeBudget(
|
||||
const PlayerId playerId,
|
||||
const GameSettingsSPtr &settings,
|
||||
const GameStateW &state) -> AITimeBudget {
|
||||
const GameStateW &state,
|
||||
const size_t numCommands) -> 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();
|
||||
@@ -72,12 +102,25 @@ auto CalculateTimeBudget(
|
||||
}
|
||||
}
|
||||
|
||||
// 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());
|
||||
// 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);
|
||||
|
||||
const auto remainingBudget = std::chrono::duration_cast<std::chrono::milliseconds>(budget);
|
||||
// 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);
|
||||
|
||||
// Get minimum depth requirement
|
||||
const size_t minDepth = settingsGetter.Backing().min_lookahead_turns();
|
||||
|
||||
@@ -36,10 +36,13 @@ 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) -> AITimeBudget;
|
||||
const GameStateW &state,
|
||||
size_t numCommands) -> AITimeBudget;
|
||||
|
||||
} // namespace shardok
|
||||
|
||||
|
||||
@@ -17,9 +17,10 @@ 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.99;
|
||||
constexpr double kAdjacentFireMultiplier = 0.80;
|
||||
constexpr double kOnIceMultiplier = 0.25;
|
||||
constexpr double kMeteorStartInRangeValue = 50;
|
||||
constexpr double kMeteorDirectTargetingEnemy = 2;
|
||||
@@ -63,7 +64,8 @@ auto ContextFreeUnitValue(const Unit *unit) -> ScoreValue {
|
||||
4.0;
|
||||
}
|
||||
|
||||
const double vigorValue = unit->has_attached_hero() ? unit->attached_hero().vigor() : 0.0;
|
||||
const double vigorValue =
|
||||
unit->has_attached_hero() ? unit->attached_hero().vigor() * kVigorScoreMultiplier : 0.0;
|
||||
|
||||
double battalionTypeMultiplier = 1.0;
|
||||
switch (unit->battalion().type()) {
|
||||
@@ -355,9 +357,7 @@ auto UnitValue(
|
||||
kCastleMultiplierBonus * (terrain->modifier().castle().integrity() + 25) / 100.0;
|
||||
}
|
||||
double onFireMultiplier = 1.0;
|
||||
if (terrain->modifier().fire().present() && (isAttacker || attackerWantsCastles)) {
|
||||
onFireMultiplier *= kOnFireMultiplier;
|
||||
}
|
||||
if (terrain->modifier().fire().present()) { onFireMultiplier *= kOnFireMultiplier; }
|
||||
{
|
||||
for (const auto adjacentCoords = HexMapUtils::GetAdjacentCoords(map, location);
|
||||
const auto &c : adjacentCoords) {
|
||||
|
||||
@@ -13,11 +13,9 @@
|
||||
#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;
|
||||
|
||||
@@ -39,8 +39,9 @@ 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:__pkg__",
|
||||
"//src/test/cpp/net/eagle0/shardok/ai:__subpackages__",
|
||||
],
|
||||
deps = [
|
||||
":ai_attack_locations",
|
||||
@@ -58,6 +59,10 @@ 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",
|
||||
@@ -96,8 +101,10 @@ 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",
|
||||
@@ -130,8 +137,10 @@ 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",
|
||||
@@ -155,7 +164,24 @@ cc_library(
|
||||
":ai_unit_score_calculator",
|
||||
"//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_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",
|
||||
],
|
||||
)
|
||||
|
||||
@@ -170,15 +196,14 @@ cc_library(
|
||||
],
|
||||
deps = [
|
||||
":ai_command_filter",
|
||||
":ai_score_calculator_interface",
|
||||
":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/api:command_descriptor_cc_proto",
|
||||
"//src/main/protobuf/net/eagle0/shardok/common:command_type_cc_proto",
|
||||
],
|
||||
)
|
||||
@@ -201,7 +226,6 @@ cc_library(
|
||||
"//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",
|
||||
],
|
||||
)
|
||||
@@ -220,59 +244,13 @@ cc_library(
|
||||
],
|
||||
)
|
||||
|
||||
cc_library(
|
||||
name = "ai_score_calculator_interface",
|
||||
hdrs = ["AIScoreCalculator.hpp"],
|
||||
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_attack_locations",
|
||||
"//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/map:coords_set",
|
||||
"//src/main/protobuf/net/eagle0/shardok/api:command_descriptor_cc_proto",
|
||||
],
|
||||
)
|
||||
|
||||
cc_library(
|
||||
name = "standard_ai_score_calculator",
|
||||
srcs = ["StandardAIScoreCalculator.cpp"],
|
||||
hdrs = ["StandardAIScoreCalculator.hpp"],
|
||||
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__",
|
||||
"//src/test/cpp/net/eagle0/shardok/ai/mcts:__pkg__",
|
||||
],
|
||||
deps = [
|
||||
":ai_score_calculator_interface",
|
||||
":ai_strategy",
|
||||
":ai_unit_score_calculator",
|
||||
":ai_victory_condition_score_calculator",
|
||||
":ai_water_crossing_calculator",
|
||||
"//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/settings:game_settings",
|
||||
],
|
||||
)
|
||||
|
||||
cc_library(
|
||||
name = "ai_strategy",
|
||||
srcs = ["AIStrategy.cpp"],
|
||||
hdrs = ["AIStrategy.hpp"],
|
||||
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:__subpackages__",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai_performance_runner:__pkg__",
|
||||
"//src/test/cpp/net/eagle0/shardok/ai:__subpackages__",
|
||||
],
|
||||
@@ -288,6 +266,7 @@ 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__",
|
||||
],
|
||||
@@ -298,28 +277,6 @@ 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_common_types",
|
||||
":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"],
|
||||
@@ -327,12 +284,13 @@ 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",
|
||||
":ai_score_calculator_interface",
|
||||
"//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",
|
||||
@@ -354,7 +312,6 @@ 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",
|
||||
],
|
||||
)
|
||||
|
||||
@@ -391,13 +348,12 @@ cc_library(
|
||||
":ai_attacker_strategy_selector",
|
||||
":ai_command_evaluator",
|
||||
":ai_defender_strategy_selector",
|
||||
":ai_score_calculator_interface",
|
||||
":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",
|
||||
],
|
||||
)
|
||||
|
||||
@@ -423,12 +379,14 @@ cc_library(
|
||||
":ai_defender_strategy_selector",
|
||||
":ai_flee_decision_calculator",
|
||||
":ai_iterative_deepening", # Direct dependency for runtime selection
|
||||
":ai_score_calculator_interface",
|
||||
":ai_time_budget",
|
||||
":ai_water_crossing_command_chooser",
|
||||
":standard_ai_score_calculator",
|
||||
"//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/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",
|
||||
|
||||
@@ -11,8 +11,8 @@
|
||||
|
||||
#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 {
|
||||
@@ -38,7 +38,7 @@ IterativeDeepeningAI::IterativeDeepeningAI(
|
||||
auto IterativeDeepeningAI::IterativeSearch(
|
||||
const GameSettingsSPtr& settings,
|
||||
const GameStateW& state,
|
||||
const std::vector<CommandProto>& commands,
|
||||
const CommandListSPtr& commands,
|
||||
const AITimeBudget& initialBudget) const -> SearchResult {
|
||||
// Make a mutable copy of the time budget to track remaining time
|
||||
AITimeBudget timeBudget = initialBudget;
|
||||
@@ -51,7 +51,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
|
||||
@@ -75,9 +75,9 @@ auto IterativeDeepeningAI::IterativeSearch(
|
||||
|
||||
// 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
|
||||
@@ -132,7 +132,8 @@ auto IterativeDeepeningAI::IterativeSearch(
|
||||
evaluatedCount++;
|
||||
|
||||
// Check if this command is not END_TURN_COMMAND
|
||||
if (commands[cmdIndex].type() != net::eagle0::shardok::common::END_TURN_COMMAND) {
|
||||
if ((*commands)[cmdIndex]->GetCommandType() !=
|
||||
net::eagle0::shardok::common::END_TURN_COMMAND) {
|
||||
allEndTurnCommands = false;
|
||||
}
|
||||
}
|
||||
@@ -143,7 +144,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];
|
||||
@@ -156,16 +157,20 @@ 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) - %s\n",
|
||||
printf(" Depth %lu best: command %zu (score %.2f) - type: %s\n",
|
||||
currentDepth - 1,
|
||||
previousBestCommand,
|
||||
scoresByDepth[previousBestCommand][currentDepth - 1],
|
||||
commands[previousBestCommand].DebugString().c_str());
|
||||
printf(" Depth %lu best: command %zu (score %.2f) - %s\n",
|
||||
net::eagle0::shardok::common::CommandType_Name(
|
||||
(*commands)[previousBestCommand]->GetCommandType())
|
||||
.c_str());
|
||||
printf(" Depth %lu best: command %zu (score %.2f) - type: %s\n",
|
||||
currentDepth,
|
||||
currentBestCommand,
|
||||
currentBestScore,
|
||||
commands[currentBestCommand].DebugString().c_str());
|
||||
net::eagle0::shardok::common::CommandType_Name(
|
||||
(*commands)[currentBestCommand]->GetCommandType())
|
||||
.c_str());
|
||||
#endif
|
||||
}
|
||||
|
||||
@@ -244,7 +249,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;
|
||||
|
||||
@@ -269,7 +274,7 @@ auto IterativeDeepeningAI::SearchCommandAtDepthWithEngine(
|
||||
const ShardokEngine& guessedEngine,
|
||||
const AIScoreCalculator& scorer,
|
||||
const int maxRepeatCount,
|
||||
const std::vector<CommandProto>& commands,
|
||||
const CommandListSPtr& commands,
|
||||
const size_t commandIndex,
|
||||
const int desiredDepth,
|
||||
const ScoreValue currentUtility,
|
||||
@@ -279,10 +284,10 @@ 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);
|
||||
|
||||
@@ -9,21 +9,20 @@
|
||||
#include <future>
|
||||
#include <vector>
|
||||
|
||||
#include "AIScoreCalculator.hpp"
|
||||
#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 CommandProto = net::eagle0::shardok::api::CommandDescriptor;
|
||||
using BattalionTypeGetter = std::function<BattalionTypeSPtr(BattalionTypeId)>;
|
||||
|
||||
/// Reason why AI evaluation completed at the achieved depth.
|
||||
@@ -70,7 +69,7 @@ public:
|
||||
[[nodiscard]] SearchResult IterativeSearch(
|
||||
const GameSettingsSPtr& settings,
|
||||
const GameStateW& state,
|
||||
const std::vector<CommandProto>& commands,
|
||||
const CommandListSPtr& commands,
|
||||
const AITimeBudget& initialBudget) const;
|
||||
|
||||
private:
|
||||
@@ -93,7 +92,7 @@ private:
|
||||
const ShardokEngine& guessedEngine,
|
||||
const AIScoreCalculator& scorer,
|
||||
int maxRepeatCount,
|
||||
const std::vector<CommandProto>& commands,
|
||||
const CommandListSPtr& commands,
|
||||
size_t commandIndex,
|
||||
int desiredDepth,
|
||||
ScoreValue currentUtility,
|
||||
|
||||
@@ -10,7 +10,15 @@
|
||||
|
||||
#define DEBUG_FLEE_DECISIONS
|
||||
|
||||
#include <google/protobuf/util/message_differencer.h>
|
||||
// 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 "AIAttackerStrategySelector.hpp"
|
||||
#include "AIConfig.hpp"
|
||||
@@ -19,8 +27,10 @@
|
||||
#include "AIScoreUtilities.hpp"
|
||||
#include "AITimeBudget.hpp"
|
||||
#include "IterativeDeepeningAI.hpp"
|
||||
#include "StandardAIScoreCalculator.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 "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"
|
||||
@@ -45,12 +55,16 @@ ShardokAIClient::ShardokAIClient(
|
||||
const bool isDefender,
|
||||
const HexMap *hexMap,
|
||||
const SettingsGetter &settings,
|
||||
const AIAlgorithmType aiAlgorithmType)
|
||||
const AIAlgorithmType aiAlgorithmType,
|
||||
const ScoringCalculatorType scoringCalculatorType,
|
||||
const mcts::MCTSConfig &mctsConfig)
|
||||
: playerId(playerId),
|
||||
isDefender(isDefender),
|
||||
aiAlgorithmType(aiAlgorithmType),
|
||||
scoringCalculatorType(scoringCalculatorType),
|
||||
alCache(std::make_unique<AttackLocationsCache>(hexMap, settings)),
|
||||
waterCrossingCommandChooser(playerId, apdCache) {
|
||||
waterCrossingCommandChooser(playerId, apdCache),
|
||||
mctsConfig(mctsConfig) {
|
||||
// Pre-generate the most common cache entries for better performance
|
||||
const auto mapId = ActionPointDistancesCache::GetMapId(hexMap);
|
||||
|
||||
@@ -74,38 +88,90 @@ ShardokAIClient::ShardokAIClient(
|
||||
apdCache->ConsolidateThreadLocalCache_Racy();
|
||||
}
|
||||
|
||||
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());
|
||||
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.
|
||||
|
||||
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");
|
||||
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");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
auto ShardokAIClient::StandardChooseCommandIndex(
|
||||
const GameSettingsSPtr &settings,
|
||||
const GameStateW &guessedState,
|
||||
const vector<CommandProto> &realAvailableCommands) const -> CommandChoiceResults {
|
||||
const CommandListSPtr &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 settings
|
||||
const auto timeBudget = CalculateTimeBudget(playerId, settings, guessedState);
|
||||
// Calculate time budget based on game situation using new dynamic per-command settings
|
||||
const auto timeBudget = CalculateTimeBudget(playerId, settings, guessedState, commandCount);
|
||||
|
||||
const auto guessedCommands = guessedEngine.GetAvailableCommandProtos(playerId, false);
|
||||
const auto commandCount = guessedCommands.size();
|
||||
// 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;
|
||||
|
||||
assert(commandCount == realAvailableCommands.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
|
||||
for (size_t i = 0; i < commandCount; i++) {
|
||||
CheckCommand(realAvailableCommands[i], guessedCommands[i]);
|
||||
CheckCommand((*realAvailableCommands)[i], (*guessedCommands)[i]);
|
||||
}
|
||||
|
||||
// Extract values directly from settings for strategy selection
|
||||
@@ -115,8 +181,18 @@ auto ShardokAIClient::StandardChooseCommandIndex(
|
||||
return settingsGetter.GetBattalionType(typeId);
|
||||
};
|
||||
|
||||
// Create scorer for actual scoring during search
|
||||
const auto scorer = MakeStandardAIScoreCalculator(settingsGetter, apdCache, alCache);
|
||||
// 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;
|
||||
}
|
||||
|
||||
// Determine strategy once for consistent scoring throughout iterative deepening
|
||||
const auto castleCoords = AllCastleCoords(guessedState->hex_map());
|
||||
@@ -142,8 +218,47 @@ 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);
|
||||
ShardokMCTSAI ai(
|
||||
playerId,
|
||||
isDefender,
|
||||
strategy,
|
||||
castleCoords,
|
||||
*scorer,
|
||||
apdCache,
|
||||
alCache,
|
||||
adjustedMCTSConfig);
|
||||
search_result = ai.Search(settings, guessedState, timeBudget);
|
||||
} else {
|
||||
// Using Iterative Deepening AI (default)
|
||||
@@ -173,7 +288,8 @@ auto ShardokAIClient::StandardChooseCommandIndex(
|
||||
result.commandCountEvaluated,
|
||||
result.availableCommandCount);
|
||||
}
|
||||
const auto chosenCommandType = realAvailableCommands[result.chosenIndex].type();
|
||||
const auto chosenCommandType =
|
||||
(*realAvailableCommands)[result.chosenIndex]->GetCommandType();
|
||||
printf("ID AI: Search complete - achieved depth %d for best command %zu (%s)\n",
|
||||
result.depthAchieved,
|
||||
result.chosenIndex,
|
||||
@@ -188,19 +304,20 @@ auto ShardokAIClient::StandardChooseCommandIndex(
|
||||
auto ShardokAIClient::LateRoundAttackerChooseCommandIndex(
|
||||
const GameSettingsSPtr &settings,
|
||||
const GameStateW &guessedState,
|
||||
const vector<CommandProto> &realAvailableCommands) const -> CommandChoiceResults {
|
||||
const CommandListSPtr &realAvailableCommands) const -> CommandChoiceResults {
|
||||
if (const auto dismissCommand = std::ranges::find_if(
|
||||
realAvailableCommands,
|
||||
[](const net::eagle0::shardok::api::CommandDescriptor &cmd) {
|
||||
return cmd.type() == net::eagle0::shardok::common::DISMISS_UNIT_COMMAND;
|
||||
*realAvailableCommands,
|
||||
[](const CommandSPtr &cmd) {
|
||||
return cmd->GetCommandType() ==
|
||||
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 =
|
||||
@@ -212,14 +329,13 @@ auto ShardokAIClient::LateRoundAttackerChooseCommandIndex(
|
||||
auto ShardokAIClient::FinalRoundAttackerChooseCommandIndex(
|
||||
const GameSettingsSPtr &settings,
|
||||
const GameStateW &guessedState,
|
||||
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;
|
||||
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;
|
||||
});
|
||||
|
||||
if (fleeCommand == realAvailableCommands.end()) {
|
||||
if (fleeCommand == realAvailableCommands->end()) {
|
||||
return LateRoundAttackerChooseCommandIndex(settings, guessedState, realAvailableCommands);
|
||||
}
|
||||
|
||||
@@ -248,7 +364,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;
|
||||
@@ -262,7 +378,7 @@ auto ShardokAIClient::FinalRoundAttackerChooseCommandIndex(
|
||||
auto ShardokAIClient::ChooseCommandIndex(
|
||||
const GameSettingsSPtr &settings,
|
||||
const GameStateView &gsv,
|
||||
const vector<CommandProto> &realAvailableCommands) const -> CommandChoiceResults {
|
||||
const CommandListSPtr &realAvailableCommands) const -> CommandChoiceResults {
|
||||
static int typeChosenCount[net::eagle0::shardok::common::CommandType_MAX + 1];
|
||||
static int totalChoices = 0;
|
||||
|
||||
@@ -281,7 +397,7 @@ auto ShardokAIClient::ChooseCommandIndex(
|
||||
results = StandardChooseCommandIndex(settings, guessedState, realAvailableCommands);
|
||||
}
|
||||
|
||||
const auto chosenType = realAvailableCommands[results.chosenIndex].type();
|
||||
const auto chosenType = (*realAvailableCommands)[results.chosenIndex]->GetCommandType();
|
||||
typeChosenCount[static_cast<int>(chosenType)]++;
|
||||
totalChoices++;
|
||||
|
||||
@@ -307,8 +423,8 @@ auto ShardokAIClient::ChooseCommandIndex(
|
||||
|
||||
auto ShardokAIClient::ChooseCommandIndex(const ShardokEngine &engine) const
|
||||
-> CommandChoiceResults {
|
||||
if (const auto &availableCommands = engine.GetAvailableCommandProtos(playerId, false);
|
||||
availableCommands.empty()) {
|
||||
if (const auto &availableCommands = engine.GetAvailableCommandsForAIPlayer(playerId);
|
||||
availableCommands->empty()) {
|
||||
printf("no commands for player %d\n", playerId);
|
||||
throw ShardokInternalErrorException(
|
||||
"Asked to choose a command, but there are none available");
|
||||
|
||||
@@ -12,11 +12,13 @@
|
||||
#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 {
|
||||
@@ -40,29 +42,33 @@ 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 vector<CommandProto>& realAvailableCommands) const -> CommandChoiceResults;
|
||||
const CommandListSPtr& realAvailableCommands) const -> CommandChoiceResults;
|
||||
[[nodiscard]] auto LateRoundAttackerChooseCommandIndex(
|
||||
const GameSettingsSPtr& settings,
|
||||
const GameStateW& guessedState,
|
||||
const vector<CommandProto>& realAvailableCommands) const -> CommandChoiceResults;
|
||||
const CommandListSPtr& realAvailableCommands) const -> CommandChoiceResults;
|
||||
[[nodiscard]] auto FinalRoundAttackerChooseCommandIndex(
|
||||
const GameSettingsSPtr& settings,
|
||||
const GameStateW& guessedState,
|
||||
const vector<CommandProto>& realAvailableCommands) const -> CommandChoiceResults;
|
||||
const CommandListSPtr& realAvailableCommands) const -> CommandChoiceResults;
|
||||
|
||||
[[nodiscard]] auto ChooseCommandIndex(
|
||||
const GameSettingsSPtr& settings,
|
||||
const net::eagle0::shardok::api::GameStateView& gsv,
|
||||
const vector<CommandProto>& realAvailableCommands) const -> CommandChoiceResults;
|
||||
const CommandListSPtr& realAvailableCommands) const -> CommandChoiceResults;
|
||||
|
||||
public:
|
||||
explicit ShardokAIClient(
|
||||
@@ -70,13 +76,19 @@ public:
|
||||
bool isDefender,
|
||||
const HexMap* hexMap,
|
||||
const SettingsGetter& settings,
|
||||
AIAlgorithmType aiAlgorithmType = AIAlgorithmType::ITERATIVE_DEEPENING);
|
||||
AIAlgorithmType aiAlgorithmType,
|
||||
ScoringCalculatorType scoringCalculatorType,
|
||||
const mcts::MCTSConfig& mctsConfig);
|
||||
~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,826 +0,0 @@
|
||||
//
|
||||
// Standard implementation of AIScoreCalculator
|
||||
//
|
||||
|
||||
#include "StandardAIScoreCalculator.hpp"
|
||||
|
||||
#include <atomic>
|
||||
#include <chrono>
|
||||
#include <future>
|
||||
#include <unordered_map>
|
||||
|
||||
#include "AIAttackGroups.hpp"
|
||||
#include "AIAttackLocations.hpp"
|
||||
#include "AIScoreCalculator.hpp"
|
||||
#include "AIScoreUtilities.hpp"
|
||||
#include "AIStrategy.hpp"
|
||||
#include "AIUnitScoreCalculator.hpp"
|
||||
#include "AIVictoryConditionScoreCalculator.hpp"
|
||||
#include "AIWaterCrossingCalculator.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/settings/GameSettings.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/api/unit_view.pb.h"
|
||||
|
||||
namespace shardok {
|
||||
|
||||
using net::eagle0::shardok::storage::fb::BattalionTypeId;
|
||||
|
||||
// Forward declare the implementation class
|
||||
class StandardAIScoreCalculator;
|
||||
|
||||
// Anonymous namespace for helper functions that don't need access to scorer
|
||||
namespace {
|
||||
|
||||
#define LOGGING_ 0
|
||||
#define PERFORMANCE_LOGGING_ 0
|
||||
|
||||
// Performance logging for AttackerScoreForState
|
||||
struct AttackerScorePerformanceLogger {
|
||||
static constexpr int LOG_INTERVAL = 100000;
|
||||
|
||||
static std::atomic<int> callCount;
|
||||
static std::atomic<double> intervalTime;
|
||||
static std::atomic<double> totalTime;
|
||||
|
||||
static void LogCall(double duration) {
|
||||
callCount.fetch_add(1);
|
||||
intervalTime.fetch_add(duration);
|
||||
totalTime.fetch_add(duration);
|
||||
|
||||
if (callCount.load() % LOG_INTERVAL == 0) {
|
||||
double intervalAvg = intervalTime.load() / LOG_INTERVAL;
|
||||
double overallAvg = totalTime.load() / callCount.load();
|
||||
printf("AttackerScoreForState: %d calls, last %d avg: %.1f µs, overall avg: %.1f µs\n",
|
||||
callCount.load(),
|
||||
LOG_INTERVAL,
|
||||
intervalAvg * 1000000.0,
|
||||
overallAvg * 1000000.0);
|
||||
intervalTime.store(0.0); // Reset for next interval
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
std::atomic<int> AttackerScorePerformanceLogger::callCount{0};
|
||||
std::atomic<double> AttackerScorePerformanceLogger::intervalTime{0.0};
|
||||
std::atomic<double> AttackerScorePerformanceLogger::totalTime{0.0};
|
||||
|
||||
// RAII timer for automatic performance logging
|
||||
class AttackerScoreTimer {
|
||||
private:
|
||||
std::chrono::high_resolution_clock::time_point startTime;
|
||||
|
||||
public:
|
||||
AttackerScoreTimer() : startTime(std::chrono::high_resolution_clock::now()) {}
|
||||
|
||||
~AttackerScoreTimer() {
|
||||
auto endTime = std::chrono::high_resolution_clock::now();
|
||||
auto duration =
|
||||
std::chrono::duration_cast<std::chrono::duration<double>>(endTime - startTime);
|
||||
AttackerScorePerformanceLogger::LogCall(duration.count());
|
||||
}
|
||||
};
|
||||
|
||||
// Memoization cache for EffectiveDistance calls
|
||||
struct EffectiveDistanceCache {
|
||||
struct CacheKey {
|
||||
UnitId unitId;
|
||||
Coords target;
|
||||
bool operator==(const CacheKey &other) const {
|
||||
return unitId == other.unitId && target == other.target;
|
||||
}
|
||||
};
|
||||
|
||||
struct CacheKeyHash {
|
||||
size_t operator()(const CacheKey &key) const {
|
||||
return std::hash<UnitId>{}(key.unitId) ^ (std::hash<int>{}(key.target.row()) << 1) ^
|
||||
(std::hash<int>{}(key.target.column()) << 2);
|
||||
}
|
||||
};
|
||||
|
||||
mutable gtl::flat_hash_map<CacheKey, DIST_T, CacheKeyHash> cache;
|
||||
|
||||
DIST_T GetOrCompute(
|
||||
const Unit *unit,
|
||||
const Coords &target,
|
||||
const ActionPointDistances *notBravingApd,
|
||||
const ActionPointDistances *bravingApd,
|
||||
const HexMap *hexMap) const {
|
||||
CacheKey key{unit->unit_id(), target};
|
||||
auto it = cache.find(key);
|
||||
if (it != cache.end()) { return it->second; }
|
||||
|
||||
CoordsSet targetSet(hexMap);
|
||||
targetSet.Add(target);
|
||||
|
||||
DIST_T result = EffectiveDistance(unit, notBravingApd, bravingApd, targetSet);
|
||||
cache[key] = result;
|
||||
return result;
|
||||
}
|
||||
};
|
||||
|
||||
#define MULTITHREAD true
|
||||
|
||||
constexpr double UNITS_BASE_MULTIPLIER = 0.05;
|
||||
|
||||
constexpr double FLEE_UNIT_SCORE = -10000;
|
||||
constexpr double FLEE_CONTROLLING_UNIT_SCORE = -10000;
|
||||
constexpr double CAPTURED_UNIT_SCORE = -10000;
|
||||
constexpr double CAPTURED_VIP_SCORE = -25000;
|
||||
|
||||
constexpr double kMaxProximityBuf = 1.5;
|
||||
constexpr double kDistanceDebufRatio = 8.0;
|
||||
|
||||
using std::async;
|
||||
using std::future;
|
||||
|
||||
using flatbuffers::FlatBufferBuilder;
|
||||
using flatbuffers::Offset;
|
||||
|
||||
using net::eagle0::shardok::api::HeroView;
|
||||
using net::eagle0::shardok::api::UnitView;
|
||||
using GameState = fb::GameState;
|
||||
using Unit = fb::Unit;
|
||||
|
||||
static auto RecursiveAttackerMultiplierForTargetDistance(
|
||||
const Unit *attackingUnit,
|
||||
vector<TargetAndAttackLocations>::const_iterator &priorityListNext,
|
||||
const vector<TargetAndAttackLocations>::const_iterator &priorityListEnd,
|
||||
const vector<const Unit *> &occupants,
|
||||
const HexMap *map,
|
||||
const BattalionTypeSPtr &battType,
|
||||
const ActionPointDistances *notBravingApd,
|
||||
const ActionPointDistances *bravingApd,
|
||||
bool isLateGame) -> double;
|
||||
|
||||
static auto RecursiveAttackerMultiplierForTargetDistance(
|
||||
const Unit *attackingUnit,
|
||||
vector<TargetAndAttackLocations>::const_iterator &priorityListNext,
|
||||
const vector<TargetAndAttackLocations>::const_iterator &priorityListEnd,
|
||||
const vector<const Unit *> &occupants,
|
||||
const HexMap *map,
|
||||
const BattalionTypeSPtr &battType,
|
||||
const ActionPointDistances *notBravingApd,
|
||||
const ActionPointDistances *bravingApd,
|
||||
const bool isLateGame) -> double {
|
||||
if (priorityListNext == priorityListEnd) return 1.0;
|
||||
|
||||
const auto &[target, attackLocations] = *priorityListNext;
|
||||
const Coords &topPriorityTarget = target;
|
||||
|
||||
// If the target is unoccupied or is occupied by this player, give the maximum multiplier, but
|
||||
// also add the bonus for the next up in the priority list
|
||||
if (const Unit *occupant = occupants
|
||||
[topPriorityTarget.row() * map->column_count() + topPriorityTarget.column()];
|
||||
!occupant || occupant->player_id() == attackingUnit->player_id()) {
|
||||
return kMaxProximityBuf + RecursiveAttackerMultiplierForTargetDistance(
|
||||
attackingUnit,
|
||||
++priorityListNext,
|
||||
priorityListEnd,
|
||||
occupants,
|
||||
map,
|
||||
battType,
|
||||
notBravingApd,
|
||||
bravingApd,
|
||||
isLateGame);
|
||||
}
|
||||
|
||||
// Use optimized EffectiveDistance with pre-computed ActionPointDistances
|
||||
// attackLocations is already the CoordsSet of attack locations for this target
|
||||
const DIST_T distance =
|
||||
EffectiveDistance(attackingUnit, notBravingApd, bravingApd, attackLocations);
|
||||
|
||||
return kMaxProximityBuf / (1 + distance / kDistanceDebufRatio);
|
||||
}
|
||||
|
||||
// Overload that accepts pre-computed ActionPointDistances
|
||||
auto AttackerMultiplierForTargetDistance(
|
||||
const Unit *attackingUnit,
|
||||
const vector<TargetAndAttackLocations> &priorityList,
|
||||
const vector<const Unit *> &occupants,
|
||||
const HexMap *map,
|
||||
const BattalionTypeSPtr &battType,
|
||||
const ActionPointDistances *notBravingApd,
|
||||
const ActionPointDistances *bravingApd,
|
||||
const bool isLateGame) -> double {
|
||||
auto iter = begin(priorityList);
|
||||
return RecursiveAttackerMultiplierForTargetDistance(
|
||||
attackingUnit,
|
||||
iter,
|
||||
end(priorityList),
|
||||
occupants,
|
||||
map,
|
||||
battType,
|
||||
notBravingApd,
|
||||
bravingApd,
|
||||
isLateGame);
|
||||
}
|
||||
|
||||
auto FleeStrategyScoreForState(const GameStateW &gameState, const PlayerId playerId) -> ScoreValue {
|
||||
ScoreValue scoreValue = 0.0;
|
||||
|
||||
const auto *gameStatePtr = gameState.Get();
|
||||
const auto *units = gameStatePtr->units();
|
||||
|
||||
for (const auto *unit : *units) {
|
||||
if (unit->status() != net::eagle0::shardok::storage::fb::UnitStatus_NORMAL_UNIT) continue;
|
||||
|
||||
if (unit->player_id() == playerId &&
|
||||
unit->battalion().type() != net::eagle0::shardok::storage::fb::BattalionTypeId_UNDEAD) {
|
||||
scoreValue += FLEE_UNIT_SCORE;
|
||||
|
||||
if (unit->has_attached_hero() &&
|
||||
unit->attached_hero().control_info().controlled_unit_id() != -1) {
|
||||
scoreValue += FLEE_CONTROLLING_UNIT_SCORE;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return scoreValue;
|
||||
}
|
||||
|
||||
} // anonymous namespace
|
||||
|
||||
/// Standard implementation of AIScoreCalculator that uses the default scoring algorithm.
|
||||
/// Stores specific setting values needed for scoring rather than the entire SettingsGetter.
|
||||
class StandardAIScoreCalculator : public AIScoreCalculator {
|
||||
public:
|
||||
StandardAIScoreCalculator(
|
||||
int maxRounds,
|
||||
ActionPoints braveWaterCost,
|
||||
int meteorRange,
|
||||
double meteorCastVigorCost,
|
||||
int minimumFleeOddsThreshold,
|
||||
int desperateFleeThreshold,
|
||||
std::unordered_map<BattalionTypeId, BattalionTypeSPtr> battalionTypes,
|
||||
const APDCache &apdCache,
|
||||
const ALCache &alCache)
|
||||
: maxRounds_(maxRounds),
|
||||
braveWaterCost_(braveWaterCost),
|
||||
meteorRange_(meteorRange),
|
||||
meteorCastVigorCost_(meteorCastVigorCost),
|
||||
minimumFleeOddsThreshold_(minimumFleeOddsThreshold),
|
||||
desperateFleeThreshold_(desperateFleeThreshold),
|
||||
battalionTypes_(std::move(battalionTypes)),
|
||||
apdCache_(apdCache),
|
||||
alCache_(alCache) {}
|
||||
|
||||
[[nodiscard]] auto GuessedStateScore(
|
||||
bool isDefender,
|
||||
const GameStateW &state,
|
||||
const AIStrategy &aiStrategy,
|
||||
const CoordsSet &allCastleCoords) const -> ScoreValue override;
|
||||
|
||||
private:
|
||||
// Internal accessor methods
|
||||
[[nodiscard]] auto GetBattalionType(BattalionTypeId typeId) const -> BattalionTypeSPtr {
|
||||
auto it = battalionTypes_.find(typeId);
|
||||
if (it == battalionTypes_.end()) {
|
||||
throw ShardokInternalErrorException("Unknown battalion type ID");
|
||||
}
|
||||
return it->second;
|
||||
}
|
||||
[[nodiscard]] auto GetApdCache() const -> const APDCache & { return apdCache_; }
|
||||
[[nodiscard]] auto GetAlCache() const -> const ALCache & { return alCache_; }
|
||||
[[nodiscard]] auto GetBraveWaterCost() const -> ActionPoints { return braveWaterCost_; }
|
||||
[[nodiscard]] auto GetMaxRounds() const -> int { return maxRounds_; }
|
||||
[[nodiscard]] auto GetMeteorRange() const -> int { return meteorRange_; }
|
||||
[[nodiscard]] auto GetMeteorCastVigorCost() const -> double { return meteorCastVigorCost_; }
|
||||
[[nodiscard]] auto GetMinimumFleeOddsThreshold() const -> int {
|
||||
return minimumFleeOddsThreshold_;
|
||||
}
|
||||
[[nodiscard]] auto GetDesperateFleeThreshold() const -> int { return desperateFleeThreshold_; }
|
||||
|
||||
// Implementation methods (converted from internal namespace functions)
|
||||
[[nodiscard]] auto AttackerUnitsScore(
|
||||
const GameStateW &gameState,
|
||||
int roundsRemaining,
|
||||
bool attackerWantsCastles,
|
||||
bool defenderShouldScatter,
|
||||
const vector<TargetPriorityList> &attackerTargetPriorities,
|
||||
const MapId &mapId) const -> ScoreValue;
|
||||
|
||||
[[nodiscard]] auto DefenderScatterStrategyScoreForState(
|
||||
const GameStateW &gameState,
|
||||
int roundsRemaining) const -> ScoreValue;
|
||||
|
||||
[[nodiscard]] auto DefenderHoldCastlesStrategyScoreForState(
|
||||
const GameStateW &gameState,
|
||||
const CoordsSet &castleCoords,
|
||||
int roundsRemaining) const -> ScoreValue;
|
||||
|
||||
[[nodiscard]] auto DefenderScoreForState(
|
||||
const GameStateW &gameState,
|
||||
const AIStrategy &defenderStrategy,
|
||||
const CoordsSet &castleCoords,
|
||||
int roundsRemaining) const -> ScoreValue;
|
||||
|
||||
[[nodiscard]] auto AttackerScoreForState(
|
||||
const GameStateW &gameState,
|
||||
const AIStrategy &attackerStrategy,
|
||||
const CoordsSet &castleCoords,
|
||||
int roundsRemaining) const -> ScoreValue;
|
||||
|
||||
// Scalar settings extracted from SettingsGetter
|
||||
int maxRounds_;
|
||||
ActionPoints braveWaterCost_;
|
||||
int meteorRange_;
|
||||
double meteorCastVigorCost_;
|
||||
int minimumFleeOddsThreshold_;
|
||||
int desperateFleeThreshold_;
|
||||
|
||||
// Battalion type lookup map
|
||||
std::unordered_map<BattalionTypeId, BattalionTypeSPtr> battalionTypes_;
|
||||
|
||||
// Caches (stored as references)
|
||||
const APDCache &apdCache_;
|
||||
const ALCache &alCache_;
|
||||
};
|
||||
|
||||
// Implementation of StandardAIScoreCalculator methods
|
||||
|
||||
auto StandardAIScoreCalculator::AttackerUnitsScore(
|
||||
const GameStateW &gameState,
|
||||
int roundsRemaining,
|
||||
bool attackerWantsCastles,
|
||||
bool defenderShouldScatter,
|
||||
const vector<TargetPriorityList> &attackerTargetPriorities,
|
||||
const MapId &mapId) const -> ScoreValue {
|
||||
// Cache frequently accessed FlatBuffer fields to avoid repeated offset calculations
|
||||
const auto *gameStateRawPtr = gameState.Get();
|
||||
const auto *cachedUnits = gameStateRawPtr->units();
|
||||
const auto *cachedHexMap = gameStateRawPtr->hex_map();
|
||||
|
||||
const int16_t cachedRowCount = cachedHexMap->row_count();
|
||||
const int16_t cachedColumnCount = cachedHexMap->column_count();
|
||||
const int cachedCurrentRound = gameStateRawPtr->current_round();
|
||||
|
||||
bool isLateGame = cachedCurrentRound > 18; // Inline IsLateGame for efficiency
|
||||
|
||||
// APDCache now has built-in thread-local caching - no need for PreCachedAPDs
|
||||
ActionPoints braveWaterCost = GetBraveWaterCost();
|
||||
|
||||
// Memoization cache for EffectiveDistance calls
|
||||
EffectiveDistanceCache distanceCache;
|
||||
|
||||
std::vector<const Unit *> attackerUnits{};
|
||||
std::vector<const Unit *> defenderUnits{};
|
||||
// Pre-allocate vectors based on estimated unit ratios to avoid reallocations
|
||||
const size_t estimatedUnitCount = cachedUnits->size();
|
||||
attackerUnits.reserve(estimatedUnitCount - 1);
|
||||
defenderUnits.reserve(estimatedUnitCount - 1);
|
||||
|
||||
double attackerUnitsValue = 0;
|
||||
double defenderUnitsValue = 0;
|
||||
|
||||
auto occupants = Occupants(*cachedUnits, cachedRowCount, cachedColumnCount);
|
||||
|
||||
for (const Unit *unit : *cachedUnits) {
|
||||
const auto *pi = PlayerInfoForPid(gameState, unit->player_id());
|
||||
if (pi == nullptr) { continue; }
|
||||
|
||||
switch (unit->status()) {
|
||||
case net::eagle0::shardok::storage::fb::UnitStatus_NORMAL_UNIT: {
|
||||
if (pi->is_defender()) {
|
||||
defenderUnits.push_back(unit);
|
||||
} else {
|
||||
attackerUnits.push_back(unit);
|
||||
}
|
||||
break;
|
||||
}
|
||||
case net::eagle0::shardok::storage::fb::UnitStatus_CAPTURED_UNIT: {
|
||||
double thisScore = unit->has_attached_hero() && unit->attached_hero().is_vip()
|
||||
? CAPTURED_VIP_SCORE
|
||||
: CAPTURED_UNIT_SCORE;
|
||||
if (pi->is_defender()) {
|
||||
defenderUnitsValue += thisScore;
|
||||
} else {
|
||||
attackerUnitsValue += thisScore;
|
||||
}
|
||||
break;
|
||||
}
|
||||
|
||||
case net::eagle0::shardok::storage::fb::UnitStatus_DESTROYED_SUMMONED_UNIT:
|
||||
case net::eagle0::shardok::storage::fb::UnitStatus_FLED_UNIT:
|
||||
case net::eagle0::shardok::storage::fb::UnitStatus_NEVER_ENTERED_UNIT:
|
||||
case net::eagle0::shardok::storage::fb::UnitStatus_OUTLAWED_UNIT:
|
||||
case net::eagle0::shardok::storage::fb::UnitStatus_RESERVE_UNIT:
|
||||
case net::eagle0::shardok::storage::fb::UnitStatus_RETREATED_UNIT:
|
||||
case net::eagle0::shardok::storage::fb::UnitStatus_RESERVED_SLOT: break;
|
||||
|
||||
case net::eagle0::shardok::storage::fb::UnitStatus_UNKNOWN_UNIT:
|
||||
throw ShardokInternalErrorException("Unknown unit status");
|
||||
}
|
||||
}
|
||||
|
||||
double defenderAdvantage = 1.0 + static_cast<double>(cachedCurrentRound) / 31.0;
|
||||
|
||||
// Can we cache this somehow, it won't usually change within your turn
|
||||
auto attackLocationsForAttacker = GetAlCache()->CachedLocations(defenderUnits, isLateGame);
|
||||
const auto &locationsCausingDanger = attackLocationsForAttacker.AllLocations();
|
||||
|
||||
// Process attacker units using cached ActionPointDistances
|
||||
for (const Unit *unit : attackerUnits) {
|
||||
const int battTypeId = unit->battalion().type();
|
||||
|
||||
const auto &priorityList = std::ranges::find_if(
|
||||
attackerTargetPriorities,
|
||||
[&unit](const TargetPriorityList &tpl) {
|
||||
return tpl.attackingUnitId == unit->unit_id();
|
||||
});
|
||||
|
||||
// If there are any tiles being targeted, give this unit a multiplier based on how close
|
||||
// they are to being able to attack it
|
||||
double distanceMultiplier =
|
||||
priorityList == end(attackerTargetPriorities)
|
||||
? 1.0
|
||||
: AttackerMultiplierForTargetDistance(
|
||||
unit,
|
||||
priorityList->priorityOrder,
|
||||
occupants,
|
||||
cachedHexMap,
|
||||
GetBattalionType(static_cast<BattalionTypeId>(battTypeId)),
|
||||
GetApdCache()->GetRaw(
|
||||
cachedHexMap,
|
||||
mapId,
|
||||
GetBattalionType(
|
||||
static_cast<BattalionTypeId>(battTypeId)),
|
||||
false),
|
||||
GetBattalionType(static_cast<BattalionTypeId>(battTypeId))
|
||||
->allowsBraveWater
|
||||
? GetApdCache()->GetRaw(
|
||||
cachedHexMap,
|
||||
mapId,
|
||||
GetBattalionType(static_cast<BattalionTypeId>(
|
||||
battTypeId)),
|
||||
true,
|
||||
braveWaterCost)
|
||||
: nullptr,
|
||||
isLateGame);
|
||||
|
||||
auto uv = UnitValue(
|
||||
unit,
|
||||
true,
|
||||
attackerUnits,
|
||||
attackerWantsCastles,
|
||||
/* includeCastleBonus=*/true,
|
||||
defenderUnits,
|
||||
cachedHexMap,
|
||||
roundsRemaining,
|
||||
attackLocationsForAttacker,
|
||||
locationsCausingDanger,
|
||||
GetApdCache()->GetRaw(
|
||||
cachedHexMap,
|
||||
mapId,
|
||||
GetBattalionType(static_cast<BattalionTypeId>(battTypeId)),
|
||||
false),
|
||||
GetMeteorRange(),
|
||||
GetMeteorCastVigorCost());
|
||||
|
||||
attackerUnitsValue += distanceMultiplier * uv;
|
||||
}
|
||||
|
||||
auto attackLocationsForDefender = GetAlCache()->CachedLocations(attackerUnits, isLateGame);
|
||||
const auto &locationsCausingDangerForAttacker = attackLocationsForDefender.AllLocations();
|
||||
|
||||
for (const Unit *unit : defenderUnits) {
|
||||
auto defenderUnitId = unit->unit_id();
|
||||
const int battTypeId = unit->battalion().type();
|
||||
|
||||
auto dv = UnitValue(
|
||||
unit,
|
||||
false,
|
||||
attackerUnits,
|
||||
attackerWantsCastles,
|
||||
/* includeCastleBonus=*/!defenderShouldScatter,
|
||||
defenderUnits,
|
||||
cachedHexMap,
|
||||
roundsRemaining,
|
||||
attackLocationsForDefender,
|
||||
locationsCausingDangerForAttacker,
|
||||
GetApdCache()->GetRaw(
|
||||
cachedHexMap,
|
||||
mapId,
|
||||
GetBattalionType(static_cast<BattalionTypeId>(battTypeId)),
|
||||
false),
|
||||
GetMeteorRange(),
|
||||
GetMeteorCastVigorCost());
|
||||
|
||||
double distanceMultiplier = 1.0;
|
||||
|
||||
// If the defender is trying to scatter, than we want to be as far away from the nearest
|
||||
// attacker as possible, AND as far away from the nearest friendly as possible
|
||||
if (unit->location().row() > -1 && defenderShouldScatter) {
|
||||
CoordsSet myLocationSet(cachedHexMap);
|
||||
myLocationSet.Add(unit->location());
|
||||
|
||||
DIST_T closestDistanceToEnemy = 999;
|
||||
for (const auto &attackerUnit : attackerUnits) {
|
||||
const int attackerBattTypeId = attackerUnit->battalion().type();
|
||||
const DIST_T thisDistance = distanceCache.GetOrCompute(
|
||||
attackerUnit,
|
||||
unit->location(),
|
||||
GetApdCache()->GetRaw(
|
||||
cachedHexMap,
|
||||
mapId,
|
||||
GetBattalionType(static_cast<BattalionTypeId>(attackerBattTypeId)),
|
||||
false),
|
||||
GetBattalionType(static_cast<BattalionTypeId>(attackerBattTypeId))
|
||||
->allowsBraveWater
|
||||
? GetApdCache()->GetRaw(
|
||||
cachedHexMap,
|
||||
mapId,
|
||||
GetBattalionType(
|
||||
static_cast<BattalionTypeId>(attackerBattTypeId)),
|
||||
true,
|
||||
braveWaterCost)
|
||||
: nullptr,
|
||||
cachedHexMap);
|
||||
if (thisDistance < closestDistanceToEnemy) {
|
||||
closestDistanceToEnemy = thisDistance;
|
||||
}
|
||||
}
|
||||
|
||||
// If the best we can do puts us very close to the enemy, and the unit is almost
|
||||
// destroyed, return a negative value; better to flee
|
||||
if (unit->can_flee() && closestDistanceToEnemy < 5 && unit->battalion().size() < 10) {
|
||||
distanceMultiplier = -1;
|
||||
} else {
|
||||
DIST_T closestDistanceToFriendly = 1;
|
||||
if (defenderUnits.size() > 1) {
|
||||
for (const auto &defenderUnit : defenderUnits) {
|
||||
if (defenderUnit->unit_id() != defenderUnitId) {
|
||||
const int defenderBattTypeId = defenderUnit->battalion().type();
|
||||
const DIST_T thisDistance = distanceCache.GetOrCompute(
|
||||
defenderUnit,
|
||||
unit->location(),
|
||||
GetApdCache()->GetRaw(
|
||||
cachedHexMap,
|
||||
mapId,
|
||||
GetBattalionType(static_cast<BattalionTypeId>(
|
||||
defenderBattTypeId)),
|
||||
false),
|
||||
GetBattalionType(
|
||||
static_cast<BattalionTypeId>(defenderBattTypeId))
|
||||
->allowsBraveWater
|
||||
? GetApdCache()->GetRaw(
|
||||
cachedHexMap,
|
||||
mapId,
|
||||
GetBattalionType(static_cast<BattalionTypeId>(
|
||||
defenderBattTypeId)),
|
||||
true,
|
||||
braveWaterCost)
|
||||
: nullptr,
|
||||
cachedHexMap);
|
||||
if (thisDistance < closestDistanceToEnemy) {
|
||||
closestDistanceToFriendly = thisDistance;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
distanceMultiplier =
|
||||
(closestDistanceToEnemy + closestDistanceToFriendly / 5.0) / 5.0;
|
||||
}
|
||||
}
|
||||
|
||||
defenderUnitsValue += distanceMultiplier * dv;
|
||||
}
|
||||
|
||||
defenderUnitsValue *= defenderAdvantage;
|
||||
|
||||
return attackerUnitsValue - defenderUnitsValue;
|
||||
}
|
||||
|
||||
auto StandardAIScoreCalculator::DefenderScatterStrategyScoreForState(
|
||||
const GameStateW &gameState,
|
||||
const int roundsRemaining) const -> ScoreValue {
|
||||
const auto *gameStatePtr = gameState.Get();
|
||||
const auto *status = gameStatePtr->status();
|
||||
|
||||
if (status->state() == net::eagle0::shardok::storage::fb::GameStatus_::State_VICTORY) {
|
||||
const auto *winningIds = status->winning_shardok_ids();
|
||||
const auto *playerInfos = gameStatePtr->player_infos();
|
||||
|
||||
for (const PlayerId winningPid : *winningIds) {
|
||||
if (winningPid < 0) continue;
|
||||
if (playerInfos->Get(winningPid)->is_defender()) return INT_MAX;
|
||||
return INT_MIN;
|
||||
}
|
||||
return INT_MAX;
|
||||
}
|
||||
if (status->state() == net::eagle0::shardok::storage::fb::GameStatus_::State_DRAW) { return 0; }
|
||||
|
||||
const auto *hexMap = gameStatePtr->hex_map();
|
||||
|
||||
const auto mapId = ActionPointDistancesCache::GetMapId(hexMap);
|
||||
|
||||
const auto unitsTotal = -AttackerUnitsScore(
|
||||
gameState,
|
||||
roundsRemaining,
|
||||
/* attackerWantsCastles=*/false,
|
||||
/* defenderShouldScatter=*/true,
|
||||
{},
|
||||
mapId);
|
||||
|
||||
return unitsTotal;
|
||||
}
|
||||
|
||||
auto StandardAIScoreCalculator::DefenderHoldCastlesStrategyScoreForState(
|
||||
const GameStateW &gameState,
|
||||
const CoordsSet &castleCoords,
|
||||
const int roundsRemaining) const -> ScoreValue {
|
||||
const auto unitsTotal = -AttackerUnitsScore(
|
||||
gameState,
|
||||
roundsRemaining,
|
||||
/* attackerWantsCastles=*/true,
|
||||
/*defenderShouldScatter=*/false,
|
||||
{},
|
||||
ActionPointDistancesCache::GetMapId(gameState->hex_map()));
|
||||
|
||||
// Check the victory condition types
|
||||
ScoreValue victoryConditionTotal = 0.0;
|
||||
|
||||
const PlayerInfo *defenderPi = nullptr;
|
||||
for (const PlayerInfo *pi : *gameState->player_infos()) {
|
||||
if (pi->is_defender()) defenderPi = pi;
|
||||
}
|
||||
|
||||
victoryConditionTotal +=
|
||||
DefenderHoldsCriticalTilesVictoryScore(gameState, castleCoords, defenderPi);
|
||||
|
||||
const double unitsMultiplier =
|
||||
static_cast<double>(roundsRemaining) / static_cast<double>(GetMaxRounds());
|
||||
|
||||
return UNITS_BASE_MULTIPLIER * unitsMultiplier * unitsTotal + victoryConditionTotal;
|
||||
}
|
||||
|
||||
auto StandardAIScoreCalculator::DefenderScoreForState(
|
||||
const GameStateW &gameState,
|
||||
const AIStrategy &defenderStrategy,
|
||||
const CoordsSet &castleCoords,
|
||||
const int roundsRemaining) const -> ScoreValue {
|
||||
if (gameState->status()->state() ==
|
||||
net::eagle0::shardok::storage::fb::GameStatus_::State_VICTORY) {
|
||||
for (const PlayerId winningPid : *gameState->status()->winning_shardok_ids()) {
|
||||
if (winningPid < 0) continue;
|
||||
if (defenderStrategy.strategyType == AIStrategy::STRATEGY_FLEE) return 0;
|
||||
if (gameState->player_infos()->Get(winningPid)->is_defender()) return INT_MAX;
|
||||
return INT_MIN;
|
||||
}
|
||||
return INT_MIN;
|
||||
}
|
||||
if (gameState->status()->state() ==
|
||||
net::eagle0::shardok::storage::fb::GameStatus_::State_DRAW) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
switch (defenderStrategy.strategyType) {
|
||||
case AIStrategy::STRATEGY_ATTACK_CASTLES:
|
||||
throw ShardokInternalErrorException("Defender cannot use AttackCastlesStrategy");
|
||||
case AIStrategy::STRATEGY_ATTACK_UNITS:
|
||||
throw ShardokInternalErrorException("Defender cannot use AttackUnitsStrategy");
|
||||
case AIStrategy::STRATEGY_CROSS_RIVERS:
|
||||
throw ShardokInternalErrorException("Defender cannot use CrossRiversStrategy");
|
||||
case AIStrategy::STRATEGY_HOLD_CASTLES:
|
||||
return DefenderHoldCastlesStrategyScoreForState(
|
||||
gameState,
|
||||
castleCoords,
|
||||
roundsRemaining);
|
||||
case AIStrategy::STRATEGY_SCATTER:
|
||||
return DefenderScatterStrategyScoreForState(gameState, roundsRemaining);
|
||||
case AIStrategy::STRATEGY_FLEE:
|
||||
for (const auto *pi : *gameState->player_infos()) {
|
||||
if (pi->is_defender()) {
|
||||
return FleeStrategyScoreForState(gameState, pi->player_id());
|
||||
}
|
||||
}
|
||||
throw ShardokInternalErrorException("Unable to find defender for FleeStrategy");
|
||||
}
|
||||
throw ShardokInternalErrorException("Escaped AIStrategy switch");
|
||||
}
|
||||
|
||||
auto StandardAIScoreCalculator::AttackerScoreForState(
|
||||
const GameStateW &gameState,
|
||||
const AIStrategy &attackerStrategy,
|
||||
const CoordsSet &castleCoords,
|
||||
const int roundsRemaining) const -> ScoreValue {
|
||||
#if PERFORMANCE_LOGGING_
|
||||
AttackerScoreTimer timer;
|
||||
#endif // # PERFORMANCE_LOGGING_
|
||||
if (gameState->status()->state() ==
|
||||
net::eagle0::shardok::storage::fb::GameStatus_::State_VICTORY) {
|
||||
for (const PlayerId winningPid : *gameState->status()->winning_shardok_ids()) {
|
||||
if (winningPid < 0) continue;
|
||||
if (attackerStrategy.strategyType == AIStrategy::STRATEGY_FLEE) return 0;
|
||||
if (gameState->player_infos()->Get(winningPid)->is_defender()) return INT_MIN;
|
||||
return INT_MAX;
|
||||
}
|
||||
return INT_MAX;
|
||||
}
|
||||
if (gameState->status()->state() ==
|
||||
net::eagle0::shardok::storage::fb::GameStatus_::State_DRAW) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
const auto mapId = ActionPointDistancesCache::GetMapId(gameState->hex_map());
|
||||
|
||||
const auto unitsTotal = AttackerUnitsScore(
|
||||
gameState,
|
||||
roundsRemaining,
|
||||
attackerStrategy.strategyType == AIStrategy::STRATEGY_HOLD_CASTLES,
|
||||
/* defenderShouldScatter=*/false,
|
||||
attackerStrategy.targetPriorities,
|
||||
mapId);
|
||||
|
||||
// Check the victory condition types
|
||||
ScoreValue victoryConditionTotal = 0.0;
|
||||
for (const PlayerInfo *pi : *gameState->player_infos()) {
|
||||
if (pi->is_defender()) continue;
|
||||
|
||||
switch (attackerStrategy.strategyType) {
|
||||
case AIStrategy::STRATEGY_CROSS_RIVERS:
|
||||
victoryConditionTotal += WaterCrossingScore(
|
||||
pi->player_id(),
|
||||
[this](BattalionTypeId typeId) { return GetBattalionType(typeId); },
|
||||
gameState,
|
||||
castleCoords,
|
||||
attackerStrategy.targetLocations,
|
||||
GetApdCache());
|
||||
break;
|
||||
|
||||
case AIStrategy::STRATEGY_ATTACK_CASTLES:
|
||||
case AIStrategy::STRATEGY_ATTACK_UNITS:
|
||||
// already factored into AttackerUnitsScore
|
||||
break;
|
||||
|
||||
case AIStrategy::STRATEGY_HOLD_CASTLES:
|
||||
victoryConditionTotal += AttackerHoldsCriticalTilesVictoryScore(
|
||||
gameState,
|
||||
castleCoords,
|
||||
pi,
|
||||
GetApdCache(),
|
||||
GetAlCache(),
|
||||
[this](BattalionTypeId typeId) { return GetBattalionType(typeId); },
|
||||
GetBraveWaterCost());
|
||||
break;
|
||||
|
||||
case AIStrategy::STRATEGY_SCATTER:
|
||||
throw ShardokInternalErrorException("Attacker cannot use ScatterStrategy");
|
||||
|
||||
case AIStrategy::STRATEGY_FLEE:
|
||||
return FleeStrategyScoreForState(gameState, pi->player_id());
|
||||
}
|
||||
}
|
||||
|
||||
const double unitsMultiplier =
|
||||
static_cast<double>(roundsRemaining) / static_cast<double>(GetMaxRounds());
|
||||
|
||||
const double finalScore =
|
||||
UNITS_BASE_MULTIPLIER * unitsMultiplier * unitsTotal + victoryConditionTotal;
|
||||
|
||||
return finalScore;
|
||||
}
|
||||
|
||||
auto StandardAIScoreCalculator::GuessedStateScore(
|
||||
const bool isDefender,
|
||||
const GameStateW &state,
|
||||
const AIStrategy &aiStrategy,
|
||||
const CoordsSet &allCastleCoords) const -> ScoreValue {
|
||||
const int roundsRemaining = GetMaxRounds() - state->current_round();
|
||||
|
||||
if (isDefender) {
|
||||
return DefenderScoreForState(state, aiStrategy, allCastleCoords, roundsRemaining);
|
||||
}
|
||||
return AttackerScoreForState(state, aiStrategy, allCastleCoords, roundsRemaining);
|
||||
}
|
||||
|
||||
// Factory function implementation
|
||||
auto MakeStandardAIScoreCalculator(
|
||||
const SettingsGetter &settingsGetter,
|
||||
const APDCache &apdCache,
|
||||
const ALCache &alCache) -> std::unique_ptr<AIScoreCalculator> {
|
||||
// Extract all battalion types
|
||||
std::unordered_map<BattalionTypeId, BattalionTypeSPtr> battalionTypes;
|
||||
for (int typeId = BattalionTypeId::BattalionTypeId_MIN;
|
||||
typeId <= BattalionTypeId::BattalionTypeId_MAX;
|
||||
typeId++) {
|
||||
auto battalionTypeId = static_cast<BattalionTypeId>(typeId);
|
||||
battalionTypes[battalionTypeId] = settingsGetter.GetBattalionType(battalionTypeId);
|
||||
}
|
||||
|
||||
return std::make_unique<StandardAIScoreCalculator>(
|
||||
settingsGetter.Backing().max_rounds(),
|
||||
settingsGetter.Backing().brave_water_action_point_cost(),
|
||||
settingsGetter.Backing().meteor_range(),
|
||||
settingsGetter.Backing().meteor_cast_vigor_cost(),
|
||||
settingsGetter.Backing().ai_minimum_flee_odds_threshold(),
|
||||
settingsGetter.Backing().ai_desperate_flee_threshold(),
|
||||
std::move(battalionTypes),
|
||||
apdCache,
|
||||
alCache);
|
||||
}
|
||||
|
||||
} // namespace shardok
|
||||
@@ -20,6 +20,5 @@ cc_library(
|
||||
"//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",
|
||||
],
|
||||
)
|
||||
|
||||
@@ -0,0 +1,813 @@
|
||||
# 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.
|
||||
@@ -7,7 +7,7 @@
|
||||
#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/AIScoreCalculator.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"
|
||||
|
||||
@@ -23,9 +23,9 @@ ShardokMCTSAI::ShardokMCTSAI(
|
||||
const ALCache& alCache,
|
||||
MCTSConfig config)
|
||||
: abstractAI_(std::make_unique<mcts::AbstractMCTSAI>(
|
||||
static_cast<mcts::MCTSPlayerId>(playerId),
|
||||
static_cast<mcts::MCTSPlayerId>(
|
||||
playerId), // Use actual player ID for correct scoring
|
||||
config)),
|
||||
playerId_(playerId),
|
||||
isDefender_(isDefender),
|
||||
strategy_(strategy),
|
||||
castleCoords_(castleCoords),
|
||||
@@ -58,12 +58,10 @@ auto ShardokMCTSAI::Search(
|
||||
// Create game engine adapter (passing critical tiles to avoid recomputation)
|
||||
auto gameEngine = mcts::ShardokMCTSFactory::createGameEngine(
|
||||
engine,
|
||||
nullptr, // commandFilter - simplified
|
||||
&scoreCalculator_, // Pass the score calculator
|
||||
settings,
|
||||
apdCache_,
|
||||
alCache_,
|
||||
playerId_,
|
||||
isDefender_,
|
||||
strategy_,
|
||||
castleCoords_,
|
||||
@@ -73,6 +71,11 @@ auto ShardokMCTSAI::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(
|
||||
|
||||
@@ -20,7 +20,6 @@
|
||||
|
||||
#pragma clang diagnostic push
|
||||
#pragma clang diagnostic ignored "-Wdeprecated-redundant-constexpr-static-def"
|
||||
#include "src/main/protobuf/net/eagle0/shardok/api/command_descriptor.pb.h"
|
||||
#pragma clang diagnostic pop
|
||||
|
||||
namespace shardok {
|
||||
@@ -32,7 +31,6 @@ class AIScoreCalculator;
|
||||
|
||||
class ShardokMCTSAI {
|
||||
public:
|
||||
using CommandProto = net::eagle0::shardok::api::CommandDescriptor;
|
||||
using SearchResult = IterativeDeepeningAI::SearchResult;
|
||||
using MCTSConfig = mcts::MCTSConfig;
|
||||
|
||||
@@ -60,7 +58,6 @@ private:
|
||||
std::unique_ptr<mcts::AbstractMCTSAI> abstractAI_;
|
||||
|
||||
// Shardok-specific context
|
||||
PlayerId playerId_;
|
||||
bool isDefender_;
|
||||
AIStrategy strategy_;
|
||||
const CoordsSet& castleCoords_;
|
||||
|
||||
@@ -11,7 +11,7 @@ cc_library(
|
||||
],
|
||||
deps = [
|
||||
"//src/main/cpp/net/eagle0/common/mcts/abstract:mcts_action",
|
||||
"//src/main/protobuf/net/eagle0/shardok/api:command_descriptor_cc_proto",
|
||||
"//src/main/cpp/net/eagle0/shardok/library:shardok_c_types",
|
||||
"//src/main/protobuf/net/eagle0/shardok/common:command_type_cc_proto",
|
||||
],
|
||||
)
|
||||
@@ -27,11 +27,12 @@ cc_library(
|
||||
],
|
||||
deps = [
|
||||
"//src/main/cpp/net/eagle0/common/mcts/abstract:mcts_game_state",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_score_calculator_interface",
|
||||
"//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",
|
||||
],
|
||||
)
|
||||
|
||||
@@ -50,7 +51,8 @@ cc_library(
|
||||
"//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_score_calculator_interface",
|
||||
"//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",
|
||||
@@ -72,10 +74,9 @@ cc_library(
|
||||
":shardok_game_engine",
|
||||
":shardok_game_state",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_command_filter",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_score_calculator_interface",
|
||||
"//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/api:command_descriptor_cc_proto",
|
||||
],
|
||||
)
|
||||
|
||||
@@ -11,61 +11,78 @@
|
||||
#include "src/main/protobuf/net/eagle0/shardok/common/command_type.pb.h"
|
||||
#pragma clang diagnostic pop
|
||||
|
||||
namespace shardok {
|
||||
namespace mcts {
|
||||
namespace shardok::mcts {
|
||||
|
||||
ShardokAction::ShardokAction(const CommandProto& command, size_t index)
|
||||
: command_(command),
|
||||
commandIndex_(index) {}
|
||||
|
||||
ShardokAction::ShardokAction(CommandProto&& command, size_t index)
|
||||
: command_(std::move(command)),
|
||||
commandIndex_(index) {}
|
||||
|
||||
int ShardokAction::getType() const { return static_cast<int>(command_.type()); }
|
||||
// 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>(command_.player()) << " ";
|
||||
ss << "P" << static_cast<int>(player_) << " ";
|
||||
|
||||
ss << net::eagle0::shardok::common::CommandType_Name(command_.type());
|
||||
ss << net::eagle0::shardok::common::CommandType_Name(type_);
|
||||
|
||||
if (command_.has_actor()) { ss << " Unit:" << command_.actor().value(); }
|
||||
if (actorId_ >= 0) { ss << " Unit:" << actorId_; }
|
||||
|
||||
if (command_.has_target()) {
|
||||
const auto& coord = command_.target();
|
||||
ss << " @(" << coord.row() << "," << coord.column() << ")";
|
||||
if (targetRow_ >= 0 && targetCol_ >= 0) {
|
||||
ss << " @(" << targetRow_ << "," << targetCol_ << ")";
|
||||
}
|
||||
|
||||
return ss.str();
|
||||
}
|
||||
|
||||
int ShardokAction::getActorId() const {
|
||||
if (command_.has_actor()) { return command_.actor().value(); }
|
||||
return -1;
|
||||
}
|
||||
|
||||
std::pair<int, int> ShardokAction::getTarget() const {
|
||||
if (command_.has_target()) {
|
||||
const auto& coord = command_.target();
|
||||
return std::make_pair(coord.row(), coord.column());
|
||||
}
|
||||
return std::make_pair(-1, -1);
|
||||
}
|
||||
|
||||
std::unique_ptr<MCTSAction> ShardokAction::clone() const {
|
||||
return std::make_unique<ShardokAction>(command_, commandIndex_);
|
||||
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; }
|
||||
|
||||
return commandIndex_ == shardokOther->commandIndex_ &&
|
||||
command_.SerializeAsString() == shardokOther->command_.SerializeAsString();
|
||||
// Compare by index only - actions from same command list are uniquely identified by index
|
||||
return commandIndex_ == shardokOther->commandIndex_;
|
||||
}
|
||||
|
||||
} // namespace mcts
|
||||
} // namespace shardok
|
||||
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
|
||||
@@ -9,42 +9,53 @@
|
||||
#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/api/command_descriptor.pb.h"
|
||||
#include "src/main/protobuf/net/eagle0/shardok/common/command_type.pb.h"
|
||||
#pragma clang diagnostic pop
|
||||
|
||||
namespace shardok {
|
||||
namespace mcts {
|
||||
namespace shardok::mcts {
|
||||
|
||||
class ShardokAction : public MCTSAction {
|
||||
public:
|
||||
using CommandProto = net::eagle0::shardok::api::CommandDescriptor;
|
||||
using CommandType = net::eagle0::shardok::common::CommandType;
|
||||
|
||||
ShardokAction(const CommandProto& command, size_t index);
|
||||
ShardokAction(CommandProto&& command, size_t index);
|
||||
// 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 methods (not part of abstract interface)
|
||||
[[nodiscard]] int getType() const;
|
||||
[[nodiscard]] int getActorId() const;
|
||||
[[nodiscard]] std::pair<int, int> getTarget() const;
|
||||
|
||||
// Shardok-specific accessor
|
||||
[[nodiscard]] const CommandProto& getCommand() const { return command_; }
|
||||
// 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:
|
||||
CommandProto command_;
|
||||
// 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 mcts
|
||||
} // namespace shardok
|
||||
} // namespace shardok::mcts
|
||||
|
||||
#endif // EAGLE0_SHARDOK_ACTION_HPP
|
||||
@@ -4,49 +4,66 @@
|
||||
|
||||
#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/AIScoreCalculator.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"
|
||||
|
||||
#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 {
|
||||
|
||||
// Static helper for average random generator
|
||||
// Note: Using nullptr for simplicity - could be improved with proper random generator
|
||||
// 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 AICommandFilter* commandFilter,
|
||||
const AIScoreCalculator* scoreCalculator,
|
||||
const GameSettingsSPtr& gameSettings,
|
||||
const APDCache* apdCache,
|
||||
const ALCache* alCache,
|
||||
PlayerId playerId,
|
||||
bool isDefender,
|
||||
const AIStrategy& strategy,
|
||||
const CoordsSet& castleCoords,
|
||||
const CoordsSet& criticalTileCoords)
|
||||
: commandFilter_(commandFilter),
|
||||
scoreCalculator_(scoreCalculator),
|
||||
: scoreCalculator_(scoreCalculator),
|
||||
gameSettings_(gameSettings),
|
||||
apdCache_(apdCache),
|
||||
alCache_(alCache),
|
||||
playerId_(playerId),
|
||||
isDefender_(isDefender),
|
||||
strategy_(strategy),
|
||||
castleCoords_(castleCoords),
|
||||
criticalTileCoords_(criticalTileCoords) {}
|
||||
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) const {
|
||||
const MCTSAction& action,
|
||||
double deterministicRoll) const {
|
||||
const auto* shardokState = dynamic_cast<const ShardokGameState*>(&state);
|
||||
const auto* shardokAction = dynamic_cast<const ShardokAction*>(&action);
|
||||
|
||||
@@ -76,10 +93,55 @@ std::unique_ptr<MCTSGameState> ShardokGameEngine::applyAction(
|
||||
engine = std::make_shared<ShardokEngine>(*engine);
|
||||
}
|
||||
|
||||
engine->PostCommand(currentPlayer, shardokAction->getIndex(), nullptr);
|
||||
// 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;
|
||||
|
||||
// Create and return the new state (don't cache the mutated engine)
|
||||
return std::make_unique<ShardokGameState>(
|
||||
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(),
|
||||
@@ -89,6 +151,15 @@ std::unique_ptr<MCTSGameState> ShardokGameEngine::applyAction(
|
||||
*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(
|
||||
@@ -125,19 +196,92 @@ void ShardokGameEngine::applyActionMutable(
|
||||
|
||||
engine->PostCommand(currentPlayer, shardokAction->getIndex(), nullptr);
|
||||
shardokState->getMutableShardokState() = engine->GetCurrentGameState();
|
||||
// Clear the cached engine since the state has been mutated
|
||||
// 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) const {
|
||||
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());
|
||||
|
||||
// Stop simulation if it's not our player's turn (turn boundary)
|
||||
if (currentPlayer != playerId_) { return {}; }
|
||||
// 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;
|
||||
@@ -174,10 +318,61 @@ std::vector<std::unique_ptr<MCTSAction>> ShardokGameEngine::getLegalActions(
|
||||
for (const size_t idx : filteredIndices) {
|
||||
if (idx < commands->size()) {
|
||||
const auto& cmd = commands->at(idx);
|
||||
actions.push_back(std::make_unique<ShardokAction>(cmd->GetCommandProto(), 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;
|
||||
}
|
||||
|
||||
@@ -189,30 +384,92 @@ double ShardokGameEngine::evaluateState(const MCTSGameState& state, MCTSPlayerId
|
||||
|
||||
std::vector<size_t> ShardokGameEngine::filterActions(
|
||||
const std::vector<std::unique_ptr<MCTSAction>>& actions,
|
||||
const MCTSGameState& state) const {
|
||||
const auto* shardokState = dynamic_cast<const ShardokGameState*>(&state);
|
||||
if (!shardokState || !commandFilter_) {
|
||||
// No filtering - return all indices
|
||||
std::vector<size_t> indices;
|
||||
indices.reserve(actions.size());
|
||||
for (size_t i = 0; i < actions.size(); ++i) { indices.push_back(i); }
|
||||
return indices;
|
||||
}
|
||||
|
||||
// For now, return all indices - proper filtering would need more work
|
||||
// to match the AICommandFilter::FilterCommands signature
|
||||
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);
|
||||
}
|
||||
|
||||
@@ -228,7 +485,7 @@ size_t ShardokGameEngine::mapFilteredIndexToOriginal(
|
||||
size_t filteredIndex,
|
||||
const MCTSGameState& state) const {
|
||||
// Get the filtered actions (uses cached engine)
|
||||
auto actions = getLegalActions(state);
|
||||
auto actions = getLegalActions(state, state.currentPlayerId(), 0, 0);
|
||||
|
||||
// Check bounds
|
||||
if (filteredIndex >= actions.size()) { return filteredIndex; }
|
||||
@@ -241,4 +498,131 @@ size_t ShardokGameEngine::mapFilteredIndexToOriginal(
|
||||
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
|
||||
|
||||
@@ -5,14 +5,15 @@
|
||||
#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/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/map/CoordsSet.hpp"
|
||||
@@ -34,12 +35,10 @@ class ShardokGameEngine : public MCTSGameEngine {
|
||||
public:
|
||||
ShardokGameEngine(
|
||||
const ShardokEngine* engine,
|
||||
const AICommandFilter* commandFilter,
|
||||
const AIScoreCalculator* scoreCalculator,
|
||||
const GameSettingsSPtr& gameSettings,
|
||||
const APDCache* apdCache,
|
||||
const ALCache* alCache,
|
||||
PlayerId playerId,
|
||||
bool isDefender,
|
||||
const AIStrategy& strategy,
|
||||
const CoordsSet& castleCoords,
|
||||
@@ -48,13 +47,17 @@ public:
|
||||
// MCTSGameEngine interface implementation
|
||||
[[nodiscard]] std::unique_ptr<MCTSGameState> applyAction(
|
||||
const MCTSGameState& state,
|
||||
const MCTSAction& action) const override;
|
||||
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) const override;
|
||||
const MCTSGameState& state,
|
||||
MCTSPlayerId rootPlayerId,
|
||||
int currentPlayerFlips,
|
||||
int maxPlayerFlips) const override;
|
||||
|
||||
[[nodiscard]] bool isTerminal(const MCTSGameState& state) const override;
|
||||
|
||||
@@ -65,6 +68,10 @@ public:
|
||||
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,
|
||||
@@ -79,18 +86,56 @@ public:
|
||||
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:
|
||||
const AICommandFilter* commandFilter_;
|
||||
// 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_;
|
||||
GameSettingsSPtr gameSettings_;
|
||||
const GameSettingsSPtr gameSettings_;
|
||||
const APDCache* apdCache_;
|
||||
const ALCache* alCache_;
|
||||
PlayerId playerId_;
|
||||
bool isDefender_;
|
||||
AIStrategy strategy_;
|
||||
const CoordsSet& castleCoords_;
|
||||
const CoordsSet castleCoords_; // Own the data to avoid dangling references
|
||||
// Computed once to avoid 8.5% overhead per engine construction
|
||||
const CoordsSet& criticalTileCoords_;
|
||||
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
|
||||
|
||||
@@ -8,7 +8,7 @@
|
||||
#include <string>
|
||||
#include <utility>
|
||||
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/AIScoreCalculator.hpp"
|
||||
#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 {
|
||||
@@ -41,8 +41,32 @@ uint64_t ShardokGameState::hash() const {
|
||||
return cachedHash_;
|
||||
}
|
||||
|
||||
double ShardokGameState::score(MCTSPlayerId /*playerId*/) const {
|
||||
return scoreCalculator_->GuessedStateScore(isDefender_, state_, strategy_, castleCoords_);
|
||||
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(); }
|
||||
|
||||
@@ -56,18 +56,24 @@ public:
|
||||
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_;
|
||||
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_;
|
||||
const CoordsSet criticalTileCoords_; // Own the data to avoid dangling references
|
||||
mutable std::shared_ptr<ShardokEngine> cachedEngine_; // Engine with cached available commands
|
||||
};
|
||||
|
||||
|
||||
@@ -14,24 +14,20 @@ namespace shardok::mcts {
|
||||
|
||||
std::unique_ptr<MCTSGameEngine> ShardokMCTSFactory::createGameEngine(
|
||||
const ShardokEngine& engine,
|
||||
const AICommandFilter* commandFilter,
|
||||
const AIScoreCalculator* scoreCalculator,
|
||||
const GameSettingsSPtr& gameSettings,
|
||||
const APDCache& apdCache,
|
||||
const ALCache& alCache,
|
||||
PlayerId playerId,
|
||||
bool isDefender,
|
||||
const AIStrategy& strategy,
|
||||
const CoordsSet& castleCoords,
|
||||
const CoordsSet& criticalTileCoords) {
|
||||
return std::make_unique<ShardokGameEngine>(
|
||||
&engine,
|
||||
commandFilter,
|
||||
scoreCalculator,
|
||||
gameSettings,
|
||||
&apdCache,
|
||||
&alCache,
|
||||
playerId,
|
||||
isDefender,
|
||||
strategy,
|
||||
castleCoords,
|
||||
@@ -60,18 +56,6 @@ std::unique_ptr<MCTSGameState> ShardokMCTSFactory::createGameState(
|
||||
criticalTileCoords);
|
||||
}
|
||||
|
||||
std::vector<std::unique_ptr<MCTSAction>> ShardokMCTSFactory::createActions(
|
||||
const std::vector<CommandProto>& commands) {
|
||||
std::vector<std::unique_ptr<MCTSAction>> actions;
|
||||
actions.reserve(commands.size());
|
||||
|
||||
for (size_t i = 0; i < commands.size(); ++i) {
|
||||
actions.push_back(std::make_unique<ShardokAction>(commands[i], i));
|
||||
}
|
||||
|
||||
return actions;
|
||||
}
|
||||
|
||||
std::vector<std::unique_ptr<MCTSAction>> ShardokMCTSFactory::createActionsFromCommandList(
|
||||
const CommandListSPtr& commands) {
|
||||
std::vector<std::unique_ptr<MCTSAction>> actions;
|
||||
@@ -80,7 +64,16 @@ std::vector<std::unique_ptr<MCTSAction>> ShardokMCTSFactory::createActionsFromCo
|
||||
actions.reserve(commands->size());
|
||||
for (size_t i = 0; i < commands->size(); ++i) {
|
||||
const auto& cmd = (*commands)[i];
|
||||
actions.push_back(std::make_unique<ShardokAction>(cmd->GetCommandProto(), 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;
|
||||
|
||||
@@ -16,7 +16,6 @@
|
||||
#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"
|
||||
|
||||
namespace shardok {
|
||||
|
||||
@@ -27,14 +26,6 @@ class AIScoreCalculator;
|
||||
class GameStateW;
|
||||
class GameSettings;
|
||||
|
||||
// Use existing type definitions to avoid conflicts
|
||||
// These are already defined in the Shardok codebase:
|
||||
// - APDCache in ActionPointDistancesCache.hpp
|
||||
// - ALCache in AIAttackLocations.hpp
|
||||
// - CommandListSPtr in ShardokCommand.hpp
|
||||
// - SettingsGetter in GameSettings.hpp
|
||||
using CommandProto = net::eagle0::shardok::api::CommandDescriptor;
|
||||
|
||||
namespace mcts {
|
||||
|
||||
// Forward declarations
|
||||
@@ -44,17 +35,13 @@ class MCTSAction;
|
||||
|
||||
class ShardokMCTSFactory {
|
||||
public:
|
||||
using CommandProto = net::eagle0::shardok::api::CommandDescriptor;
|
||||
|
||||
// Create a Shardok game engine adapter
|
||||
[[nodiscard]] static std::unique_ptr<MCTSGameEngine> createGameEngine(
|
||||
const ShardokEngine& engine,
|
||||
const AICommandFilter* commandFilter,
|
||||
const AIScoreCalculator* scoreCalculator,
|
||||
const GameSettingsSPtr& gameSettings,
|
||||
const APDCache& apdCache,
|
||||
const ALCache& alCache,
|
||||
PlayerId playerId,
|
||||
bool isDefender,
|
||||
const AIStrategy& strategy,
|
||||
const CoordsSet& castleCoords,
|
||||
@@ -72,10 +59,6 @@ public:
|
||||
const ALCache& alCache,
|
||||
const CoordsSet& criticalTileCoords);
|
||||
|
||||
// Convert Shardok commands to MCTS actions
|
||||
[[nodiscard]] static std::vector<std::unique_ptr<MCTSAction>> createActions(
|
||||
const std::vector<CommandProto>& commands);
|
||||
|
||||
// Convert from command list to MCTS actions
|
||||
[[nodiscard]] static std::vector<std::unique_ptr<MCTSAction>> createActionsFromCommandList(
|
||||
const CommandListSPtr& commands);
|
||||
|
||||
-2
@@ -12,7 +12,6 @@
|
||||
#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/api/command_descriptor.pb.h"
|
||||
|
||||
namespace shardok {
|
||||
|
||||
@@ -21,7 +20,6 @@ using std::future;
|
||||
using std::vector;
|
||||
|
||||
using ScoreValue = double;
|
||||
using CommandProto = net::eagle0::shardok::api::CommandDescriptor;
|
||||
|
||||
// Forward declarations
|
||||
class ShardokEngine;
|
||||
+2
-2
@@ -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"
|
||||
|
||||
@@ -0,0 +1,112 @@
|
||||
load("//tools:copts.bzl", "COPTS")
|
||||
|
||||
cc_library(
|
||||
name = "ai_score_calculator_interface",
|
||||
hdrs = ["AIScoreCalculator.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__",
|
||||
],
|
||||
deps = [
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_attack_locations",
|
||||
"//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/map:coords_set",
|
||||
],
|
||||
)
|
||||
|
||||
cc_library(
|
||||
name = "ai_victory_condition_score_calculator",
|
||||
srcs = ["AIVictoryConditionScoreCalculator.cpp"],
|
||||
hdrs = ["AIVictoryConditionScoreCalculator.hpp"],
|
||||
copts = COPTS,
|
||||
visibility = [
|
||||
"//src/main/cpp/net/eagle0/shardok/ai/score/private:__pkg__", # Needed by abstract base class
|
||||
"//src/main/cpp/net/eagle0/shardok/ai_performance_runner:__pkg__",
|
||||
"//src/test/cpp/net/eagle0/shardok/ai:__subpackages__",
|
||||
],
|
||||
deps = [
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_attack_groups",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_attack_locations",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_common_types",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_distance_debuf",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai: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 = "normalized_ai_score_calculator",
|
||||
srcs = ["NormalizedAIScoreCalculator.cpp"],
|
||||
hdrs = ["NormalizedAIScoreCalculator.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__",
|
||||
],
|
||||
deps = [
|
||||
":ai_score_calculator_interface",
|
||||
":ai_victory_condition_score_calculator",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_strategy",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_unit_score_calculator",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_water_crossing_calculator",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai/score/private:abstract_ai_score_calculator",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai/score/private:ai_score_calculator_shared_utilities",
|
||||
"//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/settings:game_settings",
|
||||
],
|
||||
)
|
||||
|
||||
cc_library(
|
||||
name = "standard_ai_score_calculator",
|
||||
srcs = ["StandardAIScoreCalculator.cpp"],
|
||||
hdrs = ["StandardAIScoreCalculator.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__",
|
||||
],
|
||||
deps = [
|
||||
":ai_score_calculator_interface",
|
||||
":ai_victory_condition_score_calculator",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_strategy",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_unit_score_calculator",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_water_crossing_calculator",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai/score/private:abstract_ai_score_calculator",
|
||||
"//src/main/cpp/net/eagle0/shardok/library:engine",
|
||||
"//src/main/cpp/net/eagle0/shardok/library:game_state_w",
|
||||
],
|
||||
)
|
||||
|
||||
cc_library(
|
||||
name = "mcts_optimized_ai_score_calculator",
|
||||
srcs = ["MCTSOptimizedAIScoreCalculator.cpp"],
|
||||
hdrs = ["MCTSOptimizedAIScoreCalculator.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__",
|
||||
],
|
||||
deps = [
|
||||
":ai_score_calculator_interface",
|
||||
":ai_victory_condition_score_calculator",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_strategy",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_unit_score_calculator",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_water_crossing_calculator",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai/score/private:abstract_ai_score_calculator",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai/score/private:ai_score_calculator_shared_utilities",
|
||||
"//src/main/cpp/net/eagle0/shardok/library:engine",
|
||||
"//src/main/cpp/net/eagle0/shardok/library:game_state_w",
|
||||
],
|
||||
)
|
||||
@@ -0,0 +1,248 @@
|
||||
//
|
||||
// MCTS-Optimized implementation of AIScoreCalculator
|
||||
// Uses bounded linear scoring tuned for MCTS exploration/exploitation balance
|
||||
//
|
||||
|
||||
#include "MCTSOptimizedAIScoreCalculator.hpp"
|
||||
|
||||
#include <algorithm>
|
||||
#include <unordered_map>
|
||||
|
||||
#include "private/AIScoreCalculatorSharedUtilities.hpp"
|
||||
#include "private/AbstractAIScoreCalculator.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/AIAttackLocations.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/AIScoreUtilities.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/AIStrategy.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/AIUnitScoreCalculator.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/action_point_distances/ActionPointDistancesCache.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
|
||||
|
||||
namespace shardok {
|
||||
|
||||
using net::eagle0::shardok::storage::fb::BattalionTypeId;
|
||||
|
||||
// Bring shared utilities into scope
|
||||
using score_calculator_internal::FleeStrategyScoreForState;
|
||||
using score_calculator_internal::UnitsScoreComponents;
|
||||
|
||||
// Scoring constants tuned for MCTS
|
||||
namespace {
|
||||
|
||||
// Scale constants designed to produce score differences in the range that works well for MCTS
|
||||
// With C=1.41 and typical parent visits ~10000, exploration term ≈ 0.42
|
||||
// We want:
|
||||
// - Early game tactical moves: 0.3-0.5 difference (ratio 0.7-1.2x exploration)
|
||||
// - Mid game advantages (10-30%): 4.0-8.0 difference (ratio 9-19x exploration)
|
||||
// - Late game crushing advantages: 10-40 difference (ratio 24-95x exploration)
|
||||
|
||||
// Maximum contribution from proportional unit advantage (applies to both battle sizes)
|
||||
// A 100% unit advantage (all attacker, no defender) produces ±80 score
|
||||
constexpr double UNITS_SCORE_SCALE = 80.0;
|
||||
|
||||
// Minimum reference value to avoid division by zero in edge cases
|
||||
constexpr double MIN_REFERENCE_VALUE = 1000.0;
|
||||
|
||||
// IMPORTANT: Victory condition scores are NOT normalized by army size
|
||||
// They represent absolute strategic goals (castle control, etc.) that should not
|
||||
// diminish as more units are placed. Typical range: -3000 to +3000 (raw).
|
||||
// Scaling factor of 0.01 brings them to -30 to +30 range.
|
||||
|
||||
} // anonymous namespace
|
||||
|
||||
/// MCTS-Optimized implementation of AIScoreCalculator.
|
||||
/// Produces bounded linear scores that balance MCTS exploration and exploitation.
|
||||
/// Inherits from AbstractAIScoreCalculator to share common functionality.
|
||||
class MCTSOptimizedAIScoreCalculator : public AbstractAIScoreCalculator {
|
||||
public:
|
||||
MCTSOptimizedAIScoreCalculator(
|
||||
int maxRounds,
|
||||
ActionPoints braveWaterCost,
|
||||
int meteorRange,
|
||||
double meteorCastVigorCost,
|
||||
int minimumFleeOddsThreshold,
|
||||
int desperateFleeThreshold,
|
||||
std::vector<BattalionTypeSPtr> battalionTypes,
|
||||
const APDCache &apdCache,
|
||||
const ALCache &alCache)
|
||||
: AbstractAIScoreCalculator(
|
||||
maxRounds,
|
||||
braveWaterCost,
|
||||
meteorRange,
|
||||
meteorCastVigorCost,
|
||||
minimumFleeOddsThreshold,
|
||||
desperateFleeThreshold,
|
||||
std::move(battalionTypes),
|
||||
apdCache,
|
||||
alCache) {}
|
||||
|
||||
[[nodiscard]] auto GuessedStateScore(
|
||||
bool isDefender,
|
||||
const GameStateW &state,
|
||||
const AIStrategy &aiStrategy,
|
||||
const CoordsSet &allCastleCoords) const -> ScoreValue override;
|
||||
|
||||
// Implement pure virtual methods from AbstractAIScoreCalculator
|
||||
[[nodiscard]] auto InterpretDefenderOutcome(GameOutcome outcome) const -> ScoreValue override;
|
||||
[[nodiscard]] auto InterpretAttackerOutcome(GameOutcome outcome) const -> ScoreValue override;
|
||||
|
||||
[[nodiscard]] auto AttackerFleeStrategyScoreForState(const GameStateW &gameState) const
|
||||
-> ScoreValue override;
|
||||
|
||||
[[nodiscard]] auto CombineAttackerScores(
|
||||
const UnitsScoreComponents &components,
|
||||
double victoryConditionScore,
|
||||
int roundsRemaining) const -> ScoreValue override;
|
||||
|
||||
[[nodiscard]] auto CombineDefenderScatterScores(const UnitsScoreComponents &components) const
|
||||
-> ScoreValue override;
|
||||
|
||||
[[nodiscard]] auto CombineDefenderHoldCastlesScores(
|
||||
const UnitsScoreComponents &components,
|
||||
double victoryConditionScore,
|
||||
int roundsRemaining) const -> ScoreValue override;
|
||||
|
||||
private:
|
||||
[[nodiscard]] auto DefenderFleeStrategyScoreForState(const GameStateW &gameState) const
|
||||
-> ScoreValue override;
|
||||
};
|
||||
|
||||
// Implementation of MCTSOptimizedAIScoreCalculator methods
|
||||
|
||||
auto MCTSOptimizedAIScoreCalculator::InterpretDefenderOutcome(GameOutcome outcome) const
|
||||
-> ScoreValue {
|
||||
// Use bounded values instead of INT_MAX/MIN for numerical stability
|
||||
switch (outcome) {
|
||||
case GameOutcome::DEFENDER_VICTORY: return 1000.0;
|
||||
case GameOutcome::ATTACKER_VICTORY: return -1000.0;
|
||||
case GameOutcome::DRAW: return 0.0;
|
||||
case GameOutcome::FLEE_OUTCOME: return 0.0;
|
||||
}
|
||||
throw ShardokInternalErrorException("Unknown GameOutcome");
|
||||
}
|
||||
|
||||
auto MCTSOptimizedAIScoreCalculator::InterpretAttackerOutcome(GameOutcome outcome) const
|
||||
-> ScoreValue {
|
||||
// Use bounded values instead of INT_MAX/MIN for numerical stability
|
||||
switch (outcome) {
|
||||
case GameOutcome::ATTACKER_VICTORY: return 1000.0;
|
||||
case GameOutcome::DEFENDER_VICTORY: return -1000.0;
|
||||
case GameOutcome::DRAW: return 0.0;
|
||||
case GameOutcome::FLEE_OUTCOME: return 0.0;
|
||||
}
|
||||
throw ShardokInternalErrorException("Unknown GameOutcome");
|
||||
}
|
||||
|
||||
auto MCTSOptimizedAIScoreCalculator::CombineAttackerScores(
|
||||
const UnitsScoreComponents &components,
|
||||
const double victoryConditionScore,
|
||||
const int roundsRemaining) const -> ScoreValue {
|
||||
// Use actual total army value as reference (scales with battle size)
|
||||
const double totalArmyValue = components.attackerUnitsValue + components.defenderUnitsValue;
|
||||
const double reference = std::max(totalArmyValue, MIN_REFERENCE_VALUE);
|
||||
|
||||
// Normalize proportional unit difference to approximately [-80, +80] range
|
||||
const double unitsDiff = components.attackerUnitsValue - components.defenderUnitsValue;
|
||||
const double unitsScore = (unitsDiff / reference) * UNITS_SCORE_SCALE;
|
||||
|
||||
// Victory condition score is an absolute strategic value, not normalized by army size
|
||||
// Scaling factor to bring victory scores into similar magnitude as unit scores
|
||||
const double victoryScore = victoryConditionScore * 0.01;
|
||||
|
||||
// Weight units by rounds remaining (early: units matter less, late: units dominate)
|
||||
const double unitsMultiplier =
|
||||
static_cast<double>(roundsRemaining) / static_cast<double>(GetMaxRounds());
|
||||
|
||||
return unitsMultiplier * unitsScore + victoryScore;
|
||||
}
|
||||
|
||||
auto MCTSOptimizedAIScoreCalculator::CombineDefenderScatterScores(
|
||||
const UnitsScoreComponents &components) const -> ScoreValue {
|
||||
// Use actual total army value as reference
|
||||
const double totalArmyValue = components.attackerUnitsValue + components.defenderUnitsValue;
|
||||
const double reference = std::max(totalArmyValue, MIN_REFERENCE_VALUE);
|
||||
|
||||
// For scatter strategy, just maximize proportional defender advantage
|
||||
const double unitsDiff = components.defenderUnitsValue - components.attackerUnitsValue;
|
||||
return (unitsDiff / reference) * UNITS_SCORE_SCALE;
|
||||
}
|
||||
|
||||
auto MCTSOptimizedAIScoreCalculator::CombineDefenderHoldCastlesScores(
|
||||
const UnitsScoreComponents &components,
|
||||
const double victoryConditionScore,
|
||||
const int roundsRemaining) const -> ScoreValue {
|
||||
// Use actual total army value as reference
|
||||
const double totalArmyValue = components.attackerUnitsValue + components.defenderUnitsValue;
|
||||
const double reference = std::max(totalArmyValue, MIN_REFERENCE_VALUE);
|
||||
|
||||
// Similar to attacker, but from defender's perspective
|
||||
const double unitsDiff = components.defenderUnitsValue - components.attackerUnitsValue;
|
||||
const double unitsScore = (unitsDiff / reference) * UNITS_SCORE_SCALE;
|
||||
|
||||
// Victory condition score is an absolute strategic value, not normalized by army size
|
||||
// Scaling factor to bring victory scores into similar magnitude as unit scores
|
||||
const double victoryScore = victoryConditionScore * 0.01;
|
||||
|
||||
const double unitsMultiplier =
|
||||
static_cast<double>(roundsRemaining) / static_cast<double>(GetMaxRounds());
|
||||
|
||||
return unitsMultiplier * unitsScore + victoryScore;
|
||||
}
|
||||
|
||||
auto MCTSOptimizedAIScoreCalculator::DefenderFleeStrategyScoreForState(
|
||||
const GameStateW &gameState) const -> ScoreValue {
|
||||
for (const auto *pi : *gameState->player_infos()) {
|
||||
if (pi->is_defender()) { return FleeStrategyScoreForState(gameState, pi->player_id()); }
|
||||
}
|
||||
throw ShardokInternalErrorException("Unable to find defender for FleeStrategy");
|
||||
}
|
||||
|
||||
auto MCTSOptimizedAIScoreCalculator::AttackerFleeStrategyScoreForState(
|
||||
const GameStateW &gameState) const -> ScoreValue {
|
||||
for (const PlayerInfo *pi : *gameState->player_infos()) {
|
||||
if (!pi->is_defender()) { return FleeStrategyScoreForState(gameState, pi->player_id()); }
|
||||
}
|
||||
throw ShardokInternalErrorException("Unable to find attacker for FleeStrategy");
|
||||
}
|
||||
|
||||
auto MCTSOptimizedAIScoreCalculator::GuessedStateScore(
|
||||
const bool isDefender,
|
||||
const GameStateW &state,
|
||||
const AIStrategy &aiStrategy,
|
||||
const CoordsSet &allCastleCoords) const -> ScoreValue {
|
||||
const int roundsRemaining = GetMaxRounds() - state->current_round();
|
||||
|
||||
if (isDefender) {
|
||||
return DefenderScoreForState(state, aiStrategy, allCastleCoords, roundsRemaining);
|
||||
}
|
||||
return AttackerScoreForState(state, aiStrategy, allCastleCoords, roundsRemaining);
|
||||
}
|
||||
|
||||
// Factory function implementation
|
||||
auto MakeMCTSOptimizedAIScoreCalculator(
|
||||
const SettingsGetter &settingsGetter,
|
||||
const APDCache &apdCache,
|
||||
const ALCache &alCache) -> std::unique_ptr<AIScoreCalculator> {
|
||||
// Extract all battalion types into a vector indexed by BattalionTypeId
|
||||
std::vector<BattalionTypeSPtr> battalionTypes(BattalionTypeId::BattalionTypeId_MAX + 1);
|
||||
for (int typeId = BattalionTypeId::BattalionTypeId_MIN;
|
||||
typeId <= BattalionTypeId::BattalionTypeId_MAX;
|
||||
typeId++) {
|
||||
auto battalionTypeId = static_cast<BattalionTypeId>(typeId);
|
||||
battalionTypes[battalionTypeId] = settingsGetter.GetBattalionType(battalionTypeId);
|
||||
}
|
||||
|
||||
return std::make_unique<MCTSOptimizedAIScoreCalculator>(
|
||||
settingsGetter.Backing().max_rounds(),
|
||||
settingsGetter.Backing().brave_water_action_point_cost(),
|
||||
settingsGetter.Backing().meteor_range(),
|
||||
settingsGetter.Backing().meteor_cast_vigor_cost(),
|
||||
settingsGetter.Backing().ai_minimum_flee_odds_threshold(),
|
||||
settingsGetter.Backing().ai_desperate_flee_threshold(),
|
||||
std::move(battalionTypes),
|
||||
apdCache,
|
||||
alCache);
|
||||
}
|
||||
|
||||
} // namespace shardok
|
||||
@@ -0,0 +1,32 @@
|
||||
//
|
||||
// MCTS-Optimized implementation of AIScoreCalculator
|
||||
// Uses bounded linear scoring tuned for MCTS exploration/exploitation balance
|
||||
//
|
||||
|
||||
#ifndef EAGLE0_MCTSOPTIMIZEDAISCORECALCULATOR_HPP
|
||||
#define EAGLE0_MCTSOPTIMIZEDAISCORECALCULATOR_HPP
|
||||
|
||||
#include <memory>
|
||||
|
||||
#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/settings/GameSettings.hpp"
|
||||
|
||||
namespace shardok {
|
||||
|
||||
// Forward declarations
|
||||
class AIScoreCalculator;
|
||||
|
||||
using APDCache = std::shared_ptr<ActionPointDistancesCache>;
|
||||
using ALCache = std::unique_ptr<AttackLocationsCache>;
|
||||
|
||||
/// Factory function to create an MCTSOptimizedAIScoreCalculator.
|
||||
/// Returns a unique_ptr to AIScoreCalculator to hide the implementation.
|
||||
[[nodiscard]] auto MakeMCTSOptimizedAIScoreCalculator(
|
||||
const SettingsGetter& settingsGetter,
|
||||
const APDCache& apdCache,
|
||||
const ALCache& alCache) -> std::unique_ptr<AIScoreCalculator>;
|
||||
|
||||
} // namespace shardok
|
||||
|
||||
#endif // EAGLE0_MCTSOPTIMIZEDAISCORECALCULATOR_HPP
|
||||
@@ -0,0 +1,388 @@
|
||||
# MCTS-Optimized Scoring Algorithm Design
|
||||
|
||||
## Problem Statement
|
||||
|
||||
We need a scoring algorithm that makes MCTS perform well by providing score differences in the right range:
|
||||
|
||||
- **Standard Scorer**: Returns unbounded relative scores. Small differences get amplified in MCTS UCB formula, causing over-exploitation (commits to 1-2 high-scoring nodes too early).
|
||||
- **Normalized Scorer**: Returns scores in [0,1] range with power transformation (exponent=0.1). Differences are too compressed (~0.02-0.05), causing over-exploration (all nodes explored equally, AI makes bad choices).
|
||||
|
||||
### MCTS Requirements
|
||||
|
||||
After ~100 visits, the **exploitation term** (cumulative_score / visits) should be comparable to the **exploration term** (C * sqrt(ln(parent_visits) / visits)).
|
||||
|
||||
With C=1.41 and typical parent visits ~10000:
|
||||
- Exploration term: 1.41 * sqrt(ln(10000) / 100) ≈ 0.42
|
||||
- **Target exploitation differences: 0.5 to 2.0**
|
||||
|
||||
This means individual scores should differ by **0.5 to 2.0** between meaningfully different positions.
|
||||
|
||||
## Scale Analysis from Codebase
|
||||
|
||||
### Unit Values
|
||||
- Single unit context-free value: 500-3000 (depends on battalion type, stats, size)
|
||||
- With modifiers (castle, terrain, ranged): 1000-6000 per unit
|
||||
- Full army (10 units max): 10,000-40,000
|
||||
- Typical strong army: ~20,000
|
||||
|
||||
### Victory Condition Scores
|
||||
- Castle held by defender: -(battalion_size + vigor) * distance_debuf ≈ -800 per castle
|
||||
- Distance debuf: 0.0 (adjacent) to 1.0 (unreachable), typically 0.9-0.95
|
||||
- 3 castles held by defender at medium distance: ≈ -2400
|
||||
- Range: 0 (all captured) to -3000 (all held, far away)
|
||||
|
||||
### Terminal States
|
||||
- Victory: INT_MAX (or 1.0 for normalized)
|
||||
- Defeat: INT_MIN (or 0.0 for normalized)
|
||||
- Flee/Draw: 0 (or 0.5 for normalized)
|
||||
- Captured unit: -10,000
|
||||
- Captured VIP: -25,000
|
||||
|
||||
## Proposed Algorithm: Bounded Linear Scorer
|
||||
|
||||
### Design Principles
|
||||
|
||||
1. **Normalized scale**: Map scores to approximately [-15, +15] range
|
||||
2. **Separate components**: Units and victory conditions contribute separately
|
||||
3. **Preserve relative importance**: Victory conditions dominate early, units become important as advantage grows
|
||||
4. **Round-based weighting**: Similar to Standard scorer, weight units by rounds remaining
|
||||
|
||||
### Constants
|
||||
|
||||
```cpp
|
||||
constexpr double UNITS_SCORE_SCALE = 80.0; // Max contribution from proportional unit advantage
|
||||
constexpr double VICTORY_SCORE_SCALE = 400.0; // Normalizer for victory conditions (also proportional)
|
||||
constexpr double MIN_REFERENCE_VALUE = 1000.0; // Avoid division by zero in edge cases
|
||||
```
|
||||
|
||||
**Key insights**:
|
||||
1. Both unit scores AND victory condition scores scale proportionally with battle size (victory scores use battalion.size() in their calculation). Therefore, we normalize both by the **actual total army value** rather than a fixed reference.
|
||||
|
||||
2. **MCTS requires stronger signal than minimax**: Minimax (Iterative Deepening) just picks argmax, so even tiny score differences (0.01) work fine. MCTS needs score differences comparable to the exploration term (~0.4-0.5) to guide search effectively. We use 10x larger scale constants to amplify tactical differences like positioning, distance to objectives, and incremental unit advantages.
|
||||
|
||||
### Attacker Score Formula
|
||||
|
||||
```cpp
|
||||
auto CombineAttackerScores(
|
||||
const UnitsScoreComponents &components,
|
||||
double victoryConditionScore,
|
||||
int roundsRemaining) const -> ScoreValue {
|
||||
|
||||
// Use actual total army value as reference (scales with battle size)
|
||||
const double totalArmyValue = components.attackerUnitsValue + components.defenderUnitsValue;
|
||||
const double reference = std::max(totalArmyValue, MIN_REFERENCE_VALUE);
|
||||
|
||||
// Normalize proportional unit difference to [-8, +8] range
|
||||
const double unitsDiff = components.attackerUnitsValue - components.defenderUnitsValue;
|
||||
const double unitsScore = (unitsDiff / reference) * UNITS_SCORE_SCALE;
|
||||
|
||||
// Normalize victory condition (also proportional to army size) to approximately [-10, 0] range
|
||||
const double victoryScore = (victoryConditionScore / reference) * VICTORY_SCORE_SCALE;
|
||||
|
||||
// Weight units by rounds remaining (early: units matter less, late: units dominate)
|
||||
const double unitsMultiplier =
|
||||
static_cast<double>(roundsRemaining) / static_cast<double>(GetMaxRounds());
|
||||
|
||||
return unitsMultiplier * unitsScore + victoryScore;
|
||||
}
|
||||
```
|
||||
|
||||
### Defender Score Formula
|
||||
|
||||
```cpp
|
||||
auto CombineDefenderScatterScores(
|
||||
const UnitsScoreComponents &components) const -> ScoreValue {
|
||||
|
||||
// Use actual total army value as reference
|
||||
const double totalArmyValue = components.attackerUnitsValue + components.defenderUnitsValue;
|
||||
const double reference = std::max(totalArmyValue, MIN_REFERENCE_VALUE);
|
||||
|
||||
// For scatter strategy, just maximize proportional defender advantage
|
||||
const double unitsDiff = components.defenderUnitsValue - components.attackerUnitsValue;
|
||||
return (unitsDiff / reference) * UNITS_SCORE_SCALE;
|
||||
}
|
||||
|
||||
auto CombineDefenderHoldCastlesScores(
|
||||
const UnitsScoreComponents &components,
|
||||
double victoryConditionScore,
|
||||
int roundsRemaining) const -> ScoreValue {
|
||||
|
||||
// Use actual total army value as reference
|
||||
const double totalArmyValue = components.attackerUnitsValue + components.defenderUnitsValue;
|
||||
const double reference = std::max(totalArmyValue, MIN_REFERENCE_VALUE);
|
||||
|
||||
// Similar to attacker, but from defender's perspective
|
||||
const double unitsDiff = components.defenderUnitsValue - components.attackerUnitsValue;
|
||||
const double unitsScore = (unitsDiff / reference) * UNITS_SCORE_SCALE;
|
||||
const double victoryScore = (victoryConditionScore / reference) * VICTORY_SCORE_SCALE;
|
||||
|
||||
const double unitsMultiplier =
|
||||
static_cast<double>(roundsRemaining) / static_cast<double>(GetMaxRounds());
|
||||
|
||||
return unitsMultiplier * unitsScore + victoryScore;
|
||||
}
|
||||
```
|
||||
|
||||
### Terminal States
|
||||
|
||||
```cpp
|
||||
auto InterpretAttackerOutcome(GameOutcome outcome) const -> ScoreValue {
|
||||
switch (outcome) {
|
||||
case GameOutcome::ATTACKER_VICTORY: return 1000.0; // Large but bounded
|
||||
case GameOutcome::DEFENDER_VICTORY: return -1000.0;
|
||||
case GameOutcome::DRAW: return 0.0;
|
||||
case GameOutcome::FLEE_OUTCOME: return 0.0;
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Note: Using bounded values (±1000) instead of INT_MAX/MIN ensures numerical stability in MCTS and clearer signal that these are terminal states.
|
||||
|
||||
## Example Score Traces
|
||||
|
||||
### Large Battle Scenarios (10v10, ~40000 total army)
|
||||
|
||||
#### Scenario 1: Even armies, attacker needs to capture 3 castles
|
||||
- Units: attacker 20000, defender 20000 (total: 40000)
|
||||
- Victory: -2850 (3 castles * 950 each, medium distance)
|
||||
- Rounds: 15/30 remaining
|
||||
|
||||
Score:
|
||||
- reference = 40000
|
||||
- unitsDiff = 0
|
||||
- unitsScore = 0
|
||||
- victoryScore = (-2850 / 40000) * 400.0 = -28.5
|
||||
- unitsMultiplier = 0.5
|
||||
- **total = 0.5 * 0 + (-28.5) = -28.5**
|
||||
|
||||
#### Scenario 2: Slight attacker advantage (10%)
|
||||
- Units: attacker 22000, defender 18000 (diff: +4000, total: 40000)
|
||||
- Victory: -2850
|
||||
- Rounds: 15/30
|
||||
|
||||
Score:
|
||||
- unitsScore = (4000 / 40000) * 80.0 = 8.0
|
||||
- victoryScore = -28.5
|
||||
- unitsMultiplier = 0.5
|
||||
- **total = 0.5 * 8.0 + (-28.5) = -24.5**
|
||||
- **Difference from Scenario 1: 4.0** ✓
|
||||
|
||||
#### Scenario 3: Large attacker advantage (30%)
|
||||
- Units: attacker 26000, defender 14000 (diff: +12000, total: 40000)
|
||||
- Victory: -2850
|
||||
- Rounds: 15/30
|
||||
|
||||
Score:
|
||||
- unitsScore = (12000 / 40000) * 80.0 = 24.0
|
||||
- victoryScore = -28.5
|
||||
- unitsMultiplier = 0.5
|
||||
- **total = 0.5 * 24.0 + (-28.5) = -16.5**
|
||||
- **Difference from Scenario 2: 8.0** ✓
|
||||
|
||||
### Small Battle Scenarios (2v2, ~4000 total army)
|
||||
|
||||
#### Scenario 4: Even small armies, 1 castle
|
||||
- Units: attacker 2000, defender 2000 (total: 4000)
|
||||
- Victory: -475 (1 castle * 500 * 0.95 distance)
|
||||
- Rounds: 15/30
|
||||
|
||||
Score:
|
||||
- reference = 4000
|
||||
- unitsScore = 0
|
||||
- victoryScore = (-475 / 4000) * 400.0 = -47.5
|
||||
- **total = 0.5 * 0 + (-47.5) = -47.5**
|
||||
|
||||
#### Scenario 5: Slight advantage in small battle (10%)
|
||||
- Units: attacker 2200, defender 1800 (diff: +400, total: 4000)
|
||||
- Victory: -475
|
||||
- Rounds: 15/30
|
||||
|
||||
Score:
|
||||
- unitsScore = (400 / 4000) * 80.0 = 8.0
|
||||
- victoryScore = -47.5
|
||||
- **total = 0.5 * 8.0 + (-47.5) = -43.5**
|
||||
- **Difference from Scenario 4: 4.0** ✓
|
||||
|
||||
### Early Game Scenario: Single unit movement
|
||||
|
||||
#### Scenario 6: Early game, single unit advances toward castle
|
||||
- Units: attacker 20000, defender 20000 (total: 40000)
|
||||
- Victory before: -2850 (distance debuf = 0.95)
|
||||
- Victory after: -2829 (distance debuf = 0.943, one unit moved closer)
|
||||
- Change in victory score: +21
|
||||
- Rounds: 28/30 (early game)
|
||||
|
||||
Score change:
|
||||
- victoryScoreChange = (21 / 40000) * 400.0 = 0.21
|
||||
- Additionally, the moving unit (value 2000) gets better distance multiplier:
|
||||
- Before: 2000 * 0.25 = 500
|
||||
- After: 2000 * 0.279 = 558
|
||||
- Diff = 58, normalized: (58 / 40000) * 80.0 = 0.116
|
||||
- unitsMultiplier = 28/30 = 0.933
|
||||
- **Total improvement: 0.21 + 0.933 * 0.116 = 0.32** ✓
|
||||
|
||||
With exploration term ~0.42, this gives exploitation/exploration ratio of **0.76** - still below 1.0 but much better than before (was 0.05). MCTS will slightly prefer better moves while still exploring alternatives.
|
||||
|
||||
### Scale Consistency Verification
|
||||
|
||||
Comparing **10% advantage** in both battle sizes:
|
||||
- Large battle (Scenario 2): diff = **4.0**
|
||||
- Small battle (Scenario 5): diff = **4.0**
|
||||
|
||||
**Perfect scaling!** Same proportional advantage → same score difference, regardless of battle size.
|
||||
|
||||
Early game tactical moves now produce meaningful signals (0.3-0.5 range) that guide MCTS while still allowing healthy exploration.
|
||||
|
||||
## MCTS Behavior Verification
|
||||
|
||||
After 100 visits with C=1.41, exploration term ~0.42:
|
||||
|
||||
**Early game (single unit tactical moves):**
|
||||
- Good positioning move: **0.32** (ratio 0.76x exploration)
|
||||
- MCTS explores broadly but slightly favors better moves
|
||||
|
||||
**Mid game (unit advantages matter):**
|
||||
- 10% army advantage: **4.0** (ratio 9.5x exploration)
|
||||
- 30% army advantage: **8.0** (ratio 19x exploration)
|
||||
- MCTS strongly commits to maintaining/increasing army advantage
|
||||
|
||||
**Late game (large differences):**
|
||||
- Major strategic advantages: **10-40** (ratio 24-95x exploration)
|
||||
- MCTS decisively exploits winning positions
|
||||
|
||||
This progression is ideal:
|
||||
- **Early game**: Healthy exploration (ratio < 1.0) when moves are genuinely similar
|
||||
- **Mid game**: Strong exploitation (ratio 9-19x) when clear advantages exist
|
||||
- **Late game**: Decisive exploitation (ratio > 20x) to close out wins
|
||||
|
||||
This avoids both pathologies:
|
||||
- Not over-exploiting (like Standard scorer which overcommitted to tiny early differences)
|
||||
- Not over-exploring (like Normalized scorer which explored equally even with large advantages)
|
||||
|
||||
## Why MCTS Needs Stronger Signal Than Minimax
|
||||
|
||||
**Iterative Deepening (minimax)** works fine with tiny score differences (0.01-0.1) because:
|
||||
- It explores all moves to the same depth
|
||||
- It simply picks `argmax(scores)`
|
||||
- Even a 0.01 difference causes it to prefer the better move
|
||||
|
||||
**MCTS** needs much larger differences (0.3-4.0) because:
|
||||
- It uses UCB formula: `score/visits + C*sqrt(ln(parent_visits)/visits)`
|
||||
- The exploration term (~0.4) can dominate small exploitation differences
|
||||
- With differences < 0.1, MCTS explores all moves almost equally (over-exploration)
|
||||
- With differences > 10.0, MCTS commits too early (over-exploitation)
|
||||
|
||||
**Solution**: Use 10x larger scale constants than initially designed, specifically tuned so that:
|
||||
- Early game tactical moves (positioning, distance) produce 0.3-0.5 differences
|
||||
- Mid game advantages (10-30% army strength) produce 4.0-8.0 differences
|
||||
- Late game crushing advantages produce 10-40 differences
|
||||
|
||||
This gives MCTS the right balance: explore when moves are similar, exploit when advantages are clear.
|
||||
|
||||
## Implementation Notes
|
||||
|
||||
1. **Use same calculation structure**: Inherit from AbstractAIScoreCalculator like Standard and Normalized
|
||||
2. **Reuse unit scoring**: Use existing CalculateUnitsScoreComponents and victory condition calculators
|
||||
3. **Only change combination**: Override CombineAttackerScores, CombineDefenderScores, etc.
|
||||
4. **Bounded terminals**: Use ±1000 instead of INT_MAX/MIN for numerical stability
|
||||
5. **No transformation**: Unlike Normalized, don't apply power transformation - linear scaling is sufficient
|
||||
6. **Scale constants tuned for MCTS**: 10x larger than naive normalization to provide appropriate signal strength
|
||||
|
||||
## Testing with Integration Tests
|
||||
|
||||
Before integrating with MCTS, test the new scorer with **IterativeDeepeningAI** using the AI integration test infrastructure.
|
||||
|
||||
### Integration Test Infrastructure
|
||||
|
||||
The codebase now has comprehensive AI integration tests in `src/test/cpp/net/eagle0/shardok/ai/AIIntegrationTest.cpp` that use:
|
||||
|
||||
1. **AIPerformanceTestHelpers** (`src/test/cpp/net/eagle0/shardok/library/AIPerformanceTestHelpers.{cpp,hpp}`):
|
||||
- `CreatePerfTestGameState(settings, defenderToggle)` creates a 6v6 scenario on the Alah map
|
||||
- Properly initializes units with correct battalion sizes (800 for longbowmen, capacity-based for others)
|
||||
- Handles both attacker and defender perspectives
|
||||
- Returns GameStateW in SETUP phase with 6 units per player in reserve
|
||||
|
||||
2. **ShardokAIClient** integration:
|
||||
- Tests use the full AI client interface, not just the search algorithm
|
||||
- Time budgets set to 3s for reasonable test execution time
|
||||
- Handles both setup phase placement and first turn movement
|
||||
|
||||
3. **Acceptable Position Sets** for handling AI non-determinism:
|
||||
- AI decisions may vary due to internal tie-breaking and search order
|
||||
- Tests define sets of acceptable positions for each unit
|
||||
- Example from AttackerAI_Setup_PlacesUnitsCorrectly:
|
||||
```cpp
|
||||
std::set<net::eagle0::shardok::storage::fb::Coords> acceptablePositions{
|
||||
net::eagle0::shardok::storage::fb::Coords(0, 11),
|
||||
net::eagle0::shardok::storage::fb::Coords(1, 10),
|
||||
// ... more acceptable positions
|
||||
};
|
||||
```
|
||||
|
||||
### Adding Tests for New Scorers
|
||||
|
||||
To test MCTSOptimizedAIScoreCalculator (or any new scorer) with IterativeDeepeningAI:
|
||||
|
||||
1. **Add test cases following the existing pattern** in `AIIntegrationTest.cpp`:
|
||||
```cpp
|
||||
TEST(MCTSOptimizedScorerTest, AttackerAI_Setup_PlacesUnitsCorrectly) {
|
||||
auto settings = GetDefaultGameSettingsForTest();
|
||||
auto gameStateW = CreatePerfTestGameState(settings, /*defenderToggle=*/false);
|
||||
auto hexMap = gameStateW.GetHexMap().ToProto();
|
||||
|
||||
// Use MCTSOptimizedAIScoreCalculator instead of StandardAIScoreCalculator
|
||||
auto scoreCalculator = std::make_shared<MCTSOptimizedAIScoreCalculator>(
|
||||
/*playerId=*/0, /*isDefender=*/false, hexMap, settings->GetGetter());
|
||||
|
||||
ShardokAIClient client(
|
||||
/*playerId=*/0, /*isDefender=*/false, hexMap, settings,
|
||||
scoreCalculator, std::chrono::milliseconds(3000));
|
||||
|
||||
// ... rest of test follows existing pattern
|
||||
}
|
||||
```
|
||||
|
||||
2. **Update BUILD.bazel** to add the new scorer as a dependency:
|
||||
```bazel
|
||||
deps = [
|
||||
# ... existing deps ...
|
||||
"//src/main/cpp/net/eagle0/shardok/ai/score:mcts_optimized_ai_score_calculator",
|
||||
]
|
||||
```
|
||||
|
||||
3. **Test patterns to implement**:
|
||||
- **Setup Phase Tests**: Verify AI places units in reasonable starting positions
|
||||
- `AttackerAI_Setup_PlacesUnitsCorrectly`: Attacker should place at start zone (0,11)-(1,13)
|
||||
- `DefenderAI_Setup_OccupiesCastles`: Defender should occupy castle tiles
|
||||
- **First Turn Tests**: Verify AI makes sensible initial moves
|
||||
- `AttackerAI_FirstTurn_MovesUnitsCorrectly`: Attacker should advance toward objectives
|
||||
- Use acceptable position sets to handle non-determinism
|
||||
- **Score Range Verification**: Add assertions to verify scores are in expected ranges
|
||||
```cpp
|
||||
// Example: verify scores are bounded as expected
|
||||
auto searchResult = client.GetBestCommand(gameStateW);
|
||||
EXPECT_GE(searchResult.score, -50.0); // Reasonable lower bound
|
||||
EXPECT_LE(searchResult.score, 50.0); // Reasonable upper bound
|
||||
```
|
||||
|
||||
4. **Performance Regression Testing**:
|
||||
- Run `./scripts/ai_perf_test.sh` to verify the new scorer doesn't cause performance degradation
|
||||
- Compare commands evaluated at each depth vs. StandardAIScoreCalculator
|
||||
- See CLAUDE.md "Performance Testing" section for detailed instructions
|
||||
|
||||
### Why Test with IterativeDeepeningAI First
|
||||
|
||||
The new scoring algorithm should work with **both** IterativeDeepeningAI and MCTS:
|
||||
- If it fails with IterativeDeepeningAI, the scoring logic itself is broken
|
||||
- If it passes with IterativeDeepeningAI but fails with MCTS, the issue is MCTS-specific
|
||||
- This allows incremental testing and debugging
|
||||
|
||||
Once the scorer passes integration tests with IterativeDeepeningAI, then integrate with MCTS and compare behavior.
|
||||
|
||||
## Alternative Names
|
||||
|
||||
- `BoundedLinearAIScoreCalculator`
|
||||
- `MCTSOptimizedAIScoreCalculator`
|
||||
- `LinearNormalizedAIScoreCalculator`
|
||||
|
||||
Recommend: **`MCTSOptimizedAIScoreCalculator`** to clearly indicate purpose.
|
||||
@@ -0,0 +1,252 @@
|
||||
//
|
||||
// Normalized [0,1] implementation of AIScoreCalculator
|
||||
//
|
||||
|
||||
#include "NormalizedAIScoreCalculator.hpp"
|
||||
|
||||
#include <cmath>
|
||||
#include <unordered_map>
|
||||
|
||||
#include "private/AIScoreCalculatorSharedUtilities.hpp"
|
||||
#include "private/AbstractAIScoreCalculator.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/AIAttackLocations.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/AIScoreUtilities.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/AIStrategy.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/AIUnitScoreCalculator.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/action_point_distances/ActionPointDistancesCache.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
|
||||
|
||||
namespace shardok {
|
||||
|
||||
using net::eagle0::shardok::storage::fb::BattalionTypeId;
|
||||
|
||||
// Bring shared utilities into scope
|
||||
using score_calculator_internal::UnitsScoreComponents;
|
||||
|
||||
/// Normalized implementation of AIScoreCalculator that produces scores in [0, 1] range.
|
||||
/// Inherits from AbstractAIScoreCalculator to share common functionality.
|
||||
class NormalizedAIScoreCalculator : public AbstractAIScoreCalculator {
|
||||
public:
|
||||
NormalizedAIScoreCalculator(
|
||||
int maxRounds,
|
||||
ActionPoints braveWaterCost,
|
||||
int meteorRange,
|
||||
double meteorCastVigorCost,
|
||||
int minimumFleeOddsThreshold,
|
||||
int desperateFleeThreshold,
|
||||
std::vector<BattalionTypeSPtr> battalionTypes,
|
||||
const APDCache &apdCache,
|
||||
const ALCache &alCache)
|
||||
: AbstractAIScoreCalculator(
|
||||
maxRounds,
|
||||
braveWaterCost,
|
||||
meteorRange,
|
||||
meteorCastVigorCost,
|
||||
minimumFleeOddsThreshold,
|
||||
desperateFleeThreshold,
|
||||
std::move(battalionTypes),
|
||||
apdCache,
|
||||
alCache) {}
|
||||
|
||||
[[nodiscard]] auto GuessedStateScore(
|
||||
bool isDefender,
|
||||
const GameStateW &state,
|
||||
const AIStrategy &aiStrategy,
|
||||
const CoordsSet &allCastleCoords) const -> ScoreValue override;
|
||||
|
||||
// Implement pure virtual methods from AbstractAIScoreCalculator
|
||||
[[nodiscard]] auto InterpretDefenderOutcome(GameOutcome outcome) const -> ScoreValue override;
|
||||
[[nodiscard]] auto InterpretAttackerOutcome(GameOutcome outcome) const -> ScoreValue override;
|
||||
|
||||
[[nodiscard]] auto AttackerFleeStrategyScoreForState(const GameStateW &gameState) const
|
||||
-> ScoreValue override;
|
||||
|
||||
[[nodiscard]] auto CombineAttackerScores(
|
||||
const UnitsScoreComponents &components,
|
||||
double victoryConditionScore,
|
||||
int roundsRemaining) const -> ScoreValue override;
|
||||
|
||||
[[nodiscard]] auto CombineDefenderScatterScores(const UnitsScoreComponents &components) const
|
||||
-> ScoreValue override;
|
||||
|
||||
[[nodiscard]] auto CombineDefenderHoldCastlesScores(
|
||||
const UnitsScoreComponents &components,
|
||||
double victoryConditionScore,
|
||||
int roundsRemaining) const -> ScoreValue override;
|
||||
|
||||
private:
|
||||
[[nodiscard]] auto DefenderFleeStrategyScoreForState(const GameStateW &gameState) const
|
||||
-> ScoreValue override;
|
||||
|
||||
/// Applies power transformation to spread out compressed scores for MCTS.
|
||||
/// Maps [0,1] → [0,1] but pushes values away from 0.5 toward the extremes.
|
||||
/// Terminal states (0.0, 1.0) are unchanged.
|
||||
[[nodiscard]] auto TransformForMCTS(ScoreValue score) const -> ScoreValue;
|
||||
};
|
||||
|
||||
// Implementation of NormalizedAIScoreCalculator methods
|
||||
|
||||
auto NormalizedAIScoreCalculator::InterpretDefenderOutcome(GameOutcome outcome) const
|
||||
-> ScoreValue {
|
||||
switch (outcome) {
|
||||
case GameOutcome::DEFENDER_VICTORY: return 1.0;
|
||||
case GameOutcome::ATTACKER_VICTORY: return 0.0;
|
||||
case GameOutcome::DRAW: return 0.5;
|
||||
case GameOutcome::FLEE_OUTCOME: return 0.5;
|
||||
}
|
||||
throw ShardokInternalErrorException("Unknown GameOutcome");
|
||||
}
|
||||
|
||||
auto NormalizedAIScoreCalculator::InterpretAttackerOutcome(GameOutcome outcome) const
|
||||
-> ScoreValue {
|
||||
switch (outcome) {
|
||||
case GameOutcome::ATTACKER_VICTORY: return 1.0;
|
||||
case GameOutcome::DEFENDER_VICTORY: return 0.0;
|
||||
case GameOutcome::DRAW: return 0.5;
|
||||
case GameOutcome::FLEE_OUTCOME: return 0.5;
|
||||
}
|
||||
throw ShardokInternalErrorException("Unknown GameOutcome");
|
||||
}
|
||||
|
||||
auto NormalizedAIScoreCalculator::CombineDefenderScatterScores(
|
||||
const UnitsScoreComponents &components) const -> ScoreValue {
|
||||
// For defender, we flip the perspective: defenderValue is (1), attackerValue is (2)
|
||||
const double defenderValue = components.defenderUnitsValue;
|
||||
const double attackerValue = components.attackerUnitsValue;
|
||||
|
||||
// No victory condition for scatter strategy
|
||||
const double denominator = defenderValue + attackerValue;
|
||||
if (denominator == 0.0) { return 0.5; }
|
||||
|
||||
return defenderValue / denominator;
|
||||
}
|
||||
|
||||
auto NormalizedAIScoreCalculator::CombineDefenderHoldCastlesScores(
|
||||
const UnitsScoreComponents &components,
|
||||
const double victoryConditionScore,
|
||||
const int /*roundsRemaining*/) const -> ScoreValue {
|
||||
// For defender, flip perspective
|
||||
const double defenderValue = components.defenderUnitsValue;
|
||||
const double attackerValue = components.attackerUnitsValue;
|
||||
|
||||
// Apply normalization
|
||||
double numerator;
|
||||
double denominator;
|
||||
|
||||
if (victoryConditionScore >= 0) {
|
||||
numerator = defenderValue + victoryConditionScore;
|
||||
denominator = defenderValue + attackerValue + victoryConditionScore;
|
||||
} else {
|
||||
numerator = defenderValue;
|
||||
denominator = defenderValue + attackerValue - victoryConditionScore;
|
||||
}
|
||||
|
||||
if (denominator == 0.0) { return 0.5; }
|
||||
|
||||
return numerator / denominator;
|
||||
}
|
||||
|
||||
auto NormalizedAIScoreCalculator::DefenderFleeStrategyScoreForState(
|
||||
const GameStateW & /*gameState*/) const -> ScoreValue {
|
||||
// FLEE strategy doesn't fit the [0,1] model well - return 0.5
|
||||
return 0.5;
|
||||
}
|
||||
|
||||
auto NormalizedAIScoreCalculator::AttackerFleeStrategyScoreForState(
|
||||
const GameStateW & /*gameState*/) const -> ScoreValue {
|
||||
// FLEE strategy doesn't fit the [0,1] model well - return 0.5
|
||||
return 0.5;
|
||||
}
|
||||
|
||||
auto NormalizedAIScoreCalculator::CombineAttackerScores(
|
||||
const UnitsScoreComponents &components,
|
||||
const double victoryConditionScore,
|
||||
const int /*roundsRemaining*/) const -> ScoreValue {
|
||||
const double attackerUnitsValue = components.attackerUnitsValue;
|
||||
const double defenderUnitsValue = components.defenderUnitsValue;
|
||||
|
||||
// Apply normalization formula
|
||||
double numerator;
|
||||
double denominator;
|
||||
|
||||
if (victoryConditionScore >= 0) {
|
||||
// Positive victory condition: add to numerator
|
||||
numerator = attackerUnitsValue + victoryConditionScore;
|
||||
denominator = attackerUnitsValue + defenderUnitsValue + victoryConditionScore;
|
||||
} else {
|
||||
// Negative victory condition: subtract from denominator (making it larger)
|
||||
numerator = attackerUnitsValue;
|
||||
denominator = attackerUnitsValue + defenderUnitsValue - victoryConditionScore;
|
||||
}
|
||||
|
||||
// Handle edge case of all zeros
|
||||
if (denominator == 0.0) { return 0.5; }
|
||||
|
||||
return numerator / denominator;
|
||||
}
|
||||
|
||||
auto NormalizedAIScoreCalculator::TransformForMCTS(ScoreValue score) const -> ScoreValue {
|
||||
// Power transformation exponent - lower values spread scores more toward extremes
|
||||
// Tuned for MCTS: balances exploration vs exploitation
|
||||
// - Too low (e.g., 0.3): over-exploitation like standard scorer
|
||||
// - Too high (e.g., 0.9): over-exploration like untransformed normalized
|
||||
// - 0.6-0.7: sweet spot for MCTS
|
||||
constexpr double EXPONENT = 0.1;
|
||||
|
||||
if (score > 0.5) {
|
||||
// Map [0.5, 1.0] → [0.5, 1.0] with power curve
|
||||
// (score - 0.5) * 2.0 maps to [0, 1], apply power, then scale back
|
||||
return 0.5 + 0.5 * std::pow((score - 0.5) * 2.0, EXPONENT);
|
||||
} else {
|
||||
// Map [0.0, 0.5] → [0.0, 0.5] with power curve (symmetric)
|
||||
return 0.5 - 0.5 * std::pow((0.5 - score) * 2.0, EXPONENT);
|
||||
}
|
||||
}
|
||||
|
||||
auto NormalizedAIScoreCalculator::GuessedStateScore(
|
||||
const bool isDefender,
|
||||
const GameStateW &state,
|
||||
const AIStrategy &aiStrategy,
|
||||
const CoordsSet &allCastleCoords) const -> ScoreValue {
|
||||
const int roundsRemaining = GetMaxRounds() - state->current_round();
|
||||
|
||||
ScoreValue rawScore;
|
||||
if (isDefender) {
|
||||
rawScore = DefenderScoreForState(state, aiStrategy, allCastleCoords, roundsRemaining);
|
||||
} else {
|
||||
rawScore = AttackerScoreForState(state, aiStrategy, allCastleCoords, roundsRemaining);
|
||||
}
|
||||
|
||||
// Apply power transformation to spread out scores for MCTS
|
||||
return TransformForMCTS(rawScore);
|
||||
}
|
||||
|
||||
// Factory function implementation
|
||||
auto MakeNormalizedAIScoreCalculator(
|
||||
const SettingsGetter &settingsGetter,
|
||||
const APDCache &apdCache,
|
||||
const ALCache &alCache) -> std::unique_ptr<AIScoreCalculator> {
|
||||
// Extract all battalion types into a vector indexed by BattalionTypeId
|
||||
std::vector<BattalionTypeSPtr> battalionTypes(BattalionTypeId::BattalionTypeId_MAX + 1);
|
||||
for (int typeId = BattalionTypeId::BattalionTypeId_MIN;
|
||||
typeId <= BattalionTypeId::BattalionTypeId_MAX;
|
||||
typeId++) {
|
||||
auto battalionTypeId = static_cast<BattalionTypeId>(typeId);
|
||||
battalionTypes[battalionTypeId] = settingsGetter.GetBattalionType(battalionTypeId);
|
||||
}
|
||||
|
||||
return std::make_unique<NormalizedAIScoreCalculator>(
|
||||
settingsGetter.Backing().max_rounds(),
|
||||
settingsGetter.Backing().brave_water_action_point_cost(),
|
||||
settingsGetter.Backing().meteor_range(),
|
||||
settingsGetter.Backing().meteor_cast_vigor_cost(),
|
||||
settingsGetter.Backing().ai_minimum_flee_odds_threshold(),
|
||||
settingsGetter.Backing().ai_desperate_flee_threshold(),
|
||||
std::move(battalionTypes),
|
||||
apdCache,
|
||||
alCache);
|
||||
}
|
||||
|
||||
} // namespace shardok
|
||||
@@ -0,0 +1,44 @@
|
||||
//
|
||||
// Normalized [0,1] implementation of AIScoreCalculator
|
||||
//
|
||||
|
||||
#ifndef EAGLE0_NORMALIZEDAISCORECALCULATOR_HPP
|
||||
#define EAGLE0_NORMALIZEDAISCORECALCULATOR_HPP
|
||||
|
||||
#include <memory>
|
||||
|
||||
#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/settings/GameSettings.hpp"
|
||||
|
||||
namespace shardok {
|
||||
|
||||
// Forward declarations
|
||||
class AIScoreCalculator;
|
||||
|
||||
using APDCache = std::shared_ptr<ActionPointDistancesCache>;
|
||||
using ALCache = std::unique_ptr<AttackLocationsCache>;
|
||||
|
||||
/// Factory function to create a NormalizedAIScoreCalculator.
|
||||
/// Returns a unique_ptr to AIScoreCalculator to hide the implementation.
|
||||
///
|
||||
/// The normalized scorer produces scores in the range [0, 1] where:
|
||||
/// - 0.0 = complete defender victory
|
||||
/// - 1.0 = complete attacker victory
|
||||
/// - 0.5 = neutral/draw state
|
||||
///
|
||||
/// Terminal states (victory/defeat) always return 1.0 or 0.0.
|
||||
/// Non-terminal states use asymmetric normalization:
|
||||
/// - If victory condition >= 0:
|
||||
/// score = (attackerUnits + victoryCondition) / (attackerUnits + defenderUnits +
|
||||
/// victoryCondition)
|
||||
/// - If victory condition < 0:
|
||||
/// score = attackerUnits / (attackerUnits + defenderUnits - victoryCondition)
|
||||
[[nodiscard]] auto MakeNormalizedAIScoreCalculator(
|
||||
const SettingsGetter& settingsGetter,
|
||||
const APDCache& apdCache,
|
||||
const ALCache& alCache) -> std::unique_ptr<AIScoreCalculator>;
|
||||
|
||||
} // namespace shardok
|
||||
|
||||
#endif // EAGLE0_NORMALIZEDAISCORECALCULATOR_HPP
|
||||
@@ -0,0 +1,289 @@
|
||||
//
|
||||
// Standard implementation of AIScoreCalculator
|
||||
//
|
||||
|
||||
#include "StandardAIScoreCalculator.hpp"
|
||||
|
||||
#include <atomic>
|
||||
#include <chrono>
|
||||
#include <unordered_map>
|
||||
|
||||
#include "private/AIScoreCalculatorSharedUtilities.hpp"
|
||||
#include "private/AbstractAIScoreCalculator.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/AIScoreUtilities.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/AIStrategy.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/AIUnitScoreCalculator.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/action_point_distances/ActionPointDistancesCache.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
|
||||
|
||||
namespace shardok {
|
||||
|
||||
using net::eagle0::shardok::storage::fb::BattalionTypeId;
|
||||
|
||||
// Bring shared utilities into scope
|
||||
using score_calculator_internal::AttackerMultiplierForTargetDistance;
|
||||
using score_calculator_internal::CAPTURED_UNIT_SCORE;
|
||||
using score_calculator_internal::CAPTURED_VIP_SCORE;
|
||||
using score_calculator_internal::EffectiveDistanceCache;
|
||||
using score_calculator_internal::FleeStrategyScoreForState;
|
||||
using score_calculator_internal::UNITS_BASE_MULTIPLIER;
|
||||
|
||||
// Forward declare the implementation class
|
||||
class StandardAIScoreCalculator;
|
||||
|
||||
// Anonymous namespace for helper functions that don't need access to scorer
|
||||
namespace {
|
||||
|
||||
#define LOGGING_ 0
|
||||
#define PERFORMANCE_LOGGING_ 0
|
||||
|
||||
// Performance logging for AttackerScoreForState
|
||||
struct AttackerScorePerformanceLogger {
|
||||
static constexpr int LOG_INTERVAL = 100000;
|
||||
|
||||
static std::atomic<int> callCount;
|
||||
static std::atomic<double> intervalTime;
|
||||
static std::atomic<double> totalTime;
|
||||
|
||||
static void LogCall(double duration) {
|
||||
callCount.fetch_add(1);
|
||||
intervalTime.fetch_add(duration);
|
||||
totalTime.fetch_add(duration);
|
||||
|
||||
if (callCount.load() % LOG_INTERVAL == 0) {
|
||||
double intervalAvg = intervalTime.load() / LOG_INTERVAL;
|
||||
double overallAvg = totalTime.load() / callCount.load();
|
||||
printf("AttackerScoreForState: %d calls, last %d avg: %.1f µs, overall avg: %.1f µs\n",
|
||||
callCount.load(),
|
||||
LOG_INTERVAL,
|
||||
intervalAvg * 1000000.0,
|
||||
overallAvg * 1000000.0);
|
||||
intervalTime.store(0.0); // Reset for next interval
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
std::atomic<int> AttackerScorePerformanceLogger::callCount{0};
|
||||
std::atomic<double> AttackerScorePerformanceLogger::intervalTime{0.0};
|
||||
std::atomic<double> AttackerScorePerformanceLogger::totalTime{0.0};
|
||||
|
||||
// RAII timer for automatic performance logging
|
||||
class AttackerScoreTimer {
|
||||
private:
|
||||
std::chrono::high_resolution_clock::time_point startTime;
|
||||
|
||||
public:
|
||||
AttackerScoreTimer() : startTime(std::chrono::high_resolution_clock::now()) {}
|
||||
|
||||
~AttackerScoreTimer() {
|
||||
auto endTime = std::chrono::high_resolution_clock::now();
|
||||
auto duration =
|
||||
std::chrono::duration_cast<std::chrono::duration<double>>(endTime - startTime);
|
||||
AttackerScorePerformanceLogger::LogCall(duration.count());
|
||||
}
|
||||
};
|
||||
} // anonymous namespace
|
||||
|
||||
/// Standard implementation of AIScoreCalculator that uses the default scoring algorithm.
|
||||
/// Inherits from AbstractAIScoreCalculator to share common functionality.
|
||||
class StandardAIScoreCalculator : public AbstractAIScoreCalculator {
|
||||
public:
|
||||
StandardAIScoreCalculator(
|
||||
int maxRounds,
|
||||
ActionPoints braveWaterCost,
|
||||
int meteorRange,
|
||||
double meteorCastVigorCost,
|
||||
int minimumFleeOddsThreshold,
|
||||
int desperateFleeThreshold,
|
||||
std::vector<BattalionTypeSPtr> battalionTypes,
|
||||
const APDCache &apdCache,
|
||||
const ALCache &alCache)
|
||||
: AbstractAIScoreCalculator(
|
||||
maxRounds,
|
||||
braveWaterCost,
|
||||
meteorRange,
|
||||
meteorCastVigorCost,
|
||||
minimumFleeOddsThreshold,
|
||||
desperateFleeThreshold,
|
||||
std::move(battalionTypes),
|
||||
apdCache,
|
||||
alCache) {}
|
||||
|
||||
[[nodiscard]] auto GuessedStateScore(
|
||||
bool isDefender,
|
||||
const GameStateW &state,
|
||||
const AIStrategy &aiStrategy,
|
||||
const CoordsSet &allCastleCoords) const -> ScoreValue override;
|
||||
|
||||
// Implement pure virtual methods from AbstractAIScoreCalculator
|
||||
[[nodiscard]] auto InterpretDefenderOutcome(GameOutcome outcome) const -> ScoreValue override;
|
||||
[[nodiscard]] auto InterpretAttackerOutcome(GameOutcome outcome) const -> ScoreValue override;
|
||||
|
||||
[[nodiscard]] auto AttackerFleeStrategyScoreForState(const GameStateW &gameState) const
|
||||
-> ScoreValue override;
|
||||
|
||||
[[nodiscard]] auto CombineAttackerScores(
|
||||
const UnitsScoreComponents &components,
|
||||
double victoryConditionScore,
|
||||
int roundsRemaining) const -> ScoreValue override;
|
||||
|
||||
[[nodiscard]] auto CombineDefenderScatterScores(const UnitsScoreComponents &components) const
|
||||
-> ScoreValue override;
|
||||
|
||||
[[nodiscard]] auto CombineDefenderHoldCastlesScores(
|
||||
const UnitsScoreComponents &components,
|
||||
double victoryConditionScore,
|
||||
int roundsRemaining) const -> ScoreValue override;
|
||||
|
||||
private:
|
||||
// Implementation methods (converted from internal namespace functions)
|
||||
[[nodiscard]] auto AttackerUnitsScore(
|
||||
const GameStateW &gameState,
|
||||
int roundsRemaining,
|
||||
bool attackerWantsCastles,
|
||||
bool defenderShouldScatter,
|
||||
const vector<TargetPriorityList> &attackerTargetPriorities,
|
||||
const MapId &mapId) const -> ScoreValue;
|
||||
|
||||
[[nodiscard]] auto DefenderFleeStrategyScoreForState(const GameStateW &gameState) const
|
||||
-> ScoreValue override;
|
||||
};
|
||||
|
||||
// Implementation of StandardAIScoreCalculator methods
|
||||
|
||||
auto StandardAIScoreCalculator::InterpretDefenderOutcome(GameOutcome outcome) const -> ScoreValue {
|
||||
switch (outcome) {
|
||||
case GameOutcome::DEFENDER_VICTORY: return INT_MAX;
|
||||
case GameOutcome::ATTACKER_VICTORY: return INT_MIN;
|
||||
case GameOutcome::DRAW: return 0;
|
||||
case GameOutcome::FLEE_OUTCOME: return 0;
|
||||
}
|
||||
throw ShardokInternalErrorException("Unknown GameOutcome");
|
||||
}
|
||||
|
||||
auto StandardAIScoreCalculator::InterpretAttackerOutcome(GameOutcome outcome) const -> ScoreValue {
|
||||
switch (outcome) {
|
||||
case GameOutcome::ATTACKER_VICTORY: return INT_MAX;
|
||||
case GameOutcome::DEFENDER_VICTORY: return INT_MIN;
|
||||
case GameOutcome::DRAW: return 0;
|
||||
case GameOutcome::FLEE_OUTCOME: return 0;
|
||||
}
|
||||
throw ShardokInternalErrorException("Unknown GameOutcome");
|
||||
}
|
||||
|
||||
auto StandardAIScoreCalculator::AttackerUnitsScore(
|
||||
const GameStateW &gameState,
|
||||
int roundsRemaining,
|
||||
bool attackerWantsCastles,
|
||||
bool defenderShouldScatter,
|
||||
const vector<TargetPriorityList> &attackerTargetPriorities,
|
||||
const MapId &mapId) const -> ScoreValue {
|
||||
// Use the base class implementation to get separated attacker/defender values
|
||||
const auto components = CalculateUnitsScoreComponents(
|
||||
gameState,
|
||||
roundsRemaining,
|
||||
attackerWantsCastles,
|
||||
defenderShouldScatter,
|
||||
attackerTargetPriorities,
|
||||
mapId);
|
||||
|
||||
// Standard scorer returns the difference (attacker - defender)
|
||||
return components.attackerUnitsValue - components.defenderUnitsValue;
|
||||
}
|
||||
|
||||
auto StandardAIScoreCalculator::CombineDefenderScatterScores(
|
||||
const UnitsScoreComponents &components) const -> ScoreValue {
|
||||
// For defender scatter, we want to maximize defender units and minimize attacker units
|
||||
// From defender's perspective: negate the attacker-defender difference
|
||||
return components.defenderUnitsValue - components.attackerUnitsValue;
|
||||
}
|
||||
|
||||
auto StandardAIScoreCalculator::CombineDefenderHoldCastlesScores(
|
||||
const UnitsScoreComponents &components,
|
||||
const double victoryConditionScore,
|
||||
const int roundsRemaining) const -> ScoreValue {
|
||||
(void)roundsRemaining; // Intentionally unused for now
|
||||
// From defender's perspective: negate the attacker-defender difference
|
||||
const double unitsDifference = components.defenderUnitsValue - components.attackerUnitsValue;
|
||||
// TODO: The time-decay multiplier (roundsRemaining/maxRounds) was causing END_TURN
|
||||
// to score better than tactical actions because it reduced the penalty for having
|
||||
// fewer units. Setting to constant 1.0 for now to fix tactical decision-making.
|
||||
const double unitsMultiplier = 1.0;
|
||||
// const double unitsMultiplier =
|
||||
// static_cast<double>(roundsRemaining) / static_cast<double>(GetMaxRounds());
|
||||
const double finalScore =
|
||||
UNITS_BASE_MULTIPLIER * unitsMultiplier * unitsDifference + victoryConditionScore;
|
||||
|
||||
return finalScore;
|
||||
}
|
||||
|
||||
auto StandardAIScoreCalculator::DefenderFleeStrategyScoreForState(const GameStateW &gameState) const
|
||||
-> ScoreValue {
|
||||
for (const auto *pi : *gameState->player_infos()) {
|
||||
if (pi->is_defender()) { return FleeStrategyScoreForState(gameState, pi->player_id()); }
|
||||
}
|
||||
throw ShardokInternalErrorException("Unable to find defender for FleeStrategy");
|
||||
}
|
||||
|
||||
auto StandardAIScoreCalculator::AttackerFleeStrategyScoreForState(const GameStateW &gameState) const
|
||||
-> ScoreValue {
|
||||
for (const PlayerInfo *pi : *gameState->player_infos()) {
|
||||
if (!pi->is_defender()) { return FleeStrategyScoreForState(gameState, pi->player_id()); }
|
||||
}
|
||||
throw ShardokInternalErrorException("Unable to find attacker for FleeStrategy");
|
||||
}
|
||||
|
||||
auto StandardAIScoreCalculator::CombineAttackerScores(
|
||||
const UnitsScoreComponents &components,
|
||||
const double victoryConditionScore,
|
||||
const int roundsRemaining) const -> ScoreValue {
|
||||
const double unitsDifference = components.attackerUnitsValue - components.defenderUnitsValue;
|
||||
const double unitsMultiplier =
|
||||
static_cast<double>(roundsRemaining) / static_cast<double>(GetMaxRounds());
|
||||
return UNITS_BASE_MULTIPLIER * unitsMultiplier * unitsDifference + victoryConditionScore;
|
||||
}
|
||||
|
||||
auto StandardAIScoreCalculator::GuessedStateScore(
|
||||
const bool isDefender,
|
||||
const GameStateW &state,
|
||||
const AIStrategy &aiStrategy,
|
||||
const CoordsSet &allCastleCoords) const -> ScoreValue {
|
||||
const int roundsRemaining = GetMaxRounds() - state->current_round();
|
||||
|
||||
if (isDefender) {
|
||||
return DefenderScoreForState(state, aiStrategy, allCastleCoords, roundsRemaining);
|
||||
}
|
||||
return AttackerScoreForState(state, aiStrategy, allCastleCoords, roundsRemaining);
|
||||
}
|
||||
|
||||
// Factory function implementation
|
||||
auto MakeStandardAIScoreCalculator(
|
||||
const SettingsGetter &settingsGetter,
|
||||
const APDCache &apdCache,
|
||||
const ALCache &alCache) -> std::unique_ptr<AIScoreCalculator> {
|
||||
// Extract all battalion types into a vector indexed by BattalionTypeId
|
||||
std::vector<BattalionTypeSPtr> battalionTypes(BattalionTypeId::BattalionTypeId_MAX + 1);
|
||||
for (int typeId = BattalionTypeId::BattalionTypeId_MIN;
|
||||
typeId <= BattalionTypeId::BattalionTypeId_MAX;
|
||||
typeId++) {
|
||||
auto battalionTypeId = static_cast<BattalionTypeId>(typeId);
|
||||
battalionTypes[battalionTypeId] = settingsGetter.GetBattalionType(battalionTypeId);
|
||||
}
|
||||
|
||||
return std::make_unique<StandardAIScoreCalculator>(
|
||||
settingsGetter.Backing().max_rounds(),
|
||||
settingsGetter.Backing().brave_water_action_point_cost(),
|
||||
settingsGetter.Backing().meteor_range(),
|
||||
settingsGetter.Backing().meteor_cast_vigor_cost(),
|
||||
settingsGetter.Backing().ai_minimum_flee_odds_threshold(),
|
||||
settingsGetter.Backing().ai_desperate_flee_threshold(),
|
||||
std::move(battalionTypes),
|
||||
apdCache,
|
||||
alCache);
|
||||
}
|
||||
|
||||
} // namespace shardok
|
||||
+129
@@ -0,0 +1,129 @@
|
||||
//
|
||||
// Shared utilities for AI score calculators - implementation
|
||||
//
|
||||
|
||||
#include "AIScoreCalculatorSharedUtilities.hpp"
|
||||
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/AIScoreUtilities.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/library/map/CoordsSet.hpp"
|
||||
|
||||
namespace shardok {
|
||||
namespace score_calculator_internal {
|
||||
|
||||
auto EffectiveDistanceCache::GetOrCompute(
|
||||
const Unit* unit,
|
||||
const Coords& target,
|
||||
const ActionPointDistances* notBravingApd,
|
||||
const ActionPointDistances* bravingApd,
|
||||
const HexMap* hexMap) const -> DIST_T {
|
||||
CacheKey key{unit->unit_id(), target};
|
||||
auto it = cache.find(key);
|
||||
if (it != cache.end()) { return it->second; }
|
||||
|
||||
CoordsSet targetSet(hexMap);
|
||||
targetSet.Add(target);
|
||||
|
||||
DIST_T result = EffectiveDistance(unit, notBravingApd, bravingApd, targetSet);
|
||||
cache[key] = result;
|
||||
return result;
|
||||
}
|
||||
|
||||
auto FleeStrategyScoreForState(const GameStateW& gameState, const PlayerId playerId) -> ScoreValue {
|
||||
ScoreValue scoreValue = 0.0;
|
||||
|
||||
const auto* gameStatePtr = gameState.Get();
|
||||
const auto* units = gameStatePtr->units();
|
||||
|
||||
for (const auto* unit : *units) {
|
||||
if (unit->status() != net::eagle0::shardok::storage::fb::UnitStatus_NORMAL_UNIT) continue;
|
||||
|
||||
if (unit->player_id() == playerId &&
|
||||
unit->battalion().type() != net::eagle0::shardok::storage::fb::BattalionTypeId_UNDEAD) {
|
||||
scoreValue += FLEE_UNIT_SCORE;
|
||||
|
||||
if (unit->has_attached_hero() &&
|
||||
unit->attached_hero().control_info().controlled_unit_id() != -1) {
|
||||
scoreValue += FLEE_CONTROLLING_UNIT_SCORE;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return scoreValue;
|
||||
}
|
||||
|
||||
// Forward declaration for recursive helper
|
||||
static auto RecursiveAttackerMultiplierForTargetDistance(
|
||||
const Unit* attackingUnit,
|
||||
std::vector<TargetAndAttackLocations>::const_iterator& priorityListNext,
|
||||
const std::vector<TargetAndAttackLocations>::const_iterator& priorityListEnd,
|
||||
const std::vector<const Unit*>& occupants,
|
||||
const HexMap* map,
|
||||
const BattalionTypeSPtr& battType,
|
||||
const ActionPointDistances* notBravingApd,
|
||||
const ActionPointDistances* bravingApd,
|
||||
bool isLateGame) -> double;
|
||||
|
||||
static auto RecursiveAttackerMultiplierForTargetDistance(
|
||||
const Unit* attackingUnit,
|
||||
std::vector<TargetAndAttackLocations>::const_iterator& priorityListNext,
|
||||
const std::vector<TargetAndAttackLocations>::const_iterator& priorityListEnd,
|
||||
const std::vector<const Unit*>& occupants,
|
||||
const HexMap* map,
|
||||
const BattalionTypeSPtr& battType,
|
||||
const ActionPointDistances* notBravingApd,
|
||||
const ActionPointDistances* bravingApd,
|
||||
const bool isLateGame) -> double {
|
||||
if (priorityListNext == priorityListEnd) return 1.0;
|
||||
|
||||
const auto& [target, attackLocations] = *priorityListNext;
|
||||
const Coords& topPriorityTarget = target;
|
||||
|
||||
// If the target is unoccupied or is occupied by this player, give the maximum multiplier, but
|
||||
// also add the bonus for the next up in the priority list
|
||||
if (const Unit* occupant = occupants
|
||||
[topPriorityTarget.row() * map->column_count() + topPriorityTarget.column()];
|
||||
!occupant || occupant->player_id() == attackingUnit->player_id()) {
|
||||
return kMaxProximityBuf + RecursiveAttackerMultiplierForTargetDistance(
|
||||
attackingUnit,
|
||||
++priorityListNext,
|
||||
priorityListEnd,
|
||||
occupants,
|
||||
map,
|
||||
battType,
|
||||
notBravingApd,
|
||||
bravingApd,
|
||||
isLateGame);
|
||||
}
|
||||
|
||||
// Use optimized EffectiveDistance with pre-computed ActionPointDistances
|
||||
// attackLocations is already the CoordsSet of attack locations for this target
|
||||
const DIST_T distance =
|
||||
EffectiveDistance(attackingUnit, notBravingApd, bravingApd, attackLocations);
|
||||
|
||||
return kMaxProximityBuf / (1 + distance / kDistanceDebufRatio);
|
||||
}
|
||||
|
||||
auto AttackerMultiplierForTargetDistance(
|
||||
const Unit* attackingUnit,
|
||||
const std::vector<TargetAndAttackLocations>& priorityList,
|
||||
const std::vector<const Unit*>& occupants,
|
||||
const HexMap* map,
|
||||
const BattalionTypeSPtr& battType,
|
||||
const ActionPointDistances* notBravingApd,
|
||||
const ActionPointDistances* bravingApd,
|
||||
const bool isLateGame) -> double {
|
||||
auto iter = std::begin(priorityList);
|
||||
return RecursiveAttackerMultiplierForTargetDistance(
|
||||
attackingUnit,
|
||||
iter,
|
||||
std::end(priorityList),
|
||||
occupants,
|
||||
map,
|
||||
battType,
|
||||
notBravingApd,
|
||||
bravingApd,
|
||||
isLateGame);
|
||||
}
|
||||
|
||||
} // namespace score_calculator_internal
|
||||
} // namespace shardok
|
||||
@@ -0,0 +1,95 @@
|
||||
//
|
||||
// Shared utilities for AI score calculators
|
||||
// This file is private to the ai/score package
|
||||
//
|
||||
|
||||
#ifndef EAGLE0_AI_SCORE_CALCULATOR_SHARED_UTILITIES_HPP
|
||||
#define EAGLE0_AI_SCORE_CALCULATOR_SHARED_UTILITIES_HPP
|
||||
|
||||
#include <vector>
|
||||
|
||||
#pragma GCC diagnostic push
|
||||
#pragma GCC diagnostic ignored "-Wthread-safety-analysis"
|
||||
#pragma GCC diagnostic ignored "-Wunused-result"
|
||||
#include <gtl/phmap.hpp>
|
||||
#pragma GCC diagnostic pop
|
||||
|
||||
#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/library/BattalionType.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/ActionPointDistances.hpp"
|
||||
#include "src/main/flatbuffer/net/eagle0/shardok/storage/game_state_generated.h"
|
||||
#include "src/main/flatbuffer/net/eagle0/shardok/storage/hex_map_generated.h"
|
||||
#include "src/main/flatbuffer/net/eagle0/shardok/storage/unit_generated.h"
|
||||
|
||||
namespace shardok {
|
||||
namespace score_calculator_internal {
|
||||
|
||||
using HexMap = net::eagle0::shardok::storage::fb::HexMap;
|
||||
using Unit = net::eagle0::shardok::storage::fb::Unit;
|
||||
|
||||
// Scoring constants shared across all calculators
|
||||
constexpr double UNITS_BASE_MULTIPLIER = 0.05;
|
||||
constexpr double FLEE_UNIT_SCORE = -10000;
|
||||
constexpr double FLEE_CONTROLLING_UNIT_SCORE = -10000;
|
||||
constexpr double CAPTURED_UNIT_SCORE = -10000;
|
||||
constexpr double CAPTURED_VIP_SCORE = -25000;
|
||||
constexpr double kMaxProximityBuf = 1.5;
|
||||
constexpr double kDistanceDebufRatio = 8.0;
|
||||
|
||||
/// Structure to hold separated attacker/defender unit scores
|
||||
/// Used by NormalizedAIScoreCalculator to apply asymmetric normalization
|
||||
struct UnitsScoreComponents {
|
||||
double attackerUnitsValue;
|
||||
double defenderUnitsValue;
|
||||
};
|
||||
|
||||
/// Memoization cache for EffectiveDistance calls
|
||||
struct EffectiveDistanceCache {
|
||||
struct CacheKey {
|
||||
UnitId unitId;
|
||||
Coords target;
|
||||
bool operator==(const CacheKey& other) const {
|
||||
return unitId == other.unitId && target == other.target;
|
||||
}
|
||||
};
|
||||
|
||||
struct CacheKeyHash {
|
||||
size_t operator()(const CacheKey& key) const {
|
||||
return std::hash<UnitId>{}(key.unitId) ^ (std::hash<int>{}(key.target.row()) << 1) ^
|
||||
(std::hash<int>{}(key.target.column()) << 2);
|
||||
}
|
||||
};
|
||||
|
||||
mutable gtl::flat_hash_map<CacheKey, DIST_T, CacheKeyHash> cache;
|
||||
|
||||
DIST_T GetOrCompute(
|
||||
const Unit* unit,
|
||||
const Coords& target,
|
||||
const ActionPointDistances* notBravingApd,
|
||||
const ActionPointDistances* bravingApd,
|
||||
const HexMap* hexMap) const;
|
||||
};
|
||||
|
||||
/// Calculate score for FLEE strategy
|
||||
/// Returns negative score based on fleeing units
|
||||
auto FleeStrategyScoreForState(const GameStateW& gameState, PlayerId playerId) -> ScoreValue;
|
||||
|
||||
/// Calculate attacker multiplier based on distance to priority targets
|
||||
/// This is used to weight attacker units by their proximity to objectives
|
||||
auto AttackerMultiplierForTargetDistance(
|
||||
const Unit* attackingUnit,
|
||||
const std::vector<TargetAndAttackLocations>& priorityList,
|
||||
const std::vector<const Unit*>& occupants,
|
||||
const HexMap* map,
|
||||
const BattalionTypeSPtr& battType,
|
||||
const ActionPointDistances* notBravingApd,
|
||||
const ActionPointDistances* bravingApd,
|
||||
bool isLateGame) -> double;
|
||||
|
||||
} // namespace score_calculator_internal
|
||||
} // namespace shardok
|
||||
|
||||
#endif // EAGLE0_AI_SCORE_CALCULATOR_SHARED_UTILITIES_HPP
|
||||
@@ -0,0 +1,514 @@
|
||||
//
|
||||
// Abstract base class for AI score calculator implementations
|
||||
//
|
||||
|
||||
#include "AbstractAIScoreCalculator.hpp"
|
||||
|
||||
#include "AIScoreCalculatorSharedUtilities.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/AIScoreUtilities.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/AIStrategy.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/AIUnitScoreCalculator.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/AIWaterCrossingCalculator.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/score/AIVictoryConditionScoreCalculator.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/library/util/HexMapUtils.hpp"
|
||||
|
||||
namespace shardok {
|
||||
|
||||
using net::eagle0::shardok::storage::fb::BattalionTypeId;
|
||||
using net::eagle0::shardok::storage::fb::Unit;
|
||||
|
||||
using score_calculator_internal::AttackerMultiplierForTargetDistance;
|
||||
using score_calculator_internal::CAPTURED_UNIT_SCORE;
|
||||
using score_calculator_internal::CAPTURED_VIP_SCORE;
|
||||
using score_calculator_internal::EffectiveDistanceCache;
|
||||
using score_calculator_internal::UnitsScoreComponents;
|
||||
|
||||
auto AbstractAIScoreCalculator::CalculateUnitsScoreComponents(
|
||||
const GameStateW &gameState,
|
||||
int roundsRemaining,
|
||||
bool attackerWantsCastles,
|
||||
bool defenderShouldScatter,
|
||||
const vector<TargetPriorityList> &attackerTargetPriorities,
|
||||
const MapId &mapId) const -> UnitsScoreComponents {
|
||||
// Cache frequently accessed FlatBuffer fields to avoid repeated offset calculations
|
||||
const auto *gameStateRawPtr = gameState.Get();
|
||||
const auto *cachedUnits = gameStateRawPtr->units();
|
||||
const auto *cachedHexMap = gameStateRawPtr->hex_map();
|
||||
|
||||
const int16_t cachedRowCount = cachedHexMap->row_count();
|
||||
const int16_t cachedColumnCount = cachedHexMap->column_count();
|
||||
const int cachedCurrentRound = gameStateRawPtr->current_round();
|
||||
|
||||
bool isLateGame = cachedCurrentRound > 18; // Inline IsLateGame for efficiency
|
||||
|
||||
// APDCache now has built-in thread-local caching - no need for PreCachedAPDs
|
||||
ActionPoints braveWaterCost = GetBraveWaterCost();
|
||||
|
||||
// Memoization cache for EffectiveDistance calls
|
||||
EffectiveDistanceCache distanceCache;
|
||||
|
||||
std::vector<const Unit *> attackerUnits{};
|
||||
std::vector<const Unit *> defenderUnits{};
|
||||
// Pre-allocate vectors based on estimated unit ratios to avoid reallocations
|
||||
const size_t estimatedUnitCount = cachedUnits->size();
|
||||
attackerUnits.reserve(estimatedUnitCount - 1);
|
||||
defenderUnits.reserve(estimatedUnitCount - 1);
|
||||
|
||||
double attackerUnitsValue = 0;
|
||||
double defenderUnitsValue = 0;
|
||||
|
||||
// Early return for empty game states
|
||||
if (cachedUnits->size() == 0) { return UnitsScoreComponents{0.0, 0.0}; }
|
||||
|
||||
auto occupants = Occupants(*cachedUnits, cachedRowCount, cachedColumnCount);
|
||||
|
||||
for (const Unit *unit : *cachedUnits) {
|
||||
const auto *pi = PlayerInfoForPid(gameState, unit->player_id());
|
||||
if (pi == nullptr) { continue; }
|
||||
|
||||
switch (unit->status()) {
|
||||
case net::eagle0::shardok::storage::fb::UnitStatus_NORMAL_UNIT: {
|
||||
if (pi->is_defender()) {
|
||||
defenderUnits.push_back(unit);
|
||||
} else {
|
||||
attackerUnits.push_back(unit);
|
||||
}
|
||||
break;
|
||||
}
|
||||
case net::eagle0::shardok::storage::fb::UnitStatus_CAPTURED_UNIT: {
|
||||
double thisScore = unit->has_attached_hero() && unit->attached_hero().is_vip()
|
||||
? CAPTURED_VIP_SCORE
|
||||
: CAPTURED_UNIT_SCORE;
|
||||
if (pi->is_defender()) {
|
||||
defenderUnitsValue += thisScore;
|
||||
} else {
|
||||
attackerUnitsValue += thisScore;
|
||||
}
|
||||
break;
|
||||
}
|
||||
|
||||
case net::eagle0::shardok::storage::fb::UnitStatus_DESTROYED_SUMMONED_UNIT:
|
||||
case net::eagle0::shardok::storage::fb::UnitStatus_FLED_UNIT:
|
||||
case net::eagle0::shardok::storage::fb::UnitStatus_NEVER_ENTERED_UNIT:
|
||||
case net::eagle0::shardok::storage::fb::UnitStatus_OUTLAWED_UNIT:
|
||||
case net::eagle0::shardok::storage::fb::UnitStatus_RESERVE_UNIT:
|
||||
case net::eagle0::shardok::storage::fb::UnitStatus_RETREATED_UNIT:
|
||||
case net::eagle0::shardok::storage::fb::UnitStatus_RESERVED_SLOT: break;
|
||||
|
||||
case net::eagle0::shardok::storage::fb::UnitStatus_UNKNOWN_UNIT:
|
||||
throw ShardokInternalErrorException("Unknown unit status");
|
||||
}
|
||||
}
|
||||
|
||||
double defenderAdvantage = 1.0 + static_cast<double>(cachedCurrentRound) / 31.0;
|
||||
|
||||
// Can we cache this somehow, it won't usually change within your turn
|
||||
auto attackLocationsForAttacker = GetAlCache()->CachedLocations(defenderUnits, isLateGame);
|
||||
const auto &locationsCausingDanger = attackLocationsForAttacker.AllLocations();
|
||||
|
||||
// Process attacker units using cached ActionPointDistances
|
||||
for (const Unit *unit : attackerUnits) {
|
||||
const int battTypeId = unit->battalion().type();
|
||||
// Cache battalion type reference to avoid shared_ptr atomic operations
|
||||
const auto &battalionType = GetBattalionType(static_cast<BattalionTypeId>(battTypeId));
|
||||
|
||||
// Cache APD lookups - same battalion type is used multiple times below
|
||||
const auto *notBravingApd =
|
||||
GetApdCache()->GetRaw(cachedHexMap, mapId, battalionType, false);
|
||||
const auto *bravingApd = battalionType->allowsBraveWater ? GetApdCache()->GetRaw(
|
||||
cachedHexMap,
|
||||
mapId,
|
||||
battalionType,
|
||||
true,
|
||||
braveWaterCost)
|
||||
: nullptr;
|
||||
|
||||
const auto &priorityList = std::ranges::find_if(
|
||||
attackerTargetPriorities,
|
||||
[&unit](const TargetPriorityList &tpl) {
|
||||
return tpl.attackingUnitId == unit->unit_id();
|
||||
});
|
||||
|
||||
// If there are any tiles being targeted, give this unit a multiplier based on how close
|
||||
// they are to being able to attack it
|
||||
double distanceMultiplier = priorityList == end(attackerTargetPriorities)
|
||||
? 1.0
|
||||
: AttackerMultiplierForTargetDistance(
|
||||
unit,
|
||||
priorityList->priorityOrder,
|
||||
occupants,
|
||||
cachedHexMap,
|
||||
battalionType,
|
||||
notBravingApd,
|
||||
bravingApd,
|
||||
isLateGame);
|
||||
|
||||
auto uv = UnitValue(
|
||||
unit,
|
||||
true,
|
||||
attackerUnits,
|
||||
attackerWantsCastles,
|
||||
/* includeCastleBonus=*/true,
|
||||
defenderUnits,
|
||||
cachedHexMap,
|
||||
roundsRemaining,
|
||||
attackLocationsForAttacker,
|
||||
locationsCausingDanger,
|
||||
notBravingApd,
|
||||
GetMeteorRange(),
|
||||
GetMeteorCastVigorCost());
|
||||
|
||||
attackerUnitsValue += distanceMultiplier * uv;
|
||||
}
|
||||
|
||||
auto attackLocationsForDefender = GetAlCache()->CachedLocations(attackerUnits, isLateGame);
|
||||
const auto &locationsCausingDangerForAttacker = attackLocationsForDefender.AllLocations();
|
||||
|
||||
for (const Unit *unit : defenderUnits) {
|
||||
auto defenderUnitId = unit->unit_id();
|
||||
const int battTypeId = unit->battalion().type();
|
||||
// Cache battalion type reference to avoid shared_ptr atomic operations
|
||||
const auto &battalionType = GetBattalionType(static_cast<BattalionTypeId>(battTypeId));
|
||||
|
||||
// Cache APD lookups for this defender unit
|
||||
const auto *defenderNotBravingApd =
|
||||
GetApdCache()->GetRaw(cachedHexMap, mapId, battalionType, false);
|
||||
|
||||
auto dv = UnitValue(
|
||||
unit,
|
||||
false,
|
||||
attackerUnits,
|
||||
attackerWantsCastles,
|
||||
/* includeCastleBonus=*/!defenderShouldScatter,
|
||||
defenderUnits,
|
||||
cachedHexMap,
|
||||
roundsRemaining,
|
||||
attackLocationsForDefender,
|
||||
locationsCausingDangerForAttacker,
|
||||
defenderNotBravingApd,
|
||||
GetMeteorRange(),
|
||||
GetMeteorCastVigorCost());
|
||||
|
||||
double distanceMultiplier = 1.0;
|
||||
|
||||
// If the defender is trying to scatter, than we want to be as far away from the nearest
|
||||
// attacker as possible, AND as far away from the nearest friendly as possible
|
||||
if (unit->location().row() > -1 && defenderShouldScatter) {
|
||||
CoordsSet myLocationSet(cachedHexMap);
|
||||
myLocationSet.Add(unit->location());
|
||||
|
||||
DIST_T closestDistanceToEnemy = 999;
|
||||
for (const auto &attackerUnit : attackerUnits) {
|
||||
const int attackerBattTypeId = attackerUnit->battalion().type();
|
||||
// Cache attacker battalion type reference in nested loop
|
||||
const auto &attackerBattalionType =
|
||||
GetBattalionType(static_cast<BattalionTypeId>(attackerBattTypeId));
|
||||
const DIST_T thisDistance = distanceCache.GetOrCompute(
|
||||
attackerUnit,
|
||||
unit->location(),
|
||||
GetApdCache()->GetRaw(cachedHexMap, mapId, attackerBattalionType, false),
|
||||
attackerBattalionType->allowsBraveWater ? GetApdCache()->GetRaw(
|
||||
cachedHexMap,
|
||||
mapId,
|
||||
attackerBattalionType,
|
||||
true,
|
||||
braveWaterCost)
|
||||
: nullptr,
|
||||
cachedHexMap);
|
||||
if (thisDistance < closestDistanceToEnemy) {
|
||||
closestDistanceToEnemy = thisDistance;
|
||||
}
|
||||
}
|
||||
|
||||
// If the best we can do puts us very close to the enemy, and the unit is almost
|
||||
// destroyed, return a negative value; better to flee
|
||||
if (unit->can_flee() && closestDistanceToEnemy < 5 && unit->battalion().size() < 10) {
|
||||
distanceMultiplier = -1;
|
||||
} else {
|
||||
DIST_T closestDistanceToFriendly = 1;
|
||||
if (defenderUnits.size() > 1) {
|
||||
for (const auto &defenderUnit : defenderUnits) {
|
||||
if (defenderUnit->unit_id() != defenderUnitId) {
|
||||
const int defenderBattTypeId = defenderUnit->battalion().type();
|
||||
// Cache defender battalion type reference in nested loop
|
||||
const auto &defenderBattalionType = GetBattalionType(
|
||||
static_cast<BattalionTypeId>(defenderBattTypeId));
|
||||
const DIST_T thisDistance = distanceCache.GetOrCompute(
|
||||
defenderUnit,
|
||||
unit->location(),
|
||||
GetApdCache()->GetRaw(
|
||||
cachedHexMap,
|
||||
mapId,
|
||||
defenderBattalionType,
|
||||
false),
|
||||
defenderBattalionType->allowsBraveWater
|
||||
? GetApdCache()->GetRaw(
|
||||
cachedHexMap,
|
||||
mapId,
|
||||
defenderBattalionType,
|
||||
true,
|
||||
braveWaterCost)
|
||||
: nullptr,
|
||||
cachedHexMap);
|
||||
if (thisDistance < closestDistanceToEnemy) {
|
||||
closestDistanceToFriendly = thisDistance;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
distanceMultiplier =
|
||||
(closestDistanceToEnemy + closestDistanceToFriendly / 5.0) / 5.0;
|
||||
}
|
||||
}
|
||||
|
||||
defenderUnitsValue += distanceMultiplier * dv;
|
||||
}
|
||||
|
||||
defenderUnitsValue *= defenderAdvantage;
|
||||
|
||||
return UnitsScoreComponents{attackerUnitsValue, defenderUnitsValue};
|
||||
}
|
||||
|
||||
auto AbstractAIScoreCalculator::FindDefenderPlayerInfo(const GameStateW &gameState) const
|
||||
-> const net::eagle0::shardok::storage::fb::PlayerInfo * {
|
||||
const auto *playerInfos = gameState->player_infos();
|
||||
if (playerInfos == nullptr) { return nullptr; }
|
||||
|
||||
for (const net::eagle0::shardok::storage::fb::PlayerInfo *pi : *playerInfos) {
|
||||
if (pi->is_defender()) return pi;
|
||||
}
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
auto AbstractAIScoreCalculator::CalculateAttackerVictoryConditionScore(
|
||||
const GameStateW &gameState,
|
||||
const AIStrategy &attackerStrategy,
|
||||
const CoordsSet &castleCoords) const -> ScoreValue {
|
||||
ScoreValue victoryConditionTotal = 0.0;
|
||||
|
||||
const auto *playerInfos = gameState->player_infos();
|
||||
if (playerInfos == nullptr) { return 0.0; }
|
||||
|
||||
for (const net::eagle0::shardok::storage::fb::PlayerInfo *pi : *playerInfos) {
|
||||
if (pi->is_defender()) continue;
|
||||
|
||||
switch (attackerStrategy.strategyType) {
|
||||
case AIStrategy::STRATEGY_CROSS_RIVERS:
|
||||
victoryConditionTotal += WaterCrossingScore(
|
||||
pi->player_id(),
|
||||
[this](BattalionTypeId typeId) { return GetBattalionType(typeId); },
|
||||
gameState,
|
||||
castleCoords,
|
||||
attackerStrategy.targetLocations,
|
||||
GetApdCache());
|
||||
break;
|
||||
|
||||
case AIStrategy::STRATEGY_ATTACK_CASTLES:
|
||||
case AIStrategy::STRATEGY_ATTACK_UNITS:
|
||||
// already factored into AttackerUnitsScore
|
||||
break;
|
||||
|
||||
case AIStrategy::STRATEGY_HOLD_CASTLES:
|
||||
victoryConditionTotal += AttackerHoldsCriticalTilesVictoryScore(
|
||||
gameState,
|
||||
castleCoords,
|
||||
pi,
|
||||
GetApdCache(),
|
||||
GetAlCache(),
|
||||
[this](BattalionTypeId typeId) { return GetBattalionType(typeId); },
|
||||
GetBraveWaterCost());
|
||||
break;
|
||||
|
||||
case AIStrategy::STRATEGY_SCATTER:
|
||||
throw ShardokInternalErrorException("Attacker cannot use ScatterStrategy");
|
||||
|
||||
case AIStrategy::STRATEGY_FLEE:
|
||||
// FLEE strategy is handled specially by each subclass
|
||||
// Return 0 here and let the caller handle it
|
||||
return 0.0;
|
||||
}
|
||||
}
|
||||
|
||||
return victoryConditionTotal;
|
||||
}
|
||||
|
||||
auto AbstractAIScoreCalculator::DefenderScoreForState(
|
||||
const GameStateW &gameState,
|
||||
const AIStrategy &defenderStrategy,
|
||||
const CoordsSet &castleCoords,
|
||||
const int roundsRemaining) const -> ScoreValue {
|
||||
// Check for terminal states
|
||||
if (gameState->status()->state() ==
|
||||
net::eagle0::shardok::storage::fb::GameStatus_::State_VICTORY) {
|
||||
const auto *winningIds = gameState->status()->winning_shardok_ids();
|
||||
const auto *playerInfos = gameState->player_infos();
|
||||
if (winningIds != nullptr && playerInfos != nullptr) {
|
||||
for (const PlayerId winningPid : *winningIds) {
|
||||
if (winningPid < 0) continue;
|
||||
if (defenderStrategy.strategyType == AIStrategy::STRATEGY_FLEE) {
|
||||
return InterpretDefenderOutcome(GameOutcome::FLEE_OUTCOME);
|
||||
}
|
||||
if (playerInfos->Get(winningPid)->is_defender()) {
|
||||
return InterpretDefenderOutcome(GameOutcome::DEFENDER_VICTORY);
|
||||
}
|
||||
return InterpretDefenderOutcome(GameOutcome::ATTACKER_VICTORY);
|
||||
}
|
||||
}
|
||||
return InterpretDefenderOutcome(GameOutcome::ATTACKER_VICTORY);
|
||||
}
|
||||
if (gameState->status()->state() ==
|
||||
net::eagle0::shardok::storage::fb::GameStatus_::State_DRAW) {
|
||||
return InterpretDefenderOutcome(GameOutcome::DRAW);
|
||||
}
|
||||
|
||||
// Delegate to strategy-specific methods (implemented by subclasses)
|
||||
switch (defenderStrategy.strategyType) {
|
||||
case AIStrategy::STRATEGY_ATTACK_CASTLES:
|
||||
throw ShardokInternalErrorException("Defender cannot use AttackCastlesStrategy");
|
||||
case AIStrategy::STRATEGY_ATTACK_UNITS:
|
||||
throw ShardokInternalErrorException("Defender cannot use AttackUnitsStrategy");
|
||||
case AIStrategy::STRATEGY_CROSS_RIVERS:
|
||||
throw ShardokInternalErrorException("Defender cannot use CrossRiversStrategy");
|
||||
case AIStrategy::STRATEGY_HOLD_CASTLES:
|
||||
return DefenderHoldCastlesStrategyScoreForState(
|
||||
gameState,
|
||||
castleCoords,
|
||||
roundsRemaining);
|
||||
case AIStrategy::STRATEGY_SCATTER:
|
||||
return DefenderScatterStrategyScoreForState(gameState, roundsRemaining);
|
||||
case AIStrategy::STRATEGY_FLEE: return DefenderFleeStrategyScoreForState(gameState);
|
||||
}
|
||||
throw ShardokInternalErrorException("Escaped AIStrategy switch");
|
||||
}
|
||||
|
||||
auto AbstractAIScoreCalculator::DefenderScatterStrategyScoreForState(
|
||||
const GameStateW &gameState,
|
||||
const int roundsRemaining) const -> ScoreValue {
|
||||
// Check for terminal states
|
||||
if (gameState->status()->state() ==
|
||||
net::eagle0::shardok::storage::fb::GameStatus_::State_VICTORY) {
|
||||
const auto *winningIds = gameState->status()->winning_shardok_ids();
|
||||
const auto *playerInfos = gameState->player_infos();
|
||||
if (winningIds != nullptr && playerInfos != nullptr) {
|
||||
for (const PlayerId winningPid : *winningIds) {
|
||||
if (winningPid < 0) continue;
|
||||
if (playerInfos->Get(winningPid)->is_defender()) {
|
||||
return InterpretDefenderOutcome(GameOutcome::DEFENDER_VICTORY);
|
||||
}
|
||||
return InterpretDefenderOutcome(GameOutcome::ATTACKER_VICTORY);
|
||||
}
|
||||
}
|
||||
return InterpretDefenderOutcome(GameOutcome::DEFENDER_VICTORY);
|
||||
}
|
||||
if (gameState->status()->state() ==
|
||||
net::eagle0::shardok::storage::fb::GameStatus_::State_DRAW) {
|
||||
return InterpretDefenderOutcome(GameOutcome::DRAW);
|
||||
}
|
||||
|
||||
// Get units score components
|
||||
const auto mapId = ActionPointDistancesCache::GetMapId(gameState->hex_map());
|
||||
const auto components = CalculateUnitsScoreComponents(
|
||||
gameState,
|
||||
roundsRemaining,
|
||||
/* attackerWantsCastles=*/false,
|
||||
/* defenderShouldScatter=*/true,
|
||||
{},
|
||||
mapId);
|
||||
|
||||
// Combine using subclass-specific logic
|
||||
return CombineDefenderScatterScores(components);
|
||||
}
|
||||
|
||||
auto AbstractAIScoreCalculator::DefenderHoldCastlesStrategyScoreForState(
|
||||
const GameStateW &gameState,
|
||||
const CoordsSet &castleCoords,
|
||||
const int roundsRemaining) const -> ScoreValue {
|
||||
// Get units score components
|
||||
const auto components = CalculateUnitsScoreComponents(
|
||||
gameState,
|
||||
roundsRemaining,
|
||||
/* attackerWantsCastles=*/true,
|
||||
/* defenderShouldScatter=*/false,
|
||||
{},
|
||||
ActionPointDistancesCache::GetMapId(gameState->hex_map()));
|
||||
|
||||
// Get victory condition score from defender's perspective
|
||||
const PlayerInfo *defenderPi = FindDefenderPlayerInfo(gameState);
|
||||
|
||||
// Handle null defenderPi gracefully
|
||||
if (defenderPi == nullptr) {
|
||||
// No defender player found - return neutral score using components only
|
||||
return CombineDefenderHoldCastlesScores(components, 0.0, roundsRemaining);
|
||||
}
|
||||
|
||||
const double victoryConditionScore =
|
||||
DefenderHoldsCriticalTilesVictoryScore(gameState, castleCoords, defenderPi);
|
||||
|
||||
// Combine using subclass-specific logic
|
||||
return CombineDefenderHoldCastlesScores(components, victoryConditionScore, roundsRemaining);
|
||||
}
|
||||
|
||||
auto AbstractAIScoreCalculator::AttackerScoreForState(
|
||||
const GameStateW &gameState,
|
||||
const AIStrategy &attackerStrategy,
|
||||
const CoordsSet &castleCoords,
|
||||
const int roundsRemaining) const -> ScoreValue {
|
||||
// Check for terminal states
|
||||
if (gameState->status()->state() ==
|
||||
net::eagle0::shardok::storage::fb::GameStatus_::State_VICTORY) {
|
||||
const auto *winningIds = gameState->status()->winning_shardok_ids();
|
||||
const auto *playerInfos = gameState->player_infos();
|
||||
if (winningIds != nullptr && playerInfos != nullptr) {
|
||||
for (const PlayerId winningPid : *winningIds) {
|
||||
if (winningPid < 0) continue;
|
||||
if (attackerStrategy.strategyType == AIStrategy::STRATEGY_FLEE) {
|
||||
return InterpretAttackerOutcome(GameOutcome::FLEE_OUTCOME);
|
||||
}
|
||||
if (playerInfos->Get(winningPid)->is_defender()) {
|
||||
return InterpretAttackerOutcome(GameOutcome::DEFENDER_VICTORY);
|
||||
}
|
||||
return InterpretAttackerOutcome(GameOutcome::ATTACKER_VICTORY);
|
||||
}
|
||||
}
|
||||
return InterpretAttackerOutcome(GameOutcome::ATTACKER_VICTORY);
|
||||
}
|
||||
if (gameState->status()->state() ==
|
||||
net::eagle0::shardok::storage::fb::GameStatus_::State_DRAW) {
|
||||
return InterpretAttackerOutcome(GameOutcome::DRAW);
|
||||
}
|
||||
|
||||
// Handle FLEE strategy specially
|
||||
if (attackerStrategy.strategyType == AIStrategy::STRATEGY_FLEE) {
|
||||
return AttackerFleeStrategyScoreForState(gameState);
|
||||
}
|
||||
|
||||
// Edge case: no units means call subclass method to handle neutral state
|
||||
// This is defensive - some subclasses may want special handling
|
||||
const auto *units = gameState->units();
|
||||
if (units == nullptr || units->size() == 0) {
|
||||
// Call CombineAttackerScores with all zeros
|
||||
return CombineAttackerScores(UnitsScoreComponents{0.0, 0.0}, 0.0, roundsRemaining);
|
||||
}
|
||||
|
||||
// Get units score components
|
||||
const auto mapId = ActionPointDistancesCache::GetMapId(gameState->hex_map());
|
||||
const auto components = CalculateUnitsScoreComponents(
|
||||
gameState,
|
||||
roundsRemaining,
|
||||
attackerStrategy.strategyType == AIStrategy::STRATEGY_HOLD_CASTLES,
|
||||
/* defenderShouldScatter=*/false,
|
||||
attackerStrategy.targetPriorities,
|
||||
mapId);
|
||||
|
||||
// Calculate victory condition score
|
||||
const ScoreValue victoryConditionTotal =
|
||||
CalculateAttackerVictoryConditionScore(gameState, attackerStrategy, castleCoords);
|
||||
|
||||
// Combine scores using subclass-specific logic
|
||||
// Standard: uses difference and multiplier
|
||||
// Normalized: uses normalization formula
|
||||
return CombineAttackerScores(components, victoryConditionTotal, roundsRemaining);
|
||||
}
|
||||
|
||||
} // namespace shardok
|
||||
@@ -0,0 +1,176 @@
|
||||
//
|
||||
// Abstract base class for AI score calculator implementations
|
||||
// This file is private to the ai/score package
|
||||
//
|
||||
|
||||
#ifndef EAGLE0_ABSTRACT_AI_SCORE_CALCULATOR_HPP
|
||||
#define EAGLE0_ABSTRACT_AI_SCORE_CALCULATOR_HPP
|
||||
|
||||
#include <memory>
|
||||
#include <unordered_map>
|
||||
|
||||
#include "AIScoreCalculatorSharedUtilities.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/AIStrategy.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/score/AIScoreCalculator.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/library/BattalionType.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/flatbuffer/net/eagle0/shardok/storage/game_state_generated.h"
|
||||
|
||||
namespace shardok {
|
||||
|
||||
using net::eagle0::shardok::storage::fb::BattalionTypeId;
|
||||
|
||||
using APDCache = std::shared_ptr<ActionPointDistancesCache>;
|
||||
using ALCache = std::unique_ptr<AttackLocationsCache>;
|
||||
using score_calculator_internal::UnitsScoreComponents;
|
||||
|
||||
/// Enum representing terminal game outcomes for score interpretation
|
||||
enum class GameOutcome {
|
||||
ATTACKER_VICTORY,
|
||||
DEFENDER_VICTORY,
|
||||
DRAW,
|
||||
FLEE_OUTCOME // Special outcome for FLEE strategy
|
||||
};
|
||||
|
||||
/// Abstract base class providing shared functionality for AI score calculators.
|
||||
/// Contains common member variables, accessor methods, and shared helper logic.
|
||||
/// This class is private to the ai/score package.
|
||||
class AbstractAIScoreCalculator : public AIScoreCalculator {
|
||||
protected:
|
||||
AbstractAIScoreCalculator(
|
||||
int maxRounds,
|
||||
ActionPoints braveWaterCost,
|
||||
int meteorRange,
|
||||
double meteorCastVigorCost,
|
||||
int minimumFleeOddsThreshold,
|
||||
int desperateFleeThreshold,
|
||||
std::vector<BattalionTypeSPtr> battalionTypes,
|
||||
const APDCache &apdCache,
|
||||
const ALCache &alCache)
|
||||
: maxRounds_(maxRounds),
|
||||
braveWaterCost_(braveWaterCost),
|
||||
meteorRange_(meteorRange),
|
||||
meteorCastVigorCost_(meteorCastVigorCost),
|
||||
minimumFleeOddsThreshold_(minimumFleeOddsThreshold),
|
||||
desperateFleeThreshold_(desperateFleeThreshold),
|
||||
battalionTypes_(std::move(battalionTypes)),
|
||||
apdCache_(apdCache),
|
||||
alCache_(alCache) {}
|
||||
|
||||
// Protected accessor methods available to subclasses
|
||||
[[nodiscard]] inline auto GetBattalionType(BattalionTypeId typeId) const
|
||||
-> const BattalionTypeSPtr & {
|
||||
return battalionTypes_[typeId];
|
||||
}
|
||||
|
||||
[[nodiscard]] auto GetApdCache() const -> const APDCache & { return apdCache_; }
|
||||
[[nodiscard]] auto GetAlCache() const -> const ALCache & { return alCache_; }
|
||||
[[nodiscard]] auto GetBraveWaterCost() const -> ActionPoints { return braveWaterCost_; }
|
||||
[[nodiscard]] auto GetMaxRounds() const -> int { return maxRounds_; }
|
||||
[[nodiscard]] auto GetMeteorRange() const -> int { return meteorRange_; }
|
||||
[[nodiscard]] auto GetMeteorCastVigorCost() const -> double { return meteorCastVigorCost_; }
|
||||
[[nodiscard]] auto GetMinimumFleeOddsThreshold() const -> int {
|
||||
return minimumFleeOddsThreshold_;
|
||||
}
|
||||
[[nodiscard]] auto GetDesperateFleeThreshold() const -> int { return desperateFleeThreshold_; }
|
||||
|
||||
// Protected helper method that both subclasses can use
|
||||
// Returns separate attacker and defender unit values
|
||||
[[nodiscard]] auto CalculateUnitsScoreComponents(
|
||||
const GameStateW &gameState,
|
||||
int roundsRemaining,
|
||||
bool attackerWantsCastles,
|
||||
bool defenderShouldScatter,
|
||||
const vector<TargetPriorityList> &attackerTargetPriorities,
|
||||
const MapId &mapId) const -> UnitsScoreComponents;
|
||||
|
||||
// Helper to find the defender PlayerInfo
|
||||
[[nodiscard]] auto FindDefenderPlayerInfo(const GameStateW &gameState) const
|
||||
-> const net::eagle0::shardok::storage::fb::PlayerInfo *;
|
||||
|
||||
// Calculate victory condition score for attacker strategies
|
||||
// Returns the raw victory condition score (not yet combined with units score)
|
||||
[[nodiscard]] auto CalculateAttackerVictoryConditionScore(
|
||||
const GameStateW &gameState,
|
||||
const AIStrategy &attackerStrategy,
|
||||
const CoordsSet &castleCoords) const -> ScoreValue;
|
||||
|
||||
// Pure virtual methods for subclasses to interpret game outcomes
|
||||
[[nodiscard]] virtual auto InterpretDefenderOutcome(GameOutcome outcome) const
|
||||
-> ScoreValue = 0;
|
||||
[[nodiscard]] virtual auto InterpretAttackerOutcome(GameOutcome outcome) const
|
||||
-> ScoreValue = 0;
|
||||
|
||||
// Shared implementation of DefenderScoreForState that uses InterpretDefenderOutcome
|
||||
[[nodiscard]] auto DefenderScoreForState(
|
||||
const GameStateW &gameState,
|
||||
const AIStrategy &defenderStrategy,
|
||||
const CoordsSet &castleCoords,
|
||||
int roundsRemaining) const -> ScoreValue;
|
||||
|
||||
// Shared implementations that delegate to pure virtual methods
|
||||
[[nodiscard]] auto DefenderScatterStrategyScoreForState(
|
||||
const GameStateW &gameState,
|
||||
int roundsRemaining) const -> ScoreValue;
|
||||
|
||||
[[nodiscard]] auto DefenderHoldCastlesStrategyScoreForState(
|
||||
const GameStateW &gameState,
|
||||
const CoordsSet &castleCoords,
|
||||
int roundsRemaining) const -> ScoreValue;
|
||||
|
||||
// Pure virtual methods for combining defender scores
|
||||
[[nodiscard]] virtual auto CombineDefenderScatterScores(
|
||||
const UnitsScoreComponents &components) const -> ScoreValue = 0;
|
||||
|
||||
[[nodiscard]] virtual auto CombineDefenderHoldCastlesScores(
|
||||
const UnitsScoreComponents &components,
|
||||
double victoryConditionScore,
|
||||
int roundsRemaining) const -> ScoreValue = 0;
|
||||
|
||||
[[nodiscard]] virtual auto DefenderFleeStrategyScoreForState(const GameStateW &gameState) const
|
||||
-> ScoreValue = 0;
|
||||
|
||||
// Shared implementation of AttackerScoreForState that uses InterpretAttackerOutcome
|
||||
[[nodiscard]] auto AttackerScoreForState(
|
||||
const GameStateW &gameState,
|
||||
const AIStrategy &attackerStrategy,
|
||||
const CoordsSet &castleCoords,
|
||||
int roundsRemaining) const -> ScoreValue;
|
||||
|
||||
// Pure virtual method for attacker flee strategy scoring
|
||||
[[nodiscard]] virtual auto AttackerFleeStrategyScoreForState(const GameStateW &gameState) const
|
||||
-> ScoreValue = 0;
|
||||
|
||||
// Pure virtual method for combining units score and victory condition score
|
||||
// This is where Standard and Normalized diverge in their scoring approach
|
||||
// Standard: uses difference and applies multiplier
|
||||
// Normalized: uses normalization formula with separate attacker/defender values
|
||||
[[nodiscard]] virtual auto CombineAttackerScores(
|
||||
const UnitsScoreComponents &components,
|
||||
double victoryConditionScore,
|
||||
int roundsRemaining) const -> ScoreValue = 0;
|
||||
|
||||
private:
|
||||
// Scalar settings extracted from SettingsGetter
|
||||
int maxRounds_;
|
||||
ActionPoints braveWaterCost_;
|
||||
int meteorRange_;
|
||||
double meteorCastVigorCost_;
|
||||
int minimumFleeOddsThreshold_;
|
||||
int desperateFleeThreshold_;
|
||||
|
||||
// Battalion type lookup vector (indexed by BattalionTypeId)
|
||||
std::vector<BattalionTypeSPtr> battalionTypes_;
|
||||
|
||||
// Caches (stored as references)
|
||||
const APDCache &apdCache_;
|
||||
const ALCache &alCache_;
|
||||
};
|
||||
|
||||
} // namespace shardok
|
||||
|
||||
#endif // EAGLE0_ABSTRACT_AI_SCORE_CALCULATOR_HPP
|
||||
@@ -0,0 +1,50 @@
|
||||
load("//tools:copts.bzl", "COPTS")
|
||||
|
||||
# Abstract base class for score calculators
|
||||
cc_library(
|
||||
name = "abstract_ai_score_calculator",
|
||||
srcs = ["AbstractAIScoreCalculator.cpp"],
|
||||
hdrs = ["AbstractAIScoreCalculator.hpp"],
|
||||
copts = COPTS,
|
||||
visibility = ["//src/main/cpp/net/eagle0/shardok/ai/score:__pkg__"], # Private to score package
|
||||
deps = [
|
||||
":ai_score_calculator_shared_utilities",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_attack_groups",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_attack_locations",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_score_utilities",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_strategy",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_unit_score_calculator",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_water_crossing_calculator",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai/score:ai_score_calculator_interface",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai/score:ai_victory_condition_score_calculator",
|
||||
"//src/main/cpp/net/eagle0/shardok/library:battalion_type",
|
||||
"//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/map:coords_set",
|
||||
"//src/main/cpp/net/eagle0/shardok/library/util:hex_map_utils",
|
||||
"//src/main/flatbuffer/net/eagle0/shardok/storage:game_state_cc_fbs",
|
||||
],
|
||||
)
|
||||
|
||||
# Private utilities shared between score calculators
|
||||
cc_library(
|
||||
name = "ai_score_calculator_shared_utilities",
|
||||
srcs = ["AIScoreCalculatorSharedUtilities.cpp"],
|
||||
hdrs = ["AIScoreCalculatorSharedUtilities.hpp"],
|
||||
copts = COPTS,
|
||||
visibility = ["//src/main/cpp/net/eagle0/shardok/ai/score:__pkg__"], # Private to score package
|
||||
deps = [
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_attack_groups",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_score_utilities",
|
||||
"//src/main/cpp/net/eagle0/shardok/library:battalion_type",
|
||||
"//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",
|
||||
"//src/main/cpp/net/eagle0/shardok/library/map:coords_set",
|
||||
"//src/main/flatbuffer/net/eagle0/shardok/storage:game_state_cc_fbs",
|
||||
"//src/main/flatbuffer/net/eagle0/shardok/storage:hex_map_cc_fbs",
|
||||
"//src/main/flatbuffer/net/eagle0/shardok/storage:unit_cc_fbs",
|
||||
"@gtl",
|
||||
],
|
||||
)
|
||||
@@ -7,12 +7,13 @@
|
||||
#include <iostream>
|
||||
|
||||
#include "src/main/cpp/net/eagle0/common/FilesystemUtils.hpp"
|
||||
#include "src/main/cpp/net/eagle0/common/TsvParser.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/ShardokAIClient.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai_testing_common/AIClientFactory.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai_testing_common/GamePhaseRunner.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai_testing_common/GameSettingsFactory.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/library/ShardokEngine.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/library/fb_helpers/GameStateHelpers.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/util/BattalionTypeRegistrar.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/util/MapLoader.hpp"
|
||||
#include "src/main/flatbuffer/net/eagle0/shardok/storage/unit.hpp"
|
||||
|
||||
@@ -67,20 +68,7 @@ AiBattleSimulator::AiBattleSimulator(const BattleConfigProto& config, GameSettin
|
||||
gameSettings_(std::move(gameSettings)) {
|
||||
if (!gameSettings_) {
|
||||
// Initialize default game settings
|
||||
gameSettings_ = std::make_shared<GameSettings>();
|
||||
auto setter = gameSettings_->GetSetter();
|
||||
|
||||
// Load battalion types
|
||||
BattalionTypeRegistrar::RegisterBattalionTypes(setter);
|
||||
|
||||
// Load settings from file
|
||||
const std::string settingsPath =
|
||||
FilesystemUtils::StaticShardokFilesDirectory() + "settings.tsv";
|
||||
const std::string settingsTsv = std::string(byte_vector::FromPath(settingsPath));
|
||||
|
||||
TsvParser parser;
|
||||
const auto valuesAndTypes = parser.ParseColumnEntryTsv(settingsTsv);
|
||||
setter.SetFromTypesAndValues(valuesAndTypes[1], valuesAndTypes[0]);
|
||||
gameSettings_ = ai_testing_common::GameSettingsFactory::CreateDefault();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -361,7 +349,7 @@ std::unique_ptr<ShardokAIClient> AiBattleSimulator::CreateAIClient(
|
||||
|
||||
AIAlgorithmType algorithmType = ConvertAIAlgorithmType(playerConfig.ai_algorithm());
|
||||
|
||||
return std::make_unique<ShardokAIClient>(
|
||||
return ai_testing_common::AIClientFactory::Create(
|
||||
playerId,
|
||||
isDefender,
|
||||
hexMap,
|
||||
@@ -373,44 +361,28 @@ BattleResult AiBattleSimulator::RunSetupPhase(
|
||||
ShardokEngine& engine,
|
||||
ShardokAIClient& attackerAI,
|
||||
ShardokAIClient& defenderAI) {
|
||||
int commandsExecuted = 0;
|
||||
auto getAI = [&](PlayerId playerId) -> ShardokAIClient& {
|
||||
return (playerId == ATTACKER_ID) ? attackerAI : defenderAI;
|
||||
};
|
||||
|
||||
while (engine.GetCurrentGameState()->status()->state() ==
|
||||
net::eagle0::shardok::storage::fb::GameStatus_::State_SET_UP) {
|
||||
auto currentState = engine.GetCurrentGameState();
|
||||
PlayerId currentPlayer = currentState->current_player();
|
||||
auto phaseResult = ai_testing_common::GamePhaseRunner::RunSetupPhase(engine, getAI);
|
||||
|
||||
auto availableCommands = engine.GetAvailableCommandProtos(currentPlayer, false);
|
||||
std::cout << "Setup phase complete. Commands executed: " << phaseResult.commandsExecuted
|
||||
<< "\n";
|
||||
|
||||
if (availableCommands.empty()) {
|
||||
std::cout << "No commands available during setup for player " << (int)currentPlayer
|
||||
<< "\n";
|
||||
break;
|
||||
}
|
||||
|
||||
// Choose which AI to use
|
||||
ShardokAIClient& activeAI = (currentPlayer == ATTACKER_ID) ? attackerAI : defenderAI;
|
||||
|
||||
// Get AI decision
|
||||
auto choiceResults = activeAI.ChooseCommandIndex(engine);
|
||||
|
||||
// Apply command
|
||||
engine.PostCommand(currentPlayer, choiceResults.chosenIndex);
|
||||
commandsExecuted++;
|
||||
|
||||
// Check if game ended unexpectedly
|
||||
if (engine.GameIsOver()) {
|
||||
return CreateResultFromGameState(engine.GetCurrentGameState(), 0, commandsExecuted);
|
||||
}
|
||||
// Check if game ended unexpectedly
|
||||
if (phaseResult.gameEnded) {
|
||||
return CreateResultFromGameState(
|
||||
engine.GetCurrentGameState(),
|
||||
0,
|
||||
phaseResult.commandsExecuted);
|
||||
}
|
||||
|
||||
std::cout << "Setup phase complete. Commands executed: " << commandsExecuted << "\n";
|
||||
|
||||
// Return a "not finished" result
|
||||
BattleResult result;
|
||||
result.winner = -1;
|
||||
result.totalRounds = 0;
|
||||
result.totalCommands = commandsExecuted;
|
||||
result.totalCommands = phaseResult.commandsExecuted;
|
||||
result.endReason = BattleResult::EndReason::DRAW; // Temporary placeholder
|
||||
result.description = "Setup phase completed";
|
||||
return result;
|
||||
|
||||
@@ -23,6 +23,9 @@ cc_library(
|
||||
"//src/main/cpp/net/eagle0/common:filesystem_utils",
|
||||
"//src/main/cpp/net/eagle0/common:tsv_parser",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:shardok_ai_client",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai_testing_common:ai_client_factory",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai_testing_common:game_phase_runner",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai_testing_common:game_settings_factory",
|
||||
"//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/fb_helpers:game_state_helpers",
|
||||
|
||||
@@ -13,6 +13,8 @@
|
||||
#include "src/main/cpp/net/eagle0/common/FilesystemUtils.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/AIConfig.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/ShardokAIClient.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai_testing_common/AIClientFactory.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai_testing_common/GamePhaseRunner.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/library/ShardokEngine.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/library/action_point_distances/FixedActionPointDistances.hpp"
|
||||
#include "src/main/protobuf/net/eagle0/shardok/common/command_type.pb.h"
|
||||
@@ -120,7 +122,7 @@ int main(int argc, char* argv[]) {
|
||||
const auto* hexMap = currentState->hex_map();
|
||||
const auto settingsGetter = settings->GetGetter();
|
||||
|
||||
ShardokAIClient aiClient(
|
||||
auto aiClient = ai_testing_common::AIClientFactory::Create(
|
||||
aiPlayerId,
|
||||
isDefender,
|
||||
hexMap,
|
||||
@@ -130,7 +132,7 @@ int main(int argc, char* argv[]) {
|
||||
// Create a second AI client for the human player during setup
|
||||
// This ensures consistent state handling during setup phase
|
||||
const PlayerId humanPlayerId = 1;
|
||||
ShardokAIClient humanSetupAI(
|
||||
auto humanSetupAI = ai_testing_common::AIClientFactory::Create(
|
||||
humanPlayerId,
|
||||
!isDefender,
|
||||
hexMap,
|
||||
@@ -140,29 +142,18 @@ int main(int argc, char* argv[]) {
|
||||
// Complete setup phase - AI makes intelligent placement decisions
|
||||
if (currentState->status()->state() ==
|
||||
net::eagle0::shardok::storage::fb::GameStatus_::State_SET_UP) {
|
||||
while (currentState->status()->state() ==
|
||||
net::eagle0::shardok::storage::fb::GameStatus_::State_SET_UP) {
|
||||
PlayerId currentPlayer = currentState->current_player();
|
||||
auto availableCommands = engine.GetAvailableCommandProtos(currentPlayer, false);
|
||||
auto getAI = [&](PlayerId playerId) -> ShardokAIClient& {
|
||||
return (playerId == aiPlayerId) ? *aiClient : *humanSetupAI;
|
||||
};
|
||||
|
||||
if (availableCommands.empty()) {
|
||||
std::cout << "No commands available for player "
|
||||
<< static_cast<int>(currentPlayer) << "\n";
|
||||
break;
|
||||
}
|
||||
auto setupResult = ai_testing_common::GamePhaseRunner::RunSetupPhase(engine, getAI);
|
||||
|
||||
if (currentPlayer == aiPlayerId) {
|
||||
// Let AI make intelligent placement decisions
|
||||
auto choiceResults = aiClient.ChooseCommandIndex(engine);
|
||||
engine.PostCommand(currentPlayer, choiceResults.chosenIndex);
|
||||
} else {
|
||||
// Human player: use AI for setup to ensure consistent state handling
|
||||
auto choiceResults = humanSetupAI.ChooseCommandIndex(engine);
|
||||
engine.PostCommand(currentPlayer, choiceResults.chosenIndex);
|
||||
}
|
||||
|
||||
currentState = engine.GetCurrentGameState();
|
||||
if (setupResult.gameEnded) {
|
||||
std::cout << "Game ended unexpectedly during setup phase.\n";
|
||||
return 0;
|
||||
}
|
||||
|
||||
currentState = engine.GetCurrentGameState();
|
||||
}
|
||||
|
||||
// Test AI performance for configured number of turns
|
||||
@@ -179,7 +170,7 @@ int main(int argc, char* argv[]) {
|
||||
}
|
||||
|
||||
// Get AI decision with performance metrics
|
||||
auto choiceResults = aiClient.ChooseCommandIndex(engine);
|
||||
auto choiceResults = aiClient->ChooseCommandIndex(engine);
|
||||
|
||||
std::cout << " AI chose command index: " << choiceResults.chosenIndex << "\n";
|
||||
std::cout << " Depth achieved: " << choiceResults.depthAchieved << "\n";
|
||||
|
||||
@@ -19,10 +19,13 @@ cc_binary(
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_attacker_strategy_selector",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_defender_strategy_selector",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_iterative_deepening",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_score_calculator_interface",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_time_budget",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_water_crossing_command_chooser",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:shardok_ai_client",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai/score:ai_score_calculator_interface",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai_testing_common:ai_client_factory",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai_testing_common:game_phase_runner",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai_testing_common:game_settings_factory",
|
||||
"//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",
|
||||
@@ -39,6 +42,7 @@ cc_library(
|
||||
],
|
||||
copts = COPTS,
|
||||
deps = [
|
||||
"//src/main/cpp/net/eagle0/shardok/ai_testing_common:game_settings_factory",
|
||||
"//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/fb_helpers:game_state_helpers",
|
||||
|
||||
@@ -95,7 +95,7 @@ cc_binary(
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_iterative_deepening",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_attacker_strategy_selector",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_defender_strategy_selector",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_score_calculator_interface",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai/score:ai_score_calculator_interface",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_time_budget",
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:ai_water_crossing_command_chooser",
|
||||
"//src/main/cpp/net/eagle0/shardok/library:engine",
|
||||
|
||||
+2
-17
@@ -7,11 +7,9 @@
|
||||
#include <filesystem>
|
||||
|
||||
#include "src/main/cpp/net/eagle0/common/FilesystemUtils.hpp"
|
||||
#include "src/main/cpp/net/eagle0/common/TsvParser.hpp"
|
||||
#include "src/main/cpp/net/eagle0/common/byte_vector.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai_testing_common/GameSettingsFactory.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/library/fb_helpers/GameStateHelpers.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/util/BattalionTypeRegistrar.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/util/MapLoader.hpp"
|
||||
#include "src/main/flatbuffer/net/eagle0/shardok/storage/unit.hpp"
|
||||
#include "src/main/protobuf/net/eagle0/shardok/common/player_info.pb.h"
|
||||
@@ -30,20 +28,7 @@ constexpr PlayerId HUMAN_PLAYER_ID = 1;
|
||||
} // namespace
|
||||
|
||||
auto PerformanceTestGameStateBuilder::InitializeGameSettings() -> GameSettingsSPtr {
|
||||
auto settings = std::make_shared<GameSettings>();
|
||||
auto setter = settings->GetSetter();
|
||||
|
||||
// Load battalion types
|
||||
BattalionTypeRegistrar::RegisterBattalionTypes(setter);
|
||||
|
||||
// Load complete settings from settings.tsv file
|
||||
TsvParser parser;
|
||||
const string settingsPath = FilesystemUtils::StaticShardokFilesDirectory() + "settings.tsv";
|
||||
const string settingsTsv = string(byte_vector::FromPath(settingsPath));
|
||||
const auto valuesAndTypes = parser.ParseColumnEntryTsv(settingsTsv);
|
||||
setter.SetFromTypesAndValues(valuesAndTypes[1], valuesAndTypes[0]);
|
||||
|
||||
return settings;
|
||||
return ai_testing_common::GameSettingsFactory::CreateDefault();
|
||||
}
|
||||
|
||||
auto PerformanceTestGameStateBuilder::CreatePerfTestGameState(
|
||||
|
||||
@@ -0,0 +1,469 @@
|
||||
# AI Testing Guide
|
||||
|
||||
This guide explains how to use the two AI testing tools in the Eagle0 codebase for evaluating and improving AI performance.
|
||||
|
||||
## Overview
|
||||
|
||||
The codebase provides two complementary tools for AI testing:
|
||||
|
||||
1. **AI Performance Runner** - Benchmarks AI decision-making quality over specific turns
|
||||
2. **AI Battle Simulator** - Tests AI effectiveness by running complete battles
|
||||
|
||||
Both tools support the Iterative Deepening and MCTS AI algorithms.
|
||||
|
||||
---
|
||||
|
||||
## AI Performance Runner
|
||||
|
||||
**Location**: `src/main/cpp/net/eagle0/shardok/ai_performance_runner/`
|
||||
|
||||
### Purpose
|
||||
|
||||
Measures AI decision-making performance to:
|
||||
- Identify performance regressions when making AI changes
|
||||
- Understand search depth achieved within time budgets
|
||||
- Analyze command evaluation patterns at each depth
|
||||
- Profile AI bottlenecks
|
||||
|
||||
### Building
|
||||
|
||||
```bash
|
||||
bazel build //src/main/cpp/net/eagle0/shardok/ai_performance_runner:ai_performance_runner
|
||||
```
|
||||
|
||||
### Usage
|
||||
|
||||
```bash
|
||||
bazel run //src/main/cpp/net/eagle0/shardok/ai_performance_runner:ai_performance_runner -- \
|
||||
--map=MAP_NAME \
|
||||
--turns=N \
|
||||
[--defender=true|false] \
|
||||
[--verbose]
|
||||
```
|
||||
|
||||
### Command-Line Arguments
|
||||
|
||||
| Argument | Required | Default | Description |
|
||||
|----------|----------|---------|-------------|
|
||||
| `--map` | Yes | - | Map name (e.g., "FourWay", "Bridge") |
|
||||
| `--turns` | Yes | - | Number of turns to run AI for |
|
||||
| `--defender` | No | false | If true, test defender AI; if false, test attacker AI |
|
||||
| `--verbose` | No | false | Enable verbose logging of AI decisions |
|
||||
|
||||
### Output Format
|
||||
|
||||
For each turn, outputs:
|
||||
```
|
||||
Turn 1:
|
||||
Depth achieved: 3
|
||||
Commands evaluated:
|
||||
Depth 1: 45/45 commands (100%)
|
||||
Depth 2: 127/234 commands (54%)
|
||||
Depth 3: 23/891 commands (3%)
|
||||
Completion reason: TIME_LIMIT
|
||||
Time elapsed: 5.2s
|
||||
```
|
||||
|
||||
### Performance Metrics Explained
|
||||
|
||||
- **Depth achieved**: Maximum lookahead depth the AI reached
|
||||
- **Commands evaluated**: At each depth, shows commands fully evaluated vs total available
|
||||
- Higher depth evaluations are more valuable (depth 3 > depth 2 > depth 1)
|
||||
- Completion rates show how thoroughly the AI searched each depth
|
||||
- **Completion reason**: Why the search stopped
|
||||
- `TIME_LIMIT`: Ran out of time budget (normal)
|
||||
- `COMPLETE`: Fully evaluated all possibilities (rare, usually only in simple endgames)
|
||||
- `DEPTH_LIMIT`: Hit maximum configured depth
|
||||
|
||||
### Example Workflow: Testing Performance Improvements
|
||||
|
||||
```bash
|
||||
# 1. Commit your baseline changes
|
||||
git checkout -b my-performance-improvement
|
||||
git add . && git commit -m "Baseline before optimization"
|
||||
|
||||
# 2. Run performance tests multiple times (reduce noise)
|
||||
for i in 1 2 3; do
|
||||
echo "=== Baseline Run $i ==="
|
||||
bazel run //src/main/cpp/net/eagle0/shardok/ai_performance_runner:ai_performance_runner -- \
|
||||
--map=FourWay --turns=5 --defender=true | grep -A 10 "Turn"
|
||||
done
|
||||
# Save the results
|
||||
|
||||
# 3. Make your optimization changes
|
||||
# ... edit code ...
|
||||
|
||||
# 4. Run tests again and compare
|
||||
for i in 1 2 3; do
|
||||
echo "=== Optimized Run $i ==="
|
||||
bazel run //src/main/cpp/net/eagle0/shardok/ai_performance_runner:ai_performance_runner -- \
|
||||
--map=FourWay --turns=5 --defender=true | grep -A 10 "Turn"
|
||||
done
|
||||
|
||||
# 5. Compare the results
|
||||
# Look for: increased depth, higher completion rates, more commands evaluated at deeper levels
|
||||
```
|
||||
|
||||
### When to Use Performance Runner
|
||||
|
||||
- ✅ Making changes to AI search algorithms
|
||||
- ✅ Optimizing performance-critical code paths
|
||||
- ✅ Identifying regressions in decision quality
|
||||
- ✅ Understanding where the AI spends its time
|
||||
- ❌ Testing which AI strategy wins more battles (use Battle Simulator instead)
|
||||
|
||||
---
|
||||
|
||||
## AI Battle Simulator
|
||||
|
||||
**Location**: `src/main/cpp/net/eagle0/shardok/ai_battle_simulator/`
|
||||
|
||||
### Purpose
|
||||
|
||||
Tests AI effectiveness by:
|
||||
- Running complete AI vs AI battles from start to finish
|
||||
- Measuring win rates across different AI configurations
|
||||
- Testing AI behavior with various unit compositions
|
||||
- Comparing different AI algorithms (MCTS vs Iterative Deepening)
|
||||
|
||||
### Building
|
||||
|
||||
```bash
|
||||
bazel build //src/main/cpp/net/eagle0/shardok/ai_battle_simulator:ai_battle_simulator
|
||||
```
|
||||
|
||||
### Usage
|
||||
|
||||
The battle simulator works with JSON configuration files.
|
||||
|
||||
#### Step 1: Generate a Sample Configuration
|
||||
|
||||
```bash
|
||||
bazel run //src/main/cpp/net/eagle0/shardok/ai_battle_simulator:ai_battle_simulator -- \
|
||||
--generate-config=my_battle_config.json
|
||||
```
|
||||
|
||||
This creates a sample configuration file you can customize.
|
||||
|
||||
#### Step 2: Customize the Configuration
|
||||
|
||||
Edit the generated JSON file:
|
||||
|
||||
```json
|
||||
{
|
||||
"map_name": "FourWay",
|
||||
"max_rounds": 100,
|
||||
"attacker": {
|
||||
"ai_algorithm": "ITERATIVE_DEEPENING",
|
||||
"units": [
|
||||
{
|
||||
"battalion_type": "HEAVY_INFANTRY",
|
||||
"row": 2,
|
||||
"column": 3,
|
||||
"strength": 1000,
|
||||
"has_hero": true,
|
||||
"hero_profession": "WARRIOR",
|
||||
"hero_level": 5
|
||||
},
|
||||
{
|
||||
"battalion_type": "ARCHERS",
|
||||
"row": 2,
|
||||
"column": 4,
|
||||
"strength": 800,
|
||||
"has_hero": false
|
||||
}
|
||||
]
|
||||
},
|
||||
"defender": {
|
||||
"ai_algorithm": "MCTS",
|
||||
"units": [
|
||||
{
|
||||
"battalion_type": "HEAVY_INFANTRY",
|
||||
"row": 8,
|
||||
"column": 3,
|
||||
"strength": 1000,
|
||||
"has_hero": true,
|
||||
"hero_profession": "WARRIOR",
|
||||
"hero_level": 5
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
#### Step 3: Run the Battle
|
||||
|
||||
```bash
|
||||
bazel run //src/main/cpp/net/eagle0/shardok/ai_battle_simulator:ai_battle_simulator -- \
|
||||
--config=my_battle_config.json
|
||||
```
|
||||
|
||||
### Configuration Format
|
||||
|
||||
#### Top-Level Fields
|
||||
|
||||
| Field | Type | Required | Description |
|
||||
|-------|------|----------|-------------|
|
||||
| `map_name` | string | Yes | Name of the map to use (e.g., "FourWay", "Bridge") |
|
||||
| `max_rounds` | integer | Yes | Maximum rounds before declaring a draw |
|
||||
| `attacker` | object | Yes | Attacker configuration (see below) |
|
||||
| `defender` | object | Yes | Defender configuration (see below) |
|
||||
|
||||
#### Player Configuration (attacker/defender)
|
||||
|
||||
| Field | Type | Required | Description |
|
||||
|-------|------|----------|-------------|
|
||||
| `ai_algorithm` | string | Yes | "ITERATIVE_DEEPENING" or "MCTS" |
|
||||
| `units` | array | Yes | List of unit configurations (see below) |
|
||||
|
||||
#### Unit Configuration
|
||||
|
||||
| Field | Type | Required | Default | Description |
|
||||
|-------|------|----------|---------|-------------|
|
||||
| `battalion_type` | string | Yes | - | "HEAVY_INFANTRY", "LIGHT_INFANTRY", "ARCHERS", "CAVALRY", etc. |
|
||||
| `row` | integer | Yes | - | Starting row position (0-indexed) |
|
||||
| `column` | integer | Yes | - | Starting column position (0-indexed) |
|
||||
| `strength` | integer | No | 1000 | Unit strength (max 1000) |
|
||||
| `has_hero` | boolean | No | false | Whether unit has an attached hero |
|
||||
| `hero_profession` | string | No* | - | "WARRIOR", "RANGER", "MAGE" (*required if has_hero=true) |
|
||||
| `hero_level` | integer | No | 1 | Hero level (1-10) |
|
||||
| `hero_vigor` | integer | No | 100 | Hero vigor (0-100) |
|
||||
|
||||
### Output Format
|
||||
|
||||
After running a battle, outputs:
|
||||
|
||||
```
|
||||
Battle Result:
|
||||
Winner: ATTACKER
|
||||
Total Rounds: 47
|
||||
Total Commands: 2,341
|
||||
|
||||
Attacker Survivors: 3 units
|
||||
- Heavy Infantry at (4,5): 723 strength
|
||||
- Archers at (3,6): 412 strength
|
||||
- Cavalry at (5,4): 891 strength
|
||||
|
||||
Defender Survivors: 0 units
|
||||
|
||||
Battle Duration: 14.2 seconds
|
||||
```
|
||||
|
||||
### Example Workflow: Testing AI Effectiveness
|
||||
|
||||
```bash
|
||||
# 1. Create a baseline configuration
|
||||
bazel run //src/main/cpp/net/eagle0/shardok/ai_battle_simulator:ai_battle_simulator -- \
|
||||
--generate-config=baseline.json
|
||||
|
||||
# 2. Edit baseline.json to set up your test scenario
|
||||
# Set both players to ITERATIVE_DEEPENING
|
||||
|
||||
# 3. Run multiple battles to get win rate
|
||||
for i in {1..20}; do
|
||||
echo "Battle $i:"
|
||||
bazel run //src/main/cpp/net/eagle0/shardok/ai_battle_simulator:ai_battle_simulator -- \
|
||||
--config=baseline.json | grep "Winner"
|
||||
done
|
||||
|
||||
# 4. Create a comparison configuration
|
||||
cp baseline.json mcts_comparison.json
|
||||
# Edit mcts_comparison.json: change attacker to MCTS
|
||||
|
||||
# 5. Run battles with MCTS attacker
|
||||
for i in {1..20}; do
|
||||
echo "Battle $i:"
|
||||
bazel run //src/main/cpp/net/eagle0/shardok/ai_battle_simulator:ai_battle_simulator -- \
|
||||
--config=mcts_comparison.json | grep "Winner"
|
||||
done
|
||||
|
||||
# 6. Compare win rates
|
||||
# Count ATTACKER wins in each set to see if MCTS performs better/worse
|
||||
```
|
||||
|
||||
### Advanced: Testing Specific Scenarios
|
||||
|
||||
The battle simulator is ideal for testing:
|
||||
|
||||
**Scenario 1: Hero Ability Usage**
|
||||
```json
|
||||
{
|
||||
"attacker": {
|
||||
"units": [
|
||||
{
|
||||
"battalion_type": "LIGHT_INFANTRY",
|
||||
"has_hero": true,
|
||||
"hero_profession": "RANGER",
|
||||
"hero_level": 8,
|
||||
"row": 2, "column": 3
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
```
|
||||
Tests whether AI properly uses Ranger stealth abilities.
|
||||
|
||||
**Scenario 2: Terrain Advantage**
|
||||
```json
|
||||
{
|
||||
"map_name": "BridgeChoke",
|
||||
"defender": {
|
||||
"units": [
|
||||
{"battalion_type": "HEAVY_INFANTRY", "row": 5, "column": 4}
|
||||
]
|
||||
}
|
||||
}
|
||||
```
|
||||
Tests whether AI exploits defensive terrain.
|
||||
|
||||
**Scenario 3: Numerical Superiority**
|
||||
```json
|
||||
{
|
||||
"attacker": {
|
||||
"units": [
|
||||
{"battalion_type": "ARCHERS", "row": 2, "column": 2},
|
||||
{"battalion_type": "ARCHERS", "row": 2, "column": 3},
|
||||
{"battalion_type": "ARCHERS", "row": 2, "column": 4}
|
||||
]
|
||||
},
|
||||
"defender": {
|
||||
"units": [
|
||||
{"battalion_type": "CAVALRY", "row": 8, "column": 3}
|
||||
]
|
||||
}
|
||||
}
|
||||
```
|
||||
Tests whether AI properly coordinates multiple units.
|
||||
|
||||
### When to Use Battle Simulator
|
||||
|
||||
- ✅ Testing which AI strategy wins more battles
|
||||
- ✅ Measuring AI effectiveness with different unit types
|
||||
- ✅ Comparing MCTS vs Iterative Deepening algorithms
|
||||
- ✅ Validating AI behavior in specific tactical scenarios
|
||||
- ✅ Regression testing: ensuring AI changes don't reduce win rates
|
||||
- ❌ Profiling performance bottlenecks (use Performance Runner instead)
|
||||
|
||||
---
|
||||
|
||||
## Combining Both Tools
|
||||
|
||||
For comprehensive AI evaluation:
|
||||
|
||||
1. **Make AI changes** to improve decision-making
|
||||
2. **Run Performance Runner** to verify the AI searches deeper or evaluates more commands
|
||||
3. **Run Battle Simulator** to verify the changes actually improve win rates
|
||||
4. **Iterate** if performance improves but win rate doesn't (or vice versa)
|
||||
|
||||
### Example: Evaluating a New Heuristic
|
||||
|
||||
```bash
|
||||
# 1. Baseline performance
|
||||
bazel run //src/main/cpp/net/eagle0/shardok/ai_performance_runner:ai_performance_runner -- \
|
||||
--map=FourWay --turns=3 > baseline_perf.txt
|
||||
|
||||
# 2. Baseline effectiveness
|
||||
for i in {1..10}; do
|
||||
bazel run //src/main/cpp/net/eagle0/shardok/ai_battle_simulator:ai_battle_simulator -- \
|
||||
--config=test_scenario.json | grep "Winner"
|
||||
done > baseline_wins.txt
|
||||
|
||||
# 3. Make changes to heuristic
|
||||
# ... edit code ...
|
||||
|
||||
# 4. New performance
|
||||
bazel run //src/main/cpp/net/eagle0/shardok/ai_performance_runner:ai_performance_runner -- \
|
||||
--map=FourWay --turns=3 > new_perf.txt
|
||||
|
||||
# 5. New effectiveness
|
||||
for i in {1..10}; do
|
||||
bazel run //src/main/cpp/net/eagle0/shardok/ai_battle_simulator:ai_battle_simulator -- \
|
||||
--config=test_scenario.json | grep "Winner"
|
||||
done > new_wins.txt
|
||||
|
||||
# 6. Compare results
|
||||
diff baseline_perf.txt new_perf.txt
|
||||
wc -l < baseline_wins.txt | grep ATTACKER # Count baseline wins
|
||||
wc -l < new_wins.txt | grep ATTACKER # Count new wins
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Available Maps
|
||||
|
||||
Common maps for testing:
|
||||
- **FourWay**: Open terrain with multiple approach paths
|
||||
- **Bridge**: Chokepoint scenario testing tactical positioning
|
||||
- **Forest**: Tests unit behavior in hiding terrain
|
||||
- **Mountain**: Tests pathfinding around impassable terrain
|
||||
|
||||
Find all available maps in: `src/main/resources/net/eagle0/shardok/maps/`
|
||||
|
||||
---
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
### Performance Runner Issues
|
||||
|
||||
**Issue**: "Map not found"
|
||||
```bash
|
||||
# Solution: Use exact map name without .e0mj extension
|
||||
# ✅ Correct:
|
||||
--map=FourWay
|
||||
# ❌ Incorrect:
|
||||
--map=FourWay.e0mj
|
||||
```
|
||||
|
||||
**Issue**: AI completes instantly
|
||||
```bash
|
||||
# Likely cause: Too few turns or already-won scenario
|
||||
# Solution: Increase --turns or use more complex map
|
||||
--turns=10 # Instead of --turns=1
|
||||
```
|
||||
|
||||
### Battle Simulator Issues
|
||||
|
||||
**Issue**: "Invalid unit position"
|
||||
```bash
|
||||
# Solution: Ensure positions are within map bounds and not overlapping
|
||||
# Check map size first, then place units accordingly
|
||||
```
|
||||
|
||||
**Issue**: "Battle ends immediately"
|
||||
```bash
|
||||
# Cause: Units placed too close or in invalid starting positions
|
||||
# Solution:
|
||||
# - Attacker units should start on attacker side (low row numbers)
|
||||
# - Defender units should start on defender side (high row numbers)
|
||||
# - Leave space between forces for tactical maneuvering
|
||||
```
|
||||
|
||||
**Issue**: Battle runs extremely slowly
|
||||
```bash
|
||||
# Cause: Too many units or max_rounds too high
|
||||
# Solution: Start with 2-3 units per side, max_rounds=50
|
||||
# Scale up once basic scenario works
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Best Practices
|
||||
|
||||
### For Performance Testing
|
||||
1. **Run multiple iterations** (3-5) to reduce variance
|
||||
2. **Test on consistent maps** to enable before/after comparison
|
||||
3. **Focus on depth and completion rates** rather than raw time
|
||||
4. **Test both attacker and defender** AIs (they use different strategies)
|
||||
|
||||
### For Effectiveness Testing
|
||||
1. **Run 10-20 battles** minimum for statistical significance
|
||||
2. **Use symmetric scenarios** (equal forces) to isolate AI quality
|
||||
3. **Test multiple maps** to avoid overfitting to one scenario
|
||||
4. **Document unit compositions** so tests are reproducible
|
||||
5. **Version control configs** alongside code changes
|
||||
|
||||
### For Both
|
||||
1. **Commit before testing** so you can easily revert
|
||||
2. **Document unexpected results** (AI might be correct, your intuition wrong)
|
||||
3. **Test edge cases** (one unit, heroes only, terrain heavy)
|
||||
4. **Compare algorithms** (MCTS vs Iterative Deepening) regularly
|
||||
+647
@@ -0,0 +1,647 @@
|
||||
# Proposal: Server-Based AI Effectiveness Testing
|
||||
|
||||
## Executive Summary
|
||||
|
||||
Replace proxy-based performance metrics (search depth, nodes visited) with **actual effectiveness testing** (win rates, battle outcomes) by running AI battles through the production Shardok server. This measures what matters: whether AI improvements make the AI smarter, not just faster.
|
||||
|
||||
## Problem Statement
|
||||
|
||||
### Current State: Measuring the Wrong Things
|
||||
|
||||
The existing `ai_performance_runner` measures proxy metrics:
|
||||
- Search depth achieved
|
||||
- Number of commands evaluated
|
||||
- Time spent searching
|
||||
|
||||
**Problem**: These metrics don't tell us if the AI is making good decisions.
|
||||
|
||||
**Example failure mode**:
|
||||
- AI searches to depth 4 (looks impressive!)
|
||||
- But uses terrible heuristics (all decisions are bad)
|
||||
- Result: Loses every battle despite "good" metrics
|
||||
|
||||
### What We Actually Care About
|
||||
|
||||
- **Does the AI win?** (win rate)
|
||||
- **By what margin?** (survivors, rounds taken)
|
||||
- **Is it tactically sound?** (decision quality in specific scenarios)
|
||||
- **Does it handle edge cases?** (terrain, heroes, special abilities)
|
||||
|
||||
## Proposed Solution
|
||||
|
||||
Build a **server-based AI effectiveness testing framework** that:
|
||||
1. Runs battles through the production Shardok server (real code paths)
|
||||
2. Measures actual effectiveness (win rates, outcomes)
|
||||
3. Supports repeatable test scenarios
|
||||
4. Enables comparison between AI algorithms (MCTS vs Iterative Deepening)
|
||||
5. Eventually allows Unity clients to watch battles
|
||||
|
||||
## Architecture
|
||||
|
||||
### Component Overview
|
||||
|
||||
```
|
||||
┌─────────────────────┐
|
||||
│ Test Scenarios │
|
||||
│ (JSON configs) │
|
||||
└──────────┬──────────┘
|
||||
│
|
||||
▼
|
||||
┌─────────────────────┐
|
||||
│ Go Test Client │
|
||||
│ - Loads scenarios │
|
||||
│ - Sends gRPC reqs │
|
||||
│ - Collects results │
|
||||
└──────────┬──────────┘
|
||||
│ gRPC
|
||||
▼
|
||||
┌─────────────────────┐
|
||||
│ Shardok Server │
|
||||
│ - Creates games │
|
||||
│ - Runs AI vs AI │
|
||||
│ - Returns outcomes │
|
||||
└─────────────────────┘
|
||||
```
|
||||
|
||||
### Why Go for the Client?
|
||||
|
||||
- Native gRPC support with `protoc-gen-go-grpc`
|
||||
- Easy JSON config parsing
|
||||
- Good for CLI tools
|
||||
- Fast compile times for iteration
|
||||
- Excellent concurrency for running parallel test scenarios
|
||||
|
||||
## Implementation Plan
|
||||
|
||||
### Phase 1: Core Framework (1 week)
|
||||
|
||||
**Goal**: Basic server-based effectiveness testing
|
||||
|
||||
#### 1.1 Protocol Extensions
|
||||
|
||||
Add to `src/main/protobuf/net/eagle0/common/shardok_internal_interface.proto`:
|
||||
|
||||
```protobuf
|
||||
message TestBattleRequest {
|
||||
string game_id = 1;
|
||||
string map_name = 2;
|
||||
|
||||
message PlayerSetup {
|
||||
bool is_defender = 1;
|
||||
AIAlgorithmType ai_algorithm = 2; // MCTS or ITERATIVE_DEEPENING
|
||||
ScoringCalculatorType scoring_type = 3; // STANDARD, NORMALIZED, MCTS_OPTIMIZED
|
||||
repeated UnitPlacement units = 4;
|
||||
}
|
||||
|
||||
repeated PlayerSetup players = 3;
|
||||
int32 max_rounds = 4;
|
||||
}
|
||||
|
||||
message UnitPlacement {
|
||||
string battalion_type = 1; // "HEAVY_INFANTRY", "ARCHERS", etc.
|
||||
int32 row = 2;
|
||||
int32 column = 3;
|
||||
int32 strength = 4;
|
||||
bool has_hero = 5;
|
||||
string hero_profession = 6; // "WARRIOR", "RANGER", "MAGE"
|
||||
int32 hero_level = 7;
|
||||
}
|
||||
|
||||
message TestBattleResponse {
|
||||
string game_id = 1;
|
||||
int32 winner_player_id = 2; // -1 for draw
|
||||
int32 rounds_taken = 3;
|
||||
string end_reason = 4; // "VICTORY", "DRAW", "MAX_ROUNDS"
|
||||
repeated FinalUnitStatus final_units = 5;
|
||||
}
|
||||
|
||||
message FinalUnitStatus {
|
||||
int32 player_id = 1;
|
||||
string battalion_type = 2;
|
||||
int32 strength_remaining = 3;
|
||||
int32 row = 4;
|
||||
int32 column = 5;
|
||||
}
|
||||
|
||||
// Add to ShardokInternalInterface service:
|
||||
rpc StartTestBattle(TestBattleRequest) returns (TestBattleResponse) {}
|
||||
```
|
||||
|
||||
**Estimated effort**: 1 day
|
||||
|
||||
#### 1.2 Server Implementation
|
||||
|
||||
Add handler to `src/main/cpp/net/eagle0/shardok/server/EagleInterfaceGrpcServer.cpp`:
|
||||
|
||||
```cpp
|
||||
Status EagleInterfaceImpl::StartTestBattle(
|
||||
ServerContext* context,
|
||||
const TestBattleRequest* request,
|
||||
TestBattleResponse* response) {
|
||||
// 1. Create game with specified setup
|
||||
// 2. Mark all players as is_ai=true
|
||||
// 3. Wait for game to complete (AI thread handles it)
|
||||
// 4. Collect final state and return results
|
||||
}
|
||||
```
|
||||
|
||||
**Key insight**: The infrastructure already exists! `ShardokGameController` already:
|
||||
- Creates AI clients automatically for `is_ai=true` players
|
||||
- Runs AI decisions in background thread
|
||||
- Tracks game state and completion
|
||||
|
||||
**Estimated effort**: 2 days
|
||||
|
||||
#### 1.3 Go Test Client
|
||||
|
||||
Create `src/main/go/net/eagle0/shardok/ai_effectiveness_runner/`:
|
||||
|
||||
```go
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"flag"
|
||||
"fmt"
|
||||
"log"
|
||||
|
||||
pb "eagle0/protobuf/net/eagle0/common"
|
||||
"google.golang.org/grpc"
|
||||
)
|
||||
|
||||
type ScenarioConfig struct {
|
||||
Name string `json:"name"`
|
||||
MapName string `json:"map_name"`
|
||||
MaxRounds int32 `json:"max_rounds"`
|
||||
Attacker PlayerConfig `json:"attacker"`
|
||||
Defender PlayerConfig `json:"defender"`
|
||||
}
|
||||
|
||||
type PlayerConfig struct {
|
||||
AIAlgorithm string `json:"ai_algorithm"`
|
||||
ScoringType string `json:"scoring_type"`
|
||||
Units []UnitPlacement `json:"units"`
|
||||
}
|
||||
|
||||
func runTestBattle(client pb.ShardokInternalInterfaceClient, scenario ScenarioConfig) (*pb.TestBattleResponse, error) {
|
||||
ctx := context.Background()
|
||||
|
||||
req := &pb.TestBattleRequest{
|
||||
GameId: fmt.Sprintf("test_%s_%d", scenario.Name, time.Now().Unix()),
|
||||
MapName: scenario.MapName,
|
||||
MaxRounds: scenario.MaxRounds,
|
||||
Players: []*pb.TestBattleRequest_PlayerSetup{
|
||||
{
|
||||
IsDefender: false,
|
||||
AiAlgorithm: parseAIAlgorithm(scenario.Attacker.AIAlgorithm),
|
||||
ScoringType: parseScoringType(scenario.Attacker.ScoringType),
|
||||
Units: convertUnits(scenario.Attacker.Units),
|
||||
},
|
||||
{
|
||||
IsDefender: true,
|
||||
AiAlgorithm: parseAIAlgorithm(scenario.Defender.AIAlgorithm),
|
||||
ScoringType: parseScoringType(scenario.Defender.ScoringType),
|
||||
Units: convertUnits(scenario.Defender.Units),
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
return client.StartTestBattle(ctx, req)
|
||||
}
|
||||
|
||||
func main() {
|
||||
serverAddr := flag.String("server", "localhost:50051", "Shardok server address")
|
||||
scenariosFile := flag.String("scenarios", "test_scenarios.json", "Test scenarios JSON file")
|
||||
iterations := flag.Int("iterations", 20, "Number of iterations per scenario")
|
||||
|
||||
flag.Parse()
|
||||
|
||||
// Load scenarios
|
||||
scenarios := loadScenarios(*scenariosFile)
|
||||
|
||||
// Connect to server
|
||||
conn, err := grpc.Dial(*serverAddr, grpc.WithInsecure())
|
||||
if err != nil {
|
||||
log.Fatalf("Failed to connect: %v", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
client := pb.NewShardokInternalInterfaceClient(conn)
|
||||
|
||||
// Run effectiveness tests
|
||||
results := make(map[string]*ScenarioResults)
|
||||
|
||||
for _, scenario := range scenarios {
|
||||
fmt.Printf("\n=== Testing Scenario: %s ===\n", scenario.Name)
|
||||
|
||||
scenarioResults := &ScenarioResults{
|
||||
ScenarioName: scenario.Name,
|
||||
}
|
||||
|
||||
for i := 0; i < *iterations; i++ {
|
||||
resp, err := runTestBattle(client, scenario)
|
||||
if err != nil {
|
||||
log.Printf("Battle %d failed: %v", i+1, err)
|
||||
continue
|
||||
}
|
||||
|
||||
scenarioResults.TotalBattles++
|
||||
if resp.WinnerPlayerId == 0 { // Attacker wins
|
||||
scenarioResults.AttackerWins++
|
||||
} else if resp.WinnerPlayerId == 1 { // Defender wins
|
||||
scenarioResults.DefenderWins++
|
||||
} else {
|
||||
scenarioResults.Draws++
|
||||
}
|
||||
|
||||
scenarioResults.TotalRounds += int(resp.RoundsTaken)
|
||||
scenarioResults.TotalSurvivors += len(resp.FinalUnits)
|
||||
|
||||
fmt.Printf(" Battle %d/%d: Winner=%d, Rounds=%d, Survivors=%d\n",
|
||||
i+1, *iterations, resp.WinnerPlayerId, resp.RoundsTaken, len(resp.FinalUnits))
|
||||
}
|
||||
|
||||
results[scenario.Name] = scenarioResults
|
||||
}
|
||||
|
||||
// Print summary
|
||||
printSummary(results)
|
||||
}
|
||||
```
|
||||
|
||||
**Estimated effort**: 2 days
|
||||
|
||||
#### 1.4 Test Scenario Definitions
|
||||
|
||||
Create `test_scenarios.json`:
|
||||
|
||||
```json
|
||||
{
|
||||
"scenarios": [
|
||||
{
|
||||
"name": "mcts_vs_id_balanced",
|
||||
"map_name": "FourWay",
|
||||
"max_rounds": 100,
|
||||
"attacker": {
|
||||
"ai_algorithm": "MCTS",
|
||||
"scoring_type": "MCTS_OPTIMIZED",
|
||||
"units": [
|
||||
{
|
||||
"battalion_type": "HEAVY_INFANTRY",
|
||||
"row": 2,
|
||||
"column": 3,
|
||||
"strength": 1000,
|
||||
"has_hero": true,
|
||||
"hero_profession": "WARRIOR",
|
||||
"hero_level": 5
|
||||
},
|
||||
{
|
||||
"battalion_type": "ARCHERS",
|
||||
"row": 2,
|
||||
"column": 4,
|
||||
"strength": 800,
|
||||
"has_hero": false
|
||||
}
|
||||
]
|
||||
},
|
||||
"defender": {
|
||||
"ai_algorithm": "ITERATIVE_DEEPENING",
|
||||
"scoring_type": "STANDARD",
|
||||
"units": [
|
||||
{
|
||||
"battalion_type": "HEAVY_INFANTRY",
|
||||
"row": 8,
|
||||
"column": 3,
|
||||
"strength": 1000,
|
||||
"has_hero": true,
|
||||
"hero_profession": "WARRIOR",
|
||||
"hero_level": 5
|
||||
},
|
||||
{
|
||||
"battalion_type": "ARCHERS",
|
||||
"row": 8,
|
||||
"column": 4,
|
||||
"strength": 800,
|
||||
"has_hero": false
|
||||
}
|
||||
]
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "terrain_advantage_test",
|
||||
"map_name": "Bridge",
|
||||
"max_rounds": 100,
|
||||
"attacker": {
|
||||
"ai_algorithm": "MCTS",
|
||||
"scoring_type": "MCTS_OPTIMIZED",
|
||||
"units": [
|
||||
{
|
||||
"battalion_type": "LIGHT_INFANTRY",
|
||||
"row": 2,
|
||||
"column": 5,
|
||||
"strength": 1000,
|
||||
"has_hero": false
|
||||
}
|
||||
]
|
||||
},
|
||||
"defender": {
|
||||
"ai_algorithm": "MCTS",
|
||||
"scoring_type": "MCTS_OPTIMIZED",
|
||||
"units": [
|
||||
{
|
||||
"battalion_type": "HEAVY_INFANTRY",
|
||||
"row": 10,
|
||||
"column": 5,
|
||||
"strength": 800,
|
||||
"has_hero": false
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
**Estimated effort**: 1 day (create comprehensive test suite)
|
||||
|
||||
### Phase 2: Enhanced Observability (1 week)
|
||||
|
||||
**Goal**: Stream battle updates for real-time monitoring
|
||||
|
||||
#### 2.1 Streaming Protocol
|
||||
|
||||
Extend protocol to support streaming updates:
|
||||
|
||||
```protobuf
|
||||
message TestBattleUpdate {
|
||||
enum UpdateType {
|
||||
SETUP_COMPLETE = 0;
|
||||
ROUND_START = 1;
|
||||
AI_DECISION = 2;
|
||||
ROUND_END = 3;
|
||||
BATTLE_END = 4;
|
||||
}
|
||||
|
||||
UpdateType type = 1;
|
||||
int32 current_round = 2;
|
||||
|
||||
// For AI_DECISION updates
|
||||
int32 player_id = 3;
|
||||
string command_type = 4; // "MOVE", "MELEE", "ARCHERY", etc.
|
||||
|
||||
// Optional: Include AI metrics as secondary data
|
||||
AIDecisionMetrics ai_metrics = 5;
|
||||
|
||||
// For BATTLE_END updates
|
||||
TestBattleResponse final_result = 6;
|
||||
}
|
||||
|
||||
message AIDecisionMetrics {
|
||||
int32 depth_achieved = 1;
|
||||
int32 commands_evaluated = 2;
|
||||
int64 time_spent_ms = 3;
|
||||
}
|
||||
|
||||
// Add streaming RPC:
|
||||
rpc StartTestBattleStreaming(TestBattleRequest) returns (stream TestBattleUpdate) {}
|
||||
```
|
||||
|
||||
**Benefits**:
|
||||
- Real-time battle monitoring
|
||||
- Can log/replay interesting decisions
|
||||
- Still captures proxy metrics as secondary data (if desired)
|
||||
- Enables debugging of specific scenarios
|
||||
|
||||
**Estimated effort**: 3 days
|
||||
|
||||
#### 2.2 Go Client Updates
|
||||
|
||||
Add streaming support to Go client:
|
||||
|
||||
```go
|
||||
func runTestBattleStreaming(client pb.ShardokInternalInterfaceClient, scenario ScenarioConfig) (*pb.TestBattleResponse, error) {
|
||||
ctx := context.Background()
|
||||
stream, err := client.StartTestBattleStreaming(ctx, req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
for {
|
||||
update, err := stream.Recv()
|
||||
if err == io.EOF {
|
||||
break
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
switch update.Type {
|
||||
case pb.TestBattleUpdate_AI_DECISION:
|
||||
fmt.Printf(" Round %d: Player %d chose %s\n",
|
||||
update.CurrentRound, update.PlayerId, update.CommandType)
|
||||
case pb.TestBattleUpdate_BATTLE_END:
|
||||
return update.FinalResult, nil
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**Estimated effort**: 2 days
|
||||
|
||||
### Phase 3: Unity Client Integration (2 weeks)
|
||||
|
||||
**Goal**: Enable Unity clients to watch AI battles in real-time
|
||||
|
||||
#### 3.1 Route Through Eagle
|
||||
|
||||
Extend Eagle server to accept test battle requests and create Shardok games:
|
||||
|
||||
```scala
|
||||
// In Eagle server
|
||||
def startTestBattle(request: TestBattleRequest): Unit = {
|
||||
// 1. Create game state in Eagle
|
||||
// 2. Send to Shardok via existing protocol
|
||||
// 3. Mark both players as AI
|
||||
// 4. Allow Unity clients to connect and watch
|
||||
}
|
||||
```
|
||||
|
||||
**Benefits**:
|
||||
- Tests complete production stack (Eagle + Shardok)
|
||||
- Unity clients can connect as spectators
|
||||
- Most realistic integration test possible
|
||||
|
||||
**Estimated effort**: 1 week
|
||||
|
||||
#### 3.2 Unity Spectator Mode
|
||||
|
||||
Add spectator mode to Unity client:
|
||||
|
||||
```csharp
|
||||
// In Unity client
|
||||
public class AIBattleSpectator : MonoBehaviour {
|
||||
public void ConnectToBattle(string gameId) {
|
||||
// Connect to Eagle as spectator
|
||||
// Receive battle updates
|
||||
// Render on screen
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**Estimated effort**: 1 week
|
||||
|
||||
## Usage Examples
|
||||
|
||||
### Basic Effectiveness Testing
|
||||
|
||||
```bash
|
||||
# Start Shardok server
|
||||
bazel run //src/main/cpp/net/eagle0/shardok:shardok-server
|
||||
|
||||
# Run effectiveness tests
|
||||
bazel run //src/main/go/net/eagle0/shardok/ai_effectiveness_runner:ai_effectiveness_runner -- \
|
||||
--server=localhost:50051 \
|
||||
--scenarios=test_scenarios.json \
|
||||
--iterations=20
|
||||
|
||||
# Output:
|
||||
# === Testing Scenario: mcts_vs_id_balanced ===
|
||||
# Battle 1/20: Winner=0, Rounds=47, Survivors=3
|
||||
# Battle 2/20: Winner=1, Rounds=52, Survivors=2
|
||||
# ...
|
||||
#
|
||||
# === Summary ===
|
||||
# Scenario: mcts_vs_id_balanced
|
||||
# MCTS (attacker) wins: 15/20 (75%)
|
||||
# Iterative Deepening (defender) wins: 5/20 (25%)
|
||||
# Average rounds: 48.3
|
||||
# Average survivors: 2.8
|
||||
```
|
||||
|
||||
### Comparing AI Improvements
|
||||
|
||||
```bash
|
||||
# Before optimization
|
||||
git checkout main
|
||||
bazel run //src/main/go/net/eagle0/shardok/ai_effectiveness_runner:ai_effectiveness_runner -- \
|
||||
--scenarios=regression_suite.json --iterations=50 > baseline_results.txt
|
||||
|
||||
# After optimization
|
||||
git checkout my-ai-improvement
|
||||
bazel run //src/main/go/net/eagle0/shardok/ai_effectiveness_runner:ai_effectiveness_runner -- \
|
||||
--scenarios=regression_suite.json --iterations=50 > improved_results.txt
|
||||
|
||||
# Compare
|
||||
diff baseline_results.txt improved_results.txt
|
||||
# Shows: MCTS win rate improved from 60% to 75%! ✅
|
||||
```
|
||||
|
||||
### Watching Battles in Unity
|
||||
|
||||
```bash
|
||||
# Terminal 1: Eagle server
|
||||
bazel run //src/main/scala/net/eagle0/eagle:eagle_server
|
||||
|
||||
# Terminal 2: Shardok server
|
||||
bazel run //src/main/cpp/net/eagle0/shardok:shardok-server
|
||||
|
||||
# Terminal 3: Start test battle
|
||||
bazel run //src/main/go/net/eagle0/shardok/ai_effectiveness_runner:ai_effectiveness_runner -- \
|
||||
--server=localhost:40032 \
|
||||
--scenarios=interesting_scenario.json \
|
||||
--stream
|
||||
|
||||
# Terminal 4: Unity client
|
||||
# Open Unity, connect as spectator to watch battle unfold
|
||||
```
|
||||
|
||||
## Success Metrics
|
||||
|
||||
### Immediate (Phase 1)
|
||||
- ✅ Can run 20+ battle scenarios against server
|
||||
- ✅ Measures win rates, rounds, survivors
|
||||
- ✅ Tests real production Shardok server code paths
|
||||
- ✅ Reproducible results
|
||||
|
||||
### Medium-term (Phase 2)
|
||||
- ✅ Streaming battle updates work
|
||||
- ✅ Can capture and replay interesting battles
|
||||
- ✅ Logs include both effectiveness metrics and proxy metrics
|
||||
|
||||
### Long-term (Phase 3)
|
||||
- ✅ Unity clients can watch AI battles
|
||||
- ✅ Tests complete Eagle + Shardok stack
|
||||
- ✅ Community can watch AI improvements
|
||||
|
||||
## Migration Strategy
|
||||
|
||||
### Keep Existing Tools
|
||||
|
||||
**AI Battle Simulator** (in-process, keep for development):
|
||||
- Fast iteration during development
|
||||
- Easy debugging (direct access to internals)
|
||||
- Use case: "Does this change work at all?"
|
||||
|
||||
**AI Effectiveness Runner** (server-based, new primary tool):
|
||||
- Tests production code paths
|
||||
- Measures real effectiveness
|
||||
- Use case: "Is this change actually better?"
|
||||
|
||||
### Deprecate Performance Runner
|
||||
|
||||
The current `ai_performance_runner` measures proxy metrics. Recommend:
|
||||
1. Keep it temporarily for comparison
|
||||
2. After Phase 2 (streaming + AI metrics), deprecate it
|
||||
3. Streaming effectiveness runner includes proxy metrics as secondary data
|
||||
|
||||
## Open Questions
|
||||
|
||||
1. **Server Resource Management**: Should we limit concurrent test battles on the server?
|
||||
- Proposal: Add `--max-concurrent-battles` flag to Go client
|
||||
|
||||
2. **Scenario Versioning**: How do we ensure scenarios remain valid as game evolves?
|
||||
- Proposal: Version scenarios in git, validate against server on load
|
||||
|
||||
3. **Metrics Storage**: Should we store historical effectiveness metrics?
|
||||
- Proposal: Phase 4 (future) - add database for tracking AI effectiveness over time
|
||||
|
||||
4. **Randomness Control**: How do we handle dice roll randomness in battles?
|
||||
- Current: Protocol supports `roll` override, but battles have many rolls
|
||||
- Proposal: Add `random_seed` to TestBattleRequest for reproducibility
|
||||
|
||||
## Timeline
|
||||
|
||||
- **Phase 1**: 1 week (core framework)
|
||||
- **Phase 2**: 1 week (streaming + observability)
|
||||
- **Phase 3**: 2 weeks (Unity integration)
|
||||
|
||||
**Total**: 4 weeks for complete vision
|
||||
|
||||
**Minimal viable**: Phase 1 only (1 week) provides immediate value
|
||||
|
||||
## Next Steps
|
||||
|
||||
1. Review and approve this proposal
|
||||
2. Merge current PR (#4518) - AI testing infrastructure refactor
|
||||
3. Begin Phase 1 implementation:
|
||||
- Protocol extensions (1 day)
|
||||
- Server implementation (2 days)
|
||||
- Go client (2 days)
|
||||
- Test scenarios (1 day)
|
||||
4. Validate with initial test runs
|
||||
5. Iterate based on findings
|
||||
6. Plan Phase 2 based on Phase 1 learnings
|
||||
|
||||
## Conclusion
|
||||
|
||||
This proposal shifts AI testing from measuring **proxies** (search depth) to measuring **reality** (win rates). By running battles through the production server, we:
|
||||
|
||||
- Test what matters: actual intelligence
|
||||
- Validate real code paths: gRPC, threading, server logic
|
||||
- Enable future capabilities: Unity viewing, distributed testing
|
||||
- Measure improvements objectively: win rate changes
|
||||
|
||||
The infrastructure mostly exists - ShardokGameController already handles AI vs AI battles. We just need to expose it via protocol and build a client to drive it.
|
||||
|
||||
**Recommendation**: Approve and implement Phase 1 (1 week) to immediately gain better AI effectiveness measurement.
|
||||
@@ -0,0 +1,34 @@
|
||||
//
|
||||
// AIClientFactory Implementation
|
||||
//
|
||||
|
||||
#include "AIClientFactory.hpp"
|
||||
|
||||
#include "src/main/cpp/net/eagle0/common/mcts/abstract/MCTSTypes.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/ShardokAIClient.hpp"
|
||||
|
||||
namespace shardok {
|
||||
namespace ai_testing_common {
|
||||
|
||||
auto AIClientFactory::Create(
|
||||
PlayerId playerId,
|
||||
bool isDefender,
|
||||
const net::eagle0::shardok::storage::fb::HexMap* hexMap,
|
||||
const SettingsGetter& settings,
|
||||
AIAlgorithmType algorithmType,
|
||||
ScoringCalculatorType scoringType) -> std::unique_ptr<ShardokAIClient> {
|
||||
// Use default MCTS configuration
|
||||
mcts::MCTSConfig mctsConfig{};
|
||||
|
||||
return std::make_unique<ShardokAIClient>(
|
||||
playerId,
|
||||
isDefender,
|
||||
hexMap,
|
||||
settings,
|
||||
algorithmType,
|
||||
scoringType,
|
||||
mctsConfig);
|
||||
}
|
||||
|
||||
} // namespace ai_testing_common
|
||||
} // namespace shardok
|
||||
@@ -0,0 +1,67 @@
|
||||
//
|
||||
// AIClientFactory - Unified factory for creating ShardokAIClient instances
|
||||
// Used by both AI Performance Runner and AI Battle Simulator
|
||||
//
|
||||
|
||||
#ifndef EAGLE0_AI_CLIENT_FACTORY_HPP
|
||||
#define EAGLE0_AI_CLIENT_FACTORY_HPP
|
||||
|
||||
#include <memory>
|
||||
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/AIConfig.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/library/ShardokCTypes.h"
|
||||
#include "src/main/cpp/net/eagle0/shardok/library/settings/GameSettings.hpp"
|
||||
|
||||
// Forward declarations for flatbuffer types
|
||||
namespace net {
|
||||
namespace eagle0 {
|
||||
namespace shardok {
|
||||
namespace storage {
|
||||
namespace fb {
|
||||
struct HexMap;
|
||||
}
|
||||
} // namespace storage
|
||||
} // namespace shardok
|
||||
} // namespace eagle0
|
||||
} // namespace net
|
||||
|
||||
namespace shardok {
|
||||
|
||||
// Forward declarations
|
||||
class ShardokAIClient;
|
||||
|
||||
namespace ai_testing_common {
|
||||
|
||||
/**
|
||||
* Factory class for creating ShardokAIClient instances with standard configuration.
|
||||
*
|
||||
* This centralizes the AI client creation logic that was duplicated
|
||||
* between AI Performance Runner and AI Battle Simulator.
|
||||
*/
|
||||
class AIClientFactory {
|
||||
public:
|
||||
/**
|
||||
* Create an AI client for a player.
|
||||
*
|
||||
* @param playerId The player ID for this AI
|
||||
* @param isDefender True if this AI is the defender
|
||||
* @param hexMap The game's hex map
|
||||
* @param settings The game settings to use
|
||||
* @param algorithmType The AI algorithm to use (MCTS or ITERATIVE_DEEPENING)
|
||||
* @param scoringType The scoring calculator type (defaults to STANDARD)
|
||||
* @return Unique pointer to created ShardokAIClient
|
||||
*/
|
||||
static auto Create(
|
||||
PlayerId playerId,
|
||||
bool isDefender,
|
||||
const net::eagle0::shardok::storage::fb::HexMap* hexMap,
|
||||
const SettingsGetter& settings,
|
||||
AIAlgorithmType algorithmType = AIAlgorithmType::ITERATIVE_DEEPENING,
|
||||
ScoringCalculatorType scoringType = ScoringCalculatorType::STANDARD)
|
||||
-> std::unique_ptr<ShardokAIClient>;
|
||||
};
|
||||
|
||||
} // namespace ai_testing_common
|
||||
} // namespace shardok
|
||||
|
||||
#endif // EAGLE0_AI_CLIENT_FACTORY_HPP
|
||||
@@ -0,0 +1,58 @@
|
||||
load("@rules_cc//cc:defs.bzl", "cc_library")
|
||||
|
||||
cc_library(
|
||||
name = "game_settings_factory",
|
||||
srcs = ["GameSettingsFactory.cpp"],
|
||||
hdrs = ["GameSettingsFactory.hpp"],
|
||||
visibility = ["//visibility:public"],
|
||||
deps = [
|
||||
"//src/main/cpp/net/eagle0/common:filesystem_utils",
|
||||
"//src/main/cpp/net/eagle0/common:tsv_parser",
|
||||
"//src/main/cpp/net/eagle0/shardok/library/settings:game_settings",
|
||||
"//src/main/cpp/net/eagle0/shardok/util:battalion_type_registrar",
|
||||
],
|
||||
)
|
||||
|
||||
cc_library(
|
||||
name = "ai_client_factory",
|
||||
srcs = ["AIClientFactory.cpp"],
|
||||
hdrs = ["AIClientFactory.hpp"],
|
||||
visibility = ["//visibility:public"],
|
||||
deps = [
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:shardok_ai_client",
|
||||
"//src/main/cpp/net/eagle0/shardok/library:shardok_c_types",
|
||||
"//src/main/cpp/net/eagle0/shardok/library/settings:game_settings",
|
||||
"//src/main/flatbuffer/net/eagle0/shardok/storage:hex_map_cc_fbs",
|
||||
],
|
||||
)
|
||||
|
||||
cc_library(
|
||||
name = "test_game_state_builder",
|
||||
srcs = ["TestGameStateBuilder.cpp"],
|
||||
hdrs = ["TestGameStateBuilder.hpp"],
|
||||
visibility = ["//visibility:public"],
|
||||
deps = [
|
||||
"//src/main/cpp/net/eagle0/common:filesystem_utils",
|
||||
"//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/fb_helpers:game_state_helpers",
|
||||
"//src/main/cpp/net/eagle0/shardok/library/settings:game_settings",
|
||||
"//src/main/cpp/net/eagle0/shardok/util:map_loader",
|
||||
"//src/main/flatbuffer/net/eagle0/shardok/storage:unit_fbs",
|
||||
"//src/main/protobuf/net/eagle0/shardok/common:player_info_proto",
|
||||
],
|
||||
)
|
||||
|
||||
cc_library(
|
||||
name = "game_phase_runner",
|
||||
srcs = ["GamePhaseRunner.cpp"],
|
||||
hdrs = ["GamePhaseRunner.hpp"],
|
||||
visibility = ["//visibility:public"],
|
||||
deps = [
|
||||
"//src/main/cpp/net/eagle0/shardok/ai:shardok_ai_client",
|
||||
"//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/flatbuffer/net/eagle0/shardok/storage:game_state_cc_fbs",
|
||||
],
|
||||
)
|
||||
@@ -0,0 +1,62 @@
|
||||
//
|
||||
// GamePhaseRunner Implementation
|
||||
//
|
||||
|
||||
#include "GamePhaseRunner.hpp"
|
||||
|
||||
#include "src/main/cpp/net/eagle0/shardok/ai/ShardokAIClient.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/library/ShardokEngine.hpp"
|
||||
#include "src/main/flatbuffer/net/eagle0/shardok/storage/game_state.hpp"
|
||||
|
||||
namespace shardok {
|
||||
namespace ai_testing_common {
|
||||
|
||||
auto GamePhaseRunner::RunSetupPhase(
|
||||
ShardokEngine& engine,
|
||||
const std::function<ShardokAIClient&(PlayerId)>& getAI) -> PhaseResult {
|
||||
PhaseResult result;
|
||||
|
||||
while (engine.GetCurrentGameState()->status()->state() ==
|
||||
net::eagle0::shardok::storage::fb::GameStatus_::State_SET_UP) {
|
||||
auto currentState = engine.GetCurrentGameState();
|
||||
PlayerId currentPlayer = currentState->current_player();
|
||||
|
||||
auto availableCommands = engine.GetAvailableCommandProtos(currentPlayer, false);
|
||||
|
||||
if (availableCommands.empty()) { break; }
|
||||
|
||||
// Get AI for current player and make decision
|
||||
ShardokAIClient& activeAI = getAI(currentPlayer);
|
||||
auto choiceResults = activeAI.ChooseCommandIndex(engine);
|
||||
|
||||
// Apply command
|
||||
engine.PostCommand(currentPlayer, choiceResults.chosenIndex);
|
||||
result.commandsExecuted++;
|
||||
|
||||
// Check if game ended unexpectedly
|
||||
if (engine.GameIsOver()) {
|
||||
result.gameEnded = true;
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
auto GamePhaseRunner::RunSingleTurn(ShardokEngine& engine, PlayerId playerId, ShardokAIClient& ai)
|
||||
-> bool {
|
||||
auto availableCommands = engine.GetAvailableCommandProtos(playerId, false);
|
||||
|
||||
if (availableCommands.empty()) { return false; }
|
||||
|
||||
// Get AI decision
|
||||
auto choiceResults = ai.ChooseCommandIndex(engine);
|
||||
|
||||
// Apply command
|
||||
engine.PostCommand(playerId, choiceResults.chosenIndex);
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
} // namespace ai_testing_common
|
||||
} // namespace shardok
|
||||
@@ -0,0 +1,64 @@
|
||||
//
|
||||
// GamePhaseRunner - Unified logic for running game phases with AI decision-making
|
||||
// Used by both AI Performance Runner and AI Battle Simulator
|
||||
//
|
||||
|
||||
#ifndef EAGLE0_GAME_PHASE_RUNNER_HPP
|
||||
#define EAGLE0_GAME_PHASE_RUNNER_HPP
|
||||
|
||||
#include <functional>
|
||||
|
||||
#include "src/main/cpp/net/eagle0/shardok/library/GameStateW.hpp"
|
||||
#include "src/main/cpp/net/eagle0/shardok/library/ShardokCTypes.h"
|
||||
|
||||
namespace shardok {
|
||||
|
||||
// Forward declarations
|
||||
class ShardokEngine;
|
||||
class ShardokAIClient;
|
||||
|
||||
namespace ai_testing_common {
|
||||
|
||||
/**
|
||||
* Result of running a game phase.
|
||||
*/
|
||||
struct PhaseResult {
|
||||
int commandsExecuted = 0;
|
||||
bool gameEnded = false;
|
||||
};
|
||||
|
||||
/**
|
||||
* Helper class for running setup and battle phases with AI decision-making.
|
||||
*
|
||||
* This centralizes the game phase execution logic that was duplicated
|
||||
* between AI Performance Runner and AI Battle Simulator.
|
||||
*/
|
||||
class GamePhaseRunner {
|
||||
public:
|
||||
/**
|
||||
* Run the setup phase where AIs place their units.
|
||||
*
|
||||
* @param engine The game engine
|
||||
* @param getAI Function to get the AI client for a given player ID
|
||||
* @return PhaseResult with number of commands executed and whether game ended
|
||||
*/
|
||||
static auto RunSetupPhase(
|
||||
ShardokEngine& engine,
|
||||
const std::function<ShardokAIClient&(PlayerId)>& getAI) -> PhaseResult;
|
||||
|
||||
/**
|
||||
* Run one turn of a game phase (either for AI testing or full battle).
|
||||
*
|
||||
* @param engine The game engine
|
||||
* @param playerId The player whose turn it is
|
||||
* @param ai The AI client to use for decision-making
|
||||
* @return True if the turn was executed successfully, false if no commands available
|
||||
*/
|
||||
static auto RunSingleTurn(ShardokEngine& engine, PlayerId playerId, ShardokAIClient& ai)
|
||||
-> bool;
|
||||
};
|
||||
|
||||
} // namespace ai_testing_common
|
||||
} // namespace shardok
|
||||
|
||||
#endif // EAGLE0_GAME_PHASE_RUNNER_HPP
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user