## test_training.nim — assert-based tests for training.nim ## Run: nim c tests/test_training.nim && ./tests/test_training import std/[math, random] import arraymancer import PPO_Bot/network import PPO_Bot/training template check(cond: bool, msg: string) = if not cond: quit("FAIL: " & msg, 1) # ── computeTickReward ───────────────────────────────────────────────────────── block testTickReward: # I lost 2, enemy lost 10 → reward = -2 - (-10) = 8 # + default closeness shaping 0.01*(1-0/maxDist) = 0.01 (gunBearingAbs=180 → 0) let r = computeTickReward(-2.0'f32, -10.0'f32) check abs(r - 8.01'f32) < 1e-6'f32, "computeTickReward(-2, -10) == 8.01, got " & $r # ── computeRoundReward ──────────────────────────────────────────────────────── block testRoundReward: let r = computeRoundReward(350.0'f32) check abs(r - 7.0'f32) < 1e-6'f32, "computeRoundReward(350) == 7.0, got " & $r # bounded: long-battle cumulative scores must saturate, not blow the value scale check abs(computeRoundReward(89299.0'f32) - 8.0'f32) < 1e-6'f32, "computeRoundReward(89299) == 8.0 (capped), got " & $computeRoundReward(89299.0'f32) # ── TrajectoryBuffer ────────────────────────────────────────────────────────── block testBuffer: var buf = initTrajectoryBuffer() check buf.len == 0, "empty buffer len == 0" let t1 = Transition(state: zeros[float32](STATE_DIM).stateToArr, action: zeros[float32](ACTION_DIM).actionToArr, logProb: -1.0'f32, reward: 0.5'f32, value: 0.3'f32) buf.add(t1) buf.add(t1) buf.add(t1) check buf.len == 3, "buffer len == 3 after 3 adds" buf.clear() check buf.len == 0, "buffer len == 0 after clear" # ── computeGAE — hand-calculated 3-step ────────────────────────────────────── block testGAE: # rewards = [1.0, 0.0, 1.0], values = [0.5, 0.5, 0.5], lastValue = 0.0 # gamma = 0.99, lam = 0.95 # delta_2 = 1.0 + 0.99*0.0 - 0.5 = 0.5 # adv_2 = 0.5 # delta_1 = 0.0 + 0.99*0.5 - 0.5 = -0.005 # adv_1 = -0.005 + 0.99*0.95*0.5 ≈ -0.005 + 0.47025 = 0.46525 # delta_0 = 1.0 + 0.99*0.5 - 0.5 = 0.995 # adv_0 = 0.995 + 0.99*0.95*0.46525 ≈ 0.995 + 0.43744 = 1.43244 let (adv, ret) = computeGAE( rewards = @[1.0'f32, 0.0'f32, 1.0'f32], values = @[0.5'f32, 0.5'f32, 0.5'f32], lastValue = 0.0'f32, gamma = 0.99'f32, lam = 0.95'f32 ) check abs(adv[2] - 0.5'f32) < 1e-4'f32, "adv[2] should be ~0.5, got " & $adv[2] check abs(adv[1] - 0.46525'f32) < 1e-3'f32, "adv[1] should be ~0.46525, got " & $adv[1] check abs(adv[0] - 1.43244'f32) < 1e-2'f32, "adv[0] should be ~1.43244, got " & $adv[0] # returns = adv + values check abs(ret[2] - (0.5'f32 + 0.5'f32)) < 1e-4'f32, "ret[2] = adv[2] + 0.5" check abs(ret[0] - (adv[0] + 0.5'f32)) < 1e-4'f32, "ret[0] = adv[0] + 0.5" # ── ppoUpdate runs without crash; weights change ────────────────────────────── block testPpoUpdate: randomize(42) var ac = initActorCritic() # Save a copy of w1 before update let w1Before = ac.actor.w1.clone() var buf = initTrajectoryBuffer() for _ in 0..<10: let s = randomNormalTensor[float32](STATE_DIM) let a = randomNormalTensor[float32](ACTION_DIM) let lp = ac.computeLogProb(s, a) let v = ac.criticForward(s) buf.add(Transition(state: s.stateToArr, action: a.actionToArr, logProb: lp, reward: 0.1'f32, value: v)) var adam: ACAdamStates discard ppoUpdate(ac, buf, lastValue = 0.0'f32, adamStates = adam, epochs = 2, miniBatchSize = 5) # Weights should have changed — compare flattened let n = ac.actor.w1.shape[0] * ac.actor.w1.shape[1] let w1After = ac.actor.w1.reshape(n) let w1Flat = w1Before.reshape(n) var changed = false for i in 0.. 1e-9'f32: changed = true break check changed, "actor w1 should change after ppoUpdate" # ── ppoUpdate on constant-reward trajectory: zero-variance guard ───────────── # A passive round has near-constant per-tick rewards; with constant values the # GAE advantages are identical → zero variance. The normalization must not # amplify/NaN on this — update must complete with finite losses. block testPpoUpdateConstantReward: randomize(43) var ac = initActorCritic() var buf = initTrajectoryBuffer() for _ in 0..<64: let s = randomNormalTensor[float32](STATE_DIM) let a = randomNormalTensor[float32](ACTION_DIM) let lp = ac.computeLogProb(s, a) buf.add(Transition(state: s.stateToArr, action: a.actionToArr, logProb: lp, reward: 0.05'f32, value: 0.5'f32)) # constant reward+value var adam: ACAdamStates let m = ppoUpdate(ac, buf, lastValue = 0.5'f32, adamStates = adam, epochs = 2, miniBatchSize = 16) check m.actorLoss == m.actorLoss, "actorLoss NaN on constant-reward round" check m.valueLoss == m.valueLoss, "valueLoss NaN on constant-reward round" check m.gradNorm == m.gradNorm, "gradNorm NaN on constant-reward round" # ── ppoUpdate on normal-reward trajectory: finite losses ───────────────────── block testPpoUpdateNormalReward: randomize(44) var ac = initActorCritic() var buf = initTrajectoryBuffer() for i in 0..<64: let s = randomNormalTensor[float32](STATE_DIM) let a = randomNormalTensor[float32](ACTION_DIM) let lp = ac.computeLogProb(s, a) let v = ac.criticForward(s) let r = 0.05'f32 + 0.5'f32 * sin(float32(i)) buf.add(Transition(state: s.stateToArr, action: a.actionToArr, logProb: lp, reward: r, value: v)) var adam: ACAdamStates let m = ppoUpdate(ac, buf, lastValue = 0.0'f32, adamStates = adam, epochs = 2, miniBatchSize = 16) check m.actorLoss == m.actorLoss, "actorLoss NaN on normal-reward round" check m.valueLoss == m.valueLoss, "valueLoss NaN on normal-reward round" check m.gradNorm == m.gradNorm, "gradNorm NaN on normal-reward round" # ── logStd ceiling: raw param must never drift above the collection clamp ───── # Regression for the train/collection std mismatch: logStd starting above the # ceiling (as in the warm-start snapshot) must be clamped back to [floor, ceiling] # by the first Adam step, so recomputed logP matches the acting policy's std. block testLogStdCeilingClamp: randomize(45) var ac = initActorCritic() ac.logStd = newTensor[float32](ACTION_DIM).map(proc(v: float32): float32 = 3.0'f32) var buf = initTrajectoryBuffer() for _ in 0..<16: let s = randomNormalTensor[float32](STATE_DIM) let a = randomNormalTensor[float32](ACTION_DIM) let lp = ac.computeLogProb(s, a) buf.add(Transition(state: s.stateToArr, action: a.actionToArr, logProb: lp, reward: 0.1'f32, value: 0.5'f32)) var adam: ACAdamStates discard ppoUpdate(ac, buf, lastValue = 0.0'f32, adamStates = adam, epochs = 1, miniBatchSize = 16) for v in ac.logStd: check v <= logStdCeiling + 1e-6'f32, "logStd above ceiling after ppoUpdate" check v >= logStdFloor - 1e-6'f32, "logStd below floor after ppoUpdate" echo "All tests passed"