feat(SAC_LSTM_Bot): reward module (#44)
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,49 @@
|
|||||||
|
## rewards.nim — Raw reward computation + running mean/variance normalizer.
|
||||||
|
## Welford online algorithm; safe cold-start (0 or 1 samples).
|
||||||
|
|
||||||
|
import std/math
|
||||||
|
|
||||||
|
# ── Raw reward ────────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
proc computeReward*(
|
||||||
|
damageInflicted: float64 = 0.0, # fire power p of own shot that hit
|
||||||
|
damageReceived: float64 = 0.0, # fire power p_e of enemy shot that hit
|
||||||
|
wallHitTicks: int = 0, # ticks in wall contact this step
|
||||||
|
wastedShotPower: float64 = 0.0, # fire power of shot that missed/hit wall
|
||||||
|
win: bool = false,
|
||||||
|
loss: bool = false
|
||||||
|
): float64 =
|
||||||
|
## Returns the raw (un-normalized) reward for one decision step.
|
||||||
|
## Damage formula: 4p + 2(p-1) = 6p - 2 (matches Tank Royale bullet rules).
|
||||||
|
let p = damageInflicted
|
||||||
|
let pe = damageReceived
|
||||||
|
if p > 0.0: result += 6.0 * p - 2.0
|
||||||
|
if pe > 0.0: result -= 6.0 * pe - 2.0
|
||||||
|
result -= 5.0 * wallHitTicks.float64
|
||||||
|
if wastedShotPower > 0.0: result -= 0.1 * wastedShotPower
|
||||||
|
if win: result += 20.0
|
||||||
|
if loss: result -= 10.0
|
||||||
|
|
||||||
|
# ── Running normalizer (Welford) ──────────────────────────────────────────────
|
||||||
|
|
||||||
|
const NormEps = 1e-8
|
||||||
|
|
||||||
|
type
|
||||||
|
RewardNormalizer* = object
|
||||||
|
n*: int # samples seen
|
||||||
|
mean*: float64
|
||||||
|
m2*: float64 # sum of squared deviations (Welford M2)
|
||||||
|
|
||||||
|
proc update*(rn: var RewardNormalizer; r: float64) =
|
||||||
|
rn.n += 1
|
||||||
|
let delta = r - rn.mean
|
||||||
|
rn.mean += delta / rn.n.float64
|
||||||
|
let delta2 = r - rn.mean
|
||||||
|
rn.m2 += delta * delta2
|
||||||
|
|
||||||
|
proc normalize*(rn: RewardNormalizer; r: float64): float64 =
|
||||||
|
## Returns (r - mean) / (std + eps).
|
||||||
|
## Cold start (n < 2): returns 0.0 to avoid NaN/inf.
|
||||||
|
if rn.n < 2: return 0.0
|
||||||
|
let variance = rn.m2 / rn.n.float64 # ponytail: population var; switch to n-1 if bias matters
|
||||||
|
result = (r - rn.mean) / (sqrt(variance) + NormEps)
|
||||||
@@ -0,0 +1,79 @@
|
|||||||
|
## 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: 6*1 - 2 = 4
|
||||||
|
check abs(computeReward(damageInflicted = 1.0) - 4.0) < 1e-9, "p=1 damage = +4"
|
||||||
|
# p=3: 6*3 - 2 = 16
|
||||||
|
check abs(computeReward(damageInflicted = 3.0) - 16.0) < 1e-9, "p=3 damage = +16"
|
||||||
|
|
||||||
|
block damageReceived:
|
||||||
|
# p_e=1: -(6*1 - 2) = -4
|
||||||
|
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 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:
|
||||||
|
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"
|
||||||
Reference in New Issue
Block a user