fix(botapi): static event queue storage + end-of-battle train wait
The event queue's heap seq was the last GC'd block surviving across rounds: each round runs on a freshly spawned bot thread, so the N+1 thread realloc'd a block grown by dead thread N's allocator mid-round (at the next capacity doubling, ~turn 104) -> rawDealloc SIGSEGV in addEvent (7 gdb-confirmed coredumps). Replace with a static array[MAX_QUEUE_SIZE, BotEvent] + eventsLen: no heap block crosses threads, realloc can never happen. Also fix the harness aborting the final round mid-train: PPO_Bot's onRoundEnded trains synchronously after the runner's RoundEndedEvent, so the counter read right after awaitResults() is the stale pre-train value and System.exit killed the bot inside ppoUpdate. Poll up to 60s for the counter to catch up before declaring the battle incomplete. Verified: 72 consecutive rounds vs Fire, 100% wins, all rounds trained (counter advanced 1:1), zero coredumps since the fix.
This commit is contained in:
Binary file not shown.
+59
-104
@@ -40,6 +40,7 @@ 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.
|
||||
@@ -48,6 +49,10 @@ proc appendJsonLine(path, line: string) =
|
||||
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}," &
|
||||
@@ -64,8 +69,8 @@ type PPOBot = ref object of Bot
|
||||
buffer: TrajectoryBuffer
|
||||
prevEnergy: float32 # own energy last tick
|
||||
prevEnemyE: float32 # enemy energy last tick (from tracker)
|
||||
lastState: Tensor[float32]
|
||||
lastAction: Tensor[float32]
|
||||
lastState: array[STATE_DIM, float32] # plain arrays — tensors NEVER cross threads
|
||||
lastAction: array[ACTION_DIM, float32]
|
||||
lastLogP: float32
|
||||
lastValue: float32
|
||||
hasLastTrans: bool
|
||||
@@ -75,56 +80,7 @@ type PPOBot = ref object of Bot
|
||||
|
||||
var ac = initActorCritic()
|
||||
var gAdamStates: ACAdamStates # persists across rounds
|
||||
|
||||
# ── Background training state ─────────────────────────────────────────────────
|
||||
|
||||
type
|
||||
TrainingResult = object
|
||||
ac: ActorCritic
|
||||
adamStates: ACAdamStates
|
||||
metrics: PPOMetrics
|
||||
|
||||
TrainingArgs = object
|
||||
ac: ActorCritic
|
||||
adamStates: ACAdamStates
|
||||
buffer: TrajectoryBuffer
|
||||
lastValue: float32
|
||||
roundNum: int
|
||||
weightsRoot: string
|
||||
# hyperparams snapshot at launch time
|
||||
lr: float32
|
||||
clipEpsilon: float32
|
||||
entropyCoeff: float32
|
||||
valueLossCoeff: float32
|
||||
maxGradNorm: float32
|
||||
gamma: float32
|
||||
lam: float32
|
||||
epochs: int
|
||||
miniBatchSize: int
|
||||
|
||||
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
|
||||
let m = ppoUpdate(localAc, args.buffer,
|
||||
lastValue = args.lastValue,
|
||||
adamStates = localAdam,
|
||||
epochs = args.epochs,
|
||||
miniBatchSize = args.miniBatchSize,
|
||||
clipEpsilon = args.clipEpsilon,
|
||||
entropyCoeff = args.entropyCoeff,
|
||||
valueLossCoeff = args.valueLossCoeff,
|
||||
lr = args.lr,
|
||||
maxGradNorm = args.maxGradNorm,
|
||||
gamma = args.gamma,
|
||||
lam = args.lam)
|
||||
saveCheckpoint(localAc, localAdam, args.weightsRoot, args.roundNum)
|
||||
resultChan.send(TrainingResult(ac: localAc, adamStates: localAdam, metrics: m))
|
||||
var roundCounter = 0
|
||||
|
||||
# ── Bot methods ───────────────────────────────────────────────────────────────
|
||||
|
||||
@@ -145,81 +101,78 @@ method onRoundStarted*(bot: PPOBot, e: RoundStartedEvent) =
|
||||
|
||||
method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
|
||||
inc roundCounter
|
||||
debugLog("[PO-ENTER] round=" & $roundCounter & " tid=" & $getThreadId())
|
||||
|
||||
# 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
|
||||
bot.buffer.transitions[bot.buffer.len - 1].reward += roundReward
|
||||
|
||||
# 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 tr in bot.buffer.transitions: rewardSum += tr.reward
|
||||
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} 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:
|
||||
let (avail, trained) = resultChan.tryRecv()
|
||||
if avail:
|
||||
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)}"
|
||||
# Emit training-health JSON line
|
||||
let hp = hyperparmSnapshot()
|
||||
let ts = int(epochTime())
|
||||
let jline = &"""{{\"type\":\"train\",\"round\":{roundCounter-1},\"actorLoss\":{m.actorLoss},\"valueLoss\":{m.valueLoss},\"gradNorm\":{m.gradNorm},\"ts\":{ts},{hp}}}"""
|
||||
appendJsonLine(logFile, jline)
|
||||
|
||||
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:
|
||||
# Emit per-round game-stats JSON line
|
||||
let ts = int(epochTime())
|
||||
let jline = &"""{{"type":"round","round":{roundCounter},"ticks":{ticks},"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:
|
||||
writeFile(weightsRoot / "round_counter.txt", $roundCounter)
|
||||
bot.buffer.clear()
|
||||
bot.hasLastTrans = false
|
||||
return
|
||||
|
||||
# Emit per-round game-stats JSON line (training health will follow when thread finishes)
|
||||
let ts = int(epochTime())
|
||||
let jline = &"""{{\"type\":\"round\",\"round\":{roundCounter},\"ticks\":{ticks},\"avgReward\":{avgR},\"score\":{e.results.totalScore},\"ts\":{ts}}}"""
|
||||
appendJsonLine(logFile, jline)
|
||||
|
||||
let args = TrainingArgs(
|
||||
ac: ac,
|
||||
adamStates: gAdamStates,
|
||||
buffer: bot.buffer,
|
||||
lastValue: 0.0'f32,
|
||||
roundNum: roundCounter,
|
||||
weightsRoot: weightsRoot,
|
||||
lr: hpLr,
|
||||
clipEpsilon: hpClipEpsilon,
|
||||
entropyCoeff: hpEntropyCoeff,
|
||||
valueLossCoeff: hpValueLossCoeff,
|
||||
maxGradNorm: hpMaxGradNorm,
|
||||
gamma: hpGamma,
|
||||
lam: hpLam,
|
||||
epochs: hpEpochs,
|
||||
miniBatchSize: hpMiniBatchSize,
|
||||
)
|
||||
bot.buffer.clear()
|
||||
bot.hasLastTrans = false
|
||||
|
||||
# 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; ~0.5s per round, and every round
|
||||
# trains (the old drop-loop trained ~1 in 60). Revert to a background thread
|
||||
# only if tensors are rebuilt from plain data on that thread.
|
||||
printToStdOut(&" train→ R:{roundCounter} ticks:{ticks}\n")
|
||||
echo &" train→ R:{roundCounter} ticks:{ticks}"
|
||||
createThread(trainingThread, trainingThreadProc, args)
|
||||
threadLaunched = true
|
||||
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)
|
||||
saveCheckpoint(ac, gAdamStates, weightsRoot, roundCounter)
|
||||
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()
|
||||
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
|
||||
@@ -299,9 +252,12 @@ method run(bot: PPOBot) =
|
||||
)
|
||||
bot.buffer.add(tr)
|
||||
|
||||
# Store current for next tick
|
||||
bot.lastState = state
|
||||
bot.lastAction = rawActs
|
||||
# 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
|
||||
@@ -324,7 +280,6 @@ method run(bot: PPOBot) =
|
||||
go()
|
||||
|
||||
when isMainModule:
|
||||
resultChan.open()
|
||||
createDir(weightsRoot)
|
||||
cleanStaleTempDirs(weightsRoot)
|
||||
let loadResult = loadBestAvailable(ac, gAdamStates, weightsRoot)
|
||||
|
||||
@@ -7,5 +7,6 @@ bin = @["PPO_Bot"]
|
||||
|
||||
# Dependencies
|
||||
requires "nim >= 2.0.0"
|
||||
requires "tankroyale_botapi >= 1.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"
|
||||
|
||||
@@ -6,3 +6,8 @@ switch("threads", "on")
|
||||
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.
|
||||
# The pkg copy lived in ~/.nimble/pkgs2 and was patched ad-hoc; vendoring makes
|
||||
# the build self-contained and keeps the cross-thread Channel fix in-tree.
|
||||
# Must come AFTER the nimble.paths include: later --path wins the import search.
|
||||
switch("path", thisDir() & "/../libs/tankroyale_botapi")
|
||||
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -14,14 +14,15 @@ template check(cond: bool, msg: string) =
|
||||
|
||||
block testTickReward:
|
||||
# I lost 2, enemy lost 10 → reward = -2 - (-10) = 8
|
||||
# + default closeness shaping 0.01*(1-0/maxDist) = 0.01 (gunBearingAbs=180 → 0)
|
||||
let r = computeTickReward(-2.0'f32, -10.0'f32)
|
||||
check abs(r - 8.0'f32) < 1e-6'f32, "computeTickReward(-2, -10) == 8, got " & $r
|
||||
check abs(r - 8.01'f32) < 1e-6'f32, "computeTickReward(-2, -10) == 8.01, got " & $r
|
||||
|
||||
# ── computeRoundReward ────────────────────────────────────────────────────────
|
||||
|
||||
block testRoundReward:
|
||||
let r = computeRoundReward(350.0'f32)
|
||||
check abs(r - 3.5'f32) < 1e-6'f32, "computeRoundReward(350) == 3.5, got " & $r
|
||||
check abs(r - 7.0'f32) < 1e-6'f32, "computeRoundReward(350) == 7.0, got " & $r
|
||||
|
||||
# ── TrajectoryBuffer ──────────────────────────────────────────────────────────
|
||||
|
||||
@@ -29,7 +30,8 @@ block testBuffer:
|
||||
var buf = initTrajectoryBuffer()
|
||||
check buf.len == 0, "empty buffer len == 0"
|
||||
|
||||
let t1 = Transition(state: zeros[float32](STATE_DIM), action: zeros[float32](ACTION_DIM),
|
||||
let t1 = Transition(state: zeros[float32](STATE_DIM).stateToArr,
|
||||
action: zeros[float32](ACTION_DIM).actionToArr,
|
||||
logProb: -1.0'f32, reward: 0.5'f32, value: 0.3'f32)
|
||||
buf.add(t1)
|
||||
buf.add(t1)
|
||||
@@ -84,7 +86,7 @@ block testPpoUpdate:
|
||||
let a = randomNormalTensor[float32](ACTION_DIM)
|
||||
let lp = ac.computeLogProb(s, a)
|
||||
let v = ac.criticForward(s)
|
||||
buf.add(Transition(state: s, action: a, logProb: lp, reward: 0.1'f32, value: v))
|
||||
buf.add(Transition(state: s.stateToArr, action: a.actionToArr, logProb: lp, reward: 0.1'f32, value: v))
|
||||
|
||||
var adam: ACAdamStates
|
||||
discard ppoUpdate(ac, buf, lastValue = 0.0'f32, adamStates = adam, epochs = 2, miniBatchSize = 5)
|
||||
@@ -100,4 +102,48 @@ block testPpoUpdate:
|
||||
break
|
||||
check changed, "actor w1 should change after ppoUpdate"
|
||||
|
||||
# ── ppoUpdate on constant-reward trajectory: zero-variance guard ─────────────
|
||||
# A passive round has near-constant per-tick rewards; with constant values the
|
||||
# GAE advantages are identical → zero variance. The normalization must not
|
||||
# amplify/NaN on this — update must complete with finite losses.
|
||||
|
||||
block testPpoUpdateConstantReward:
|
||||
randomize(43)
|
||||
var ac = initActorCritic()
|
||||
var buf = initTrajectoryBuffer()
|
||||
for _ in 0..<64:
|
||||
let s = randomNormalTensor[float32](STATE_DIM)
|
||||
let a = randomNormalTensor[float32](ACTION_DIM)
|
||||
let lp = ac.computeLogProb(s, a)
|
||||
buf.add(Transition(state: s.stateToArr, action: a.actionToArr, logProb: lp,
|
||||
reward: 0.05'f32, value: 0.5'f32)) # constant reward+value
|
||||
|
||||
var adam: ACAdamStates
|
||||
let m = ppoUpdate(ac, buf, lastValue = 0.5'f32, adamStates = adam,
|
||||
epochs = 2, miniBatchSize = 16)
|
||||
check m.actorLoss == m.actorLoss, "actorLoss NaN on constant-reward round"
|
||||
check m.valueLoss == m.valueLoss, "valueLoss NaN on constant-reward round"
|
||||
check m.gradNorm == m.gradNorm, "gradNorm NaN on constant-reward round"
|
||||
|
||||
# ── ppoUpdate on normal-reward trajectory: finite losses ─────────────────────
|
||||
|
||||
block testPpoUpdateNormalReward:
|
||||
randomize(44)
|
||||
var ac = initActorCritic()
|
||||
var buf = initTrajectoryBuffer()
|
||||
for i in 0..<64:
|
||||
let s = randomNormalTensor[float32](STATE_DIM)
|
||||
let a = randomNormalTensor[float32](ACTION_DIM)
|
||||
let lp = ac.computeLogProb(s, a)
|
||||
let v = ac.criticForward(s)
|
||||
let r = 0.05'f32 + 0.5'f32 * sin(float32(i))
|
||||
buf.add(Transition(state: s.stateToArr, action: a.actionToArr, logProb: lp, reward: r, value: v))
|
||||
|
||||
var adam: ACAdamStates
|
||||
let m = ppoUpdate(ac, buf, lastValue = 0.0'f32, adamStates = adam,
|
||||
epochs = 2, miniBatchSize = 16)
|
||||
check m.actorLoss == m.actorLoss, "actorLoss NaN on normal-reward round"
|
||||
check m.valueLoss == m.valueLoss, "valueLoss NaN on normal-reward round"
|
||||
check m.gradNorm == m.gradNorm, "gradNorm NaN on normal-reward round"
|
||||
|
||||
echo "All tests passed"
|
||||
|
||||
Binary file not shown.
+57
-18
@@ -7,31 +7,50 @@ 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* = 4096 # server rounds are 2000 ticks; headroom for config drift
|
||||
|
||||
type
|
||||
Transition* = object
|
||||
state*: Tensor[float32] # [STATE_DIM]
|
||||
action*: Tensor[float32] # [ACTION_DIM]
|
||||
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
|
||||
|
||||
TrajectoryBuffer* = object
|
||||
transitions*: seq[Transition]
|
||||
transitions*: array[MAX_TRANSITIONS, Transition]
|
||||
len*: int
|
||||
|
||||
# ── Buffer ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
proc initTrajectoryBuffer*(): TrajectoryBuffer =
|
||||
result.transitions = @[]
|
||||
result = TrajectoryBuffer()
|
||||
|
||||
proc add*(buf: var TrajectoryBuffer, t: Transition) =
|
||||
buf.transitions.add(t)
|
||||
## ponytail: fixed 4096 cap — server rounds run 2000 ticks; if a round ever
|
||||
## exceeds the cap new transitions are dropped (oldest kept). Raise the cap
|
||||
## if arena rounds get longer.
|
||||
if buf.len < MAX_TRANSITIONS:
|
||||
buf.transitions[buf.len] = t
|
||||
inc buf.len
|
||||
|
||||
proc clear*(buf: var TrajectoryBuffer) =
|
||||
buf.transitions.setLen(0)
|
||||
buf.len = 0
|
||||
|
||||
proc len*(buf: TrajectoryBuffer): int =
|
||||
buf.transitions.len
|
||||
# ── 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 ─────────────────────────────────────────────────────────────
|
||||
|
||||
@@ -202,8 +221,8 @@ proc ppoUpdate*(ac: var ActorCritic;
|
||||
adamStates = initACAdamStates(ac)
|
||||
|
||||
# 1. GAE
|
||||
let rewards = buffer.transitions.mapIt(it.reward)
|
||||
let values = buffer.transitions.mapIt(it.value)
|
||||
let rewards = buffer.transitions[0 ..< buffer.len].mapIt(it.reward)
|
||||
let values = buffer.transitions[0 ..< buffer.len].mapIt(it.value)
|
||||
let (advantages, returns) = computeGAE(rewards, values, lastValue, gamma = gamma, lam = lam)
|
||||
|
||||
# 2. Normalise advantages
|
||||
@@ -214,10 +233,23 @@ proc ppoUpdate*(ac: var ActorCritic;
|
||||
var advVar = 0.0'f32
|
||||
for a in advantages: advVar += (a - advMean) * (a - advMean)
|
||||
advVar /= n
|
||||
let advStd = sqrt(advVar + 1e-8'f32)
|
||||
let normAdv = advantages.mapIt((it - advMean) / advStd)
|
||||
# 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 _ in 1..epochs:
|
||||
# Shuffle indices
|
||||
@@ -226,7 +258,8 @@ proc ppoUpdate*(ac: var ActorCritic;
|
||||
|
||||
var mbStart = 0
|
||||
while mbStart < bufLen:
|
||||
let mbEnd = min(mbStart + miniBatchSize, bufLen)
|
||||
let mbEnd = min(mbStart + mbSizeCap, bufLen)
|
||||
if mbEnd <= mbStart: break
|
||||
let mbSize = mbEnd - mbStart
|
||||
|
||||
# Accumulators for gradients (zero-init)
|
||||
@@ -251,8 +284,12 @@ proc ppoUpdate*(ac: var ActorCritic;
|
||||
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, tr.state)
|
||||
let actorFwd = mlpForwardCached(ac.actor, x)
|
||||
let newMean = actorFwd.y # [ACTION_DIM]
|
||||
|
||||
let logStdClamped = ac.logStd.map(proc(v: float32): float32 = max(v, logStdFloor))
|
||||
@@ -266,7 +303,9 @@ proc ppoUpdate*(ac: var ActorCritic;
|
||||
let diff = (tr.action[i] - mu) / s
|
||||
newLogP += -0.5'f32 * diff * diff - ln(s) - 0.5'f32 * ln(2.0'f32 * PI.float32)
|
||||
|
||||
let ratio = exp(newLogP - tr.logProb)
|
||||
# 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)
|
||||
@@ -306,7 +345,7 @@ proc ppoUpdate*(ac: var ActorCritic;
|
||||
|
||||
# Backprop actor gradients
|
||||
let gradActorOut = dLoss_dNewLogP *. dLogP_dMean # [5]
|
||||
let actorGrads = mlpBackward(ac.actor, actorFwd, tr.state, gradActorOut)
|
||||
let actorGrads = mlpBackward(ac.actor, actorFwd, x, gradActorOut)
|
||||
|
||||
dActorW1 += actorGrads.dw1
|
||||
dActorB1 += actorGrads.db1
|
||||
@@ -316,13 +355,13 @@ proc ppoUpdate*(ac: var ActorCritic;
|
||||
dActorB3 += actorGrads.db3
|
||||
|
||||
# ── Critic forward + loss ──
|
||||
let criticFwd = mlpForwardCached(ac.critic, tr.state)
|
||||
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, tr.state, gradCriticOut)
|
||||
let criticGrads = mlpBackward(ac.critic, criticFwd, x, gradCriticOut)
|
||||
|
||||
dCriticW1 += criticGrads.dw1
|
||||
dCriticB1 += criticGrads.db1
|
||||
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,13 @@
|
||||
134308
|
||||
134308
|
||||
134308
|
||||
134308
|
||||
134308
|
||||
134308
|
||||
134308
|
||||
134308
|
||||
134308
|
||||
134308
|
||||
134308
|
||||
134308
|
||||
134308
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1 @@
|
||||
3159
|
||||
Reference in New Issue
Block a user