feat(PPO_Bot): command abstraction layer — goto/aimTo controllers (#24)

- Add gotoTick/aimToTick controller functions (#25)
- Update network dims: actor 5→6, state 42→44 (#26)
- Rewrite mapActions for 6-dim command space (#27)
- Delete stale weight files (shape mismatch)
- Fix existing tests for new signatures

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-08-17 19:08:31 +02:00
parent 56e0b306c9
commit cdde60d79f
39 changed files with 490 additions and 78 deletions
+3
View File
@@ -0,0 +1,3 @@
nimble.develop
nimble.paths
nimbledeps
BIN
View File
Binary file not shown.
+11
View File
@@ -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"
}
+95
View File
@@ -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)
+10
View File
@@ -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"
+4
View File
@@ -0,0 +1,4 @@
# begin Nimble config (version 2)
when withDir(thisDir(), system.fileExists("nimble.paths")):
include "nimble.paths"
# end Nimble config
BIN
View File
Binary file not shown.
+22 -4
View File
@@ -24,6 +24,7 @@ type PPOBot = ref object of Bot
lastLogP: float32 lastLogP: float32
lastValue: float32 lastValue: float32
hasLastTrans: bool hasLastTrans: bool
lastActions: BotActions # previous tick's decoded actions (for state vector)
roundRewardSum: float32 # cumulative reward this round (for live display) roundRewardSum: float32 # cumulative reward this round (for live display)
roundTicks: int # ticks this round roundTicks: int # ticks this round
@@ -36,6 +37,7 @@ type
TrainingResult = object TrainingResult = object
ac: ActorCritic ac: ActorCritic
adamStates: ACAdamStates adamStates: ACAdamStates
metrics: PPOMetrics
TrainingArgs = object TrainingArgs = object
ac: ActorCritic ac: ActorCritic
@@ -54,9 +56,9 @@ var
proc trainingThreadProc(args: TrainingArgs) {.thread.} = proc trainingThreadProc(args: TrainingArgs) {.thread.} =
var localAc = args.ac var localAc = args.ac
var localAdam = args.adamStates var localAdam = args.adamStates
ppoUpdate(localAc, args.buffer, lastValue = args.lastValue, adamStates = localAdam) let m = ppoUpdate(localAc, args.buffer, lastValue = args.lastValue, adamStates = localAdam)
saveCheckpoint(localAc, args.weightsRoot, args.roundNum) saveCheckpoint(localAc, args.weightsRoot, args.roundNum)
resultChan.send(TrainingResult(ac: localAc, adamStates: localAdam)) resultChan.send(TrainingResult(ac: localAc, adamStates: localAdam, metrics: m))
# ── Bot methods ─────────────────────────────────────────────────────────────── # ── Bot methods ───────────────────────────────────────────────────────────────
@@ -91,6 +93,7 @@ method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
let avgR = rewardSum / ticks.float32 let avgR = rewardSum / ticks.float32
let avgRStr = formatFloat(avgR.float, ffDecimal, 3) let avgRStr = formatFloat(avgR.float, ffDecimal, 3)
printToStdOut(&"R:{roundCounter} ticks:{ticks} avgR:{avgRStr} score:{e.results.totalScore}\n") printToStdOut(&"R:{roundCounter} ticks:{ticks} avgR:{avgRStr} score:{e.results.totalScore}\n")
echo &"R:{roundCounter} ticks:{ticks} avgR:{avgRStr} score:{e.results.totalScore}"
# Pick up result from previous training thread if available; channel IS the sync # Pick up result from previous training thread if available; channel IS the sync
if threadLaunched: if threadLaunched:
@@ -99,6 +102,9 @@ method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
ac = trained.ac ac = trained.ac
gAdamStates = trained.adamStates gAdamStates = trained.adamStates
threadLaunched = false threadLaunched = false
let m = trained.metrics
printToStdOut(&" trained R:{roundCounter-1} 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-1} aLoss:{formatFloat(m.actorLoss.float, ffDecimal, 4)} vLoss:{formatFloat(m.valueLoss.float, ffDecimal, 4)} gNorm:{formatFloat(m.gradNorm.float, ffDecimal, 3)}"
if bot.buffer.len == 0: if bot.buffer.len == 0:
bot.hasLastTrans = false bot.hasLastTrans = false
@@ -122,6 +128,8 @@ method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
bot.buffer.clear() bot.buffer.clear()
bot.hasLastTrans = false bot.hasLastTrans = false
printToStdOut(&" train→ R:{roundCounter} ticks:{ticks}\n")
echo &" train→ R:{roundCounter} ticks:{ticks}"
createThread(trainingThread, trainingThreadProc, args) createThread(trainingThread, trainingThreadProc, args)
threadLaunched = true threadLaunched = true
@@ -147,10 +155,19 @@ method run(bot: PPOBot) =
arenaHeight: float64(getArenaHeight()), arenaHeight: float64(getArenaHeight()),
) )
let state = buildStateVector(botData, bot.tracker) 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)
let (rawActs, logP) = ac.actorForward(state) let (rawActs, logP) = ac.actorForward(state)
let value = ac.criticForward(state) let value = ac.criticForward(state)
let acts = mapActions(rawActs, getSpeed().float32, getGunHeat().float32) let acts = mapActions(rawActs,
getGunHeat().float,
botData.arenaWidth, botData.arenaHeight,
botData.x, botData.y,
botData.direction, botData.speed, botData.gunDirection)
# Compute tick reward from energy deltas # Compute tick reward from energy deltas
let curEnergy = getEnergy().float32 let curEnergy = getEnergy().float32
@@ -182,6 +199,7 @@ method run(bot: PPOBot) =
bot.prevEnergy = curEnergy bot.prevEnergy = curEnergy
bot.prevEnemyE = curEnemyE bot.prevEnemyE = curEnemyE
bot.hasLastTrans = true bot.hasLastTrans = true
bot.lastActions = acts
setTargetSpeed(acts.targetSpeed.float) setTargetSpeed(acts.targetSpeed.float)
setTurnRate(acts.turnRate.float) setTurnRate(acts.turnRate.float)
+35 -23
View File
@@ -2,32 +2,44 @@
import arraymancer import arraymancer
import std/math import std/math
import ./controllers
func sigmoid(x: float): float = 1.0 / (1.0 + exp(-x))
type type
BotActions* = object BotActions* = object
targetSpeed*: float32 targetSpeed*: float
turnRate*: float32 turnRate*: float
gunTurnRate*: float32 gunTurnRate*: float
shouldFire*: bool shouldFire*: bool
firePower*: float32 firePower*: float
gotoX*: float
gotoY*: float
aimToX*: float
aimToY*: float
proc mapActions*(rawActions: Tensor[float32], currentSpeed: float32, gunHeat: float32): BotActions = proc mapActions*(rawActions: Tensor[float32],
## rawActions: [5] tensor from actorForward. gunHeat: float,
## Ranges: arenaWidth, arenaHeight: float,
## targetSpeed ∈ [-8, 8] botX, botY, direction, speed, gunDirection: float): BotActions =
## turnRate ∈ [-(10 - 0.75|v|), (10 - 0.75|v|)] ## rawActions: [6] tensor from actorForward.
## gunTurnRate ∈ [-20, 20] ## Dims 0–1: goto x/y, 2–3: aimTo x/y, 4: fire decision, 5: fire power.
## firePower ∈ [0.1, 3.0] let gotoX = sigmoid(rawActions[0].float) * arenaWidth
let r0 = tanh(rawActions[0].float64).float32 let gotoY = sigmoid(rawActions[1].float) * arenaHeight
let r1 = tanh(rawActions[1].float64).float32 let aimToX = sigmoid(rawActions[2].float) * arenaWidth
let r2 = tanh(rawActions[2].float64).float32 let aimToY = sigmoid(rawActions[3].float) * arenaHeight
let r3 = tanh(rawActions[3].float64).float32 let fireDec = tanh(rawActions[4].float)
let r4 = rawActions[4] let fp = sigmoid(rawActions[5].float) * 2.9 + 0.1
result.targetSpeed = r0 * 8.0'f32 let (ts, tr) = gotoTick(gotoX, gotoY, botX, botY, direction, speed)
result.turnRate = r1 * (10.0'f32 - 0.75'f32 * abs(currentSpeed)) let gtr = aimToTick(aimToX, aimToY, botX, botY, gunDirection)
result.gunTurnRate = r2 * 20.0'f32
let fireDecision = r3 result.gotoX = gotoX
result.shouldFire = fireDecision >= 0.0'f32 and gunHeat <= 0.0'f32 result.gotoY = gotoY
# sigmoid(r4) * 2.9 + 0.1 → [0.1, 3.0] result.aimToX = aimToX
result.firePower = (1.0'f32 / (1.0'f32 + exp(-r4.float64).float32)) * 2.9'f32 + 0.1'f32 result.aimToY = aimToY
result.targetSpeed = ts
result.turnRate = tr
result.gunTurnRate = gtr
result.shouldFire = fireDec >= 0.0 and gunHeat <= 0.0
result.firePower = fp
+24
View File
@@ -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)
+13 -9
View File
@@ -3,6 +3,10 @@
import arraymancer import arraymancer
import std/[math, random] import std/[math, random]
const
STATE_DIM* = 44
ACTION_DIM* = 6
type type
MLP* = object MLP* = object
w1*, b1*: Tensor[float32] # [hidden, input], [hidden] w1*, b1*: Tensor[float32] # [hidden, input], [hidden]
@@ -12,7 +16,7 @@ type
ActorCritic* = object ActorCritic* = object
actor*: MLP actor*: MLP
critic*: MLP critic*: MLP
logStd*: Tensor[float32] # [5] — one per action dim logStd*: Tensor[float32] # [ACTION_DIM] — one per action dim
proc initMLP*(inputDim, hiddenDim, outputDim: int): MLP = proc initMLP*(inputDim, hiddenDim, outputDim: int): MLP =
# Xavier/He-style init: scale weights by sqrt(2/fan_in) # Xavier/He-style init: scale weights by sqrt(2/fan_in)
@@ -24,9 +28,9 @@ proc initMLP*(inputDim, hiddenDim, outputDim: int): MLP =
result.b3 = zeros[float32](outputDim) result.b3 = zeros[float32](outputDim)
proc initActorCritic*(): ActorCritic = proc initActorCritic*(): ActorCritic =
result.actor = initMLP(42, 64, 5) result.actor = initMLP(STATE_DIM, 64, ACTION_DIM)
result.critic = initMLP(42, 64, 1) result.critic = initMLP(STATE_DIM, 64, 1)
result.logStd = zeros[float32](5) # init to 0 → std=1 result.logStd = zeros[float32](ACTION_DIM) # init to 0 → std=1
proc forward*(mlp: MLP, x: Tensor[float32]): Tensor[float32] = proc forward*(mlp: MLP, x: Tensor[float32]): Tensor[float32] =
## x shape: [inputDim] (1D vector) ## x shape: [inputDim] (1D vector)
@@ -35,15 +39,15 @@ proc forward*(mlp: MLP, x: Tensor[float32]): Tensor[float32] =
result = mlp.w3 * h2 + mlp.b3 result = mlp.w3 * h2 + mlp.b3
proc actorForward*(ac: ActorCritic, state: Tensor[float32]): tuple[actions: Tensor[float32], logProb: float32] = proc actorForward*(ac: ActorCritic, state: Tensor[float32]): tuple[actions: Tensor[float32], logProb: float32] =
## state: [42]. Returns sampled actions [5] and sum log-prob. ## state: [STATE_DIM]. Returns sampled actions [ACTION_DIM] and sum log-prob.
let mean = ac.actor.forward(state) let mean = ac.actor.forward(state)
# Floor logStd at -3 before exp → min std ≈ 0.05 # Floor logStd at -3 before exp → min std ≈ 0.05
var clampedLogStd = ac.logStd.map(proc(v: float32): float32 = max(v, -3.0'f32)) var clampedLogStd = ac.logStd.map(proc(v: float32): float32 = max(v, -3.0'f32))
let std = clampedLogStd.map(proc(v: float32): float32 = exp(v)) let std = clampedLogStd.map(proc(v: float32): float32 = exp(v))
var actions = newTensor[float32](5) var actions = newTensor[float32](ACTION_DIM)
var logP = 0.0'f32 var logP = 0.0'f32
for i in 0..<5: for i in 0..<ACTION_DIM:
let mu = mean[i] let mu = mean[i]
let s = std[i] let s = std[i]
let z = gauss(0.0'f64, 1.0'f64).float32 let z = gauss(0.0'f64, 1.0'f64).float32
@@ -55,7 +59,7 @@ proc actorForward*(ac: ActorCritic, state: Tensor[float32]): tuple[actions: Tens
result = (actions: actions, logProb: logP) result = (actions: actions, logProb: logP)
proc criticForward*(ac: ActorCritic, state: Tensor[float32]): float32 = proc criticForward*(ac: ActorCritic, state: Tensor[float32]): float32 =
## state: [42]. Returns scalar value estimate. ## state: [STATE_DIM]. Returns scalar value estimate.
let val = ac.critic.forward(state) let val = ac.critic.forward(state)
result = val[0] result = val[0]
@@ -65,7 +69,7 @@ proc computeLogProb*(ac: ActorCritic, state, action: Tensor[float32]): float32 =
let logStdClamped = ac.logStd.map(proc(v: float32): float32 = max(v, -3.0'f32)) let logStdClamped = ac.logStd.map(proc(v: float32): float32 = max(v, -3.0'f32))
let std = logStdClamped.map(proc(v: float32): float32 = exp(v)) let std = logStdClamped.map(proc(v: float32): float32 = exp(v))
var logP = 0.0'f32 var logP = 0.0'f32
for i in 0..<5: for i in 0..<ACTION_DIM:
let mu = mean[i] let mu = mean[i]
let s = std[i] let s = std[i]
let diff = (action[i] - mu) / s let diff = (action[i] - mu) / s
+10 -4
View File
@@ -1,4 +1,4 @@
## State vector builder — produces 42-float normalized tensor for PPO policy. ## State vector builder — produces 44-float normalized tensor for PPO policy.
## No bot API imports; takes plain BotState + EnemyTracker structs. ## No bot API imports; takes plain BotState + EnemyTracker structs.
import std/math import std/math
@@ -16,10 +16,12 @@ type
gunHeat*: float64 gunHeat*: float64
arenaWidth*, arenaHeight*: float64 arenaWidth*, arenaHeight*: float64
proc buildStateVector*(bot: BotStateData; enemy: EnemyTracker): Tensor[float32] = proc buildStateVector*(bot: BotStateData; enemy: EnemyTracker;
## Build the 42-float normalized state tensor. remainingGotoDistance: float64 = 0.0;
remainingGunAngle: float64 = 0.0): Tensor[float32] =
## Build the 44-float normalized state tensor.
## All values clipped to roughly [-1, 1] via division by physical maxima. ## All values clipped to roughly [-1, 1] via division by physical maxima.
result = zeros[float32](42) result = zeros[float32](44)
let aW = bot.arenaWidth let aW = bot.arenaWidth
let aH = bot.arenaHeight let aH = bot.arenaHeight
@@ -85,3 +87,7 @@ proc buildStateVector*(bot: BotStateData; enemy: EnemyTracker): Tensor[float32]
result[base + 2] = float32(enemy.history[i].direction / 360.0) result[base + 2] = float32(enemy.history[i].direction / 360.0)
result[base + 3] = float32(enemy.history[i].speed / 8.0) result[base + 3] = float32(enemy.history[i].speed / 8.0)
# else: remain 0.0 (pad) # else: remain 0.0 (pad)
# --- Goto controller inputs (indices 42-43) ---
result[42] = float32(remainingGotoDistance / diag)
result[43] = float32(remainingGunAngle / 180.0)
BIN
View File
Binary file not shown.
+69
View File
@@ -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)
# 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)
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)
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)
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)
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)
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)
let actsMax = mapActions(rawMax, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.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"
BIN
View File
Binary file not shown.
+54
View File
@@ -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"
+11 -14
View File
@@ -11,16 +11,16 @@ func isNaNF(x: float32): bool = classify(x) == fcNan
when isMainModule: when isMainModule:
# ---- MLP forward shape ---- # ---- MLP forward shape ----
let mlp = initMLP(42, 64, 5) let mlp = initMLP(STATE_DIM, 64, ACTION_DIM)
let inp = zeros[float32](42) let inp = zeros[float32](STATE_DIM)
let mlpOut = mlp.forward(inp) let mlpOut = mlp.forward(inp)
assert mlpOut.shape[0] == 5, &"MLP output shape wrong: {mlpOut.shape}" assert mlpOut.shape[0] == ACTION_DIM, &"MLP output shape wrong: {mlpOut.shape}"
# ---- ActorCritic actorForward ---- # ---- ActorCritic actorForward ----
let ac = initActorCritic() let ac = initActorCritic()
let state = zeros[float32](42) let state = zeros[float32](STATE_DIM)
let (acts, logP) = ac.actorForward(state) let (acts, logP) = ac.actorForward(state)
assert acts.shape[0] == 5, &"actorForward actions shape wrong: {acts.shape}" assert acts.shape[0] == ACTION_DIM, &"actorForward actions shape wrong: {acts.shape}"
assert not isNaNF(logP), "logProb is NaN" assert not isNaNF(logP), "logProb is NaN"
assert isFiniteF(logP), &"logProb not finite: {logP}" assert isFiniteF(logP), &"logProb not finite: {logP}"
@@ -31,23 +31,20 @@ when isMainModule:
# ---- logStd floor: collapsing logStd should not break actorForward ---- # ---- logStd floor: collapsing logStd should not break actorForward ----
var ac2 = initActorCritic() var ac2 = initActorCritic()
for i in 0..<5: ac2.logStd[i] = -10.0'f32 for i in 0..<ACTION_DIM: ac2.logStd[i] = -10.0'f32
let (acts2, logP2) = ac2.actorForward(state) let (acts2, logP2) = ac2.actorForward(state)
assert acts2.shape[0] == 5, "acts2 shape wrong after logStd=-10" assert acts2.shape[0] == ACTION_DIM, "acts2 shape wrong after logStd=-10"
assert isFiniteF(logP2), &"logP2 not finite with floored logStd: {logP2}" assert isFiniteF(logP2), &"logP2 not finite with floored logStd: {logP2}"
# ---- action mapping ranges ---- # ---- action mapping ranges ----
let raw = randomNormalTensor[float32](5) let raw = randomNormalTensor[float32](ACTION_DIM)
let speed = 4.0'f32 let speed = 4.0'f32
let botActs = mapActions(raw, speed, 0.0'f32) # gunHeat=0 → fire allowed # 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) # gunHeat=0 → fire allowed
assert botActs.targetSpeed >= -8.0'f32 and botActs.targetSpeed <= 8.0'f32, assert botActs.targetSpeed >= -8.0'f32 and botActs.targetSpeed <= 8.0'f32,
&"targetSpeed out of range: {botActs.targetSpeed}" &"targetSpeed out of range: {botActs.targetSpeed}"
let maxTurn = 10.0'f32 - 0.75'f32 * abs(speed) # = 7.0
assert botActs.turnRate >= -maxTurn and botActs.turnRate <= maxTurn,
&"turnRate out of range: {botActs.turnRate}"
assert botActs.gunTurnRate >= -20.0'f32 and botActs.gunTurnRate <= 20.0'f32, assert botActs.gunTurnRate >= -20.0'f32 and botActs.gunTurnRate <= 20.0'f32,
&"gunTurnRate out of range: {botActs.gunTurnRate}" &"gunTurnRate out of range: {botActs.gunTurnRate}"
@@ -55,7 +52,7 @@ when isMainModule:
&"firePower out of range: {botActs.firePower}" &"firePower out of range: {botActs.firePower}"
# shouldFire=false when gunHeat > 0 # shouldFire=false when gunHeat > 0
let noFire = mapActions(raw, speed, 1.0'f32) let noFire = mapActions(raw, 1.0'f32, 800.0, 600.0, 400.0, 300.0, 0.0, speed.float, 0.0)
assert not noFire.shouldFire, "shouldFire should be false when gunHeat > 0" assert not noFire.shouldFire, "shouldFire should be false when gunHeat > 0"
echo "All tests passed" echo "All tests passed"
+4 -4
View File
@@ -10,16 +10,16 @@ template check(cond: bool, msg: string) =
# Bot at arena center; stationary enemy due north (same x, higher y). # Bot at arena center; stationary enemy due north (same x, higher y).
# Tank Royale: y increases northward. # Tank Royale: y increases northward.
# True bearing to enemy = 0° (north = 0° in game coords). # Radar uses math convention (0°=east, CCW+); north = 90° in that system.
let botX = 400.0 let botX = 400.0
let botY = 300.0 let botY = 300.0
let enemyX = 400.0 # same x → dx = 0 let enemyX = 400.0 # same x → dx = 0
let enemyY = 500.0 # north of bot → dy > 0 let enemyY = 500.0 # north of bot → dy > 0
let trueBearing = 0.0 # north let trueBearing = 90.0 # north in math convention (0=east, CCW+)
# Radar starts pointing at the enemy (radarDirection = 0°, due north). # Radar starts pointing at the enemy (radarDirection = 90°, due north in math convention).
var radarDir = 0.0 var radarDir = 90.0
var tracker = initEnemyTracker() var tracker = initEnemyTracker()
# Prime with contact at the known position # Prime with contact at the known position
Binary file not shown.
+5 -2
View File
@@ -84,7 +84,7 @@ block testStateVectorLength:
arenaWidth: 800.0, arenaHeight: 600.0, arenaWidth: 800.0, arenaHeight: 600.0,
) )
let sv = buildStateVector(bot, t) let sv = buildStateVector(bot, t)
check sv.shape == [42], "state vector has 42 elements" check sv.shape == [44], "state vector has 44 elements"
block testStateVectorRange: block testStateVectorRange:
var t = initEnemyTracker() var t = initEnemyTracker()
@@ -95,7 +95,7 @@ block testStateVectorRange:
arenaWidth: 800.0, arenaHeight: 600.0, arenaWidth: 800.0, arenaHeight: 600.0,
) )
let sv = buildStateVector(bot, t) let sv = buildStateVector(bot, t)
for i in 0 ..< 42: for i in 0 ..< 44:
check sv[i] >= -2.0f32 and sv[i] <= 2.0f32, check sv[i] >= -2.0f32 and sv[i] <= 2.0f32,
&"sv[{i}]={sv[i]} out of [-2,2] range" &"sv[{i}]={sv[i]} out of [-2,2] range"
@@ -155,5 +155,8 @@ block testHistoryPaddedWhenEmpty:
let sv = buildStateVector(bot, t) let sv = buildStateVector(bot, t)
for i in 22 ..< 42: for i in 22 ..< 42:
check sv[i] == 0.0f32, &"history slot {i} should be 0 when no contact" 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"
echo "All tests passed" echo "All tests passed"
+4 -4
View File
@@ -29,7 +29,7 @@ block testBuffer:
var buf = initTrajectoryBuffer() var buf = initTrajectoryBuffer()
check buf.len == 0, "empty buffer len == 0" check buf.len == 0, "empty buffer len == 0"
let t1 = Transition(state: zeros[float32](42), action: zeros[float32](5), let t1 = Transition(state: zeros[float32](STATE_DIM), action: zeros[float32](ACTION_DIM),
logProb: -1.0'f32, reward: 0.5'f32, value: 0.3'f32) logProb: -1.0'f32, reward: 0.5'f32, value: 0.3'f32)
buf.add(t1) buf.add(t1)
buf.add(t1) buf.add(t1)
@@ -80,14 +80,14 @@ block testPpoUpdate:
var buf = initTrajectoryBuffer() var buf = initTrajectoryBuffer()
for _ in 0..<10: for _ in 0..<10:
let s = randomNormalTensor[float32](42) let s = randomNormalTensor[float32](STATE_DIM)
let a = randomNormalTensor[float32](5) let a = randomNormalTensor[float32](ACTION_DIM)
let lp = ac.computeLogProb(s, a) let lp = ac.computeLogProb(s, a)
let v = ac.criticForward(s) let v = ac.criticForward(s)
buf.add(Transition(state: s, action: a, logProb: lp, reward: 0.1'f32, value: v)) buf.add(Transition(state: s, action: a, logProb: lp, reward: 0.1'f32, value: v))
var adam: ACAdamStates var adam: ACAdamStates
ppoUpdate(ac, buf, lastValue = 0.0'f32, adamStates = adam, epochs = 2, miniBatchSize = 5) discard ppoUpdate(ac, buf, lastValue = 0.0'f32, adamStates = adam, epochs = 2, miniBatchSize = 5)
# Weights should have changed — compare flattened # Weights should have changed — compare flattened
let n = ac.actor.w1.shape[0] * ac.actor.w1.shape[1] let n = ac.actor.w1.shape[0] * ac.actor.w1.shape[1]
+30 -8
View File
@@ -10,8 +10,8 @@ import ./network
type type
Transition* = object Transition* = object
state*: Tensor[float32] # [42] state*: Tensor[float32] # [STATE_DIM]
action*: Tensor[float32] # [5] action*: Tensor[float32] # [ACTION_DIM]
logProb*: float32 logProb*: float32
reward*: float32 reward*: float32
value*: float32 # critic estimate at collection time value*: float32 # critic estimate at collection time
@@ -151,6 +151,13 @@ proc initACAdamStates*(ac: ActorCritic): ACAdamStates =
result.logStd = initAdamState(ac.logStd) result.logStd = initAdamState(ac.logStd)
result.initialized = true result.initialized = true
# ── Training metrics ──────────────────────────────────────────────────────────
type PPOMetrics* = object
actorLoss*: float32
valueLoss*: float32
gradNorm*: float32
# ── Gradient clipping ───────────────────────────────────────────────────────── # ── Gradient clipping ─────────────────────────────────────────────────────────
proc globalNorm(grads: varargs[Tensor[float32]]): float32 = proc globalNorm(grads: varargs[Tensor[float32]]): float32 =
@@ -171,9 +178,14 @@ proc ppoUpdate*(ac: var ActorCritic;
entropyCoeff: float32 = 0.01'f32; entropyCoeff: float32 = 0.01'f32;
valueLossCoeff: float32 = 0.5'f32; valueLossCoeff: float32 = 0.5'f32;
lr: float32 = 3e-4'f32; lr: float32 = 3e-4'f32;
maxGradNorm: float32 = 0.5'f32) {.gcsafe.} = maxGradNorm: float32 = 0.5'f32): PPOMetrics {.gcsafe.} =
if buffer.len == 0: return 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 # Initialise Adam states once; caller persists them across rounds
if not adamStates.initialized: if not adamStates.initialized:
adamStates = initACAdamStates(ac) adamStates = initACAdamStates(ac)
@@ -230,14 +242,14 @@ proc ppoUpdate*(ac: var ActorCritic;
# ── Actor forward ── # ── Actor forward ──
let actorFwd = mlpForwardCached(ac.actor, tr.state) let actorFwd = mlpForwardCached(ac.actor, tr.state)
let newMean = actorFwd.y # [5] let newMean = actorFwd.y # [ACTION_DIM]
let logStdClamped = ac.logStd.map(proc(v: float32): float32 = max(v, -3.0'f32)) let logStdClamped = ac.logStd.map(proc(v: float32): float32 = max(v, -3.0'f32))
let std = logStdClamped.map(proc(v: float32): float32 = exp(v)) let std = logStdClamped.map(proc(v: float32): float32 = exp(v))
# New log prob # New log prob
var newLogP = 0.0'f32 var newLogP = 0.0'f32
for i in 0..<5: for i in 0..<ACTION_DIM:
let mu = newMean[i] let mu = newMean[i]
let s = std[i] let s = std[i]
let diff = (tr.action[i] - mu) / s let diff = (tr.action[i] - mu) / s
@@ -250,6 +262,7 @@ proc ppoUpdate*(ac: var ActorCritic;
let surr1 = ratio * adv let surr1 = ratio * adv
let surr2 = ratioClipped * adv let surr2 = ratioClipped * adv
# Actor loss per sample = -min(surr1, surr2) # Actor loss per sample = -min(surr1, surr2)
totalActorLoss += -min(surr1, surr2)
# Which branch is active? # Which branch is active?
let useClipped = (surr2 < surr1) let useClipped = (surr2 < surr1)
let dLoss_dSurr = -1.0'f32 / mbSize.float32 # d(-mean(min))/d(min) = -1/N let dLoss_dSurr = -1.0'f32 / mbSize.float32 # d(-mean(min))/d(min) = -1/N
@@ -260,8 +273,8 @@ proc ppoUpdate*(ac: var ActorCritic;
let dLoss_dNewLogP = dLoss_dRatio * ratio let dLoss_dNewLogP = dLoss_dRatio * ratio
# d(newLogP)/d(mean[i]) = (action[i] - mean[i]) / std[i]^2 # d(newLogP)/d(mean[i]) = (action[i] - mean[i]) / std[i]^2
var dLogP_dMean = newTensor[float32](5) var dLogP_dMean = newTensor[float32](ACTION_DIM)
for i in 0..<5: for i in 0..<ACTION_DIM:
let s = std[i] let s = std[i]
dLogP_dMean[i] = (tr.action[i] - newMean[i]) / (s * s) dLogP_dMean[i] = (tr.action[i] - newMean[i]) / (s * s)
@@ -271,7 +284,7 @@ proc ppoUpdate*(ac: var ActorCritic;
# total loss gradient w.r.t. logStd: -entropyCoeff * d(entropy)/d(logStd) # total loss gradient w.r.t. logStd: -entropyCoeff * d(entropy)/d(logStd)
# Also: d(newLogP)/d(logStd_i) when logStd_i > -3: # Also: d(newLogP)/d(logStd_i) when logStd_i > -3:
# = (action_i - mean_i)^2/std_i^2 - 1 # = (action_i - mean_i)^2/std_i^2 - 1
for i in 0..<5: for i in 0..<ACTION_DIM:
let isClamped = (ac.logStd[i] <= -3.0'f32) let isClamped = (ac.logStd[i] <= -3.0'f32)
if not isClamped: if not isClamped:
let s = std[i] let s = std[i]
@@ -295,6 +308,7 @@ proc ppoUpdate*(ac: var ActorCritic;
let criticFwd = mlpForwardCached(ac.critic, tr.state) let criticFwd = mlpForwardCached(ac.critic, tr.state)
let newVal = criticFwd.y[0] let newVal = criticFwd.y[0]
# Value loss = (newVal - ret)^2; d/d(newVal) = 2*(newVal-ret) # 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 dVLoss_dVal = valueLossCoeff * 2.0'f32 * (newVal - ret) / mbSize.float32
let gradCriticOut = [dVLoss_dVal].toTensor() # [1] let gradCriticOut = [dVLoss_dVal].toTensor() # [1]
let criticGrads = mlpBackward(ac.critic, criticFwd, tr.state, gradCriticOut) let criticGrads = mlpBackward(ac.critic, criticFwd, tr.state, gradCriticOut)
@@ -314,6 +328,8 @@ proc ppoUpdate*(ac: var ActorCritic;
dCriticW1, dCriticB1, dCriticW2, dCriticB2, dCriticW3, dCriticB3 dCriticW1, dCriticB1, dCriticW2, dCriticB2, dCriticW3, dCriticB3
] ]
let norm = globalNorm(allGrads) let norm = globalNorm(allGrads)
totalGradNorm += norm
inc totalMiniBatches
if norm > maxGradNorm: if norm > maxGradNorm:
let scale = maxGradNorm / norm let scale = maxGradNorm / norm
for g in allGrads.mitems: g = g *. scale for g in allGrads.mitems: g = g *. scale
@@ -343,3 +359,9 @@ proc ppoUpdate*(ac: var ActorCritic;
adamStep(ac.critic.b3, dCriticB3, adamStates.cb3, lr) adamStep(ac.critic.b3, dCriticB3, adamStates.cb3, lr)
mbStart = mbEnd 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
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.
+10 -6
View File
@@ -40,9 +40,9 @@ public class RunBattle {
idToName.put(p.getId(), p.getName()); idToName.put(p.getId(), p.getName());
System.out.printf(" #%d %s%n", p.getId(), p.getName()); System.out.printf(" #%d %s%n", p.getId(), p.getName());
} }
System.out.printf("%-6s %-4s %-30s %-30s %-20s%n", System.out.printf("%-6s %-12s %-21s %-12s %-10s %-10s%n",
"Turn", "Id", "Name", "RadarDir", "Event"); "Turn", "Name", "Position", "Direction", "Speed", "RadarDir");
System.out.println("-".repeat(95)); System.out.println("-".repeat(80));
}); });
handle.getOnTickEvent().on(owner, event -> { handle.getOnTickEvent().on(owner, event -> {
@@ -52,12 +52,16 @@ public class RunBattle {
for (var state : event.getBotStates()) for (var state : event.getBotStates())
if (state.getName() != null) idToName.putIfAbsent(state.getId(), state.getName()); if (state.getName() != null) idToName.putIfAbsent(state.getId(), state.getName());
// Print radar direction for every bot // Print position/speed/direction for every bot
for (var state : event.getBotStates()) { for (var state : event.getBotStates()) {
String name = idToName.getOrDefault(state.getId(), String name = idToName.getOrDefault(state.getId(),
state.getName() != null ? state.getName() : "#" + state.getId()); state.getName() != null ? state.getName() : "#" + state.getId());
System.out.printf("%-6d %-4d %-30s %-30.2f%n", System.out.printf("%-6d %-12s pos=(%-6.1f,%-6.1f) dir=%-7.2f spd=%-6.2f radar=%-7.2f%n",
turn, state.getId(), name, state.getRadarDirection()); turn, name,
state.getX(), state.getY(),
state.getDirection(),
state.getSpeed(),
state.getRadarDirection());
} }
// Print any scan events on the same tick // Print any scan events on the same tick
Binary file not shown.
+76
View File
@@ -0,0 +1,76 @@
import dev.robocode.tankroyale.runner.*;
import dev.robocode.tankroyale.client.model.*;
import java.util.List;
import java.util.Map;
import java.util.concurrent.ConcurrentHashMap;
import java.util.logging.Level;
import java.util.logging.Logger;
/**
* GotoTest battle: GotoTest vs Target (1 round, 500 turns max via observation).
* Watches bot positions to verify diamond completion including wall-smash north waypoint.
*/
public class RunGotoTest {
static final Map<Integer, String> idToName = new ConcurrentHashMap<>();
public static void main(String[] args) {
Logger.getLogger("dev.robocode.tankroyale").setLevel(Level.WARNING);
String gotoDir = requireEnv("GOTO_BOT_DIR");
String sampleBots = requireEnv("SAMPLE_BOTS_DIR");
try (var runner = BattleRunner.create(b -> b.embeddedServer().suppressServerOutput())) {
var setup = BattleSetup.classic(s -> s.setNumberOfRounds(1));
var bots = List.of(
BotEntry.of(gotoDir),
BotEntry.of(sampleBots + "/Target")
);
var owner = new Object();
try (var handle = runner.startBattleAsync(setup, bots)) {
handle.getOnGameStarted().on(owner, event -> {
System.out.println("=== GAME STARTED ===");
for (var p : event.getParticipants()) {
idToName.put(p.getId(), p.getName());
System.out.printf(" #%d %s%n", p.getId(), p.getName());
}
});
handle.getOnTickEvent().on(owner, event -> {
int turn = event.getTurnNumber();
for (var state : event.getBotStates()) {
if (state.getName() != null) idToName.putIfAbsent(state.getId(), state.getName());
String name = idToName.getOrDefault(state.getId(), "#" + state.getId());
if (name.equals("GotoTest")) {
System.out.printf("T%-5d pos=(%-6.1f,%-6.1f) dir=%-7.2f spd=%-5.2f%n",
turn, state.getX(), state.getY(),
state.getDirection(), state.getSpeed());
}
}
});
handle.getOnRoundEnded().on(owner, event ->
System.out.printf("%n=== ROUND %d ENDED (turn %d) ===%n",
event.getRoundNumber(), event.getTurnNumber()));
var results = handle.awaitResults();
System.out.printf("%n=== RESULTS ===%n");
for (var r : results.getResults()) {
System.out.printf(" #%d %-25s %d pts%n",
r.getRank(), r.getName(), r.getTotalScore());
}
}
}
}
static String requireEnv(String name) {
var v = System.getenv(name);
if (v == null || v.isBlank()) {
System.err.println("Error: " + name + " env var not set");
System.exit(1);
}
return v;
}
}