feat(PPO_Bot): command abstraction layer — goto/aimTo controllers (#24)
- Add gotoTick/aimToTick controller functions (#25) - Update network dims: actor 5→6, state 42→44 (#26) - Rewrite mapActions for 6-dim command space (#27) - Delete stale weight files (shape mismatch) - Fix existing tests for new signatures Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
+22
-4
@@ -24,6 +24,7 @@ type PPOBot = ref object of Bot
|
||||
lastLogP: float32
|
||||
lastValue: float32
|
||||
hasLastTrans: bool
|
||||
lastActions: BotActions # previous tick's decoded actions (for state vector)
|
||||
roundRewardSum: float32 # cumulative reward this round (for live display)
|
||||
roundTicks: int # ticks this round
|
||||
|
||||
@@ -36,6 +37,7 @@ type
|
||||
TrainingResult = object
|
||||
ac: ActorCritic
|
||||
adamStates: ACAdamStates
|
||||
metrics: PPOMetrics
|
||||
|
||||
TrainingArgs = object
|
||||
ac: ActorCritic
|
||||
@@ -54,9 +56,9 @@ var
|
||||
proc trainingThreadProc(args: TrainingArgs) {.thread.} =
|
||||
var localAc = args.ac
|
||||
var localAdam = args.adamStates
|
||||
ppoUpdate(localAc, args.buffer, lastValue = args.lastValue, adamStates = localAdam)
|
||||
let m = ppoUpdate(localAc, args.buffer, lastValue = args.lastValue, adamStates = localAdam)
|
||||
saveCheckpoint(localAc, args.weightsRoot, args.roundNum)
|
||||
resultChan.send(TrainingResult(ac: localAc, adamStates: localAdam))
|
||||
resultChan.send(TrainingResult(ac: localAc, adamStates: localAdam, metrics: m))
|
||||
|
||||
# ── Bot methods ───────────────────────────────────────────────────────────────
|
||||
|
||||
@@ -91,6 +93,7 @@ method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
|
||||
let avgR = rewardSum / ticks.float32
|
||||
let avgRStr = formatFloat(avgR.float, ffDecimal, 3)
|
||||
printToStdOut(&"R:{roundCounter} ticks:{ticks} avgR:{avgRStr} score:{e.results.totalScore}\n")
|
||||
echo &"R:{roundCounter} ticks:{ticks} avgR:{avgRStr} score:{e.results.totalScore}"
|
||||
|
||||
# Pick up result from previous training thread if available; channel IS the sync
|
||||
if threadLaunched:
|
||||
@@ -99,6 +102,9 @@ method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
|
||||
ac = trained.ac
|
||||
gAdamStates = trained.adamStates
|
||||
threadLaunched = false
|
||||
let m = trained.metrics
|
||||
printToStdOut(&" trained R:{roundCounter-1} aLoss:{formatFloat(m.actorLoss.float, ffDecimal, 4)} vLoss:{formatFloat(m.valueLoss.float, ffDecimal, 4)} gNorm:{formatFloat(m.gradNorm.float, ffDecimal, 3)}\n")
|
||||
echo &" trained R:{roundCounter-1} aLoss:{formatFloat(m.actorLoss.float, ffDecimal, 4)} vLoss:{formatFloat(m.valueLoss.float, ffDecimal, 4)} gNorm:{formatFloat(m.gradNorm.float, ffDecimal, 3)}"
|
||||
|
||||
if bot.buffer.len == 0:
|
||||
bot.hasLastTrans = false
|
||||
@@ -122,6 +128,8 @@ method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
|
||||
bot.buffer.clear()
|
||||
bot.hasLastTrans = false
|
||||
|
||||
printToStdOut(&" train→ R:{roundCounter} ticks:{ticks}\n")
|
||||
echo &" train→ R:{roundCounter} ticks:{ticks}"
|
||||
createThread(trainingThread, trainingThreadProc, args)
|
||||
threadLaunched = true
|
||||
|
||||
@@ -147,10 +155,19 @@ method run(bot: PPOBot) =
|
||||
arenaHeight: float64(getArenaHeight()),
|
||||
)
|
||||
|
||||
let state = buildStateVector(botData, bot.tracker)
|
||||
let remainingGotoDistance = hypot(bot.lastActions.gotoX - botData.x,
|
||||
bot.lastActions.gotoY - botData.y)
|
||||
let remainingGunAngle = abs(normalizeRelativeAngle(
|
||||
directionTo(botData.x, botData.y, bot.lastActions.aimToX, bot.lastActions.aimToY) -
|
||||
botData.gunDirection))
|
||||
let state = buildStateVector(botData, bot.tracker, remainingGotoDistance, remainingGunAngle)
|
||||
let (rawActs, logP) = ac.actorForward(state)
|
||||
let value = ac.criticForward(state)
|
||||
let acts = mapActions(rawActs, getSpeed().float32, getGunHeat().float32)
|
||||
let acts = mapActions(rawActs,
|
||||
getGunHeat().float,
|
||||
botData.arenaWidth, botData.arenaHeight,
|
||||
botData.x, botData.y,
|
||||
botData.direction, botData.speed, botData.gunDirection)
|
||||
|
||||
# Compute tick reward from energy deltas
|
||||
let curEnergy = getEnergy().float32
|
||||
@@ -182,6 +199,7 @@ method run(bot: PPOBot) =
|
||||
bot.prevEnergy = curEnergy
|
||||
bot.prevEnemyE = curEnemyE
|
||||
bot.hasLastTrans = true
|
||||
bot.lastActions = acts
|
||||
|
||||
setTargetSpeed(acts.targetSpeed.float)
|
||||
setTurnRate(acts.turnRate.float)
|
||||
|
||||
Reference in New Issue
Block a user