a8ee2a86e3
Per-tick SVG drawText overlay above the bot showing round number and running average reward (e.g. "R:42 avg:3.50"). Per-round summary also printed to the UI console via printToStdOut with tick count and score. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
209 lines
6.8 KiB
Nim
209 lines
6.8 KiB
Nim
## PPO_Bot — enemy tracker + state vector wired into the game loop.
|
|
## Training: trajectory collected per tick, PPO update in background thread.
|
|
|
|
import std/[os, strformat, strutils, math]
|
|
import arraymancer
|
|
import tankroyale_botapi
|
|
import network
|
|
import actions
|
|
import training
|
|
import weights
|
|
import ./enemy_tracker
|
|
import ./state_vector
|
|
|
|
const botJsonPath = currentSourcePath().parentDir / "PPO_Bot.json"
|
|
const weightsRoot = currentSourcePath().parentDir / "weights"
|
|
|
|
type PPOBot = ref object of Bot
|
|
tracker: EnemyTracker
|
|
buffer: TrajectoryBuffer
|
|
prevEnergy: float32 # own energy last tick
|
|
prevEnemyE: float32 # enemy energy last tick (from tracker)
|
|
lastState: Tensor[float32]
|
|
lastAction: Tensor[float32]
|
|
lastLogP: float32
|
|
lastValue: float32
|
|
hasLastTrans: bool
|
|
roundRewardSum: float32 # cumulative reward this round (for live display)
|
|
roundTicks: int # ticks this round
|
|
|
|
var ac = initActorCritic()
|
|
var gAdamStates: ACAdamStates # persists across rounds
|
|
|
|
# ── Background training state ─────────────────────────────────────────────────
|
|
|
|
type
|
|
TrainingResult = object
|
|
ac: ActorCritic
|
|
adamStates: ACAdamStates
|
|
|
|
TrainingArgs = object
|
|
ac: ActorCritic
|
|
adamStates: ACAdamStates
|
|
buffer: TrajectoryBuffer
|
|
lastValue: float32
|
|
roundNum: int
|
|
weightsRoot: string
|
|
|
|
var
|
|
trainingThread: Thread[TrainingArgs]
|
|
resultChan: Channel[TrainingResult]
|
|
threadLaunched: bool = false # true while training thread is running
|
|
roundCounter: int = 0
|
|
|
|
proc trainingThreadProc(args: TrainingArgs) {.thread.} =
|
|
var localAc = args.ac
|
|
var localAdam = args.adamStates
|
|
ppoUpdate(localAc, args.buffer, lastValue = args.lastValue, adamStates = localAdam)
|
|
saveCheckpoint(localAc, args.weightsRoot, args.roundNum)
|
|
resultChan.send(TrainingResult(ac: localAc, adamStates: localAdam))
|
|
|
|
# ── Bot methods ───────────────────────────────────────────────────────────────
|
|
|
|
method onScannedBot*(bot: PPOBot, e: ScannedBotEvent) =
|
|
bot.tracker.update(e.x, e.y, e.direction, e.speed, e.energy)
|
|
|
|
method onRoundStarted*(bot: PPOBot, e: RoundStartedEvent) =
|
|
bot.tracker = initEnemyTracker()
|
|
bot.buffer = initTrajectoryBuffer()
|
|
bot.prevEnergy = 0.0'f32
|
|
bot.prevEnemyE = 0.0'f32
|
|
bot.hasLastTrans = false
|
|
bot.roundRewardSum = 0.0'f32
|
|
bot.roundTicks = 0
|
|
|
|
method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
|
|
inc roundCounter
|
|
|
|
# Add round-end score bonus to last transition (if any)
|
|
let roundReward = computeRoundReward(e.results.totalScore.float32)
|
|
if bot.hasLastTrans and bot.buffer.len > 0:
|
|
bot.buffer.transitions[^1].reward += roundReward
|
|
|
|
# Training progress display — one line per round in the UI console
|
|
let ticks = bot.buffer.len
|
|
if ticks > 0:
|
|
var rewardSum = 0.0'f32
|
|
for tr in bot.buffer.transitions: rewardSum += tr.reward
|
|
let avgR = rewardSum / ticks.float32
|
|
let avgRStr = formatFloat(avgR.float, ffDecimal, 3)
|
|
printToStdOut(&"R:{roundCounter} ticks:{ticks} avgR:{avgRStr} score:{e.results.totalScore}\n")
|
|
|
|
# Pick up result from previous training thread if available; channel IS the sync
|
|
if threadLaunched:
|
|
let (avail, trained) = resultChan.tryRecv()
|
|
if avail:
|
|
ac = trained.ac
|
|
gAdamStates = trained.adamStates
|
|
threadLaunched = false
|
|
|
|
if bot.buffer.len == 0:
|
|
bot.hasLastTrans = false
|
|
return
|
|
|
|
# If last thread still running, drop this pass — start fresh with newer data
|
|
# ponytail: simple drop; queue if every round must train
|
|
if threadLaunched:
|
|
bot.buffer.clear()
|
|
bot.hasLastTrans = false
|
|
return
|
|
|
|
let args = TrainingArgs(
|
|
ac: ac,
|
|
adamStates: gAdamStates,
|
|
buffer: bot.buffer,
|
|
lastValue: 0.0'f32,
|
|
roundNum: roundCounter,
|
|
weightsRoot: weightsRoot,
|
|
)
|
|
bot.buffer.clear()
|
|
bot.hasLastTrans = false
|
|
|
|
createThread(trainingThread, trainingThreadProc, args)
|
|
threadLaunched = true
|
|
|
|
method run(bot: PPOBot) =
|
|
# Seed energy on first tick
|
|
bot.prevEnergy = getEnergy().float32
|
|
bot.prevEnemyE = if bot.tracker.hasContact: bot.tracker.current.energy.float32 else: 0.0'f32
|
|
|
|
while isRunning():
|
|
bot.tracker.deadReckon()
|
|
|
|
setRadarTurnRate(bot.tracker.getRadarTurnRate(
|
|
getX(), getY(), getDirection(), getRadarDirection()))
|
|
|
|
let botData = BotStateData(
|
|
x: getX(),
|
|
y: getY(),
|
|
direction: getDirection(),
|
|
speed: getSpeed(),
|
|
energy: getEnergy(),
|
|
gunDirection: getGunDirection(),
|
|
gunHeat: getGunHeat(),
|
|
arenaWidth: float64(getArenaWidth()),
|
|
arenaHeight: float64(getArenaHeight()),
|
|
)
|
|
|
|
let state = buildStateVector(botData, bot.tracker)
|
|
let (rawActs, logP) = ac.actorForward(state)
|
|
let value = ac.criticForward(state)
|
|
let acts = mapActions(rawActs, getSpeed().float32, getGunHeat().float32)
|
|
|
|
# Compute tick reward from energy deltas
|
|
let curEnergy = getEnergy().float32
|
|
let curEnemyE = if bot.tracker.hasContact: bot.tracker.current.energy.float32 else: bot.prevEnemyE
|
|
let myDelta = curEnergy - bot.prevEnergy
|
|
let enemyDelta = curEnemyE - bot.prevEnemyE
|
|
let tickReward = computeTickReward(myDelta, enemyDelta)
|
|
|
|
# Track running reward for in-game display
|
|
bot.roundRewardSum += tickReward
|
|
inc bot.roundTicks
|
|
|
|
# Finalise previous transition with the reward from this tick's state change
|
|
if bot.hasLastTrans:
|
|
let tr = Transition(
|
|
state: bot.lastState,
|
|
action: bot.lastAction,
|
|
logProb: bot.lastLogP,
|
|
reward: tickReward,
|
|
value: bot.lastValue,
|
|
)
|
|
bot.buffer.add(tr)
|
|
|
|
# Store current for next tick
|
|
bot.lastState = state
|
|
bot.lastAction = rawActs
|
|
bot.lastLogP = logP
|
|
bot.lastValue = value
|
|
bot.prevEnergy = curEnergy
|
|
bot.prevEnemyE = curEnemyE
|
|
bot.hasLastTrans = true
|
|
|
|
setTargetSpeed(acts.targetSpeed.float)
|
|
setTurnRate(acts.turnRate.float)
|
|
setGunTurnRate(acts.gunTurnRate.float)
|
|
if acts.shouldFire:
|
|
discard setFire(acts.firePower.float)
|
|
|
|
# In-game training progress overlay
|
|
let avgR = if bot.roundTicks > 0: bot.roundRewardSum / bot.roundTicks.float32
|
|
else: 0.0'f32
|
|
let avgRStr = formatFloat(avgR.float, ffDecimal, 2)
|
|
drawText(&"R:{roundCounter} avg:{avgRStr}", getX(), getY() - 40.0)
|
|
|
|
go()
|
|
|
|
when isMainModule:
|
|
resultChan.open()
|
|
createDir(weightsRoot)
|
|
cleanStaleTempDirs(weightsRoot)
|
|
discard loadBestAvailable(ac, weightsRoot)
|
|
|
|
var bot = PPOBot(
|
|
tracker: initEnemyTracker(),
|
|
buffer: initTrajectoryBuffer(),
|
|
)
|
|
start(bot, botJsonPath)
|