Files
SirRoboGarage/SAC_LSTM_Bot_garage/tests/test_rewards.nim

110 lines
4.7 KiB
Nim

## Assert-based tests for rewards.nim.
## Run: nim c -r tests/test_rewards.nim
import std/[math, strformat]
import SAC_LSTM_Bot/rewards
template check(cond: bool, msg: string) =
if not cond:
quit("FAIL: " & msg, 1)
# ── computeReward ─────────────────────────────────────────────────────────────
block damageInflicted:
# p=1: 1.25 * (6*1 - 2) = 5 (lever 2 aggression mult, #59)
check abs(computeReward(damageInflicted = 1.0) - 5.0) < 1e-9, "p=1 damage = +5"
# p=3: 1.25 * (6*3 - 2) = 20
check abs(computeReward(damageInflicted = 3.0) - 20.0) < 1e-9, "p=3 damage = +20"
# low-power spam stays unprofitable: 1.25*(6*0.1-2) < 0
check computeReward(damageInflicted = 0.1) < 0.0, "p=0.1 spam still negative"
block damageReceived:
# p_e=1: -(6*1 - 2) = -4 (unchanged by lever 2)
check abs(computeReward(damageReceived = 1.0) - (-4.0)) < 1e-9, "p_e=1 received = -4"
# p_e=3: -(6*3 - 2) = -16
check abs(computeReward(damageReceived = 3.0) - (-16.0)) < 1e-9, "p_e=3 received = -16"
block hitBonus:
# flat +0.5 per landed shot: p=1 hit -> 5.0 + 0.5
check abs(computeReward(damageInflicted = 1.0, hitCount = 1) - 5.5) < 1e-9,
"p=1 hit = +5.5"
# two hits in one step: 1.25*(6*2-2) + 2*0.5 = 12.5 + 1.0 = 13.5
check abs(computeReward(damageInflicted = 2.0, hitCount = 2) - 13.5) < 1e-9,
"two hits = +13.5"
block ramTaken:
# flat per victim collision (#59)
check abs(computeReward(ramTakenCount = 1) - (-3.0)) < 1e-9, "ram taken x1 = -3"
check abs(computeReward(ramTakenCount = 2) - (-6.0)) < 1e-9, "ram taken x2 = -6"
block chargeDeterrent:
# zero-damage case at half threshold depth: -2 * (1 - 0.06/0.12) = -1
let rHalf = computeReward(enemyDistFrac = 0.06)
check abs(rHalf - (-1.0)) < 1e-9, "charge at frac 0.06 = -1"
# at zero distance: full ceiling
check abs(computeReward(enemyDistFrac = 0.0) - (-2.0)) < 1e-9, "charge at frac 0 = -2"
# at/beyond threshold and no-contact sentinel: no penalty
check abs(computeReward(enemyDistFrac = 0.12)) < 1e-9, "at threshold = 0"
check abs(computeReward(enemyDistFrac = 0.5)) < 1e-9, "beyond threshold = 0"
check abs(computeReward(enemyDistFrac = 2.0)) < 1e-9, "no-contact sentinel = 0"
# suppressed while dealing damage that step (fighting back at close range is fine)
let rFight = computeReward(damageInflicted = 1.0, enemyDistFrac = 0.06)
check abs(rFight - 5.0) < 1e-9, "dealing damage cancels charge penalty"
block wallHit:
check abs(computeReward(wallHitTicks = 1) - (-5.0)) < 1e-9, "1 wall tick = -5"
block wastedShot:
# p=2: -0.1 * 2 = -0.2
check abs(computeReward(wastedShotPower = 2.0) - (-0.2)) < 1e-9, "missed p=2 = -0.2"
block winLoss:
# terminal terms stay dominant over shaping (#59 scale discipline)
check abs(computeReward(win = true) - 20.0) < 1e-9, "win = +20"
check abs(computeReward(loss = true) - (-10.0)) < 1e-9, "loss = -10"
# ── RewardNormalizer cold start ───────────────────────────────────────────────
block coldStart:
var rn: RewardNormalizer
# 0 samples
let v0 = rn.normalize(99.0)
check not isNaN(v0), "0 samples: not NaN"
check classify(v0) != fcInf and classify(v0) != fcNegInf, "0 samples: not inf"
check abs(v0) < 1e-9, "0 samples: returns 0"
# 1 sample (variance undefined)
rn.update(5.0)
let v1 = rn.normalize(5.0)
check not isNaN(v1), "1 sample: not NaN"
check classify(v1) != fcInf and classify(v1) != fcNegInf, "1 sample: not inf"
check abs(v1) < 1e-9, "1 sample: returns 0"
# ── Running normalization convergence ─────────────────────────────────────────
block convergence:
var rn: RewardNormalizer
# Feed 1000 identical samples of 5.0 — mean=5.0, std=0 → normalizer returns ~0
for _ in 0 ..< 1000:
rn.update(5.0)
let v = rn.normalize(5.0)
check not isNaN(v), "convergence: not NaN"
check classify(v) != fcInf and classify(v) != fcNegInf, "convergence: not inf"
# (5 - 5) / (0 + eps) = 0
check abs(v) < 1e-6, "convergence to mean: normalized ≈ 0"
block knownMeanStd:
# Insert samples -1 and +1 repeatedly → mean=0, std=1
var rn: RewardNormalizer
for _ in 0 ..< 500:
rn.update(-1.0)
rn.update( 1.0)
# normalize(1.0) ≈ (1 - 0) / (1 + eps) ≈ 1
let vPos = rn.normalize(1.0)
check abs(vPos - 1.0) < 1e-4, &"normalize(+1) ≈ +1, got {vPos}"
let vNeg = rn.normalize(-1.0)
check abs(vNeg - (-1.0)) < 1e-4, &"normalize(-1) ≈ -1, got {vNeg}"
let vMid = rn.normalize(0.0)
check abs(vMid) < 1e-4, &"normalize(0) ≈ 0, got {vMid}"
echo "test_rewards: all passed"