Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| b7f49071ba |
+59
@@ -0,0 +1,59 @@
|
|||||||
|
# Evo_Bot
|
||||||
|
|
||||||
|
Robocode Tank Royale bot with a modular gun system where evolved neural networks learn to predict enemy dodge behavior.
|
||||||
|
|
||||||
|
## Language
|
||||||
|
|
||||||
|
**Evo_Bot**:
|
||||||
|
The bot itself — handles movement and firing discipline. Guns are pluggable modules.
|
||||||
|
_Avoid_: robot, tank
|
||||||
|
|
||||||
|
**Guess Factor (GF)**:
|
||||||
|
A value from -1 to +1 representing where on the maximum escape angle arc the enemy is. 0 = directly ahead, -1 = full left dodge, +1 = full right dodge. The gun's prediction target.
|
||||||
|
_Avoid_: aim offset, dodge index
|
||||||
|
|
||||||
|
**Max Escape Angle (MEA)**:
|
||||||
|
The widest angle the enemy can reach before a bullet arrives, computed from distance and bullet speed. Guess factor is multiplied by MEA to get the aim offset.
|
||||||
|
|
||||||
|
**Lateral Velocity**:
|
||||||
|
Enemy speed projected perpendicular to the line between you and them. The primary signal for guess factor prediction.
|
||||||
|
_Avoid_: tangential speed, sideways velocity
|
||||||
|
|
||||||
|
**Sliding Window**:
|
||||||
|
The last N ticks (default 30) of enemy state fed as input to the network. Each tick contains lateral velocity, heading delta, and wall distance ahead.
|
||||||
|
_Avoid_: observation buffer, input history
|
||||||
|
|
||||||
|
**Replay Tape**:
|
||||||
|
Rolling buffer of recorded enemy states (~2000 ticks). The evolution thread evaluates gun fitness against this tape.
|
||||||
|
_Avoid_: experience buffer, replay buffer
|
||||||
|
|
||||||
|
**Virtual Gun**:
|
||||||
|
A gun that runs in parallel without actually firing. It tracks where it would have aimed and whether a simulated bullet would have hit. Used to compare gun variants.
|
||||||
|
|
||||||
|
**Virtual Bullet**:
|
||||||
|
A simulated bullet fired by a virtual gun. Never actually sent to the game engine.
|
||||||
|
|
||||||
|
**TOPO_Gun**:
|
||||||
|
Fixed-topology ANN gun evolved by GA. Network shape is predetermined (e.g., 91-8-1), only weights are evolved.
|
||||||
|
_Avoid_: static gun, fixed gun
|
||||||
|
|
||||||
|
**NEAT_Gun**:
|
||||||
|
Variable-topology ANN gun where evolution can add/remove neurons and connections (NEAT algorithm). Deferred — only built if TOPO_Gun hits a ceiling.
|
||||||
|
|
||||||
|
**Population**:
|
||||||
|
The set of candidate networks (default 64-200) being evolved. Each member is a complete set of ANN weights.
|
||||||
|
|
||||||
|
**Champion**:
|
||||||
|
The best-performing network in the current GA population. The champion's weights are pushed to the inference side when it beats the current best.
|
||||||
|
_Avoid_: best, winner, elite
|
||||||
|
|
||||||
|
**Fitness**:
|
||||||
|
Hit count when a network's aim predictions are evaluated as virtual bullets against sampled ticks from the replay tape.
|
||||||
|
_Avoid_: score, reward
|
||||||
|
|
||||||
|
**Cold Start**:
|
||||||
|
The first-ever battle with no saved weights. The gun does not fire until the evolution thread produces its first champion. Subsequent battles load persisted weights.
|
||||||
|
|
||||||
|
**Weight Persistence**:
|
||||||
|
Saving evolved weights to disk. Load order: per-opponent file, then global fallback, then random initialization.
|
||||||
|
_Avoid_: model saving, checkpointing
|
||||||
@@ -1,3 +0,0 @@
|
|||||||
nimble.develop
|
|
||||||
nimble.paths
|
|
||||||
nimbledeps
|
|
||||||
Binary file not shown.
@@ -1,11 +0,0 @@
|
|||||||
{
|
|
||||||
"name": "GotoTest",
|
|
||||||
"version": "0.1.0",
|
|
||||||
"authors": ["Davide Cappellini"],
|
|
||||||
"description": "Throwaway goto(x,y) controller prototype",
|
|
||||||
"homepage": "",
|
|
||||||
"countryCodes": ["IT"],
|
|
||||||
"gameTypes": ["classic", "melee", "1v1"],
|
|
||||||
"platform": "Nim",
|
|
||||||
"programmingLang": "Nim"
|
|
||||||
}
|
|
||||||
@@ -1,95 +0,0 @@
|
|||||||
## GotoTest — throwaway prototype to validate goto(x,y) + aimTo(x,y) controllers.
|
|
||||||
## Diamond pattern with wall-smash north waypoint to test stuck recovery.
|
|
||||||
## Tank Royale: Y=0 is south, Y increases northward. Arena 800x600.
|
|
||||||
|
|
||||||
import std/[strformat, os]
|
|
||||||
import tankroyale_botapi
|
|
||||||
|
|
||||||
const botJsonPath = currentSourcePath().parentDir / "GotoTest.json"
|
|
||||||
|
|
||||||
# Diamond waypoints. (400,600) is AT the north wall — physically impossible, tests wall recovery.
|
|
||||||
const waypoints = [
|
|
||||||
(400.0, 300.0), # center
|
|
||||||
(400.0, 600.0), # north wall — AT wall, unreachable, tests wall-stuck recovery
|
|
||||||
( 36.0, 300.0), # west wall (left)
|
|
||||||
(400.0, 36.0), # south wall (bottom)
|
|
||||||
(764.0, 300.0), # east wall (right)
|
|
||||||
(400.0, 300.0), # center
|
|
||||||
]
|
|
||||||
|
|
||||||
type GotoBot = ref object of Bot
|
|
||||||
waypointIdx: int
|
|
||||||
oldX, oldY: float # position 3 ticks ago for stuck detection
|
|
||||||
tickCount: int # ticks since oldX/oldY was last updated
|
|
||||||
stuckTicks: int # remaining ticks of reverse override
|
|
||||||
|
|
||||||
|
|
||||||
method onRoundStarted*(bot: GotoBot, e: RoundStartedEvent) =
|
|
||||||
setAdjustGunForBodyTurn(true)
|
|
||||||
bot.waypointIdx = 0
|
|
||||||
bot.oldX = 0.0; bot.oldY = 0.0
|
|
||||||
bot.tickCount = 0; bot.stuckTicks = 0
|
|
||||||
|
|
||||||
method run(bot: GotoBot) =
|
|
||||||
while isRunning():
|
|
||||||
let bx = getX()
|
|
||||||
let by = getY()
|
|
||||||
let (tx, ty) = waypoints[bot.waypointIdx]
|
|
||||||
let dist = distanceTo(bx, by, tx, ty)
|
|
||||||
|
|
||||||
if dist < 40.0:
|
|
||||||
bot.waypointIdx = (bot.waypointIdx + 1) mod waypoints.len
|
|
||||||
|
|
||||||
# goto controller — explicit forward/reverse proportional steering
|
|
||||||
# ponytail: no speed-taper on heading error; add when overshooting observed
|
|
||||||
let rawBearing = normalizeRelativeAngle(directionTo(tx, ty) - getDirection())
|
|
||||||
let (dirSign, effBearing) =
|
|
||||||
if abs(rawBearing) > 90.0:
|
|
||||||
(-1.0, normalizeRelativeAngle(rawBearing + 180.0))
|
|
||||||
else:
|
|
||||||
(1.0, rawBearing)
|
|
||||||
let maxTurn = calcMaxTurnRate(getSpeed())
|
|
||||||
setTurnRate(effBearing.clamp(-maxTurn, maxTurn))
|
|
||||||
|
|
||||||
# stuck detector — ponytail: 3-tick sample, upgrade to wall-nav if needed
|
|
||||||
inc bot.tickCount
|
|
||||||
if bot.tickCount >= 3:
|
|
||||||
let moved = distanceTo(bot.oldX, bot.oldY, bx, by)
|
|
||||||
if moved < 1.0 and dist >= 30.0:
|
|
||||||
bot.stuckTicks = 6
|
|
||||||
echo &"STUCK at ({bx:.1f},{by:.1f}) moved={moved:.2f} — reversing"
|
|
||||||
bot.oldX = bx; bot.oldY = by
|
|
||||||
bot.tickCount = 0
|
|
||||||
|
|
||||||
if bot.stuckTicks > 0:
|
|
||||||
# reverse current speed direction to unstick
|
|
||||||
let targetSpd = if getSpeed() >= 0.0: -8.0 else: 8.0
|
|
||||||
setTargetSpeed(targetSpd)
|
|
||||||
dec bot.stuckTicks
|
|
||||||
else:
|
|
||||||
let rawSpeed = getNewTargetSpeed(MAX_SPEED, abs(getSpeed()), dist)
|
|
||||||
setTargetSpeed(dirSign * rawSpeed)
|
|
||||||
|
|
||||||
# aimTo arena center — radarBearingTo reuses getX/getY/getGunDirection internally
|
|
||||||
let cx = getArenaWidth().float / 2.0
|
|
||||||
let cy = getArenaHeight().float / 2.0
|
|
||||||
let gunTurnNeeded = normalizeRelativeAngle(directionTo(cx, cy) - getGunDirection())
|
|
||||||
setGunTurnRate(gunTurnNeeded.clamp(-MAX_GUN_TURN_RATE, MAX_GUN_TURN_RATE))
|
|
||||||
|
|
||||||
echo &"pos=({bx:.1f},{by:.1f}) target=({tx:.1f},{ty:.1f}) dist={dist:.1f} spd={getSpeed():.2f} stuck={bot.stuckTicks}"
|
|
||||||
|
|
||||||
# Debug graphics: waypoint circles + line to current target
|
|
||||||
for i, (wx, wy) in waypoints:
|
|
||||||
if i == bot.waypointIdx:
|
|
||||||
setStrokeColor(RED)
|
|
||||||
else:
|
|
||||||
setStrokeColor(GRAY)
|
|
||||||
drawCircle(wx, wy, 10.0)
|
|
||||||
setStrokeColor(YELLOW)
|
|
||||||
drawLine(bx, by, tx, ty)
|
|
||||||
|
|
||||||
go()
|
|
||||||
|
|
||||||
when isMainModule:
|
|
||||||
var bot = GotoBot()
|
|
||||||
start(bot, botJsonPath)
|
|
||||||
@@ -1,10 +0,0 @@
|
|||||||
# Package
|
|
||||||
version = "0.1.0"
|
|
||||||
author = "Davide Cappellini"
|
|
||||||
description = "GotoTest — throwaway goto(x,y) controller prototype"
|
|
||||||
license = "MIT"
|
|
||||||
bin = @["GotoTest"]
|
|
||||||
|
|
||||||
# Dependencies
|
|
||||||
requires "nim >= 2.0.0"
|
|
||||||
requires "tankroyale_botapi >= 1.0.0"
|
|
||||||
@@ -1,4 +0,0 @@
|
|||||||
# begin Nimble config (version 2)
|
|
||||||
when withDir(thisDir(), system.fileExists("nimble.paths")):
|
|
||||||
include "nimble.paths"
|
|
||||||
# end Nimble config
|
|
||||||
@@ -1,11 +1,11 @@
|
|||||||
{
|
{
|
||||||
"name": "PPO_Bot",
|
"name": "OscillatorBot",
|
||||||
"version": "0.1.0",
|
"version": "0.1.0",
|
||||||
"authors": ["Davide Cappellini"],
|
"authors": ["Davide Cappellini"],
|
||||||
"description": "PPO-trained RL bot",
|
"description": "Predictable zigzag sparring partner for gun testing",
|
||||||
"homepage": "",
|
"homepage": "",
|
||||||
"countryCodes": ["IT"],
|
"countryCodes": ["IT"],
|
||||||
"gameTypes": ["classic", "melee", "1v1"],
|
"gameTypes": ["classic", "1v1"],
|
||||||
"platform": "Nim",
|
"platform": "Nim",
|
||||||
"programmingLang": "Nim"
|
"programmingLang": "Nim"
|
||||||
}
|
}
|
||||||
@@ -0,0 +1,48 @@
|
|||||||
|
# OscillatorBot — predictable zigzag sparring partner for GA gun testing.
|
||||||
|
# Reverses direction + turn every PERIOD ticks. Head-on targeting only.
|
||||||
|
|
||||||
|
import std/[math, os]
|
||||||
|
import tankroyale_botapi
|
||||||
|
|
||||||
|
const botJsonPath = currentSourcePath().parentDir / "OscillatorBot.json"
|
||||||
|
|
||||||
|
const
|
||||||
|
SPEED = 7.0 # forward/backward speed
|
||||||
|
PERIOD = 25 # ticks between direction reversals
|
||||||
|
|
||||||
|
type OscillatorBot = ref object of Bot
|
||||||
|
tickCount: int
|
||||||
|
moveSign: float # +1 forward, -1 backward
|
||||||
|
turnSign: float # +1 right, -1 left
|
||||||
|
|
||||||
|
method onRoundStarted*(bot: OscillatorBot, e: RoundStartedEvent) =
|
||||||
|
setAdjustGunForBodyTurn(true)
|
||||||
|
setAdjustRadarForBodyTurn(true)
|
||||||
|
setAdjustRadarForGunTurn(true)
|
||||||
|
bot.tickCount = 0
|
||||||
|
bot.moveSign = 1.0
|
||||||
|
bot.turnSign = 1.0
|
||||||
|
|
||||||
|
method onScannedBot*(bot: OscillatorBot, e: ScannedBotEvent) =
|
||||||
|
# Head-on targeting: aim gun directly at enemy, fire medium power
|
||||||
|
let bearing = directionTo(getX(), getY(), e.x, e.y)
|
||||||
|
let gunDelta = normalizeRelativeAngle(bearing - getGunDirection())
|
||||||
|
setGunTurnRate(gunDelta.clamp(-MAX_GUN_TURN_RATE, MAX_GUN_TURN_RATE))
|
||||||
|
if abs(gunDelta) < 10.0 and getGunHeat() <= 0.0:
|
||||||
|
discard setFire(2.0)
|
||||||
|
|
||||||
|
method run*(bot: OscillatorBot) =
|
||||||
|
while isRunning():
|
||||||
|
inc bot.tickCount
|
||||||
|
if bot.tickCount mod PERIOD == 0:
|
||||||
|
bot.moveSign *= -1.0
|
||||||
|
bot.turnSign *= -1.0
|
||||||
|
|
||||||
|
setTargetSpeed(bot.moveSign * SPEED)
|
||||||
|
setTurnRate(bot.turnSign * 4.0)
|
||||||
|
setRadarTurnRate(45.0) # spin radar to keep scanning
|
||||||
|
go()
|
||||||
|
|
||||||
|
when isMainModule:
|
||||||
|
var bot = OscillatorBot(moveSign: 1.0, turnSign: 1.0)
|
||||||
|
start(bot, botJsonPath)
|
||||||
Executable
+3
@@ -0,0 +1,3 @@
|
|||||||
|
#!/bin/sh
|
||||||
|
cd "$(dirname "$0")"
|
||||||
|
exec ./OscillatorBot 2>> /tmp/oscillatorbot_stderr.log
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
--path:"../libs"
|
||||||
@@ -1,3 +0,0 @@
|
|||||||
nimble.develop
|
|
||||||
nimble.paths
|
|
||||||
nimbledeps
|
|
||||||
Binary file not shown.
@@ -1,370 +0,0 @@
|
|||||||
## PPO_Bot — enemy tracker + state vector wired into the game loop.
|
|
||||||
## Training: trajectory collected per tick, PPO update in background thread.
|
|
||||||
|
|
||||||
import std/[os, strformat, strutils, math, times, algorithm]
|
|
||||||
import arraymancer
|
|
||||||
import tankroyale_botapi
|
|
||||||
import network
|
|
||||||
import actions
|
|
||||||
import training
|
|
||||||
import weights
|
|
||||||
import ./enemy_tracker
|
|
||||||
import ./state_vector
|
|
||||||
|
|
||||||
# ── Hyperparameters from env vars (PPOB_ prefix) ─────────────────────────────
|
|
||||||
# All optional; defaults match ppoUpdate signature in training.nim.
|
|
||||||
|
|
||||||
proc getEnvFloat(name: string, default: float32): float32 =
|
|
||||||
let v = getEnv(name)
|
|
||||||
if v.len == 0: default else: parseFloat(v).float32
|
|
||||||
|
|
||||||
proc getEnvInt(name: string, default: int): int =
|
|
||||||
let v = getEnv(name)
|
|
||||||
if v.len == 0: default else: parseInt(v)
|
|
||||||
|
|
||||||
var
|
|
||||||
hpLr: float32 = getEnvFloat("PPOB_LR", 3e-4'f32)
|
|
||||||
hpClipEpsilon: float32 = getEnvFloat("PPOB_CLIP_EPSILON", 0.2'f32)
|
|
||||||
hpEntropyCoeff: float32 = getEnvFloat("PPOB_ENTROPY_COEFF", 0.01'f32)
|
|
||||||
hpValueLossCoeff: float32 = getEnvFloat("PPOB_VALUE_LOSS_COEFF", 0.5'f32)
|
|
||||||
hpMaxGradNorm: float32 = getEnvFloat("PPOB_MAX_GRAD_NORM", 0.5'f32)
|
|
||||||
hpGamma: float32 = getEnvFloat("PPOB_GAMMA", 0.99'f32)
|
|
||||||
hpLam: float32 = getEnvFloat("PPOB_LAM", 0.95'f32)
|
|
||||||
hpEpochs: int = getEnvInt("PPOB_EPOCHS", 4)
|
|
||||||
hpMiniBatchSize: int = getEnvInt("PPOB_MINI_BATCH_SIZE", 64)
|
|
||||||
hpUpdateInterval: int = getEnvInt("PPOB_UPDATE_INTERVAL", 10)
|
|
||||||
|
|
||||||
# Wire logStd tunable params into network module vars (read before initActorCritic)
|
|
||||||
logStdFloor = getEnvFloat("PPOB_LOG_STD_FLOOR", -3.0'f32)
|
|
||||||
logStdCeiling = getEnvFloat("PPOB_LOG_STD_CEILING", 0.5'f32)
|
|
||||||
initialLogStd = getEnvFloat("PPOB_INITIAL_LOG_STD", 0.0'f32)
|
|
||||||
|
|
||||||
# ── Structured log output ─────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
let logFile = getEnv("PPOB_LOG_FILE") # empty → no JSON logging
|
|
||||||
let evalOnly = getEnv("PPOB_EVAL_ONLY") == "1" # freeze training (pure evaluation)
|
|
||||||
|
|
||||||
proc appendJsonLine(path, line: string) =
|
|
||||||
## Append a JSON line to path; no-op if path is empty.
|
|
||||||
if path.len == 0: return
|
|
||||||
let f = open(path, fmAppend)
|
|
||||||
f.writeLine(line)
|
|
||||||
f.close()
|
|
||||||
|
|
||||||
proc jsonFloat(v: float32): string =
|
|
||||||
## Serialize a float for JSON; non-finite → null (keeps JSONL parseable).
|
|
||||||
if v == v and v > -1e30'f32 and v < 1e30'f32: $v else: "null"
|
|
||||||
|
|
||||||
proc hyperparmSnapshot(): string =
|
|
||||||
## Compact JSON object of current hyperparams (no outer braces).
|
|
||||||
&"\"lr\":{hpLr},\"clipEpsilon\":{hpClipEpsilon}," &
|
|
||||||
&"\"entropyCoeff\":{hpEntropyCoeff},\"valueLossCoeff\":{hpValueLossCoeff}," &
|
|
||||||
&"\"maxGradNorm\":{hpMaxGradNorm},\"gamma\":{hpGamma},\"lam\":{hpLam}," &
|
|
||||||
&"\"epochs\":{hpEpochs},\"miniBatchSize\":{hpMiniBatchSize}," &
|
|
||||||
&"\"logStdFloor\":{logStdFloor},\"logStdCeiling\":{logStdCeiling},\"initialLogStd\":{initialLogStd}"
|
|
||||||
|
|
||||||
const botJsonPath = currentSourcePath().parentDir / "PPO_Bot.json"
|
|
||||||
const weightsRoot = currentSourcePath().parentDir / "weights"
|
|
||||||
|
|
||||||
type
|
|
||||||
InFlightBullet = object
|
|
||||||
x, y: float64 # current position
|
|
||||||
vx, vy: float64 # velocity (pixels/tick)
|
|
||||||
power: float64
|
|
||||||
|
|
||||||
type PPOBot = ref object of Bot
|
|
||||||
tracker: EnemyTracker
|
|
||||||
buffer: TrajectoryBuffer
|
|
||||||
prevEnergy: float32 # own energy last tick
|
|
||||||
prevEnemyE: float32 # enemy energy last tick (from tracker)
|
|
||||||
lastState: array[STATE_DIM, float32] # plain arrays — tensors NEVER cross threads
|
|
||||||
lastAction: array[ACTION_DIM, float32]
|
|
||||||
lastLogP: float32
|
|
||||||
lastValue: float32
|
|
||||||
hasLastTrans: bool
|
|
||||||
lastActions: BotActions # previous tick's decoded actions (for state vector)
|
|
||||||
roundRewardSum: float32 # cumulative reward this round (for live display)
|
|
||||||
roundTicks: int # ticks this round
|
|
||||||
# Fixed-size bullet buffer — NO heap on the shared bot object. A seq here is
|
|
||||||
# allocated by the per-round bot thread and freed by the next round's thread
|
|
||||||
# (bot.bullets = @[] on round start) → foreign-heap free under --threads:on +
|
|
||||||
# ORC → SIGSEGV. N=4: the state vector only consumes the closest 3 slots.
|
|
||||||
bullets: array[4, InFlightBullet]
|
|
||||||
bulletCount: int
|
|
||||||
|
|
||||||
var ac = initActorCritic()
|
|
||||||
var gAdamStates: ACAdamStates # persists across rounds
|
|
||||||
var roundCounter = 0
|
|
||||||
var roundsSinceUpdate = 0
|
|
||||||
|
|
||||||
# ── Bot methods ───────────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
method onScannedBot*(bot: PPOBot, e: ScannedBotEvent) =
|
|
||||||
bot.tracker.update(e.x, e.y, e.direction, e.speed, e.energy)
|
|
||||||
|
|
||||||
method onRoundStarted*(bot: PPOBot, e: RoundStartedEvent) =
|
|
||||||
setAdjustGunForBodyTurn(true)
|
|
||||||
setAdjustRadarForBodyTurn(true)
|
|
||||||
setAdjustRadarForGunTurn(true)
|
|
||||||
bot.tracker = initEnemyTracker()
|
|
||||||
# Do NOT clear bot.buffer here — transitions accumulate across rounds
|
|
||||||
# until hpUpdateInterval rounds have passed (cleared in onRoundEnded).
|
|
||||||
bot.prevEnergy = 0.0'f32
|
|
||||||
bot.prevEnemyE = 0.0'f32
|
|
||||||
bot.hasLastTrans = false
|
|
||||||
bot.roundRewardSum = 0.0'f32
|
|
||||||
bot.roundTicks = 0
|
|
||||||
bot.bulletCount = 0
|
|
||||||
|
|
||||||
method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
|
|
||||||
inc roundCounter
|
|
||||||
inc roundsSinceUpdate
|
|
||||||
debugLog("[PO-ENTER] round=" & $roundCounter & " tid=" & $getThreadId())
|
|
||||||
|
|
||||||
# Add round-end score bonus to last transition and mark it as episode boundary
|
|
||||||
let roundReward = computeRoundReward(e.results.totalScore.float32)
|
|
||||||
if bot.hasLastTrans and bot.buffer.len > 0:
|
|
||||||
let lastIdx = bot.buffer.len - 1
|
|
||||||
bot.buffer.transitions[lastIdx].reward += roundReward
|
|
||||||
bot.buffer.transitions[lastIdx].done = true
|
|
||||||
|
|
||||||
# Training progress display — one line per round in the UI console
|
|
||||||
let ticks = bot.buffer.len
|
|
||||||
var avgR = 0.0'f32
|
|
||||||
if ticks > 0:
|
|
||||||
var rewardSum = 0.0'f32
|
|
||||||
for i in 0 ..< bot.buffer.len: rewardSum += bot.buffer.transitions[i].reward
|
|
||||||
avgR = rewardSum / ticks.float32
|
|
||||||
let avgRStr = formatFloat(avgR.float, ffDecimal, 3)
|
|
||||||
printToStdOut(&"R:{roundCounter} ticks:{ticks} buf:{bot.buffer.len} sinceUpd:{roundsSinceUpdate}/{hpUpdateInterval} avgR:{avgRStr} score:{e.results.totalScore}\n")
|
|
||||||
echo &"R:{roundCounter} ticks:{ticks} buf:{bot.buffer.len} sinceUpd:{roundsSinceUpdate}/{hpUpdateInterval} avgR:{avgRStr} score:{e.results.totalScore}"
|
|
||||||
|
|
||||||
if bot.buffer.len == 0:
|
|
||||||
bot.hasLastTrans = false
|
|
||||||
return
|
|
||||||
|
|
||||||
# Emit per-round game-stats JSON line
|
|
||||||
let ts = int(epochTime())
|
|
||||||
let jline = &"""{{"type":"round","round":{roundCounter},"ticks":{ticks},"bufLen":{bot.buffer.len},"avgReward":{avgR},"score":{e.results.totalScore},"ts":{ts}}}"""
|
|
||||||
appendJsonLine(logFile, jline)
|
|
||||||
|
|
||||||
# PPOB_EVAL_ONLY=1 → freeze training (pure evaluation): skip ppoUpdate and
|
|
||||||
# checkpoint save, but keep advancing/writing round_counter.txt so run.sh's
|
|
||||||
# remaining-rounds bookkeeping still works, and keep the game line above.
|
|
||||||
if evalOnly:
|
|
||||||
writeFile(weightsRoot / "round_counter.txt", $roundCounter)
|
|
||||||
bot.buffer.clear()
|
|
||||||
bot.hasLastTrans = false
|
|
||||||
roundsSinceUpdate = 0
|
|
||||||
return
|
|
||||||
|
|
||||||
# Always save weights every round so round_counter.txt stays current.
|
|
||||||
saveCheckpoint(ac, gAdamStates, weightsRoot, roundCounter)
|
|
||||||
|
|
||||||
# Only update policy every hpUpdateInterval rounds (~3000 transitions).
|
|
||||||
if roundsSinceUpdate >= hpUpdateInterval:
|
|
||||||
# ponytail: synchronous update. Arraymancer tensors can't cross threads under
|
|
||||||
# ORC — training-thread ppoUpdate frees/rebinds tensors owned by the bot
|
|
||||||
# thread's heap (SIGSEGV, reproduced with a lone trainer thread on a fixed
|
|
||||||
# buffer; save/channel/forward exonerated). The bot API runs events on one
|
|
||||||
# bot thread, so inline is single-threaded. Revert to a background thread
|
|
||||||
# only if tensors are rebuilt from plain data on that thread.
|
|
||||||
printToStdOut(&" train→ R:{roundCounter} bufLen:{bot.buffer.len}\n")
|
|
||||||
echo &" train→ R:{roundCounter} bufLen:{bot.buffer.len}"
|
|
||||||
let m = ppoUpdate(ac, bot.buffer,
|
|
||||||
lastValue = 0.0'f32,
|
|
||||||
adamStates = gAdamStates,
|
|
||||||
epochs = hpEpochs,
|
|
||||||
miniBatchSize = hpMiniBatchSize,
|
|
||||||
clipEpsilon = hpClipEpsilon,
|
|
||||||
entropyCoeff = hpEntropyCoeff,
|
|
||||||
valueLossCoeff = hpValueLossCoeff,
|
|
||||||
lr = hpLr,
|
|
||||||
maxGradNorm = hpMaxGradNorm,
|
|
||||||
gamma = hpGamma,
|
|
||||||
lam = hpLam)
|
|
||||||
printToStdOut(&" trained R:{roundCounter} aLoss:{formatFloat(m.actorLoss.float, ffDecimal, 4)} vLoss:{formatFloat(m.valueLoss.float, ffDecimal, 4)} gNorm:{formatFloat(m.gradNorm.float, ffDecimal, 3)}\n")
|
|
||||||
echo &" trained R:{roundCounter} aLoss:{formatFloat(m.actorLoss.float, ffDecimal, 4)} vLoss:{formatFloat(m.valueLoss.float, ffDecimal, 4)} gNorm:{formatFloat(m.gradNorm.float, ffDecimal, 3)}"
|
|
||||||
# Emit training-health JSON line
|
|
||||||
let hp = hyperparmSnapshot()
|
|
||||||
let ts2 = int(epochTime())
|
|
||||||
let jline2 = &"""{{"type":"train","round":{roundCounter},"actorLoss":{jsonFloat(m.actorLoss)},"valueLoss":{jsonFloat(m.valueLoss)},"gradNorm":{jsonFloat(m.gradNorm)},"ts":{ts2},{hp}}}"""
|
|
||||||
appendJsonLine(logFile, jline2)
|
|
||||||
bot.buffer.clear()
|
|
||||||
roundsSinceUpdate = 0
|
|
||||||
|
|
||||||
bot.hasLastTrans = false
|
|
||||||
debugLog("[PO-EXIT] round=" & $roundCounter & " tid=" & $getThreadId())
|
|
||||||
|
|
||||||
method run(bot: PPOBot) =
|
|
||||||
debugLog("[RUN-ENTER] tid=" & $getThreadId())
|
|
||||||
# Seed energy and goto/aimTo targets on first tick (remainingDistance = 0 initially)
|
|
||||||
bot.prevEnergy = getEnergy().float32
|
|
||||||
bot.prevEnemyE = if bot.tracker.hasContact: bot.tracker.current.energy.float32 else: 0.0'f32
|
|
||||||
bot.lastActions.gotoX = getX()
|
|
||||||
bot.lastActions.gotoY = getY()
|
|
||||||
bot.lastActions.aimToX = getX()
|
|
||||||
bot.lastActions.aimToY = getY()
|
|
||||||
|
|
||||||
while isRunning():
|
|
||||||
bot.tracker.deadReckon()
|
|
||||||
|
|
||||||
# Spawn bullet when enemy fired last scan tick. Edge-triggered: the flag is
|
|
||||||
# consumed later (after the state build) so a radar gap (deadReckon ticks)
|
|
||||||
# can't spawn k phantom bullets from one shot — exactly one bullet, and
|
|
||||||
# state index 12 still pulses 1 on the detection tick.
|
|
||||||
if bot.tracker.hasContact and bot.tracker.current.hasFired:
|
|
||||||
if bot.bulletCount < 4:
|
|
||||||
let power = bot.tracker.current.lastFirePower
|
|
||||||
let speed = 20.0 - 3.0 * power
|
|
||||||
# Approximate gun direction: bearing from enemy toward our position
|
|
||||||
let myX = getX(); let myY = getY()
|
|
||||||
let ang = arctan2(myY - bot.tracker.current.y, myX - bot.tracker.current.x)
|
|
||||||
bot.bullets[bot.bulletCount] = InFlightBullet(
|
|
||||||
x: bot.tracker.current.x,
|
|
||||||
y: bot.tracker.current.y,
|
|
||||||
vx: speed * cos(ang),
|
|
||||||
vy: speed * sin(ang),
|
|
||||||
power: power,
|
|
||||||
)
|
|
||||||
inc bot.bulletCount
|
|
||||||
|
|
||||||
# Advance in-flight bullets and prune those off-arena (in-place compaction
|
|
||||||
# into the fixed buffer — no per-tick heap churn).
|
|
||||||
let aW = float64(getArenaWidth()); let aH = float64(getArenaHeight())
|
|
||||||
var n = 0
|
|
||||||
for i in 0 ..< bot.bulletCount:
|
|
||||||
let b = bot.bullets[i]
|
|
||||||
let nx = b.x + b.vx; let ny = b.y + b.vy
|
|
||||||
if nx >= 0.0 and nx <= aW and ny >= 0.0 and ny <= aH:
|
|
||||||
bot.bullets[n] = InFlightBullet(x: nx, y: ny, vx: b.vx, vy: b.vy, power: b.power)
|
|
||||||
inc n
|
|
||||||
bot.bulletCount = n
|
|
||||||
|
|
||||||
# Convert to BulletData for state vector (closest 3 by distance to us)
|
|
||||||
let myX2 = getX(); let myY2 = getY()
|
|
||||||
var bulletData: array[4, BulletData]
|
|
||||||
var bulletDataCount = 0
|
|
||||||
for i in 0 ..< bot.bulletCount:
|
|
||||||
let b = bot.bullets[i]
|
|
||||||
bulletData[bulletDataCount] = BulletData(x: b.x, y: b.y, power: b.power)
|
|
||||||
inc bulletDataCount
|
|
||||||
# sort ascending by distance so the nearest threats fill slots 0-2
|
|
||||||
if bulletDataCount > 1:
|
|
||||||
bulletData.toOpenArray(0, bulletDataCount - 1).sort(proc(a, b: BulletData): int =
|
|
||||||
let da = hypot(a.x - myX2, a.y - myY2)
|
|
||||||
let db = hypot(b.x - myX2, b.y - myY2)
|
|
||||||
cmp(da, db))
|
|
||||||
|
|
||||||
setRadarTurnRate(bot.tracker.getRadarTurnRate(getX(), getY(), getDirection(), getRadarDirection()))
|
|
||||||
|
|
||||||
let botData = BotStateData(
|
|
||||||
x: getX(),
|
|
||||||
y: getY(),
|
|
||||||
direction: getDirection(),
|
|
||||||
speed: getSpeed(),
|
|
||||||
energy: getEnergy(),
|
|
||||||
gunDirection: getGunDirection(),
|
|
||||||
gunHeat: getGunHeat(),
|
|
||||||
arenaWidth: float64(getArenaWidth()),
|
|
||||||
arenaHeight: float64(getArenaHeight()),
|
|
||||||
)
|
|
||||||
|
|
||||||
let remainingGotoDistance = hypot(bot.lastActions.gotoX - botData.x,
|
|
||||||
bot.lastActions.gotoY - botData.y)
|
|
||||||
let remainingGunAngle = abs(normalizeRelativeAngle(
|
|
||||||
directionTo(botData.x, botData.y, bot.lastActions.aimToX, bot.lastActions.aimToY) -
|
|
||||||
botData.gunDirection))
|
|
||||||
let state = buildStateVector(botData, bot.tracker, remainingGotoDistance, remainingGunAngle, bulletData, bulletDataCount)
|
|
||||||
# Consume the fired flag AFTER the state build: state index 12 saw the
|
|
||||||
# detection-tick pulse, and the next iteration's spawn check sees false —
|
|
||||||
# one shot → exactly one bullet, even across deadReckon gaps.
|
|
||||||
bot.tracker.current.hasFired = false
|
|
||||||
let (rawActs, logP) = ac.actorForward(state)
|
|
||||||
let value = ac.criticForward(state)
|
|
||||||
let ex = if bot.tracker.hasContact: bot.tracker.current.x else: botData.arenaWidth / 2.0
|
|
||||||
let ey = if bot.tracker.hasContact: bot.tracker.current.y else: botData.arenaHeight / 2.0
|
|
||||||
let acts = mapActions(rawActs,
|
|
||||||
getGunHeat().float,
|
|
||||||
botData.arenaWidth, botData.arenaHeight,
|
|
||||||
botData.x, botData.y,
|
|
||||||
botData.direction, botData.speed, botData.gunDirection,
|
|
||||||
ex, ey)
|
|
||||||
|
|
||||||
# Compute tick reward from energy deltas + dense shaping
|
|
||||||
let curEnergy = getEnergy().float32
|
|
||||||
let curEnemyE = if bot.tracker.hasContact: bot.tracker.current.energy.float32 else: bot.prevEnemyE
|
|
||||||
let myDelta = curEnergy - bot.prevEnergy
|
|
||||||
let enemyDelta = curEnemyE - bot.prevEnemyE
|
|
||||||
let arenaDiag = float32(sqrt(botData.arenaWidth * botData.arenaWidth +
|
|
||||||
botData.arenaHeight * botData.arenaHeight))
|
|
||||||
let distEnemy = if bot.tracker.hasContact:
|
|
||||||
float32(hypot(bot.tracker.current.x - botData.x,
|
|
||||||
bot.tracker.current.y - botData.y))
|
|
||||||
else: arenaDiag
|
|
||||||
let gunToEnemy = if bot.tracker.hasContact:
|
|
||||||
abs(normalizeRelativeAngle(
|
|
||||||
arctan2(bot.tracker.current.y - botData.y,
|
|
||||||
bot.tracker.current.x - botData.x) * 180.0 / PI -
|
|
||||||
botData.gunDirection)).float32
|
|
||||||
else: 180.0'f32
|
|
||||||
let tickReward = computeTickReward(myDelta, enemyDelta,
|
|
||||||
distToEnemy = distEnemy,
|
|
||||||
maxDist = arenaDiag,
|
|
||||||
gunBearingAbs = gunToEnemy)
|
|
||||||
|
|
||||||
# Track running reward for in-game display
|
|
||||||
bot.roundRewardSum += tickReward
|
|
||||||
inc bot.roundTicks
|
|
||||||
|
|
||||||
# Finalise previous transition with the reward from this tick's state change
|
|
||||||
if bot.hasLastTrans:
|
|
||||||
let tr = Transition(
|
|
||||||
state: bot.lastState,
|
|
||||||
action: bot.lastAction,
|
|
||||||
logProb: bot.lastLogP,
|
|
||||||
reward: tickReward,
|
|
||||||
value: bot.lastValue,
|
|
||||||
done: false, # episode boundary set in onRoundEnded
|
|
||||||
)
|
|
||||||
bot.buffer.add(tr)
|
|
||||||
|
|
||||||
# Store current for next tick — plain arrays only. `state`/`rawActs` tensors
|
|
||||||
# live and die on this thread; a NEW bot thread runs each round, so storing
|
|
||||||
# tensors in the shared bot object would free round-N's heap memory from
|
|
||||||
# round N+1's thread (SIGSEGV; confirmed empirically).
|
|
||||||
bot.lastState = stateToArr(state)
|
|
||||||
bot.lastAction = actionToArr(rawActs)
|
|
||||||
bot.lastLogP = logP
|
|
||||||
bot.lastValue = value
|
|
||||||
bot.prevEnergy = curEnergy
|
|
||||||
bot.prevEnemyE = curEnemyE
|
|
||||||
bot.hasLastTrans = true
|
|
||||||
bot.lastActions = acts
|
|
||||||
|
|
||||||
setTargetSpeed(acts.targetSpeed.float)
|
|
||||||
setTurnRate(acts.turnRate.float)
|
|
||||||
setGunTurnRate(acts.gunTurnRate.float)
|
|
||||||
if acts.shouldFire:
|
|
||||||
discard setFire(acts.firePower.float)
|
|
||||||
|
|
||||||
# In-game training progress overlay
|
|
||||||
let avgR = if bot.roundTicks > 0: bot.roundRewardSum / bot.roundTicks.float32
|
|
||||||
else: 0.0'f32
|
|
||||||
let avgRStr = formatFloat(avgR.float, ffDecimal, 2)
|
|
||||||
drawText(&"R:{roundCounter} avg:{avgRStr}", getX(), getY() - 40.0)
|
|
||||||
|
|
||||||
go()
|
|
||||||
|
|
||||||
when isMainModule:
|
|
||||||
createDir(weightsRoot)
|
|
||||||
cleanStaleTempDirs(weightsRoot)
|
|
||||||
let loadResult = loadBestAvailable(ac, gAdamStates, weightsRoot)
|
|
||||||
if loadResult.loaded:
|
|
||||||
roundCounter = loadResult.roundNum
|
|
||||||
|
|
||||||
var bot = PPOBot(
|
|
||||||
tracker: initEnemyTracker(),
|
|
||||||
buffer: initTrajectoryBuffer(),
|
|
||||||
)
|
|
||||||
start(bot, botJsonPath)
|
|
||||||
@@ -1,12 +0,0 @@
|
|||||||
# Package
|
|
||||||
version = "0.1.0"
|
|
||||||
author = "Davide Cappellini"
|
|
||||||
description = "PPO-trained Tank Royale bot"
|
|
||||||
license = "MIT"
|
|
||||||
bin = @["PPO_Bot"]
|
|
||||||
|
|
||||||
# Dependencies
|
|
||||||
requires "nim >= 2.0.0"
|
|
||||||
# tankroyale_botapi is vendored in-tree (libs/tankroyale_botapi) and wired via
|
|
||||||
# config.nims --path; no nimble dependency so builds never touch ~/.nimble/pkgs2.
|
|
||||||
requires "arraymancer >= 0.7.0"
|
|
||||||
@@ -1,7 +0,0 @@
|
|||||||
#!/bin/sh
|
|
||||||
# PPO_Bot — PPO-trained RL bot (compiled native binary)
|
|
||||||
# OPENBLAS_NUM_THREADS=1: prevent OpenBLAS from spawning worker threads,
|
|
||||||
# which deadlock when called from within a multi-threaded Nim bot process.
|
|
||||||
export OPENBLAS_NUM_THREADS=1
|
|
||||||
cd -- "$(dirname -- "$0")"
|
|
||||||
exec "./PPO_Bot"
|
|
||||||
@@ -1,49 +0,0 @@
|
|||||||
## actions.nim — map raw network output to Tank Royale bot commands.
|
|
||||||
|
|
||||||
import arraymancer
|
|
||||||
import std/math
|
|
||||||
import ./controllers
|
|
||||||
|
|
||||||
func sigmoid(x: float): float = 1.0 / (1.0 + exp(-x))
|
|
||||||
|
|
||||||
type
|
|
||||||
BotActions* = object
|
|
||||||
targetSpeed*: float
|
|
||||||
turnRate*: float
|
|
||||||
gunTurnRate*: float
|
|
||||||
shouldFire*: bool
|
|
||||||
firePower*: float
|
|
||||||
gotoX*: float
|
|
||||||
gotoY*: float
|
|
||||||
aimToX*: float
|
|
||||||
aimToY*: float
|
|
||||||
|
|
||||||
proc mapActions*(rawActions: Tensor[float32],
|
|
||||||
gunHeat: float,
|
|
||||||
arenaWidth, arenaHeight: float,
|
|
||||||
botX, botY, direction, speed, gunDirection: float,
|
|
||||||
enemyX, enemyY: float): BotActions =
|
|
||||||
## rawActions: [6] tensor from actorForward.
|
|
||||||
## Dims 0–1: goto x/y offset from enemy, 2–3: aimTo x/y offset from enemy,
|
|
||||||
## 4: fire decision, 5: fire power.
|
|
||||||
## Enemy-centred mapping: tanh gives [-1,1]; scale by arena/4 (goto) and
|
|
||||||
## arena/8 (aimTo) so zero-init defaults the bot toward the enemy.
|
|
||||||
let gotoX = clamp(enemyX + tanh(rawActions[0].float) * arenaWidth * 0.25, 0.0, arenaWidth)
|
|
||||||
let gotoY = clamp(enemyY + tanh(rawActions[1].float) * arenaHeight * 0.25, 0.0, arenaHeight)
|
|
||||||
let aimToX = clamp(enemyX + tanh(rawActions[2].float) * arenaWidth * 0.125, 0.0, arenaWidth)
|
|
||||||
let aimToY = clamp(enemyY + tanh(rawActions[3].float) * arenaHeight * 0.125, 0.0, arenaHeight)
|
|
||||||
let fireDec = tanh(rawActions[4].float)
|
|
||||||
let fp = sigmoid(rawActions[5].float) * 2.9 + 0.1
|
|
||||||
|
|
||||||
let (ts, tr) = gotoTick(gotoX, gotoY, botX, botY, direction, speed)
|
|
||||||
let gtr = aimToTick(aimToX, aimToY, botX, botY, gunDirection)
|
|
||||||
|
|
||||||
result.gotoX = gotoX
|
|
||||||
result.gotoY = gotoY
|
|
||||||
result.aimToX = aimToX
|
|
||||||
result.aimToY = aimToY
|
|
||||||
result.targetSpeed = ts
|
|
||||||
result.turnRate = tr
|
|
||||||
result.gunTurnRate = gtr
|
|
||||||
result.shouldFire = fireDec >= 0.0 and gunHeat <= 0.0
|
|
||||||
result.firePower = fp
|
|
||||||
@@ -1,13 +0,0 @@
|
|||||||
# Static-link OpenBLAS for portable deployment
|
|
||||||
# ponytail: adjust path per machine, or use pkg-config
|
|
||||||
switch("passL", "-L/nix/store/v07svn2y92bvzjl51aj7c9ca1cwg7rw7-openblas-0.3.32/lib -lopenblas")
|
|
||||||
switch("threads", "on")
|
|
||||||
# begin Nimble config (version 2)
|
|
||||||
when withDir(thisDir(), system.fileExists("nimble.paths")):
|
|
||||||
include "nimble.paths"
|
|
||||||
# end Nimble config
|
|
||||||
# Use the repo-vendored Tank Royale bot API (libs/) instead of the nimble pkg.
|
|
||||||
# The pkg copy lived in ~/.nimble/pkgs2 and was patched ad-hoc; vendoring makes
|
|
||||||
# the build self-contained and keeps the cross-thread Channel fix in-tree.
|
|
||||||
# Must come AFTER the nimble.paths include: later --path wins the import search.
|
|
||||||
switch("path", thisDir() & "/../libs/tankroyale_botapi")
|
|
||||||
@@ -1,24 +0,0 @@
|
|||||||
## Pure tick-level controllers for goto(x,y) and aimTo(x,y).
|
|
||||||
## No bot object needed — all inputs are explicit parameters.
|
|
||||||
|
|
||||||
import tankroyale_botapi
|
|
||||||
|
|
||||||
proc gotoTick*(targetX, targetY, botX, botY, direction, speed: float): (float, float) =
|
|
||||||
## Returns (targetSpeed, turnRate) to drive toward (targetX, targetY).
|
|
||||||
## Selects forward or reverse automatically based on bearing.
|
|
||||||
let bearing = normalizeRelativeAngle(directionTo(botX, botY, targetX, targetY) - direction)
|
|
||||||
let (dirSign, effBearing) =
|
|
||||||
if abs(bearing) > 90.0:
|
|
||||||
(-1.0, normalizeRelativeAngle(bearing + 180.0))
|
|
||||||
else:
|
|
||||||
(1.0, bearing)
|
|
||||||
let dist = distanceTo(botX, botY, targetX, targetY)
|
|
||||||
let maxTurn = calcMaxTurnRate(speed)
|
|
||||||
let turnRate = effBearing.clamp(-maxTurn, maxTurn)
|
|
||||||
let rawSpeed = getNewTargetSpeed(MAX_SPEED, abs(speed), dist)
|
|
||||||
(dirSign * rawSpeed, turnRate)
|
|
||||||
|
|
||||||
proc aimToTick*(targetX, targetY, botX, botY, gunDirection: float): float =
|
|
||||||
## Returns gunTurnRate (clamped to ±MAX_GUN_TURN_RATE) to rotate gun toward target.
|
|
||||||
let bearing = normalizeRelativeAngle(directionTo(botX, botY, targetX, targetY) - gunDirection)
|
|
||||||
bearing.clamp(-MAX_GUN_TURN_RATE, MAX_GUN_TURN_RATE)
|
|
||||||
@@ -1,100 +0,0 @@
|
|||||||
## Enemy tracker — deterministic radar lock + dead reckoning for PPO_Bot.
|
|
||||||
## No bot API imports; takes plain floats.
|
|
||||||
|
|
||||||
import std/math
|
|
||||||
|
|
||||||
type
|
|
||||||
EnemyState* = object
|
|
||||||
x*, y*: float64
|
|
||||||
direction*: float64
|
|
||||||
speed*: float64
|
|
||||||
energy*: float64
|
|
||||||
ticksSinceLastScan*: int
|
|
||||||
hasFired*: bool
|
|
||||||
lastFirePower*: float64
|
|
||||||
|
|
||||||
EnemyTracker* = object
|
|
||||||
current*: EnemyState
|
|
||||||
history*: array[5, tuple[x, y, direction, speed: float64]] # sliding window
|
|
||||||
historyCount*: int # valid entries 0-5
|
|
||||||
prevEnergy*: float64
|
|
||||||
hasContact*: bool
|
|
||||||
|
|
||||||
proc initEnemyTracker*(): EnemyTracker = discard
|
|
||||||
|
|
||||||
proc update*(tracker: var EnemyTracker;
|
|
||||||
scanX, scanY, scanDir, scanSpeed, scanEnergy: float64) =
|
|
||||||
## Call on ScannedBotEvent. Detects enemy fire from energy delta.
|
|
||||||
|
|
||||||
# Shift history window
|
|
||||||
if tracker.historyCount == 0:
|
|
||||||
# Cold-start: pre-fill all slots with the incoming scan so indices 22-41
|
|
||||||
# are never zero-padded on tick 1. Accel/turn-rate correctly stay 0 (no delta yet).
|
|
||||||
for i in 0 ..< 5:
|
|
||||||
tracker.history[i] = (scanX, scanY, scanDir, scanSpeed)
|
|
||||||
tracker.historyCount = 5
|
|
||||||
else:
|
|
||||||
for i in countdown(min(tracker.historyCount, 4), 1):
|
|
||||||
tracker.history[i] = tracker.history[i - 1]
|
|
||||||
tracker.history[0] = (tracker.current.x, tracker.current.y,
|
|
||||||
tracker.current.direction, tracker.current.speed)
|
|
||||||
if tracker.historyCount < 5:
|
|
||||||
inc tracker.historyCount
|
|
||||||
|
|
||||||
# Detect firing: energy drop in [0.1, 3.0] means enemy fired
|
|
||||||
let delta = tracker.prevEnergy - scanEnergy
|
|
||||||
if tracker.hasContact and delta >= 0.1 and delta <= 3.0:
|
|
||||||
tracker.current.hasFired = true
|
|
||||||
tracker.current.lastFirePower = delta
|
|
||||||
else:
|
|
||||||
tracker.current.hasFired = false
|
|
||||||
|
|
||||||
tracker.prevEnergy = scanEnergy
|
|
||||||
tracker.current.x = scanX
|
|
||||||
tracker.current.y = scanY
|
|
||||||
tracker.current.direction = scanDir
|
|
||||||
tracker.current.speed = scanSpeed
|
|
||||||
tracker.current.energy = scanEnergy
|
|
||||||
tracker.current.ticksSinceLastScan = 0
|
|
||||||
tracker.hasContact = true
|
|
||||||
|
|
||||||
proc deadReckon*(tracker: var EnemyTracker) =
|
|
||||||
## Call on missed ticks. Predict position from last known velocity.
|
|
||||||
if not tracker.hasContact:
|
|
||||||
return
|
|
||||||
let rad = tracker.current.direction * PI / 180.0
|
|
||||||
tracker.current.x += tracker.current.speed * sin(rad)
|
|
||||||
tracker.current.y += tracker.current.speed * cos(rad)
|
|
||||||
inc tracker.current.ticksSinceLastScan
|
|
||||||
|
|
||||||
proc normalizeRelative(angle: float64): float64 {.inline.} =
|
|
||||||
result = angle mod 360.0
|
|
||||||
if result >= 180.0: result -= 360.0
|
|
||||||
elif result < -180.0: result += 360.0
|
|
||||||
|
|
||||||
proc getRadarTurnRate*(tracker: var EnemyTracker;
|
|
||||||
botX, botY, botDirection, radarDirection: float64): float64 =
|
|
||||||
## Returns radar turn rate (degrees/tick, positive = right).
|
|
||||||
## Before contact: full 45° sweep.
|
|
||||||
## After contact: lock with overshoot; widen if stale.
|
|
||||||
if not tracker.hasContact:
|
|
||||||
return 45.0
|
|
||||||
|
|
||||||
if tracker.current.ticksSinceLastScan >= 8:
|
|
||||||
# Lost lock — widen sweep
|
|
||||||
return 45.0
|
|
||||||
|
|
||||||
# Bearing from radar to enemy.
|
|
||||||
# Tank Royale radar directions use standard math convention (0=east, CCW+).
|
|
||||||
# arctan2(dy, dx) gives the standard math angle matching radarDirection units.
|
|
||||||
let dx = tracker.current.x - botX
|
|
||||||
let dy = tracker.current.y - botY
|
|
||||||
let absoluteDir = (180.0 * arctan2(dy, dx) / PI + 360.0) mod 360.0
|
|
||||||
var radarTurn = normalizeRelative(absoluteDir - radarDirection)
|
|
||||||
|
|
||||||
# Width Lock: overshoot proportional to arctan(36 / distance)
|
|
||||||
let distance = sqrt(dx * dx + dy * dy)
|
|
||||||
let extraTurn = min(arctan(36.0 / distance) * 180.0 / PI, 45.0)
|
|
||||||
if radarTurn < 0.0: radarTurn -= extraTurn
|
|
||||||
else: radarTurn += extraTurn
|
|
||||||
result = radarTurn.clamp(-45.0, 45.0)
|
|
||||||
@@ -1,82 +0,0 @@
|
|||||||
## network.nim — MLP and ActorCritic forward pass (inference only, no autograd).
|
|
||||||
|
|
||||||
import arraymancer
|
|
||||||
import std/[math, random]
|
|
||||||
|
|
||||||
const
|
|
||||||
STATE_DIM* = 57
|
|
||||||
ACTION_DIM* = 6
|
|
||||||
|
|
||||||
var
|
|
||||||
logStdFloor*: float32 = -3.0'f32 # overridden by PPOB_LOG_STD_FLOOR
|
|
||||||
logStdCeiling*: float32 = 0.5'f32 # overridden by PPOB_LOG_STD_CEILING
|
|
||||||
initialLogStd*: float32 = 0.0'f32 # overridden by PPOB_INITIAL_LOG_STD
|
|
||||||
|
|
||||||
type
|
|
||||||
MLP* = object
|
|
||||||
w1*, b1*: Tensor[float32] # [hidden, input], [hidden]
|
|
||||||
w2*, b2*: Tensor[float32] # [hidden, hidden], [hidden]
|
|
||||||
w3*, b3*: Tensor[float32] # [output, hidden], [output]
|
|
||||||
|
|
||||||
ActorCritic* = object
|
|
||||||
actor*: MLP
|
|
||||||
critic*: MLP
|
|
||||||
logStd*: Tensor[float32] # [ACTION_DIM] — one per action dim
|
|
||||||
|
|
||||||
proc initMLP*(inputDim, hiddenDim, outputDim: int): MLP =
|
|
||||||
# Xavier/He-style init: scale weights by sqrt(2/fan_in)
|
|
||||||
result.w1 = randomNormalTensor[float32]([hiddenDim, inputDim]) *. sqrt(2.0'f32 / inputDim.float32)
|
|
||||||
result.b1 = zeros[float32](hiddenDim)
|
|
||||||
result.w2 = randomNormalTensor[float32]([hiddenDim, hiddenDim]) *. sqrt(2.0'f32 / hiddenDim.float32)
|
|
||||||
result.b2 = zeros[float32](hiddenDim)
|
|
||||||
result.w3 = randomNormalTensor[float32]([outputDim, hiddenDim]) *. sqrt(1.0'f32 / hiddenDim.float32)
|
|
||||||
result.b3 = zeros[float32](outputDim)
|
|
||||||
|
|
||||||
proc initActorCritic*(): ActorCritic =
|
|
||||||
result.actor = initMLP(STATE_DIM, 64, ACTION_DIM)
|
|
||||||
result.critic = initMLP(STATE_DIM, 64, 1)
|
|
||||||
result.logStd = newTensor[float32](ACTION_DIM).map(proc(v: float32): float32 = initialLogStd)
|
|
||||||
|
|
||||||
proc forward*(mlp: MLP, x: Tensor[float32]): Tensor[float32] =
|
|
||||||
## x shape: [inputDim] (1D vector)
|
|
||||||
let h1 = tanh(mlp.w1 * x + mlp.b1)
|
|
||||||
let h2 = tanh(mlp.w2 * h1 + mlp.b2)
|
|
||||||
result = mlp.w3 * h2 + mlp.b3
|
|
||||||
|
|
||||||
proc actorForward*(ac: ActorCritic, state: Tensor[float32]): tuple[actions: Tensor[float32], logProb: float32] =
|
|
||||||
## state: [STATE_DIM]. Returns sampled actions [ACTION_DIM] and sum log-prob.
|
|
||||||
let mean = ac.actor.forward(state)
|
|
||||||
# Floor logStd at logStdFloor before exp → min std ≈ exp(logStdFloor)
|
|
||||||
var clampedLogStd = ac.logStd.map(proc(v: float32): float32 = clamp(v, logStdFloor, logStdCeiling))
|
|
||||||
let std = clampedLogStd.map(proc(v: float32): float32 = exp(v))
|
|
||||||
|
|
||||||
var actions = newTensor[float32](ACTION_DIM)
|
|
||||||
var logP = 0.0'f32
|
|
||||||
for i in 0..<ACTION_DIM:
|
|
||||||
let mu = mean[i]
|
|
||||||
let s = std[i]
|
|
||||||
let z = gauss(0.0'f64, 1.0'f64).float32
|
|
||||||
actions[i] = mu + s * z
|
|
||||||
# log N(a; mu, s) = -0.5*((a-mu)/s)^2 - log(s) - 0.5*log(2π)
|
|
||||||
let diff = (actions[i] - mu) / s
|
|
||||||
logP += -0.5'f32 * diff * diff - ln(s) - 0.5'f32 * ln(2.0'f32 * PI.float32)
|
|
||||||
|
|
||||||
result = (actions: actions, logProb: logP)
|
|
||||||
|
|
||||||
proc criticForward*(ac: ActorCritic, state: Tensor[float32]): float32 =
|
|
||||||
## state: [STATE_DIM]. Returns scalar value estimate.
|
|
||||||
let val = ac.critic.forward(state)
|
|
||||||
result = val[0]
|
|
||||||
|
|
||||||
proc computeLogProb*(ac: ActorCritic, state, action: Tensor[float32]): float32 =
|
|
||||||
## Log-probability of action under current policy (no sampling).
|
|
||||||
let mean = ac.actor.forward(state)
|
|
||||||
let logStdClamped = ac.logStd.map(proc(v: float32): float32 = clamp(v, logStdFloor, logStdCeiling))
|
|
||||||
let std = logStdClamped.map(proc(v: float32): float32 = exp(v))
|
|
||||||
var logP = 0.0'f32
|
|
||||||
for i in 0..<ACTION_DIM:
|
|
||||||
let mu = mean[i]
|
|
||||||
let s = std[i]
|
|
||||||
let diff = (action[i] - mu) / s
|
|
||||||
logP += -0.5'f32 * diff * diff - ln(s) - 0.5'f32 * ln(2.0'f32 * PI.float32)
|
|
||||||
result = logP
|
|
||||||
@@ -1,7 +0,0 @@
|
|||||||
{ pkgs ? import <nixpkgs> {} }:
|
|
||||||
|
|
||||||
pkgs.mkShell {
|
|
||||||
buildInputs = with pkgs; [
|
|
||||||
openblas
|
|
||||||
];
|
|
||||||
}
|
|
||||||
@@ -1,121 +0,0 @@
|
|||||||
## State vector builder — produces 57-float normalized tensor for PPO policy.
|
|
||||||
## No bot API imports; takes plain BotState + EnemyTracker structs.
|
|
||||||
|
|
||||||
import std/math
|
|
||||||
import arraymancer
|
|
||||||
import ./enemy_tracker
|
|
||||||
|
|
||||||
type
|
|
||||||
BotStateData* = object
|
|
||||||
## Plain data mirror of the bot's observable state.
|
|
||||||
x*, y*: float64
|
|
||||||
direction*: float64
|
|
||||||
speed*: float64
|
|
||||||
energy*: float64
|
|
||||||
gunDirection*: float64
|
|
||||||
gunHeat*: float64
|
|
||||||
arenaWidth*, arenaHeight*: float64
|
|
||||||
|
|
||||||
BulletData* = object
|
|
||||||
## Enemy bullet in flight (absolute arena coords + fire power).
|
|
||||||
x*, y*: float64
|
|
||||||
power*: float64 # fire power in [0.1, 3.0]; speed = 20 - 3*power
|
|
||||||
|
|
||||||
proc buildStateVector*(bot: BotStateData; enemy: EnemyTracker;
|
|
||||||
remainingGotoDistance: float64 = 0.0;
|
|
||||||
remainingGunAngle: float64 = 0.0;
|
|
||||||
bullets: openArray[BulletData] = [];
|
|
||||||
bulletCount: int = 0): Tensor[float32] =
|
|
||||||
## Build the 57-float normalized state tensor.
|
|
||||||
## Indices 0-43: existing features. Indices 44-55: up to 3 bullet slots (4 floats each).
|
|
||||||
## Index 56: scan staleness (ticksSinceLastScan / 30, clamped to 1).
|
|
||||||
## All values clipped to roughly [-1, 1] via division by physical maxima.
|
|
||||||
result = zeros[float32](57)
|
|
||||||
|
|
||||||
let aW = bot.arenaWidth
|
|
||||||
let aH = bot.arenaHeight
|
|
||||||
let diag = sqrt(aW * aW + aH * aH) # ≈ 1700 for 1200×800
|
|
||||||
let wallMax = max(aW, aH)
|
|
||||||
|
|
||||||
# --- Current tick: own bot (indices 0-6) ---
|
|
||||||
result[0] = float32(bot.x / aW)
|
|
||||||
result[1] = float32(bot.y / aH)
|
|
||||||
result[2] = float32(bot.direction / 360.0)
|
|
||||||
result[3] = float32(bot.speed / 8.0)
|
|
||||||
result[4] = float32(bot.energy / 100.0)
|
|
||||||
result[5] = float32(bot.gunDirection / 360.0)
|
|
||||||
result[6] = float32(bot.gunHeat / 1.8)
|
|
||||||
|
|
||||||
# --- Current tick: enemy (indices 7-13) ---
|
|
||||||
if enemy.hasContact:
|
|
||||||
result[7] = float32(enemy.current.x / aW)
|
|
||||||
result[8] = float32(enemy.current.y / aH)
|
|
||||||
result[9] = float32(enemy.current.direction / 360.0)
|
|
||||||
result[10] = float32(enemy.current.speed / 8.0)
|
|
||||||
result[11] = float32(enemy.current.energy / 100.0)
|
|
||||||
result[12] = float32(if enemy.current.hasFired: 1.0 else: 0.0)
|
|
||||||
result[13] = float32(enemy.current.lastFirePower / 3.0)
|
|
||||||
# else: remain 0.0
|
|
||||||
|
|
||||||
# --- Derived features (indices 14-21) ---
|
|
||||||
if enemy.hasContact:
|
|
||||||
# Enemy acceleration (speed delta from last history entry)
|
|
||||||
# Max delta is ±8 (stopped ↔ full speed); divide by 8 to normalize
|
|
||||||
if enemy.historyCount >= 1:
|
|
||||||
result[14] = float32((enemy.current.speed - enemy.history[0].speed) / 8.0)
|
|
||||||
# Enemy turn rate (direction delta from last history entry, normalized to [-1,1])
|
|
||||||
# Divide by 180 (max possible relative rotation) rather than 10 (max body turn rate)
|
|
||||||
# ponytail: /180 covers all cases; use /10 if you want sensitivity to small turns
|
|
||||||
if enemy.historyCount >= 1:
|
|
||||||
let dirDelta = ((enemy.current.direction - enemy.history[0].direction) + 540.0) mod 360.0 - 180.0
|
|
||||||
result[15] = float32(dirDelta / 180.0)
|
|
||||||
# Relative bearing to enemy (signed, from bot perspective)
|
|
||||||
let dx = enemy.current.x - bot.x
|
|
||||||
let dy = enemy.current.y - bot.y
|
|
||||||
let absDir = (180.0 * arctan2(dx, dy) / PI + 360.0) mod 360.0
|
|
||||||
let relBearing = ((absDir - bot.direction) + 540.0) mod 360.0 - 180.0
|
|
||||||
result[16] = float32(relBearing / 180.0)
|
|
||||||
# Distance to enemy
|
|
||||||
let dist = sqrt(dx * dx + dy * dy)
|
|
||||||
result[17] = float32(dist / diag)
|
|
||||||
|
|
||||||
# Wall distances (indices 18-21): top, bottom, left, right
|
|
||||||
# top = distance from bot to top wall (y=aH), bottom = distance to bottom (y=0)
|
|
||||||
# left = distance to left (x=0), right = distance to right (x=aW)
|
|
||||||
result[18] = float32((aH - bot.y) / wallMax) # top
|
|
||||||
result[19] = float32(bot.y / wallMax) # bottom
|
|
||||||
result[20] = float32(bot.x / wallMax) # left
|
|
||||||
result[21] = float32((aW - bot.x) / wallMax) # right
|
|
||||||
|
|
||||||
# --- History: 5 ticks × 4 floats = 20 floats (indices 22-41) ---
|
|
||||||
for i in 0 ..< 5:
|
|
||||||
let base = 22 + i * 4
|
|
||||||
if i < enemy.historyCount:
|
|
||||||
result[base + 0] = float32(enemy.history[i].x / aW)
|
|
||||||
result[base + 1] = float32(enemy.history[i].y / aH)
|
|
||||||
result[base + 2] = float32(enemy.history[i].direction / 360.0)
|
|
||||||
result[base + 3] = float32(enemy.history[i].speed / 8.0)
|
|
||||||
# else: remain 0.0 (pad)
|
|
||||||
|
|
||||||
# --- Goto controller inputs (indices 42-43) ---
|
|
||||||
result[42] = float32(remainingGotoDistance / diag)
|
|
||||||
result[43] = float32(remainingGunAngle / 180.0)
|
|
||||||
|
|
||||||
# --- Bullet tracking (indices 44-55): up to 3 enemy bullets, 4 floats each ---
|
|
||||||
# Per bullet: relX/aW, relY/aH, speed/20, ticksToImpact/diag
|
|
||||||
# Positions are relative to bot (useful for dodging). Slots beyond bulletCount stay 0.
|
|
||||||
for i in 0 ..< min(bulletCount, 3):
|
|
||||||
let b = bullets[i]
|
|
||||||
let bSpeed = 20.0 - 3.0 * b.power # Tank Royale bullet speed formula
|
|
||||||
let bdx = b.x - bot.x
|
|
||||||
let bdy = b.y - bot.y
|
|
||||||
let bdist = sqrt(bdx * bdx + bdy * bdy)
|
|
||||||
let ticks = if bSpeed > 0.0: bdist / bSpeed else: 0.0
|
|
||||||
let base = 44 + i * 4
|
|
||||||
result[base + 0] = float32(bdx / bot.arenaWidth)
|
|
||||||
result[base + 1] = float32(bdy / bot.arenaHeight)
|
|
||||||
result[base + 2] = float32(bSpeed / 20.0)
|
|
||||||
result[base + 3] = float32(ticks / diag)
|
|
||||||
|
|
||||||
# --- Scan staleness (index 56) ---
|
|
||||||
result[56] = float32(min(enemy.current.ticksSinceLastScan.float64 / 30.0, 1.0))
|
|
||||||
Binary file not shown.
@@ -1,69 +0,0 @@
|
|||||||
## Assert-based tests for actions.nim.
|
|
||||||
## Run: nim c -r tests/test_actions.nim
|
|
||||||
|
|
||||||
import arraymancer
|
|
||||||
import std/math
|
|
||||||
import "../actions"
|
|
||||||
|
|
||||||
template check(cond: bool, msg: string) =
|
|
||||||
if not cond:
|
|
||||||
quit("FAIL: " & msg, 1)
|
|
||||||
|
|
||||||
let arenaW = 1200.0
|
|
||||||
let arenaH = 800.0
|
|
||||||
|
|
||||||
# Build a 6-element zero tensor and a helper to set individual values
|
|
||||||
proc makeRaw(vals: array[6, float32]): Tensor[float32] =
|
|
||||||
result = zeros[float32](6)
|
|
||||||
for i in 0 ..< 6: result[i] = vals[i]
|
|
||||||
|
|
||||||
# --- 6-dim input produces a valid BotActions ---
|
|
||||||
block basicDecode:
|
|
||||||
let raw = makeRaw([0.0'f32, 0.0, 0.0, 0.0, 0.0, 0.0])
|
|
||||||
let acts = mapActions(raw, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
|
||||||
# sigmoid(0)*arenaW = 0.5*1200 = 600, sigmoid(0)*arenaH = 0.5*800 = 400
|
|
||||||
check abs(acts.gotoX - 600.0) < 1e-6, "gotoX = sigmoid(0)*arenaW"
|
|
||||||
check abs(acts.gotoY - 400.0) < 1e-6, "gotoY = sigmoid(0)*arenaH"
|
|
||||||
check abs(acts.aimToX - 600.0) < 1e-6, "aimToX = sigmoid(0)*arenaW"
|
|
||||||
check abs(acts.aimToY - 400.0) < 1e-6, "aimToY = sigmoid(0)*arenaH"
|
|
||||||
|
|
||||||
# --- Coordinates bounded to arena size ---
|
|
||||||
block coordBounds:
|
|
||||||
# Large positive raw → sigmoid ≈ 1 → close to arenaW/arenaH
|
|
||||||
let rawHigh = makeRaw([100.0'f32, 100.0, 100.0, 100.0, 0.0, 0.0])
|
|
||||||
let actsHigh = mapActions(rawHigh, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
|
||||||
check actsHigh.gotoX <= arenaW + 1e-9, "gotoX <= arenaWidth"
|
|
||||||
check actsHigh.gotoY <= arenaH + 1e-9, "gotoY <= arenaHeight"
|
|
||||||
check actsHigh.gotoX >= 0.0, "gotoX >= 0"
|
|
||||||
# Large negative raw → sigmoid ≈ 0 → close to 0
|
|
||||||
let rawLow = makeRaw([-100.0'f32, -100.0, -100.0, -100.0, 0.0, 0.0])
|
|
||||||
let actsLow = mapActions(rawLow, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
|
||||||
check actsLow.gotoX >= -1e-9, "gotoX >= 0 (low raw)"
|
|
||||||
check actsLow.gotoY >= -1e-9, "gotoY >= 0 (low raw)"
|
|
||||||
|
|
||||||
# --- Fire triggers correctly ---
|
|
||||||
block fireTrigger:
|
|
||||||
# tanh(positive) >= 0 → fire when gunHeat = 0
|
|
||||||
let rawFire = makeRaw([0.0'f32, 0.0, 0.0, 0.0, 1.0, 0.0])
|
|
||||||
let actsFire = mapActions(rawFire, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
|
||||||
check actsFire.shouldFire, "positive tanh → should fire when gun cool"
|
|
||||||
|
|
||||||
# tanh(negative) < 0 → no fire
|
|
||||||
let rawNoFire = makeRaw([0.0'f32, 0.0, 0.0, 0.0, -1.0, 0.0])
|
|
||||||
let actsNoFire = mapActions(rawNoFire, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
|
||||||
check not actsNoFire.shouldFire, "negative tanh → no fire"
|
|
||||||
|
|
||||||
# gunHeat > 0 → no fire even with positive decision
|
|
||||||
let actsHot = mapActions(rawFire, 0.5, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
|
||||||
check not actsHot.shouldFire, "positive tanh but gun hot → no fire"
|
|
||||||
|
|
||||||
# --- Fire power in [0.1, 3.0] ---
|
|
||||||
block firePowerRange:
|
|
||||||
let rawMin = makeRaw([0.0'f32, 0.0, 0.0, 0.0, 0.0, -100.0])
|
|
||||||
let rawMax = makeRaw([0.0'f32, 0.0, 0.0, 0.0, 0.0, 100.0])
|
|
||||||
let actsMin = mapActions(rawMin, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
|
||||||
let actsMax = mapActions(rawMax, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
|
||||||
check actsMin.firePower >= 0.1 - 1e-6, "firePower >= 0.1"
|
|
||||||
check actsMax.firePower <= 3.0 + 1e-6, "firePower <= 3.0"
|
|
||||||
|
|
||||||
echo "test_actions: all passed"
|
|
||||||
Binary file not shown.
@@ -1,54 +0,0 @@
|
|||||||
## Assert-based tests for controllers.nim.
|
|
||||||
## Run: nim c -r tests/test_controllers.nim
|
|
||||||
|
|
||||||
import std/math
|
|
||||||
import "../controllers"
|
|
||||||
|
|
||||||
template check(cond: bool, msg: string) =
|
|
||||||
if not cond:
|
|
||||||
quit("FAIL: " & msg, 1)
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# gotoTick tests
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
block forwardMovement:
|
|
||||||
# Bot at (0,0) facing East (geometric 0°), target due east at (100,0) → bearing=0 → forward
|
|
||||||
let (spd2, turn2) = gotoTick(100.0, 0.0, 0.0, 0.0, 0.0, 0.0)
|
|
||||||
check spd2 > 0.0, "forward: targetSpeed should be positive"
|
|
||||||
check abs(turn2) < 1e-9, "forward: no turn needed when already aimed"
|
|
||||||
|
|
||||||
block reverseMovement:
|
|
||||||
# Bot at (100,0) facing East (direction=0), target at (0,0) — directly behind
|
|
||||||
# bearing = normalizeRelativeAngle(180 - 0) = 180 → |bearing|>90 → reverse
|
|
||||||
let (spd, _) = gotoTick(0.0, 0.0, 100.0, 0.0, 0.0, 0.0)
|
|
||||||
check spd < 0.0, "reverse: targetSpeed should be negative when target is behind"
|
|
||||||
|
|
||||||
block turnRateClamping:
|
|
||||||
# Bot at (0,0) facing North (game north = geometric 90°, so direction=90 in geometric)
|
|
||||||
# Target at (100,0) = East. bearing = normalizeRelativeAngle(0 - 90) = -90 → still ≤90
|
|
||||||
# Use large perpendicular target so bearing is 89°, and high speed → small maxTurn
|
|
||||||
# At speed=8, calcMaxTurnRate = 10 - 0.75*8 = 4°
|
|
||||||
# Bot facing East (0°), target at angle 89° bearing (just under 90)
|
|
||||||
let (_, turn) = gotoTick(100.0 * cos(89.0 * PI / 180.0), 100.0 * sin(89.0 * PI / 180.0), 0.0, 0.0, 0.0, 8.0)
|
|
||||||
check abs(turn) <= 4.0 + 1e-9, "turn rate clamped to calcMaxTurnRate at speed=8 (max 4°)"
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# aimToTick tests
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
block gunShortestArc:
|
|
||||||
# Gun facing East (0°), target due north (geometric 90°) → turn left +90° but clamped to 20
|
|
||||||
let rate = aimToTick(0.0, 100.0, 0.0, 0.0, 0.0)
|
|
||||||
check rate > 0.0, "gun shortest arc: should turn toward target"
|
|
||||||
# Gun facing East (0°), target due south (geometric 270° → normalised -90°)
|
|
||||||
let rate2 = aimToTick(0.0, -100.0, 0.0, 0.0, 0.0)
|
|
||||||
check rate2 < 0.0, "gun shortest arc: should turn other way for target behind"
|
|
||||||
|
|
||||||
block gunTurnRateClamping:
|
|
||||||
# 180° away → clamped to ±20
|
|
||||||
let rate = aimToTick(-100.0, 0.0, 0.0, 0.0, 0.0) # target west, gun east
|
|
||||||
check abs(rate) <= 20.0 + 1e-9, "gun turn rate clamped to ±MAX_GUN_TURN_RATE"
|
|
||||||
check abs(abs(rate) - 20.0) < 1e-9, "gun turn rate at max when 180° away"
|
|
||||||
|
|
||||||
echo "test_controllers: all passed"
|
|
||||||
Binary file not shown.
@@ -1,58 +0,0 @@
|
|||||||
## test_network.nim — assert-based tests for network.nim and actions.nim.
|
|
||||||
## Run: nim c --threads:on tests/test_network.nim && ./tests/test_network
|
|
||||||
|
|
||||||
import arraymancer
|
|
||||||
import std/[math, strformat]
|
|
||||||
import ../network
|
|
||||||
import ../actions
|
|
||||||
|
|
||||||
func isFiniteF(x: float32): bool = classify(x) notin {fcNan, fcInf, fcNegInf}
|
|
||||||
func isNaNF(x: float32): bool = classify(x) == fcNan
|
|
||||||
|
|
||||||
when isMainModule:
|
|
||||||
# ---- MLP forward shape ----
|
|
||||||
let mlp = initMLP(STATE_DIM, 64, ACTION_DIM)
|
|
||||||
let inp = zeros[float32](STATE_DIM)
|
|
||||||
let mlpOut = mlp.forward(inp)
|
|
||||||
assert mlpOut.shape[0] == ACTION_DIM, &"MLP output shape wrong: {mlpOut.shape}"
|
|
||||||
|
|
||||||
# ---- ActorCritic actorForward ----
|
|
||||||
let ac = initActorCritic()
|
|
||||||
let state = zeros[float32](STATE_DIM)
|
|
||||||
let (acts, logP) = ac.actorForward(state)
|
|
||||||
assert acts.shape[0] == ACTION_DIM, &"actorForward actions shape wrong: {acts.shape}"
|
|
||||||
assert not isNaNF(logP), "logProb is NaN"
|
|
||||||
assert isFiniteF(logP), &"logProb not finite: {logP}"
|
|
||||||
|
|
||||||
# ---- criticForward ----
|
|
||||||
let v = ac.criticForward(state)
|
|
||||||
assert not isNaNF(v), "critic value is NaN"
|
|
||||||
assert isFiniteF(v), &"critic value not finite: {v}"
|
|
||||||
|
|
||||||
# ---- logStd floor: collapsing logStd should not break actorForward ----
|
|
||||||
var ac2 = initActorCritic()
|
|
||||||
for i in 0..<ACTION_DIM: ac2.logStd[i] = -10.0'f32
|
|
||||||
let (acts2, logP2) = ac2.actorForward(state)
|
|
||||||
assert acts2.shape[0] == ACTION_DIM, "acts2 shape wrong after logStd=-10"
|
|
||||||
assert isFiniteF(logP2), &"logP2 not finite with floored logStd: {logP2}"
|
|
||||||
|
|
||||||
# ---- action mapping ranges ----
|
|
||||||
let raw = randomNormalTensor[float32](ACTION_DIM)
|
|
||||||
let speed = 4.0'f32
|
|
||||||
# arena 800×600, bot at centre, heading north, gun north
|
|
||||||
let botActs = mapActions(raw, 0.0'f32, 800.0, 600.0, 400.0, 300.0, 0.0, speed.float, 0.0, 400.0, 300.0) # gunHeat=0 → fire allowed
|
|
||||||
|
|
||||||
assert botActs.targetSpeed >= -8.0'f32 and botActs.targetSpeed <= 8.0'f32,
|
|
||||||
&"targetSpeed out of range: {botActs.targetSpeed}"
|
|
||||||
|
|
||||||
assert botActs.gunTurnRate >= -20.0'f32 and botActs.gunTurnRate <= 20.0'f32,
|
|
||||||
&"gunTurnRate out of range: {botActs.gunTurnRate}"
|
|
||||||
|
|
||||||
assert botActs.firePower >= 0.1'f32 and botActs.firePower <= 3.0'f32,
|
|
||||||
&"firePower out of range: {botActs.firePower}"
|
|
||||||
|
|
||||||
# shouldFire=false when gunHeat > 0
|
|
||||||
let noFire = mapActions(raw, 1.0'f32, 800.0, 600.0, 400.0, 300.0, 0.0, speed.float, 0.0, 400.0, 300.0)
|
|
||||||
assert not noFire.shouldFire, "shouldFire should be false when gunHeat > 0"
|
|
||||||
|
|
||||||
echo "All tests passed"
|
|
||||||
@@ -1,42 +0,0 @@
|
|||||||
## Regression test: radar lock must hold on a stationary target.
|
|
||||||
## Run: nim c -r tests/test_radar_lock.nim
|
|
||||||
|
|
||||||
import std/[math, strformat]
|
|
||||||
import "../enemy_tracker"
|
|
||||||
|
|
||||||
template check(cond: bool, msg: string) =
|
|
||||||
if not cond:
|
|
||||||
quit("FAIL: " & msg, 1)
|
|
||||||
|
|
||||||
# Bot at arena center; stationary enemy due north (same x, higher y).
|
|
||||||
# Tank Royale: y increases northward.
|
|
||||||
# Radar uses math convention (0°=east, CCW+); north = 90° in that system.
|
|
||||||
|
|
||||||
let botX = 400.0
|
|
||||||
let botY = 300.0
|
|
||||||
let enemyX = 400.0 # same x → dx = 0
|
|
||||||
let enemyY = 500.0 # north of bot → dy > 0
|
|
||||||
let trueBearing = 90.0 # north in math convention (0=east, CCW+)
|
|
||||||
|
|
||||||
# Radar starts pointing at the enemy (radarDirection = 90°, due north in math convention).
|
|
||||||
var radarDir = 90.0
|
|
||||||
|
|
||||||
var tracker = initEnemyTracker()
|
|
||||||
# Prime with contact at the known position
|
|
||||||
tracker.update(enemyX, enemyY, 0.0, 0.0, 100.0)
|
|
||||||
|
|
||||||
echo "Tick | radarDir | trueBearing | bearingErr"
|
|
||||||
for tick in 1 .. 20:
|
|
||||||
let rate = tracker.getRadarTurnRate(botX, botY, 0.0, radarDir)
|
|
||||||
radarDir = (radarDir + rate + 360.0) mod 360.0
|
|
||||||
|
|
||||||
# Simulate a successful scan every tick (enemy is stationary)
|
|
||||||
tracker.update(enemyX, enemyY, 0.0, 0.0, 100.0)
|
|
||||||
|
|
||||||
# Bearing error: signed difference, wrapped to [-180, 180]
|
|
||||||
let err = ((radarDir - trueBearing) + 540.0) mod 360.0 - 180.0
|
|
||||||
echo &" {tick:2d} | {radarDir:8.3f}° | {trueBearing:8.3f}° | {err:+.3f}°"
|
|
||||||
|
|
||||||
check abs(err) <= 15.0, &"tick {tick}: radar {radarDir:.1f}° drifted > 15° from target {trueBearing:.1f}°"
|
|
||||||
|
|
||||||
echo "All radar lock tests passed"
|
|
||||||
Binary file not shown.
@@ -1,187 +0,0 @@
|
|||||||
## Assert-based tests for EnemyTracker and StateVector.
|
|
||||||
## Run: nim c -r tests/test_state.nim
|
|
||||||
|
|
||||||
import std/[math, strformat]
|
|
||||||
import arraymancer
|
|
||||||
|
|
||||||
# Import from parent dir
|
|
||||||
import "../enemy_tracker"
|
|
||||||
import "../state_vector"
|
|
||||||
|
|
||||||
template check(cond: bool, msg: string) =
|
|
||||||
if not cond:
|
|
||||||
quit("FAIL: " & msg, 1)
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# EnemyTracker tests
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
block testBasicUpdate:
|
|
||||||
var t = initEnemyTracker()
|
|
||||||
t.update(200.0, 300.0, 90.0, 5.0, 80.0)
|
|
||||||
check t.hasContact, "hasContact after update"
|
|
||||||
check t.current.x == 200.0, "x after update"
|
|
||||||
check t.current.y == 300.0, "y after update"
|
|
||||||
check t.current.direction == 90.0, "direction after update"
|
|
||||||
check t.current.speed == 5.0, "speed after update"
|
|
||||||
check t.current.energy == 80.0, "energy after update"
|
|
||||||
check t.current.ticksSinceLastScan == 0, "ticksSinceLastScan reset"
|
|
||||||
|
|
||||||
block testFireDetection:
|
|
||||||
var t = initEnemyTracker()
|
|
||||||
# First update sets prevEnergy
|
|
||||||
t.update(100.0, 100.0, 0.0, 0.0, 100.0)
|
|
||||||
# Second update: energy drop of 3.0 → enemy fired power 3.0
|
|
||||||
t.update(100.0, 100.0, 0.0, 0.0, 97.0)
|
|
||||||
check t.current.hasFired, "hasFired when energy drops by 3.0"
|
|
||||||
check abs(t.current.lastFirePower - 3.0) < 0.001, "lastFirePower == 3.0"
|
|
||||||
|
|
||||||
block testNoFireOnSmallDrop:
|
|
||||||
var t = initEnemyTracker()
|
|
||||||
t.update(100.0, 100.0, 0.0, 0.0, 100.0)
|
|
||||||
# Drop of 0.05 — below MIN_FIRE_POWER threshold
|
|
||||||
t.update(100.0, 100.0, 0.0, 0.0, 99.95)
|
|
||||||
check not t.current.hasFired, "no fire on small energy drop"
|
|
||||||
|
|
||||||
block testDeadReckoning:
|
|
||||||
var t = initEnemyTracker()
|
|
||||||
# direction=0° in Tank Royale means north (y increases)
|
|
||||||
t.update(100.0, 100.0, 0.0, 5.0, 100.0)
|
|
||||||
t.deadReckon()
|
|
||||||
# x unchanged (sin 0° = 0), y increases by speed (cos 0° = 1)
|
|
||||||
check abs(t.current.x - 100.0) < 0.001, "dead reckon: x unchanged for dir=0"
|
|
||||||
check abs(t.current.y - 105.0) < 0.001, "dead reckon: y += speed for dir=0"
|
|
||||||
check t.current.ticksSinceLastScan == 1, "ticksSinceLastScan incremented"
|
|
||||||
|
|
||||||
block testDeadReckonEast:
|
|
||||||
var t = initEnemyTracker()
|
|
||||||
# direction=90° → east (sin 90° = 1, cos 90° = 0)
|
|
||||||
t.update(100.0, 100.0, 90.0, 5.0, 100.0)
|
|
||||||
t.deadReckon()
|
|
||||||
check abs(t.current.x - 105.0) < 0.001, "dead reckon east: x += speed"
|
|
||||||
check abs(t.current.y - 100.0) < 0.001, "dead reckon east: y unchanged"
|
|
||||||
|
|
||||||
block testHistoryWindow:
|
|
||||||
var t = initEnemyTracker()
|
|
||||||
# Feed 6 updates — history should hold last 5
|
|
||||||
for i in 1 .. 6:
|
|
||||||
t.update(float64(i) * 10.0, float64(i) * 20.0, 0.0, float64(i), 100.0)
|
|
||||||
check t.historyCount == 5, "historyCount capped at 5"
|
|
||||||
# history[0] should be the second-to-last scan (i=5)
|
|
||||||
check abs(t.history[0].x - 50.0) < 0.001, "history[0].x == 50 (i=5)"
|
|
||||||
check abs(t.history[4].x - 10.0) < 0.001, "history[4].x == 10 (i=1)"
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# StateVector tests
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
block testStateVectorLength:
|
|
||||||
var t = initEnemyTracker()
|
|
||||||
t.update(400.0, 300.0, 45.0, 3.0, 80.0)
|
|
||||||
let bot = BotStateData(
|
|
||||||
x: 200.0, y: 200.0, direction: 90.0, speed: 4.0, energy: 50.0,
|
|
||||||
gunDirection: 180.0, gunHeat: 0.5,
|
|
||||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
|
||||||
)
|
|
||||||
let sv = buildStateVector(bot, t)
|
|
||||||
check sv.shape == [57], "state vector has 57 elements"
|
|
||||||
|
|
||||||
block testStateVectorRange:
|
|
||||||
var t = initEnemyTracker()
|
|
||||||
t.update(400.0, 300.0, 180.0, 8.0, 100.0)
|
|
||||||
let bot = BotStateData(
|
|
||||||
x: 800.0, y: 600.0, direction: 360.0, speed: 8.0, energy: 100.0,
|
|
||||||
gunDirection: 360.0, gunHeat: 1.8,
|
|
||||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
|
||||||
)
|
|
||||||
let sv = buildStateVector(bot, t)
|
|
||||||
for i in 0 ..< 57:
|
|
||||||
check sv[i] >= -2.0f32 and sv[i] <= 2.0f32,
|
|
||||||
&"sv[{i}]={sv[i]} out of [-2,2] range"
|
|
||||||
|
|
||||||
block testWallDistances:
|
|
||||||
# Bot at (100, 200) in 800×600 arena
|
|
||||||
# wallMax = max(800, 600) = 800
|
|
||||||
# top = (600 - 200) / 800 = 400/800 = 0.5
|
|
||||||
# bottom = 200 / 800 = 0.25
|
|
||||||
# left = 100 / 800 = 0.125
|
|
||||||
# right = (800 - 100) / 800 = 700/800 = 0.875
|
|
||||||
var t = initEnemyTracker()
|
|
||||||
let bot = BotStateData(
|
|
||||||
x: 100.0, y: 200.0, direction: 0.0, speed: 0.0, energy: 100.0,
|
|
||||||
gunDirection: 0.0, gunHeat: 0.0,
|
|
||||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
|
||||||
)
|
|
||||||
let sv = buildStateVector(bot, t)
|
|
||||||
check abs(sv[18] - 0.5f32) < 0.001f32, "top wall = 0.5"
|
|
||||||
check abs(sv[19] - 0.25f32) < 0.001f32, "bottom wall = 0.25"
|
|
||||||
check abs(sv[20] - 0.125f32) < 0.001f32, "left wall = 0.125"
|
|
||||||
check abs(sv[21] - 0.875f32) < 0.001f32, "right wall = 0.875"
|
|
||||||
|
|
||||||
block testRelativeBearing:
|
|
||||||
# Game convention: north=0°, CW. arctan2(dx,dy) used.
|
|
||||||
# Bot at (0,0) dir=0°. Enemy at (0,100) → due north → absDir=0°.
|
|
||||||
# relBearing = (0 - 0 + 540) mod 360 - 180 = 0°. normalized = 0/180 = 0.0
|
|
||||||
var t = initEnemyTracker()
|
|
||||||
t.update(0.0, 100.0, 0.0, 0.0, 100.0)
|
|
||||||
let bot = BotStateData(
|
|
||||||
x: 0.0, y: 0.0, direction: 0.0, speed: 0.0, energy: 100.0,
|
|
||||||
gunDirection: 0.0, gunHeat: 0.0,
|
|
||||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
|
||||||
)
|
|
||||||
let sv = buildStateVector(bot, t)
|
|
||||||
check abs(sv[16] - 0.0f32) < 0.01f32, "relative bearing = 0.0 (due north), got " & $sv[16]
|
|
||||||
|
|
||||||
block testRelativeBearingEast:
|
|
||||||
# Enemy at (100,0) → due east → absDir=90°.
|
|
||||||
# relBearing = (90 - 0 + 540) mod 360 - 180 = 90°. normalized = 90/180 = 0.5
|
|
||||||
var t = initEnemyTracker()
|
|
||||||
t.update(100.0, 0.0, 0.0, 0.0, 100.0)
|
|
||||||
let bot = BotStateData(
|
|
||||||
x: 0.0, y: 0.0, direction: 0.0, speed: 0.0, energy: 100.0,
|
|
||||||
gunDirection: 0.0, gunHeat: 0.0,
|
|
||||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
|
||||||
)
|
|
||||||
let sv = buildStateVector(bot, t)
|
|
||||||
check abs(sv[16] - 0.5f32) < 0.01f32, "relative bearing east = 0.5, got " & $sv[16]
|
|
||||||
|
|
||||||
block testHistoryPaddedWhenEmpty:
|
|
||||||
var t = initEnemyTracker()
|
|
||||||
let bot = BotStateData(
|
|
||||||
x: 400.0, y: 300.0, direction: 0.0, speed: 0.0, energy: 100.0,
|
|
||||||
gunDirection: 0.0, gunHeat: 0.0,
|
|
||||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
|
||||||
)
|
|
||||||
let sv = buildStateVector(bot, t)
|
|
||||||
for i in 22 ..< 42:
|
|
||||||
check sv[i] == 0.0f32, &"history slot {i} should be 0 when no contact"
|
|
||||||
# indices 42-43 (goto inputs) default to 0 when not provided
|
|
||||||
check sv[42] == 0.0f32, "sv[42] (remainingGotoDistance) should default to 0"
|
|
||||||
check sv[43] == 0.0f32, "sv[43] (remainingGunAngle) should default to 0"
|
|
||||||
# indices 44-55 (bullet slots) default to 0 when no bullets provided
|
|
||||||
for i in 44 ..< 56:
|
|
||||||
check sv[i] == 0.0f32, &"bullet slot {i} should be 0 when no bullets"
|
|
||||||
# index 56 (staleness): no contact so ticksSinceLastScan=0 → 0/30 = 0
|
|
||||||
check sv[56] == 0.0f32, "sv[56] staleness should be 0 when no contact"
|
|
||||||
|
|
||||||
block testBulletSlots:
|
|
||||||
var t = initEnemyTracker()
|
|
||||||
t.update(400.0, 300.0, 0.0, 0.0, 100.0) # enemy at (400,300)
|
|
||||||
let bot = BotStateData(
|
|
||||||
x: 200.0, y: 200.0, direction: 0.0, speed: 0.0, energy: 100.0,
|
|
||||||
gunDirection: 0.0, gunHeat: 0.0,
|
|
||||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
|
||||||
)
|
|
||||||
# Bullet at (300,250), power 1.0 → speed = 20-3 = 17
|
|
||||||
# relX = 300-200 = 100, relY = 250-200 = 50 (relative to bot, not enemy)
|
|
||||||
# dist to bot = sqrt(100^2+50^2) ≈ 111.8, ticks ≈ 111.8/17 ≈ 6.58
|
|
||||||
let b = BulletData(x: 300.0, y: 250.0, power: 1.0)
|
|
||||||
let sv = buildStateVector(bot, t, 0.0, 0.0, [b], 1)
|
|
||||||
check abs(sv[44] - (100.0/800.0).float32) < 0.001f32, "bullet relX"
|
|
||||||
check abs(sv[45] - (50.0/600.0).float32) < 0.001f32, "bullet relY"
|
|
||||||
check abs(sv[46] - (17.0/20.0).float32) < 0.001f32, "bullet speed norm"
|
|
||||||
# second slot should be zero-padded
|
|
||||||
for i in 48 ..< 57:
|
|
||||||
check sv[i] == 0.0f32, &"unused bullet slot {i} should be 0"
|
|
||||||
|
|
||||||
echo "All tests passed"
|
|
||||||
Binary file not shown.
@@ -1,175 +0,0 @@
|
|||||||
## test_training.nim — assert-based tests for training.nim
|
|
||||||
## Run: nim c tests/test_training.nim && ./tests/test_training
|
|
||||||
|
|
||||||
import std/[math, random]
|
|
||||||
import arraymancer
|
|
||||||
import "../network"
|
|
||||||
import "../training"
|
|
||||||
|
|
||||||
template check(cond: bool, msg: string) =
|
|
||||||
if not cond:
|
|
||||||
quit("FAIL: " & msg, 1)
|
|
||||||
|
|
||||||
# ── computeTickReward ─────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
block testTickReward:
|
|
||||||
# I lost 2, enemy lost 10 → reward = -2 - (-10) = 8
|
|
||||||
# + default closeness shaping 0.01*(1-0/maxDist) = 0.01 (gunBearingAbs=180 → 0)
|
|
||||||
let r = computeTickReward(-2.0'f32, -10.0'f32)
|
|
||||||
check abs(r - 8.01'f32) < 1e-6'f32, "computeTickReward(-2, -10) == 8.01, got " & $r
|
|
||||||
|
|
||||||
# ── computeRoundReward ────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
block testRoundReward:
|
|
||||||
let r = computeRoundReward(350.0'f32)
|
|
||||||
check abs(r - 7.0'f32) < 1e-6'f32, "computeRoundReward(350) == 7.0, got " & $r
|
|
||||||
# bounded: long-battle cumulative scores must saturate, not blow the value scale
|
|
||||||
check abs(computeRoundReward(89299.0'f32) - 8.0'f32) < 1e-6'f32,
|
|
||||||
"computeRoundReward(89299) == 8.0 (capped), got " & $computeRoundReward(89299.0'f32)
|
|
||||||
|
|
||||||
# ── TrajectoryBuffer ──────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
block testBuffer:
|
|
||||||
var buf = initTrajectoryBuffer()
|
|
||||||
check buf.len == 0, "empty buffer len == 0"
|
|
||||||
|
|
||||||
let t1 = Transition(state: zeros[float32](STATE_DIM).stateToArr,
|
|
||||||
action: zeros[float32](ACTION_DIM).actionToArr,
|
|
||||||
logProb: -1.0'f32, reward: 0.5'f32, value: 0.3'f32)
|
|
||||||
buf.add(t1)
|
|
||||||
buf.add(t1)
|
|
||||||
buf.add(t1)
|
|
||||||
check buf.len == 3, "buffer len == 3 after 3 adds"
|
|
||||||
|
|
||||||
buf.clear()
|
|
||||||
check buf.len == 0, "buffer len == 0 after clear"
|
|
||||||
|
|
||||||
# ── computeGAE — hand-calculated 3-step ──────────────────────────────────────
|
|
||||||
|
|
||||||
block testGAE:
|
|
||||||
# rewards = [1.0, 0.0, 1.0], values = [0.5, 0.5, 0.5], lastValue = 0.0
|
|
||||||
# gamma = 0.99, lam = 0.95
|
|
||||||
# delta_2 = 1.0 + 0.99*0.0 - 0.5 = 0.5
|
|
||||||
# adv_2 = 0.5
|
|
||||||
# delta_1 = 0.0 + 0.99*0.5 - 0.5 = -0.005
|
|
||||||
# adv_1 = -0.005 + 0.99*0.95*0.5 ≈ -0.005 + 0.47025 = 0.46525
|
|
||||||
# delta_0 = 1.0 + 0.99*0.5 - 0.5 = 0.995
|
|
||||||
# adv_0 = 0.995 + 0.99*0.95*0.46525 ≈ 0.995 + 0.43744 = 1.43244
|
|
||||||
let (adv, ret) = computeGAE(
|
|
||||||
rewards = @[1.0'f32, 0.0'f32, 1.0'f32],
|
|
||||||
values = @[0.5'f32, 0.5'f32, 0.5'f32],
|
|
||||||
lastValue = 0.0'f32,
|
|
||||||
gamma = 0.99'f32,
|
|
||||||
lam = 0.95'f32
|
|
||||||
)
|
|
||||||
|
|
||||||
check abs(adv[2] - 0.5'f32) < 1e-4'f32,
|
|
||||||
"adv[2] should be ~0.5, got " & $adv[2]
|
|
||||||
check abs(adv[1] - 0.46525'f32) < 1e-3'f32,
|
|
||||||
"adv[1] should be ~0.46525, got " & $adv[1]
|
|
||||||
check abs(adv[0] - 1.43244'f32) < 1e-2'f32,
|
|
||||||
"adv[0] should be ~1.43244, got " & $adv[0]
|
|
||||||
|
|
||||||
# returns = adv + values
|
|
||||||
check abs(ret[2] - (0.5'f32 + 0.5'f32)) < 1e-4'f32, "ret[2] = adv[2] + 0.5"
|
|
||||||
check abs(ret[0] - (adv[0] + 0.5'f32)) < 1e-4'f32, "ret[0] = adv[0] + 0.5"
|
|
||||||
|
|
||||||
# ── ppoUpdate runs without crash; weights change ──────────────────────────────
|
|
||||||
|
|
||||||
block testPpoUpdate:
|
|
||||||
randomize(42)
|
|
||||||
var ac = initActorCritic()
|
|
||||||
|
|
||||||
# Save a copy of w1 before update
|
|
||||||
let w1Before = ac.actor.w1.clone()
|
|
||||||
|
|
||||||
var buf = initTrajectoryBuffer()
|
|
||||||
for _ in 0..<10:
|
|
||||||
let s = randomNormalTensor[float32](STATE_DIM)
|
|
||||||
let a = randomNormalTensor[float32](ACTION_DIM)
|
|
||||||
let lp = ac.computeLogProb(s, a)
|
|
||||||
let v = ac.criticForward(s)
|
|
||||||
buf.add(Transition(state: s.stateToArr, action: a.actionToArr, logProb: lp, reward: 0.1'f32, value: v))
|
|
||||||
|
|
||||||
var adam: ACAdamStates
|
|
||||||
discard ppoUpdate(ac, buf, lastValue = 0.0'f32, adamStates = adam, epochs = 2, miniBatchSize = 5)
|
|
||||||
|
|
||||||
# Weights should have changed — compare flattened
|
|
||||||
let n = ac.actor.w1.shape[0] * ac.actor.w1.shape[1]
|
|
||||||
let w1After = ac.actor.w1.reshape(n)
|
|
||||||
let w1Flat = w1Before.reshape(n)
|
|
||||||
var changed = false
|
|
||||||
for i in 0..<n:
|
|
||||||
if abs(w1After[i] - w1Flat[i]) > 1e-9'f32:
|
|
||||||
changed = true
|
|
||||||
break
|
|
||||||
check changed, "actor w1 should change after ppoUpdate"
|
|
||||||
|
|
||||||
# ── ppoUpdate on constant-reward trajectory: zero-variance guard ─────────────
|
|
||||||
# A passive round has near-constant per-tick rewards; with constant values the
|
|
||||||
# GAE advantages are identical → zero variance. The normalization must not
|
|
||||||
# amplify/NaN on this — update must complete with finite losses.
|
|
||||||
|
|
||||||
block testPpoUpdateConstantReward:
|
|
||||||
randomize(43)
|
|
||||||
var ac = initActorCritic()
|
|
||||||
var buf = initTrajectoryBuffer()
|
|
||||||
for _ in 0..<64:
|
|
||||||
let s = randomNormalTensor[float32](STATE_DIM)
|
|
||||||
let a = randomNormalTensor[float32](ACTION_DIM)
|
|
||||||
let lp = ac.computeLogProb(s, a)
|
|
||||||
buf.add(Transition(state: s.stateToArr, action: a.actionToArr, logProb: lp,
|
|
||||||
reward: 0.05'f32, value: 0.5'f32)) # constant reward+value
|
|
||||||
|
|
||||||
var adam: ACAdamStates
|
|
||||||
let m = ppoUpdate(ac, buf, lastValue = 0.5'f32, adamStates = adam,
|
|
||||||
epochs = 2, miniBatchSize = 16)
|
|
||||||
check m.actorLoss == m.actorLoss, "actorLoss NaN on constant-reward round"
|
|
||||||
check m.valueLoss == m.valueLoss, "valueLoss NaN on constant-reward round"
|
|
||||||
check m.gradNorm == m.gradNorm, "gradNorm NaN on constant-reward round"
|
|
||||||
|
|
||||||
# ── ppoUpdate on normal-reward trajectory: finite losses ─────────────────────
|
|
||||||
|
|
||||||
block testPpoUpdateNormalReward:
|
|
||||||
randomize(44)
|
|
||||||
var ac = initActorCritic()
|
|
||||||
var buf = initTrajectoryBuffer()
|
|
||||||
for i in 0..<64:
|
|
||||||
let s = randomNormalTensor[float32](STATE_DIM)
|
|
||||||
let a = randomNormalTensor[float32](ACTION_DIM)
|
|
||||||
let lp = ac.computeLogProb(s, a)
|
|
||||||
let v = ac.criticForward(s)
|
|
||||||
let r = 0.05'f32 + 0.5'f32 * sin(float32(i))
|
|
||||||
buf.add(Transition(state: s.stateToArr, action: a.actionToArr, logProb: lp, reward: r, value: v))
|
|
||||||
|
|
||||||
var adam: ACAdamStates
|
|
||||||
let m = ppoUpdate(ac, buf, lastValue = 0.0'f32, adamStates = adam,
|
|
||||||
epochs = 2, miniBatchSize = 16)
|
|
||||||
check m.actorLoss == m.actorLoss, "actorLoss NaN on normal-reward round"
|
|
||||||
check m.valueLoss == m.valueLoss, "valueLoss NaN on normal-reward round"
|
|
||||||
check m.gradNorm == m.gradNorm, "gradNorm NaN on normal-reward round"
|
|
||||||
|
|
||||||
# ── logStd ceiling: raw param must never drift above the collection clamp ─────
|
|
||||||
# Regression for the train/collection std mismatch: logStd starting above the
|
|
||||||
# ceiling (as in the warm-start snapshot) must be clamped back to [floor, ceiling]
|
|
||||||
# by the first Adam step, so recomputed logP matches the acting policy's std.
|
|
||||||
|
|
||||||
block testLogStdCeilingClamp:
|
|
||||||
randomize(45)
|
|
||||||
var ac = initActorCritic()
|
|
||||||
ac.logStd = newTensor[float32](ACTION_DIM).map(proc(v: float32): float32 = 3.0'f32)
|
|
||||||
var buf = initTrajectoryBuffer()
|
|
||||||
for _ in 0..<16:
|
|
||||||
let s = randomNormalTensor[float32](STATE_DIM)
|
|
||||||
let a = randomNormalTensor[float32](ACTION_DIM)
|
|
||||||
let lp = ac.computeLogProb(s, a)
|
|
||||||
buf.add(Transition(state: s.stateToArr, action: a.actionToArr, logProb: lp,
|
|
||||||
reward: 0.1'f32, value: 0.5'f32))
|
|
||||||
var adam: ACAdamStates
|
|
||||||
discard ppoUpdate(ac, buf, lastValue = 0.0'f32, adamStates = adam,
|
|
||||||
epochs = 1, miniBatchSize = 16)
|
|
||||||
for v in ac.logStd:
|
|
||||||
check v <= logStdCeiling + 1e-6'f32, "logStd above ceiling after ppoUpdate"
|
|
||||||
check v >= logStdFloor - 1e-6'f32, "logStd below floor after ppoUpdate"
|
|
||||||
|
|
||||||
echo "All tests passed"
|
|
||||||
Binary file not shown.
@@ -1,105 +0,0 @@
|
|||||||
## test_weights.nim — assert-based tests for weights.nim
|
|
||||||
## Run: nim c tests/test_weights.nim && ./tests/test_weights
|
|
||||||
|
|
||||||
import std/[os, math]
|
|
||||||
import arraymancer
|
|
||||||
import "../network"
|
|
||||||
import "../training"
|
|
||||||
import "../weights"
|
|
||||||
|
|
||||||
template check(cond: bool, msg: string) =
|
|
||||||
if not cond:
|
|
||||||
quit("FAIL: " & msg, 1)
|
|
||||||
|
|
||||||
const tmpBase = "/tmp/test_weights_nim"
|
|
||||||
|
|
||||||
# ── saveWeights / loadWeights roundtrip ───────────────────────────────────────
|
|
||||||
|
|
||||||
block testRoundtrip:
|
|
||||||
let dir = tmpBase & "_roundtrip"
|
|
||||||
removeDir(dir)
|
|
||||||
|
|
||||||
let ac1 = initActorCritic()
|
|
||||||
saveWeights(ac1, dir)
|
|
||||||
|
|
||||||
var ac2 = initActorCritic()
|
|
||||||
loadWeights(ac2, dir)
|
|
||||||
|
|
||||||
# Verify a sample of tensors
|
|
||||||
template tensorEq(a, b: Tensor[float32]) =
|
|
||||||
check a.shape == b.shape, "shape mismatch"
|
|
||||||
let diff = abs(a - b)
|
|
||||||
var maxDiff = 0.0'f32
|
|
||||||
for v in diff: maxDiff = max(maxDiff, v)
|
|
||||||
check maxDiff < 1e-6'f32, "tensor values differ by " & $maxDiff
|
|
||||||
|
|
||||||
tensorEq(ac1.actor.w1, ac2.actor.w1)
|
|
||||||
tensorEq(ac1.actor.b1, ac2.actor.b1)
|
|
||||||
tensorEq(ac1.actor.w3, ac2.actor.w3)
|
|
||||||
tensorEq(ac1.critic.w1, ac2.critic.w1)
|
|
||||||
tensorEq(ac1.critic.b3, ac2.critic.b3)
|
|
||||||
tensorEq(ac1.logStd, ac2.logStd)
|
|
||||||
|
|
||||||
removeDir(dir)
|
|
||||||
|
|
||||||
# ── saveWeightsAtomic ─────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
block testAtomic:
|
|
||||||
let dir = tmpBase & "_atomic"
|
|
||||||
removeDir(dir)
|
|
||||||
|
|
||||||
let ac = initActorCritic()
|
|
||||||
saveWeightsAtomic(ac, dir)
|
|
||||||
|
|
||||||
check dirExists(dir), "targetDir should exist after atomic save"
|
|
||||||
for f in ["actor_w1.npy", "critic_w1.npy", "log_std.npy"]:
|
|
||||||
check fileExists(dir / f), "missing file: " & f
|
|
||||||
|
|
||||||
removeDir(dir)
|
|
||||||
|
|
||||||
# ── loadBestAvailable ─────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
block testLoadBest:
|
|
||||||
let root = tmpBase & "_loadbest"
|
|
||||||
removeDir(root)
|
|
||||||
createDir(root)
|
|
||||||
|
|
||||||
let ac0 = initActorCritic()
|
|
||||||
|
|
||||||
# Try with no weights — should return false
|
|
||||||
var acEmpty = initActorCritic()
|
|
||||||
var adamEmpty: ACAdamStates
|
|
||||||
check not loadBestAvailable(acEmpty, adamEmpty, root).loaded, "should return false with no weights"
|
|
||||||
|
|
||||||
# Save to latest/; should load
|
|
||||||
saveWeights(ac0, root / "latest")
|
|
||||||
var ac1 = initActorCritic()
|
|
||||||
var adam1: ACAdamStates
|
|
||||||
check loadBestAvailable(ac1, adam1, root).loaded, "should load from latest/"
|
|
||||||
|
|
||||||
# Remove latest/, save to checkpoint_1/ — should fall back
|
|
||||||
removeDir(root / "latest")
|
|
||||||
saveWeights(ac0, root / "checkpoint_1")
|
|
||||||
var ac2 = initActorCritic()
|
|
||||||
var adam2: ACAdamStates
|
|
||||||
check loadBestAvailable(ac2, adam2, root).loaded, "should load from checkpoint_1/"
|
|
||||||
|
|
||||||
removeDir(root)
|
|
||||||
|
|
||||||
# ── cleanStaleTempDirs ────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
block testClean:
|
|
||||||
let root = tmpBase & "_clean"
|
|
||||||
removeDir(root)
|
|
||||||
createDir(root)
|
|
||||||
|
|
||||||
let stale = root / "latest_tmp_12345"
|
|
||||||
createDir(stale)
|
|
||||||
check dirExists(stale), "stale dir should exist before clean"
|
|
||||||
|
|
||||||
cleanStaleTempDirs(root)
|
|
||||||
check not dirExists(stale), "stale dir should be gone after clean"
|
|
||||||
|
|
||||||
removeDir(root)
|
|
||||||
|
|
||||||
echo "All tests passed"
|
|
||||||
Binary file not shown.
@@ -1,449 +0,0 @@
|
|||||||
## training.nim — Trajectory buffer, GAE, and PPO training loop.
|
|
||||||
## Uses manual backprop through the 3-layer tanh MLP + manual Adam.
|
|
||||||
## No external autograd dependencies — pure Arraymancer Tensor math.
|
|
||||||
|
|
||||||
import arraymancer
|
|
||||||
import std/[math, random, sequtils]
|
|
||||||
import ./network
|
|
||||||
|
|
||||||
# ── Types ─────────────────────────────────────────────────────────────────────
|
|
||||||
# ponytail: transitions hold PLAIN fixed-size arrays, never Arraymancer tensors.
|
|
||||||
# Tensors crossing the bot-thread → main-thread boundary get freed on the wrong
|
|
||||||
# thread's heap under ORC (bot thread SIGSEGVs mid-round in addEvent — 5 matching
|
|
||||||
# coredumps). Plain arrays are value types: no heap, no GC, safe to move across
|
|
||||||
# threads. Tensors are rebuilt from the arrays on the consuming (training) thread.
|
|
||||||
|
|
||||||
const
|
|
||||||
MAX_TRANSITIONS* = 8192 # 10 rounds × ~300 ticks + headroom
|
|
||||||
|
|
||||||
type
|
|
||||||
Transition* = object
|
|
||||||
state*: array[STATE_DIM, float32] # plain copy, rebuilt as tensor in ppoUpdate
|
|
||||||
action*: array[ACTION_DIM, float32]
|
|
||||||
logProb*: float32
|
|
||||||
reward*: float32
|
|
||||||
value*: float32 # critic estimate at collection time
|
|
||||||
done*: bool # true at episode (round) boundary
|
|
||||||
|
|
||||||
TrajectoryBuffer* = object
|
|
||||||
transitions*: array[MAX_TRANSITIONS, Transition]
|
|
||||||
len*: int
|
|
||||||
|
|
||||||
# ── Buffer ─────────────────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
proc initTrajectoryBuffer*(): TrajectoryBuffer =
|
|
||||||
result = TrajectoryBuffer()
|
|
||||||
|
|
||||||
proc add*(buf: var TrajectoryBuffer, t: Transition) =
|
|
||||||
## ponytail: fixed 8192 cap — 10 rounds × ~300 ticks with headroom. Drops
|
|
||||||
## new transitions when full. Raise cap if accumulation window grows.
|
|
||||||
if buf.len < MAX_TRANSITIONS:
|
|
||||||
buf.transitions[buf.len] = t
|
|
||||||
inc buf.len
|
|
||||||
|
|
||||||
proc clear*(buf: var TrajectoryBuffer) =
|
|
||||||
buf.len = 0
|
|
||||||
|
|
||||||
# ── Tensor → plain array (same-thread use; tensors never cross threads) ─────
|
|
||||||
|
|
||||||
proc stateToArr*(t: Tensor[float32]): array[STATE_DIM, float32] =
|
|
||||||
for i in 0..<STATE_DIM: result[i] = t[i]
|
|
||||||
|
|
||||||
proc actionToArr*(t: Tensor[float32]): array[ACTION_DIM, float32] =
|
|
||||||
for i in 0..<ACTION_DIM: result[i] = t[i]
|
|
||||||
|
|
||||||
# ── Reward helpers ─────────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
proc computeTickReward*(myEnergyDelta, enemyEnergyDelta: float32;
|
|
||||||
distToEnemy: float32 = 0.0'f32;
|
|
||||||
maxDist: float32 = 1.0'f32;
|
|
||||||
gunBearingAbs: float32 = 180.0'f32): float32 =
|
|
||||||
## Positive when we deal more damage than we receive.
|
|
||||||
## Dense shaping: closeness (0-0.01/tick) + aim quality (0-0.02/tick).
|
|
||||||
## ponytail: magnitudes 10x smaller than original to keep shaping as a nudge,
|
|
||||||
## not the dominant signal. Increase if bot ignores positioning entirely.
|
|
||||||
let sparseReward = myEnergyDelta - enemyEnergyDelta
|
|
||||||
let distReward = 0.01'f32 * (1.0'f32 - distToEnemy / maxDist)
|
|
||||||
let aimReward = 0.02'f32 * (1.0'f32 - gunBearingAbs / 180.0'f32)
|
|
||||||
result = sparseReward + distReward + aimReward
|
|
||||||
|
|
||||||
proc computeRoundReward*(roundScore: float32): float32 =
|
|
||||||
## Normalise round-end score to a rough ±6 range (doubled win signal).
|
|
||||||
## ponytail: cap cumulative score at 400 before /50 — bounded terminal bonus
|
|
||||||
## keeps critic value scale stable across battle boundaries and long battles.
|
|
||||||
result = min(roundScore, 400.0'f32) / 50.0'f32
|
|
||||||
|
|
||||||
# ── GAE ───────────────────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
proc computeGAE*(rewards, values: seq[float32];
|
|
||||||
dones: seq[bool];
|
|
||||||
lastValue: float32;
|
|
||||||
gamma: float32 = 0.99'f32;
|
|
||||||
lam: float32 = 0.95'f32):
|
|
||||||
tuple[advantages: seq[float32], returns: seq[float32]] =
|
|
||||||
## Generalised Advantage Estimation — reverse sweep with episode boundaries.
|
|
||||||
## When done=true on transition t, bootstrap value and accumulated GAE are
|
|
||||||
## reset to 0 at that boundary (terminal state has no future value).
|
|
||||||
let n = rewards.len
|
|
||||||
var advantages = newSeq[float32](n)
|
|
||||||
var lastGae = 0.0'f32
|
|
||||||
|
|
||||||
for t in countdown(n - 1, 0):
|
|
||||||
let nextVal: float32 =
|
|
||||||
if t == n - 1 or dones[t]: 0.0'f32
|
|
||||||
else: values[t + 1]
|
|
||||||
if t == n - 1 or dones[t]:
|
|
||||||
lastGae = 0.0'f32
|
|
||||||
let delta = rewards[t] + gamma * nextVal - values[t]
|
|
||||||
lastGae = delta + gamma * lam * lastGae
|
|
||||||
advantages[t] = lastGae
|
|
||||||
|
|
||||||
var returns = newSeq[float32](n)
|
|
||||||
for t in 0..<n:
|
|
||||||
returns[t] = advantages[t] + values[t]
|
|
||||||
|
|
||||||
result = (advantages: advantages, returns: returns)
|
|
||||||
|
|
||||||
# ── Manual Adam state ─────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
type
|
|
||||||
AdamState* = object
|
|
||||||
m*, v*: Tensor[float32]
|
|
||||||
t*: int
|
|
||||||
|
|
||||||
proc initAdamState(like: Tensor[float32]): AdamState =
|
|
||||||
result.m = zeros_like(like)
|
|
||||||
result.v = zeros_like(like)
|
|
||||||
result.t = 0
|
|
||||||
|
|
||||||
proc adamStep(param: var Tensor[float32];
|
|
||||||
grad: Tensor[float32];
|
|
||||||
state: var AdamState;
|
|
||||||
lr: float32 = 3e-4'f32;
|
|
||||||
beta1: float32 = 0.9'f32;
|
|
||||||
beta2: float32 = 0.999'f32;
|
|
||||||
eps: float32 = 1e-8'f32) =
|
|
||||||
inc state.t
|
|
||||||
state.m = beta1 *. state.m + (1.0'f32 - beta1) *. grad
|
|
||||||
state.v = beta2 *. state.v + (1.0'f32 - beta2) *. (grad *. grad)
|
|
||||||
let mHat = state.m /. (1.0'f32 - beta1 ^ state.t.float32)
|
|
||||||
let vHat = state.v /. (1.0'f32 - beta2 ^ state.t.float32)
|
|
||||||
param -= lr *. mHat /. (vHat.map(proc(x: float32): float32 = sqrt(x) + eps))
|
|
||||||
|
|
||||||
# ── MLP forward with cached activations (for backprop) ────────────────────────
|
|
||||||
|
|
||||||
type MLPFwd = object
|
|
||||||
h1, h2, y: Tensor[float32] # activations (h1=layer1, h2=layer2, y=output)
|
|
||||||
|
|
||||||
proc mlpForwardCached(mlp: MLP; x: Tensor[float32]): MLPFwd =
|
|
||||||
## Forward pass saving intermediate activations needed for backprop.
|
|
||||||
result.h1 = tanh(mlp.w1 * x + mlp.b1)
|
|
||||||
result.h2 = tanh(mlp.w2 * result.h1 + mlp.b2)
|
|
||||||
result.y = mlp.w3 * result.h2 + mlp.b3
|
|
||||||
|
|
||||||
proc mlpBackward(mlp: MLP; fwd: MLPFwd; x: Tensor[float32];
|
|
||||||
gradOut: Tensor[float32]):
|
|
||||||
tuple[dw1, db1, dw2, db2, dw3, db3: Tensor[float32]] =
|
|
||||||
## Chain-rule through 3-layer tanh MLP.
|
|
||||||
## gradOut: [outputDim] — d_loss / d_out
|
|
||||||
# Layer 3
|
|
||||||
let dw3 = gradOut.unsqueeze(1) * fwd.h2.unsqueeze(0) # [out, hidden]
|
|
||||||
let db3 = gradOut
|
|
||||||
let dh2 = mlp.w3.transpose * gradOut # [hidden]
|
|
||||||
# tanh backward: d/dx tanh(x) = 1 - tanh²(x)
|
|
||||||
let dpre2 = dh2 *. (ones[float32](fwd.h2.shape) - fwd.h2 *. fwd.h2)
|
|
||||||
# Layer 2
|
|
||||||
let dw2 = dpre2.unsqueeze(1) * fwd.h1.unsqueeze(0) # [hidden, hidden]
|
|
||||||
let db2 = dpre2
|
|
||||||
let dh1 = mlp.w2.transpose * dpre2 # [hidden]
|
|
||||||
let dpre1 = dh1 *. (ones[float32](fwd.h1.shape) - fwd.h1 *. fwd.h1)
|
|
||||||
# Layer 1
|
|
||||||
let dw1 = dpre1.unsqueeze(1) * x.unsqueeze(0) # [hidden, input]
|
|
||||||
let db1 = dpre1
|
|
||||||
result = (dw1: dw1, db1: db1, dw2: dw2, db2: db2, dw3: dw3, db3: db3)
|
|
||||||
|
|
||||||
# ── Adam states for ActorCritic parameters ───────────────────────────────────
|
|
||||||
|
|
||||||
type ACAdamStates* = object
|
|
||||||
## One AdamState per learnable tensor in ActorCritic.
|
|
||||||
aw1*, ab1*, aw2*, ab2*, aw3*, ab3*: AdamState # actor MLP
|
|
||||||
cw1*, cb1*, cw2*, cb2*, cw3*, cb3*: AdamState # critic MLP
|
|
||||||
logStd*: AdamState
|
|
||||||
initialized*: bool
|
|
||||||
|
|
||||||
proc initACAdamStates*(ac: ActorCritic): ACAdamStates =
|
|
||||||
result.aw1 = initAdamState(ac.actor.w1)
|
|
||||||
result.ab1 = initAdamState(ac.actor.b1)
|
|
||||||
result.aw2 = initAdamState(ac.actor.w2)
|
|
||||||
result.ab2 = initAdamState(ac.actor.b2)
|
|
||||||
result.aw3 = initAdamState(ac.actor.w3)
|
|
||||||
result.ab3 = initAdamState(ac.actor.b3)
|
|
||||||
result.cw1 = initAdamState(ac.critic.w1)
|
|
||||||
result.cb1 = initAdamState(ac.critic.b1)
|
|
||||||
result.cw2 = initAdamState(ac.critic.w2)
|
|
||||||
result.cb2 = initAdamState(ac.critic.b2)
|
|
||||||
result.cw3 = initAdamState(ac.critic.w3)
|
|
||||||
result.cb3 = initAdamState(ac.critic.b3)
|
|
||||||
result.logStd = initAdamState(ac.logStd)
|
|
||||||
result.initialized = true
|
|
||||||
|
|
||||||
# ── Training metrics ──────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
type PPOMetrics* = object
|
|
||||||
actorLoss*: float32
|
|
||||||
valueLoss*: float32
|
|
||||||
gradNorm*: float32
|
|
||||||
|
|
||||||
# ── Gradient clipping ─────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
proc globalNorm(grads: varargs[Tensor[float32]]): float32 =
|
|
||||||
var sumSq = 0.0'f32
|
|
||||||
for g in grads:
|
|
||||||
for v in g: sumSq += v * v
|
|
||||||
result = sqrt(sumSq)
|
|
||||||
|
|
||||||
# ── PPO update ────────────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
proc ppoUpdate*(ac: var ActorCritic;
|
|
||||||
buffer: TrajectoryBuffer;
|
|
||||||
lastValue: float32;
|
|
||||||
adamStates: var ACAdamStates;
|
|
||||||
epochs: int = 4;
|
|
||||||
miniBatchSize: int = 64;
|
|
||||||
clipEpsilon: float32 = 0.2'f32;
|
|
||||||
entropyCoeff: float32 = 0.01'f32;
|
|
||||||
valueLossCoeff: float32 = 0.5'f32;
|
|
||||||
lr: float32 = 3e-4'f32;
|
|
||||||
maxGradNorm: float32 = 0.5'f32;
|
|
||||||
gamma: float32 = 0.99'f32;
|
|
||||||
lam: float32 = 0.95'f32): PPOMetrics {.gcsafe.} =
|
|
||||||
if buffer.len == 0: return
|
|
||||||
|
|
||||||
var totalActorLoss = 0.0'f32
|
|
||||||
var totalValueLoss = 0.0'f32
|
|
||||||
var totalGradNorm = 0.0'f32
|
|
||||||
var totalMiniBatches = 0
|
|
||||||
|
|
||||||
# Initialise Adam states once; caller persists them across rounds.
|
|
||||||
# Also reinit if aw1.m has wrong shape (e.g. loaded from old checkpoint with
|
|
||||||
# different STATE_DIM, leaving a (0,) placeholder after shape-mismatch skip).
|
|
||||||
if not adamStates.initialized or
|
|
||||||
adamStates.aw1.m.shape.len == 0 or
|
|
||||||
adamStates.aw1.m.shape != ac.actor.w1.shape:
|
|
||||||
adamStates = initACAdamStates(ac)
|
|
||||||
|
|
||||||
# 1. GAE
|
|
||||||
let rewards = buffer.transitions[0 ..< buffer.len].mapIt(it.reward)
|
|
||||||
let values = buffer.transitions[0 ..< buffer.len].mapIt(it.value)
|
|
||||||
let dones = buffer.transitions[0 ..< buffer.len].mapIt(it.done)
|
|
||||||
let (advantages, returns) = computeGAE(rewards, values, dones, lastValue, gamma = gamma, lam = lam)
|
|
||||||
|
|
||||||
# 2. Normalise advantages
|
|
||||||
let n = advantages.len.float32
|
|
||||||
var advMean = 0.0'f32
|
|
||||||
for a in advantages: advMean += a
|
|
||||||
advMean /= n
|
|
||||||
var advVar = 0.0'f32
|
|
||||||
for a in advantages: advVar += (a - advMean) * (a - advMean)
|
|
||||||
advVar /= n
|
|
||||||
# ponytail: float32 adv noise ~1e-12; advVar < 1e-8 = constant-reward
|
|
||||||
# (passive) round — dividing by that amplifies noise ~1e4+ and drifts the
|
|
||||||
# policy into exp() overflow. Center-only, skip the divide.
|
|
||||||
var normAdv: seq[float32]
|
|
||||||
if advantages.allIt(it == it and abs(it) < 1e30'f32):
|
|
||||||
if advVar < 1e-8'f32:
|
|
||||||
normAdv = advantages.mapIt(it - advMean)
|
|
||||||
else:
|
|
||||||
let advStd = sqrt(advVar + 1e-8'f32)
|
|
||||||
normAdv = advantages.mapIt((it - advMean) / advStd)
|
|
||||||
else:
|
|
||||||
normAdv = newSeq[float32](advantages.len) # poisoned input → zero advantages, no-op update
|
|
||||||
|
|
||||||
let bufLen = buffer.len
|
|
||||||
# ponytail: minibatch size <= 0 would make mbEnd == mbStart forever and spin.
|
|
||||||
# Treat as full-batch; breaks the loop unconditionally.
|
|
||||||
let mbSizeCap = if miniBatchSize > 0: miniBatchSize else: bufLen
|
|
||||||
|
|
||||||
for epochNum in 1..epochs:
|
|
||||||
# Shuffle indices
|
|
||||||
var indices = toSeq(0..<bufLen)
|
|
||||||
shuffle(indices)
|
|
||||||
|
|
||||||
var mbStart = 0
|
|
||||||
while mbStart < bufLen:
|
|
||||||
let mbEnd = min(mbStart + mbSizeCap, bufLen)
|
|
||||||
if mbEnd <= mbStart: break
|
|
||||||
let mbSize = mbEnd - mbStart
|
|
||||||
# Accumulators for gradients (zero-init)
|
|
||||||
var dActorW1 = zeros[float32](ac.actor.w1.shape)
|
|
||||||
var dActorB1 = zeros[float32](ac.actor.b1.shape)
|
|
||||||
var dActorW2 = zeros[float32](ac.actor.w2.shape)
|
|
||||||
var dActorB2 = zeros[float32](ac.actor.b2.shape)
|
|
||||||
var dActorW3 = zeros[float32](ac.actor.w3.shape)
|
|
||||||
var dActorB3 = zeros[float32](ac.actor.b3.shape)
|
|
||||||
var dLogStd = zeros[float32](ac.logStd.shape)
|
|
||||||
|
|
||||||
var dCriticW1 = zeros[float32](ac.critic.w1.shape)
|
|
||||||
var dCriticB1 = zeros[float32](ac.critic.b1.shape)
|
|
||||||
var dCriticW2 = zeros[float32](ac.critic.w2.shape)
|
|
||||||
var dCriticB2 = zeros[float32](ac.critic.b2.shape)
|
|
||||||
var dCriticW3 = zeros[float32](ac.critic.w3.shape)
|
|
||||||
var dCriticB3 = zeros[float32](ac.critic.b3.shape)
|
|
||||||
|
|
||||||
for j in mbStart..<mbEnd:
|
|
||||||
let idx = indices[j]
|
|
||||||
let tr = buffer.transitions[idx]
|
|
||||||
let adv = normAdv[idx]
|
|
||||||
let ret = returns[idx].float32
|
|
||||||
|
|
||||||
# Rebuild the state tensor on this (training) thread — transitions hold
|
|
||||||
# plain arrays so no tensor ever crosses a thread boundary.
|
|
||||||
let x = tr.state.toTensor()
|
|
||||||
|
|
||||||
# ── Actor forward ──
|
|
||||||
let actorFwd = mlpForwardCached(ac.actor, x)
|
|
||||||
# Critic forward
|
|
||||||
|
|
||||||
let newMean = actorFwd.y # [ACTION_DIM]
|
|
||||||
|
|
||||||
# Same clamp as collection (network.nim actorForward): train-time std must
|
|
||||||
# exactly match the std the acting policy used, or ratios are distorted.
|
|
||||||
let logStdClamped = ac.logStd.map(proc(v: float32): float32 = clamp(v, logStdFloor, logStdCeiling))
|
|
||||||
let std = logStdClamped.map(proc(v: float32): float32 = exp(v))
|
|
||||||
|
|
||||||
# New log prob
|
|
||||||
var newLogP = 0.0'f32
|
|
||||||
for i in 0..<ACTION_DIM:
|
|
||||||
let mu = newMean[i]
|
|
||||||
let s = std[i]
|
|
||||||
let diff = (tr.action[i] - mu) / s
|
|
||||||
newLogP += -0.5'f32 * diff * diff - ln(s) - 0.5'f32 * ln(2.0'f32 * PI.float32)
|
|
||||||
|
|
||||||
# ponytail: float32 exp overflows at ±88; ±20 is deep in clipped-ratio
|
|
||||||
# territory, so loss/grad are identical to the true ratio
|
|
||||||
let ratio = exp(clamp(newLogP - tr.logProb, -20.0'f32, 20.0'f32))
|
|
||||||
|
|
||||||
# Clipped surrogate
|
|
||||||
let ratioClipped = clamp(ratio, 1.0'f32 - clipEpsilon, 1.0'f32 + clipEpsilon)
|
|
||||||
let surr1 = ratio * adv
|
|
||||||
let surr2 = ratioClipped * adv
|
|
||||||
# Actor loss per sample = -min(surr1, surr2)
|
|
||||||
totalActorLoss += -min(surr1, surr2)
|
|
||||||
# Which branch is active?
|
|
||||||
let useClipped = (surr2 < surr1)
|
|
||||||
let dLoss_dSurr = -1.0'f32 / mbSize.float32 # d(-mean(min))/d(min) = -1/N
|
|
||||||
|
|
||||||
# d(actor_loss)/d(ratio): only the non-clipped branch passes gradient
|
|
||||||
let dLoss_dRatio = if useClipped: 0.0'f32 else: dLoss_dSurr * adv
|
|
||||||
# d(ratio)/d(newLogP) = ratio
|
|
||||||
let dLoss_dNewLogP = dLoss_dRatio * ratio
|
|
||||||
|
|
||||||
# d(newLogP)/d(mean[i]) = (action[i] - mean[i]) / std[i]^2
|
|
||||||
var dLogP_dMean = newTensor[float32](ACTION_DIM)
|
|
||||||
for i in 0..<ACTION_DIM:
|
|
||||||
let s = std[i]
|
|
||||||
dLogP_dMean[i] = (tr.action[i] - newMean[i]) / (s * s)
|
|
||||||
|
|
||||||
# Entropy gradient for logStd:
|
|
||||||
# entropy = sum_i [ logStd_i + 0.5*(1+ln(2π)) ]
|
|
||||||
# d(entropy)/d(logStd_i) = 1 (for clamped logStd_i > -3, else 0)
|
|
||||||
# total loss gradient w.r.t. logStd: -entropyCoeff * d(entropy)/d(logStd)
|
|
||||||
# Also: d(newLogP)/d(logStd_i) when logStd_i > -3:
|
|
||||||
# = (action_i - mean_i)^2/std_i^2 - 1
|
|
||||||
for i in 0..<ACTION_DIM:
|
|
||||||
let isFloorClamped = (ac.logStd[i] <= logStdFloor)
|
|
||||||
let isCeilingClamped = (ac.logStd[i] >= logStdCeiling)
|
|
||||||
if not isFloorClamped:
|
|
||||||
let s = std[i]
|
|
||||||
let diff = (tr.action[i] - newMean[i]) / s
|
|
||||||
let dLogP_dLogStdI = diff * diff - 1.0'f32
|
|
||||||
let ppoGrad = dLoss_dNewLogP * dLogP_dLogStdI
|
|
||||||
# Entropy term pushes logStd up (update = param - lr*grad, grad is -entropyCoeff < 0).
|
|
||||||
# Gate it off at the ceiling to prevent runaway logStd.
|
|
||||||
let entropyGrad = if isCeilingClamped: 0.0'f32
|
|
||||||
else: -entropyCoeff / mbSize.float32
|
|
||||||
dLogStd[i] += ppoGrad + entropyGrad
|
|
||||||
|
|
||||||
# Backprop actor gradients
|
|
||||||
let gradActorOut = dLoss_dNewLogP *. dLogP_dMean # [5]
|
|
||||||
let actorGrads = mlpBackward(ac.actor, actorFwd, x, gradActorOut)
|
|
||||||
|
|
||||||
dActorW1 += actorGrads.dw1
|
|
||||||
dActorB1 += actorGrads.db1
|
|
||||||
dActorW2 += actorGrads.dw2
|
|
||||||
dActorB2 += actorGrads.db2
|
|
||||||
dActorW3 += actorGrads.dw3
|
|
||||||
dActorB3 += actorGrads.db3
|
|
||||||
|
|
||||||
# ── Critic forward + loss ──
|
|
||||||
let criticFwd = mlpForwardCached(ac.critic, x)
|
|
||||||
let newVal = criticFwd.y[0]
|
|
||||||
# Value loss = (newVal - ret)^2; d/d(newVal) = 2*(newVal-ret)
|
|
||||||
totalValueLoss += (newVal - ret) * (newVal - ret)
|
|
||||||
let dVLoss_dVal = valueLossCoeff * 2.0'f32 * (newVal - ret) / mbSize.float32
|
|
||||||
let gradCriticOut = [dVLoss_dVal].toTensor() # [1]
|
|
||||||
let criticGrads = mlpBackward(ac.critic, criticFwd, x, gradCriticOut)
|
|
||||||
|
|
||||||
dCriticW1 += criticGrads.dw1
|
|
||||||
dCriticB1 += criticGrads.db1
|
|
||||||
dCriticW2 += criticGrads.dw2
|
|
||||||
dCriticB2 += criticGrads.db2
|
|
||||||
dCriticW3 += criticGrads.dw3
|
|
||||||
dCriticB3 += criticGrads.db3
|
|
||||||
|
|
||||||
# ── Gradient clipping ──
|
|
||||||
# Collect all grads into a seq for norm computation
|
|
||||||
var allGrads: seq[Tensor[float32]] = @[
|
|
||||||
dActorW1, dActorB1, dActorW2, dActorB2, dActorW3, dActorB3,
|
|
||||||
dLogStd,
|
|
||||||
dCriticW1, dCriticB1, dCriticW2, dCriticB2, dCriticW3, dCriticB3
|
|
||||||
]
|
|
||||||
let norm = globalNorm(allGrads)
|
|
||||||
# ponytail: NaN/Inf grad norm means something exploded this minibatch
|
|
||||||
# (extreme logprob ratios, poisoned initial weights, etc.). Skip the
|
|
||||||
# Adam update entirely — no-op is safer than writing NaN into weights,
|
|
||||||
# which corrupts all future inference and hangs the bot.
|
|
||||||
if norm != norm or norm > 1e15'f32:
|
|
||||||
mbStart = mbEnd
|
|
||||||
continue
|
|
||||||
totalGradNorm += norm
|
|
||||||
inc totalMiniBatches
|
|
||||||
if norm > maxGradNorm:
|
|
||||||
let scale = maxGradNorm / norm
|
|
||||||
for g in allGrads.mitems: g = g *. scale
|
|
||||||
|
|
||||||
# Unpack clipped grads
|
|
||||||
dActorW1 = allGrads[0]; dActorB1 = allGrads[1]
|
|
||||||
dActorW2 = allGrads[2]; dActorB2 = allGrads[3]
|
|
||||||
dActorW3 = allGrads[4]; dActorB3 = allGrads[5]
|
|
||||||
dLogStd = allGrads[6]
|
|
||||||
dCriticW1 = allGrads[7]; dCriticB1 = allGrads[8]
|
|
||||||
dCriticW2 = allGrads[9]; dCriticB2 = allGrads[10]
|
|
||||||
dCriticW3 = allGrads[11]; dCriticB3 = allGrads[12]
|
|
||||||
|
|
||||||
# ── Adam updates ──
|
|
||||||
adamStep(ac.actor.w1, dActorW1, adamStates.aw1, lr)
|
|
||||||
adamStep(ac.actor.b1, dActorB1, adamStates.ab1, lr)
|
|
||||||
adamStep(ac.actor.w2, dActorW2, adamStates.aw2, lr)
|
|
||||||
adamStep(ac.actor.b2, dActorB2, adamStates.ab2, lr)
|
|
||||||
adamStep(ac.actor.w3, dActorW3, adamStates.aw3, lr)
|
|
||||||
adamStep(ac.actor.b3, dActorB3, adamStates.ab3, lr)
|
|
||||||
adamStep(ac.logStd, dLogStd, adamStates.logStd, lr)
|
|
||||||
# Collection clamps logStd to [floor, ceiling] at inference; clamp the raw
|
|
||||||
# param after the step so it can't drift above the ceiling (the old code
|
|
||||||
# only clamped at collection → train-time recompute used a bigger std than
|
|
||||||
# the policy that actually acted → distorted importance ratios).
|
|
||||||
ac.logStd = ac.logStd.map(proc(v: float32): float32 = clamp(v, logStdFloor, logStdCeiling))
|
|
||||||
adamStep(ac.critic.w1, dCriticW1, adamStates.cw1, lr)
|
|
||||||
adamStep(ac.critic.b1, dCriticB1, adamStates.cb1, lr)
|
|
||||||
adamStep(ac.critic.w2, dCriticW2, adamStates.cw2, lr)
|
|
||||||
adamStep(ac.critic.b2, dCriticB2, adamStates.cb2, lr)
|
|
||||||
adamStep(ac.critic.w3, dCriticW3, adamStates.cw3, lr)
|
|
||||||
adamStep(ac.critic.b3, dCriticB3, adamStates.cb3, lr)
|
|
||||||
mbStart = mbEnd
|
|
||||||
|
|
||||||
let totalSamples = (epochs * bufLen).float32
|
|
||||||
result.actorLoss = totalActorLoss / totalSamples
|
|
||||||
result.valueLoss = totalValueLoss / totalSamples
|
|
||||||
result.gradNorm = if totalMiniBatches > 0: totalGradNorm / totalMiniBatches.float32
|
|
||||||
else: 0.0'f32
|
|
||||||
@@ -1,211 +0,0 @@
|
|||||||
## weights.nim — save/load ActorCritic weights as .npy files.
|
|
||||||
|
|
||||||
import std/[os, times, strutils, algorithm, sequtils]
|
|
||||||
import arraymancer
|
|
||||||
import ./network
|
|
||||||
import ./training
|
|
||||||
|
|
||||||
# ── Tensor names — order must match save/load ─────────────────────────────────
|
|
||||||
|
|
||||||
const weightFiles = [
|
|
||||||
"actor_w1.npy", "actor_b1.npy", "actor_w2.npy", "actor_b2.npy",
|
|
||||||
"actor_w3.npy", "actor_b3.npy",
|
|
||||||
"critic_w1.npy", "critic_b1.npy", "critic_w2.npy", "critic_b2.npy",
|
|
||||||
"critic_w3.npy", "critic_b3.npy",
|
|
||||||
"log_std.npy",
|
|
||||||
]
|
|
||||||
|
|
||||||
proc saveWeights*(ac: ActorCritic, dir: string) =
|
|
||||||
## Write all weight tensors to dir/ as .npy files.
|
|
||||||
createDir(dir)
|
|
||||||
ac.actor.w1.write_npy(dir / "actor_w1.npy")
|
|
||||||
ac.actor.b1.write_npy(dir / "actor_b1.npy")
|
|
||||||
ac.actor.w2.write_npy(dir / "actor_w2.npy")
|
|
||||||
ac.actor.b2.write_npy(dir / "actor_b2.npy")
|
|
||||||
ac.actor.w3.write_npy(dir / "actor_w3.npy")
|
|
||||||
ac.actor.b3.write_npy(dir / "actor_b3.npy")
|
|
||||||
ac.critic.w1.write_npy(dir / "critic_w1.npy")
|
|
||||||
ac.critic.b1.write_npy(dir / "critic_b1.npy")
|
|
||||||
ac.critic.w2.write_npy(dir / "critic_w2.npy")
|
|
||||||
ac.critic.b2.write_npy(dir / "critic_b2.npy")
|
|
||||||
ac.critic.w3.write_npy(dir / "critic_w3.npy")
|
|
||||||
ac.critic.b3.write_npy(dir / "critic_b3.npy")
|
|
||||||
ac.logStd.write_npy(dir / "log_std.npy")
|
|
||||||
|
|
||||||
proc loadWeights*(ac: var ActorCritic, dir: string) =
|
|
||||||
## Load all weight tensors from dir/.
|
|
||||||
## If a tensor's shape doesn't match (e.g. STATE_DIM changed), keep the
|
|
||||||
## freshly-initialised value and print a warning — other tensors still load.
|
|
||||||
template loadOrSkip(dest: untyped, path: string) =
|
|
||||||
let loaded = read_npy[float32](path)
|
|
||||||
if loaded.shape == dest.shape:
|
|
||||||
dest = loaded
|
|
||||||
else:
|
|
||||||
echo "weights: shape mismatch for " & path &
|
|
||||||
" (got " & $loaded.shape & " want " & $dest.shape & ") — keeping fresh init"
|
|
||||||
|
|
||||||
loadOrSkip(ac.actor.w1, dir / "actor_w1.npy")
|
|
||||||
loadOrSkip(ac.actor.b1, dir / "actor_b1.npy")
|
|
||||||
loadOrSkip(ac.actor.w2, dir / "actor_w2.npy")
|
|
||||||
loadOrSkip(ac.actor.b2, dir / "actor_b2.npy")
|
|
||||||
loadOrSkip(ac.actor.w3, dir / "actor_w3.npy")
|
|
||||||
loadOrSkip(ac.actor.b3, dir / "actor_b3.npy")
|
|
||||||
loadOrSkip(ac.critic.w1, dir / "critic_w1.npy")
|
|
||||||
loadOrSkip(ac.critic.b1, dir / "critic_b1.npy")
|
|
||||||
loadOrSkip(ac.critic.w2, dir / "critic_w2.npy")
|
|
||||||
loadOrSkip(ac.critic.b2, dir / "critic_b2.npy")
|
|
||||||
loadOrSkip(ac.critic.w3, dir / "critic_w3.npy")
|
|
||||||
loadOrSkip(ac.critic.b3, dir / "critic_b3.npy")
|
|
||||||
loadOrSkip(ac.logStd, dir / "log_std.npy")
|
|
||||||
|
|
||||||
proc saveWeightsAtomic*(ac: ActorCritic, targetDir: string) =
|
|
||||||
## Write to a temp dir, then rename atomically over targetDir.
|
|
||||||
let tmpDir = targetDir & "_tmp_" & $int(epochTime())
|
|
||||||
saveWeights(ac, tmpDir)
|
|
||||||
if dirExists(targetDir):
|
|
||||||
removeDir(targetDir)
|
|
||||||
moveDir(tmpDir, targetDir)
|
|
||||||
|
|
||||||
# ── Adam state file names ──────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
const adamMFiles = [
|
|
||||||
"adam_aw1_m.npy", "adam_ab1_m.npy", "adam_aw2_m.npy", "adam_ab2_m.npy",
|
|
||||||
"adam_aw3_m.npy", "adam_ab3_m.npy",
|
|
||||||
"adam_cw1_m.npy", "adam_cb1_m.npy", "adam_cw2_m.npy", "adam_cb2_m.npy",
|
|
||||||
"adam_cw3_m.npy", "adam_cb3_m.npy",
|
|
||||||
"adam_logstd_m.npy",
|
|
||||||
]
|
|
||||||
const adamVFiles = [
|
|
||||||
"adam_aw1_v.npy", "adam_ab1_v.npy", "adam_aw2_v.npy", "adam_ab2_v.npy",
|
|
||||||
"adam_aw3_v.npy", "adam_ab3_v.npy",
|
|
||||||
"adam_cw1_v.npy", "adam_cb1_v.npy", "adam_cw2_v.npy", "adam_cb2_v.npy",
|
|
||||||
"adam_cw3_v.npy", "adam_cb3_v.npy",
|
|
||||||
"adam_logstd_v.npy",
|
|
||||||
]
|
|
||||||
|
|
||||||
proc saveAdamStates*(adam: ACAdamStates, dir: string) =
|
|
||||||
## Write Adam m/v tensors and t counters to dir/.
|
|
||||||
adam.aw1.m.write_npy(dir / "adam_aw1_m.npy"); adam.aw1.v.write_npy(dir / "adam_aw1_v.npy")
|
|
||||||
adam.ab1.m.write_npy(dir / "adam_ab1_m.npy"); adam.ab1.v.write_npy(dir / "adam_ab1_v.npy")
|
|
||||||
adam.aw2.m.write_npy(dir / "adam_aw2_m.npy"); adam.aw2.v.write_npy(dir / "adam_aw2_v.npy")
|
|
||||||
adam.ab2.m.write_npy(dir / "adam_ab2_m.npy"); adam.ab2.v.write_npy(dir / "adam_ab2_v.npy")
|
|
||||||
adam.aw3.m.write_npy(dir / "adam_aw3_m.npy"); adam.aw3.v.write_npy(dir / "adam_aw3_v.npy")
|
|
||||||
adam.ab3.m.write_npy(dir / "adam_ab3_m.npy"); adam.ab3.v.write_npy(dir / "adam_ab3_v.npy")
|
|
||||||
adam.cw1.m.write_npy(dir / "adam_cw1_m.npy"); adam.cw1.v.write_npy(dir / "adam_cw1_v.npy")
|
|
||||||
adam.cb1.m.write_npy(dir / "adam_cb1_m.npy"); adam.cb1.v.write_npy(dir / "adam_cb1_v.npy")
|
|
||||||
adam.cw2.m.write_npy(dir / "adam_cw2_m.npy"); adam.cw2.v.write_npy(dir / "adam_cw2_v.npy")
|
|
||||||
adam.cb2.m.write_npy(dir / "adam_cb2_m.npy"); adam.cb2.v.write_npy(dir / "adam_cb2_v.npy")
|
|
||||||
adam.cw3.m.write_npy(dir / "adam_cw3_m.npy"); adam.cw3.v.write_npy(dir / "adam_cw3_v.npy")
|
|
||||||
adam.cb3.m.write_npy(dir / "adam_cb3_m.npy"); adam.cb3.v.write_npy(dir / "adam_cb3_v.npy")
|
|
||||||
adam.logStd.m.write_npy(dir / "adam_logstd_m.npy")
|
|
||||||
adam.logStd.v.write_npy(dir / "adam_logstd_v.npy")
|
|
||||||
# t counters (all stepped in lockstep; store each for safety)
|
|
||||||
writeFile(dir / "adam_t.txt",
|
|
||||||
[$adam.aw1.t, $adam.ab1.t, $adam.aw2.t, $adam.ab2.t,
|
|
||||||
$adam.aw3.t, $adam.ab3.t, $adam.cw1.t, $adam.cb1.t,
|
|
||||||
$adam.cw2.t, $adam.cb2.t, $adam.cw3.t, $adam.cb3.t,
|
|
||||||
$adam.logStd.t].join("\n"))
|
|
||||||
|
|
||||||
proc loadAdamStates*(adam: var ACAdamStates, dir: string) =
|
|
||||||
## Load Adam m/v tensors and t counters from dir/. Called only when files exist.
|
|
||||||
## Shape mismatch (e.g. STATE_DIM changed) → keep zero-initialised state (safe fresh start).
|
|
||||||
template lm(dest: untyped, path: string) =
|
|
||||||
let loaded = read_npy[float32](path)
|
|
||||||
if loaded.shape == dest.shape:
|
|
||||||
dest = loaded
|
|
||||||
else:
|
|
||||||
echo "weights: Adam shape mismatch for " & path &
|
|
||||||
" (got " & $loaded.shape & " want " & $dest.shape & ") — resetting Adam state"
|
|
||||||
lm(adam.aw1.m, dir / "adam_aw1_m.npy"); lm(adam.aw1.v, dir / "adam_aw1_v.npy")
|
|
||||||
lm(adam.ab1.m, dir / "adam_ab1_m.npy"); lm(adam.ab1.v, dir / "adam_ab1_v.npy")
|
|
||||||
lm(adam.aw2.m, dir / "adam_aw2_m.npy"); lm(adam.aw2.v, dir / "adam_aw2_v.npy")
|
|
||||||
lm(adam.ab2.m, dir / "adam_ab2_m.npy"); lm(adam.ab2.v, dir / "adam_ab2_v.npy")
|
|
||||||
lm(adam.aw3.m, dir / "adam_aw3_m.npy"); lm(adam.aw3.v, dir / "adam_aw3_v.npy")
|
|
||||||
lm(adam.ab3.m, dir / "adam_ab3_m.npy"); lm(adam.ab3.v, dir / "adam_ab3_v.npy")
|
|
||||||
lm(adam.cw1.m, dir / "adam_cw1_m.npy"); lm(adam.cw1.v, dir / "adam_cw1_v.npy")
|
|
||||||
lm(adam.cb1.m, dir / "adam_cb1_m.npy"); lm(adam.cb1.v, dir / "adam_cb1_v.npy")
|
|
||||||
lm(adam.cw2.m, dir / "adam_cw2_m.npy"); lm(adam.cw2.v, dir / "adam_cw2_v.npy")
|
|
||||||
lm(adam.cb2.m, dir / "adam_cb2_m.npy"); lm(adam.cb2.v, dir / "adam_cb2_v.npy")
|
|
||||||
lm(adam.cw3.m, dir / "adam_cw3_m.npy"); lm(adam.cw3.v, dir / "adam_cw3_v.npy")
|
|
||||||
lm(adam.cb3.m, dir / "adam_cb3_m.npy"); lm(adam.cb3.v, dir / "adam_cb3_v.npy")
|
|
||||||
lm(adam.logStd.m, dir / "adam_logstd_m.npy"); lm(adam.logStd.v, dir / "adam_logstd_v.npy")
|
|
||||||
let ts = readFile(dir / "adam_t.txt").strip().splitLines()
|
|
||||||
if ts.len >= 13:
|
|
||||||
adam.aw1.t = parseInt(ts[0]); adam.ab1.t = parseInt(ts[1])
|
|
||||||
adam.aw2.t = parseInt(ts[2]); adam.ab2.t = parseInt(ts[3])
|
|
||||||
adam.aw3.t = parseInt(ts[4]); adam.ab3.t = parseInt(ts[5])
|
|
||||||
adam.cw1.t = parseInt(ts[6]); adam.cb1.t = parseInt(ts[7])
|
|
||||||
adam.cw2.t = parseInt(ts[8]); adam.cb2.t = parseInt(ts[9])
|
|
||||||
adam.cw3.t = parseInt(ts[10]); adam.cb3.t = parseInt(ts[11])
|
|
||||||
adam.logStd.t = parseInt(ts[12])
|
|
||||||
adam.initialized = true
|
|
||||||
|
|
||||||
proc adamStateFilesExist(dir: string): bool =
|
|
||||||
## Check that the minimum set of Adam files is present.
|
|
||||||
for f in adamMFiles:
|
|
||||||
if not fileExists(dir / f): return false
|
|
||||||
for f in adamVFiles:
|
|
||||||
if not fileExists(dir / f): return false
|
|
||||||
fileExists(dir / "adam_t.txt")
|
|
||||||
|
|
||||||
proc saveCheckpoint*(ac: ActorCritic, adam: ACAdamStates,
|
|
||||||
weightsRoot: string, roundNum: int) =
|
|
||||||
## Always saves weights + Adam state to weightsRoot/latest/.
|
|
||||||
## Every 50 rounds also saves to checkpoint_{1,2,3} in round-robin.
|
|
||||||
## Round counter is saved to weightsRoot/round_counter.txt (outside checkpoint dirs).
|
|
||||||
let latestDir = weightsRoot / "latest"
|
|
||||||
saveWeightsAtomic(ac, latestDir)
|
|
||||||
if adam.initialized:
|
|
||||||
saveAdamStates(adam, latestDir)
|
|
||||||
writeFile(weightsRoot / "round_counter.txt", $roundNum)
|
|
||||||
if roundNum mod 50 == 0:
|
|
||||||
let slot = ((roundNum div 50 - 1) mod 3) + 1 # 50→1, 100→2, 150→3, 200→1, …
|
|
||||||
let ckDir = weightsRoot / ("checkpoint_" & $slot)
|
|
||||||
saveWeightsAtomic(ac, ckDir)
|
|
||||||
if adam.initialized:
|
|
||||||
saveAdamStates(adam, ckDir)
|
|
||||||
|
|
||||||
proc loadBestAvailable*(ac: var ActorCritic, adam: var ACAdamStates,
|
|
||||||
weightsRoot: string): tuple[loaded: bool, roundNum: int] =
|
|
||||||
## Try latest/ first, then checkpoints sorted newest-first by mtime.
|
|
||||||
## Returns (true, roundNum) if weights loaded, (false, 0) if all fail.
|
|
||||||
## Adam state is loaded if present alongside weights; otherwise left uninitialised.
|
|
||||||
## Round counter is read from weightsRoot/round_counter.txt if present.
|
|
||||||
let checkpoints = [weightsRoot / "checkpoint_1",
|
|
||||||
weightsRoot / "checkpoint_2",
|
|
||||||
weightsRoot / "checkpoint_3"]
|
|
||||||
# Sort checkpoints newest-first by modification time
|
|
||||||
var existing: seq[tuple[mtime: Time, path: string]]
|
|
||||||
for p in checkpoints:
|
|
||||||
if dirExists(p):
|
|
||||||
existing.add((getLastModificationTime(p), p))
|
|
||||||
existing.sort(proc(a, b: tuple[mtime: Time, path: string]): int =
|
|
||||||
cmp(b.mtime, a.mtime)) # descending
|
|
||||||
|
|
||||||
let candidates = @[weightsRoot / "latest"] & existing.mapIt(it.path)
|
|
||||||
for candidate in candidates:
|
|
||||||
if dirExists(candidate):
|
|
||||||
var ok = true
|
|
||||||
for f in weightFiles:
|
|
||||||
if not fileExists(candidate / f):
|
|
||||||
ok = false
|
|
||||||
break
|
|
||||||
if ok:
|
|
||||||
ac.loadWeights(candidate)
|
|
||||||
if adamStateFilesExist(candidate):
|
|
||||||
adam.loadAdamStates(candidate)
|
|
||||||
let rcPath = weightsRoot / "round_counter.txt"
|
|
||||||
# Torn/empty file (e.g. after a crash) must not abort startup → treat as 0
|
|
||||||
let roundNum = if fileExists(rcPath):
|
|
||||||
try: parseInt(readFile(rcPath).strip())
|
|
||||||
except ValueError: 0
|
|
||||||
else: 0
|
|
||||||
return (loaded: true, roundNum: roundNum)
|
|
||||||
result = (loaded: false, roundNum: 0)
|
|
||||||
|
|
||||||
proc cleanStaleTempDirs*(weightsRoot: string) =
|
|
||||||
## Delete any dirs inside weightsRoot whose name contains "_tmp_".
|
|
||||||
if not dirExists(weightsRoot): return
|
|
||||||
for kind, path in walkDir(weightsRoot):
|
|
||||||
if kind == pcDir and "_tmp_" in lastPathPart(path):
|
|
||||||
removeDir(path)
|
|
||||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -1,13 +0,0 @@
|
|||||||
324564
|
|
||||||
324564
|
|
||||||
324564
|
|
||||||
324564
|
|
||||||
324564
|
|
||||||
324564
|
|
||||||
324564
|
|
||||||
324564
|
|
||||||
324564
|
|
||||||
324564
|
|
||||||
324564
|
|
||||||
324564
|
|
||||||
324564
|
|
||||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -1 +0,0 @@
|
|||||||
0
|
|
||||||
@@ -0,0 +1,46 @@
|
|||||||
|
# Neuroevolution gun with fixed-topology ANN evolved by GA
|
||||||
|
|
||||||
|
Evo_Bot needs a gun that adapts to each opponent's dodge patterns during a match, finds nonlinear movement patterns that histograms miss, and is original. We chose a fixed-topology feedforward ANN (91->8->1, 745 weights) whose weights are evolved by a mutation-only GA running on a parallel thread. This beats the alternatives (RL too slow to adapt in-match, Q-learning collapses to histogram for single-shot decisions, guess-factor histograms are unoriginal, transformer/LLM-style prediction is data-starved at ~14k ticks per match) while keeping implementation risk low by deferring topology evolution (NEAT) until the fixed network hits its ceiling.
|
||||||
|
|
||||||
|
## Considered Options
|
||||||
|
|
||||||
|
- **PPO / SAC (end-to-end RL):** Too slow -- thousands of rounds to converge, cannot adapt mid-match. Explored in other bots in this repo.
|
||||||
|
- **Q-learning for aiming:** Collapses to a histogram. Single-shot aiming has no sequential decision structure for Q-learning to exploit.
|
||||||
|
- **Guess-factor histogram:** Proven and fast to converge (~15 ticks), but unoriginal -- 20 years of community tuning.
|
||||||
|
- **Transformer / LLM-style sequence prediction:** Data-starved. ~14k ticks per match vs billions needed for attention-based models.
|
||||||
|
- **GA with crossover:** Literature uniformly shows crossover is harmful for ANN weight evolution -- it breaks co-adapted weight configurations. Every modern neuroevolution paper (Uber Deep GA, NRA, OpenAI ES) drops it.
|
||||||
|
- **CMA-ES:** Ideal at d=750 weights (self-adapts sigma and covariance). More complex to implement; upgrade path from simple GA when needed.
|
||||||
|
|
||||||
|
## Decision
|
||||||
|
|
||||||
|
**Architecture:**
|
||||||
|
- Evo_Bot (1v1) with modular gun interface: `feed(state)` / `aim() -> (angle, power)`
|
||||||
|
- Gun owns its evolution thread (parallel, never blocks inference)
|
||||||
|
|
||||||
|
**TOPO_Gun (first gun implementation):**
|
||||||
|
- Network: 91->8->1 (hidden size configurable), 745 weights
|
||||||
|
- Input: 30 ticks x (lateral_vel, delta_heading, wall_distance_ahead) + current distance = 91
|
||||||
|
- Output: guess factor (-1 to +1)
|
||||||
|
- Bullet power: deterministic distance-based formula (not learned)
|
||||||
|
- Evolution: population 300, clone loaded champion + small Gaussian mutations (ALL weights, sigma=0.005-0.01), NO crossover, single elite preserved unchanged
|
||||||
|
- Fitness: virtual bullet hits on 100 randomly sampled replay tape ticks, using real distance-based power
|
||||||
|
- Replay tape: rolling window ~2000 ticks (configurable)
|
||||||
|
- Weight persistence: per-opponent -> global fallback -> random init (load order)
|
||||||
|
- Cold start: first-ever run, don't fire until champion emerges; subsequent runs load weights, fire from tick 1
|
||||||
|
- Push new champion weights to inference when it hits better than current
|
||||||
|
|
||||||
|
**Deferred:**
|
||||||
|
- NEAT_Gun: deferred until TOPO_Gun hits its ceiling
|
||||||
|
- Virtual Guns: run multiple guns in parallel, fire whichever has best virtual hit rate
|
||||||
|
|
||||||
|
**Boundaries:**
|
||||||
|
- Bot controls firing discipline (when to shoot, energy management); gun always returns best aim
|
||||||
|
|
||||||
|
## Consequences
|
||||||
|
|
||||||
|
- GA+ANN finds nonlinear patterns histograms miss, but needs more data (~50+ ticks vs ~15 for guess-factor histogram) before it outperforms
|
||||||
|
- Fixed topology before NEAT reduces implementation risk
|
||||||
|
- Mutation-only evolution simplifies implementation (no crossover logic)
|
||||||
|
- Parallel evolution thread reuses the pattern from SAC_LSTM_Bot (#48)
|
||||||
|
- Per-opponent weight persistence eliminates cold start after first encounter
|
||||||
|
- CMA-ES is the natural upgrade path if simple GA convergence is too slow (d=750 is CMA-ES sweet spot)
|
||||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,173 @@
|
|||||||
|
# GA/ES Parameters for ~750-Weight Neuroevolution
|
||||||
|
|
||||||
|
Research for issue #65. Concrete parameter recommendations extracted from 13 papers in `docs/papers/neuroevolution/`.
|
||||||
|
|
||||||
|
## Context
|
||||||
|
|
||||||
|
Evo_Bot's TOPO_Gun: fixed-topology feedforward ANN, 745-1489 weights depending on layer sizes. Task is predicting enemy dodge behavior from a 30-tick sliding window, outputting a guess factor. Online evolution during Robocode matches.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 1. Population Size
|
||||||
|
|
||||||
|
**Recommendation: 64-256, not 300.**
|
||||||
|
|
||||||
|
| Source | Network size | Population | Notes |
|
||||||
|
|--------|-------------|------------|-------|
|
||||||
|
| Uber Deep GA (Such et al. 2017) | 4M params (Atari), 167k (Humanoid) | 1,000 | Massive networks, distributed; overkill for ~750 weights |
|
||||||
|
| OpenAI ES (Salimans et al. 2017) | 1.7M params | 720-1,440 workers | NES-style, not a population GA |
|
||||||
|
| Canonical ES (Chrabaszcz et al. 2018) | 1.7M params | 798 (lambda), mu=50 | mu=50 selected as best across games |
|
||||||
|
| World Models CMA-ES (Ha & Schmidhuber 2018) | 867-1,088 params (controller) | 64 | CMA-ES with 16 evals per individual |
|
||||||
|
| Challenges paper (Muller & Glasmachers 2018) | 1,352-2,349 weights | default CMA-ES lambda | LM-MA-ES for ~1k-2k weights |
|
||||||
|
| Evolving Generalists (Triebold & Yaman 2023) | 246-728 weights | xNES default: 4+floor(3*ln(d)) | For d=745 -> ~24; for d=1489 -> ~26 |
|
||||||
|
| NRA (Le Clei & Bellec 2022) | 3-322 params (dynamic) | 8-512 | 256+elitism best for ~300-param tasks; 16 enough for <100 params |
|
||||||
|
| Playing Atari 6 Neurons (Cuccu et al. 2018) | ~3k connections, 6-18 neurons | xNES default | Small networks, 100 generations sufficient |
|
||||||
|
|
||||||
|
**Key finding:** For ~750 weights, CMA-ES default lambda = 4+floor(3*ln(745)) = ~24 is a starting floor. The World Models paper (867-1,088 params, closest to our size) used 64 with CMA-ES and solved CarRacing. NRA used 256+elitism for ~300-param Pendulum networks. For a simple truncation-selection GA (not CMA-ES), 100-200 is a reasonable population; 300 is slightly wasteful but not harmful if evaluation is cheap.
|
||||||
|
|
||||||
|
**Verdict: Start with 100-200 for a truncation GA. If using CMA-ES or xNES, use their defaults (~24-26). 300 is too large for the weight count but acceptable if per-evaluation cost is low (Robocode rounds are fast).**
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 2. Mutation Rate and Distribution
|
||||||
|
|
||||||
|
**Recommendation: Additive Gaussian on ALL weights, sigma=0.002-0.02, not 5% at sigma=0.1.**
|
||||||
|
|
||||||
|
| Source | Mutation scheme | Notes |
|
||||||
|
|--------|----------------|-------|
|
||||||
|
| Uber Deep GA (Such et al. 2017) | theta' = theta + sigma * epsilon, epsilon ~ N(0,I). Sigma determined empirically per task. | Mutates ALL weights every generation, not a fraction. No "mutation rate" — every weight gets noise. |
|
||||||
|
| OpenAI ES (Salimans et al. 2017) | sigma = fixed hyperparameter (not adapted). Perturbation on full parameter vector. | Full-vector Gaussian perturbation, sigma tuned. |
|
||||||
|
| Canonical ES (Chrabaszcz et al. 2018) | N(0, sigma^2) added to all params. Network init from N(0, 0.05). | sigma is the step-size, adapted or fixed. |
|
||||||
|
| World Models (Ha & Schmidhuber 2018) | CMA-ES adapts sigma and full covariance matrix. | Self-adapting sigma — no manual sigma needed. |
|
||||||
|
| Challenges paper (Muller & Glasmachers 2018) | CSA (cumulative step-size adaptation) essential. Fixed sigma converges as slowly as random search. | Step-size adaptation is critical; fixed sigma is a known failure mode. |
|
||||||
|
| NRA (Le Clei & Bellec 2022) | N(0, 0.01) perturbation to all weights and biases. Top 50% selection. | sigma=0.01 for small networks. |
|
||||||
|
|
||||||
|
**Key finding:** No paper uses a "5% mutation rate" (mutating only 5% of weights). ALL papers mutate ALL weights simultaneously with small additive Gaussian noise. The "mutation rate" concept from traditional GAs (flip probability per gene) does not apply to real-valued neuroevolution. Instead, the noise magnitude (sigma) controls exploration.
|
||||||
|
|
||||||
|
For ~750 weights:
|
||||||
|
- NRA uses sigma=0.01 for networks up to ~300 params
|
||||||
|
- Uber GA uses sigma empirically per task (typical range 0.002-0.02 for Atari)
|
||||||
|
- CMA-ES/xNES adapt sigma automatically
|
||||||
|
|
||||||
|
**Verdict: Mutate ALL weights every generation. Use sigma=0.005-0.01 as starting point. If using CMA-ES, sigma self-adapts. The "5% of weights mutated" approach is non-standard and likely harmful — it under-explores the search space.**
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 3. Selection Pressure
|
||||||
|
|
||||||
|
**Recommendation: Top 10-50% (truncation) or top mu out of lambda. 20% is reasonable.**
|
||||||
|
|
||||||
|
| Source | Selection | Notes |
|
||||||
|
|--------|-----------|-------|
|
||||||
|
| Uber Deep GA (Such et al. 2017) | Truncation selection, top T individuals become parents. T not specified as percentage — varies. | Parents chosen uniformly at random from top T. |
|
||||||
|
| Canonical ES (Chrabaszcz et al. 2018) | Top mu=50 out of lambda=798 (~6%). Weighted mean of top mu. | mu=50 found optimal across games; tested mu in {10,20,50,100,200,400}. |
|
||||||
|
| NRA (Le Clei & Bellec 2022) | Top 50% duplicated, bottom 50% replaced. | Simple and effective for small populations. |
|
||||||
|
| Evolving Generalists (Triebold & Yaman 2023) | xNES default selection. | NES uses weighted rank-based update. |
|
||||||
|
|
||||||
|
**Key finding:** Selection pressure varies widely. Canonical ES uses ~6% (mu=50 out of 798). NRA uses 50%. Standard CMA-ES uses mu = lambda/2 (50%). Top 20% is in the middle range and is fine for a truncation GA.
|
||||||
|
|
||||||
|
**Verdict: 20% is reasonable. For small populations (64-100), 50% (top half) may work better. For larger populations (200+), stricter selection (10-20%) is appropriate. The Canonical ES result suggests mu=50 works well regardless of lambda for Atari-scale problems.**
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 4. Crossover
|
||||||
|
|
||||||
|
**Recommendation: No crossover. Mutation-only.**
|
||||||
|
|
||||||
|
| Source | Crossover? | Notes |
|
||||||
|
|--------|-----------|-------|
|
||||||
|
| Uber Deep GA (Such et al. 2017) | **No crossover.** "Historically, GAs often involve crossover, but for simplicity we did not include it." | Explicitly dropped crossover for DNN weights. |
|
||||||
|
| OpenAI ES (Salimans et al. 2017) | No crossover. | ES-style: mean update, not recombination of individuals. |
|
||||||
|
| NRA (Le Clei & Bellec 2022) | **No crossover.** "stripping down many mechanisms popular in traditional evolutionary methods, like agent crossover and speciation" | Crossover explicitly excluded. |
|
||||||
|
| CMA-ES/xNES | No crossover in the traditional sense. | Weighted recombination of top individuals into distribution mean — not pairwise crossover. |
|
||||||
|
| NEAT (Stanley & Miikkulainen 2011) | Has crossover via innovation numbers. | But NEAT is for topology evolution, not fixed-topology weight-only GA. |
|
||||||
|
|
||||||
|
**Key finding:** Every modern neuroevolution paper that works with fixed-topology networks drops crossover. Fogel & Stayton (1994, cited by Such et al.) showed crossover is often ineffective for simulated evolutionary optimization. For ANN weight vectors, crossover tends to be destructive because individual weights are not independent genes — they form functional units (layers, pathways) where mixing two different solutions creates non-functional hybrids.
|
||||||
|
|
||||||
|
**Verdict: No crossover. Mutation-only. Uniform crossover of ANN weights is harmful — it breaks co-adapted weight configurations. If recombination is desired, use CMA-ES/xNES-style weighted mean of top solutions, which is mathematically sound.**
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 5. Elitism
|
||||||
|
|
||||||
|
**Recommendation: Yes, keep top 1 unchanged (single elite).**
|
||||||
|
|
||||||
|
| Source | Elitism? | Notes |
|
||||||
|
|--------|---------|-------|
|
||||||
|
| Uber Deep GA (Such et al. 2017) | **Yes, 1 elite.** "The Nth individual is an unmodified copy of the best individual from the previous generation." Additionally, top 10 re-evaluated 30 times to find the true elite. | Single elite with robust re-evaluation. |
|
||||||
|
| NRA (Le Clei & Bellec 2022) | **Yes, elitism tested and beneficial.** Population sizes labeled "(elite)" consistently outperform non-elite variants in all figures. | Elitism was the single most impactful improvement for small populations. |
|
||||||
|
| CMA-ES | Elitist variants exist (mu+lambda). Standard CMA-ES is (mu,lambda) — non-elitist. | Non-elitist CMA-ES relies on distribution adaptation, not individual survival. |
|
||||||
|
|
||||||
|
**Key finding:** For simple truncation GAs, elitism (keeping top 1) prevents regression and is universally recommended. The NRA paper shows that adding elitism to even a population of 16 dramatically improves results. Uber's Deep GA uses elitism with robust re-evaluation (30 episodes to confirm the elite).
|
||||||
|
|
||||||
|
**Verdict: Keep top 1 elite unchanged. In noisy evaluation environments (Robocode), re-evaluate the top few candidates multiple times to find the true elite, following Uber's approach.**
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 6. Generations to Convergence
|
||||||
|
|
||||||
|
**Recommendation: 100-1,500 generations for ~750 weights, depending on the algorithm.**
|
||||||
|
|
||||||
|
| Source | Network size | Generations | Notes |
|
||||||
|
|--------|-------------|-------------|-------|
|
||||||
|
| Uber Deep GA (Such et al. 2017) | 4M params | 348-1,834 gens (at 1k pop) | Many games: best-in-run found in 1-29 gens |
|
||||||
|
| World Models CMA-ES (Ha & Schmidhuber 2018) | 867 params | ~1,800 gens | CMA-ES, pop=64, solved CarRacing |
|
||||||
|
| NRA (Le Clei & Bellec 2022) | 3-322 params (dynamic) | 100-5,000 gens | Simple tasks: <300 gens. Complex (Ant/Humanoid): 5,000+ |
|
||||||
|
| Playing Atari 6 Neurons (Cuccu et al. 2018) | ~3k connections | 100 gens | Extremely tight budget, still achieved competitive results |
|
||||||
|
| Challenges paper (Muller & Glasmachers 2018) | 1,352-2,349 weights | 100k-300k evals | LM-MA-ES, ~150k evals for bipedal walker convergence |
|
||||||
|
| Evolving Generalists (Triebold & Yaman 2023) | 728 weights (Ant) | 5,000 gens max | xNES, some tasks solved in <100 gens |
|
||||||
|
|
||||||
|
**Key finding:** For ~750 weights with a simple truncation GA (pop=100), expect 200-500 generations for a well-tuned sigma. CMA-ES/xNES may converge faster in generations but each generation is more expensive. The Challenges paper warns that halving the distance to the optimum requires O(d) samples, so for d=750, expect ~750 evaluations per halving step.
|
||||||
|
|
||||||
|
For Evo_Bot's online evolution during matches: each Robocode round can evaluate one individual. With 35-round matches (typical), ~5 generations of pop=7 per match, or ~2 generations of pop=15. Convergence within a single match is unlikely; evolution must persist across matches via weight persistence.
|
||||||
|
|
||||||
|
**Verdict: Budget 500-2,000 generations. With pop=100, that's 50k-200k evaluations. Online evolution will need many matches to converge — weight persistence is essential.**
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 7. CMA-ES vs Simple GA vs Tournament Selection
|
||||||
|
|
||||||
|
**Recommendation: CMA-ES or xNES for ~750 weights. Simple GA as a simpler fallback.**
|
||||||
|
|
||||||
|
| Algorithm | Sweet spot | Pros | Cons | Source |
|
||||||
|
|-----------|-----------|------|------|--------|
|
||||||
|
| **CMA-ES** | d <= 1,000 (ideal), up to ~2,000 (practical) | Self-adapts sigma and covariance, best convergence rate, handles ill-conditioned landscapes | O(d^2) memory/time per generation, needs O(d^2) evals for full covariance learning | Muller & Glasmachers 2018, Ha & Schmidhuber 2018 |
|
||||||
|
| **LM-MA-ES** | d = 1,000-10,000 | O(d) complexity, adapts fastest-evolving subspace, strong on ~2k weights | More complex to implement | Muller & Glasmachers 2018 |
|
||||||
|
| **xNES** | d <= 1,000 | Natural gradient, self-adapting, elegant | Similar scaling limits to CMA-ES | Cuccu et al. 2018, Triebold & Yaman 2023 |
|
||||||
|
| **Simple truncation GA** | Any d | Dead simple, trivially parallel, no internal state beyond population | Needs manual sigma tuning, no adaptation, converges slowly | Such et al. 2017, Le Clei & Bellec 2022 |
|
||||||
|
| **Canonical (mu,lambda)-ES** | Any d | Step-size adaptation via CSA, simple | mu tuning matters; mu=50 worked well in Chrabaszcz 2018 | Chrabaszcz et al. 2018 |
|
||||||
|
| **Tournament selection** | Traditional GA context | Tunable selection pressure | No advantage over truncation for ANN weights | Not specifically tested in any of the 13 papers |
|
||||||
|
|
||||||
|
**Key finding at d=750:** CMA-ES is in its sweet spot. The World Models paper (Ha & Schmidhuber 2018) used CMA-ES with pop=64 on 867-1,088 params and solved complex control tasks. The Evolving Generalists paper (Triebold & Yaman 2023) used xNES on 728 weights (Ant controller) with default population sizes. The Challenges paper (Muller & Glasmachers 2018) explicitly shows CMA-ES and LM-MA-ES outperforming simple ES on problems with 769-2,738 weights.
|
||||||
|
|
||||||
|
However, CMA-ES requires O(d^2) = O(560k) memory for the covariance matrix at d=750. This is manageable but not trivial for an online Robocode bot. A simpler option is a (mu,lambda)-ES with CSA for step-size adaptation.
|
||||||
|
|
||||||
|
**Verdict: CMA-ES or xNES is the best fit for 750 weights. If implementation complexity is a concern, a truncation GA with adaptive sigma (or even fixed sigma=0.005) is the pragmatic choice. Tournament selection offers no advantage.**
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Summary: Recommended Parameters for Evo_Bot TOPO_Gun
|
||||||
|
|
||||||
|
| Parameter | Current assumption | Recommendation | Rationale |
|
||||||
|
|-----------|-------------------|----------------|-----------|
|
||||||
|
| Population | 300 | 64-200 | 300 is oversized for ~750 weights; 64 (CMA-ES) to 200 (truncation GA) |
|
||||||
|
| Mutation | 5% of weights, gaussian sigma=0.1 | ALL weights, sigma=0.005-0.01 | Every paper mutates all weights. sigma=0.1 is too large. |
|
||||||
|
| Selection | Top 20% | Top 20-50% | 20% is fine; 50% if pop is small |
|
||||||
|
| Crossover | Uniform | None | Uniformly dropped in all modern neuroevolution papers |
|
||||||
|
| Elitism | Not specified | Top 1, re-evaluated | Single elite prevents regression; re-evaluate to handle noise |
|
||||||
|
| Algorithm | Simple GA | CMA-ES or truncation GA + CSA | CMA-ES is in its sweet spot at d=750; simple GA works but converges slower |
|
||||||
|
| Generations | Not specified | 500-2,000 | Online evolution needs many matches for convergence |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Sources
|
||||||
|
|
||||||
|
1. Such et al. 2017 — "Deep Neuroevolution: Genetic Algorithms are a Competitive Alternative" (Uber AI Labs)
|
||||||
|
2. Salimans et al. 2017 — "Evolution Strategies as a Scalable Alternative to Reinforcement Learning" (OpenAI)
|
||||||
|
3. Chrabaszcz et al. 2018 — "Back to Basics: Benchmarking Canonical Evolution Strategies for Playing Atari"
|
||||||
|
4. Muller & Glasmachers 2018 — "Challenges in High-dimensional Reinforcement Learning with Evolution Strategies"
|
||||||
|
5. Ha & Schmidhuber 2018 — "Recurrent World Models Facilitate Policy Evolution" (World Models)
|
||||||
|
6. Cuccu et al. 2018 — "Playing Atari with Six Neurons"
|
||||||
|
7. Le Clei & Bellec 2022 — "Neuroevolution of Recurrent Architectures on Control Tasks"
|
||||||
|
8. Triebold & Yaman 2023 — "Evolving Generalist Controllers to Handle a Wide Range of Morphological Variations"
|
||||||
|
9. Stanley & Miikkulainen 2011 — "Competitive Coevolution through Evolutionary Complexification" (NEAT)
|
||||||
@@ -1,224 +0,0 @@
|
|||||||
# Goto Controller Algorithm — Research
|
|
||||||
|
|
||||||
**Issue:** #20
|
|
||||||
**Branch:** research/goto-controller
|
|
||||||
**Date:** 2026-08-17
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Problem Statement
|
|
||||||
|
|
||||||
The PPO network will output a target position `(x, y)`. A goto controller must
|
|
||||||
translate that into per-tick `setTargetSpeed` and `setTurnRate` commands for
|
|
||||||
the Tank Royale Nim bot API.
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Codebase Findings
|
|
||||||
|
|
||||||
### Tank Royale Nim API — no built-in goto
|
|
||||||
|
|
||||||
The library (`tankroyale_botapi` v1.0.1) provides:
|
|
||||||
|
|
||||||
- `setTargetSpeed(speed: float)` — desired speed, clamped to ±8 units/tick.
|
|
||||||
Server auto-manages acceleration/deceleration via `getNewTargetSpeed`.
|
|
||||||
- `setTurnRate(rate: float)` — desired turn rate, clamped to `calcMaxTurnRate(speed) = 10 - 0.75 * abs(speed)`.
|
|
||||||
- `setForward(distance)` / `setBack(distance)` — blocking helpers that use
|
|
||||||
`gDistanceRemaining` + the server's deceleration model. These are blocking
|
|
||||||
(call `go()` internally) and therefore cannot be used in the non-blocking
|
|
||||||
per-tick run loop used by PPO_Bot.
|
|
||||||
|
|
||||||
There is **no built-in `setDistanceRemaining`-style goto**. The controller must
|
|
||||||
be written from scratch.
|
|
||||||
|
|
||||||
### Physics constants (from `constants.nim` / `utils.nim`)
|
|
||||||
|
|
||||||
| Constant | Value |
|
|
||||||
|---|---|
|
|
||||||
| Max speed | 8 units/tick |
|
|
||||||
| Acceleration | +1 unit/tick² |
|
|
||||||
| Deceleration | −2 units/tick² (braking is twice as fast) |
|
|
||||||
| Max turn rate | `10 − 0.75 × |speed|` deg/tick |
|
|
||||||
| Min turn rate (at max speed) | `10 − 0.75 × 8 = 4` deg/tick |
|
|
||||||
| `getNewTargetSpeed(maxSpeed, speed, dist)` | already implemented in utils.nim |
|
|
||||||
|
|
||||||
Key implication: **you can turn faster while slow**. Turn-then-drive lets the
|
|
||||||
bot use full 10°/tick turn rate, but wastes ticks stopped. Driving-while-turning
|
|
||||||
is smooth but limited to 4°/tick at top speed.
|
|
||||||
|
|
||||||
### Coordinate system
|
|
||||||
|
|
||||||
North = 0°, clockwise. `directionTo` in `utils.nim` returns a bearing in
|
|
||||||
`[0, 360)`. `bearingTo` returns a signed relative bearing in `(-180, 180]`.
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Approaches Considered
|
|
||||||
|
|
||||||
### A — Turn-then-drive (sequential)
|
|
||||||
|
|
||||||
Stop → turn to face target → drive full speed → brake.
|
|
||||||
|
|
||||||
- Simple to implement.
|
|
||||||
- Very slow: wastes ticks turning at zero speed then decelerating.
|
|
||||||
- Produces jerky, non-smooth movement — bad as a controller layer.
|
|
||||||
|
|
||||||
### B — Proportional navigation (continuous per-tick)
|
|
||||||
|
|
||||||
Each tick: compute bearing to target, set turn rate proportional to bearing
|
|
||||||
error, set speed based on distance remaining.
|
|
||||||
|
|
||||||
- Standard Robocode idiom. Very common in published bots.
|
|
||||||
- Does not make the forward-vs-reverse decision optimally.
|
|
||||||
- Can overshoot if gains are too high; can be sluggish if too low.
|
|
||||||
|
|
||||||
### C — Arc/pursuit steering (proportional + speed-dependent turn limit)
|
|
||||||
|
|
||||||
Like B, but explicitly clamps turn rate to `calcMaxTurnRate(currentSpeed)` and
|
|
||||||
scales speed down when the heading error is large (so the bot slows to increase
|
|
||||||
turn authority).
|
|
||||||
|
|
||||||
- Handles Tank Royale's speed-dependent turn rate correctly.
|
|
||||||
- Naturally smooth.
|
|
||||||
- Still needs explicit forward/reverse decision.
|
|
||||||
|
|
||||||
### D — Forward-vs-reverse decision + proportional steering (recommended)
|
|
||||||
|
|
||||||
Extend C with the classic Robocode "should I go backward?" heuristic:
|
|
||||||
if `|bearingError| > 90°`, it is faster to reverse and face the target with
|
|
||||||
the rear than to turn more than 90° forward. Flip target speed sign and add
|
|
||||||
180° to the bearing before computing turn rate.
|
|
||||||
|
|
||||||
This is the approach used by high-quality Robocode 1 bots (e.g. RaikoMX,
|
|
||||||
Aristocles) and it trivially maps to Tank Royale's API.
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Recommended Algorithm
|
|
||||||
|
|
||||||
### Decision: forward or reverse?
|
|
||||||
|
|
||||||
```
|
|
||||||
bearing = normalizeRelativeAngle(directionTo(x, y) - direction)
|
|
||||||
if abs(bearing) > 90.0:
|
|
||||||
# Going backward is cheaper
|
|
||||||
direction_sign = -1
|
|
||||||
effective_bearing = normalizeRelativeAngle(bearing + 180.0)
|
|
||||||
else:
|
|
||||||
direction_sign = +1
|
|
||||||
effective_bearing = bearing
|
|
||||||
```
|
|
||||||
|
|
||||||
### Turn rate
|
|
||||||
|
|
||||||
Apply full proportional turn rate toward the effective bearing:
|
|
||||||
|
|
||||||
```
|
|
||||||
max_turn = 10.0 - 0.75 * abs(currentSpeed)
|
|
||||||
turnRate = clamp(effective_bearing, -max_turn, max_turn)
|
|
||||||
```
|
|
||||||
|
|
||||||
`effective_bearing` acts as both direction and magnitude: if the error is
|
|
||||||
small, the turn rate is small (smooth approach); if large, it clamps to max
|
|
||||||
(fastest possible turn).
|
|
||||||
|
|
||||||
### Target speed
|
|
||||||
|
|
||||||
Use `getNewTargetSpeed` (already in `utils.nim`) to determine the speed
|
|
||||||
that will arrive at the target with zero velocity:
|
|
||||||
|
|
||||||
```
|
|
||||||
dist = distanceTo(x, y)
|
|
||||||
raw_speed = getNewTargetSpeed(MAX_SPEED, currentSpeed, dist)
|
|
||||||
targetSpeed = direction_sign * raw_speed
|
|
||||||
```
|
|
||||||
|
|
||||||
This reuses the exact deceleration model the server uses, so the bot always
|
|
||||||
brakes at the right time with no overshoot.
|
|
||||||
|
|
||||||
### Stop condition
|
|
||||||
|
|
||||||
```
|
|
||||||
if dist < ARRIVAL_THRESHOLD: # e.g. 18.0 (= BOT_RADIUS)
|
|
||||||
targetSpeed = 0.0
|
|
||||||
turnRate = 0.0
|
|
||||||
```
|
|
||||||
|
|
||||||
### Full pseudocode (one tick)
|
|
||||||
|
|
||||||
```nim
|
|
||||||
proc gotoTick*(tx, ty, x, y, direction, currentSpeed: float):
|
|
||||||
tuple[targetSpeed, turnRate: float] =
|
|
||||||
|
|
||||||
let dist = distanceTo(x, y, tx, ty)
|
|
||||||
|
|
||||||
if dist < ARRIVAL_THRESHOLD:
|
|
||||||
return (0.0, 0.0)
|
|
||||||
|
|
||||||
let rawBearing = normalizeRelativeAngle(directionTo(x, y, tx, ty) - direction)
|
|
||||||
|
|
||||||
let (dirSign, effBearing) =
|
|
||||||
if abs(rawBearing) > 90.0:
|
|
||||||
(-1.0, normalizeRelativeAngle(rawBearing + 180.0))
|
|
||||||
else:
|
|
||||||
(1.0, rawBearing)
|
|
||||||
|
|
||||||
let maxTurn = 10.0 - 0.75 * abs(currentSpeed)
|
|
||||||
let turnRate = effBearing.clamp(-maxTurn, maxTurn)
|
|
||||||
|
|
||||||
let rawSpeed = getNewTargetSpeed(MAX_SPEED, abs(currentSpeed), dist)
|
|
||||||
let targetSpeed = dirSign * rawSpeed
|
|
||||||
|
|
||||||
return (targetSpeed, turnRate)
|
|
||||||
```
|
|
||||||
|
|
||||||
Call once per tick from the `run` loop, pass results to `setTargetSpeed` /
|
|
||||||
`setTurnRate`.
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Why not pure proportional navigation (option B)?
|
|
||||||
|
|
||||||
Option B without the speed-dependent turn clamp will attempt to command more
|
|
||||||
turn rate than the server will honor at high speed — it does the right thing
|
|
||||||
emergently but wastes the gap. Explicitly scaling turn rate with
|
|
||||||
`calcMaxTurnRate(speed)` is more intentional and matches the physics exactly.
|
|
||||||
This is already coded in `actions.nim` (`r1 * (10.0 - 0.75 * abs(currentSpeed))`),
|
|
||||||
so the pattern is established in the codebase.
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Why reuse `getNewTargetSpeed` from utils.nim?
|
|
||||||
|
|
||||||
It already encodes the exact asymmetric acceleration/deceleration model
|
|
||||||
(accel +1, decel −2 per tick). Reimplementing distance-based speed management
|
|
||||||
from scratch would duplicate this and risk drift. Import it directly.
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Forward/Reverse optimality
|
|
||||||
|
|
||||||
The 90° threshold is the exact breakeven point:
|
|
||||||
|
|
||||||
- Turning 91° forward takes ≥10 ticks at slow speed + travel time.
|
|
||||||
- Reversing 89° (i.e. 180−91=89° effective turn) takes fewer ticks total
|
|
||||||
for any distance large enough to matter.
|
|
||||||
- For very short distances (< ~36 units) the bot will decelerate before the
|
|
||||||
turn completes anyway; the threshold still works because the speed penalty
|
|
||||||
applies equally to both cases.
|
|
||||||
|
|
||||||
For a controller layer that feeds a neural network's goto target, sub-optimal
|
|
||||||
behavior on very short hops is acceptable — the network will learn to avoid
|
|
||||||
issuing tiny hops.
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Sources / References
|
|
||||||
|
|
||||||
- Tank Royale Nim API source: `tankroyale_botapi/utils.nim`, `bot.nim`,
|
|
||||||
`constants.nim` (v1.0.1, installed at `~/.nimble/pkgs2/`).
|
|
||||||
- Robocode wiki — "Proportional navigation" and "Should I go backward?"
|
|
||||||
heuristic: widely documented in the Robocode community (e.g. RoboWiki
|
|
||||||
`BasicSurfer`, `RaikoMX` source).
|
|
||||||
- Tank Royale physics spec: confirmed against `ACCELERATION = 1.0`,
|
|
||||||
`ABS_DECELERATION = 2.0` in `constants.nim`.
|
|
||||||
@@ -1,258 +0,0 @@
|
|||||||
## Main entry-point module for Robocode Tank Royale Nim bot API.
|
|
||||||
##
|
|
||||||
## Usage:
|
|
||||||
## import tankroyale_botapi
|
|
||||||
##
|
|
||||||
## type MyBot = ref object of Bot
|
|
||||||
## method run(bot: MyBot) =
|
|
||||||
## forward(100)
|
|
||||||
## ...
|
|
||||||
##
|
|
||||||
## var bot = MyBot()
|
|
||||||
## start(bot, "MyBot.json")
|
|
||||||
|
|
||||||
import std/[os, json]
|
|
||||||
|
|
||||||
import ./tankroyale_botapi/constants
|
|
||||||
import ./tankroyale_botapi/color
|
|
||||||
import ./tankroyale_botapi/schemas
|
|
||||||
import ./tankroyale_botapi/utils
|
|
||||||
import ./tankroyale_botapi/bot_info
|
|
||||||
import ./tankroyale_botapi/ws_client
|
|
||||||
import ./tankroyale_botapi/json_parse
|
|
||||||
import ./tankroyale_botapi/event_queue
|
|
||||||
import ./tankroyale_botapi/bot
|
|
||||||
import ./tankroyale_botapi/graphics
|
|
||||||
|
|
||||||
export constants
|
|
||||||
export color
|
|
||||||
export schemas
|
|
||||||
export utils
|
|
||||||
export bot_info
|
|
||||||
export json_parse
|
|
||||||
export event_queue
|
|
||||||
export bot
|
|
||||||
export graphics
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# WebSocket receive loop (main thread)
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
proc handleServerHandshake(ws: SyncWebSocket; node: JsonNode; info: BotInfo; secret: string) =
|
|
||||||
let sessionId = node{"sessionId"}.getStr
|
|
||||||
setServerInfo(node{"variant"}.getStr, node{"version"}.getStr)
|
|
||||||
|
|
||||||
# Build bot handshake
|
|
||||||
var h = newJObject()
|
|
||||||
h["type"] = %"BotHandshake"
|
|
||||||
h["sessionId"] = %sessionId
|
|
||||||
h["name"] = %info.name
|
|
||||||
h["version"] = %info.version
|
|
||||||
h["authors"] = %info.authors
|
|
||||||
h["description"] = %info.description
|
|
||||||
h["homepage"] = %info.homepage
|
|
||||||
h["countryCodes"] = %info.countryCodes
|
|
||||||
h["gameTypes"] = %info.gameTypes
|
|
||||||
h["platform"] = %info.platform
|
|
||||||
h["programmingLang"]= %info.programmingLang
|
|
||||||
h["isDroid"] = %info.isDroid
|
|
||||||
if secret.len > 0:
|
|
||||||
h["secret"] = %secret
|
|
||||||
let ip = info.initialPosition
|
|
||||||
if ip.x != 0.0 or ip.y != 0.0 or ip.direction != 0.0:
|
|
||||||
var ipObj = newJObject()
|
|
||||||
if ip.x != 0.0: ipObj["x"] = %ip.x
|
|
||||||
if ip.y != 0.0: ipObj["y"] = %ip.y
|
|
||||||
if ip.direction != 0.0: ipObj["direction"] = %ip.direction
|
|
||||||
h["initialPosition"] = ipObj
|
|
||||||
ws.send($h)
|
|
||||||
|
|
||||||
proc parseGameSetup(node: JsonNode): GameSetup =
|
|
||||||
if node.isNil: return
|
|
||||||
result.gameType = node{"gameType"}.getStr("classic")
|
|
||||||
result.arenaWidth = node{"arenaWidth"}.getInt(800)
|
|
||||||
result.isArenaWidthLocked = node{"isArenaWidthLocked"}.getBool(false)
|
|
||||||
result.arenaHeight = node{"arenaHeight"}.getInt(600)
|
|
||||||
result.isArenaHeightLocked = node{"isArenaHeightLocked"}.getBool(false)
|
|
||||||
result.numberOfRounds = node{"numberOfRounds"}.getInt(10)
|
|
||||||
result.isNumberOfRoundsLocked = node{"isNumberOfRoundsLocked"}.getBool(false)
|
|
||||||
result.minNumberOfParticipants = node{"minNumberOfParticipants"}.getInt(2)
|
|
||||||
result.isMinNumberOfParticipantsLocked = node{"isMinNumberOfParticipantsLocked"}.getBool(false)
|
|
||||||
result.maxNumberOfParticipants = node{"maxNumberOfParticipants"}.getInt(10)
|
|
||||||
result.isMaxNumberOfParticipantsLocked = node{"isMaxNumberOfParticipantsLocked"}.getBool(false)
|
|
||||||
result.gunCoolingRate = node{"gunCoolingRate"}.getFloat(0.1)
|
|
||||||
result.isGunCoolingRateLocked = node{"isGunCoolingRateLocked"}.getBool(false)
|
|
||||||
result.maxInactivityTurns = node{"maxInactivityTurns"}.getInt(450)
|
|
||||||
result.isMaxInactivityTurnsLocked = node{"isMaxInactivityTurnsLocked"}.getBool(false)
|
|
||||||
result.turnTimeout = node{"turnTimeout"}.getInt(30000)
|
|
||||||
result.isTurnTimeoutLocked = node{"isTurnTimeoutLocked"}.getBool(false)
|
|
||||||
result.readyTimeout = node{"readyTimeout"}.getInt(1000000)
|
|
||||||
result.isReadyTimeoutLocked = node{"isReadyTimeoutLocked"}.getBool(false)
|
|
||||||
result.defaultTurnsPerSecond = node{"defaultTurnsPerSecond"}.getInt(30)
|
|
||||||
|
|
||||||
proc handleGameStarted(ws: SyncWebSocket; node: JsonNode) =
|
|
||||||
let setup = parseGameSetup(node{"gameSetup"})
|
|
||||||
|
|
||||||
var teammateIds: seq[int] = @[]
|
|
||||||
if not node{"teammateIds"}.isNil and node["teammateIds"].kind == JArray:
|
|
||||||
for id in node["teammateIds"]: teammateIds.add id.getInt
|
|
||||||
|
|
||||||
let myId = node{"myId"}.getInt
|
|
||||||
setGameStarted(myId, setup, teammateIds)
|
|
||||||
|
|
||||||
# Build event object manually — GameStartedEventForBot has no turnNumber
|
|
||||||
let e = GameStartedEventForBot(
|
|
||||||
`type`: "GameStartedEventForBot",
|
|
||||||
myId: myId,
|
|
||||||
startX: node{"startX"}.getFloat(0.0),
|
|
||||||
startY: node{"startY"}.getFloat(0.0),
|
|
||||||
startDirection: node{"startDirection"}.getFloat(0.0),
|
|
||||||
teammateIds: teammateIds,
|
|
||||||
gameSetup: setup
|
|
||||||
)
|
|
||||||
gBot.onGameStarted(e)
|
|
||||||
|
|
||||||
# Send BotReady
|
|
||||||
ws.send("""{"type":"BotReady"}""")
|
|
||||||
|
|
||||||
proc handleTick(node: JsonNode) =
|
|
||||||
# Build TickEventForBot manually to handle optional fields safely
|
|
||||||
var tick: TickEventForBot
|
|
||||||
tick.`type` = "TickEventForBot"
|
|
||||||
tick.turnNumber = node{"turnNumber"}.getInt(0)
|
|
||||||
tick.roundNumber = node{"roundNumber"}.getInt(0)
|
|
||||||
tick.botState = parseBotState(node{"botState"})
|
|
||||||
tick.bulletStates = @[]
|
|
||||||
if not node{"bulletStates"}.isNil and node["bulletStates"].kind == JArray:
|
|
||||||
for bs in node["bulletStates"]:
|
|
||||||
tick.bulletStates.add parseBulletState(bs)
|
|
||||||
tick.events = @[] # sub-events parsed separately into typed BotEvent
|
|
||||||
|
|
||||||
# Parse embedded events into typed BotEvent for priority-based dispatch
|
|
||||||
var events: seq[BotEvent] = @[]
|
|
||||||
let myId = getMyId()
|
|
||||||
if node.hasKey("events") and node["events"].kind == JArray:
|
|
||||||
for ev in node["events"]:
|
|
||||||
events.add parseBotEvent(ev, myId)
|
|
||||||
|
|
||||||
signalTick(tick, events) # update shared state
|
|
||||||
processTickOnMainThread() # motion tracking (while bot is blocked)
|
|
||||||
wakeBotThread() # wake bot — state + motion ready
|
|
||||||
|
|
||||||
proc runReceiveLoop*(ws: SyncWebSocket; info: BotInfo; secret: string; serverUrl: string) =
|
|
||||||
## Main WebSocket receive loop. Blocks until disconnected.
|
|
||||||
while ws.connected:
|
|
||||||
var msg: string
|
|
||||||
try:
|
|
||||||
msg = ws.receive()
|
|
||||||
except Exception as e:
|
|
||||||
stderr.writeLine "[ws] receive error: " & e.msg
|
|
||||||
gBot.onConnectionError(ConnectionErrorEvent(serverUrl: serverUrl, error: e.msg))
|
|
||||||
break
|
|
||||||
|
|
||||||
if msg.len == 0:
|
|
||||||
break # connection closed
|
|
||||||
|
|
||||||
var node: JsonNode
|
|
||||||
try:
|
|
||||||
node = parseJson(msg)
|
|
||||||
except Exception as e:
|
|
||||||
stderr.writeLine "[ws] json parse error: " & e.msg
|
|
||||||
continue
|
|
||||||
|
|
||||||
let msgType = node{"type"}.getStr
|
|
||||||
try:
|
|
||||||
case msgType
|
|
||||||
of "ServerHandshake":
|
|
||||||
handleServerHandshake(ws, node, info, secret)
|
|
||||||
of "GameStartedEventForBot":
|
|
||||||
handleGameStarted(ws, node)
|
|
||||||
of "RoundStartedEvent":
|
|
||||||
let e = node.to(RoundStartedEvent)
|
|
||||||
debugLog("[NS-ENTER] round=" & $e.roundNumber & " tid=" & $getThreadId())
|
|
||||||
# Start (or restart) the bot thread each round
|
|
||||||
startRound()
|
|
||||||
startBotThread()
|
|
||||||
gBot.onRoundStarted(e)
|
|
||||||
debugLog("[NS-EXIT] round=" & $e.roundNumber & " tid=" & $getThreadId())
|
|
||||||
of "TickEventForBot":
|
|
||||||
handleTick(node)
|
|
||||||
of "RoundEndedEventForBot":
|
|
||||||
setRunning(false)
|
|
||||||
let e = node.to(RoundEndedEventForBot)
|
|
||||||
debugLog("[RE-ENTER] round=" & $e.roundNumber & " tid=" & $getThreadId())
|
|
||||||
signalStop() # unblock bot thread blocked in go()
|
|
||||||
debugLog("[WT-ENTER] round=" & $e.roundNumber & " tid=" & $getThreadId())
|
|
||||||
waitForBotThread()
|
|
||||||
debugLog("[WT-EXIT] round=" & $e.roundNumber & " tid=" & $getThreadId())
|
|
||||||
debugLog("[DR-ENTER] round=" & $e.roundNumber & " tid=" & $getThreadId())
|
|
||||||
drainTickChan() # drain stop signal if bot exited via isRunning() check
|
|
||||||
drainIntentChan() # drain AFTER thread joined — no more writes possible
|
|
||||||
drainEventChan() # drop any unconsumed tick events (stale into next round)
|
|
||||||
debugLog("[DR-EXIT] round=" & $e.roundNumber & " tid=" & $getThreadId())
|
|
||||||
debugLog("[ONRE-ENTER] round=" & $e.roundNumber & " tid=" & $getThreadId())
|
|
||||||
gBot.onRoundEnded(e)
|
|
||||||
debugLog("[ONRE-EXIT] round=" & $e.roundNumber & " tid=" & $getThreadId())
|
|
||||||
of "GameEndedEventForBot":
|
|
||||||
setRunning(false)
|
|
||||||
let e = node.to(GameEndedEventForBot)
|
|
||||||
drainIntentChan()
|
|
||||||
gBot.onGameEnded(e)
|
|
||||||
of "GameAbortedEvent":
|
|
||||||
setRunning(false)
|
|
||||||
signalStop() # unblock bot thread (game aborted mid-round)
|
|
||||||
waitForBotThread()
|
|
||||||
drainTickChan() # drain stop signal if bot exited via isRunning() check
|
|
||||||
drainIntentChan() # drain AFTER thread joined — no more writes possible
|
|
||||||
drainEventChan() # drop any unconsumed tick events (stale into next round)
|
|
||||||
gBot.onGameAborted()
|
|
||||||
of "SkippedTurnEvent":
|
|
||||||
let e = node.to(SkippedTurnEvent)
|
|
||||||
gBot.onSkippedTurn(e)
|
|
||||||
else:
|
|
||||||
discard # unknown message type — ignore
|
|
||||||
except Exception as e:
|
|
||||||
# A raised handler/callback (e.g. an OSError from sync training inside
|
|
||||||
# onRoundEnded) must not kill the receive loop — that is the silent
|
|
||||||
# corpse path (process lives, no intents ever again). Log and continue.
|
|
||||||
stderr.writeLine "[ws] handler error (" & msgType & "): " & e.msg
|
|
||||||
debugLog("[WS-HANDLER-ERR] " & msgType & ": " & e.msg)
|
|
||||||
|
|
||||||
# Loop exited: server disconnected or ws error. Make sure the bot thread is
|
|
||||||
# stopped and joined so the process exits cleanly instead of hanging forever
|
|
||||||
# with a blocked bot (corpse). ponytail: signalStop + join; the bot's go()
|
|
||||||
# consumes the stop as a non-tick and exits via the isRunning() check.
|
|
||||||
if isRunning():
|
|
||||||
debugLog("[DBG] receive loop exited while running — stopping bot thread")
|
|
||||||
setRunning(false)
|
|
||||||
signalStop()
|
|
||||||
waitForBotThread()
|
|
||||||
drainTickChan()
|
|
||||||
drainIntentChan()
|
|
||||||
drainEventChan()
|
|
||||||
gBot.onDisconnected(DisconnectedEvent(serverUrl: serverUrl))
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# Public start() procedure
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
proc start*(bot: Bot; jsonFile: string = "") =
|
|
||||||
## Connect to the server and start the bot.
|
|
||||||
## jsonFile: path to bot JSON profile (optional; falls back to env vars).
|
|
||||||
gBot = bot
|
|
||||||
gBotInfo = loadBotInfo(jsonFile)
|
|
||||||
initGlobals()
|
|
||||||
|
|
||||||
let serverUrl = getEnv("SERVER_URL", "ws://localhost:7654")
|
|
||||||
let serverSecret = getEnv("SERVER_SECRET", "")
|
|
||||||
|
|
||||||
try:
|
|
||||||
gWs = newSyncWebSocket(serverUrl)
|
|
||||||
except Exception as e:
|
|
||||||
stderr.writeLine "[start] Cannot connect to " & serverUrl & ": " & e.msg
|
|
||||||
quit(1)
|
|
||||||
|
|
||||||
bot.onConnected(ConnectedEvent(serverUrl: serverUrl))
|
|
||||||
startSenderThread()
|
|
||||||
runReceiveLoop(gWs, gBotInfo, serverSecret, serverUrl)
|
|
||||||
stopSenderThread()
|
|
||||||
@@ -1,11 +0,0 @@
|
|||||||
# Package
|
|
||||||
version = "1.0.1"
|
|
||||||
author = "Davide Cappellini"
|
|
||||||
description = "Nim bot API for Robocode Tank Royale"
|
|
||||||
license = "Apache-2.0"
|
|
||||||
srcDir = "src"
|
|
||||||
skipDirs = @["sample_bots"]
|
|
||||||
|
|
||||||
# Dependencies
|
|
||||||
requires "nim >= 2.0.0"
|
|
||||||
requires "jsony >= 1.1.5"
|
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -1,79 +0,0 @@
|
|||||||
## BotInfo: bot identification loaded from a JSON file or environment variables.
|
|
||||||
|
|
||||||
import std/[os, json, strutils, sequtils]
|
|
||||||
import ./schemas
|
|
||||||
|
|
||||||
type
|
|
||||||
BotInfo* = object
|
|
||||||
name*: string
|
|
||||||
version*: string
|
|
||||||
authors*: seq[string]
|
|
||||||
description*: string
|
|
||||||
homepage*: string
|
|
||||||
countryCodes*: seq[string]
|
|
||||||
gameTypes*: seq[string]
|
|
||||||
platform*: string
|
|
||||||
programmingLang*: string
|
|
||||||
initialPosition*: InitialPosition
|
|
||||||
isDroid*: bool
|
|
||||||
|
|
||||||
proc botInfoFromJson*(path: string): BotInfo =
|
|
||||||
let data = parseJson(readFile(path))
|
|
||||||
result.name = data{"name"}.getStr
|
|
||||||
result.version = data{"version"}.getStr
|
|
||||||
if data.hasKey("authors"):
|
|
||||||
for a in data["authors"]: result.authors.add a.getStr
|
|
||||||
result.description = data{"description"}.getStr
|
|
||||||
result.homepage = data{"homepage"}.getStr
|
|
||||||
if data.hasKey("countryCodes"):
|
|
||||||
for c in data["countryCodes"]: result.countryCodes.add c.getStr
|
|
||||||
if data.hasKey("gameTypes"):
|
|
||||||
for g in data["gameTypes"]: result.gameTypes.add g.getStr
|
|
||||||
result.platform = data{"platform"}.getStr("Nim " & NimVersion)
|
|
||||||
result.programmingLang = data{"programmingLang"}.getStr("Nim")
|
|
||||||
if data.hasKey("initialPosition"):
|
|
||||||
let ip = data["initialPosition"]
|
|
||||||
result.initialPosition.x = ip{"x"}.getFloat
|
|
||||||
result.initialPosition.y = ip{"y"}.getFloat
|
|
||||||
result.initialPosition.direction = ip{"direction"}.getFloat
|
|
||||||
result.isDroid = data{"isDroid"}.getBool(false)
|
|
||||||
|
|
||||||
proc botInfoFromEnv*(): BotInfo =
|
|
||||||
## Fall back to environment variables when no JSON file is given.
|
|
||||||
result.name = getEnv("BOT_NAME", "Unnamed Bot")
|
|
||||||
result.version = getEnv("BOT_VERSION", "1.0")
|
|
||||||
let authorsStr = getEnv("BOT_AUTHORS", "Unknown")
|
|
||||||
result.authors = authorsStr.split(',').mapIt(it.strip)
|
|
||||||
result.description = getEnv("BOT_DESCRIPTION", "")
|
|
||||||
result.homepage = getEnv("BOT_HOMEPAGE", "")
|
|
||||||
let ccStr = getEnv("BOT_COUNTRY_CODES", "")
|
|
||||||
if ccStr.len > 0:
|
|
||||||
result.countryCodes = ccStr.split(',').mapIt(it.strip)
|
|
||||||
let gtStr = getEnv("BOT_GAME_TYPES", "classic,melee,1v1")
|
|
||||||
result.gameTypes = gtStr.split(',').mapIt(it.strip)
|
|
||||||
result.platform = getEnv("BOT_PLATFORM", "Nim " & NimVersion)
|
|
||||||
result.programmingLang = getEnv("BOT_PROGRAMMING_LANG", "Nim")
|
|
||||||
result.isDroid = getEnv("BOT_IS_DROID", "false").toLowerAscii == "true"
|
|
||||||
|
|
||||||
proc loadBotInfo*(jsonFile: string = ""): BotInfo =
|
|
||||||
var resolved = ""
|
|
||||||
if jsonFile.len > 0:
|
|
||||||
if fileExists(jsonFile):
|
|
||||||
resolved = jsonFile
|
|
||||||
else:
|
|
||||||
# Try alongside the executable
|
|
||||||
let appPath = getAppDir() / jsonFile
|
|
||||||
if fileExists(appPath):
|
|
||||||
resolved = appPath
|
|
||||||
|
|
||||||
if resolved.len > 0:
|
|
||||||
result = botInfoFromJson(resolved)
|
|
||||||
else:
|
|
||||||
result = botInfoFromEnv()
|
|
||||||
# Ensure gameTypes has at least one entry
|
|
||||||
if result.gameTypes.len == 0:
|
|
||||||
result.gameTypes = @["classic", "melee", "1v1"]
|
|
||||||
if result.platform.len == 0:
|
|
||||||
result.platform = "Nim " & NimVersion
|
|
||||||
if result.programmingLang.len == 0:
|
|
||||||
result.programmingLang = "Nim"
|
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user