feat(PPO_Bot): network forward pass + action mapping (#14)
Two-hidden-layer MLP actor-critic (42→64→64→5/1) with stochastic actorForward, logStd floor at -3, and BotAction mapper wired into the run() loop. Assert-based test suite covers shapes, finiteness, logStd collapse, and all action range bounds. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,57 @@
|
|||||||
|
## PPO_Bot — enemy tracker + state vector wired into the game loop.
|
||||||
|
## Forward pass uses a random ActorCritic policy (weights not yet trained).
|
||||||
|
|
||||||
|
import std/os
|
||||||
|
import arraymancer
|
||||||
|
import tankroyale_botapi
|
||||||
|
import network
|
||||||
|
import actions
|
||||||
|
import ./enemy_tracker
|
||||||
|
import ./state_vector
|
||||||
|
|
||||||
|
const botJsonPath = currentSourcePath().parentDir / "PPO_Bot.json"
|
||||||
|
|
||||||
|
type PPOBot = ref object of Bot
|
||||||
|
tracker: EnemyTracker
|
||||||
|
|
||||||
|
var ac = initActorCritic()
|
||||||
|
|
||||||
|
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) =
|
||||||
|
bot.tracker = initEnemyTracker()
|
||||||
|
|
||||||
|
method run(bot: PPOBot) =
|
||||||
|
while isRunning():
|
||||||
|
bot.tracker.deadReckon()
|
||||||
|
|
||||||
|
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 state = buildStateVector(botData, bot.tracker)
|
||||||
|
let (rawActs, _) = ac.actorForward(state)
|
||||||
|
let acts = mapActions(rawActs, getSpeed().float32, getGunHeat().float32)
|
||||||
|
|
||||||
|
setTargetSpeed(acts.targetSpeed.float)
|
||||||
|
setTurnRate(acts.turnRate.float)
|
||||||
|
setGunTurnRate(acts.gunTurnRate.float)
|
||||||
|
if acts.shouldFire:
|
||||||
|
discard setFire(acts.firePower.float)
|
||||||
|
go()
|
||||||
|
|
||||||
|
when isMainModule:
|
||||||
|
var bot = PPOBot(tracker: initEnemyTracker())
|
||||||
|
start(bot, botJsonPath)
|
||||||
@@ -0,0 +1,33 @@
|
|||||||
|
## actions.nim — map raw network output to Tank Royale bot commands.
|
||||||
|
|
||||||
|
import arraymancer
|
||||||
|
import std/math
|
||||||
|
|
||||||
|
type
|
||||||
|
BotActions* = object
|
||||||
|
targetSpeed*: float32
|
||||||
|
turnRate*: float32
|
||||||
|
gunTurnRate*: float32
|
||||||
|
shouldFire*: bool
|
||||||
|
firePower*: float32
|
||||||
|
|
||||||
|
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]
|
||||||
|
|
||||||
|
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
|
||||||
@@ -0,0 +1,60 @@
|
|||||||
|
## network.nim — MLP and ActorCritic forward pass (inference only, no autograd).
|
||||||
|
|
||||||
|
import arraymancer
|
||||||
|
import std/[math, random]
|
||||||
|
|
||||||
|
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] # [5] — 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(42, 64, 5)
|
||||||
|
result.critic = initMLP(42, 64, 1)
|
||||||
|
result.logStd = zeros[float32](5) # init to 0 → std=1
|
||||||
|
|
||||||
|
proc forward*(mlp: MLP, x: Tensor[float32]): Tensor[float32] =
|
||||||
|
## x shape: [inputDim] (1D vector)
|
||||||
|
let h1 = tanh(mlp.w1 * x + mlp.b1)
|
||||||
|
let h2 = tanh(mlp.w2 * h1 + mlp.b2)
|
||||||
|
result = mlp.w3 * h2 + mlp.b3
|
||||||
|
|
||||||
|
proc actorForward*(ac: ActorCritic, state: Tensor[float32]): tuple[actions: Tensor[float32], logProb: float32] =
|
||||||
|
## state: [42]. Returns sampled actions [5] 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 logP = 0.0'f32
|
||||||
|
for i in 0..<5:
|
||||||
|
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: [42]. Returns scalar value estimate.
|
||||||
|
let val = ac.critic.forward(state)
|
||||||
|
result = val[0]
|
||||||
@@ -0,0 +1,61 @@
|
|||||||
|
## 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(42, 64, 5)
|
||||||
|
let inp = zeros[float32](42)
|
||||||
|
let mlpOut = mlp.forward(inp)
|
||||||
|
assert mlpOut.shape[0] == 5, &"MLP output shape wrong: {mlpOut.shape}"
|
||||||
|
|
||||||
|
# ---- ActorCritic actorForward ----
|
||||||
|
let ac = initActorCritic()
|
||||||
|
let state = zeros[float32](42)
|
||||||
|
let (acts, logP) = ac.actorForward(state)
|
||||||
|
assert acts.shape[0] == 5, &"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..<5: ac2.logStd[i] = -10.0'f32
|
||||||
|
let (acts2, logP2) = ac2.actorForward(state)
|
||||||
|
assert acts2.shape[0] == 5, "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 speed = 4.0'f32
|
||||||
|
let botActs = mapActions(raw, speed, 0.0'f32) # 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}"
|
||||||
|
|
||||||
|
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, speed, 1.0'f32)
|
||||||
|
assert not noFire.shouldFire, "shouldFire should be false when gunHeat > 0"
|
||||||
|
|
||||||
|
echo "All tests passed"
|
||||||
Reference in New Issue
Block a user