feat(PPO_Bot): show training progress in game UI
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>
This commit is contained in:
+39
-15
@@ -1,7 +1,7 @@
|
|||||||
## PPO_Bot — enemy tracker + state vector wired into the game loop.
|
## PPO_Bot — enemy tracker + state vector wired into the game loop.
|
||||||
## Training: trajectory collected per tick, PPO update in background thread.
|
## Training: trajectory collected per tick, PPO update in background thread.
|
||||||
|
|
||||||
import std/os
|
import std/[os, strformat, strutils, math]
|
||||||
import arraymancer
|
import arraymancer
|
||||||
import tankroyale_botapi
|
import tankroyale_botapi
|
||||||
import network
|
import network
|
||||||
@@ -15,15 +15,17 @@ const botJsonPath = currentSourcePath().parentDir / "PPO_Bot.json"
|
|||||||
const weightsRoot = currentSourcePath().parentDir / "weights"
|
const weightsRoot = currentSourcePath().parentDir / "weights"
|
||||||
|
|
||||||
type PPOBot = ref object of Bot
|
type PPOBot = ref object of Bot
|
||||||
tracker: EnemyTracker
|
tracker: EnemyTracker
|
||||||
buffer: TrajectoryBuffer
|
buffer: TrajectoryBuffer
|
||||||
prevEnergy: float32 # own energy last tick
|
prevEnergy: float32 # own energy last tick
|
||||||
prevEnemyE: float32 # enemy energy last tick (from tracker)
|
prevEnemyE: float32 # enemy energy last tick (from tracker)
|
||||||
lastState: Tensor[float32]
|
lastState: Tensor[float32]
|
||||||
lastAction: Tensor[float32]
|
lastAction: Tensor[float32]
|
||||||
lastLogP: float32
|
lastLogP: float32
|
||||||
lastValue: float32
|
lastValue: float32
|
||||||
hasLastTrans: bool
|
hasLastTrans: bool
|
||||||
|
roundRewardSum: float32 # cumulative reward this round (for live display)
|
||||||
|
roundTicks: int # ticks this round
|
||||||
|
|
||||||
var ac = initActorCritic()
|
var ac = initActorCritic()
|
||||||
var gAdamStates: ACAdamStates # persists across rounds
|
var gAdamStates: ACAdamStates # persists across rounds
|
||||||
@@ -62,11 +64,13 @@ method onScannedBot*(bot: PPOBot, e: ScannedBotEvent) =
|
|||||||
bot.tracker.update(e.x, e.y, e.direction, e.speed, e.energy)
|
bot.tracker.update(e.x, e.y, e.direction, e.speed, e.energy)
|
||||||
|
|
||||||
method onRoundStarted*(bot: PPOBot, e: RoundStartedEvent) =
|
method onRoundStarted*(bot: PPOBot, e: RoundStartedEvent) =
|
||||||
bot.tracker = initEnemyTracker()
|
bot.tracker = initEnemyTracker()
|
||||||
bot.buffer = initTrajectoryBuffer()
|
bot.buffer = initTrajectoryBuffer()
|
||||||
bot.prevEnergy = 0.0'f32
|
bot.prevEnergy = 0.0'f32
|
||||||
bot.prevEnemyE = 0.0'f32
|
bot.prevEnemyE = 0.0'f32
|
||||||
bot.hasLastTrans = false
|
bot.hasLastTrans = false
|
||||||
|
bot.roundRewardSum = 0.0'f32
|
||||||
|
bot.roundTicks = 0
|
||||||
|
|
||||||
method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
|
method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
|
||||||
inc roundCounter
|
inc roundCounter
|
||||||
@@ -76,6 +80,15 @@ method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
|
|||||||
if bot.hasLastTrans and bot.buffer.len > 0:
|
if bot.hasLastTrans and bot.buffer.len > 0:
|
||||||
bot.buffer.transitions[^1].reward += roundReward
|
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
|
# Pick up result from previous training thread if available; channel IS the sync
|
||||||
if threadLaunched:
|
if threadLaunched:
|
||||||
let (avail, trained) = resultChan.tryRecv()
|
let (avail, trained) = resultChan.tryRecv()
|
||||||
@@ -144,6 +157,10 @@ method run(bot: PPOBot) =
|
|||||||
let enemyDelta = curEnemyE - bot.prevEnemyE
|
let enemyDelta = curEnemyE - bot.prevEnemyE
|
||||||
let tickReward = computeTickReward(myDelta, enemyDelta)
|
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
|
# Finalise previous transition with the reward from this tick's state change
|
||||||
if bot.hasLastTrans:
|
if bot.hasLastTrans:
|
||||||
let tr = Transition(
|
let tr = Transition(
|
||||||
@@ -169,6 +186,13 @@ method run(bot: PPOBot) =
|
|||||||
setGunTurnRate(acts.gunTurnRate.float)
|
setGunTurnRate(acts.gunTurnRate.float)
|
||||||
if acts.shouldFire:
|
if acts.shouldFire:
|
||||||
discard setFire(acts.firePower.float)
|
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()
|
go()
|
||||||
|
|
||||||
when isMainModule:
|
when isMainModule:
|
||||||
|
|||||||
Reference in New Issue
Block a user