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:
Binary file not shown.
+22
-4
@@ -24,6 +24,7 @@ type PPOBot = ref object of Bot
|
||||
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
|
||||
|
||||
@@ -36,6 +37,7 @@ type
|
||||
TrainingResult = object
|
||||
ac: ActorCritic
|
||||
adamStates: ACAdamStates
|
||||
metrics: PPOMetrics
|
||||
|
||||
TrainingArgs = object
|
||||
ac: ActorCritic
|
||||
@@ -54,9 +56,9 @@ var
|
||||
proc trainingThreadProc(args: TrainingArgs) {.thread.} =
|
||||
var localAc = args.ac
|
||||
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)
|
||||
resultChan.send(TrainingResult(ac: localAc, adamStates: localAdam))
|
||||
resultChan.send(TrainingResult(ac: localAc, adamStates: localAdam, metrics: m))
|
||||
|
||||
# ── Bot methods ───────────────────────────────────────────────────────────────
|
||||
|
||||
@@ -91,6 +93,7 @@ method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
|
||||
let avgR = rewardSum / ticks.float32
|
||||
let avgRStr = formatFloat(avgR.float, ffDecimal, 3)
|
||||
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
|
||||
if threadLaunched:
|
||||
@@ -99,6 +102,9 @@ method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
|
||||
ac = trained.ac
|
||||
gAdamStates = trained.adamStates
|
||||
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:
|
||||
bot.hasLastTrans = false
|
||||
@@ -122,6 +128,8 @@ method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
|
||||
bot.buffer.clear()
|
||||
bot.hasLastTrans = false
|
||||
|
||||
printToStdOut(&" train→ R:{roundCounter} ticks:{ticks}\n")
|
||||
echo &" train→ R:{roundCounter} ticks:{ticks}"
|
||||
createThread(trainingThread, trainingThreadProc, args)
|
||||
threadLaunched = true
|
||||
|
||||
@@ -147,10 +155,19 @@ method run(bot: PPOBot) =
|
||||
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 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
|
||||
let curEnergy = getEnergy().float32
|
||||
@@ -182,6 +199,7 @@ method run(bot: PPOBot) =
|
||||
bot.prevEnergy = curEnergy
|
||||
bot.prevEnemyE = curEnemyE
|
||||
bot.hasLastTrans = true
|
||||
bot.lastActions = acts
|
||||
|
||||
setTargetSpeed(acts.targetSpeed.float)
|
||||
setTurnRate(acts.turnRate.float)
|
||||
|
||||
+35
-23
@@ -2,32 +2,44 @@
|
||||
|
||||
import arraymancer
|
||||
import std/math
|
||||
import ./controllers
|
||||
|
||||
func sigmoid(x: float): float = 1.0 / (1.0 + exp(-x))
|
||||
|
||||
type
|
||||
BotActions* = object
|
||||
targetSpeed*: float32
|
||||
turnRate*: float32
|
||||
gunTurnRate*: float32
|
||||
targetSpeed*: float
|
||||
turnRate*: float
|
||||
gunTurnRate*: float
|
||||
shouldFire*: bool
|
||||
firePower*: float32
|
||||
firePower*: float
|
||||
gotoX*: float
|
||||
gotoY*: float
|
||||
aimToX*: float
|
||||
aimToY*: float
|
||||
|
||||
proc mapActions*(rawActions: Tensor[float32], currentSpeed: float32, gunHeat: float32): BotActions =
|
||||
## rawActions: [5] tensor from actorForward.
|
||||
## Ranges:
|
||||
## targetSpeed ∈ [-8, 8]
|
||||
## turnRate ∈ [-(10 - 0.75|v|), (10 - 0.75|v|)]
|
||||
## gunTurnRate ∈ [-20, 20]
|
||||
## firePower ∈ [0.1, 3.0]
|
||||
let r0 = tanh(rawActions[0].float64).float32
|
||||
let r1 = tanh(rawActions[1].float64).float32
|
||||
let r2 = tanh(rawActions[2].float64).float32
|
||||
let r3 = tanh(rawActions[3].float64).float32
|
||||
let r4 = rawActions[4]
|
||||
proc mapActions*(rawActions: Tensor[float32],
|
||||
gunHeat: float,
|
||||
arenaWidth, arenaHeight: float,
|
||||
botX, botY, direction, speed, gunDirection: float): BotActions =
|
||||
## rawActions: [6] tensor from actorForward.
|
||||
## Dims 0–1: goto x/y, 2–3: aimTo x/y, 4: fire decision, 5: fire power.
|
||||
let gotoX = sigmoid(rawActions[0].float) * arenaWidth
|
||||
let gotoY = sigmoid(rawActions[1].float) * arenaHeight
|
||||
let aimToX = sigmoid(rawActions[2].float) * arenaWidth
|
||||
let aimToY = sigmoid(rawActions[3].float) * arenaHeight
|
||||
let fireDec = tanh(rawActions[4].float)
|
||||
let fp = sigmoid(rawActions[5].float) * 2.9 + 0.1
|
||||
|
||||
result.targetSpeed = r0 * 8.0'f32
|
||||
result.turnRate = r1 * (10.0'f32 - 0.75'f32 * abs(currentSpeed))
|
||||
result.gunTurnRate = r2 * 20.0'f32
|
||||
let fireDecision = r3
|
||||
result.shouldFire = fireDecision >= 0.0'f32 and gunHeat <= 0.0'f32
|
||||
# sigmoid(r4) * 2.9 + 0.1 → [0.1, 3.0]
|
||||
result.firePower = (1.0'f32 / (1.0'f32 + exp(-r4.float64).float32)) * 2.9'f32 + 0.1'f32
|
||||
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,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
@@ -3,6 +3,10 @@
|
||||
import arraymancer
|
||||
import std/[math, random]
|
||||
|
||||
const
|
||||
STATE_DIM* = 44
|
||||
ACTION_DIM* = 6
|
||||
|
||||
type
|
||||
MLP* = object
|
||||
w1*, b1*: Tensor[float32] # [hidden, input], [hidden]
|
||||
@@ -12,7 +16,7 @@ type
|
||||
ActorCritic* = object
|
||||
actor*: 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 =
|
||||
# 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)
|
||||
|
||||
proc initActorCritic*(): ActorCritic =
|
||||
result.actor = initMLP(42, 64, 5)
|
||||
result.critic = initMLP(42, 64, 1)
|
||||
result.logStd = zeros[float32](5) # init to 0 → std=1
|
||||
result.actor = initMLP(STATE_DIM, 64, ACTION_DIM)
|
||||
result.critic = initMLP(STATE_DIM, 64, 1)
|
||||
result.logStd = zeros[float32](ACTION_DIM) # init to 0 → std=1
|
||||
|
||||
proc forward*(mlp: MLP, x: Tensor[float32]): Tensor[float32] =
|
||||
## x shape: [inputDim] (1D vector)
|
||||
@@ -35,15 +39,15 @@ proc forward*(mlp: MLP, x: Tensor[float32]): Tensor[float32] =
|
||||
result = mlp.w3 * h2 + mlp.b3
|
||||
|
||||
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)
|
||||
# Floor logStd at -3 before exp → min std ≈ 0.05
|
||||
var clampedLogStd = ac.logStd.map(proc(v: float32): float32 = max(v, -3.0'f32))
|
||||
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
|
||||
for i in 0..<5:
|
||||
for i in 0..<ACTION_DIM:
|
||||
let mu = mean[i]
|
||||
let s = std[i]
|
||||
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)
|
||||
|
||||
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)
|
||||
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 std = logStdClamped.map(proc(v: float32): float32 = exp(v))
|
||||
var logP = 0.0'f32
|
||||
for i in 0..<5:
|
||||
for i in 0..<ACTION_DIM:
|
||||
let mu = mean[i]
|
||||
let s = std[i]
|
||||
let diff = (action[i] - mu) / s
|
||||
|
||||
@@ -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.
|
||||
|
||||
import std/math
|
||||
@@ -16,10 +16,12 @@ type
|
||||
gunHeat*: float64
|
||||
arenaWidth*, arenaHeight*: float64
|
||||
|
||||
proc buildStateVector*(bot: BotStateData; enemy: EnemyTracker): Tensor[float32] =
|
||||
## Build the 42-float normalized state tensor.
|
||||
proc buildStateVector*(bot: BotStateData; enemy: EnemyTracker;
|
||||
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.
|
||||
result = zeros[float32](42)
|
||||
result = zeros[float32](44)
|
||||
|
||||
let aW = bot.arenaWidth
|
||||
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 + 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)
|
||||
|
||||
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)
|
||||
# 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"
|
||||
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"
|
||||
@@ -11,16 +11,16 @@ func isNaNF(x: float32): bool = classify(x) == fcNan
|
||||
|
||||
when isMainModule:
|
||||
# ---- MLP forward shape ----
|
||||
let mlp = initMLP(42, 64, 5)
|
||||
let inp = zeros[float32](42)
|
||||
let mlp = initMLP(STATE_DIM, 64, ACTION_DIM)
|
||||
let inp = zeros[float32](STATE_DIM)
|
||||
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 ----
|
||||
let ac = initActorCritic()
|
||||
let state = zeros[float32](42)
|
||||
let state = zeros[float32](STATE_DIM)
|
||||
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 isFiniteF(logP), &"logProb not finite: {logP}"
|
||||
|
||||
@@ -31,23 +31,20 @@ when isMainModule:
|
||||
|
||||
# ---- logStd floor: collapsing logStd should not break actorForward ----
|
||||
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)
|
||||
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}"
|
||||
|
||||
# ---- action mapping ranges ----
|
||||
let raw = randomNormalTensor[float32](5)
|
||||
let raw = randomNormalTensor[float32](ACTION_DIM)
|
||||
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,
|
||||
&"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,
|
||||
&"gunTurnRate out of range: {botActs.gunTurnRate}"
|
||||
|
||||
@@ -55,7 +52,7 @@ when isMainModule:
|
||||
&"firePower out of range: {botActs.firePower}"
|
||||
|
||||
# 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"
|
||||
|
||||
echo "All tests passed"
|
||||
|
||||
@@ -10,16 +10,16 @@ template check(cond: bool, msg: string) =
|
||||
|
||||
# Bot at arena center; stationary enemy due north (same x, higher y).
|
||||
# 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 botY = 300.0
|
||||
let enemyX = 400.0 # same x → dx = 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).
|
||||
var radarDir = 0.0
|
||||
# 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
|
||||
|
||||
Binary file not shown.
@@ -84,7 +84,7 @@ block testStateVectorLength:
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
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:
|
||||
var t = initEnemyTracker()
|
||||
@@ -95,7 +95,7 @@ block testStateVectorRange:
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
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,
|
||||
&"sv[{i}]={sv[i]} out of [-2,2] range"
|
||||
|
||||
@@ -155,5 +155,8 @@ block testHistoryPaddedWhenEmpty:
|
||||
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"
|
||||
|
||||
echo "All tests passed"
|
||||
|
||||
@@ -29,7 +29,7 @@ block testBuffer:
|
||||
var buf = initTrajectoryBuffer()
|
||||
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)
|
||||
buf.add(t1)
|
||||
buf.add(t1)
|
||||
@@ -80,14 +80,14 @@ block testPpoUpdate:
|
||||
|
||||
var buf = initTrajectoryBuffer()
|
||||
for _ in 0..<10:
|
||||
let s = randomNormalTensor[float32](42)
|
||||
let a = randomNormalTensor[float32](5)
|
||||
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, action: a, logProb: lp, reward: 0.1'f32, value: v))
|
||||
|
||||
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
|
||||
let n = ac.actor.w1.shape[0] * ac.actor.w1.shape[1]
|
||||
|
||||
+30
-8
@@ -10,8 +10,8 @@ import ./network
|
||||
|
||||
type
|
||||
Transition* = object
|
||||
state*: Tensor[float32] # [42]
|
||||
action*: Tensor[float32] # [5]
|
||||
state*: Tensor[float32] # [STATE_DIM]
|
||||
action*: Tensor[float32] # [ACTION_DIM]
|
||||
logProb*: float32
|
||||
reward*: float32
|
||||
value*: float32 # critic estimate at collection time
|
||||
@@ -151,6 +151,13 @@ proc initACAdamStates*(ac: ActorCritic): ACAdamStates =
|
||||
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 =
|
||||
@@ -171,9 +178,14 @@ proc ppoUpdate*(ac: var ActorCritic;
|
||||
entropyCoeff: float32 = 0.01'f32;
|
||||
valueLossCoeff: float32 = 0.5'f32;
|
||||
lr: float32 = 3e-4'f32;
|
||||
maxGradNorm: float32 = 0.5'f32) {.gcsafe.} =
|
||||
maxGradNorm: float32 = 0.5'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
|
||||
if not adamStates.initialized:
|
||||
adamStates = initACAdamStates(ac)
|
||||
@@ -230,14 +242,14 @@ proc ppoUpdate*(ac: var ActorCritic;
|
||||
|
||||
# ── Actor forward ──
|
||||
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 std = logStdClamped.map(proc(v: float32): float32 = exp(v))
|
||||
|
||||
# New log prob
|
||||
var newLogP = 0.0'f32
|
||||
for i in 0..<5:
|
||||
for i in 0..<ACTION_DIM:
|
||||
let mu = newMean[i]
|
||||
let s = std[i]
|
||||
let diff = (tr.action[i] - mu) / s
|
||||
@@ -250,6 +262,7 @@ proc ppoUpdate*(ac: var ActorCritic;
|
||||
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
|
||||
@@ -260,8 +273,8 @@ proc ppoUpdate*(ac: var ActorCritic;
|
||||
let dLoss_dNewLogP = dLoss_dRatio * ratio
|
||||
|
||||
# d(newLogP)/d(mean[i]) = (action[i] - mean[i]) / std[i]^2
|
||||
var dLogP_dMean = newTensor[float32](5)
|
||||
for i in 0..<5:
|
||||
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)
|
||||
|
||||
@@ -271,7 +284,7 @@ proc ppoUpdate*(ac: var ActorCritic;
|
||||
# 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..<5:
|
||||
for i in 0..<ACTION_DIM:
|
||||
let isClamped = (ac.logStd[i] <= -3.0'f32)
|
||||
if not isClamped:
|
||||
let s = std[i]
|
||||
@@ -295,6 +308,7 @@ proc ppoUpdate*(ac: var ActorCritic;
|
||||
let criticFwd = mlpForwardCached(ac.critic, tr.state)
|
||||
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, tr.state, gradCriticOut)
|
||||
@@ -314,6 +328,8 @@ proc ppoUpdate*(ac: var ActorCritic;
|
||||
dCriticW1, dCriticB1, dCriticW2, dCriticB2, dCriticW3, dCriticB3
|
||||
]
|
||||
let norm = globalNorm(allGrads)
|
||||
totalGradNorm += norm
|
||||
inc totalMiniBatches
|
||||
if norm > maxGradNorm:
|
||||
let scale = maxGradNorm / norm
|
||||
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)
|
||||
|
||||
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.
Reference in New Issue
Block a user