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:
2026-08-17 19:08:31 +02:00
parent 56e0b306c9
commit cdde60d79f
39 changed files with 490 additions and 78 deletions
+30 -8
View File
@@ -10,8 +10,8 @@ import ./network
type
Transition* = object
state*: Tensor[float32] # [42]
action*: Tensor[float32] # [5]
state*: Tensor[float32] # [STATE_DIM]
action*: Tensor[float32] # [ACTION_DIM]
logProb*: float32
reward*: float32
value*: float32 # critic estimate at collection time
@@ -151,6 +151,13 @@ proc initACAdamStates*(ac: ActorCritic): ACAdamStates =
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 =
@@ -171,9 +178,14 @@ proc ppoUpdate*(ac: var ActorCritic;
entropyCoeff: float32 = 0.01'f32;
valueLossCoeff: float32 = 0.5'f32;
lr: float32 = 3e-4'f32;
maxGradNorm: float32 = 0.5'f32) {.gcsafe.} =
maxGradNorm: float32 = 0.5'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
if not adamStates.initialized:
adamStates = initACAdamStates(ac)
@@ -230,14 +242,14 @@ proc ppoUpdate*(ac: var ActorCritic;
# ── Actor forward ──
let actorFwd = mlpForwardCached(ac.actor, tr.state)
let newMean = actorFwd.y # [5]
let newMean = actorFwd.y # [ACTION_DIM]
let logStdClamped = ac.logStd.map(proc(v: float32): float32 = max(v, -3.0'f32))
let std = logStdClamped.map(proc(v: float32): float32 = exp(v))
# New log prob
var newLogP = 0.0'f32
for i in 0..<5:
for i in 0..<ACTION_DIM:
let mu = newMean[i]
let s = std[i]
let diff = (tr.action[i] - mu) / s
@@ -250,6 +262,7 @@ proc ppoUpdate*(ac: var ActorCritic;
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
@@ -260,8 +273,8 @@ proc ppoUpdate*(ac: var ActorCritic;
let dLoss_dNewLogP = dLoss_dRatio * ratio
# d(newLogP)/d(mean[i]) = (action[i] - mean[i]) / std[i]^2
var dLogP_dMean = newTensor[float32](5)
for i in 0..<5:
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)
@@ -271,7 +284,7 @@ proc ppoUpdate*(ac: var ActorCritic;
# 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..<5:
for i in 0..<ACTION_DIM:
let isClamped = (ac.logStd[i] <= -3.0'f32)
if not isClamped:
let s = std[i]
@@ -295,6 +308,7 @@ proc ppoUpdate*(ac: var ActorCritic;
let criticFwd = mlpForwardCached(ac.critic, tr.state)
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)
@@ -314,6 +328,8 @@ proc ppoUpdate*(ac: var ActorCritic;
dCriticW1, dCriticB1, dCriticW2, dCriticB2, dCriticW3, dCriticB3
]
let norm = globalNorm(allGrads)
totalGradNorm += norm
inc totalMiniBatches
if norm > maxGradNorm:
let scale = maxGradNorm / norm
for g in allGrads.mitems: g = g *. scale
@@ -343,3 +359,9 @@ proc ppoUpdate*(ac: var ActorCritic;
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