chore: standardize internal bot dir structure (src/, tests/, out/)
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,11 @@
|
||||
{
|
||||
"name": "PPO_Bot",
|
||||
"version": "0.1.0",
|
||||
"authors": ["Davide Cappellini"],
|
||||
"description": "PPO-trained RL bot",
|
||||
"homepage": "",
|
||||
"countryCodes": ["IT"],
|
||||
"gameTypes": ["classic", "melee", "1v1"],
|
||||
"platform": "Nim",
|
||||
"programmingLang": "Nim"
|
||||
}
|
||||
@@ -0,0 +1,371 @@
|
||||
## 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, times, algorithm]
|
||||
import arraymancer
|
||||
import tankroyale_botapi
|
||||
import PPO_Bot/network
|
||||
import PPO_Bot/actions
|
||||
import PPO_Bot/training
|
||||
import PPO_Bot/weights
|
||||
import PPO_Bot/enemy_tracker
|
||||
import PPO_Bot/state_vector
|
||||
|
||||
# ── Hyperparameters from env vars (PPOB_ prefix) ─────────────────────────────
|
||||
# All optional; defaults match ppoUpdate signature in training.nim.
|
||||
|
||||
proc getEnvFloat(name: string, default: float32): float32 =
|
||||
let v = getEnv(name)
|
||||
if v.len == 0: default else: parseFloat(v).float32
|
||||
|
||||
proc getEnvInt(name: string, default: int): int =
|
||||
let v = getEnv(name)
|
||||
if v.len == 0: default else: parseInt(v)
|
||||
|
||||
var
|
||||
hpLr: float32 = getEnvFloat("PPOB_LR", 3e-4'f32)
|
||||
hpClipEpsilon: float32 = getEnvFloat("PPOB_CLIP_EPSILON", 0.2'f32)
|
||||
hpEntropyCoeff: float32 = getEnvFloat("PPOB_ENTROPY_COEFF", 0.01'f32)
|
||||
hpValueLossCoeff: float32 = getEnvFloat("PPOB_VALUE_LOSS_COEFF", 0.5'f32)
|
||||
hpMaxGradNorm: float32 = getEnvFloat("PPOB_MAX_GRAD_NORM", 0.5'f32)
|
||||
hpGamma: float32 = getEnvFloat("PPOB_GAMMA", 0.99'f32)
|
||||
hpLam: float32 = getEnvFloat("PPOB_LAM", 0.95'f32)
|
||||
hpEpochs: int = getEnvInt("PPOB_EPOCHS", 4)
|
||||
hpMiniBatchSize: int = getEnvInt("PPOB_MINI_BATCH_SIZE", 64)
|
||||
hpUpdateInterval: int = getEnvInt("PPOB_UPDATE_INTERVAL", 10)
|
||||
|
||||
# Wire logStd tunable params into network module vars (read before initActorCritic)
|
||||
logStdFloor = getEnvFloat("PPOB_LOG_STD_FLOOR", -3.0'f32)
|
||||
logStdCeiling = getEnvFloat("PPOB_LOG_STD_CEILING", 0.5'f32)
|
||||
initialLogStd = getEnvFloat("PPOB_INITIAL_LOG_STD", 0.0'f32)
|
||||
|
||||
# ── Structured log output ─────────────────────────────────────────────────────
|
||||
|
||||
let logFile = getEnv("PPOB_LOG_FILE") # empty → no JSON logging
|
||||
let evalOnly = getEnv("PPOB_EVAL_ONLY") == "1" # freeze training (pure evaluation)
|
||||
|
||||
proc appendJsonLine(path, line: string) =
|
||||
## Append a JSON line to path; no-op if path is empty.
|
||||
if path.len == 0: return
|
||||
let f = open(path, fmAppend)
|
||||
f.writeLine(line)
|
||||
f.close()
|
||||
|
||||
proc jsonFloat(v: float32): string =
|
||||
## Serialize a float for JSON; non-finite → null (keeps JSONL parseable).
|
||||
if v == v and v > -1e30'f32 and v < 1e30'f32: $v else: "null"
|
||||
|
||||
proc hyperparmSnapshot(): string =
|
||||
## Compact JSON object of current hyperparams (no outer braces).
|
||||
&"\"lr\":{hpLr},\"clipEpsilon\":{hpClipEpsilon}," &
|
||||
&"\"entropyCoeff\":{hpEntropyCoeff},\"valueLossCoeff\":{hpValueLossCoeff}," &
|
||||
&"\"maxGradNorm\":{hpMaxGradNorm},\"gamma\":{hpGamma},\"lam\":{hpLam}," &
|
||||
&"\"epochs\":{hpEpochs},\"miniBatchSize\":{hpMiniBatchSize}," &
|
||||
&"\"logStdFloor\":{logStdFloor},\"logStdCeiling\":{logStdCeiling},\"initialLogStd\":{initialLogStd}"
|
||||
|
||||
const botJsonPath = currentSourcePath().parentDir / "PPO_Bot.json"
|
||||
const weightsRoot = currentSourcePath().parentDir / "../weights"
|
||||
|
||||
type
|
||||
InFlightBullet = object
|
||||
x, y: float64 # current position
|
||||
vx, vy: float64 # velocity (pixels/tick)
|
||||
power: float64
|
||||
|
||||
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: array[STATE_DIM, float32] # plain arrays — tensors NEVER cross threads
|
||||
lastAction: array[ACTION_DIM, float32]
|
||||
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
|
||||
# Fixed-size bullet buffer — NO heap on the shared bot object. A seq here is
|
||||
# allocated by the per-round bot thread and freed by the next round's thread
|
||||
# (bot.bullets = @[] on round start) → foreign-heap free under --threads:on +
|
||||
# ORC → SIGSEGV. N=4: the state vector only consumes the closest 3 slots.
|
||||
bullets: array[4, InFlightBullet]
|
||||
bulletCount: int
|
||||
|
||||
var ac = initActorCritic()
|
||||
var gAdamStates: ACAdamStates # persists across rounds
|
||||
var roundCounter = 0
|
||||
var roundsSinceUpdate = 0
|
||||
|
||||
# ── 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) =
|
||||
setAdjustGunForBodyTurn(true)
|
||||
setAdjustRadarForBodyTurn(true)
|
||||
setAdjustRadarForGunTurn(true)
|
||||
bot.tracker = initEnemyTracker()
|
||||
# Do NOT clear bot.buffer here — transitions accumulate across rounds
|
||||
# until hpUpdateInterval rounds have passed (cleared in onRoundEnded).
|
||||
bot.prevEnergy = 0.0'f32
|
||||
bot.prevEnemyE = 0.0'f32
|
||||
bot.hasLastTrans = false
|
||||
bot.roundRewardSum = 0.0'f32
|
||||
bot.roundTicks = 0
|
||||
bot.bulletCount = 0
|
||||
|
||||
method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
|
||||
inc roundCounter
|
||||
inc roundsSinceUpdate
|
||||
debugLog("[PO-ENTER] round=" & $roundCounter & " tid=" & $getThreadId())
|
||||
|
||||
# Add round-end score bonus to last transition and mark it as episode boundary
|
||||
let roundReward = computeRoundReward(e.results.totalScore.float32)
|
||||
if bot.hasLastTrans and bot.buffer.len > 0:
|
||||
let lastIdx = bot.buffer.len - 1
|
||||
bot.buffer.transitions[lastIdx].reward += roundReward
|
||||
bot.buffer.transitions[lastIdx].done = true
|
||||
|
||||
# Training progress display — one line per round in the UI console
|
||||
let ticks = bot.buffer.len
|
||||
var avgR = 0.0'f32
|
||||
if ticks > 0:
|
||||
var rewardSum = 0.0'f32
|
||||
for i in 0 ..< bot.buffer.len: rewardSum += bot.buffer.transitions[i].reward
|
||||
avgR = rewardSum / ticks.float32
|
||||
let avgRStr = formatFloat(avgR.float, ffDecimal, 3)
|
||||
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}"
|
||||
|
||||
writeFile(weightsRoot / "round_counter.txt", $roundCounter)
|
||||
|
||||
if bot.buffer.len == 0:
|
||||
bot.hasLastTrans = false
|
||||
return
|
||||
|
||||
# Emit per-round game-stats JSON line
|
||||
let ts = int(epochTime())
|
||||
let jline = &"""{{"type":"round","round":{roundCounter},"ticks":{ticks},"bufLen":{bot.buffer.len},"avgReward":{avgR},"score":{e.results.totalScore},"ts":{ts}}}"""
|
||||
appendJsonLine(logFile, jline)
|
||||
|
||||
# PPOB_EVAL_ONLY=1 → freeze training (pure evaluation): skip ppoUpdate and
|
||||
# checkpoint save, but keep advancing/writing round_counter.txt so run.sh's
|
||||
# remaining-rounds bookkeeping still works, and keep the game line above.
|
||||
if evalOnly:
|
||||
bot.buffer.clear()
|
||||
bot.hasLastTrans = false
|
||||
roundsSinceUpdate = 0
|
||||
return
|
||||
|
||||
# Always save weights every round so round_counter.txt stays current.
|
||||
saveCheckpoint(ac, gAdamStates, weightsRoot, roundCounter)
|
||||
|
||||
# Only update policy every hpUpdateInterval rounds (~3000 transitions).
|
||||
if roundsSinceUpdate >= hpUpdateInterval:
|
||||
# ponytail: synchronous update. Arraymancer tensors can't cross threads under
|
||||
# ORC — training-thread ppoUpdate frees/rebinds tensors owned by the bot
|
||||
# thread's heap (SIGSEGV, reproduced with a lone trainer thread on a fixed
|
||||
# buffer; save/channel/forward exonerated). The bot API runs events on one
|
||||
# bot thread, so inline is single-threaded. Revert to a background thread
|
||||
# only if tensors are rebuilt from plain data on that thread.
|
||||
printToStdOut(&" train→ R:{roundCounter} bufLen:{bot.buffer.len}\n")
|
||||
echo &" train→ R:{roundCounter} bufLen:{bot.buffer.len}"
|
||||
let m = ppoUpdate(ac, bot.buffer,
|
||||
lastValue = 0.0'f32,
|
||||
adamStates = gAdamStates,
|
||||
epochs = hpEpochs,
|
||||
miniBatchSize = hpMiniBatchSize,
|
||||
clipEpsilon = hpClipEpsilon,
|
||||
entropyCoeff = hpEntropyCoeff,
|
||||
valueLossCoeff = hpValueLossCoeff,
|
||||
lr = hpLr,
|
||||
maxGradNorm = hpMaxGradNorm,
|
||||
gamma = hpGamma,
|
||||
lam = hpLam)
|
||||
printToStdOut(&" trained R:{roundCounter} 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} aLoss:{formatFloat(m.actorLoss.float, ffDecimal, 4)} vLoss:{formatFloat(m.valueLoss.float, ffDecimal, 4)} gNorm:{formatFloat(m.gradNorm.float, ffDecimal, 3)}"
|
||||
# Emit training-health JSON line
|
||||
let hp = hyperparmSnapshot()
|
||||
let ts2 = int(epochTime())
|
||||
let jline2 = &"""{{"type":"train","round":{roundCounter},"actorLoss":{jsonFloat(m.actorLoss)},"valueLoss":{jsonFloat(m.valueLoss)},"gradNorm":{jsonFloat(m.gradNorm)},"ts":{ts2},{hp}}}"""
|
||||
appendJsonLine(logFile, jline2)
|
||||
bot.buffer.clear()
|
||||
roundsSinceUpdate = 0
|
||||
|
||||
bot.hasLastTrans = false
|
||||
debugLog("[PO-EXIT] round=" & $roundCounter & " tid=" & $getThreadId())
|
||||
|
||||
method run(bot: PPOBot) =
|
||||
debugLog("[RUN-ENTER] tid=" & $getThreadId())
|
||||
# Seed energy and goto/aimTo targets on first tick (remainingDistance = 0 initially)
|
||||
bot.prevEnergy = getEnergy().float32
|
||||
bot.prevEnemyE = if bot.tracker.hasContact: bot.tracker.current.energy.float32 else: 0.0'f32
|
||||
bot.lastActions.gotoX = getX()
|
||||
bot.lastActions.gotoY = getY()
|
||||
bot.lastActions.aimToX = getX()
|
||||
bot.lastActions.aimToY = getY()
|
||||
|
||||
while isRunning():
|
||||
bot.tracker.deadReckon()
|
||||
|
||||
# Spawn bullet when enemy fired last scan tick. Edge-triggered: the flag is
|
||||
# consumed later (after the state build) so a radar gap (deadReckon ticks)
|
||||
# can't spawn k phantom bullets from one shot — exactly one bullet, and
|
||||
# state index 12 still pulses 1 on the detection tick.
|
||||
if bot.tracker.hasContact and bot.tracker.current.hasFired:
|
||||
if bot.bulletCount < 4:
|
||||
let power = bot.tracker.current.lastFirePower
|
||||
let speed = 20.0 - 3.0 * power
|
||||
# Approximate gun direction: bearing from enemy toward our position
|
||||
let myX = getX(); let myY = getY()
|
||||
let ang = arctan2(myY - bot.tracker.current.y, myX - bot.tracker.current.x)
|
||||
bot.bullets[bot.bulletCount] = InFlightBullet(
|
||||
x: bot.tracker.current.x,
|
||||
y: bot.tracker.current.y,
|
||||
vx: speed * cos(ang),
|
||||
vy: speed * sin(ang),
|
||||
power: power,
|
||||
)
|
||||
inc bot.bulletCount
|
||||
|
||||
# Advance in-flight bullets and prune those off-arena (in-place compaction
|
||||
# into the fixed buffer — no per-tick heap churn).
|
||||
let aW = float64(getArenaWidth()); let aH = float64(getArenaHeight())
|
||||
var n = 0
|
||||
for i in 0 ..< bot.bulletCount:
|
||||
let b = bot.bullets[i]
|
||||
let nx = b.x + b.vx; let ny = b.y + b.vy
|
||||
if nx >= 0.0 and nx <= aW and ny >= 0.0 and ny <= aH:
|
||||
bot.bullets[n] = InFlightBullet(x: nx, y: ny, vx: b.vx, vy: b.vy, power: b.power)
|
||||
inc n
|
||||
bot.bulletCount = n
|
||||
|
||||
# Convert to BulletData for state vector (closest 3 by distance to us)
|
||||
let myX2 = getX(); let myY2 = getY()
|
||||
var bulletData: array[4, BulletData]
|
||||
var bulletDataCount = 0
|
||||
for i in 0 ..< bot.bulletCount:
|
||||
let b = bot.bullets[i]
|
||||
bulletData[bulletDataCount] = BulletData(x: b.x, y: b.y, power: b.power)
|
||||
inc bulletDataCount
|
||||
# sort ascending by distance so the nearest threats fill slots 0-2
|
||||
if bulletDataCount > 1:
|
||||
bulletData.toOpenArray(0, bulletDataCount - 1).sort(proc(a, b: BulletData): int =
|
||||
let da = hypot(a.x - myX2, a.y - myY2)
|
||||
let db = hypot(b.x - myX2, b.y - myY2)
|
||||
cmp(da, db))
|
||||
|
||||
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 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, bulletData, bulletDataCount)
|
||||
# Consume the fired flag AFTER the state build: state index 12 saw the
|
||||
# detection-tick pulse, and the next iteration's spawn check sees false —
|
||||
# one shot → exactly one bullet, even across deadReckon gaps.
|
||||
bot.tracker.current.hasFired = false
|
||||
let (rawActs, logP) = ac.actorForward(state, deterministic = evalOnly)
|
||||
let value = ac.criticForward(state)
|
||||
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 acts = mapActions(rawActs,
|
||||
getGunHeat().float,
|
||||
botData.arenaWidth, botData.arenaHeight,
|
||||
botData.x, botData.y,
|
||||
botData.direction, botData.speed, botData.gunDirection,
|
||||
ex, ey)
|
||||
|
||||
# Compute tick reward from energy deltas + dense shaping
|
||||
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 arenaDiag = float32(sqrt(botData.arenaWidth * botData.arenaWidth +
|
||||
botData.arenaHeight * botData.arenaHeight))
|
||||
let distEnemy = if bot.tracker.hasContact:
|
||||
float32(hypot(bot.tracker.current.x - botData.x,
|
||||
bot.tracker.current.y - botData.y))
|
||||
else: arenaDiag
|
||||
let gunToEnemy = if bot.tracker.hasContact:
|
||||
abs(normalizeRelativeAngle(
|
||||
arctan2(bot.tracker.current.y - botData.y,
|
||||
bot.tracker.current.x - botData.x) * 180.0 / PI -
|
||||
botData.gunDirection)).float32
|
||||
else: 180.0'f32
|
||||
let tickReward = computeTickReward(myDelta, enemyDelta,
|
||||
distToEnemy = distEnemy,
|
||||
maxDist = arenaDiag,
|
||||
gunBearingAbs = gunToEnemy)
|
||||
|
||||
# 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,
|
||||
done: false, # episode boundary set in onRoundEnded
|
||||
)
|
||||
bot.buffer.add(tr)
|
||||
|
||||
# Store current for next tick — plain arrays only. `state`/`rawActs` tensors
|
||||
# live and die on this thread; a NEW bot thread runs each round, so storing
|
||||
# tensors in the shared bot object would free round-N's heap memory from
|
||||
# round N+1's thread (SIGSEGV; confirmed empirically).
|
||||
bot.lastState = stateToArr(state)
|
||||
bot.lastAction = actionToArr(rawActs)
|
||||
bot.lastLogP = logP
|
||||
bot.lastValue = value
|
||||
bot.prevEnergy = curEnergy
|
||||
bot.prevEnemyE = curEnemyE
|
||||
bot.hasLastTrans = true
|
||||
bot.lastActions = acts
|
||||
|
||||
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:
|
||||
createDir(weightsRoot)
|
||||
cleanStaleTempDirs(weightsRoot)
|
||||
let loadResult = loadBestAvailable(ac, gAdamStates, weightsRoot)
|
||||
if loadResult.loaded:
|
||||
roundCounter = loadResult.roundNum
|
||||
|
||||
var bot = PPOBot(
|
||||
tracker: initEnemyTracker(),
|
||||
buffer: initTrajectoryBuffer(),
|
||||
)
|
||||
start(bot, botJsonPath)
|
||||
@@ -0,0 +1,49 @@
|
||||
## actions.nim — map raw network output to Tank Royale bot commands.
|
||||
|
||||
import arraymancer
|
||||
import std/math
|
||||
import ./controllers
|
||||
|
||||
func sigmoid(x: float): float = 1.0 / (1.0 + exp(-x))
|
||||
|
||||
type
|
||||
BotActions* = object
|
||||
targetSpeed*: float
|
||||
turnRate*: float
|
||||
gunTurnRate*: float
|
||||
shouldFire*: bool
|
||||
firePower*: float
|
||||
gotoX*: float
|
||||
gotoY*: float
|
||||
aimToX*: float
|
||||
aimToY*: float
|
||||
|
||||
proc mapActions*(rawActions: Tensor[float32],
|
||||
gunHeat: float,
|
||||
arenaWidth, arenaHeight: float,
|
||||
botX, botY, direction, speed, gunDirection: float,
|
||||
enemyX, enemyY: float): BotActions =
|
||||
## rawActions: [6] tensor from actorForward.
|
||||
## Dims 0–1: goto x/y offset from enemy, 2–3: aimTo x/y offset from enemy,
|
||||
## 4: fire decision, 5: fire power.
|
||||
## Enemy-centred mapping: tanh gives [-1,1]; scale by arena/4 (goto) and
|
||||
## arena/8 (aimTo) so zero-init defaults the bot toward the enemy.
|
||||
let gotoX = clamp(enemyX + tanh(rawActions[0].float) * arenaWidth * 0.25, 0.0, arenaWidth)
|
||||
let gotoY = clamp(enemyY + tanh(rawActions[1].float) * arenaHeight * 0.25, 0.0, arenaHeight)
|
||||
let aimToX = clamp(enemyX + tanh(rawActions[2].float) * arenaWidth * 0.125, 0.0, arenaWidth)
|
||||
let aimToY = clamp(enemyY + tanh(rawActions[3].float) * arenaHeight * 0.125, 0.0, arenaHeight)
|
||||
let fireDec = tanh(rawActions[4].float)
|
||||
let fp = sigmoid(rawActions[5].float) * 2.9 + 0.1
|
||||
|
||||
let (ts, tr) = gotoTick(gotoX, gotoY, botX, botY, direction, speed)
|
||||
let gtr = aimToTick(aimToX, aimToY, botX, botY, gunDirection)
|
||||
|
||||
result.gotoX = gotoX
|
||||
result.gotoY = gotoY
|
||||
result.aimToX = aimToX
|
||||
result.aimToY = aimToY
|
||||
result.targetSpeed = ts
|
||||
result.turnRate = tr
|
||||
result.gunTurnRate = gtr
|
||||
result.shouldFire = fireDec >= 0.0 and gunHeat <= 0.0
|
||||
result.firePower = fp
|
||||
@@ -0,0 +1,24 @@
|
||||
## Pure tick-level controllers for goto(x,y) and aimTo(x,y).
|
||||
## No bot object needed — all inputs are explicit parameters.
|
||||
|
||||
import tankroyale_botapi
|
||||
|
||||
proc gotoTick*(targetX, targetY, botX, botY, direction, speed: float): (float, float) =
|
||||
## Returns (targetSpeed, turnRate) to drive toward (targetX, targetY).
|
||||
## Selects forward or reverse automatically based on bearing.
|
||||
let bearing = normalizeRelativeAngle(directionTo(botX, botY, targetX, targetY) - direction)
|
||||
let (dirSign, effBearing) =
|
||||
if abs(bearing) > 90.0:
|
||||
(-1.0, normalizeRelativeAngle(bearing + 180.0))
|
||||
else:
|
||||
(1.0, bearing)
|
||||
let dist = distanceTo(botX, botY, targetX, targetY)
|
||||
let maxTurn = calcMaxTurnRate(speed)
|
||||
let turnRate = effBearing.clamp(-maxTurn, maxTurn)
|
||||
let rawSpeed = getNewTargetSpeed(MAX_SPEED, abs(speed), dist)
|
||||
(dirSign * rawSpeed, turnRate)
|
||||
|
||||
proc aimToTick*(targetX, targetY, botX, botY, gunDirection: float): float =
|
||||
## Returns gunTurnRate (clamped to ±MAX_GUN_TURN_RATE) to rotate gun toward target.
|
||||
let bearing = normalizeRelativeAngle(directionTo(botX, botY, targetX, targetY) - gunDirection)
|
||||
bearing.clamp(-MAX_GUN_TURN_RATE, MAX_GUN_TURN_RATE)
|
||||
@@ -0,0 +1,100 @@
|
||||
## Enemy tracker — deterministic radar lock + dead reckoning for PPO_Bot.
|
||||
## No bot API imports; takes plain floats.
|
||||
|
||||
import std/math
|
||||
|
||||
type
|
||||
EnemyState* = object
|
||||
x*, y*: float64
|
||||
direction*: float64
|
||||
speed*: float64
|
||||
energy*: float64
|
||||
ticksSinceLastScan*: int
|
||||
hasFired*: bool
|
||||
lastFirePower*: float64
|
||||
|
||||
EnemyTracker* = object
|
||||
current*: EnemyState
|
||||
history*: array[5, tuple[x, y, direction, speed: float64]] # sliding window
|
||||
historyCount*: int # valid entries 0-5
|
||||
prevEnergy*: float64
|
||||
hasContact*: bool
|
||||
|
||||
proc initEnemyTracker*(): EnemyTracker = discard
|
||||
|
||||
proc update*(tracker: var EnemyTracker;
|
||||
scanX, scanY, scanDir, scanSpeed, scanEnergy: float64) =
|
||||
## Call on ScannedBotEvent. Detects enemy fire from energy delta.
|
||||
|
||||
# Shift history window
|
||||
if tracker.historyCount == 0:
|
||||
# Cold-start: pre-fill all slots with the incoming scan so indices 22-41
|
||||
# are never zero-padded on tick 1. Accel/turn-rate correctly stay 0 (no delta yet).
|
||||
for i in 0 ..< 5:
|
||||
tracker.history[i] = (scanX, scanY, scanDir, scanSpeed)
|
||||
tracker.historyCount = 5
|
||||
else:
|
||||
for i in countdown(min(tracker.historyCount, 4), 1):
|
||||
tracker.history[i] = tracker.history[i - 1]
|
||||
tracker.history[0] = (tracker.current.x, tracker.current.y,
|
||||
tracker.current.direction, tracker.current.speed)
|
||||
if tracker.historyCount < 5:
|
||||
inc tracker.historyCount
|
||||
|
||||
# Detect firing: energy drop in [0.1, 3.0] means enemy fired
|
||||
let delta = tracker.prevEnergy - scanEnergy
|
||||
if tracker.hasContact and delta >= 0.1 and delta <= 3.0:
|
||||
tracker.current.hasFired = true
|
||||
tracker.current.lastFirePower = delta
|
||||
else:
|
||||
tracker.current.hasFired = false
|
||||
|
||||
tracker.prevEnergy = scanEnergy
|
||||
tracker.current.x = scanX
|
||||
tracker.current.y = scanY
|
||||
tracker.current.direction = scanDir
|
||||
tracker.current.speed = scanSpeed
|
||||
tracker.current.energy = scanEnergy
|
||||
tracker.current.ticksSinceLastScan = 0
|
||||
tracker.hasContact = true
|
||||
|
||||
proc deadReckon*(tracker: var EnemyTracker) =
|
||||
## Call on missed ticks. Predict position from last known velocity.
|
||||
if not tracker.hasContact:
|
||||
return
|
||||
let rad = tracker.current.direction * PI / 180.0
|
||||
tracker.current.x += tracker.current.speed * sin(rad)
|
||||
tracker.current.y += tracker.current.speed * cos(rad)
|
||||
inc tracker.current.ticksSinceLastScan
|
||||
|
||||
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 getRadarTurnRate*(tracker: var EnemyTracker;
|
||||
botX, botY, botDirection, radarDirection: float64): float64 =
|
||||
## Returns radar turn rate (degrees/tick, positive = right).
|
||||
## Before contact: full 45° sweep.
|
||||
## After contact: lock with overshoot; widen if stale.
|
||||
if not tracker.hasContact:
|
||||
return 45.0
|
||||
|
||||
if tracker.current.ticksSinceLastScan >= 8:
|
||||
# Lost lock — widen sweep
|
||||
return 45.0
|
||||
|
||||
# Bearing from radar to enemy.
|
||||
# Tank Royale radar directions use standard math convention (0=east, CCW+).
|
||||
# arctan2(dy, dx) gives the standard math angle matching radarDirection units.
|
||||
let dx = tracker.current.x - botX
|
||||
let dy = tracker.current.y - botY
|
||||
let absoluteDir = (180.0 * arctan2(dy, dx) / PI + 360.0) mod 360.0
|
||||
var radarTurn = normalizeRelative(absoluteDir - radarDirection)
|
||||
|
||||
# Width Lock: overshoot proportional to arctan(36 / distance)
|
||||
let distance = sqrt(dx * dx + dy * dy)
|
||||
let extraTurn = min(arctan(36.0 / distance) * 180.0 / PI, 45.0)
|
||||
if radarTurn < 0.0: radarTurn -= extraTurn
|
||||
else: radarTurn += extraTurn
|
||||
result = radarTurn.clamp(-45.0, 45.0)
|
||||
@@ -0,0 +1,86 @@
|
||||
## network.nim — MLP and ActorCritic forward pass (inference only, no autograd).
|
||||
|
||||
import arraymancer
|
||||
import std/[math, random]
|
||||
|
||||
const
|
||||
STATE_DIM* = 57
|
||||
ACTION_DIM* = 6
|
||||
|
||||
var
|
||||
logStdFloor*: float32 = -3.0'f32 # overridden by PPOB_LOG_STD_FLOOR
|
||||
logStdCeiling*: float32 = 0.5'f32 # overridden by PPOB_LOG_STD_CEILING
|
||||
initialLogStd*: float32 = 0.0'f32 # overridden by PPOB_INITIAL_LOG_STD
|
||||
|
||||
type
|
||||
MLP* = object
|
||||
w1*, b1*: Tensor[float32] # [hidden, input], [hidden]
|
||||
w2*, b2*: Tensor[float32] # [hidden, hidden], [hidden]
|
||||
w3*, b3*: Tensor[float32] # [output, hidden], [output]
|
||||
|
||||
ActorCritic* = object
|
||||
actor*: MLP
|
||||
critic*: MLP
|
||||
logStd*: Tensor[float32] # [ACTION_DIM] — one per action dim
|
||||
|
||||
proc initMLP*(inputDim, hiddenDim, outputDim: int): MLP =
|
||||
# Xavier/He-style init: scale weights by sqrt(2/fan_in)
|
||||
result.w1 = randomNormalTensor[float32]([hiddenDim, inputDim]) *. sqrt(2.0'f32 / inputDim.float32)
|
||||
result.b1 = zeros[float32](hiddenDim)
|
||||
result.w2 = randomNormalTensor[float32]([hiddenDim, hiddenDim]) *. sqrt(2.0'f32 / hiddenDim.float32)
|
||||
result.b2 = zeros[float32](hiddenDim)
|
||||
result.w3 = randomNormalTensor[float32]([outputDim, hiddenDim]) *. sqrt(1.0'f32 / hiddenDim.float32)
|
||||
result.b3 = zeros[float32](outputDim)
|
||||
|
||||
proc initActorCritic*(): ActorCritic =
|
||||
result.actor = initMLP(STATE_DIM, 64, ACTION_DIM)
|
||||
result.critic = initMLP(STATE_DIM, 64, 1)
|
||||
result.logStd = newTensor[float32](ACTION_DIM).map(proc(v: float32): float32 = initialLogStd)
|
||||
|
||||
proc forward*(mlp: MLP, x: Tensor[float32]): Tensor[float32] =
|
||||
## x shape: [inputDim] (1D vector)
|
||||
let h1 = tanh(mlp.w1 * x + mlp.b1)
|
||||
let h2 = tanh(mlp.w2 * h1 + mlp.b2)
|
||||
result = mlp.w3 * h2 + mlp.b3
|
||||
|
||||
proc actorForward*(ac: ActorCritic, state: Tensor[float32], deterministic = false): tuple[actions: Tensor[float32], logProb: float32] =
|
||||
## 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)
|
||||
if deterministic:
|
||||
return (actions: mean, logProb: 0.0'f32)
|
||||
|
||||
# Floor logStd at logStdFloor before exp → min std ≈ exp(logStdFloor)
|
||||
var clampedLogStd = ac.logStd.map(proc(v: float32): float32 = clamp(v, logStdFloor, logStdCeiling))
|
||||
let std = clampedLogStd.map(proc(v: float32): float32 = exp(v))
|
||||
|
||||
var actions = newTensor[float32](ACTION_DIM)
|
||||
var logP = 0.0'f32
|
||||
for i in 0..<ACTION_DIM:
|
||||
let mu = mean[i]
|
||||
let s = std[i]
|
||||
let z = gauss(0.0'f64, 1.0'f64).float32
|
||||
actions[i] = mu + s * z
|
||||
# log N(a; mu, s) = -0.5*((a-mu)/s)^2 - log(s) - 0.5*log(2π)
|
||||
let diff = (actions[i] - mu) / s
|
||||
logP += -0.5'f32 * diff * diff - ln(s) - 0.5'f32 * ln(2.0'f32 * PI.float32)
|
||||
|
||||
result = (actions: actions, logProb: logP)
|
||||
|
||||
proc criticForward*(ac: ActorCritic, state: Tensor[float32]): float32 =
|
||||
## state: [STATE_DIM]. Returns scalar value estimate.
|
||||
let val = ac.critic.forward(state)
|
||||
result = val[0]
|
||||
|
||||
proc computeLogProb*(ac: ActorCritic, state, action: Tensor[float32]): float32 =
|
||||
## Log-probability of action under current policy (no sampling).
|
||||
let mean = ac.actor.forward(state)
|
||||
let logStdClamped = ac.logStd.map(proc(v: float32): float32 = clamp(v, logStdFloor, logStdCeiling))
|
||||
let std = logStdClamped.map(proc(v: float32): float32 = exp(v))
|
||||
var logP = 0.0'f32
|
||||
for i in 0..<ACTION_DIM:
|
||||
let mu = mean[i]
|
||||
let s = std[i]
|
||||
let diff = (action[i] - mu) / s
|
||||
logP += -0.5'f32 * diff * diff - ln(s) - 0.5'f32 * ln(2.0'f32 * PI.float32)
|
||||
result = logP
|
||||
@@ -0,0 +1,121 @@
|
||||
## State vector builder — produces 57-float normalized tensor for PPO policy.
|
||||
## No bot API imports; takes plain BotState + EnemyTracker structs.
|
||||
|
||||
import std/math
|
||||
import arraymancer
|
||||
import ./enemy_tracker
|
||||
|
||||
type
|
||||
BotStateData* = object
|
||||
## Plain data mirror of the bot's observable state.
|
||||
x*, y*: float64
|
||||
direction*: float64
|
||||
speed*: float64
|
||||
energy*: float64
|
||||
gunDirection*: float64
|
||||
gunHeat*: float64
|
||||
arenaWidth*, arenaHeight*: float64
|
||||
|
||||
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
|
||||
|
||||
proc buildStateVector*(bot: BotStateData; enemy: EnemyTracker;
|
||||
remainingGotoDistance: float64 = 0.0;
|
||||
remainingGunAngle: float64 = 0.0;
|
||||
bullets: openArray[BulletData] = [];
|
||||
bulletCount: int = 0): Tensor[float32] =
|
||||
## Build the 57-float normalized state tensor.
|
||||
## Indices 0-43: existing features. Indices 44-55: up to 3 bullet slots (4 floats each).
|
||||
## Index 56: scan staleness (ticksSinceLastScan / 30, clamped to 1).
|
||||
## All values clipped to roughly [-1, 1] via division by physical maxima.
|
||||
result = zeros[float32](57)
|
||||
|
||||
let aW = bot.arenaWidth
|
||||
let aH = bot.arenaHeight
|
||||
let diag = sqrt(aW * aW + aH * aH) # ≈ 1700 for 1200×800
|
||||
let wallMax = max(aW, aH)
|
||||
|
||||
# --- Current tick: own bot (indices 0-6) ---
|
||||
result[0] = float32(bot.x / aW)
|
||||
result[1] = float32(bot.y / aH)
|
||||
result[2] = float32(bot.direction / 360.0)
|
||||
result[3] = float32(bot.speed / 8.0)
|
||||
result[4] = float32(bot.energy / 100.0)
|
||||
result[5] = float32(bot.gunDirection / 360.0)
|
||||
result[6] = float32(bot.gunHeat / 1.8)
|
||||
|
||||
# --- Current tick: enemy (indices 7-13) ---
|
||||
if enemy.hasContact:
|
||||
result[7] = float32(enemy.current.x / aW)
|
||||
result[8] = float32(enemy.current.y / aH)
|
||||
result[9] = float32(enemy.current.direction / 360.0)
|
||||
result[10] = float32(enemy.current.speed / 8.0)
|
||||
result[11] = float32(enemy.current.energy / 100.0)
|
||||
result[12] = float32(if enemy.current.hasFired: 1.0 else: 0.0)
|
||||
result[13] = float32(enemy.current.lastFirePower / 3.0)
|
||||
# else: remain 0.0
|
||||
|
||||
# --- Derived features (indices 14-21) ---
|
||||
if enemy.hasContact:
|
||||
# Enemy acceleration (speed delta from last history entry)
|
||||
# Max delta is ±8 (stopped ↔ full speed); divide by 8 to normalize
|
||||
if enemy.historyCount >= 1:
|
||||
result[14] = float32((enemy.current.speed - enemy.history[0].speed) / 8.0)
|
||||
# Enemy turn rate (direction delta from last history entry, normalized to [-1,1])
|
||||
# Divide by 180 (max possible relative rotation) rather than 10 (max body turn rate)
|
||||
# ponytail: /180 covers all cases; use /10 if you want sensitivity to small turns
|
||||
if enemy.historyCount >= 1:
|
||||
let dirDelta = ((enemy.current.direction - enemy.history[0].direction) + 540.0) mod 360.0 - 180.0
|
||||
result[15] = float32(dirDelta / 180.0)
|
||||
# Relative bearing to enemy (signed, from bot perspective)
|
||||
let dx = enemy.current.x - bot.x
|
||||
let dy = enemy.current.y - bot.y
|
||||
let absDir = (180.0 * arctan2(dx, dy) / PI + 360.0) mod 360.0
|
||||
let relBearing = ((absDir - bot.direction) + 540.0) mod 360.0 - 180.0
|
||||
result[16] = float32(relBearing / 180.0)
|
||||
# Distance to enemy
|
||||
let dist = sqrt(dx * dx + dy * dy)
|
||||
result[17] = float32(dist / diag)
|
||||
|
||||
# Wall distances (indices 18-21): top, bottom, left, right
|
||||
# top = distance from bot to top wall (y=aH), bottom = distance to bottom (y=0)
|
||||
# left = distance to left (x=0), right = distance to right (x=aW)
|
||||
result[18] = float32((aH - bot.y) / wallMax) # top
|
||||
result[19] = float32(bot.y / wallMax) # bottom
|
||||
result[20] = float32(bot.x / wallMax) # left
|
||||
result[21] = float32((aW - bot.x) / wallMax) # right
|
||||
|
||||
# --- History: 5 ticks × 4 floats = 20 floats (indices 22-41) ---
|
||||
for i in 0 ..< 5:
|
||||
let base = 22 + i * 4
|
||||
if i < enemy.historyCount:
|
||||
result[base + 0] = float32(enemy.history[i].x / aW)
|
||||
result[base + 1] = float32(enemy.history[i].y / aH)
|
||||
result[base + 2] = float32(enemy.history[i].direction / 360.0)
|
||||
result[base + 3] = float32(enemy.history[i].speed / 8.0)
|
||||
# else: remain 0.0 (pad)
|
||||
|
||||
# --- Goto controller inputs (indices 42-43) ---
|
||||
result[42] = float32(remainingGotoDistance / diag)
|
||||
result[43] = float32(remainingGunAngle / 180.0)
|
||||
|
||||
# --- Bullet tracking (indices 44-55): up to 3 enemy bullets, 4 floats each ---
|
||||
# Per bullet: relX/aW, relY/aH, speed/20, ticksToImpact/diag
|
||||
# Positions are relative to bot (useful for dodging). Slots beyond bulletCount stay 0.
|
||||
for i in 0 ..< min(bulletCount, 3):
|
||||
let b = bullets[i]
|
||||
let bSpeed = 20.0 - 3.0 * b.power # Tank Royale bullet speed formula
|
||||
let bdx = b.x - bot.x
|
||||
let bdy = b.y - bot.y
|
||||
let bdist = sqrt(bdx * bdx + bdy * bdy)
|
||||
let ticks = if bSpeed > 0.0: bdist / bSpeed else: 0.0
|
||||
let base = 44 + i * 4
|
||||
result[base + 0] = float32(bdx / bot.arenaWidth)
|
||||
result[base + 1] = float32(bdy / bot.arenaHeight)
|
||||
result[base + 2] = float32(bSpeed / 20.0)
|
||||
result[base + 3] = float32(ticks / diag)
|
||||
|
||||
# --- Scan staleness (index 56) ---
|
||||
result[56] = float32(min(enemy.current.ticksSinceLastScan.float64 / 30.0, 1.0))
|
||||
@@ -0,0 +1,449 @@
|
||||
## training.nim — Trajectory buffer, GAE, and PPO training loop.
|
||||
## Uses manual backprop through the 3-layer tanh MLP + manual Adam.
|
||||
## No external autograd dependencies — pure Arraymancer Tensor math.
|
||||
|
||||
import arraymancer
|
||||
import std/[math, random, sequtils]
|
||||
import ./network
|
||||
|
||||
# ── Types ─────────────────────────────────────────────────────────────────────
|
||||
# ponytail: transitions hold PLAIN fixed-size arrays, never Arraymancer tensors.
|
||||
# Tensors crossing the bot-thread → main-thread boundary get freed on the wrong
|
||||
# thread's heap under ORC (bot thread SIGSEGVs mid-round in addEvent — 5 matching
|
||||
# coredumps). Plain arrays are value types: no heap, no GC, safe to move across
|
||||
# threads. Tensors are rebuilt from the arrays on the consuming (training) thread.
|
||||
|
||||
const
|
||||
MAX_TRANSITIONS* = 8192 # 10 rounds × ~300 ticks + headroom
|
||||
|
||||
type
|
||||
Transition* = object
|
||||
state*: array[STATE_DIM, float32] # plain copy, rebuilt as tensor in ppoUpdate
|
||||
action*: array[ACTION_DIM, float32]
|
||||
logProb*: float32
|
||||
reward*: float32
|
||||
value*: float32 # critic estimate at collection time
|
||||
done*: bool # true at episode (round) boundary
|
||||
|
||||
TrajectoryBuffer* = object
|
||||
transitions*: array[MAX_TRANSITIONS, Transition]
|
||||
len*: int
|
||||
|
||||
# ── Buffer ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
proc initTrajectoryBuffer*(): TrajectoryBuffer =
|
||||
result = TrajectoryBuffer()
|
||||
|
||||
proc add*(buf: var TrajectoryBuffer, t: Transition) =
|
||||
## ponytail: fixed 8192 cap — 10 rounds × ~300 ticks with headroom. Drops
|
||||
## new transitions when full. Raise cap if accumulation window grows.
|
||||
if buf.len < MAX_TRANSITIONS:
|
||||
buf.transitions[buf.len] = t
|
||||
inc buf.len
|
||||
|
||||
proc clear*(buf: var TrajectoryBuffer) =
|
||||
buf.len = 0
|
||||
|
||||
# ── Tensor → plain array (same-thread use; tensors never cross threads) ─────
|
||||
|
||||
proc stateToArr*(t: Tensor[float32]): array[STATE_DIM, float32] =
|
||||
for i in 0..<STATE_DIM: result[i] = t[i]
|
||||
|
||||
proc actionToArr*(t: Tensor[float32]): array[ACTION_DIM, float32] =
|
||||
for i in 0..<ACTION_DIM: result[i] = t[i]
|
||||
|
||||
# ── Reward helpers ─────────────────────────────────────────────────────────────
|
||||
|
||||
proc computeTickReward*(myEnergyDelta, enemyEnergyDelta: float32;
|
||||
distToEnemy: float32 = 0.0'f32;
|
||||
maxDist: float32 = 1.0'f32;
|
||||
gunBearingAbs: float32 = 180.0'f32): float32 =
|
||||
## Positive when we deal more damage than we receive.
|
||||
## Dense shaping: closeness (0-0.01/tick) + aim quality (0-0.02/tick).
|
||||
## ponytail: magnitudes 10x smaller than original to keep shaping as a nudge,
|
||||
## not the dominant signal. Increase if bot ignores positioning entirely.
|
||||
let sparseReward = myEnergyDelta - enemyEnergyDelta
|
||||
let distReward = 0.01'f32 * (1.0'f32 - distToEnemy / maxDist)
|
||||
let aimReward = 0.02'f32 * (1.0'f32 - gunBearingAbs / 180.0'f32)
|
||||
result = sparseReward + distReward + aimReward
|
||||
|
||||
proc computeRoundReward*(roundScore: float32): float32 =
|
||||
## Normalise round-end score to a rough ±6 range (doubled win signal).
|
||||
## ponytail: cap cumulative score at 400 before /50 — bounded terminal bonus
|
||||
## keeps critic value scale stable across battle boundaries and long battles.
|
||||
result = min(roundScore, 400.0'f32) / 50.0'f32
|
||||
|
||||
# ── GAE ───────────────────────────────────────────────────────────────────────
|
||||
|
||||
proc computeGAE*(rewards, values: seq[float32];
|
||||
dones: seq[bool];
|
||||
lastValue: float32;
|
||||
gamma: float32 = 0.99'f32;
|
||||
lam: float32 = 0.95'f32):
|
||||
tuple[advantages: seq[float32], returns: seq[float32]] =
|
||||
## Generalised Advantage Estimation — reverse sweep with episode boundaries.
|
||||
## When done=true on transition t, bootstrap value and accumulated GAE are
|
||||
## reset to 0 at that boundary (terminal state has no future value).
|
||||
let n = rewards.len
|
||||
var advantages = newSeq[float32](n)
|
||||
var lastGae = 0.0'f32
|
||||
|
||||
for t in countdown(n - 1, 0):
|
||||
let nextVal: float32 =
|
||||
if t == n - 1 or dones[t]: 0.0'f32
|
||||
else: values[t + 1]
|
||||
if t == n - 1 or dones[t]:
|
||||
lastGae = 0.0'f32
|
||||
let delta = rewards[t] + gamma * nextVal - values[t]
|
||||
lastGae = delta + gamma * lam * lastGae
|
||||
advantages[t] = lastGae
|
||||
|
||||
var returns = newSeq[float32](n)
|
||||
for t in 0..<n:
|
||||
returns[t] = advantages[t] + values[t]
|
||||
|
||||
result = (advantages: advantages, returns: returns)
|
||||
|
||||
# ── Manual Adam state ─────────────────────────────────────────────────────────
|
||||
|
||||
type
|
||||
AdamState* = object
|
||||
m*, v*: Tensor[float32]
|
||||
t*: int
|
||||
|
||||
proc initAdamState(like: Tensor[float32]): AdamState =
|
||||
result.m = zeros_like(like)
|
||||
result.v = zeros_like(like)
|
||||
result.t = 0
|
||||
|
||||
proc adamStep(param: var Tensor[float32];
|
||||
grad: Tensor[float32];
|
||||
state: var AdamState;
|
||||
lr: float32 = 3e-4'f32;
|
||||
beta1: float32 = 0.9'f32;
|
||||
beta2: float32 = 0.999'f32;
|
||||
eps: float32 = 1e-8'f32) =
|
||||
inc state.t
|
||||
state.m = beta1 *. state.m + (1.0'f32 - beta1) *. grad
|
||||
state.v = beta2 *. state.v + (1.0'f32 - beta2) *. (grad *. grad)
|
||||
let mHat = state.m /. (1.0'f32 - beta1 ^ state.t.float32)
|
||||
let vHat = state.v /. (1.0'f32 - beta2 ^ state.t.float32)
|
||||
param -= lr *. mHat /. (vHat.map(proc(x: float32): float32 = sqrt(x) + eps))
|
||||
|
||||
# ── MLP forward with cached activations (for backprop) ────────────────────────
|
||||
|
||||
type MLPFwd = object
|
||||
h1, h2, y: Tensor[float32] # activations (h1=layer1, h2=layer2, y=output)
|
||||
|
||||
proc mlpForwardCached(mlp: MLP; x: Tensor[float32]): MLPFwd =
|
||||
## Forward pass saving intermediate activations needed for backprop.
|
||||
result.h1 = tanh(mlp.w1 * x + mlp.b1)
|
||||
result.h2 = tanh(mlp.w2 * result.h1 + mlp.b2)
|
||||
result.y = mlp.w3 * result.h2 + mlp.b3
|
||||
|
||||
proc mlpBackward(mlp: MLP; fwd: MLPFwd; x: Tensor[float32];
|
||||
gradOut: Tensor[float32]):
|
||||
tuple[dw1, db1, dw2, db2, dw3, db3: Tensor[float32]] =
|
||||
## Chain-rule through 3-layer tanh MLP.
|
||||
## gradOut: [outputDim] — d_loss / d_out
|
||||
# Layer 3
|
||||
let dw3 = gradOut.unsqueeze(1) * fwd.h2.unsqueeze(0) # [out, hidden]
|
||||
let db3 = gradOut
|
||||
let dh2 = mlp.w3.transpose * gradOut # [hidden]
|
||||
# tanh backward: d/dx tanh(x) = 1 - tanh²(x)
|
||||
let dpre2 = dh2 *. (ones[float32](fwd.h2.shape) - fwd.h2 *. fwd.h2)
|
||||
# Layer 2
|
||||
let dw2 = dpre2.unsqueeze(1) * fwd.h1.unsqueeze(0) # [hidden, hidden]
|
||||
let db2 = dpre2
|
||||
let dh1 = mlp.w2.transpose * dpre2 # [hidden]
|
||||
let dpre1 = dh1 *. (ones[float32](fwd.h1.shape) - fwd.h1 *. fwd.h1)
|
||||
# Layer 1
|
||||
let dw1 = dpre1.unsqueeze(1) * x.unsqueeze(0) # [hidden, input]
|
||||
let db1 = dpre1
|
||||
result = (dw1: dw1, db1: db1, dw2: dw2, db2: db2, dw3: dw3, db3: db3)
|
||||
|
||||
# ── Adam states for ActorCritic parameters ───────────────────────────────────
|
||||
|
||||
type ACAdamStates* = object
|
||||
## One AdamState per learnable tensor in ActorCritic.
|
||||
aw1*, ab1*, aw2*, ab2*, aw3*, ab3*: AdamState # actor MLP
|
||||
cw1*, cb1*, cw2*, cb2*, cw3*, cb3*: AdamState # critic MLP
|
||||
logStd*: AdamState
|
||||
initialized*: bool
|
||||
|
||||
proc initACAdamStates*(ac: ActorCritic): ACAdamStates =
|
||||
result.aw1 = initAdamState(ac.actor.w1)
|
||||
result.ab1 = initAdamState(ac.actor.b1)
|
||||
result.aw2 = initAdamState(ac.actor.w2)
|
||||
result.ab2 = initAdamState(ac.actor.b2)
|
||||
result.aw3 = initAdamState(ac.actor.w3)
|
||||
result.ab3 = initAdamState(ac.actor.b3)
|
||||
result.cw1 = initAdamState(ac.critic.w1)
|
||||
result.cb1 = initAdamState(ac.critic.b1)
|
||||
result.cw2 = initAdamState(ac.critic.w2)
|
||||
result.cb2 = initAdamState(ac.critic.b2)
|
||||
result.cw3 = initAdamState(ac.critic.w3)
|
||||
result.cb3 = initAdamState(ac.critic.b3)
|
||||
result.logStd = initAdamState(ac.logStd)
|
||||
result.initialized = true
|
||||
|
||||
# ── Training metrics ──────────────────────────────────────────────────────────
|
||||
|
||||
type PPOMetrics* = object
|
||||
actorLoss*: float32
|
||||
valueLoss*: float32
|
||||
gradNorm*: float32
|
||||
|
||||
# ── Gradient clipping ─────────────────────────────────────────────────────────
|
||||
|
||||
proc globalNorm(grads: varargs[Tensor[float32]]): float32 =
|
||||
var sumSq = 0.0'f32
|
||||
for g in grads:
|
||||
for v in g: sumSq += v * v
|
||||
result = sqrt(sumSq)
|
||||
|
||||
# ── PPO update ────────────────────────────────────────────────────────────────
|
||||
|
||||
proc ppoUpdate*(ac: var ActorCritic;
|
||||
buffer: TrajectoryBuffer;
|
||||
lastValue: float32;
|
||||
adamStates: var ACAdamStates;
|
||||
epochs: int = 4;
|
||||
miniBatchSize: int = 64;
|
||||
clipEpsilon: float32 = 0.2'f32;
|
||||
entropyCoeff: float32 = 0.01'f32;
|
||||
valueLossCoeff: float32 = 0.5'f32;
|
||||
lr: float32 = 3e-4'f32;
|
||||
maxGradNorm: float32 = 0.5'f32;
|
||||
gamma: float32 = 0.99'f32;
|
||||
lam: float32 = 0.95'f32): PPOMetrics {.gcsafe.} =
|
||||
if buffer.len == 0: return
|
||||
|
||||
var totalActorLoss = 0.0'f32
|
||||
var totalValueLoss = 0.0'f32
|
||||
var totalGradNorm = 0.0'f32
|
||||
var totalMiniBatches = 0
|
||||
|
||||
# Initialise Adam states once; caller persists them across rounds.
|
||||
# Also reinit if aw1.m has wrong shape (e.g. loaded from old checkpoint with
|
||||
# different STATE_DIM, leaving a (0,) placeholder after shape-mismatch skip).
|
||||
if not adamStates.initialized or
|
||||
adamStates.aw1.m.shape.len == 0 or
|
||||
adamStates.aw1.m.shape != ac.actor.w1.shape:
|
||||
adamStates = initACAdamStates(ac)
|
||||
|
||||
# 1. GAE
|
||||
let rewards = buffer.transitions[0 ..< buffer.len].mapIt(it.reward)
|
||||
let values = buffer.transitions[0 ..< buffer.len].mapIt(it.value)
|
||||
let dones = buffer.transitions[0 ..< buffer.len].mapIt(it.done)
|
||||
let (advantages, returns) = computeGAE(rewards, values, dones, lastValue, gamma = gamma, lam = lam)
|
||||
|
||||
# 2. Normalise advantages
|
||||
let n = advantages.len.float32
|
||||
var advMean = 0.0'f32
|
||||
for a in advantages: advMean += a
|
||||
advMean /= n
|
||||
var advVar = 0.0'f32
|
||||
for a in advantages: advVar += (a - advMean) * (a - advMean)
|
||||
advVar /= n
|
||||
# ponytail: float32 adv noise ~1e-12; advVar < 1e-8 = constant-reward
|
||||
# (passive) round — dividing by that amplifies noise ~1e4+ and drifts the
|
||||
# policy into exp() overflow. Center-only, skip the divide.
|
||||
var normAdv: seq[float32]
|
||||
if advantages.allIt(it == it and abs(it) < 1e30'f32):
|
||||
if advVar < 1e-8'f32:
|
||||
normAdv = advantages.mapIt(it - advMean)
|
||||
else:
|
||||
let advStd = sqrt(advVar + 1e-8'f32)
|
||||
normAdv = advantages.mapIt((it - advMean) / advStd)
|
||||
else:
|
||||
normAdv = newSeq[float32](advantages.len) # poisoned input → zero advantages, no-op update
|
||||
|
||||
let bufLen = buffer.len
|
||||
# ponytail: minibatch size <= 0 would make mbEnd == mbStart forever and spin.
|
||||
# Treat as full-batch; breaks the loop unconditionally.
|
||||
let mbSizeCap = if miniBatchSize > 0: miniBatchSize else: bufLen
|
||||
|
||||
for epochNum in 1..epochs:
|
||||
# Shuffle indices
|
||||
var indices = toSeq(0..<bufLen)
|
||||
shuffle(indices)
|
||||
|
||||
var mbStart = 0
|
||||
while mbStart < bufLen:
|
||||
let mbEnd = min(mbStart + mbSizeCap, bufLen)
|
||||
if mbEnd <= mbStart: break
|
||||
let mbSize = mbEnd - mbStart
|
||||
# Accumulators for gradients (zero-init)
|
||||
var dActorW1 = zeros[float32](ac.actor.w1.shape)
|
||||
var dActorB1 = zeros[float32](ac.actor.b1.shape)
|
||||
var dActorW2 = zeros[float32](ac.actor.w2.shape)
|
||||
var dActorB2 = zeros[float32](ac.actor.b2.shape)
|
||||
var dActorW3 = zeros[float32](ac.actor.w3.shape)
|
||||
var dActorB3 = zeros[float32](ac.actor.b3.shape)
|
||||
var dLogStd = zeros[float32](ac.logStd.shape)
|
||||
|
||||
var dCriticW1 = zeros[float32](ac.critic.w1.shape)
|
||||
var dCriticB1 = zeros[float32](ac.critic.b1.shape)
|
||||
var dCriticW2 = zeros[float32](ac.critic.w2.shape)
|
||||
var dCriticB2 = zeros[float32](ac.critic.b2.shape)
|
||||
var dCriticW3 = zeros[float32](ac.critic.w3.shape)
|
||||
var dCriticB3 = zeros[float32](ac.critic.b3.shape)
|
||||
|
||||
for j in mbStart..<mbEnd:
|
||||
let idx = indices[j]
|
||||
let tr = buffer.transitions[idx]
|
||||
let adv = normAdv[idx]
|
||||
let ret = returns[idx].float32
|
||||
|
||||
# Rebuild the state tensor on this (training) thread — transitions hold
|
||||
# plain arrays so no tensor ever crosses a thread boundary.
|
||||
let x = tr.state.toTensor()
|
||||
|
||||
# ── Actor forward ──
|
||||
let actorFwd = mlpForwardCached(ac.actor, x)
|
||||
# Critic forward
|
||||
|
||||
let newMean = actorFwd.y # [ACTION_DIM]
|
||||
|
||||
# Same clamp as collection (network.nim actorForward): train-time std must
|
||||
# exactly match the std the acting policy used, or ratios are distorted.
|
||||
let logStdClamped = ac.logStd.map(proc(v: float32): float32 = clamp(v, logStdFloor, logStdCeiling))
|
||||
let std = logStdClamped.map(proc(v: float32): float32 = exp(v))
|
||||
|
||||
# New log prob
|
||||
var newLogP = 0.0'f32
|
||||
for i in 0..<ACTION_DIM:
|
||||
let mu = newMean[i]
|
||||
let s = std[i]
|
||||
let diff = (tr.action[i] - mu) / s
|
||||
newLogP += -0.5'f32 * diff * diff - ln(s) - 0.5'f32 * ln(2.0'f32 * PI.float32)
|
||||
|
||||
# ponytail: float32 exp overflows at ±88; ±20 is deep in clipped-ratio
|
||||
# territory, so loss/grad are identical to the true ratio
|
||||
let ratio = exp(clamp(newLogP - tr.logProb, -20.0'f32, 20.0'f32))
|
||||
|
||||
# Clipped surrogate
|
||||
let ratioClipped = clamp(ratio, 1.0'f32 - clipEpsilon, 1.0'f32 + clipEpsilon)
|
||||
let surr1 = ratio * adv
|
||||
let surr2 = ratioClipped * adv
|
||||
# Actor loss per sample = -min(surr1, surr2)
|
||||
totalActorLoss += -min(surr1, surr2)
|
||||
# Which branch is active?
|
||||
let useClipped = (surr2 < surr1)
|
||||
let dLoss_dSurr = -1.0'f32 / mbSize.float32 # d(-mean(min))/d(min) = -1/N
|
||||
|
||||
# d(actor_loss)/d(ratio): only the non-clipped branch passes gradient
|
||||
let dLoss_dRatio = if useClipped: 0.0'f32 else: dLoss_dSurr * adv
|
||||
# d(ratio)/d(newLogP) = ratio
|
||||
let dLoss_dNewLogP = dLoss_dRatio * ratio
|
||||
|
||||
# d(newLogP)/d(mean[i]) = (action[i] - mean[i]) / std[i]^2
|
||||
var dLogP_dMean = newTensor[float32](ACTION_DIM)
|
||||
for i in 0..<ACTION_DIM:
|
||||
let s = std[i]
|
||||
dLogP_dMean[i] = (tr.action[i] - newMean[i]) / (s * s)
|
||||
|
||||
# Entropy gradient for logStd:
|
||||
# entropy = sum_i [ logStd_i + 0.5*(1+ln(2π)) ]
|
||||
# d(entropy)/d(logStd_i) = 1 (for clamped logStd_i > -3, else 0)
|
||||
# total loss gradient w.r.t. logStd: -entropyCoeff * d(entropy)/d(logStd)
|
||||
# Also: d(newLogP)/d(logStd_i) when logStd_i > -3:
|
||||
# = (action_i - mean_i)^2/std_i^2 - 1
|
||||
for i in 0..<ACTION_DIM:
|
||||
let isFloorClamped = (ac.logStd[i] <= logStdFloor)
|
||||
let isCeilingClamped = (ac.logStd[i] >= logStdCeiling)
|
||||
if not isFloorClamped:
|
||||
let s = std[i]
|
||||
let diff = (tr.action[i] - newMean[i]) / s
|
||||
let dLogP_dLogStdI = diff * diff - 1.0'f32
|
||||
let ppoGrad = dLoss_dNewLogP * dLogP_dLogStdI
|
||||
# Entropy term pushes logStd up (update = param - lr*grad, grad is -entropyCoeff < 0).
|
||||
# Gate it off at the ceiling to prevent runaway logStd.
|
||||
let entropyGrad = if isCeilingClamped: 0.0'f32
|
||||
else: -entropyCoeff / mbSize.float32
|
||||
dLogStd[i] += ppoGrad + entropyGrad
|
||||
|
||||
# Backprop actor gradients
|
||||
let gradActorOut = dLoss_dNewLogP *. dLogP_dMean # [5]
|
||||
let actorGrads = mlpBackward(ac.actor, actorFwd, x, gradActorOut)
|
||||
|
||||
dActorW1 += actorGrads.dw1
|
||||
dActorB1 += actorGrads.db1
|
||||
dActorW2 += actorGrads.dw2
|
||||
dActorB2 += actorGrads.db2
|
||||
dActorW3 += actorGrads.dw3
|
||||
dActorB3 += actorGrads.db3
|
||||
|
||||
# ── Critic forward + loss ──
|
||||
let criticFwd = mlpForwardCached(ac.critic, x)
|
||||
let newVal = criticFwd.y[0]
|
||||
# Value loss = (newVal - ret)^2; d/d(newVal) = 2*(newVal-ret)
|
||||
totalValueLoss += (newVal - ret) * (newVal - ret)
|
||||
let dVLoss_dVal = valueLossCoeff * 2.0'f32 * (newVal - ret) / mbSize.float32
|
||||
let gradCriticOut = [dVLoss_dVal].toTensor() # [1]
|
||||
let criticGrads = mlpBackward(ac.critic, criticFwd, x, gradCriticOut)
|
||||
|
||||
dCriticW1 += criticGrads.dw1
|
||||
dCriticB1 += criticGrads.db1
|
||||
dCriticW2 += criticGrads.dw2
|
||||
dCriticB2 += criticGrads.db2
|
||||
dCriticW3 += criticGrads.dw3
|
||||
dCriticB3 += criticGrads.db3
|
||||
|
||||
# ── Gradient clipping ──
|
||||
# Collect all grads into a seq for norm computation
|
||||
var allGrads: seq[Tensor[float32]] = @[
|
||||
dActorW1, dActorB1, dActorW2, dActorB2, dActorW3, dActorB3,
|
||||
dLogStd,
|
||||
dCriticW1, dCriticB1, dCriticW2, dCriticB2, dCriticW3, dCriticB3
|
||||
]
|
||||
let norm = globalNorm(allGrads)
|
||||
# ponytail: NaN/Inf grad norm means something exploded this minibatch
|
||||
# (extreme logprob ratios, poisoned initial weights, etc.). Skip the
|
||||
# Adam update entirely — no-op is safer than writing NaN into weights,
|
||||
# which corrupts all future inference and hangs the bot.
|
||||
if norm != norm or norm > 1e15'f32:
|
||||
mbStart = mbEnd
|
||||
continue
|
||||
totalGradNorm += norm
|
||||
inc totalMiniBatches
|
||||
if norm > maxGradNorm:
|
||||
let scale = maxGradNorm / norm
|
||||
for g in allGrads.mitems: g = g *. scale
|
||||
|
||||
# Unpack clipped grads
|
||||
dActorW1 = allGrads[0]; dActorB1 = allGrads[1]
|
||||
dActorW2 = allGrads[2]; dActorB2 = allGrads[3]
|
||||
dActorW3 = allGrads[4]; dActorB3 = allGrads[5]
|
||||
dLogStd = allGrads[6]
|
||||
dCriticW1 = allGrads[7]; dCriticB1 = allGrads[8]
|
||||
dCriticW2 = allGrads[9]; dCriticB2 = allGrads[10]
|
||||
dCriticW3 = allGrads[11]; dCriticB3 = allGrads[12]
|
||||
|
||||
# ── Adam updates ──
|
||||
adamStep(ac.actor.w1, dActorW1, adamStates.aw1, lr)
|
||||
adamStep(ac.actor.b1, dActorB1, adamStates.ab1, lr)
|
||||
adamStep(ac.actor.w2, dActorW2, adamStates.aw2, lr)
|
||||
adamStep(ac.actor.b2, dActorB2, adamStates.ab2, lr)
|
||||
adamStep(ac.actor.w3, dActorW3, adamStates.aw3, lr)
|
||||
adamStep(ac.actor.b3, dActorB3, adamStates.ab3, lr)
|
||||
adamStep(ac.logStd, dLogStd, adamStates.logStd, lr)
|
||||
# Collection clamps logStd to [floor, ceiling] at inference; clamp the raw
|
||||
# param after the step so it can't drift above the ceiling (the old code
|
||||
# only clamped at collection → train-time recompute used a bigger std than
|
||||
# the policy that actually acted → distorted importance ratios).
|
||||
ac.logStd = ac.logStd.map(proc(v: float32): float32 = clamp(v, logStdFloor, logStdCeiling))
|
||||
adamStep(ac.critic.w1, dCriticW1, adamStates.cw1, lr)
|
||||
adamStep(ac.critic.b1, dCriticB1, adamStates.cb1, lr)
|
||||
adamStep(ac.critic.w2, dCriticW2, adamStates.cw2, lr)
|
||||
adamStep(ac.critic.b2, dCriticB2, adamStates.cb2, lr)
|
||||
adamStep(ac.critic.w3, dCriticW3, adamStates.cw3, lr)
|
||||
adamStep(ac.critic.b3, dCriticB3, adamStates.cb3, lr)
|
||||
mbStart = mbEnd
|
||||
|
||||
let totalSamples = (epochs * bufLen).float32
|
||||
result.actorLoss = totalActorLoss / totalSamples
|
||||
result.valueLoss = totalValueLoss / totalSamples
|
||||
result.gradNorm = if totalMiniBatches > 0: totalGradNorm / totalMiniBatches.float32
|
||||
else: 0.0'f32
|
||||
@@ -0,0 +1,211 @@
|
||||
## weights.nim — save/load ActorCritic weights as .npy files.
|
||||
|
||||
import std/[os, times, strutils, algorithm, sequtils]
|
||||
import arraymancer
|
||||
import ./network
|
||||
import ./training
|
||||
|
||||
# ── Tensor names — order must match save/load ─────────────────────────────────
|
||||
|
||||
const weightFiles = [
|
||||
"actor_w1.npy", "actor_b1.npy", "actor_w2.npy", "actor_b2.npy",
|
||||
"actor_w3.npy", "actor_b3.npy",
|
||||
"critic_w1.npy", "critic_b1.npy", "critic_w2.npy", "critic_b2.npy",
|
||||
"critic_w3.npy", "critic_b3.npy",
|
||||
"log_std.npy",
|
||||
]
|
||||
|
||||
proc saveWeights*(ac: ActorCritic, dir: string) =
|
||||
## Write all weight tensors to dir/ as .npy files.
|
||||
createDir(dir)
|
||||
ac.actor.w1.write_npy(dir / "actor_w1.npy")
|
||||
ac.actor.b1.write_npy(dir / "actor_b1.npy")
|
||||
ac.actor.w2.write_npy(dir / "actor_w2.npy")
|
||||
ac.actor.b2.write_npy(dir / "actor_b2.npy")
|
||||
ac.actor.w3.write_npy(dir / "actor_w3.npy")
|
||||
ac.actor.b3.write_npy(dir / "actor_b3.npy")
|
||||
ac.critic.w1.write_npy(dir / "critic_w1.npy")
|
||||
ac.critic.b1.write_npy(dir / "critic_b1.npy")
|
||||
ac.critic.w2.write_npy(dir / "critic_w2.npy")
|
||||
ac.critic.b2.write_npy(dir / "critic_b2.npy")
|
||||
ac.critic.w3.write_npy(dir / "critic_w3.npy")
|
||||
ac.critic.b3.write_npy(dir / "critic_b3.npy")
|
||||
ac.logStd.write_npy(dir / "log_std.npy")
|
||||
|
||||
proc loadWeights*(ac: var ActorCritic, dir: string) =
|
||||
## Load all weight tensors from dir/.
|
||||
## If a tensor's shape doesn't match (e.g. STATE_DIM changed), keep the
|
||||
## freshly-initialised value and print a warning — other tensors still load.
|
||||
template loadOrSkip(dest: untyped, path: string) =
|
||||
let loaded = read_npy[float32](path)
|
||||
if loaded.shape == dest.shape:
|
||||
dest = loaded
|
||||
else:
|
||||
echo "weights: shape mismatch for " & path &
|
||||
" (got " & $loaded.shape & " want " & $dest.shape & ") — keeping fresh init"
|
||||
|
||||
loadOrSkip(ac.actor.w1, dir / "actor_w1.npy")
|
||||
loadOrSkip(ac.actor.b1, dir / "actor_b1.npy")
|
||||
loadOrSkip(ac.actor.w2, dir / "actor_w2.npy")
|
||||
loadOrSkip(ac.actor.b2, dir / "actor_b2.npy")
|
||||
loadOrSkip(ac.actor.w3, dir / "actor_w3.npy")
|
||||
loadOrSkip(ac.actor.b3, dir / "actor_b3.npy")
|
||||
loadOrSkip(ac.critic.w1, dir / "critic_w1.npy")
|
||||
loadOrSkip(ac.critic.b1, dir / "critic_b1.npy")
|
||||
loadOrSkip(ac.critic.w2, dir / "critic_w2.npy")
|
||||
loadOrSkip(ac.critic.b2, dir / "critic_b2.npy")
|
||||
loadOrSkip(ac.critic.w3, dir / "critic_w3.npy")
|
||||
loadOrSkip(ac.critic.b3, dir / "critic_b3.npy")
|
||||
loadOrSkip(ac.logStd, dir / "log_std.npy")
|
||||
|
||||
proc saveWeightsAtomic*(ac: ActorCritic, targetDir: string) =
|
||||
## Write to a temp dir, then rename atomically over targetDir.
|
||||
let tmpDir = targetDir & "_tmp_" & $int(epochTime())
|
||||
saveWeights(ac, tmpDir)
|
||||
if dirExists(targetDir):
|
||||
removeDir(targetDir)
|
||||
moveDir(tmpDir, targetDir)
|
||||
|
||||
# ── Adam state file names ──────────────────────────────────────────────────────
|
||||
|
||||
const adamMFiles = [
|
||||
"adam_aw1_m.npy", "adam_ab1_m.npy", "adam_aw2_m.npy", "adam_ab2_m.npy",
|
||||
"adam_aw3_m.npy", "adam_ab3_m.npy",
|
||||
"adam_cw1_m.npy", "adam_cb1_m.npy", "adam_cw2_m.npy", "adam_cb2_m.npy",
|
||||
"adam_cw3_m.npy", "adam_cb3_m.npy",
|
||||
"adam_logstd_m.npy",
|
||||
]
|
||||
const adamVFiles = [
|
||||
"adam_aw1_v.npy", "adam_ab1_v.npy", "adam_aw2_v.npy", "adam_ab2_v.npy",
|
||||
"adam_aw3_v.npy", "adam_ab3_v.npy",
|
||||
"adam_cw1_v.npy", "adam_cb1_v.npy", "adam_cw2_v.npy", "adam_cb2_v.npy",
|
||||
"adam_cw3_v.npy", "adam_cb3_v.npy",
|
||||
"adam_logstd_v.npy",
|
||||
]
|
||||
|
||||
proc saveAdamStates*(adam: ACAdamStates, dir: string) =
|
||||
## Write Adam m/v tensors and t counters to dir/.
|
||||
adam.aw1.m.write_npy(dir / "adam_aw1_m.npy"); adam.aw1.v.write_npy(dir / "adam_aw1_v.npy")
|
||||
adam.ab1.m.write_npy(dir / "adam_ab1_m.npy"); adam.ab1.v.write_npy(dir / "adam_ab1_v.npy")
|
||||
adam.aw2.m.write_npy(dir / "adam_aw2_m.npy"); adam.aw2.v.write_npy(dir / "adam_aw2_v.npy")
|
||||
adam.ab2.m.write_npy(dir / "adam_ab2_m.npy"); adam.ab2.v.write_npy(dir / "adam_ab2_v.npy")
|
||||
adam.aw3.m.write_npy(dir / "adam_aw3_m.npy"); adam.aw3.v.write_npy(dir / "adam_aw3_v.npy")
|
||||
adam.ab3.m.write_npy(dir / "adam_ab3_m.npy"); adam.ab3.v.write_npy(dir / "adam_ab3_v.npy")
|
||||
adam.cw1.m.write_npy(dir / "adam_cw1_m.npy"); adam.cw1.v.write_npy(dir / "adam_cw1_v.npy")
|
||||
adam.cb1.m.write_npy(dir / "adam_cb1_m.npy"); adam.cb1.v.write_npy(dir / "adam_cb1_v.npy")
|
||||
adam.cw2.m.write_npy(dir / "adam_cw2_m.npy"); adam.cw2.v.write_npy(dir / "adam_cw2_v.npy")
|
||||
adam.cb2.m.write_npy(dir / "adam_cb2_m.npy"); adam.cb2.v.write_npy(dir / "adam_cb2_v.npy")
|
||||
adam.cw3.m.write_npy(dir / "adam_cw3_m.npy"); adam.cw3.v.write_npy(dir / "adam_cw3_v.npy")
|
||||
adam.cb3.m.write_npy(dir / "adam_cb3_m.npy"); adam.cb3.v.write_npy(dir / "adam_cb3_v.npy")
|
||||
adam.logStd.m.write_npy(dir / "adam_logstd_m.npy")
|
||||
adam.logStd.v.write_npy(dir / "adam_logstd_v.npy")
|
||||
# t counters (all stepped in lockstep; store each for safety)
|
||||
writeFile(dir / "adam_t.txt",
|
||||
[$adam.aw1.t, $adam.ab1.t, $adam.aw2.t, $adam.ab2.t,
|
||||
$adam.aw3.t, $adam.ab3.t, $adam.cw1.t, $adam.cb1.t,
|
||||
$adam.cw2.t, $adam.cb2.t, $adam.cw3.t, $adam.cb3.t,
|
||||
$adam.logStd.t].join("\n"))
|
||||
|
||||
proc loadAdamStates*(adam: var ACAdamStates, dir: string) =
|
||||
## Load Adam m/v tensors and t counters from dir/. Called only when files exist.
|
||||
## Shape mismatch (e.g. STATE_DIM changed) → keep zero-initialised state (safe fresh start).
|
||||
template lm(dest: untyped, path: string) =
|
||||
let loaded = read_npy[float32](path)
|
||||
if loaded.shape == dest.shape:
|
||||
dest = loaded
|
||||
else:
|
||||
echo "weights: Adam shape mismatch for " & path &
|
||||
" (got " & $loaded.shape & " want " & $dest.shape & ") — resetting Adam state"
|
||||
lm(adam.aw1.m, dir / "adam_aw1_m.npy"); lm(adam.aw1.v, dir / "adam_aw1_v.npy")
|
||||
lm(adam.ab1.m, dir / "adam_ab1_m.npy"); lm(adam.ab1.v, dir / "adam_ab1_v.npy")
|
||||
lm(adam.aw2.m, dir / "adam_aw2_m.npy"); lm(adam.aw2.v, dir / "adam_aw2_v.npy")
|
||||
lm(adam.ab2.m, dir / "adam_ab2_m.npy"); lm(adam.ab2.v, dir / "adam_ab2_v.npy")
|
||||
lm(adam.aw3.m, dir / "adam_aw3_m.npy"); lm(adam.aw3.v, dir / "adam_aw3_v.npy")
|
||||
lm(adam.ab3.m, dir / "adam_ab3_m.npy"); lm(adam.ab3.v, dir / "adam_ab3_v.npy")
|
||||
lm(adam.cw1.m, dir / "adam_cw1_m.npy"); lm(adam.cw1.v, dir / "adam_cw1_v.npy")
|
||||
lm(adam.cb1.m, dir / "adam_cb1_m.npy"); lm(adam.cb1.v, dir / "adam_cb1_v.npy")
|
||||
lm(adam.cw2.m, dir / "adam_cw2_m.npy"); lm(adam.cw2.v, dir / "adam_cw2_v.npy")
|
||||
lm(adam.cb2.m, dir / "adam_cb2_m.npy"); lm(adam.cb2.v, dir / "adam_cb2_v.npy")
|
||||
lm(adam.cw3.m, dir / "adam_cw3_m.npy"); lm(adam.cw3.v, dir / "adam_cw3_v.npy")
|
||||
lm(adam.cb3.m, dir / "adam_cb3_m.npy"); lm(adam.cb3.v, dir / "adam_cb3_v.npy")
|
||||
lm(adam.logStd.m, dir / "adam_logstd_m.npy"); lm(adam.logStd.v, dir / "adam_logstd_v.npy")
|
||||
let ts = readFile(dir / "adam_t.txt").strip().splitLines()
|
||||
if ts.len >= 13:
|
||||
adam.aw1.t = parseInt(ts[0]); adam.ab1.t = parseInt(ts[1])
|
||||
adam.aw2.t = parseInt(ts[2]); adam.ab2.t = parseInt(ts[3])
|
||||
adam.aw3.t = parseInt(ts[4]); adam.ab3.t = parseInt(ts[5])
|
||||
adam.cw1.t = parseInt(ts[6]); adam.cb1.t = parseInt(ts[7])
|
||||
adam.cw2.t = parseInt(ts[8]); adam.cb2.t = parseInt(ts[9])
|
||||
adam.cw3.t = parseInt(ts[10]); adam.cb3.t = parseInt(ts[11])
|
||||
adam.logStd.t = parseInt(ts[12])
|
||||
adam.initialized = true
|
||||
|
||||
proc adamStateFilesExist(dir: string): bool =
|
||||
## Check that the minimum set of Adam files is present.
|
||||
for f in adamMFiles:
|
||||
if not fileExists(dir / f): return false
|
||||
for f in adamVFiles:
|
||||
if not fileExists(dir / f): return false
|
||||
fileExists(dir / "adam_t.txt")
|
||||
|
||||
proc saveCheckpoint*(ac: ActorCritic, adam: ACAdamStates,
|
||||
weightsRoot: string, roundNum: int) =
|
||||
## Always saves weights + Adam state to weightsRoot/latest/.
|
||||
## Every 50 rounds also saves to checkpoint_{1,2,3} in round-robin.
|
||||
## Round counter is saved to weightsRoot/round_counter.txt (outside checkpoint dirs).
|
||||
let latestDir = weightsRoot / "latest"
|
||||
saveWeightsAtomic(ac, latestDir)
|
||||
if adam.initialized:
|
||||
saveAdamStates(adam, latestDir)
|
||||
writeFile(weightsRoot / "round_counter.txt", $roundNum)
|
||||
if roundNum mod 50 == 0:
|
||||
let slot = ((roundNum div 50 - 1) mod 3) + 1 # 50→1, 100→2, 150→3, 200→1, …
|
||||
let ckDir = weightsRoot / ("checkpoint_" & $slot)
|
||||
saveWeightsAtomic(ac, ckDir)
|
||||
if adam.initialized:
|
||||
saveAdamStates(adam, ckDir)
|
||||
|
||||
proc loadBestAvailable*(ac: var ActorCritic, adam: var ACAdamStates,
|
||||
weightsRoot: string): tuple[loaded: bool, roundNum: int] =
|
||||
## Try latest/ first, then checkpoints sorted newest-first by mtime.
|
||||
## Returns (true, roundNum) if weights loaded, (false, 0) if all fail.
|
||||
## Adam state is loaded if present alongside weights; otherwise left uninitialised.
|
||||
## Round counter is read from weightsRoot/round_counter.txt if present.
|
||||
let checkpoints = [weightsRoot / "checkpoint_1",
|
||||
weightsRoot / "checkpoint_2",
|
||||
weightsRoot / "checkpoint_3"]
|
||||
# Sort checkpoints newest-first by modification time
|
||||
var existing: seq[tuple[mtime: Time, path: string]]
|
||||
for p in checkpoints:
|
||||
if dirExists(p):
|
||||
existing.add((getLastModificationTime(p), p))
|
||||
existing.sort(proc(a, b: tuple[mtime: Time, path: string]): int =
|
||||
cmp(b.mtime, a.mtime)) # descending
|
||||
|
||||
let candidates = @[weightsRoot / "latest"] & existing.mapIt(it.path)
|
||||
for candidate in candidates:
|
||||
if dirExists(candidate):
|
||||
var ok = true
|
||||
for f in weightFiles:
|
||||
if not fileExists(candidate / f):
|
||||
ok = false
|
||||
break
|
||||
if ok:
|
||||
ac.loadWeights(candidate)
|
||||
if adamStateFilesExist(candidate):
|
||||
adam.loadAdamStates(candidate)
|
||||
let rcPath = weightsRoot / "round_counter.txt"
|
||||
# Torn/empty file (e.g. after a crash) must not abort startup → treat as 0
|
||||
let roundNum = if fileExists(rcPath):
|
||||
try: parseInt(readFile(rcPath).strip())
|
||||
except ValueError: 0
|
||||
else: 0
|
||||
return (loaded: true, roundNum: roundNum)
|
||||
result = (loaded: false, roundNum: 0)
|
||||
|
||||
proc cleanStaleTempDirs*(weightsRoot: string) =
|
||||
## Delete any dirs inside weightsRoot whose name contains "_tmp_".
|
||||
if not dirExists(weightsRoot): return
|
||||
for kind, path in walkDir(weightsRoot):
|
||||
if kind == pcDir and "_tmp_" in lastPathPart(path):
|
||||
removeDir(path)
|
||||
Reference in New Issue
Block a user