## 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.. -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, 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) 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