## 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] # [STATE_DIM] action*: Tensor[float32] # [ACTION_DIM] 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; 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). result = roundScore / 50.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.. 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 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