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,371 +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}"
|
|
||||||
|
|
||||||
writeFile(weightsRoot / "round_counter.txt", $roundCounter)
|
|
||||||
|
|
||||||
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:
|
|
||||||
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, deterministic = evalOnly)
|
|
||||||
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,86 +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], deterministic = false): tuple[actions: Tensor[float32], logProb: float32] =
|
|
||||||
## state: [STATE_DIM]. Returns actions [ACTION_DIM] and sum log-prob.
|
|
||||||
## deterministic=true: return mean only (no noise), logProb=0.
|
|
||||||
let mean = ac.actor.forward(state)
|
|
||||||
if deterministic:
|
|
||||||
return (actions: mean, logProb: 0.0'f32)
|
|
||||||
|
|
||||||
# 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
|
|
||||||
@@ -1,13 +0,0 @@
|
|||||||
# Package
|
|
||||||
version = "0.1.0"
|
|
||||||
author = "Davide Cappellini"
|
|
||||||
description = "SAC+LSTM-trained Tank Royale bot"
|
|
||||||
license = "MIT"
|
|
||||||
srcDir = "src"
|
|
||||||
bin = @["SAC_LSTM_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,12 +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.
|
|
||||||
# Must come AFTER the nimble.paths include: later --path wins the import search.
|
|
||||||
switch("path", thisDir() & "/../libs/tankroyale_botapi")
|
|
||||||
switch("path", thisDir() & "/../libs/radar_lock")
|
|
||||||
@@ -1,11 +0,0 @@
|
|||||||
{
|
|
||||||
"name": "Recurrent Royalty",
|
|
||||||
"version": "0.1.0",
|
|
||||||
"authors": ["Davide Cappellini"],
|
|
||||||
"description": "SAC+LSTM Tank Royale bot — skeleton with radar lock",
|
|
||||||
"homepage": "",
|
|
||||||
"countryCodes": ["IT"],
|
|
||||||
"gameTypes": ["classic", "melee", "1v1"],
|
|
||||||
"platform": "Nim",
|
|
||||||
"programmingLang": "Nim"
|
|
||||||
}
|
|
||||||
@@ -1,60 +0,0 @@
|
|||||||
## SAC_LSTM_Bot — skeleton: radar lock + "Recurrent Royalty" color scheme.
|
|
||||||
## No RL yet. Connects, sets colors, locks radar onto enemy.
|
|
||||||
|
|
||||||
import std/os
|
|
||||||
import tankroyale_botapi
|
|
||||||
import radar_lock
|
|
||||||
|
|
||||||
const botJsonPath = currentSourcePath().parentDir / "SAC_LSTM_Bot.json"
|
|
||||||
|
|
||||||
# ── Colors (Recurrent Royalty palette) ───────────────────────────────────────
|
|
||||||
const
|
|
||||||
ColBody = fromHex("#7B2FBE")
|
|
||||||
ColTurret = fromHex("#FFD700")
|
|
||||||
ColGun = fromHex("#4A0E6B")
|
|
||||||
ColRadar = fromHex("#FFD700")
|
|
||||||
ColScan = fromHex("#FFB000")
|
|
||||||
ColBullet = fromHex("#FFC125")
|
|
||||||
ColTracks = fromHex("#2C2C34")
|
|
||||||
|
|
||||||
proc applyColors() =
|
|
||||||
setBodyColor(ColBody)
|
|
||||||
setTurretColor(ColTurret)
|
|
||||||
setGunColor(ColGun)
|
|
||||||
setRadarColor(ColRadar)
|
|
||||||
setScanColor(ColScan)
|
|
||||||
setBulletColor(ColBullet)
|
|
||||||
setTracksColor(ColTracks)
|
|
||||||
|
|
||||||
# ── Bot type ──────────────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
type SacBot = ref object of Bot
|
|
||||||
enemyBearing: float # last known absolute bearing to enemy
|
|
||||||
|
|
||||||
# ── Event handlers ────────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
method onRoundStarted*(bot: SacBot, e: RoundStartedEvent) =
|
|
||||||
setAdjustRadarForBodyTurn(true)
|
|
||||||
setAdjustRadarForGunTurn(true)
|
|
||||||
radar_lock.init()
|
|
||||||
applyColors()
|
|
||||||
|
|
||||||
method onScannedBot*(bot: SacBot, e: ScannedBotEvent) =
|
|
||||||
bot.enemyBearing = directionTo(getX(), getY(), e.x, e.y)
|
|
||||||
# Same-tick radar lock: apply turn rate immediately so it takes effect this tick.
|
|
||||||
setRadarTurnRate(radar_lock.doRadar(getRadarDirection(), bot.enemyBearing))
|
|
||||||
|
|
||||||
# ── Run loop ──────────────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
method run(bot: SacBot) =
|
|
||||||
while isRunning():
|
|
||||||
# Spin radar when no enemy is visible (full sweep).
|
|
||||||
if bot.enemyBearing == 0.0:
|
|
||||||
setRadarTurnRate(45.0)
|
|
||||||
go()
|
|
||||||
|
|
||||||
# ── Entry point ───────────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
when isMainModule:
|
|
||||||
var bot = SacBot()
|
|
||||||
start(bot, botJsonPath)
|
|
||||||
@@ -1,49 +0,0 @@
|
|||||||
## rewards.nim — Raw reward computation + running mean/variance normalizer.
|
|
||||||
## Welford online algorithm; safe cold-start (0 or 1 samples).
|
|
||||||
|
|
||||||
import std/math
|
|
||||||
|
|
||||||
# ── Raw reward ────────────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
proc computeReward*(
|
|
||||||
damageInflicted: float64 = 0.0, # fire power p of own shot that hit
|
|
||||||
damageReceived: float64 = 0.0, # fire power p_e of enemy shot that hit
|
|
||||||
wallHitTicks: int = 0, # ticks in wall contact this step
|
|
||||||
wastedShotPower: float64 = 0.0, # fire power of shot that missed/hit wall
|
|
||||||
win: bool = false,
|
|
||||||
loss: bool = false
|
|
||||||
): float64 =
|
|
||||||
## Returns the raw (un-normalized) reward for one decision step.
|
|
||||||
## Damage formula: 4p + 2(p-1) = 6p - 2 (matches Tank Royale bullet rules).
|
|
||||||
let p = damageInflicted
|
|
||||||
let pe = damageReceived
|
|
||||||
if p > 0.0: result += 6.0 * p - 2.0
|
|
||||||
if pe > 0.0: result -= 6.0 * pe - 2.0
|
|
||||||
result -= 5.0 * wallHitTicks.float64
|
|
||||||
if wastedShotPower > 0.0: result -= 0.1 * wastedShotPower
|
|
||||||
if win: result += 20.0
|
|
||||||
if loss: result -= 10.0
|
|
||||||
|
|
||||||
# ── Running normalizer (Welford) ──────────────────────────────────────────────
|
|
||||||
|
|
||||||
const NormEps = 1e-8
|
|
||||||
|
|
||||||
type
|
|
||||||
RewardNormalizer* = object
|
|
||||||
n*: int # samples seen
|
|
||||||
mean*: float64
|
|
||||||
m2*: float64 # sum of squared deviations (Welford M2)
|
|
||||||
|
|
||||||
proc update*(rn: var RewardNormalizer; r: float64) =
|
|
||||||
rn.n += 1
|
|
||||||
let delta = r - rn.mean
|
|
||||||
rn.mean += delta / rn.n.float64
|
|
||||||
let delta2 = r - rn.mean
|
|
||||||
rn.m2 += delta * delta2
|
|
||||||
|
|
||||||
proc normalize*(rn: RewardNormalizer; r: float64): float64 =
|
|
||||||
## Returns (r - mean) / (std + eps).
|
|
||||||
## Cold start (n < 2): returns 0.0 to avoid NaN/inf.
|
|
||||||
if rn.n < 2: return 0.0
|
|
||||||
let variance = rn.m2 / rn.n.float64 # ponytail: population var; switch to n-1 if bias matters
|
|
||||||
result = (r - rn.mean) / (sqrt(variance) + NormEps)
|
|
||||||
@@ -1,117 +0,0 @@
|
|||||||
## State vector module — produces a 35-dimensional normalized tensor for SAC+LSTM policy.
|
|
||||||
## No bot API imports; takes plain data structs populated from game events.
|
|
||||||
## The LSTM handles temporal context, so no explicit history window here.
|
|
||||||
|
|
||||||
import std/math
|
|
||||||
import arraymancer
|
|
||||||
|
|
||||||
const STATE_DIM* = 35
|
|
||||||
|
|
||||||
type
|
|
||||||
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
|
|
||||||
|
|
||||||
EnemyData* = object
|
|
||||||
## Current enemy state, from the most recent onScannedBot event.
|
|
||||||
x*, y*: float64
|
|
||||||
direction*: float64
|
|
||||||
speed*: float64
|
|
||||||
energy*: float64
|
|
||||||
hasFired*: bool
|
|
||||||
lastFirePower*: float64
|
|
||||||
prevSpeed*: float64 # speed from the previous scan (for acceleration)
|
|
||||||
prevDirection*: float64 # direction from the previous scan (for turn rate)
|
|
||||||
hasPrevScan*: bool # true once we have at least two scans
|
|
||||||
|
|
||||||
GameState* = object
|
|
||||||
## Accumulates data from bot events. Populate fields before calling buildState.
|
|
||||||
# Own bot
|
|
||||||
x*, y*: float64
|
|
||||||
direction*: float64
|
|
||||||
speed*: float64
|
|
||||||
energy*: float64
|
|
||||||
gunDirection*: float64
|
|
||||||
gunHeat*: float64
|
|
||||||
arenaWidth*, arenaHeight*: float64
|
|
||||||
# Enemy
|
|
||||||
hasContact*: bool
|
|
||||||
enemy*: EnemyData
|
|
||||||
ticksSinceLastScan*: int
|
|
||||||
# Bullets in flight (up to 3 tracked)
|
|
||||||
bullets*: array[3, BulletData]
|
|
||||||
bulletCount*: int
|
|
||||||
|
|
||||||
proc buildState*(gs: GameState): Tensor[float32] =
|
|
||||||
## Build the 35-float normalized state tensor.
|
|
||||||
##
|
|
||||||
## Layout:
|
|
||||||
## [0-6] own bot: x/aW, y/aH, dir/360, speed/8, energy/100, gunDir/360, gunHeat/1.8
|
|
||||||
## [7-13] enemy: x/aW, y/aH, dir/360, speed/8, energy/100, hasFired, lastFirePower/3
|
|
||||||
## [14-17] derived: enemyAccel/8, enemyTurnRate/180, relBearing/180, distance/diag
|
|
||||||
## [18-21] walls: top, bottom, left, right — each / max(aW,aH)
|
|
||||||
## [22-33] bullets: up to 3 × (relX/aW, relY/aH, speed/20, ticksToImpact clamped to 1)
|
|
||||||
## [34] scan staleness: ticksSinceLastScan/30 clamped to 1
|
|
||||||
result = zeros[float32](STATE_DIM)
|
|
||||||
|
|
||||||
let aW = gs.arenaWidth
|
|
||||||
let aH = gs.arenaHeight
|
|
||||||
let diag = sqrt(aW * aW + aH * aH)
|
|
||||||
let wMax = max(aW, aH)
|
|
||||||
|
|
||||||
# --- Own bot (0-6) ---
|
|
||||||
result[0] = float32(gs.x / aW)
|
|
||||||
result[1] = float32(gs.y / aH)
|
|
||||||
result[2] = float32(gs.direction / 360.0)
|
|
||||||
result[3] = float32(gs.speed / 8.0)
|
|
||||||
result[4] = float32(gs.energy / 100.0)
|
|
||||||
result[5] = float32(gs.gunDirection / 360.0)
|
|
||||||
result[6] = float32(gs.gunHeat / 1.8)
|
|
||||||
|
|
||||||
# --- Enemy current (7-13) ---
|
|
||||||
if gs.hasContact:
|
|
||||||
result[7] = float32(gs.enemy.x / aW)
|
|
||||||
result[8] = float32(gs.enemy.y / aH)
|
|
||||||
result[9] = float32(gs.enemy.direction / 360.0)
|
|
||||||
result[10] = float32(gs.enemy.speed / 8.0)
|
|
||||||
result[11] = float32(gs.enemy.energy / 100.0)
|
|
||||||
result[12] = float32(if gs.enemy.hasFired: 1.0 else: 0.0)
|
|
||||||
result[13] = float32(gs.enemy.lastFirePower / 3.0)
|
|
||||||
|
|
||||||
# --- Derived (14-17) ---
|
|
||||||
if gs.hasContact:
|
|
||||||
if gs.enemy.hasPrevScan:
|
|
||||||
result[14] = float32((gs.enemy.speed - gs.enemy.prevSpeed) / 8.0)
|
|
||||||
let dDir = ((gs.enemy.direction - gs.enemy.prevDirection) + 540.0) mod 360.0 - 180.0
|
|
||||||
result[15] = float32(dDir / 180.0)
|
|
||||||
let dx = gs.enemy.x - gs.x
|
|
||||||
let dy = gs.enemy.y - gs.y
|
|
||||||
let absDir = (180.0 * arctan2(dx, dy) / PI + 360.0) mod 360.0
|
|
||||||
let relBearing = ((absDir - gs.direction) + 540.0) mod 360.0 - 180.0
|
|
||||||
result[16] = float32(relBearing / 180.0)
|
|
||||||
result[17] = float32(sqrt(dx * dx + dy * dy) / diag)
|
|
||||||
|
|
||||||
# --- Wall distances (18-21): top, bottom, left, right ---
|
|
||||||
result[18] = float32((aH - gs.y) / wMax)
|
|
||||||
result[19] = float32(gs.y / wMax)
|
|
||||||
result[20] = float32(gs.x / wMax)
|
|
||||||
result[21] = float32((aW - gs.x) / wMax)
|
|
||||||
|
|
||||||
# --- Bullet tracking (22-33): up to 3 bullets × 4 floats ---
|
|
||||||
# Per slot: relX/aW, relY/aH, speed/20, ticksToImpact/diag (clamped to 1)
|
|
||||||
for i in 0 ..< min(gs.bulletCount, 3):
|
|
||||||
let b = gs.bullets[i]
|
|
||||||
let bSpd = 20.0 - 3.0 * b.power
|
|
||||||
let bdx = b.x - gs.x
|
|
||||||
let bdy = b.y - gs.y
|
|
||||||
let bdist = sqrt(bdx * bdx + bdy * bdy)
|
|
||||||
let ticks = if bSpd > 0.0: min(bdist / bSpd / diag, 1.0) else: 0.0
|
|
||||||
let base = 22 + i * 4
|
|
||||||
result[base + 0] = float32(bdx / aW)
|
|
||||||
result[base + 1] = float32(bdy / aH)
|
|
||||||
result[base + 2] = float32(bSpd / 20.0)
|
|
||||||
result[base + 3] = float32(ticks)
|
|
||||||
|
|
||||||
# --- Scan staleness (34) ---
|
|
||||||
result[34] = float32(min(gs.ticksSinceLastScan.float64 / 30.0, 1.0))
|
|
||||||
@@ -1,2 +0,0 @@
|
|||||||
switch("path", "../src")
|
|
||||||
switch("path", "../../libs")
|
|
||||||
@@ -1,79 +0,0 @@
|
|||||||
## Assert-based tests for rewards.nim.
|
|
||||||
## Run: nim c -r tests/test_rewards.nim
|
|
||||||
|
|
||||||
import std/[math, strformat]
|
|
||||||
import SAC_LSTM_Bot/rewards
|
|
||||||
|
|
||||||
template check(cond: bool, msg: string) =
|
|
||||||
if not cond:
|
|
||||||
quit("FAIL: " & msg, 1)
|
|
||||||
|
|
||||||
# ── computeReward ─────────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
block damageInflicted:
|
|
||||||
# p=1: 6*1 - 2 = 4
|
|
||||||
check abs(computeReward(damageInflicted = 1.0) - 4.0) < 1e-9, "p=1 damage = +4"
|
|
||||||
# p=3: 6*3 - 2 = 16
|
|
||||||
check abs(computeReward(damageInflicted = 3.0) - 16.0) < 1e-9, "p=3 damage = +16"
|
|
||||||
|
|
||||||
block damageReceived:
|
|
||||||
# p_e=1: -(6*1 - 2) = -4
|
|
||||||
check abs(computeReward(damageReceived = 1.0) - (-4.0)) < 1e-9, "p_e=1 received = -4"
|
|
||||||
# p_e=3: -(6*3 - 2) = -16
|
|
||||||
check abs(computeReward(damageReceived = 3.0) - (-16.0)) < 1e-9, "p_e=3 received = -16"
|
|
||||||
|
|
||||||
block wallHit:
|
|
||||||
check abs(computeReward(wallHitTicks = 1) - (-5.0)) < 1e-9, "1 wall tick = -5"
|
|
||||||
|
|
||||||
block wastedShot:
|
|
||||||
# p=2: -0.1 * 2 = -0.2
|
|
||||||
check abs(computeReward(wastedShotPower = 2.0) - (-0.2)) < 1e-9, "missed p=2 = -0.2"
|
|
||||||
|
|
||||||
block winLoss:
|
|
||||||
check abs(computeReward(win = true) - 20.0) < 1e-9, "win = +20"
|
|
||||||
check abs(computeReward(loss = true) - (-10.0)) < 1e-9, "loss = -10"
|
|
||||||
|
|
||||||
# ── RewardNormalizer cold start ───────────────────────────────────────────────
|
|
||||||
|
|
||||||
block coldStart:
|
|
||||||
var rn: RewardNormalizer
|
|
||||||
# 0 samples
|
|
||||||
let v0 = rn.normalize(99.0)
|
|
||||||
check not isNaN(v0), "0 samples: not NaN"
|
|
||||||
check classify(v0) != fcInf and classify(v0) != fcNegInf, "0 samples: not inf"
|
|
||||||
check abs(v0) < 1e-9, "0 samples: returns 0"
|
|
||||||
# 1 sample (variance undefined)
|
|
||||||
rn.update(5.0)
|
|
||||||
let v1 = rn.normalize(5.0)
|
|
||||||
check not isNaN(v1), "1 sample: not NaN"
|
|
||||||
check classify(v1) != fcInf and classify(v1) != fcNegInf, "1 sample: not inf"
|
|
||||||
check abs(v1) < 1e-9, "1 sample: returns 0"
|
|
||||||
|
|
||||||
# ── Running normalization convergence ─────────────────────────────────────────
|
|
||||||
|
|
||||||
block convergence:
|
|
||||||
var rn: RewardNormalizer
|
|
||||||
# Feed 1000 identical samples of 5.0 — mean=5.0, std=0 → normalizer returns ~0
|
|
||||||
for _ in 0 ..< 1000:
|
|
||||||
rn.update(5.0)
|
|
||||||
let v = rn.normalize(5.0)
|
|
||||||
check not isNaN(v), "convergence: not NaN"
|
|
||||||
check classify(v) != fcInf and classify(v) != fcNegInf, "convergence: not inf"
|
|
||||||
# (5 - 5) / (0 + eps) = 0
|
|
||||||
check abs(v) < 1e-6, "convergence to mean: normalized ≈ 0"
|
|
||||||
|
|
||||||
block knownMeanStd:
|
|
||||||
# Insert samples -1 and +1 repeatedly → mean=0, std=1
|
|
||||||
var rn: RewardNormalizer
|
|
||||||
for _ in 0 ..< 500:
|
|
||||||
rn.update(-1.0)
|
|
||||||
rn.update( 1.0)
|
|
||||||
# normalize(1.0) ≈ (1 - 0) / (1 + eps) ≈ 1
|
|
||||||
let vPos = rn.normalize(1.0)
|
|
||||||
check abs(vPos - 1.0) < 1e-4, &"normalize(+1) ≈ +1, got {vPos}"
|
|
||||||
let vNeg = rn.normalize(-1.0)
|
|
||||||
check abs(vNeg - (-1.0)) < 1e-4, &"normalize(-1) ≈ -1, got {vNeg}"
|
|
||||||
let vMid = rn.normalize(0.0)
|
|
||||||
check abs(vMid) < 1e-4, &"normalize(0) ≈ 0, got {vMid}"
|
|
||||||
|
|
||||||
echo "test_rewards: all passed"
|
|
||||||
@@ -1,87 +0,0 @@
|
|||||||
## Tests for state.nim — assert-based, no framework.
|
|
||||||
|
|
||||||
import std/math
|
|
||||||
import arraymancer
|
|
||||||
import SAC_LSTM_Bot/state
|
|
||||||
|
|
||||||
proc makeBase(): GameState =
|
|
||||||
result.arenaWidth = 1200.0
|
|
||||||
result.arenaHeight = 800.0
|
|
||||||
result.x = 600.0; result.y = 400.0
|
|
||||||
result.direction = 90.0; result.speed = 4.0
|
|
||||||
result.energy = 50.0
|
|
||||||
result.gunDirection = 90.0; result.gunHeat = 0.5
|
|
||||||
|
|
||||||
proc allInRange(t: Tensor[float32]): bool =
|
|
||||||
for v in t:
|
|
||||||
if v < -1.01f32 or v > 1.01f32: return false
|
|
||||||
true
|
|
||||||
|
|
||||||
proc hasNaN(t: Tensor[float32]): bool =
|
|
||||||
for v in t:
|
|
||||||
if v.float64.isNaN: return true
|
|
||||||
false
|
|
||||||
|
|
||||||
# 1. Correct shape
|
|
||||||
block:
|
|
||||||
let gs = makeBase()
|
|
||||||
let t = buildState(gs)
|
|
||||||
assert t.shape[0] == STATE_DIM, "wrong dim: " & $t.shape[0]
|
|
||||||
echo "PASS shape"
|
|
||||||
|
|
||||||
# 2. All values in [-1, 1] for typical input
|
|
||||||
block:
|
|
||||||
var gs = makeBase()
|
|
||||||
gs.hasContact = true
|
|
||||||
gs.enemy = EnemyData(x: 800.0, y: 300.0, direction: 45.0, speed: 6.0,
|
|
||||||
energy: 80.0, hasFired: true, lastFirePower: 2.0,
|
|
||||||
prevSpeed: 5.0, prevDirection: 40.0, hasPrevScan: true)
|
|
||||||
gs.bulletCount = 1
|
|
||||||
gs.bullets[0] = BulletData(x: 610.0, y: 410.0, power: 1.0)
|
|
||||||
gs.ticksSinceLastScan = 10
|
|
||||||
let t = buildState(gs)
|
|
||||||
assert not hasNaN(t), "NaN in tensor"
|
|
||||||
assert allInRange(t), "value out of [-1,1]"
|
|
||||||
echo "PASS range"
|
|
||||||
|
|
||||||
# 3. Missing scan → valid tensor, no NaN, zeros for enemy fields
|
|
||||||
block:
|
|
||||||
let gs = makeBase() # hasContact = false
|
|
||||||
let t = buildState(gs)
|
|
||||||
assert not hasNaN(t), "NaN with no scan"
|
|
||||||
for i in 7 .. 17:
|
|
||||||
assert t[i] == 0.0f32, "enemy slot " & $i & " should be 0"
|
|
||||||
echo "PASS no-scan zeros"
|
|
||||||
|
|
||||||
# 4. Bullet tracking: 0, 1, 2, 3 bullets
|
|
||||||
block:
|
|
||||||
for n in 0 .. 3:
|
|
||||||
var gs = makeBase()
|
|
||||||
gs.bulletCount = n
|
|
||||||
for i in 0 ..< n:
|
|
||||||
gs.bullets[i] = BulletData(x: float64(600 + i * 10), y: float64(400 + i * 5), power: 1.0)
|
|
||||||
let t = buildState(gs)
|
|
||||||
assert not hasNaN(t), "NaN with " & $n & " bullets"
|
|
||||||
# slots beyond bulletCount must be 0
|
|
||||||
for i in n ..< 3:
|
|
||||||
let base = 22 + i * 4
|
|
||||||
for j in 0 ..< 4:
|
|
||||||
assert t[base + j] == 0.0f32, "bullet slot " & $i & " slot " & $j & " should be 0"
|
|
||||||
echo "PASS bullet tracking 0-3"
|
|
||||||
|
|
||||||
# 5. Scan staleness increments and clamps
|
|
||||||
block:
|
|
||||||
var gs = makeBase()
|
|
||||||
gs.hasContact = true
|
|
||||||
gs.ticksSinceLastScan = 0
|
|
||||||
assert buildState(gs)[34] == 0.0f32, "staleness at 0"
|
|
||||||
gs.ticksSinceLastScan = 15
|
|
||||||
let mid = buildState(gs)[34]
|
|
||||||
assert mid > 0.0f32 and mid < 1.0f32, "staleness mid not in (0,1)"
|
|
||||||
gs.ticksSinceLastScan = 30
|
|
||||||
assert buildState(gs)[34] == 1.0f32, "staleness at 30"
|
|
||||||
gs.ticksSinceLastScan = 60
|
|
||||||
assert buildState(gs)[34] == 1.0f32, "staleness clamp at 60"
|
|
||||||
echo "PASS staleness"
|
|
||||||
|
|
||||||
echo "ALL TESTS PASSED"
|
|
||||||
@@ -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.
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user