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