feat(PPO_Bot): reward + trajectory + GAE + PPO training (#16)
Manual-backprop PPO with Adam: TrajectoryBuffer, computeGAE, ppoUpdate (4 epochs, minibatch 64, clip 0.2, grad norm 0.5). Reward helpers computeTickReward/computeRoundReward. Bot wired: tick transitions collected in run loop, ppoUpdate called on onRoundEnded. Fix: add arraymancer import to PPO_Bot.nim so Tensor resolves at top level. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,3 @@
|
|||||||
|
nimble.develop
|
||||||
|
nimble.paths
|
||||||
|
nimbledeps
|
||||||
@@ -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"
|
||||||
|
}
|
||||||
+66
-5
@@ -1,17 +1,27 @@
|
|||||||
## PPO_Bot — enemy tracker + state vector wired into the game loop.
|
## PPO_Bot — enemy tracker + state vector wired into the game loop.
|
||||||
## Forward pass uses a random ActorCritic policy (weights not yet trained).
|
## Training: trajectory collected per tick, PPO update on round end.
|
||||||
|
|
||||||
import std/os
|
import std/os
|
||||||
|
import arraymancer
|
||||||
import tankroyale_botapi
|
import tankroyale_botapi
|
||||||
import network
|
import network
|
||||||
import actions
|
import actions
|
||||||
|
import training
|
||||||
import ./enemy_tracker
|
import ./enemy_tracker
|
||||||
import ./state_vector
|
import ./state_vector
|
||||||
|
|
||||||
const botJsonPath = currentSourcePath().parentDir / "PPO_Bot.json"
|
const botJsonPath = currentSourcePath().parentDir / "PPO_Bot.json"
|
||||||
|
|
||||||
type PPOBot = ref object of Bot
|
type PPOBot = ref object of Bot
|
||||||
tracker: EnemyTracker
|
tracker: EnemyTracker
|
||||||
|
buffer: TrajectoryBuffer
|
||||||
|
prevEnergy: float32 # own energy last tick
|
||||||
|
prevEnemyE: float32 # enemy energy last tick (from tracker)
|
||||||
|
lastState: Tensor[float32]
|
||||||
|
lastAction: Tensor[float32]
|
||||||
|
lastLogP: float32
|
||||||
|
lastValue: float32
|
||||||
|
hasLastTrans: bool
|
||||||
|
|
||||||
var ac = initActorCritic()
|
var ac = initActorCritic()
|
||||||
|
|
||||||
@@ -19,9 +29,29 @@ method onScannedBot*(bot: PPOBot, e: ScannedBotEvent) =
|
|||||||
bot.tracker.update(e.x, e.y, e.direction, e.speed, e.energy)
|
bot.tracker.update(e.x, e.y, e.direction, e.speed, e.energy)
|
||||||
|
|
||||||
method onRoundStarted*(bot: PPOBot, e: RoundStartedEvent) =
|
method onRoundStarted*(bot: PPOBot, e: RoundStartedEvent) =
|
||||||
bot.tracker = initEnemyTracker()
|
bot.tracker = initEnemyTracker()
|
||||||
|
bot.buffer = initTrajectoryBuffer()
|
||||||
|
bot.prevEnergy = 0.0'f32
|
||||||
|
bot.prevEnemyE = 0.0'f32
|
||||||
|
bot.hasLastTrans = false
|
||||||
|
|
||||||
|
method onRoundEnded*(bot: PPOBot, e: RoundEndedEventForBot) =
|
||||||
|
# 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
|
||||||
|
|
||||||
|
# PPO update synchronous; background thread is issue #17
|
||||||
|
# ponytail: blocking update per round; move to thread pool when #17 lands
|
||||||
|
ppoUpdate(ac, bot.buffer, lastValue = 0.0'f32)
|
||||||
|
bot.buffer.clear()
|
||||||
|
bot.hasLastTrans = false
|
||||||
|
|
||||||
method run(bot: PPOBot) =
|
method run(bot: PPOBot) =
|
||||||
|
# Seed energy on first tick
|
||||||
|
bot.prevEnergy = getEnergy().float32
|
||||||
|
bot.prevEnemyE = if bot.tracker.hasContact: bot.tracker.current.energy.float32 else: 0.0'f32
|
||||||
|
|
||||||
while isRunning():
|
while isRunning():
|
||||||
bot.tracker.deadReckon()
|
bot.tracker.deadReckon()
|
||||||
|
|
||||||
@@ -41,9 +71,37 @@ method run(bot: PPOBot) =
|
|||||||
)
|
)
|
||||||
|
|
||||||
let state = buildStateVector(botData, bot.tracker)
|
let state = buildStateVector(botData, bot.tracker)
|
||||||
let (rawActs, _) = ac.actorForward(state)
|
let (rawActs, logP) = ac.actorForward(state)
|
||||||
|
let value = ac.criticForward(state)
|
||||||
let acts = mapActions(rawActs, getSpeed().float32, getGunHeat().float32)
|
let acts = mapActions(rawActs, getSpeed().float32, getGunHeat().float32)
|
||||||
|
|
||||||
|
# Compute tick reward from energy deltas
|
||||||
|
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 tickReward = computeTickReward(myDelta, enemyDelta)
|
||||||
|
|
||||||
|
# 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,
|
||||||
|
)
|
||||||
|
bot.buffer.add(tr)
|
||||||
|
|
||||||
|
# Store current for next tick
|
||||||
|
bot.lastState = state
|
||||||
|
bot.lastAction = rawActs
|
||||||
|
bot.lastLogP = logP
|
||||||
|
bot.lastValue = value
|
||||||
|
bot.prevEnergy = curEnergy
|
||||||
|
bot.prevEnemyE = curEnemyE
|
||||||
|
bot.hasLastTrans = true
|
||||||
|
|
||||||
setTargetSpeed(acts.targetSpeed.float)
|
setTargetSpeed(acts.targetSpeed.float)
|
||||||
setTurnRate(acts.turnRate.float)
|
setTurnRate(acts.turnRate.float)
|
||||||
setGunTurnRate(acts.gunTurnRate.float)
|
setGunTurnRate(acts.gunTurnRate.float)
|
||||||
@@ -52,5 +110,8 @@ method run(bot: PPOBot) =
|
|||||||
go()
|
go()
|
||||||
|
|
||||||
when isMainModule:
|
when isMainModule:
|
||||||
var bot = PPOBot(tracker: initEnemyTracker())
|
var bot = PPOBot(
|
||||||
|
tracker: initEnemyTracker(),
|
||||||
|
buffer: initTrajectoryBuffer(),
|
||||||
|
)
|
||||||
start(bot, botJsonPath)
|
start(bot, botJsonPath)
|
||||||
|
|||||||
@@ -0,0 +1,11 @@
|
|||||||
|
# Package
|
||||||
|
version = "0.1.0"
|
||||||
|
author = "Davide Cappellini"
|
||||||
|
description = "PPO-trained Tank Royale bot"
|
||||||
|
license = "MIT"
|
||||||
|
bin = @["PPO_Bot"]
|
||||||
|
|
||||||
|
# Dependencies
|
||||||
|
requires "nim >= 2.0.0"
|
||||||
|
requires "tankroyale_botapi >= 1.0.0"
|
||||||
|
requires "arraymancer >= 0.7.0"
|
||||||
Executable
+4
@@ -0,0 +1,4 @@
|
|||||||
|
#!/bin/sh
|
||||||
|
# PPO_Bot — PPO-trained RL bot (compiled native binary)
|
||||||
|
cd -- "$(dirname -- "$0")"
|
||||||
|
exec "./PPO_Bot"
|
||||||
@@ -0,0 +1,8 @@
|
|||||||
|
# Static-link OpenBLAS for portable deployment
|
||||||
|
# ponytail: adjust path per machine, or use pkg-config
|
||||||
|
switch("passL", "-lopenblas")
|
||||||
|
switch("threads", "on")
|
||||||
|
# begin Nimble config (version 2)
|
||||||
|
when withDir(thisDir(), system.fileExists("nimble.paths")):
|
||||||
|
include "nimble.paths"
|
||||||
|
# end Nimble config
|
||||||
@@ -58,3 +58,16 @@ proc criticForward*(ac: ActorCritic, state: Tensor[float32]): float32 =
|
|||||||
## state: [42]. Returns scalar value estimate.
|
## state: [42]. Returns scalar value estimate.
|
||||||
let val = ac.critic.forward(state)
|
let val = ac.critic.forward(state)
|
||||||
result = val[0]
|
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 = max(v, -3.0'f32))
|
||||||
|
let std = logStdClamped.map(proc(v: float32): float32 = exp(v))
|
||||||
|
var logP = 0.0'f32
|
||||||
|
for i in 0..<5:
|
||||||
|
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,7 @@
|
|||||||
|
{ pkgs ? import <nixpkgs> {} }:
|
||||||
|
|
||||||
|
pkgs.mkShell {
|
||||||
|
buildInputs = with pkgs; [
|
||||||
|
openblas
|
||||||
|
];
|
||||||
|
}
|
||||||
@@ -0,0 +1,102 @@
|
|||||||
|
## test_training.nim — assert-based tests for training.nim
|
||||||
|
## Run: nim c tests/test_training.nim && ./tests/test_training
|
||||||
|
|
||||||
|
import std/[math, random]
|
||||||
|
import arraymancer
|
||||||
|
import "../network"
|
||||||
|
import "../training"
|
||||||
|
|
||||||
|
template check(cond: bool, msg: string) =
|
||||||
|
if not cond:
|
||||||
|
quit("FAIL: " & msg, 1)
|
||||||
|
|
||||||
|
# ── computeTickReward ─────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
block testTickReward:
|
||||||
|
# I lost 2, enemy lost 10 → reward = -2 - (-10) = 8
|
||||||
|
let r = computeTickReward(-2.0'f32, -10.0'f32)
|
||||||
|
check abs(r - 8.0'f32) < 1e-6'f32, "computeTickReward(-2, -10) == 8, 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
|
||||||
|
|
||||||
|
# ── TrajectoryBuffer ──────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
block testBuffer:
|
||||||
|
var buf = initTrajectoryBuffer()
|
||||||
|
check buf.len == 0, "empty buffer len == 0"
|
||||||
|
|
||||||
|
let t1 = Transition(state: zeros[float32](42), action: zeros[float32](5),
|
||||||
|
logProb: -1.0'f32, reward: 0.5'f32, value: 0.3'f32)
|
||||||
|
buf.add(t1)
|
||||||
|
buf.add(t1)
|
||||||
|
buf.add(t1)
|
||||||
|
check buf.len == 3, "buffer len == 3 after 3 adds"
|
||||||
|
|
||||||
|
buf.clear()
|
||||||
|
check buf.len == 0, "buffer len == 0 after clear"
|
||||||
|
|
||||||
|
# ── computeGAE — hand-calculated 3-step ──────────────────────────────────────
|
||||||
|
|
||||||
|
block testGAE:
|
||||||
|
# rewards = [1.0, 0.0, 1.0], values = [0.5, 0.5, 0.5], lastValue = 0.0
|
||||||
|
# gamma = 0.99, lam = 0.95
|
||||||
|
# delta_2 = 1.0 + 0.99*0.0 - 0.5 = 0.5
|
||||||
|
# adv_2 = 0.5
|
||||||
|
# delta_1 = 0.0 + 0.99*0.5 - 0.5 = -0.005
|
||||||
|
# adv_1 = -0.005 + 0.99*0.95*0.5 ≈ -0.005 + 0.47025 = 0.46525
|
||||||
|
# delta_0 = 1.0 + 0.99*0.5 - 0.5 = 0.995
|
||||||
|
# adv_0 = 0.995 + 0.99*0.95*0.46525 ≈ 0.995 + 0.43744 = 1.43244
|
||||||
|
let (adv, ret) = computeGAE(
|
||||||
|
rewards = @[1.0'f32, 0.0'f32, 1.0'f32],
|
||||||
|
values = @[0.5'f32, 0.5'f32, 0.5'f32],
|
||||||
|
lastValue = 0.0'f32,
|
||||||
|
gamma = 0.99'f32,
|
||||||
|
lam = 0.95'f32
|
||||||
|
)
|
||||||
|
|
||||||
|
check abs(adv[2] - 0.5'f32) < 1e-4'f32,
|
||||||
|
"adv[2] should be ~0.5, got " & $adv[2]
|
||||||
|
check abs(adv[1] - 0.46525'f32) < 1e-3'f32,
|
||||||
|
"adv[1] should be ~0.46525, got " & $adv[1]
|
||||||
|
check abs(adv[0] - 1.43244'f32) < 1e-2'f32,
|
||||||
|
"adv[0] should be ~1.43244, got " & $adv[0]
|
||||||
|
|
||||||
|
# returns = adv + values
|
||||||
|
check abs(ret[2] - (0.5'f32 + 0.5'f32)) < 1e-4'f32, "ret[2] = adv[2] + 0.5"
|
||||||
|
check abs(ret[0] - (adv[0] + 0.5'f32)) < 1e-4'f32, "ret[0] = adv[0] + 0.5"
|
||||||
|
|
||||||
|
# ── ppoUpdate runs without crash; weights change ──────────────────────────────
|
||||||
|
|
||||||
|
block testPpoUpdate:
|
||||||
|
randomize(42)
|
||||||
|
var ac = initActorCritic()
|
||||||
|
|
||||||
|
# Save a copy of w1 before update
|
||||||
|
let w1Before = ac.actor.w1.clone()
|
||||||
|
|
||||||
|
var buf = initTrajectoryBuffer()
|
||||||
|
for _ in 0..<10:
|
||||||
|
let s = randomNormalTensor[float32](42)
|
||||||
|
let a = randomNormalTensor[float32](5)
|
||||||
|
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))
|
||||||
|
|
||||||
|
ppoUpdate(ac, buf, lastValue = 0.0'f32, epochs = 2, miniBatchSize = 5)
|
||||||
|
|
||||||
|
# Weights should have changed — compare flattened
|
||||||
|
let n = ac.actor.w1.shape[0] * ac.actor.w1.shape[1]
|
||||||
|
let w1After = ac.actor.w1.reshape(n)
|
||||||
|
let w1Flat = w1Before.reshape(n)
|
||||||
|
var changed = false
|
||||||
|
for i in 0..<n:
|
||||||
|
if abs(w1After[i] - w1Flat[i]) > 1e-9'f32:
|
||||||
|
changed = true
|
||||||
|
break
|
||||||
|
check changed, "actor w1 should change after ppoUpdate"
|
||||||
|
|
||||||
|
echo "All tests passed"
|
||||||
@@ -0,0 +1,347 @@
|
|||||||
|
## 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 ─────────────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
type
|
||||||
|
Transition* = object
|
||||||
|
state*: Tensor[float32] # [42]
|
||||||
|
action*: Tensor[float32] # [5]
|
||||||
|
logProb*: float32
|
||||||
|
reward*: float32
|
||||||
|
value*: float32 # critic estimate at collection time
|
||||||
|
|
||||||
|
TrajectoryBuffer* = object
|
||||||
|
transitions*: seq[Transition]
|
||||||
|
|
||||||
|
# ── Buffer ─────────────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
proc initTrajectoryBuffer*(): TrajectoryBuffer =
|
||||||
|
result.transitions = @[]
|
||||||
|
|
||||||
|
proc add*(buf: var TrajectoryBuffer, t: Transition) =
|
||||||
|
buf.transitions.add(t)
|
||||||
|
|
||||||
|
proc clear*(buf: var TrajectoryBuffer) =
|
||||||
|
buf.transitions.setLen(0)
|
||||||
|
|
||||||
|
proc len*(buf: TrajectoryBuffer): int =
|
||||||
|
buf.transitions.len
|
||||||
|
|
||||||
|
# ── Reward helpers ─────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
proc computeTickReward*(myEnergyDelta, enemyEnergyDelta: float32): float32 =
|
||||||
|
## Positive when we deal more damage than we receive.
|
||||||
|
result = myEnergyDelta - enemyEnergyDelta
|
||||||
|
|
||||||
|
proc computeRoundReward*(roundScore: float32): float32 =
|
||||||
|
## Normalise round-end score to a rough ±3 range.
|
||||||
|
result = roundScore / 100.0'f32
|
||||||
|
|
||||||
|
# ── GAE ───────────────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
proc computeGAE*(rewards, values: seq[float32];
|
||||||
|
lastValue: float32;
|
||||||
|
gamma: float32 = 0.99'f32;
|
||||||
|
lam: float32 = 0.95'f32):
|
||||||
|
tuple[advantages: seq[float32], returns: seq[float32]] =
|
||||||
|
## Generalised Advantage Estimation — reverse sweep.
|
||||||
|
## lastValue = 0 for natural episode end (death/win).
|
||||||
|
let n = rewards.len
|
||||||
|
var advantages = newSeq[float32](n)
|
||||||
|
var gaeAcc = 0.0'f32
|
||||||
|
|
||||||
|
for t in countdown(n - 1, 0):
|
||||||
|
let nextVal = if t == n - 1: lastValue else: values[t + 1]
|
||||||
|
let delta = rewards[t] + gamma * nextVal - values[t]
|
||||||
|
gaeAcc = delta + gamma * lam * gaeAcc
|
||||||
|
advantages[t] = gaeAcc
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
# Persistent Adam state — survives across ppoUpdate calls (lives in training module)
|
||||||
|
var gAdamStates: ACAdamStates
|
||||||
|
var gAdamInit = false
|
||||||
|
|
||||||
|
# ── 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;
|
||||||
|
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) =
|
||||||
|
if buffer.len == 0: return
|
||||||
|
|
||||||
|
# Initialise Adam states once (persists across rounds)
|
||||||
|
if not gAdamInit:
|
||||||
|
gAdamStates = initACAdamStates(ac)
|
||||||
|
gAdamInit = true
|
||||||
|
|
||||||
|
# 1. GAE
|
||||||
|
let rewards = buffer.transitions.mapIt(it.reward)
|
||||||
|
let values = buffer.transitions.mapIt(it.value)
|
||||||
|
let (advantages, returns) = computeGAE(rewards, values, lastValue)
|
||||||
|
|
||||||
|
# 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
|
||||||
|
let advStd = sqrt(advVar + 1e-8'f32)
|
||||||
|
let normAdv = advantages.mapIt((it - advMean) / advStd)
|
||||||
|
|
||||||
|
let bufLen = buffer.len
|
||||||
|
|
||||||
|
for _ in 1..epochs:
|
||||||
|
# Shuffle indices
|
||||||
|
var indices = toSeq(0..<bufLen)
|
||||||
|
shuffle(indices)
|
||||||
|
|
||||||
|
var mbStart = 0
|
||||||
|
while mbStart < bufLen:
|
||||||
|
let mbEnd = min(mbStart + miniBatchSize, bufLen)
|
||||||
|
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
|
||||||
|
|
||||||
|
# ── Actor forward ──
|
||||||
|
let actorFwd = mlpForwardCached(ac.actor, tr.state)
|
||||||
|
let newMean = actorFwd.y # [5]
|
||||||
|
|
||||||
|
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:
|
||||||
|
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)
|
||||||
|
|
||||||
|
let ratio = exp(newLogP - tr.logProb)
|
||||||
|
|
||||||
|
# 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)
|
||||||
|
# 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](5)
|
||||||
|
for i in 0..<5:
|
||||||
|
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..<5:
|
||||||
|
let isClamped = (ac.logStd[i] <= -3.0'f32)
|
||||||
|
if not isClamped:
|
||||||
|
let s = std[i]
|
||||||
|
let diff = (tr.action[i] - newMean[i]) / s
|
||||||
|
let dLogP_dLogStdI = diff * diff - 1.0'f32
|
||||||
|
dLogStd[i] += dLoss_dNewLogP * dLogP_dLogStdI -
|
||||||
|
entropyCoeff / mbSize.float32 # entropy: d(-entropyCoeff*H)/d(logStd_i) = -entropyCoeff
|
||||||
|
|
||||||
|
# Backprop actor gradients
|
||||||
|
let gradActorOut = dLoss_dNewLogP *. dLogP_dMean # [5]
|
||||||
|
let actorGrads = mlpBackward(ac.actor, actorFwd, tr.state, 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, tr.state)
|
||||||
|
let newVal = criticFwd.y[0]
|
||||||
|
# Value loss = (newVal - ret)^2; d/d(newVal) = 2*(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)
|
||||||
|
|
||||||
|
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)
|
||||||
|
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, gAdamStates.aw1, lr)
|
||||||
|
adamStep(ac.actor.b1, dActorB1, gAdamStates.ab1, lr)
|
||||||
|
adamStep(ac.actor.w2, dActorW2, gAdamStates.aw2, lr)
|
||||||
|
adamStep(ac.actor.b2, dActorB2, gAdamStates.ab2, lr)
|
||||||
|
adamStep(ac.actor.w3, dActorW3, gAdamStates.aw3, lr)
|
||||||
|
adamStep(ac.actor.b3, dActorB3, gAdamStates.ab3, lr)
|
||||||
|
adamStep(ac.logStd, dLogStd, gAdamStates.logStd, lr)
|
||||||
|
adamStep(ac.critic.w1, dCriticW1, gAdamStates.cw1, lr)
|
||||||
|
adamStep(ac.critic.b1, dCriticB1, gAdamStates.cb1, lr)
|
||||||
|
adamStep(ac.critic.w2, dCriticW2, gAdamStates.cw2, lr)
|
||||||
|
adamStep(ac.critic.b2, dCriticB2, gAdamStates.cb2, lr)
|
||||||
|
adamStep(ac.critic.w3, dCriticW3, gAdamStates.cw3, lr)
|
||||||
|
adamStep(ac.critic.b3, dCriticB3, gAdamStates.cb3, lr)
|
||||||
|
|
||||||
|
mbStart = mbEnd
|
||||||
Reference in New Issue
Block a user