Compare commits
35 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 23c65c9ac6 | |||
| f130bf1254 | |||
| a0a3840980 | |||
| 7d73d32c85 | |||
| df3bbbd14e | |||
| 4258d364b9 | |||
| f88580b157 | |||
| ca3e3d2272 | |||
| 0d35646dc9 | |||
| c834d2cbee | |||
| fedab54bc0 | |||
| 82eeb53e5c | |||
| 12624d3069 | |||
| 6ad51148f4 | |||
| 75e32e3315 | |||
| db99153f65 | |||
| da2f825ad8 | |||
| 186e005a96 | |||
| a4e830531b | |||
| 64697f917e | |||
| 766b9e03ee | |||
| 0b17430735 | |||
| cdde60d79f | |||
| 56e0b306c9 | |||
| e5609a7d9b | |||
| a8ee2a86e3 | |||
| bbc9e51166 | |||
| fd22535f5b | |||
| f27b0238f0 | |||
| aea0724d3a | |||
| 473d67f644 | |||
| eadd177d3b | |||
| 588c9ebc2f | |||
| aa4bc77068 | |||
| 30cda871cc |
@@ -0,0 +1,3 @@
|
||||
nimble.develop
|
||||
nimble.paths
|
||||
nimbledeps
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,11 @@
|
||||
{
|
||||
"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"
|
||||
}
|
||||
@@ -0,0 +1,95 @@
|
||||
## 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)
|
||||
@@ -0,0 +1,10 @@
|
||||
# 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"
|
||||
@@ -0,0 +1,4 @@
|
||||
# begin Nimble config (version 2)
|
||||
when withDir(thisDir(), system.fileExists("nimble.paths")):
|
||||
include "nimble.paths"
|
||||
# end Nimble config
|
||||
@@ -0,0 +1,3 @@
|
||||
nimble.develop
|
||||
nimble.paths
|
||||
nimbledeps
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,11 @@
|
||||
{
|
||||
"name": "PPO_Bot",
|
||||
"version": "0.1.0",
|
||||
"authors": ["Davide Cappellini"],
|
||||
"description": "PPO-trained RL bot",
|
||||
"homepage": "",
|
||||
"countryCodes": ["IT"],
|
||||
"gameTypes": ["classic", "melee", "1v1"],
|
||||
"platform": "Nim",
|
||||
"programmingLang": "Nim"
|
||||
}
|
||||
@@ -0,0 +1,371 @@
|
||||
## 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)
|
||||
@@ -0,0 +1,12 @@
|
||||
# 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"
|
||||
Executable
+7
@@ -0,0 +1,7 @@
|
||||
#!/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"
|
||||
@@ -0,0 +1,49 @@
|
||||
## 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
|
||||
@@ -0,0 +1,13 @@
|
||||
# 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")
|
||||
@@ -0,0 +1,24 @@
|
||||
## 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)
|
||||
@@ -0,0 +1,100 @@
|
||||
## 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)
|
||||
@@ -0,0 +1,86 @@
|
||||
## 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
|
||||
@@ -0,0 +1,7 @@
|
||||
{ pkgs ? import <nixpkgs> {} }:
|
||||
|
||||
pkgs.mkShell {
|
||||
buildInputs = with pkgs; [
|
||||
openblas
|
||||
];
|
||||
}
|
||||
@@ -0,0 +1,121 @@
|
||||
## 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))
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,69 @@
|
||||
## 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"
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,54 @@
|
||||
## 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"
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,58 @@
|
||||
## 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"
|
||||
@@ -0,0 +1,42 @@
|
||||
## 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"
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,187 @@
|
||||
## 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"
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,175 @@
|
||||
## 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"
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,105 @@
|
||||
## 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"
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,449 @@
|
||||
## 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
|
||||
@@ -0,0 +1,211 @@
|
||||
## 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.
@@ -0,0 +1,13 @@
|
||||
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.
@@ -0,0 +1 @@
|
||||
0
|
||||
@@ -0,0 +1,13 @@
|
||||
# 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"
|
||||
@@ -0,0 +1,12 @@
|
||||
# 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")
|
||||
@@ -0,0 +1,11 @@
|
||||
{
|
||||
"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"
|
||||
}
|
||||
@@ -0,0 +1,60 @@
|
||||
## 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)
|
||||
@@ -0,0 +1,49 @@
|
||||
## 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)
|
||||
@@ -0,0 +1,117 @@
|
||||
## 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))
|
||||
@@ -0,0 +1,2 @@
|
||||
switch("path", "../src")
|
||||
switch("path", "../../libs")
|
||||
@@ -0,0 +1,79 @@
|
||||
## 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"
|
||||
@@ -0,0 +1,87 @@
|
||||
## 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,224 @@
|
||||
# Goto Controller Algorithm — Research
|
||||
|
||||
**Issue:** #20
|
||||
**Branch:** research/goto-controller
|
||||
**Date:** 2026-08-17
|
||||
|
||||
---
|
||||
|
||||
## Problem Statement
|
||||
|
||||
The PPO network will output a target position `(x, y)`. A goto controller must
|
||||
translate that into per-tick `setTargetSpeed` and `setTurnRate` commands for
|
||||
the Tank Royale Nim bot API.
|
||||
|
||||
---
|
||||
|
||||
## Codebase Findings
|
||||
|
||||
### Tank Royale Nim API — no built-in goto
|
||||
|
||||
The library (`tankroyale_botapi` v1.0.1) provides:
|
||||
|
||||
- `setTargetSpeed(speed: float)` — desired speed, clamped to ±8 units/tick.
|
||||
Server auto-manages acceleration/deceleration via `getNewTargetSpeed`.
|
||||
- `setTurnRate(rate: float)` — desired turn rate, clamped to `calcMaxTurnRate(speed) = 10 - 0.75 * abs(speed)`.
|
||||
- `setForward(distance)` / `setBack(distance)` — blocking helpers that use
|
||||
`gDistanceRemaining` + the server's deceleration model. These are blocking
|
||||
(call `go()` internally) and therefore cannot be used in the non-blocking
|
||||
per-tick run loop used by PPO_Bot.
|
||||
|
||||
There is **no built-in `setDistanceRemaining`-style goto**. The controller must
|
||||
be written from scratch.
|
||||
|
||||
### Physics constants (from `constants.nim` / `utils.nim`)
|
||||
|
||||
| Constant | Value |
|
||||
|---|---|
|
||||
| Max speed | 8 units/tick |
|
||||
| Acceleration | +1 unit/tick² |
|
||||
| Deceleration | −2 units/tick² (braking is twice as fast) |
|
||||
| Max turn rate | `10 − 0.75 × |speed|` deg/tick |
|
||||
| Min turn rate (at max speed) | `10 − 0.75 × 8 = 4` deg/tick |
|
||||
| `getNewTargetSpeed(maxSpeed, speed, dist)` | already implemented in utils.nim |
|
||||
|
||||
Key implication: **you can turn faster while slow**. Turn-then-drive lets the
|
||||
bot use full 10°/tick turn rate, but wastes ticks stopped. Driving-while-turning
|
||||
is smooth but limited to 4°/tick at top speed.
|
||||
|
||||
### Coordinate system
|
||||
|
||||
North = 0°, clockwise. `directionTo` in `utils.nim` returns a bearing in
|
||||
`[0, 360)`. `bearingTo` returns a signed relative bearing in `(-180, 180]`.
|
||||
|
||||
---
|
||||
|
||||
## Approaches Considered
|
||||
|
||||
### A — Turn-then-drive (sequential)
|
||||
|
||||
Stop → turn to face target → drive full speed → brake.
|
||||
|
||||
- Simple to implement.
|
||||
- Very slow: wastes ticks turning at zero speed then decelerating.
|
||||
- Produces jerky, non-smooth movement — bad as a controller layer.
|
||||
|
||||
### B — Proportional navigation (continuous per-tick)
|
||||
|
||||
Each tick: compute bearing to target, set turn rate proportional to bearing
|
||||
error, set speed based on distance remaining.
|
||||
|
||||
- Standard Robocode idiom. Very common in published bots.
|
||||
- Does not make the forward-vs-reverse decision optimally.
|
||||
- Can overshoot if gains are too high; can be sluggish if too low.
|
||||
|
||||
### C — Arc/pursuit steering (proportional + speed-dependent turn limit)
|
||||
|
||||
Like B, but explicitly clamps turn rate to `calcMaxTurnRate(currentSpeed)` and
|
||||
scales speed down when the heading error is large (so the bot slows to increase
|
||||
turn authority).
|
||||
|
||||
- Handles Tank Royale's speed-dependent turn rate correctly.
|
||||
- Naturally smooth.
|
||||
- Still needs explicit forward/reverse decision.
|
||||
|
||||
### D — Forward-vs-reverse decision + proportional steering (recommended)
|
||||
|
||||
Extend C with the classic Robocode "should I go backward?" heuristic:
|
||||
if `|bearingError| > 90°`, it is faster to reverse and face the target with
|
||||
the rear than to turn more than 90° forward. Flip target speed sign and add
|
||||
180° to the bearing before computing turn rate.
|
||||
|
||||
This is the approach used by high-quality Robocode 1 bots (e.g. RaikoMX,
|
||||
Aristocles) and it trivially maps to Tank Royale's API.
|
||||
|
||||
---
|
||||
|
||||
## Recommended Algorithm
|
||||
|
||||
### Decision: forward or reverse?
|
||||
|
||||
```
|
||||
bearing = normalizeRelativeAngle(directionTo(x, y) - direction)
|
||||
if abs(bearing) > 90.0:
|
||||
# Going backward is cheaper
|
||||
direction_sign = -1
|
||||
effective_bearing = normalizeRelativeAngle(bearing + 180.0)
|
||||
else:
|
||||
direction_sign = +1
|
||||
effective_bearing = bearing
|
||||
```
|
||||
|
||||
### Turn rate
|
||||
|
||||
Apply full proportional turn rate toward the effective bearing:
|
||||
|
||||
```
|
||||
max_turn = 10.0 - 0.75 * abs(currentSpeed)
|
||||
turnRate = clamp(effective_bearing, -max_turn, max_turn)
|
||||
```
|
||||
|
||||
`effective_bearing` acts as both direction and magnitude: if the error is
|
||||
small, the turn rate is small (smooth approach); if large, it clamps to max
|
||||
(fastest possible turn).
|
||||
|
||||
### Target speed
|
||||
|
||||
Use `getNewTargetSpeed` (already in `utils.nim`) to determine the speed
|
||||
that will arrive at the target with zero velocity:
|
||||
|
||||
```
|
||||
dist = distanceTo(x, y)
|
||||
raw_speed = getNewTargetSpeed(MAX_SPEED, currentSpeed, dist)
|
||||
targetSpeed = direction_sign * raw_speed
|
||||
```
|
||||
|
||||
This reuses the exact deceleration model the server uses, so the bot always
|
||||
brakes at the right time with no overshoot.
|
||||
|
||||
### Stop condition
|
||||
|
||||
```
|
||||
if dist < ARRIVAL_THRESHOLD: # e.g. 18.0 (= BOT_RADIUS)
|
||||
targetSpeed = 0.0
|
||||
turnRate = 0.0
|
||||
```
|
||||
|
||||
### Full pseudocode (one tick)
|
||||
|
||||
```nim
|
||||
proc gotoTick*(tx, ty, x, y, direction, currentSpeed: float):
|
||||
tuple[targetSpeed, turnRate: float] =
|
||||
|
||||
let dist = distanceTo(x, y, tx, ty)
|
||||
|
||||
if dist < ARRIVAL_THRESHOLD:
|
||||
return (0.0, 0.0)
|
||||
|
||||
let rawBearing = normalizeRelativeAngle(directionTo(x, y, tx, ty) - direction)
|
||||
|
||||
let (dirSign, effBearing) =
|
||||
if abs(rawBearing) > 90.0:
|
||||
(-1.0, normalizeRelativeAngle(rawBearing + 180.0))
|
||||
else:
|
||||
(1.0, rawBearing)
|
||||
|
||||
let maxTurn = 10.0 - 0.75 * abs(currentSpeed)
|
||||
let turnRate = effBearing.clamp(-maxTurn, maxTurn)
|
||||
|
||||
let rawSpeed = getNewTargetSpeed(MAX_SPEED, abs(currentSpeed), dist)
|
||||
let targetSpeed = dirSign * rawSpeed
|
||||
|
||||
return (targetSpeed, turnRate)
|
||||
```
|
||||
|
||||
Call once per tick from the `run` loop, pass results to `setTargetSpeed` /
|
||||
`setTurnRate`.
|
||||
|
||||
---
|
||||
|
||||
## Why not pure proportional navigation (option B)?
|
||||
|
||||
Option B without the speed-dependent turn clamp will attempt to command more
|
||||
turn rate than the server will honor at high speed — it does the right thing
|
||||
emergently but wastes the gap. Explicitly scaling turn rate with
|
||||
`calcMaxTurnRate(speed)` is more intentional and matches the physics exactly.
|
||||
This is already coded in `actions.nim` (`r1 * (10.0 - 0.75 * abs(currentSpeed))`),
|
||||
so the pattern is established in the codebase.
|
||||
|
||||
---
|
||||
|
||||
## Why reuse `getNewTargetSpeed` from utils.nim?
|
||||
|
||||
It already encodes the exact asymmetric acceleration/deceleration model
|
||||
(accel +1, decel −2 per tick). Reimplementing distance-based speed management
|
||||
from scratch would duplicate this and risk drift. Import it directly.
|
||||
|
||||
---
|
||||
|
||||
## Forward/Reverse optimality
|
||||
|
||||
The 90° threshold is the exact breakeven point:
|
||||
|
||||
- Turning 91° forward takes ≥10 ticks at slow speed + travel time.
|
||||
- Reversing 89° (i.e. 180−91=89° effective turn) takes fewer ticks total
|
||||
for any distance large enough to matter.
|
||||
- For very short distances (< ~36 units) the bot will decelerate before the
|
||||
turn completes anyway; the threshold still works because the speed penalty
|
||||
applies equally to both cases.
|
||||
|
||||
For a controller layer that feeds a neural network's goto target, sub-optimal
|
||||
behavior on very short hops is acceptable — the network will learn to avoid
|
||||
issuing tiny hops.
|
||||
|
||||
---
|
||||
|
||||
## Sources / References
|
||||
|
||||
- Tank Royale Nim API source: `tankroyale_botapi/utils.nim`, `bot.nim`,
|
||||
`constants.nim` (v1.0.1, installed at `~/.nimble/pkgs2/`).
|
||||
- Robocode wiki — "Proportional navigation" and "Should I go backward?"
|
||||
heuristic: widely documented in the Robocode community (e.g. RoboWiki
|
||||
`BasicSurfer`, `RaikoMX` source).
|
||||
- Tank Royale physics spec: confirmed against `ACCELERATION = 1.0`,
|
||||
`ABS_DECELERATION = 2.0` in `constants.nim`.
|
||||
@@ -0,0 +1,38 @@
|
||||
## Standalone radar lock module for Robocode Tank Royale (1v1).
|
||||
## No bot-specific imports — takes plain floats, returns radarTurnRate.
|
||||
##
|
||||
## Tank Royale radar uses standard math convention: 0° = east, CCW positive.
|
||||
## Angles are in degrees.
|
||||
|
||||
import std/math
|
||||
|
||||
const
|
||||
MaxRadarTurn* = 45.0
|
||||
DefaultOvershootDeg* = 5.0
|
||||
|
||||
var overShootDeg* = DefaultOvershootDeg
|
||||
|
||||
proc init*() =
|
||||
## Reset module to defaults.
|
||||
## In your bot's constructor set:
|
||||
## adjustRadarForBodyTurn = true
|
||||
## adjustRadarForGunTurn = true
|
||||
overShootDeg = DefaultOvershootDeg
|
||||
|
||||
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 doRadar*(currentRadarHeading, enemyBearing: float64): float64 =
|
||||
## Returns radarTurnRate (degrees/tick, positive = clockwise).
|
||||
##
|
||||
## currentRadarHeading: current radar direction in degrees (0=east, CCW+).
|
||||
## enemyBearing: absolute bearing to enemy in same coordinate system.
|
||||
##
|
||||
## Handles angle wrapping, clamps to [-45, +45], adds overshoot sweep.
|
||||
var turn = normalizeRelative(enemyBearing - currentRadarHeading)
|
||||
# Add overshoot in the same direction as the turn to maintain lock
|
||||
if turn < 0.0: turn -= overShootDeg
|
||||
else: turn += overShootDeg
|
||||
result = turn.clamp(-MaxRadarTurn, MaxRadarTurn)
|
||||
@@ -0,0 +1,8 @@
|
||||
# Package
|
||||
version = "1.0.0"
|
||||
author = "Davide Cappellini"
|
||||
description = "Standalone radar lock module for Robocode Tank Royale"
|
||||
license = "Apache-2.0"
|
||||
|
||||
# Dependencies
|
||||
requires "nim >= 2.0.0"
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,74 @@
|
||||
import std/unittest
|
||||
import std/math
|
||||
import ../radar_lock
|
||||
|
||||
suite "radar_lock":
|
||||
|
||||
setup:
|
||||
init()
|
||||
|
||||
# Basic lock: radar already on enemy — overshoot pushes it slightly
|
||||
test "radar on enemy — returns overshoot only":
|
||||
let rate = doRadar(90.0, 90.0)
|
||||
check rate == DefaultOvershootDeg
|
||||
|
||||
# Cardinal directions
|
||||
test "lock from 0 degrees":
|
||||
let rate = doRadar(0.0, 0.0)
|
||||
check rate == DefaultOvershootDeg
|
||||
|
||||
test "lock from 90 degrees":
|
||||
let rate = doRadar(90.0, 90.0)
|
||||
check rate == DefaultOvershootDeg
|
||||
|
||||
test "lock from 180 degrees":
|
||||
let rate = doRadar(180.0, 180.0)
|
||||
check rate == DefaultOvershootDeg
|
||||
|
||||
test "lock from 270 degrees":
|
||||
let rate = doRadar(270.0, 270.0)
|
||||
check rate == DefaultOvershootDeg
|
||||
|
||||
# Angle wrapping: radar at 350°, enemy at 10° → shortest path is +20°
|
||||
test "wrapping 350 to 10 — turns right 20 + overshoot":
|
||||
let rate = doRadar(350.0, 10.0)
|
||||
check abs(rate - (20.0 + DefaultOvershootDeg)) < 1e-9
|
||||
|
||||
# Angle wrapping: radar at 10°, enemy at 350° → shortest path is -20°
|
||||
test "wrapping 10 to 350 — turns left 20 + overshoot":
|
||||
let rate = doRadar(10.0, 350.0)
|
||||
check abs(rate - (-20.0 - DefaultOvershootDeg)) < 1e-9
|
||||
|
||||
# Clamping: enemy 100° away → raw turn+overshoot > 45°, must clamp
|
||||
test "clamp positive — large gap":
|
||||
let rate = doRadar(0.0, 100.0)
|
||||
check rate == 45.0
|
||||
|
||||
test "clamp negative — large gap":
|
||||
let rate = doRadar(100.0, 0.0)
|
||||
check rate == -45.0
|
||||
|
||||
# Overshoot direction: turn right → positive overshoot
|
||||
test "overshoot direction right":
|
||||
let rate = doRadar(0.0, 30.0) # needs +30°, overshoot adds +5°
|
||||
check abs(rate - 35.0) < 1e-9
|
||||
|
||||
# Overshoot direction: turn left → negative overshoot
|
||||
test "overshoot direction left":
|
||||
let rate = doRadar(30.0, 0.0) # needs -30°, overshoot adds -5°
|
||||
check abs(rate - (-35.0)) < 1e-9
|
||||
|
||||
# Output never exceeds ±45
|
||||
test "output bounded above":
|
||||
let rate = doRadar(0.0, 179.0)
|
||||
check rate <= 45.0
|
||||
|
||||
test "output bounded below":
|
||||
let rate = doRadar(179.0, 0.0)
|
||||
check rate >= -45.0
|
||||
|
||||
# init() resets overShootDeg
|
||||
test "init resets overshoot":
|
||||
overShootDeg = 20.0
|
||||
init()
|
||||
check overShootDeg == DefaultOvershootDeg
|
||||
@@ -0,0 +1,258 @@
|
||||
## Main entry-point module for Robocode Tank Royale Nim bot API.
|
||||
##
|
||||
## Usage:
|
||||
## import tankroyale_botapi
|
||||
##
|
||||
## type MyBot = ref object of Bot
|
||||
## method run(bot: MyBot) =
|
||||
## forward(100)
|
||||
## ...
|
||||
##
|
||||
## var bot = MyBot()
|
||||
## start(bot, "MyBot.json")
|
||||
|
||||
import std/[os, json]
|
||||
|
||||
import ./tankroyale_botapi/constants
|
||||
import ./tankroyale_botapi/color
|
||||
import ./tankroyale_botapi/schemas
|
||||
import ./tankroyale_botapi/utils
|
||||
import ./tankroyale_botapi/bot_info
|
||||
import ./tankroyale_botapi/ws_client
|
||||
import ./tankroyale_botapi/json_parse
|
||||
import ./tankroyale_botapi/event_queue
|
||||
import ./tankroyale_botapi/bot
|
||||
import ./tankroyale_botapi/graphics
|
||||
|
||||
export constants
|
||||
export color
|
||||
export schemas
|
||||
export utils
|
||||
export bot_info
|
||||
export json_parse
|
||||
export event_queue
|
||||
export bot
|
||||
export graphics
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# WebSocket receive loop (main thread)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
proc handleServerHandshake(ws: SyncWebSocket; node: JsonNode; info: BotInfo; secret: string) =
|
||||
let sessionId = node{"sessionId"}.getStr
|
||||
setServerInfo(node{"variant"}.getStr, node{"version"}.getStr)
|
||||
|
||||
# Build bot handshake
|
||||
var h = newJObject()
|
||||
h["type"] = %"BotHandshake"
|
||||
h["sessionId"] = %sessionId
|
||||
h["name"] = %info.name
|
||||
h["version"] = %info.version
|
||||
h["authors"] = %info.authors
|
||||
h["description"] = %info.description
|
||||
h["homepage"] = %info.homepage
|
||||
h["countryCodes"] = %info.countryCodes
|
||||
h["gameTypes"] = %info.gameTypes
|
||||
h["platform"] = %info.platform
|
||||
h["programmingLang"]= %info.programmingLang
|
||||
h["isDroid"] = %info.isDroid
|
||||
if secret.len > 0:
|
||||
h["secret"] = %secret
|
||||
let ip = info.initialPosition
|
||||
if ip.x != 0.0 or ip.y != 0.0 or ip.direction != 0.0:
|
||||
var ipObj = newJObject()
|
||||
if ip.x != 0.0: ipObj["x"] = %ip.x
|
||||
if ip.y != 0.0: ipObj["y"] = %ip.y
|
||||
if ip.direction != 0.0: ipObj["direction"] = %ip.direction
|
||||
h["initialPosition"] = ipObj
|
||||
ws.send($h)
|
||||
|
||||
proc parseGameSetup(node: JsonNode): GameSetup =
|
||||
if node.isNil: return
|
||||
result.gameType = node{"gameType"}.getStr("classic")
|
||||
result.arenaWidth = node{"arenaWidth"}.getInt(800)
|
||||
result.isArenaWidthLocked = node{"isArenaWidthLocked"}.getBool(false)
|
||||
result.arenaHeight = node{"arenaHeight"}.getInt(600)
|
||||
result.isArenaHeightLocked = node{"isArenaHeightLocked"}.getBool(false)
|
||||
result.numberOfRounds = node{"numberOfRounds"}.getInt(10)
|
||||
result.isNumberOfRoundsLocked = node{"isNumberOfRoundsLocked"}.getBool(false)
|
||||
result.minNumberOfParticipants = node{"minNumberOfParticipants"}.getInt(2)
|
||||
result.isMinNumberOfParticipantsLocked = node{"isMinNumberOfParticipantsLocked"}.getBool(false)
|
||||
result.maxNumberOfParticipants = node{"maxNumberOfParticipants"}.getInt(10)
|
||||
result.isMaxNumberOfParticipantsLocked = node{"isMaxNumberOfParticipantsLocked"}.getBool(false)
|
||||
result.gunCoolingRate = node{"gunCoolingRate"}.getFloat(0.1)
|
||||
result.isGunCoolingRateLocked = node{"isGunCoolingRateLocked"}.getBool(false)
|
||||
result.maxInactivityTurns = node{"maxInactivityTurns"}.getInt(450)
|
||||
result.isMaxInactivityTurnsLocked = node{"isMaxInactivityTurnsLocked"}.getBool(false)
|
||||
result.turnTimeout = node{"turnTimeout"}.getInt(30000)
|
||||
result.isTurnTimeoutLocked = node{"isTurnTimeoutLocked"}.getBool(false)
|
||||
result.readyTimeout = node{"readyTimeout"}.getInt(1000000)
|
||||
result.isReadyTimeoutLocked = node{"isReadyTimeoutLocked"}.getBool(false)
|
||||
result.defaultTurnsPerSecond = node{"defaultTurnsPerSecond"}.getInt(30)
|
||||
|
||||
proc handleGameStarted(ws: SyncWebSocket; node: JsonNode) =
|
||||
let setup = parseGameSetup(node{"gameSetup"})
|
||||
|
||||
var teammateIds: seq[int] = @[]
|
||||
if not node{"teammateIds"}.isNil and node["teammateIds"].kind == JArray:
|
||||
for id in node["teammateIds"]: teammateIds.add id.getInt
|
||||
|
||||
let myId = node{"myId"}.getInt
|
||||
setGameStarted(myId, setup, teammateIds)
|
||||
|
||||
# Build event object manually — GameStartedEventForBot has no turnNumber
|
||||
let e = GameStartedEventForBot(
|
||||
`type`: "GameStartedEventForBot",
|
||||
myId: myId,
|
||||
startX: node{"startX"}.getFloat(0.0),
|
||||
startY: node{"startY"}.getFloat(0.0),
|
||||
startDirection: node{"startDirection"}.getFloat(0.0),
|
||||
teammateIds: teammateIds,
|
||||
gameSetup: setup
|
||||
)
|
||||
gBot.onGameStarted(e)
|
||||
|
||||
# Send BotReady
|
||||
ws.send("""{"type":"BotReady"}""")
|
||||
|
||||
proc handleTick(node: JsonNode) =
|
||||
# Build TickEventForBot manually to handle optional fields safely
|
||||
var tick: TickEventForBot
|
||||
tick.`type` = "TickEventForBot"
|
||||
tick.turnNumber = node{"turnNumber"}.getInt(0)
|
||||
tick.roundNumber = node{"roundNumber"}.getInt(0)
|
||||
tick.botState = parseBotState(node{"botState"})
|
||||
tick.bulletStates = @[]
|
||||
if not node{"bulletStates"}.isNil and node["bulletStates"].kind == JArray:
|
||||
for bs in node["bulletStates"]:
|
||||
tick.bulletStates.add parseBulletState(bs)
|
||||
tick.events = @[] # sub-events parsed separately into typed BotEvent
|
||||
|
||||
# Parse embedded events into typed BotEvent for priority-based dispatch
|
||||
var events: seq[BotEvent] = @[]
|
||||
let myId = getMyId()
|
||||
if node.hasKey("events") and node["events"].kind == JArray:
|
||||
for ev in node["events"]:
|
||||
events.add parseBotEvent(ev, myId)
|
||||
|
||||
signalTick(tick, events) # update shared state
|
||||
processTickOnMainThread() # motion tracking (while bot is blocked)
|
||||
wakeBotThread() # wake bot — state + motion ready
|
||||
|
||||
proc runReceiveLoop*(ws: SyncWebSocket; info: BotInfo; secret: string; serverUrl: string) =
|
||||
## Main WebSocket receive loop. Blocks until disconnected.
|
||||
while ws.connected:
|
||||
var msg: string
|
||||
try:
|
||||
msg = ws.receive()
|
||||
except Exception as e:
|
||||
stderr.writeLine "[ws] receive error: " & e.msg
|
||||
gBot.onConnectionError(ConnectionErrorEvent(serverUrl: serverUrl, error: e.msg))
|
||||
break
|
||||
|
||||
if msg.len == 0:
|
||||
break # connection closed
|
||||
|
||||
var node: JsonNode
|
||||
try:
|
||||
node = parseJson(msg)
|
||||
except Exception as e:
|
||||
stderr.writeLine "[ws] json parse error: " & e.msg
|
||||
continue
|
||||
|
||||
let msgType = node{"type"}.getStr
|
||||
try:
|
||||
case msgType
|
||||
of "ServerHandshake":
|
||||
handleServerHandshake(ws, node, info, secret)
|
||||
of "GameStartedEventForBot":
|
||||
handleGameStarted(ws, node)
|
||||
of "RoundStartedEvent":
|
||||
let e = node.to(RoundStartedEvent)
|
||||
debugLog("[NS-ENTER] round=" & $e.roundNumber & " tid=" & $getThreadId())
|
||||
# Start (or restart) the bot thread each round
|
||||
startRound()
|
||||
startBotThread()
|
||||
gBot.onRoundStarted(e)
|
||||
debugLog("[NS-EXIT] round=" & $e.roundNumber & " tid=" & $getThreadId())
|
||||
of "TickEventForBot":
|
||||
handleTick(node)
|
||||
of "RoundEndedEventForBot":
|
||||
setRunning(false)
|
||||
let e = node.to(RoundEndedEventForBot)
|
||||
debugLog("[RE-ENTER] round=" & $e.roundNumber & " tid=" & $getThreadId())
|
||||
signalStop() # unblock bot thread blocked in go()
|
||||
debugLog("[WT-ENTER] round=" & $e.roundNumber & " tid=" & $getThreadId())
|
||||
waitForBotThread()
|
||||
debugLog("[WT-EXIT] round=" & $e.roundNumber & " tid=" & $getThreadId())
|
||||
debugLog("[DR-ENTER] round=" & $e.roundNumber & " tid=" & $getThreadId())
|
||||
drainTickChan() # drain stop signal if bot exited via isRunning() check
|
||||
drainIntentChan() # drain AFTER thread joined — no more writes possible
|
||||
drainEventChan() # drop any unconsumed tick events (stale into next round)
|
||||
debugLog("[DR-EXIT] round=" & $e.roundNumber & " tid=" & $getThreadId())
|
||||
debugLog("[ONRE-ENTER] round=" & $e.roundNumber & " tid=" & $getThreadId())
|
||||
gBot.onRoundEnded(e)
|
||||
debugLog("[ONRE-EXIT] round=" & $e.roundNumber & " tid=" & $getThreadId())
|
||||
of "GameEndedEventForBot":
|
||||
setRunning(false)
|
||||
let e = node.to(GameEndedEventForBot)
|
||||
drainIntentChan()
|
||||
gBot.onGameEnded(e)
|
||||
of "GameAbortedEvent":
|
||||
setRunning(false)
|
||||
signalStop() # unblock bot thread (game aborted mid-round)
|
||||
waitForBotThread()
|
||||
drainTickChan() # drain stop signal if bot exited via isRunning() check
|
||||
drainIntentChan() # drain AFTER thread joined — no more writes possible
|
||||
drainEventChan() # drop any unconsumed tick events (stale into next round)
|
||||
gBot.onGameAborted()
|
||||
of "SkippedTurnEvent":
|
||||
let e = node.to(SkippedTurnEvent)
|
||||
gBot.onSkippedTurn(e)
|
||||
else:
|
||||
discard # unknown message type — ignore
|
||||
except Exception as e:
|
||||
# A raised handler/callback (e.g. an OSError from sync training inside
|
||||
# onRoundEnded) must not kill the receive loop — that is the silent
|
||||
# corpse path (process lives, no intents ever again). Log and continue.
|
||||
stderr.writeLine "[ws] handler error (" & msgType & "): " & e.msg
|
||||
debugLog("[WS-HANDLER-ERR] " & msgType & ": " & e.msg)
|
||||
|
||||
# Loop exited: server disconnected or ws error. Make sure the bot thread is
|
||||
# stopped and joined so the process exits cleanly instead of hanging forever
|
||||
# with a blocked bot (corpse). ponytail: signalStop + join; the bot's go()
|
||||
# consumes the stop as a non-tick and exits via the isRunning() check.
|
||||
if isRunning():
|
||||
debugLog("[DBG] receive loop exited while running — stopping bot thread")
|
||||
setRunning(false)
|
||||
signalStop()
|
||||
waitForBotThread()
|
||||
drainTickChan()
|
||||
drainIntentChan()
|
||||
drainEventChan()
|
||||
gBot.onDisconnected(DisconnectedEvent(serverUrl: serverUrl))
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public start() procedure
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
proc start*(bot: Bot; jsonFile: string = "") =
|
||||
## Connect to the server and start the bot.
|
||||
## jsonFile: path to bot JSON profile (optional; falls back to env vars).
|
||||
gBot = bot
|
||||
gBotInfo = loadBotInfo(jsonFile)
|
||||
initGlobals()
|
||||
|
||||
let serverUrl = getEnv("SERVER_URL", "ws://localhost:7654")
|
||||
let serverSecret = getEnv("SERVER_SECRET", "")
|
||||
|
||||
try:
|
||||
gWs = newSyncWebSocket(serverUrl)
|
||||
except Exception as e:
|
||||
stderr.writeLine "[start] Cannot connect to " & serverUrl & ": " & e.msg
|
||||
quit(1)
|
||||
|
||||
bot.onConnected(ConnectedEvent(serverUrl: serverUrl))
|
||||
startSenderThread()
|
||||
runReceiveLoop(gWs, gBotInfo, serverSecret, serverUrl)
|
||||
stopSenderThread()
|
||||
@@ -0,0 +1,11 @@
|
||||
# Package
|
||||
version = "1.0.1"
|
||||
author = "Davide Cappellini"
|
||||
description = "Nim bot API for Robocode Tank Royale"
|
||||
license = "Apache-2.0"
|
||||
srcDir = "src"
|
||||
skipDirs = @["sample_bots"]
|
||||
|
||||
# Dependencies
|
||||
requires "nim >= 2.0.0"
|
||||
requires "jsony >= 1.1.5"
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,79 @@
|
||||
## BotInfo: bot identification loaded from a JSON file or environment variables.
|
||||
|
||||
import std/[os, json, strutils, sequtils]
|
||||
import ./schemas
|
||||
|
||||
type
|
||||
BotInfo* = object
|
||||
name*: string
|
||||
version*: string
|
||||
authors*: seq[string]
|
||||
description*: string
|
||||
homepage*: string
|
||||
countryCodes*: seq[string]
|
||||
gameTypes*: seq[string]
|
||||
platform*: string
|
||||
programmingLang*: string
|
||||
initialPosition*: InitialPosition
|
||||
isDroid*: bool
|
||||
|
||||
proc botInfoFromJson*(path: string): BotInfo =
|
||||
let data = parseJson(readFile(path))
|
||||
result.name = data{"name"}.getStr
|
||||
result.version = data{"version"}.getStr
|
||||
if data.hasKey("authors"):
|
||||
for a in data["authors"]: result.authors.add a.getStr
|
||||
result.description = data{"description"}.getStr
|
||||
result.homepage = data{"homepage"}.getStr
|
||||
if data.hasKey("countryCodes"):
|
||||
for c in data["countryCodes"]: result.countryCodes.add c.getStr
|
||||
if data.hasKey("gameTypes"):
|
||||
for g in data["gameTypes"]: result.gameTypes.add g.getStr
|
||||
result.platform = data{"platform"}.getStr("Nim " & NimVersion)
|
||||
result.programmingLang = data{"programmingLang"}.getStr("Nim")
|
||||
if data.hasKey("initialPosition"):
|
||||
let ip = data["initialPosition"]
|
||||
result.initialPosition.x = ip{"x"}.getFloat
|
||||
result.initialPosition.y = ip{"y"}.getFloat
|
||||
result.initialPosition.direction = ip{"direction"}.getFloat
|
||||
result.isDroid = data{"isDroid"}.getBool(false)
|
||||
|
||||
proc botInfoFromEnv*(): BotInfo =
|
||||
## Fall back to environment variables when no JSON file is given.
|
||||
result.name = getEnv("BOT_NAME", "Unnamed Bot")
|
||||
result.version = getEnv("BOT_VERSION", "1.0")
|
||||
let authorsStr = getEnv("BOT_AUTHORS", "Unknown")
|
||||
result.authors = authorsStr.split(',').mapIt(it.strip)
|
||||
result.description = getEnv("BOT_DESCRIPTION", "")
|
||||
result.homepage = getEnv("BOT_HOMEPAGE", "")
|
||||
let ccStr = getEnv("BOT_COUNTRY_CODES", "")
|
||||
if ccStr.len > 0:
|
||||
result.countryCodes = ccStr.split(',').mapIt(it.strip)
|
||||
let gtStr = getEnv("BOT_GAME_TYPES", "classic,melee,1v1")
|
||||
result.gameTypes = gtStr.split(',').mapIt(it.strip)
|
||||
result.platform = getEnv("BOT_PLATFORM", "Nim " & NimVersion)
|
||||
result.programmingLang = getEnv("BOT_PROGRAMMING_LANG", "Nim")
|
||||
result.isDroid = getEnv("BOT_IS_DROID", "false").toLowerAscii == "true"
|
||||
|
||||
proc loadBotInfo*(jsonFile: string = ""): BotInfo =
|
||||
var resolved = ""
|
||||
if jsonFile.len > 0:
|
||||
if fileExists(jsonFile):
|
||||
resolved = jsonFile
|
||||
else:
|
||||
# Try alongside the executable
|
||||
let appPath = getAppDir() / jsonFile
|
||||
if fileExists(appPath):
|
||||
resolved = appPath
|
||||
|
||||
if resolved.len > 0:
|
||||
result = botInfoFromJson(resolved)
|
||||
else:
|
||||
result = botInfoFromEnv()
|
||||
# Ensure gameTypes has at least one entry
|
||||
if result.gameTypes.len == 0:
|
||||
result.gameTypes = @["classic", "melee", "1v1"]
|
||||
if result.platform.len == 0:
|
||||
result.platform = "Nim " & NimVersion
|
||||
if result.programmingLang.len == 0:
|
||||
result.programmingLang = "Nim"
|
||||
@@ -0,0 +1,205 @@
|
||||
## Color type for Robocode Tank Royale — RGBA packed as uint32 (R<<24|G<<16|B<<8|A),
|
||||
## matching the layout of the Java Color class.
|
||||
|
||||
import std/strutils
|
||||
|
||||
type Color* = distinct uint32
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Factory procs
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
proc fromRgb*(r, g, b: uint8): Color {.inline.} =
|
||||
Color((r.uint32 shl 24) or (g.uint32 shl 16) or (b.uint32 shl 8) or 0xFF)
|
||||
|
||||
proc fromRgba*(r, g, b, a: uint8): Color {.inline.} =
|
||||
Color((r.uint32 shl 24) or (g.uint32 shl 16) or (b.uint32 shl 8) or a.uint32)
|
||||
|
||||
proc fromHex*(s: string): Color =
|
||||
## Parse "#RRGGBB" or "#RRGGBBAA". Raises ValueError on bad input.
|
||||
let h = if s.len > 0 and s[0] == '#': s[1..^1] else: s
|
||||
case h.len
|
||||
of 6:
|
||||
let v = parseHexInt(h)
|
||||
result = Color((v.uint32 shl 8) or 0xFF)
|
||||
of 8:
|
||||
result = Color(parseHexInt(h).uint32)
|
||||
else:
|
||||
raise newException(ValueError, "invalid color string: " & s)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Accessors
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
proc r*(c: Color): uint8 {.inline.} = uint8(c.uint32 shr 24)
|
||||
proc g*(c: Color): uint8 {.inline.} = uint8((c.uint32 shr 16) and 0xFF)
|
||||
proc b*(c: Color): uint8 {.inline.} = uint8((c.uint32 shr 8) and 0xFF)
|
||||
proc a*(c: Color): uint8 {.inline.} = uint8(c.uint32 and 0xFF)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Serialisation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
proc toHex*(c: Color): string =
|
||||
## Returns "#RRGGBB" when alpha=255, "#RRGGBBAA" otherwise.
|
||||
if c.a == 0xFF:
|
||||
result = '#' & toHex(c.r.int, 2) & toHex(c.g.int, 2) & toHex(c.b.int, 2)
|
||||
else:
|
||||
result = '#' & toHex(c.r.int, 2) & toHex(c.g.int, 2) & toHex(c.b.int, 2) & toHex(c.a.int, 2)
|
||||
|
||||
proc `$`*(c: Color): string = c.toHex
|
||||
|
||||
proc `==`*(a, b: Color): bool {.borrow.}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Backward compat: implicit conversion from string literal / variable
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
converter toColor*(s: string): Color = fromHex(s)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Named constants (all 141 from Java Color class)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
const
|
||||
TRANSPARENT* = fromRgba(255, 255, 255, 0)
|
||||
ALICE_BLUE* = fromRgb(240, 248, 255)
|
||||
ANTIQUE_WHITE* = fromRgb(250, 235, 215)
|
||||
AQUA* = fromRgb(0, 255, 255)
|
||||
AQUAMARINE* = fromRgb(127, 255, 212)
|
||||
AZURE* = fromRgb(240, 255, 255)
|
||||
BEIGE* = fromRgb(245, 245, 220)
|
||||
BISQUE* = fromRgb(255, 228, 196)
|
||||
BLACK* = fromRgb(0, 0, 0)
|
||||
BLANCHED_ALMOND* = fromRgb(255, 235, 205)
|
||||
BLUE* = fromRgb(0, 0, 255)
|
||||
BLUE_VIOLET* = fromRgb(138, 43, 226)
|
||||
BROWN* = fromRgb(165, 42, 42)
|
||||
BURLY_WOOD* = fromRgb(222, 184, 135)
|
||||
CADET_BLUE* = fromRgb(95, 158, 160)
|
||||
CHARTREUSE* = fromRgb(127, 255, 0)
|
||||
CHOCOLATE* = fromRgb(210, 105, 30)
|
||||
CORAL* = fromRgb(255, 127, 80)
|
||||
CORNFLOWER_BLUE* = fromRgb(100, 149, 237)
|
||||
CORNSILK* = fromRgb(255, 248, 220)
|
||||
CRIMSON* = fromRgb(220, 20, 60)
|
||||
CYAN* = fromRgb(0, 255, 255)
|
||||
DARK_BLUE* = fromRgb(0, 0, 139)
|
||||
DARK_CYAN* = fromRgb(0, 139, 139)
|
||||
DARK_GOLDENROD* = fromRgb(184, 134, 11)
|
||||
DARK_GRAY* = fromRgb(169, 169, 169)
|
||||
DARK_GREEN* = fromRgb(0, 100, 0)
|
||||
DARK_KHAKI* = fromRgb(189, 183, 107)
|
||||
DARK_MAGENTA* = fromRgb(139, 0, 139)
|
||||
DARK_OLIVE_GREEN* = fromRgb(85, 107, 47)
|
||||
DARK_ORANGE* = fromRgb(255, 140, 0)
|
||||
DARK_ORCHID* = fromRgb(153, 50, 204)
|
||||
DARK_RED* = fromRgb(139, 0, 0)
|
||||
DARK_SALMON* = fromRgb(233, 150, 122)
|
||||
DARK_SEA_GREEN* = fromRgb(143, 188, 139)
|
||||
DARK_SLATE_BLUE* = fromRgb(72, 61, 139)
|
||||
DARK_SLATE_GRAY* = fromRgb(47, 79, 79)
|
||||
DARK_TURQUOISE* = fromRgb(0, 206, 209)
|
||||
DARK_VIOLET* = fromRgb(148, 0, 211)
|
||||
DEEP_PINK* = fromRgb(255, 20, 147)
|
||||
DEEP_SKY_BLUE* = fromRgb(0, 191, 255)
|
||||
DIM_GRAY* = fromRgb(105, 105, 105)
|
||||
DODGER_BLUE* = fromRgb(30, 144, 255)
|
||||
FIREBRICK* = fromRgb(178, 34, 34)
|
||||
FLORAL_WHITE* = fromRgb(255, 250, 240)
|
||||
FOREST_GREEN* = fromRgb(34, 139, 34)
|
||||
FUCHSIA* = fromRgb(255, 0, 255)
|
||||
GAINSBORO* = fromRgb(220, 220, 220)
|
||||
GHOST_WHITE* = fromRgb(248, 248, 255)
|
||||
GOLD* = fromRgb(255, 215, 0)
|
||||
GOLDENROD* = fromRgb(218, 165, 32)
|
||||
GRAY* = fromRgb(128, 128, 128)
|
||||
GREEN* = fromRgb(0, 128, 0)
|
||||
GREEN_YELLOW* = fromRgb(173, 255, 47)
|
||||
HONEYDEW* = fromRgb(240, 255, 240)
|
||||
HOT_PINK* = fromRgb(255, 105, 180)
|
||||
INDIAN_RED* = fromRgb(205, 92, 92)
|
||||
INDIGO* = fromRgb(75, 0, 130)
|
||||
IVORY* = fromRgb(255, 255, 240)
|
||||
KHAKI* = fromRgb(240, 230, 140)
|
||||
LAVENDER* = fromRgb(230, 230, 250)
|
||||
LAVENDER_BLUSH* = fromRgb(255, 240, 245)
|
||||
LAWN_GREEN* = fromRgb(124, 252, 0)
|
||||
LEMON_CHIFFON* = fromRgb(255, 250, 205)
|
||||
LIGHT_BLUE* = fromRgb(173, 216, 230)
|
||||
LIGHT_CORAL* = fromRgb(240, 128, 128)
|
||||
LIGHT_CYAN* = fromRgb(224, 255, 255)
|
||||
LIGHT_GOLDENROD_YELLOW* = fromRgb(250, 250, 210)
|
||||
LIGHT_GRAY* = fromRgb(211, 211, 211)
|
||||
LIGHT_GREEN* = fromRgb(144, 238, 144)
|
||||
LIGHT_PINK* = fromRgb(255, 182, 193)
|
||||
LIGHT_SALMON* = fromRgb(255, 160, 122)
|
||||
LIGHT_SEA_GREEN* = fromRgb(32, 178, 170)
|
||||
LIGHT_SKY_BLUE* = fromRgb(135, 206, 250)
|
||||
LIGHT_SLATE_GRAY* = fromRgb(119, 136, 153)
|
||||
LIGHT_STEEL_BLUE* = fromRgb(176, 196, 222)
|
||||
LIGHT_YELLOW* = fromRgb(255, 255, 224)
|
||||
LIME* = fromRgb(0, 255, 0)
|
||||
LIME_GREEN* = fromRgb(50, 205, 50)
|
||||
LINEN* = fromRgb(250, 240, 230)
|
||||
MAGENTA* = fromRgb(255, 0, 255)
|
||||
MAROON* = fromRgb(128, 0, 0)
|
||||
MEDIUM_AQUAMARINE* = fromRgb(102, 205, 170)
|
||||
MEDIUM_BLUE* = fromRgb(0, 0, 205)
|
||||
MEDIUM_ORCHID* = fromRgb(186, 85, 211)
|
||||
MEDIUM_PURPLE* = fromRgb(147, 112, 219)
|
||||
MEDIUM_SEA_GREEN* = fromRgb(60, 179, 113)
|
||||
MEDIUM_SLATE_BLUE* = fromRgb(123, 104, 238)
|
||||
MEDIUM_SPRING_GREEN* = fromRgb(0, 250, 154)
|
||||
MEDIUM_TURQUOISE* = fromRgb(72, 209, 204)
|
||||
MEDIUM_VIOLET_RED* = fromRgb(199, 21, 133)
|
||||
MIDNIGHT_BLUE* = fromRgb(25, 25, 112)
|
||||
MINT_CREAM* = fromRgb(245, 255, 250)
|
||||
MISTY_ROSE* = fromRgb(255, 228, 225)
|
||||
MOCCASIN* = fromRgb(255, 228, 181)
|
||||
NAVAJO_WHITE* = fromRgb(255, 222, 173)
|
||||
NAVY* = fromRgb(0, 0, 128)
|
||||
OLD_LACE* = fromRgb(253, 245, 230)
|
||||
OLIVE* = fromRgb(128, 128, 0)
|
||||
OLIVE_DRAB* = fromRgb(107, 142, 35)
|
||||
ORANGE* = fromRgb(255, 165, 0)
|
||||
ORANGE_RED* = fromRgb(255, 69, 0)
|
||||
ORCHID* = fromRgb(218, 112, 214)
|
||||
PALE_GOLDENROD* = fromRgb(238, 232, 170)
|
||||
PALE_GREEN* = fromRgb(152, 251, 152)
|
||||
PALE_TURQUOISE* = fromRgb(175, 238, 238)
|
||||
PALE_VIOLET_RED* = fromRgb(219, 112, 147)
|
||||
PAPAYA_WHIP* = fromRgb(255, 239, 213)
|
||||
PEACH_PUFF* = fromRgb(255, 218, 185)
|
||||
PERU* = fromRgb(205, 133, 63)
|
||||
PINK* = fromRgb(255, 192, 203)
|
||||
PLUM* = fromRgb(221, 160, 221)
|
||||
POWDER_BLUE* = fromRgb(176, 224, 230)
|
||||
PURPLE* = fromRgb(128, 0, 128)
|
||||
RED* = fromRgb(255, 0, 0)
|
||||
ROSY_BROWN* = fromRgb(188, 143, 143)
|
||||
ROYAL_BLUE* = fromRgb(65, 105, 225)
|
||||
SADDLE_BROWN* = fromRgb(139, 69, 19)
|
||||
SALMON* = fromRgb(250, 128, 114)
|
||||
SANDY_BROWN* = fromRgb(244, 164, 96)
|
||||
SEA_GREEN* = fromRgb(46, 139, 87)
|
||||
SEA_SHELL* = fromRgb(255, 245, 238)
|
||||
SIENNA* = fromRgb(160, 82, 45)
|
||||
SILVER* = fromRgb(192, 192, 192)
|
||||
SKY_BLUE* = fromRgb(135, 206, 235)
|
||||
SLATE_BLUE* = fromRgb(106, 90, 205)
|
||||
SLATE_GRAY* = fromRgb(112, 128, 144)
|
||||
SNOW* = fromRgb(255, 250, 250)
|
||||
SPRING_GREEN* = fromRgb(0, 255, 127)
|
||||
STEEL_BLUE* = fromRgb(70, 130, 180)
|
||||
TAN* = fromRgb(210, 180, 140)
|
||||
TEAL* = fromRgb(0, 128, 128)
|
||||
THISTLE* = fromRgb(216, 191, 216)
|
||||
TOMATO* = fromRgb(255, 99, 71)
|
||||
TURQUOISE* = fromRgb(64, 224, 208)
|
||||
VIOLET* = fromRgb(238, 130, 238)
|
||||
WHEAT* = fromRgb(245, 222, 179)
|
||||
WHITE* = fromRgb(255, 255, 255)
|
||||
WHITE_SMOKE* = fromRgb(245, 245, 245)
|
||||
YELLOW* = fromRgb(255, 255, 0)
|
||||
YELLOW_GREEN* = fromRgb(154, 205, 50)
|
||||
@@ -0,0 +1,44 @@
|
||||
## Game constants for Robocode Tank Royale Nim bot API
|
||||
|
||||
const
|
||||
# Infinity helpers
|
||||
POSITIVE_INFINITY* = high(float)
|
||||
NEGATIVE_INFINITY* = low(float)
|
||||
|
||||
# Event queue limits
|
||||
MAX_QUEUE_SIZE* = 256
|
||||
MAX_EVENTS_AGE* = 2
|
||||
MIN_VALUE* = low(int32)
|
||||
|
||||
# Event priorities (higher = processed first)
|
||||
PRIORITY_WON_ROUND* = 150
|
||||
PRIORITY_SKIPPED_TURN* = 140
|
||||
PRIORITY_TICK* = 130
|
||||
PRIORITY_CUSTOM* = 120
|
||||
PRIORITY_TEAM_MESSAGE* = 110
|
||||
PRIORITY_BOT_DEATH* = 100
|
||||
PRIORITY_BULLET_HIT_WALL* = 90
|
||||
PRIORITY_BULLET_HIT_BULLET* = 80
|
||||
PRIORITY_BULLET_HIT_BOT* = 70
|
||||
PRIORITY_BULLET_FIRED* = 60
|
||||
PRIORITY_HIT_BY_BULLET* = 50
|
||||
PRIORITY_HIT_WALL* = 40
|
||||
PRIORITY_HIT_BOT* = 30
|
||||
PRIORITY_SCANNED_BOT* = 20
|
||||
PRIORITY_DEATH* = 10
|
||||
|
||||
# Physics
|
||||
ACCELERATION* = 1.0
|
||||
DECELERATION* = -2.0
|
||||
ABS_DECELERATION* = 2.0
|
||||
|
||||
MAX_SPEED* = 8.0
|
||||
MAX_TURN_RATE* = 10.0
|
||||
MAX_GUN_TURN_RATE* = 20.0
|
||||
MAX_RADAR_TURN_RATE* = 45.0
|
||||
|
||||
MAX_FIRE_POWER* = 3.0
|
||||
MIN_FIRE_POWER* = 0.1
|
||||
|
||||
BOT_RADIUS* = 18.0
|
||||
RADAR_RADIUS* = 1200.0
|
||||
@@ -0,0 +1,167 @@
|
||||
## Priority-based event queue for Robocode Tank Royale bot API.
|
||||
## Typed BotEvent variants wrapping schema types, priority-sorted dispatch.
|
||||
|
||||
import std/[algorithm, tables]
|
||||
import ./constants
|
||||
import ./schemas
|
||||
|
||||
type
|
||||
EventKind* = enum
|
||||
ekTick
|
||||
ekSkippedTurn
|
||||
ekBotDeath
|
||||
ekDeath ## self-death; isCritical=true
|
||||
ekBulletFired
|
||||
ekBulletHitBot
|
||||
ekBulletHitBullet
|
||||
ekBulletHitWall
|
||||
ekHitByBullet
|
||||
ekHitBot
|
||||
ekHitWall
|
||||
ekScannedBot
|
||||
ekWonRound
|
||||
ekTeamMessage
|
||||
ekCustom
|
||||
|
||||
Condition* = object
|
||||
name*: string
|
||||
test*: proc(): bool {.closure.}
|
||||
|
||||
BotEvent* = object
|
||||
turnNumber*: int
|
||||
case kind*: EventKind
|
||||
of ekTick: tick*: TickEventForBot
|
||||
of ekSkippedTurn: skippedTurn*: SkippedTurnEvent
|
||||
of ekBotDeath: botDeath*: BotDeathEvent
|
||||
of ekDeath: death*: BotDeathEvent
|
||||
of ekBulletFired: bulletFired*: BulletFiredEvent
|
||||
of ekBulletHitBot: bulletHitBot*: BulletHitBotEvent
|
||||
of ekBulletHitBullet: bulletHitBullet*: BulletHitBulletEvent
|
||||
of ekBulletHitWall: bulletHitWall*: BulletHitWallEvent
|
||||
of ekHitByBullet: hitByBullet*: HitByBulletEvent
|
||||
of ekHitBot: hitBot*: BotHitBotEvent
|
||||
of ekHitWall: hitWall*: BotHitWallEvent
|
||||
of ekScannedBot: scannedBot*: ScannedBotEvent
|
||||
of ekWonRound: wonRound*: WonRoundEvent
|
||||
of ekTeamMessage: teamMessage*: TeamMessageEvent
|
||||
of ekCustom: condition*: Condition
|
||||
|
||||
EventQueue* = object
|
||||
# ponytail: fixed static storage instead of a seq. The queue outlives bot
|
||||
# threads (a fresh thread runs each round), so a heap seq's backing array
|
||||
# is realloc'd by a *different* dead thread's allocator mid-round ->
|
||||
# rawDealloc SIGSEGV in addEvent (7 gdb-confirmed dumps). Static array:
|
||||
# no heap block crosses threads, realloc can never happen.
|
||||
events*: array[MAX_QUEUE_SIZE, BotEvent]
|
||||
eventsLen*: int
|
||||
priorities: Table[EventKind, int] ## runtime-mutable overrides
|
||||
interruptible*: set[EventKind]
|
||||
currentTopEventKind*: EventKind
|
||||
currentTopPriority*: int
|
||||
conditions*: seq[Condition]
|
||||
|
||||
proc priorityOf*(kind: EventKind): int =
|
||||
case kind
|
||||
of ekTick: PRIORITY_TICK
|
||||
of ekSkippedTurn: PRIORITY_SKIPPED_TURN
|
||||
of ekBotDeath: PRIORITY_BOT_DEATH
|
||||
of ekDeath: PRIORITY_DEATH
|
||||
of ekBulletFired: PRIORITY_BULLET_FIRED
|
||||
of ekBulletHitBot: PRIORITY_BULLET_HIT_BOT
|
||||
of ekBulletHitBullet: PRIORITY_BULLET_HIT_BULLET
|
||||
of ekBulletHitWall: PRIORITY_BULLET_HIT_WALL
|
||||
of ekHitByBullet: PRIORITY_HIT_BY_BULLET
|
||||
of ekHitBot: PRIORITY_HIT_BOT
|
||||
of ekHitWall: PRIORITY_HIT_WALL
|
||||
of ekScannedBot: PRIORITY_SCANNED_BOT
|
||||
of ekWonRound: PRIORITY_WON_ROUND
|
||||
of ekTeamMessage: PRIORITY_TEAM_MESSAGE
|
||||
of ekCustom: PRIORITY_CUSTOM
|
||||
|
||||
proc isCritical*(e: BotEvent): bool =
|
||||
e.kind in {ekDeath, ekWonRound, ekSkippedTurn}
|
||||
|
||||
proc initEventQueue*(): EventQueue =
|
||||
result.currentTopPriority = MIN_VALUE
|
||||
|
||||
proc getPriority*(eq: EventQueue; kind: EventKind): int =
|
||||
eq.priorities.getOrDefault(kind, priorityOf(kind))
|
||||
|
||||
proc setPriority*(eq: var EventQueue; kind: EventKind; p: int) =
|
||||
eq.priorities[kind] = p
|
||||
|
||||
proc addEvent*(eq: var EventQueue; e: BotEvent) =
|
||||
if eq.eventsLen < MAX_QUEUE_SIZE:
|
||||
eq.events[eq.eventsLen] = e
|
||||
inc eq.eventsLen
|
||||
|
||||
proc clear*(eq: var EventQueue) =
|
||||
for i in 0 ..< eq.eventsLen:
|
||||
eq.events[i].reset # destroy refcounted payloads before len drops to 0
|
||||
eq.eventsLen = 0
|
||||
eq.currentTopPriority = MIN_VALUE
|
||||
|
||||
proc clearEvents*(eq: var EventQueue) =
|
||||
clear(eq)
|
||||
|
||||
proc removeOldEvents*(eq: var EventQueue; turnNumber: int) =
|
||||
var i = 0
|
||||
while i < eq.eventsLen:
|
||||
if eq.events[i].turnNumber < turnNumber - MAX_EVENTS_AGE and
|
||||
not eq.events[i].isCritical:
|
||||
for j in i ..< eq.eventsLen - 1:
|
||||
eq.events[j] = eq.events[j + 1]
|
||||
dec eq.eventsLen
|
||||
eq.events[eq.eventsLen].reset
|
||||
else:
|
||||
inc i
|
||||
|
||||
proc popFirst*(eq: var EventQueue): BotEvent =
|
||||
## Remove and return the head element (replaces seq delete(0)).
|
||||
if eq.eventsLen == 0: return
|
||||
result = eq.events[0]
|
||||
for i in 0 ..< eq.eventsLen - 1:
|
||||
eq.events[i] = eq.events[i + 1]
|
||||
dec eq.eventsLen
|
||||
eq.events[eq.eventsLen].reset
|
||||
|
||||
proc addCustomEvents*(eq: var EventQueue; turnNumber: int) =
|
||||
for c in eq.conditions:
|
||||
try:
|
||||
if c.test():
|
||||
eq.addEvent(BotEvent(kind: ekCustom, turnNumber: turnNumber, condition: c))
|
||||
except: discard
|
||||
|
||||
proc sortEvents*(eq: var EventQueue) =
|
||||
# ponytail: copy priorities table for closure capture (cheap, overrides are rare)
|
||||
let prio = eq.priorities
|
||||
if eq.eventsLen > 1:
|
||||
eq.events.toOpenArray(0, eq.eventsLen - 1).sort(proc(a, b: BotEvent): int =
|
||||
let dc = b.isCritical.int - a.isCritical.int
|
||||
if dc != 0: return dc
|
||||
let dt = a.turnNumber - b.turnNumber
|
||||
if dt != 0: return dt
|
||||
let pa = prio.getOrDefault(a.kind, priorityOf(a.kind))
|
||||
let pb = prio.getOrDefault(b.kind, priorityOf(b.kind))
|
||||
pb - pa
|
||||
)
|
||||
|
||||
proc setInterruptible*(eq: var EventQueue; kind: EventKind; v: bool) =
|
||||
if v: eq.interruptible.incl kind
|
||||
else: eq.interruptible.excl kind
|
||||
|
||||
proc isInterruptible*(eq: EventQueue; kind: EventKind): bool =
|
||||
kind in eq.interruptible
|
||||
|
||||
proc addCondition*(eq: var EventQueue; c: Condition) =
|
||||
eq.conditions.add c
|
||||
|
||||
proc removeConditionByName*(eq: var EventQueue; name: string) =
|
||||
for i in countdown(eq.conditions.high, 0):
|
||||
if eq.conditions[i].name == name:
|
||||
eq.conditions.del i
|
||||
return
|
||||
|
||||
proc getEvents*(eq: EventQueue): seq[BotEvent] =
|
||||
for i in 0 ..< eq.eventsLen:
|
||||
result.add eq.events[i]
|
||||
@@ -0,0 +1,108 @@
|
||||
## SVG debug graphics API for Robocode Tank Royale Nim bot API.
|
||||
## Module-level procs write SVG into a buffer that is flushed into
|
||||
## BotIntent.debugGraphics each tick and cleared afterward.
|
||||
|
||||
import std/strformat
|
||||
import ./color
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Module-level state (single bot per process)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# ponytail: static char array + length instead of a heap string. A fresh bot
|
||||
# thread runs each round; a module-level string grown by thread N and cleared
|
||||
# ("") by thread N+1 free/reallocs a dead thread's allocator block ->
|
||||
# rawDealloc SIGSEGV (same crash class as the event queue seq; gdb-confirmed
|
||||
# in drawText mid-campaign). Static storage: no heap block crosses threads.
|
||||
const SVG_BUFFER_CAP = 16384
|
||||
|
||||
var gSvgLen: int
|
||||
var gSvgBuffer: array[SVG_BUFFER_CAP, char]
|
||||
var gStrokeColor: Color = WHITE
|
||||
var gFillColor: Color = WHITE
|
||||
var gStrokeWidth: float = 1.0
|
||||
var gFontFamily: string = "Arial" # never rebound at runtime (setFont unused)
|
||||
var gFontSize: float = 12.0
|
||||
|
||||
proc appendSvg(s: string) =
|
||||
## Append an SVG fragment, dropping anything past the static cap.
|
||||
let room = SVG_BUFFER_CAP - gSvgLen
|
||||
if room > 0:
|
||||
let n = min(room, s.len)
|
||||
for i in 0 ..< n: gSvgBuffer[gSvgLen + i] = s[i]
|
||||
inc gSvgLen, n
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
proc svgAttrs(): string =
|
||||
## Current stroke/fill/width as SVG attribute string.
|
||||
&"stroke=\"{gStrokeColor.toHex}\" fill=\"{gFillColor.toHex}\" stroke-width=\"{gStrokeWidth}\""
|
||||
|
||||
proc svgOutput*(): string =
|
||||
## Returns the SVG fragment for this tick, or "" if nothing was drawn.
|
||||
if gSvgLen == 0: return ""
|
||||
"<g>" & $gSvgBuffer[0 ..< gSvgLen] & "</g>"
|
||||
|
||||
proc clearGraphics*() =
|
||||
## Reset buffer and all style globals to defaults. Called after each tick.
|
||||
gSvgLen = 0
|
||||
gStrokeColor = WHITE
|
||||
gFillColor = WHITE
|
||||
gStrokeWidth = 1.0
|
||||
gFontFamily = "Arial"
|
||||
gFontSize = 12.0
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# State setters
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
proc setStrokeColor*(c: Color) = gStrokeColor = c
|
||||
proc setFillColor*(c: Color) = gFillColor = c
|
||||
proc setStrokeWidth*(w: float) = gStrokeWidth = w
|
||||
proc setFont*(family: string; size: float) =
|
||||
gFontFamily = family
|
||||
gFontSize = size
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Draw procs — append raw SVG elements
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
proc drawLine*(x1, y1, x2, y2: float) =
|
||||
appendSvg(&"<line x1=\"{x1}\" y1=\"{y1}\" x2=\"{x2}\" y2=\"{y2}\" {svgAttrs()}/>")
|
||||
|
||||
proc drawRectangle*(x, y, w, h: float) =
|
||||
let attrs = &"stroke=\"{gStrokeColor.toHex}\" fill=\"none\" stroke-width=\"{gStrokeWidth}\""
|
||||
appendSvg(&"<rect x=\"{x}\" y=\"{y}\" width=\"{w}\" height=\"{h}\" {attrs}/>")
|
||||
|
||||
proc fillRectangle*(x, y, w, h: float) =
|
||||
let attrs = &"stroke=\"none\" fill=\"{gFillColor.toHex}\""
|
||||
appendSvg(&"<rect x=\"{x}\" y=\"{y}\" width=\"{w}\" height=\"{h}\" {attrs}/>")
|
||||
|
||||
proc drawCircle*(x, y, r: float) =
|
||||
let attrs = &"stroke=\"{gStrokeColor.toHex}\" fill=\"none\" stroke-width=\"{gStrokeWidth}\""
|
||||
appendSvg(&"<circle cx=\"{x}\" cy=\"{y}\" r=\"{r}\" {attrs}/>")
|
||||
|
||||
proc fillCircle*(x, y, r: float) =
|
||||
let attrs = &"stroke=\"none\" fill=\"{gFillColor.toHex}\""
|
||||
appendSvg(&"<circle cx=\"{x}\" cy=\"{y}\" r=\"{r}\" {attrs}/>")
|
||||
|
||||
proc drawText*(text: string; x, y: float) =
|
||||
appendSvg(&"<text x=\"{x}\" y=\"{y}\" font-family=\"{gFontFamily}\" font-size=\"{gFontSize}\">{text}</text>")
|
||||
|
||||
proc drawPolygon*(points: seq[(float, float)]) =
|
||||
var pts = ""
|
||||
for (px, py) in points:
|
||||
if pts.len > 0: pts.add ' '
|
||||
pts.add &"{px},{py}"
|
||||
let attrs = &"stroke=\"{gStrokeColor.toHex}\" fill=\"none\" stroke-width=\"{gStrokeWidth}\""
|
||||
appendSvg(&"<polygon points=\"{pts}\" {attrs}/>")
|
||||
|
||||
proc fillPolygon*(points: seq[(float, float)]) =
|
||||
var pts = ""
|
||||
for (px, py) in points:
|
||||
if pts.len > 0: pts.add ' '
|
||||
pts.add &"{px},{py}"
|
||||
let attrs = &"stroke=\"none\" fill=\"{gFillColor.toHex}\""
|
||||
appendSvg(&"<polygon points=\"{pts}\" {attrs}/>")
|
||||
@@ -0,0 +1,106 @@
|
||||
## Safe JSON-to-type parsing for Robocode Tank Royale protocol.
|
||||
##
|
||||
## Uses the {} accessor (returns nil on missing keys) and typed getters with
|
||||
## default values so that optional schema fields never raise KeyError.
|
||||
|
||||
import std/json
|
||||
import ./schemas
|
||||
import ./color
|
||||
import ./event_queue
|
||||
|
||||
proc parseBulletState*(node: JsonNode): BulletState =
|
||||
## Parse a BulletState from JSON; missing optional fields default to zero.
|
||||
if node.isNil: return
|
||||
result.bulletId = node{"bulletId"}.getInt(0)
|
||||
result.ownerId = node{"ownerId"}.getInt(0)
|
||||
result.power = node{"power"}.getFloat(0.0)
|
||||
result.x = node{"x"}.getFloat(0.0)
|
||||
result.y = node{"y"}.getFloat(0.0)
|
||||
result.direction = node{"direction"}.getFloat(0.0)
|
||||
let bulletColorStr = node{"color"}.getStr("")
|
||||
result.color = if bulletColorStr.len > 0: fromHex(bulletColorStr) else: Color(0)
|
||||
|
||||
proc parseBotState*(node: JsonNode): BotState =
|
||||
## Parse a BotState from JSON; optional colour/flag fields default to empty/false.
|
||||
if node.isNil: return
|
||||
result.isDroid = node{"isDroid"}.getBool(false)
|
||||
result.energy = node{"energy"}.getFloat(0.0)
|
||||
result.x = node{"x"}.getFloat(0.0)
|
||||
result.y = node{"y"}.getFloat(0.0)
|
||||
result.direction = node{"direction"}.getFloat(0.0)
|
||||
result.gunDirection = node{"gunDirection"}.getFloat(0.0)
|
||||
result.radarDirection = node{"radarDirection"}.getFloat(0.0)
|
||||
result.radarSweep = node{"radarSweep"}.getFloat(0.0)
|
||||
result.speed = node{"speed"}.getFloat(0.0)
|
||||
result.turnRate = node{"turnRate"}.getFloat(0.0)
|
||||
result.gunTurnRate = node{"gunTurnRate"}.getFloat(0.0)
|
||||
result.radarTurnRate = node{"radarTurnRate"}.getFloat(0.0)
|
||||
result.gunHeat = node{"gunHeat"}.getFloat(0.0)
|
||||
result.enemyCount = node{"enemyCount"}.getInt(0)
|
||||
template parseColor(field: untyped) =
|
||||
let s = node{astToStr(field)}.getStr("")
|
||||
result.field = if s.len > 0: fromHex(s) else: Color(0)
|
||||
parseColor(bodyColor)
|
||||
parseColor(turretColor)
|
||||
parseColor(radarColor)
|
||||
parseColor(bulletColor)
|
||||
parseColor(scanColor)
|
||||
parseColor(tracksColor)
|
||||
parseColor(gunColor)
|
||||
|
||||
proc parseBotEvent*(node: JsonNode; myId: int): BotEvent =
|
||||
## Parse a JSON event node into a typed BotEvent.
|
||||
let typeStr = node{"type"}.getStr
|
||||
let tn = node{"turnNumber"}.getInt(0)
|
||||
case typeStr
|
||||
of "BotDeathEvent":
|
||||
let victimId = node{"victimId"}.getInt(0)
|
||||
if victimId == myId:
|
||||
result = BotEvent(kind: ekDeath, turnNumber: tn,
|
||||
death: BotDeathEvent(`type`: typeStr, turnNumber: tn, victimId: victimId))
|
||||
else:
|
||||
result = BotEvent(kind: ekBotDeath, turnNumber: tn,
|
||||
botDeath: BotDeathEvent(`type`: typeStr, turnNumber: tn, victimId: victimId))
|
||||
of "BulletFiredEvent":
|
||||
result = BotEvent(kind: ekBulletFired, turnNumber: tn,
|
||||
bulletFired: BulletFiredEvent(`type`: typeStr, turnNumber: tn,
|
||||
bullet: parseBulletState(node{"bullet"})))
|
||||
of "BulletHitBotEvent":
|
||||
let victimId = node{"victimId"}.getInt(0)
|
||||
let bullet = parseBulletState(node{"bullet"})
|
||||
let damage = node{"damage"}.getFloat(0.0)
|
||||
let energy = node{"energy"}.getFloat(0.0)
|
||||
if victimId == myId:
|
||||
result = BotEvent(kind: ekHitByBullet, turnNumber: tn,
|
||||
hitByBullet: HitByBulletEvent(`type`: "HitByBulletEvent", turnNumber: tn,
|
||||
bullet: bullet, damage: damage, energy: energy))
|
||||
else:
|
||||
result = BotEvent(kind: ekBulletHitBot, turnNumber: tn,
|
||||
bulletHitBot: BulletHitBotEvent(`type`: typeStr, turnNumber: tn,
|
||||
victimId: victimId, bullet: bullet, damage: damage, energy: energy))
|
||||
of "BulletHitBulletEvent":
|
||||
result = BotEvent(kind: ekBulletHitBullet, turnNumber: tn,
|
||||
bulletHitBullet: BulletHitBulletEvent(`type`: typeStr, turnNumber: tn,
|
||||
bullet: parseBulletState(node{"bullet"}),
|
||||
hitBullet: parseBulletState(node{"hitBullet"})))
|
||||
of "BulletHitWallEvent":
|
||||
result = BotEvent(kind: ekBulletHitWall, turnNumber: tn,
|
||||
bulletHitWall: BulletHitWallEvent(`type`: typeStr, turnNumber: tn,
|
||||
bullet: parseBulletState(node{"bullet"})))
|
||||
of "BotHitBotEvent":
|
||||
result = BotEvent(kind: ekHitBot, turnNumber: tn,
|
||||
hitBot: node.to(BotHitBotEvent))
|
||||
of "BotHitWallEvent":
|
||||
result = BotEvent(kind: ekHitWall, turnNumber: tn,
|
||||
hitWall: node.to(BotHitWallEvent))
|
||||
of "ScannedBotEvent":
|
||||
result = BotEvent(kind: ekScannedBot, turnNumber: tn,
|
||||
scannedBot: node.to(ScannedBotEvent))
|
||||
of "WonRoundEvent":
|
||||
result = BotEvent(kind: ekWonRound, turnNumber: tn,
|
||||
wonRound: WonRoundEvent(`type`: typeStr, turnNumber: tn))
|
||||
of "TeamMessageEvent":
|
||||
result = BotEvent(kind: ekTeamMessage, turnNumber: tn,
|
||||
teamMessage: node.to(TeamMessageEvent))
|
||||
else:
|
||||
discard
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user