Files
SirRoboGarage/PPO_Bot/training.nim
T
SirStone fedab54bc0 feat(PPO_Bot): multi-round transition accumulation (UPDATE_INTERVAL=10)
- Accumulate transitions across 10 rounds (~3000) before PPO update
  (was per-round ~300 — gradient estimates were far too noisy)
- training.nim: MAX_TRANSITIONS 4096→8192, done flag on transitions,
  GAE handles episode boundaries correctly
- PPO_Bot.nim: buffer persists across rounds, update every N rounds
- training.env: lr 5e-5→1e-4, entropy 0.001, UPDATE_INTERVAL=10
2026-08-20 15:06:00 +02:00

450 lines
20 KiB
Nim
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
## 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 ─────────────────────────────────────────────────────────────────────
# ponytail: transitions hold PLAIN fixed-size arrays, never Arraymancer tensors.
# Tensors crossing the bot-thread → main-thread boundary get freed on the wrong
# thread's heap under ORC (bot thread SIGSEGVs mid-round in addEvent — 5 matching
# coredumps). Plain arrays are value types: no heap, no GC, safe to move across
# threads. Tensors are rebuilt from the arrays on the consuming (training) thread.
const
MAX_TRANSITIONS* = 8192 # 10 rounds × ~300 ticks + headroom
type
Transition* = object
state*: array[STATE_DIM, float32] # plain copy, rebuilt as tensor in ppoUpdate
action*: array[ACTION_DIM, float32]
logProb*: float32
reward*: float32
value*: float32 # critic estimate at collection time
done*: bool # true at episode (round) boundary
TrajectoryBuffer* = object
transitions*: array[MAX_TRANSITIONS, Transition]
len*: int
# ── Buffer ─────────────────────────────────────────────────────────────────────
proc initTrajectoryBuffer*(): TrajectoryBuffer =
result = TrajectoryBuffer()
proc add*(buf: var TrajectoryBuffer, t: Transition) =
## ponytail: fixed 8192 cap — 10 rounds × ~300 ticks with headroom. Drops
## new transitions when full. Raise cap if accumulation window grows.
if buf.len < MAX_TRANSITIONS:
buf.transitions[buf.len] = t
inc buf.len
proc clear*(buf: var TrajectoryBuffer) =
buf.len = 0
# ── Tensor → plain array (same-thread use; tensors never cross threads) ─────
proc stateToArr*(t: Tensor[float32]): array[STATE_DIM, float32] =
for i in 0..<STATE_DIM: result[i] = t[i]
proc actionToArr*(t: Tensor[float32]): array[ACTION_DIM, float32] =
for i in 0..<ACTION_DIM: result[i] = t[i]
# ── Reward helpers ─────────────────────────────────────────────────────────────
proc computeTickReward*(myEnergyDelta, enemyEnergyDelta: float32;
distToEnemy: float32 = 0.0'f32;
maxDist: float32 = 1.0'f32;
gunBearingAbs: float32 = 180.0'f32): float32 =
## Positive when we deal more damage than we receive.
## Dense shaping: closeness (0-0.01/tick) + aim quality (0-0.02/tick).
## ponytail: magnitudes 10x smaller than original to keep shaping as a nudge,
## not the dominant signal. Increase if bot ignores positioning entirely.
let sparseReward = myEnergyDelta - enemyEnergyDelta
let distReward = 0.01'f32 * (1.0'f32 - distToEnemy / maxDist)
let aimReward = 0.02'f32 * (1.0'f32 - gunBearingAbs / 180.0'f32)
result = sparseReward + distReward + aimReward
proc computeRoundReward*(roundScore: float32): float32 =
## Normalise round-end score to a rough ±6 range (doubled win signal).
## ponytail: cap cumulative score at 400 before /50 — bounded terminal bonus
## keeps critic value scale stable across battle boundaries and long battles.
result = min(roundScore, 400.0'f32) / 50.0'f32
# ── GAE ───────────────────────────────────────────────────────────────────────
proc computeGAE*(rewards, values: seq[float32];
dones: seq[bool];
lastValue: float32;
gamma: float32 = 0.99'f32;
lam: float32 = 0.95'f32):
tuple[advantages: seq[float32], returns: seq[float32]] =
## Generalised Advantage Estimation — reverse sweep with episode boundaries.
## When done=true on transition t, bootstrap value and accumulated GAE are
## reset to 0 at that boundary (terminal state has no future value).
let n = rewards.len
var advantages = newSeq[float32](n)
var lastGae = 0.0'f32
for t in countdown(n - 1, 0):
let nextVal: float32 =
if t == n - 1 or dones[t]: 0.0'f32
else: values[t + 1]
if t == n - 1 or dones[t]:
lastGae = 0.0'f32
let delta = rewards[t] + gamma * nextVal - values[t]
lastGae = delta + gamma * lam * lastGae
advantages[t] = lastGae
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
initialized*: bool
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)
result.initialized = true
# ── Training metrics ──────────────────────────────────────────────────────────
type PPOMetrics* = object
actorLoss*: float32
valueLoss*: float32
gradNorm*: float32
# ── 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;
adamStates: var ACAdamStates;
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;
gamma: float32 = 0.99'f32;
lam: float32 = 0.95'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.
# Also reinit if aw1.m has wrong shape (e.g. loaded from old checkpoint with
# different STATE_DIM, leaving a (0,) placeholder after shape-mismatch skip).
if not adamStates.initialized or
adamStates.aw1.m.shape.len == 0 or
adamStates.aw1.m.shape != ac.actor.w1.shape:
adamStates = initACAdamStates(ac)
# 1. GAE
let rewards = buffer.transitions[0 ..< buffer.len].mapIt(it.reward)
let values = buffer.transitions[0 ..< buffer.len].mapIt(it.value)
let dones = buffer.transitions[0 ..< buffer.len].mapIt(it.done)
let (advantages, returns) = computeGAE(rewards, values, dones, lastValue, gamma = gamma, lam = lam)
# 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
# ponytail: float32 adv noise ~1e-12; advVar < 1e-8 = constant-reward
# (passive) round — dividing by that amplifies noise ~1e4+ and drifts the
# policy into exp() overflow. Center-only, skip the divide.
var normAdv: seq[float32]
if advantages.allIt(it == it and abs(it) < 1e30'f32):
if advVar < 1e-8'f32:
normAdv = advantages.mapIt(it - advMean)
else:
let advStd = sqrt(advVar + 1e-8'f32)
normAdv = advantages.mapIt((it - advMean) / advStd)
else:
normAdv = newSeq[float32](advantages.len) # poisoned input → zero advantages, no-op update
let bufLen = buffer.len
# ponytail: minibatch size <= 0 would make mbEnd == mbStart forever and spin.
# Treat as full-batch; breaks the loop unconditionally.
let mbSizeCap = if miniBatchSize > 0: miniBatchSize else: bufLen
for epochNum in 1..epochs:
# Shuffle indices
var indices = toSeq(0..<bufLen)
shuffle(indices)
var mbStart = 0
while mbStart < bufLen:
let mbEnd = min(mbStart + mbSizeCap, bufLen)
if mbEnd <= mbStart: break
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
# Rebuild the state tensor on this (training) thread — transitions hold
# plain arrays so no tensor ever crosses a thread boundary.
let x = tr.state.toTensor()
# ── Actor forward ──
let actorFwd = mlpForwardCached(ac.actor, x)
# Critic forward
let newMean = actorFwd.y # [ACTION_DIM]
# Same clamp as collection (network.nim actorForward): train-time std must
# exactly match the std the acting policy used, or ratios are distorted.
let logStdClamped = ac.logStd.map(proc(v: float32): float32 = clamp(v, logStdFloor, logStdCeiling))
let std = logStdClamped.map(proc(v: float32): float32 = exp(v))
# New log prob
var newLogP = 0.0'f32
for i in 0..<ACTION_DIM:
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)
# ponytail: float32 exp overflows at ±88; ±20 is deep in clipped-ratio
# territory, so loss/grad are identical to the true ratio
let ratio = exp(clamp(newLogP - tr.logProb, -20.0'f32, 20.0'f32))
# 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)
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
# 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](ACTION_DIM)
for i in 0..<ACTION_DIM:
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..<ACTION_DIM:
let isFloorClamped = (ac.logStd[i] <= logStdFloor)
let isCeilingClamped = (ac.logStd[i] >= logStdCeiling)
if not isFloorClamped:
let s = std[i]
let diff = (tr.action[i] - newMean[i]) / s
let dLogP_dLogStdI = diff * diff - 1.0'f32
let ppoGrad = dLoss_dNewLogP * dLogP_dLogStdI
# Entropy term pushes logStd up (update = param - lr*grad, grad is -entropyCoeff < 0).
# Gate it off at the ceiling to prevent runaway logStd.
let entropyGrad = if isCeilingClamped: 0.0'f32
else: -entropyCoeff / mbSize.float32
dLogStd[i] += ppoGrad + entropyGrad
# Backprop actor gradients
let gradActorOut = dLoss_dNewLogP *. dLogP_dMean # [5]
let actorGrads = mlpBackward(ac.actor, actorFwd, x, 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, x)
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, x, 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)
# ponytail: NaN/Inf grad norm means something exploded this minibatch
# (extreme logprob ratios, poisoned initial weights, etc.). Skip the
# Adam update entirely — no-op is safer than writing NaN into weights,
# which corrupts all future inference and hangs the bot.
if norm != norm or norm > 1e15'f32:
mbStart = mbEnd
continue
totalGradNorm += norm
inc totalMiniBatches
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, adamStates.aw1, lr)
adamStep(ac.actor.b1, dActorB1, adamStates.ab1, lr)
adamStep(ac.actor.w2, dActorW2, adamStates.aw2, lr)
adamStep(ac.actor.b2, dActorB2, adamStates.ab2, lr)
adamStep(ac.actor.w3, dActorW3, adamStates.aw3, lr)
adamStep(ac.actor.b3, dActorB3, adamStates.ab3, lr)
adamStep(ac.logStd, dLogStd, adamStates.logStd, lr)
# Collection clamps logStd to [floor, ceiling] at inference; clamp the raw
# param after the step so it can't drift above the ceiling (the old code
# only clamped at collection → train-time recompute used a bigger std than
# the policy that actually acted → distorted importance ratios).
ac.logStd = ac.logStd.map(proc(v: float32): float32 = clamp(v, logStdFloor, logStdCeiling))
adamStep(ac.critic.w1, dCriticW1, adamStates.cw1, lr)
adamStep(ac.critic.b1, dCriticB1, adamStates.cb1, lr)
adamStep(ac.critic.w2, dCriticW2, adamStates.cw2, lr)
adamStep(ac.critic.b2, dCriticB2, adamStates.cb2, lr)
adamStep(ac.critic.w3, dCriticW3, adamStates.cw3, lr)
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