## 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* = 4096 # server rounds are 2000 ticks; headroom for config drift 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 TrajectoryBuffer* = object transitions*: array[MAX_TRANSITIONS, Transition] len*: int # ── Buffer ───────────────────────────────────────────────────────────────────── proc initTrajectoryBuffer*(): TrajectoryBuffer = result = TrajectoryBuffer() proc add*(buf: var TrajectoryBuffer, t: Transition) = ## ponytail: fixed 4096 cap — server rounds run 2000 ticks; if a round ever ## exceeds the cap new transitions are dropped (oldest kept). Raise the cap ## if arena rounds get longer. 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.. 0: miniBatchSize else: bufLen for epochNum in 1..epochs: # Shuffle indices var indices = toSeq(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..= 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