Compare commits
10 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 23c65c9ac6 | |||
| f130bf1254 | |||
| a0a3840980 | |||
| 7d73d32c85 | |||
| df3bbbd14e | |||
| 4258d364b9 | |||
| f88580b157 | |||
| ca3e3d2272 | |||
| 0d35646dc9 | |||
| c834d2cbee |
+3
-2
@@ -139,6 +139,8 @@ method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
|
|||||||
printToStdOut(&"R:{roundCounter} ticks:{ticks} buf:{bot.buffer.len} sinceUpd:{roundsSinceUpdate}/{hpUpdateInterval} avgR:{avgRStr} score:{e.results.totalScore}\n")
|
printToStdOut(&"R:{roundCounter} ticks:{ticks} buf:{bot.buffer.len} sinceUpd:{roundsSinceUpdate}/{hpUpdateInterval} avgR:{avgRStr} score:{e.results.totalScore}\n")
|
||||||
echo &"R:{roundCounter} ticks:{ticks} buf:{bot.buffer.len} sinceUpd:{roundsSinceUpdate}/{hpUpdateInterval} avgR:{avgRStr} score:{e.results.totalScore}"
|
echo &"R:{roundCounter} ticks:{ticks} buf:{bot.buffer.len} sinceUpd:{roundsSinceUpdate}/{hpUpdateInterval} avgR:{avgRStr} score:{e.results.totalScore}"
|
||||||
|
|
||||||
|
writeFile(weightsRoot / "round_counter.txt", $roundCounter)
|
||||||
|
|
||||||
if bot.buffer.len == 0:
|
if bot.buffer.len == 0:
|
||||||
bot.hasLastTrans = false
|
bot.hasLastTrans = false
|
||||||
return
|
return
|
||||||
@@ -152,7 +154,6 @@ method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
|
|||||||
# checkpoint save, but keep advancing/writing round_counter.txt so run.sh's
|
# checkpoint save, but keep advancing/writing round_counter.txt so run.sh's
|
||||||
# remaining-rounds bookkeeping still works, and keep the game line above.
|
# remaining-rounds bookkeeping still works, and keep the game line above.
|
||||||
if evalOnly:
|
if evalOnly:
|
||||||
writeFile(weightsRoot / "round_counter.txt", $roundCounter)
|
|
||||||
bot.buffer.clear()
|
bot.buffer.clear()
|
||||||
bot.hasLastTrans = false
|
bot.hasLastTrans = false
|
||||||
roundsSinceUpdate = 0
|
roundsSinceUpdate = 0
|
||||||
@@ -280,7 +281,7 @@ method run(bot: PPOBot) =
|
|||||||
# detection-tick pulse, and the next iteration's spawn check sees false —
|
# detection-tick pulse, and the next iteration's spawn check sees false —
|
||||||
# one shot → exactly one bullet, even across deadReckon gaps.
|
# one shot → exactly one bullet, even across deadReckon gaps.
|
||||||
bot.tracker.current.hasFired = false
|
bot.tracker.current.hasFired = false
|
||||||
let (rawActs, logP) = ac.actorForward(state)
|
let (rawActs, logP) = ac.actorForward(state, deterministic = evalOnly)
|
||||||
let value = ac.criticForward(state)
|
let value = ac.criticForward(state)
|
||||||
let ex = if bot.tracker.hasContact: bot.tracker.current.x else: botData.arenaWidth / 2.0
|
let ex = if bot.tracker.hasContact: bot.tracker.current.x else: botData.arenaWidth / 2.0
|
||||||
let ey = if bot.tracker.hasContact: bot.tracker.current.y else: botData.arenaHeight / 2.0
|
let ey = if bot.tracker.hasContact: bot.tracker.current.y else: botData.arenaHeight / 2.0
|
||||||
|
|||||||
+6
-2
@@ -43,9 +43,13 @@ proc forward*(mlp: MLP, x: Tensor[float32]): Tensor[float32] =
|
|||||||
let h2 = tanh(mlp.w2 * h1 + mlp.b2)
|
let h2 = tanh(mlp.w2 * h1 + mlp.b2)
|
||||||
result = mlp.w3 * h2 + mlp.b3
|
result = mlp.w3 * h2 + mlp.b3
|
||||||
|
|
||||||
proc actorForward*(ac: ActorCritic, state: Tensor[float32]): tuple[actions: Tensor[float32], logProb: float32] =
|
proc actorForward*(ac: ActorCritic, state: Tensor[float32], deterministic = false): tuple[actions: Tensor[float32], logProb: float32] =
|
||||||
## state: [STATE_DIM]. Returns sampled actions [ACTION_DIM] and sum log-prob.
|
## state: [STATE_DIM]. Returns actions [ACTION_DIM] and sum log-prob.
|
||||||
|
## deterministic=true: return mean only (no noise), logProb=0.
|
||||||
let mean = ac.actor.forward(state)
|
let mean = ac.actor.forward(state)
|
||||||
|
if deterministic:
|
||||||
|
return (actions: mean, logProb: 0.0'f32)
|
||||||
|
|
||||||
# Floor logStd at logStdFloor before exp → min std ≈ exp(logStdFloor)
|
# Floor logStd at logStdFloor before exp → min std ≈ exp(logStdFloor)
|
||||||
var clampedLogStd = ac.logStd.map(proc(v: float32): float32 = clamp(v, logStdFloor, logStdCeiling))
|
var clampedLogStd = ac.logStd.map(proc(v: float32): float32 = clamp(v, logStdFloor, logStdCeiling))
|
||||||
let std = clampedLogStd.map(proc(v: float32): float32 = exp(v))
|
let std = clampedLogStd.map(proc(v: float32): float32 = exp(v))
|
||||||
|
|||||||
@@ -0,0 +1,13 @@
|
|||||||
|
# Package
|
||||||
|
version = "0.1.0"
|
||||||
|
author = "Davide Cappellini"
|
||||||
|
description = "SAC+LSTM-trained Tank Royale bot"
|
||||||
|
license = "MIT"
|
||||||
|
srcDir = "src"
|
||||||
|
bin = @["SAC_LSTM_Bot"]
|
||||||
|
|
||||||
|
# Dependencies
|
||||||
|
requires "nim >= 2.0.0"
|
||||||
|
# tankroyale_botapi is vendored in-tree (libs/tankroyale_botapi) and wired via
|
||||||
|
# config.nims --path; no nimble dependency so builds never touch ~/.nimble/pkgs2.
|
||||||
|
requires "arraymancer >= 0.7.0"
|
||||||
@@ -0,0 +1,12 @@
|
|||||||
|
# Static-link OpenBLAS for portable deployment
|
||||||
|
# ponytail: adjust path per machine, or use pkg-config
|
||||||
|
switch("passL", "-L/nix/store/v07svn2y92bvzjl51aj7c9ca1cwg7rw7-openblas-0.3.32/lib -lopenblas")
|
||||||
|
switch("threads", "on")
|
||||||
|
# begin Nimble config (version 2)
|
||||||
|
when withDir(thisDir(), system.fileExists("nimble.paths")):
|
||||||
|
include "nimble.paths"
|
||||||
|
# end Nimble config
|
||||||
|
# Use the repo-vendored Tank Royale bot API (libs/) instead of the nimble pkg.
|
||||||
|
# Must come AFTER the nimble.paths include: later --path wins the import search.
|
||||||
|
switch("path", thisDir() & "/../libs/tankroyale_botapi")
|
||||||
|
switch("path", thisDir() & "/../libs/radar_lock")
|
||||||
@@ -0,0 +1,11 @@
|
|||||||
|
{
|
||||||
|
"name": "Recurrent Royalty",
|
||||||
|
"version": "0.1.0",
|
||||||
|
"authors": ["Davide Cappellini"],
|
||||||
|
"description": "SAC+LSTM Tank Royale bot — skeleton with radar lock",
|
||||||
|
"homepage": "",
|
||||||
|
"countryCodes": ["IT"],
|
||||||
|
"gameTypes": ["classic", "melee", "1v1"],
|
||||||
|
"platform": "Nim",
|
||||||
|
"programmingLang": "Nim"
|
||||||
|
}
|
||||||
@@ -0,0 +1,60 @@
|
|||||||
|
## SAC_LSTM_Bot — skeleton: radar lock + "Recurrent Royalty" color scheme.
|
||||||
|
## No RL yet. Connects, sets colors, locks radar onto enemy.
|
||||||
|
|
||||||
|
import std/os
|
||||||
|
import tankroyale_botapi
|
||||||
|
import radar_lock
|
||||||
|
|
||||||
|
const botJsonPath = currentSourcePath().parentDir / "SAC_LSTM_Bot.json"
|
||||||
|
|
||||||
|
# ── Colors (Recurrent Royalty palette) ───────────────────────────────────────
|
||||||
|
const
|
||||||
|
ColBody = fromHex("#7B2FBE")
|
||||||
|
ColTurret = fromHex("#FFD700")
|
||||||
|
ColGun = fromHex("#4A0E6B")
|
||||||
|
ColRadar = fromHex("#FFD700")
|
||||||
|
ColScan = fromHex("#FFB000")
|
||||||
|
ColBullet = fromHex("#FFC125")
|
||||||
|
ColTracks = fromHex("#2C2C34")
|
||||||
|
|
||||||
|
proc applyColors() =
|
||||||
|
setBodyColor(ColBody)
|
||||||
|
setTurretColor(ColTurret)
|
||||||
|
setGunColor(ColGun)
|
||||||
|
setRadarColor(ColRadar)
|
||||||
|
setScanColor(ColScan)
|
||||||
|
setBulletColor(ColBullet)
|
||||||
|
setTracksColor(ColTracks)
|
||||||
|
|
||||||
|
# ── Bot type ──────────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
type SacBot = ref object of Bot
|
||||||
|
enemyBearing: float # last known absolute bearing to enemy
|
||||||
|
|
||||||
|
# ── Event handlers ────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
method onRoundStarted*(bot: SacBot, e: RoundStartedEvent) =
|
||||||
|
setAdjustRadarForBodyTurn(true)
|
||||||
|
setAdjustRadarForGunTurn(true)
|
||||||
|
radar_lock.init()
|
||||||
|
applyColors()
|
||||||
|
|
||||||
|
method onScannedBot*(bot: SacBot, e: ScannedBotEvent) =
|
||||||
|
bot.enemyBearing = directionTo(getX(), getY(), e.x, e.y)
|
||||||
|
# Same-tick radar lock: apply turn rate immediately so it takes effect this tick.
|
||||||
|
setRadarTurnRate(radar_lock.doRadar(getRadarDirection(), bot.enemyBearing))
|
||||||
|
|
||||||
|
# ── Run loop ──────────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
method run(bot: SacBot) =
|
||||||
|
while isRunning():
|
||||||
|
# Spin radar when no enemy is visible (full sweep).
|
||||||
|
if bot.enemyBearing == 0.0:
|
||||||
|
setRadarTurnRate(45.0)
|
||||||
|
go()
|
||||||
|
|
||||||
|
# ── Entry point ───────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
when isMainModule:
|
||||||
|
var bot = SacBot()
|
||||||
|
start(bot, botJsonPath)
|
||||||
@@ -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,117 @@
|
|||||||
|
## State vector module — produces a 35-dimensional normalized tensor for SAC+LSTM policy.
|
||||||
|
## No bot API imports; takes plain data structs populated from game events.
|
||||||
|
## The LSTM handles temporal context, so no explicit history window here.
|
||||||
|
|
||||||
|
import std/math
|
||||||
|
import arraymancer
|
||||||
|
|
||||||
|
const STATE_DIM* = 35
|
||||||
|
|
||||||
|
type
|
||||||
|
BulletData* = object
|
||||||
|
## Enemy bullet in flight (absolute arena coords + fire power).
|
||||||
|
x*, y*: float64
|
||||||
|
power*: float64 # fire power in [0.1, 3.0]; speed = 20 - 3*power
|
||||||
|
|
||||||
|
EnemyData* = object
|
||||||
|
## Current enemy state, from the most recent onScannedBot event.
|
||||||
|
x*, y*: float64
|
||||||
|
direction*: float64
|
||||||
|
speed*: float64
|
||||||
|
energy*: float64
|
||||||
|
hasFired*: bool
|
||||||
|
lastFirePower*: float64
|
||||||
|
prevSpeed*: float64 # speed from the previous scan (for acceleration)
|
||||||
|
prevDirection*: float64 # direction from the previous scan (for turn rate)
|
||||||
|
hasPrevScan*: bool # true once we have at least two scans
|
||||||
|
|
||||||
|
GameState* = object
|
||||||
|
## Accumulates data from bot events. Populate fields before calling buildState.
|
||||||
|
# Own bot
|
||||||
|
x*, y*: float64
|
||||||
|
direction*: float64
|
||||||
|
speed*: float64
|
||||||
|
energy*: float64
|
||||||
|
gunDirection*: float64
|
||||||
|
gunHeat*: float64
|
||||||
|
arenaWidth*, arenaHeight*: float64
|
||||||
|
# Enemy
|
||||||
|
hasContact*: bool
|
||||||
|
enemy*: EnemyData
|
||||||
|
ticksSinceLastScan*: int
|
||||||
|
# Bullets in flight (up to 3 tracked)
|
||||||
|
bullets*: array[3, BulletData]
|
||||||
|
bulletCount*: int
|
||||||
|
|
||||||
|
proc buildState*(gs: GameState): Tensor[float32] =
|
||||||
|
## Build the 35-float normalized state tensor.
|
||||||
|
##
|
||||||
|
## Layout:
|
||||||
|
## [0-6] own bot: x/aW, y/aH, dir/360, speed/8, energy/100, gunDir/360, gunHeat/1.8
|
||||||
|
## [7-13] enemy: x/aW, y/aH, dir/360, speed/8, energy/100, hasFired, lastFirePower/3
|
||||||
|
## [14-17] derived: enemyAccel/8, enemyTurnRate/180, relBearing/180, distance/diag
|
||||||
|
## [18-21] walls: top, bottom, left, right — each / max(aW,aH)
|
||||||
|
## [22-33] bullets: up to 3 × (relX/aW, relY/aH, speed/20, ticksToImpact clamped to 1)
|
||||||
|
## [34] scan staleness: ticksSinceLastScan/30 clamped to 1
|
||||||
|
result = zeros[float32](STATE_DIM)
|
||||||
|
|
||||||
|
let aW = gs.arenaWidth
|
||||||
|
let aH = gs.arenaHeight
|
||||||
|
let diag = sqrt(aW * aW + aH * aH)
|
||||||
|
let wMax = max(aW, aH)
|
||||||
|
|
||||||
|
# --- Own bot (0-6) ---
|
||||||
|
result[0] = float32(gs.x / aW)
|
||||||
|
result[1] = float32(gs.y / aH)
|
||||||
|
result[2] = float32(gs.direction / 360.0)
|
||||||
|
result[3] = float32(gs.speed / 8.0)
|
||||||
|
result[4] = float32(gs.energy / 100.0)
|
||||||
|
result[5] = float32(gs.gunDirection / 360.0)
|
||||||
|
result[6] = float32(gs.gunHeat / 1.8)
|
||||||
|
|
||||||
|
# --- Enemy current (7-13) ---
|
||||||
|
if gs.hasContact:
|
||||||
|
result[7] = float32(gs.enemy.x / aW)
|
||||||
|
result[8] = float32(gs.enemy.y / aH)
|
||||||
|
result[9] = float32(gs.enemy.direction / 360.0)
|
||||||
|
result[10] = float32(gs.enemy.speed / 8.0)
|
||||||
|
result[11] = float32(gs.enemy.energy / 100.0)
|
||||||
|
result[12] = float32(if gs.enemy.hasFired: 1.0 else: 0.0)
|
||||||
|
result[13] = float32(gs.enemy.lastFirePower / 3.0)
|
||||||
|
|
||||||
|
# --- Derived (14-17) ---
|
||||||
|
if gs.hasContact:
|
||||||
|
if gs.enemy.hasPrevScan:
|
||||||
|
result[14] = float32((gs.enemy.speed - gs.enemy.prevSpeed) / 8.0)
|
||||||
|
let dDir = ((gs.enemy.direction - gs.enemy.prevDirection) + 540.0) mod 360.0 - 180.0
|
||||||
|
result[15] = float32(dDir / 180.0)
|
||||||
|
let dx = gs.enemy.x - gs.x
|
||||||
|
let dy = gs.enemy.y - gs.y
|
||||||
|
let absDir = (180.0 * arctan2(dx, dy) / PI + 360.0) mod 360.0
|
||||||
|
let relBearing = ((absDir - gs.direction) + 540.0) mod 360.0 - 180.0
|
||||||
|
result[16] = float32(relBearing / 180.0)
|
||||||
|
result[17] = float32(sqrt(dx * dx + dy * dy) / diag)
|
||||||
|
|
||||||
|
# --- Wall distances (18-21): top, bottom, left, right ---
|
||||||
|
result[18] = float32((aH - gs.y) / wMax)
|
||||||
|
result[19] = float32(gs.y / wMax)
|
||||||
|
result[20] = float32(gs.x / wMax)
|
||||||
|
result[21] = float32((aW - gs.x) / wMax)
|
||||||
|
|
||||||
|
# --- Bullet tracking (22-33): up to 3 bullets × 4 floats ---
|
||||||
|
# Per slot: relX/aW, relY/aH, speed/20, ticksToImpact/diag (clamped to 1)
|
||||||
|
for i in 0 ..< min(gs.bulletCount, 3):
|
||||||
|
let b = gs.bullets[i]
|
||||||
|
let bSpd = 20.0 - 3.0 * b.power
|
||||||
|
let bdx = b.x - gs.x
|
||||||
|
let bdy = b.y - gs.y
|
||||||
|
let bdist = sqrt(bdx * bdx + bdy * bdy)
|
||||||
|
let ticks = if bSpd > 0.0: min(bdist / bSpd / diag, 1.0) else: 0.0
|
||||||
|
let base = 22 + i * 4
|
||||||
|
result[base + 0] = float32(bdx / aW)
|
||||||
|
result[base + 1] = float32(bdy / aH)
|
||||||
|
result[base + 2] = float32(bSpd / 20.0)
|
||||||
|
result[base + 3] = float32(ticks)
|
||||||
|
|
||||||
|
# --- Scan staleness (34) ---
|
||||||
|
result[34] = float32(min(gs.ticksSinceLastScan.float64 / 30.0, 1.0))
|
||||||
@@ -0,0 +1,2 @@
|
|||||||
|
switch("path", "../src")
|
||||||
|
switch("path", "../../libs")
|
||||||
@@ -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"
|
||||||
@@ -0,0 +1,87 @@
|
|||||||
|
## Tests for state.nim — assert-based, no framework.
|
||||||
|
|
||||||
|
import std/math
|
||||||
|
import arraymancer
|
||||||
|
import SAC_LSTM_Bot/state
|
||||||
|
|
||||||
|
proc makeBase(): GameState =
|
||||||
|
result.arenaWidth = 1200.0
|
||||||
|
result.arenaHeight = 800.0
|
||||||
|
result.x = 600.0; result.y = 400.0
|
||||||
|
result.direction = 90.0; result.speed = 4.0
|
||||||
|
result.energy = 50.0
|
||||||
|
result.gunDirection = 90.0; result.gunHeat = 0.5
|
||||||
|
|
||||||
|
proc allInRange(t: Tensor[float32]): bool =
|
||||||
|
for v in t:
|
||||||
|
if v < -1.01f32 or v > 1.01f32: return false
|
||||||
|
true
|
||||||
|
|
||||||
|
proc hasNaN(t: Tensor[float32]): bool =
|
||||||
|
for v in t:
|
||||||
|
if v.float64.isNaN: return true
|
||||||
|
false
|
||||||
|
|
||||||
|
# 1. Correct shape
|
||||||
|
block:
|
||||||
|
let gs = makeBase()
|
||||||
|
let t = buildState(gs)
|
||||||
|
assert t.shape[0] == STATE_DIM, "wrong dim: " & $t.shape[0]
|
||||||
|
echo "PASS shape"
|
||||||
|
|
||||||
|
# 2. All values in [-1, 1] for typical input
|
||||||
|
block:
|
||||||
|
var gs = makeBase()
|
||||||
|
gs.hasContact = true
|
||||||
|
gs.enemy = EnemyData(x: 800.0, y: 300.0, direction: 45.0, speed: 6.0,
|
||||||
|
energy: 80.0, hasFired: true, lastFirePower: 2.0,
|
||||||
|
prevSpeed: 5.0, prevDirection: 40.0, hasPrevScan: true)
|
||||||
|
gs.bulletCount = 1
|
||||||
|
gs.bullets[0] = BulletData(x: 610.0, y: 410.0, power: 1.0)
|
||||||
|
gs.ticksSinceLastScan = 10
|
||||||
|
let t = buildState(gs)
|
||||||
|
assert not hasNaN(t), "NaN in tensor"
|
||||||
|
assert allInRange(t), "value out of [-1,1]"
|
||||||
|
echo "PASS range"
|
||||||
|
|
||||||
|
# 3. Missing scan → valid tensor, no NaN, zeros for enemy fields
|
||||||
|
block:
|
||||||
|
let gs = makeBase() # hasContact = false
|
||||||
|
let t = buildState(gs)
|
||||||
|
assert not hasNaN(t), "NaN with no scan"
|
||||||
|
for i in 7 .. 17:
|
||||||
|
assert t[i] == 0.0f32, "enemy slot " & $i & " should be 0"
|
||||||
|
echo "PASS no-scan zeros"
|
||||||
|
|
||||||
|
# 4. Bullet tracking: 0, 1, 2, 3 bullets
|
||||||
|
block:
|
||||||
|
for n in 0 .. 3:
|
||||||
|
var gs = makeBase()
|
||||||
|
gs.bulletCount = n
|
||||||
|
for i in 0 ..< n:
|
||||||
|
gs.bullets[i] = BulletData(x: float64(600 + i * 10), y: float64(400 + i * 5), power: 1.0)
|
||||||
|
let t = buildState(gs)
|
||||||
|
assert not hasNaN(t), "NaN with " & $n & " bullets"
|
||||||
|
# slots beyond bulletCount must be 0
|
||||||
|
for i in n ..< 3:
|
||||||
|
let base = 22 + i * 4
|
||||||
|
for j in 0 ..< 4:
|
||||||
|
assert t[base + j] == 0.0f32, "bullet slot " & $i & " slot " & $j & " should be 0"
|
||||||
|
echo "PASS bullet tracking 0-3"
|
||||||
|
|
||||||
|
# 5. Scan staleness increments and clamps
|
||||||
|
block:
|
||||||
|
var gs = makeBase()
|
||||||
|
gs.hasContact = true
|
||||||
|
gs.ticksSinceLastScan = 0
|
||||||
|
assert buildState(gs)[34] == 0.0f32, "staleness at 0"
|
||||||
|
gs.ticksSinceLastScan = 15
|
||||||
|
let mid = buildState(gs)[34]
|
||||||
|
assert mid > 0.0f32 and mid < 1.0f32, "staleness mid not in (0,1)"
|
||||||
|
gs.ticksSinceLastScan = 30
|
||||||
|
assert buildState(gs)[34] == 1.0f32, "staleness at 30"
|
||||||
|
gs.ticksSinceLastScan = 60
|
||||||
|
assert buildState(gs)[34] == 1.0f32, "staleness clamp at 60"
|
||||||
|
echo "PASS staleness"
|
||||||
|
|
||||||
|
echo "ALL TESTS PASSED"
|
||||||
@@ -0,0 +1,38 @@
|
|||||||
|
## Standalone radar lock module for Robocode Tank Royale (1v1).
|
||||||
|
## No bot-specific imports — takes plain floats, returns radarTurnRate.
|
||||||
|
##
|
||||||
|
## Tank Royale radar uses standard math convention: 0° = east, CCW positive.
|
||||||
|
## Angles are in degrees.
|
||||||
|
|
||||||
|
import std/math
|
||||||
|
|
||||||
|
const
|
||||||
|
MaxRadarTurn* = 45.0
|
||||||
|
DefaultOvershootDeg* = 5.0
|
||||||
|
|
||||||
|
var overShootDeg* = DefaultOvershootDeg
|
||||||
|
|
||||||
|
proc init*() =
|
||||||
|
## Reset module to defaults.
|
||||||
|
## In your bot's constructor set:
|
||||||
|
## adjustRadarForBodyTurn = true
|
||||||
|
## adjustRadarForGunTurn = true
|
||||||
|
overShootDeg = DefaultOvershootDeg
|
||||||
|
|
||||||
|
proc normalizeRelative(angle: float64): float64 {.inline.} =
|
||||||
|
result = angle mod 360.0
|
||||||
|
if result >= 180.0: result -= 360.0
|
||||||
|
elif result < -180.0: result += 360.0
|
||||||
|
|
||||||
|
proc doRadar*(currentRadarHeading, enemyBearing: float64): float64 =
|
||||||
|
## Returns radarTurnRate (degrees/tick, positive = clockwise).
|
||||||
|
##
|
||||||
|
## currentRadarHeading: current radar direction in degrees (0=east, CCW+).
|
||||||
|
## enemyBearing: absolute bearing to enemy in same coordinate system.
|
||||||
|
##
|
||||||
|
## Handles angle wrapping, clamps to [-45, +45], adds overshoot sweep.
|
||||||
|
var turn = normalizeRelative(enemyBearing - currentRadarHeading)
|
||||||
|
# Add overshoot in the same direction as the turn to maintain lock
|
||||||
|
if turn < 0.0: turn -= overShootDeg
|
||||||
|
else: turn += overShootDeg
|
||||||
|
result = turn.clamp(-MaxRadarTurn, MaxRadarTurn)
|
||||||
@@ -0,0 +1,8 @@
|
|||||||
|
# Package
|
||||||
|
version = "1.0.0"
|
||||||
|
author = "Davide Cappellini"
|
||||||
|
description = "Standalone radar lock module for Robocode Tank Royale"
|
||||||
|
license = "Apache-2.0"
|
||||||
|
|
||||||
|
# Dependencies
|
||||||
|
requires "nim >= 2.0.0"
|
||||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,74 @@
|
|||||||
|
import std/unittest
|
||||||
|
import std/math
|
||||||
|
import ../radar_lock
|
||||||
|
|
||||||
|
suite "radar_lock":
|
||||||
|
|
||||||
|
setup:
|
||||||
|
init()
|
||||||
|
|
||||||
|
# Basic lock: radar already on enemy — overshoot pushes it slightly
|
||||||
|
test "radar on enemy — returns overshoot only":
|
||||||
|
let rate = doRadar(90.0, 90.0)
|
||||||
|
check rate == DefaultOvershootDeg
|
||||||
|
|
||||||
|
# Cardinal directions
|
||||||
|
test "lock from 0 degrees":
|
||||||
|
let rate = doRadar(0.0, 0.0)
|
||||||
|
check rate == DefaultOvershootDeg
|
||||||
|
|
||||||
|
test "lock from 90 degrees":
|
||||||
|
let rate = doRadar(90.0, 90.0)
|
||||||
|
check rate == DefaultOvershootDeg
|
||||||
|
|
||||||
|
test "lock from 180 degrees":
|
||||||
|
let rate = doRadar(180.0, 180.0)
|
||||||
|
check rate == DefaultOvershootDeg
|
||||||
|
|
||||||
|
test "lock from 270 degrees":
|
||||||
|
let rate = doRadar(270.0, 270.0)
|
||||||
|
check rate == DefaultOvershootDeg
|
||||||
|
|
||||||
|
# Angle wrapping: radar at 350°, enemy at 10° → shortest path is +20°
|
||||||
|
test "wrapping 350 to 10 — turns right 20 + overshoot":
|
||||||
|
let rate = doRadar(350.0, 10.0)
|
||||||
|
check abs(rate - (20.0 + DefaultOvershootDeg)) < 1e-9
|
||||||
|
|
||||||
|
# Angle wrapping: radar at 10°, enemy at 350° → shortest path is -20°
|
||||||
|
test "wrapping 10 to 350 — turns left 20 + overshoot":
|
||||||
|
let rate = doRadar(10.0, 350.0)
|
||||||
|
check abs(rate - (-20.0 - DefaultOvershootDeg)) < 1e-9
|
||||||
|
|
||||||
|
# Clamping: enemy 100° away → raw turn+overshoot > 45°, must clamp
|
||||||
|
test "clamp positive — large gap":
|
||||||
|
let rate = doRadar(0.0, 100.0)
|
||||||
|
check rate == 45.0
|
||||||
|
|
||||||
|
test "clamp negative — large gap":
|
||||||
|
let rate = doRadar(100.0, 0.0)
|
||||||
|
check rate == -45.0
|
||||||
|
|
||||||
|
# Overshoot direction: turn right → positive overshoot
|
||||||
|
test "overshoot direction right":
|
||||||
|
let rate = doRadar(0.0, 30.0) # needs +30°, overshoot adds +5°
|
||||||
|
check abs(rate - 35.0) < 1e-9
|
||||||
|
|
||||||
|
# Overshoot direction: turn left → negative overshoot
|
||||||
|
test "overshoot direction left":
|
||||||
|
let rate = doRadar(30.0, 0.0) # needs -30°, overshoot adds -5°
|
||||||
|
check abs(rate - (-35.0)) < 1e-9
|
||||||
|
|
||||||
|
# Output never exceeds ±45
|
||||||
|
test "output bounded above":
|
||||||
|
let rate = doRadar(0.0, 179.0)
|
||||||
|
check rate <= 45.0
|
||||||
|
|
||||||
|
test "output bounded below":
|
||||||
|
let rate = doRadar(179.0, 0.0)
|
||||||
|
check rate >= -45.0
|
||||||
|
|
||||||
|
# init() resets overShootDeg
|
||||||
|
test "init resets overshoot":
|
||||||
|
overShootDeg = 20.0
|
||||||
|
init()
|
||||||
|
check overShootDeg == DefaultOvershootDeg
|
||||||
@@ -4,7 +4,7 @@ PPOB_LOG_FILE=/home/davide/Projects/SirRoboGarage/tools/training_runner/logs/fir
|
|||||||
PPOB_LR=1e-4
|
PPOB_LR=1e-4
|
||||||
PPOB_UPDATE_INTERVAL=10
|
PPOB_UPDATE_INTERVAL=10
|
||||||
PPOB_CLIP_EPSILON=0.2
|
PPOB_CLIP_EPSILON=0.2
|
||||||
PPOB_ENTROPY_COEFF=0.001
|
PPOB_ENTROPY_COEFF=0.0
|
||||||
PPOB_VALUE_LOSS_COEFF=0.5
|
PPOB_VALUE_LOSS_COEFF=0.5
|
||||||
PPOB_MAX_GRAD_NORM=0.5
|
PPOB_MAX_GRAD_NORM=0.5
|
||||||
PPOB_GAMMA=0.99
|
PPOB_GAMMA=0.99
|
||||||
@@ -12,7 +12,7 @@ PPOB_LAM=0.95
|
|||||||
PPOB_EPOCHS=4
|
PPOB_EPOCHS=4
|
||||||
PPOB_MINI_BATCH_SIZE=64
|
PPOB_MINI_BATCH_SIZE=64
|
||||||
PPOB_LOG_STD_FLOOR=-3.0
|
PPOB_LOG_STD_FLOOR=-3.0
|
||||||
PPOB_LOG_STD_CEILING=0.0
|
PPOB_LOG_STD_CEILING=-1.0
|
||||||
PPOB_INITIAL_LOG_STD=-0.5
|
PPOB_INITIAL_LOG_STD=-0.5
|
||||||
|
|
||||||
# PPOB_EVAL_ONLY=1 → freeze training (pure evaluation)
|
# PPOB_EVAL_ONLY=1 → freeze training (pure evaluation)
|
||||||
|
|||||||
+5
-1
@@ -31,13 +31,17 @@ for prefix in ("actor", "critic"):
|
|||||||
unchanged = [
|
unchanged = [
|
||||||
"actor_w2", "actor_w3", "actor_b1", "actor_b2", "actor_b3",
|
"actor_w2", "actor_w3", "actor_b1", "actor_b2", "actor_b3",
|
||||||
"critic_w2", "critic_w3", "critic_b1", "critic_b2", "critic_b3",
|
"critic_w2", "critic_w3", "critic_b1", "critic_b2", "critic_b3",
|
||||||
"log_std",
|
|
||||||
]
|
]
|
||||||
for name in unchanged:
|
for name in unchanged:
|
||||||
data = np.load(SRC / f"{name}.npy")
|
data = np.load(SRC / f"{name}.npy")
|
||||||
np.save(DST / f"{name}.npy", data)
|
np.save(DST / f"{name}.npy", data)
|
||||||
print(f" {name}: {data.shape} copied")
|
print(f" {name}: {data.shape} copied")
|
||||||
|
|
||||||
|
# Initialize log_std to -2.0 (std ≈ 0.135) — tighter than -1.0, proven workable
|
||||||
|
log_std = np.full(6, -2.0, dtype=np.float32)
|
||||||
|
np.save(DST / "log_std.npy", log_std)
|
||||||
|
print(f" log_std: initialized to -2.0 (std≈0.135), shape={log_std.shape}")
|
||||||
|
|
||||||
# Copy unchanged Adam moments (all except w1, which were handled above)
|
# Copy unchanged Adam moments (all except w1, which were handled above)
|
||||||
unchanged_adam = [
|
unchanged_adam = [
|
||||||
"adam_aw2", "adam_cw2",
|
"adam_aw2", "adam_cw2",
|
||||||
|
|||||||
Reference in New Issue
Block a user