Files
SirRoboGarage/PPO_Bot/training.nim
T
SirStone eadd177d3b feat(PPO_Bot): reward + trajectory + GAE + PPO training (#16)
Manual-backprop PPO with Adam: TrajectoryBuffer, computeGAE, ppoUpdate
(4 epochs, minibatch 64, clip 0.2, grad norm 0.5). Reward helpers
computeTickReward/computeRoundReward. Bot wired: tick transitions
collected in run loop, ppoUpdate called on onRoundEnded. Fix: add
arraymancer import to PPO_Bot.nim so Tensor resolves at top level.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-08-16 15:27:15 +02:00

348 lines
14 KiB
Nim

## 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 ─────────────────────────────────────────────────────────────────────
type
Transition* = object
state*: Tensor[float32] # [42]
action*: Tensor[float32] # [5]
logProb*: float32
reward*: float32
value*: float32 # critic estimate at collection time
TrajectoryBuffer* = object
transitions*: seq[Transition]
# ── Buffer ─────────────────────────────────────────────────────────────────────
proc initTrajectoryBuffer*(): TrajectoryBuffer =
result.transitions = @[]
proc add*(buf: var TrajectoryBuffer, t: Transition) =
buf.transitions.add(t)
proc clear*(buf: var TrajectoryBuffer) =
buf.transitions.setLen(0)
proc len*(buf: TrajectoryBuffer): int =
buf.transitions.len
# ── Reward helpers ─────────────────────────────────────────────────────────────
proc computeTickReward*(myEnergyDelta, enemyEnergyDelta: float32): float32 =
## Positive when we deal more damage than we receive.
result = myEnergyDelta - enemyEnergyDelta
proc computeRoundReward*(roundScore: float32): float32 =
## Normalise round-end score to a rough ±3 range.
result = roundScore / 100.0'f32
# ── GAE ───────────────────────────────────────────────────────────────────────
proc computeGAE*(rewards, values: seq[float32];
lastValue: float32;
gamma: float32 = 0.99'f32;
lam: float32 = 0.95'f32):
tuple[advantages: seq[float32], returns: seq[float32]] =
## Generalised Advantage Estimation — reverse sweep.
## lastValue = 0 for natural episode end (death/win).
let n = rewards.len
var advantages = newSeq[float32](n)
var gaeAcc = 0.0'f32
for t in countdown(n - 1, 0):
let nextVal = if t == n - 1: lastValue else: values[t + 1]
let delta = rewards[t] + gamma * nextVal - values[t]
gaeAcc = delta + gamma * lam * gaeAcc
advantages[t] = gaeAcc
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
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)
# Persistent Adam state — survives across ppoUpdate calls (lives in training module)
var gAdamStates: ACAdamStates
var gAdamInit = false
# ── 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;
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) =
if buffer.len == 0: return
# Initialise Adam states once (persists across rounds)
if not gAdamInit:
gAdamStates = initACAdamStates(ac)
gAdamInit = true
# 1. GAE
let rewards = buffer.transitions.mapIt(it.reward)
let values = buffer.transitions.mapIt(it.value)
let (advantages, returns) = computeGAE(rewards, values, lastValue)
# 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
let advStd = sqrt(advVar + 1e-8'f32)
let normAdv = advantages.mapIt((it - advMean) / advStd)
let bufLen = buffer.len
for _ in 1..epochs:
# Shuffle indices
var indices = toSeq(0..<bufLen)
shuffle(indices)
var mbStart = 0
while mbStart < bufLen:
let mbEnd = min(mbStart + miniBatchSize, bufLen)
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
# ── Actor forward ──
let actorFwd = mlpForwardCached(ac.actor, tr.state)
let newMean = actorFwd.y # [5]
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:
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)
let ratio = exp(newLogP - tr.logProb)
# 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)
# 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](5)
for i in 0..<5:
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..<5:
let isClamped = (ac.logStd[i] <= -3.0'f32)
if not isClamped:
let s = std[i]
let diff = (tr.action[i] - newMean[i]) / s
let dLogP_dLogStdI = diff * diff - 1.0'f32
dLogStd[i] += dLoss_dNewLogP * dLogP_dLogStdI -
entropyCoeff / mbSize.float32 # entropy: d(-entropyCoeff*H)/d(logStd_i) = -entropyCoeff
# Backprop actor gradients
let gradActorOut = dLoss_dNewLogP *. dLogP_dMean # [5]
let actorGrads = mlpBackward(ac.actor, actorFwd, tr.state, 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, tr.state)
let newVal = criticFwd.y[0]
# Value loss = (newVal - ret)^2; d/d(newVal) = 2*(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)
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)
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, gAdamStates.aw1, lr)
adamStep(ac.actor.b1, dActorB1, gAdamStates.ab1, lr)
adamStep(ac.actor.w2, dActorW2, gAdamStates.aw2, lr)
adamStep(ac.actor.b2, dActorB2, gAdamStates.ab2, lr)
adamStep(ac.actor.w3, dActorW3, gAdamStates.aw3, lr)
adamStep(ac.actor.b3, dActorB3, gAdamStates.ab3, lr)
adamStep(ac.logStd, dLogStd, gAdamStates.logStd, lr)
adamStep(ac.critic.w1, dCriticW1, gAdamStates.cw1, lr)
adamStep(ac.critic.b1, dCriticB1, gAdamStates.cb1, lr)
adamStep(ac.critic.w2, dCriticW2, gAdamStates.cw2, lr)
adamStep(ac.critic.b2, dCriticB2, gAdamStates.cb2, lr)
adamStep(ac.critic.w3, dCriticW3, gAdamStates.cw3, lr)
adamStep(ac.critic.b3, dCriticB3, gAdamStates.cb3, lr)
mbStart = mbEnd