Compare commits
74 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 96065597c0 | |||
| eae6fc15a2 | |||
| 81718e3a4c | |||
| 04c149ea28 | |||
| 533a146342 | |||
| 03a1853e57 | |||
| 002a568c7e | |||
| f3ac0888bb | |||
| 0b0933294a | |||
| 781e41595e | |||
| 40e074e5a3 | |||
| 167bcc4ce5 | |||
| 1619b86f25 | |||
| 19f34abf0c | |||
| bd58794b4c | |||
| 6fc01eb4e5 | |||
| a07e5305f5 | |||
| 2619ba06fc | |||
| f45e8f2717 | |||
| b7492f1080 | |||
| f1962c7506 | |||
| 05929d2dbd | |||
| 4b64bf18ac | |||
| 26536713ba | |||
| 2f49cb243f | |||
| 6a294ad7ad | |||
| edf26aa45d | |||
| df256b4d3e | |||
| 7104645f5d | |||
| 32b71d9fc8 | |||
| 62a6cc8ccf | |||
| 415d4e3738 | |||
| 717ef3ead8 | |||
| 54b8139b11 | |||
| 55ef22ff8b | |||
| 5ea57bcae3 | |||
| cb33551621 | |||
| 23c65c9ac6 | |||
| 4ee0d8272c | |||
| add3e34926 | |||
| f130bf1254 | |||
| a0a3840980 | |||
| 7d73d32c85 | |||
| df3bbbd14e | |||
| 4258d364b9 | |||
| f88580b157 | |||
| ca3e3d2272 | |||
| 0d35646dc9 | |||
| c834d2cbee | |||
| fedab54bc0 | |||
| 82eeb53e5c | |||
| 12624d3069 | |||
| 6ad51148f4 | |||
| 75e32e3315 | |||
| db99153f65 | |||
| da2f825ad8 | |||
| 186e005a96 | |||
| a4e830531b | |||
| 64697f917e | |||
| 766b9e03ee | |||
| 0b17430735 | |||
| cdde60d79f | |||
| 56e0b306c9 | |||
| e5609a7d9b | |||
| a8ee2a86e3 | |||
| bbc9e51166 | |||
| fd22535f5b | |||
| f27b0238f0 | |||
| aea0724d3a | |||
| 473d67f644 | |||
| eadd177d3b | |||
| 588c9ebc2f | |||
| aa4bc77068 | |||
| 30cda871cc |
-59
@@ -1,59 +0,0 @@
|
||||
# Evo_Bot
|
||||
|
||||
Robocode Tank Royale bot with a modular gun system where evolved neural networks learn to predict enemy dodge behavior.
|
||||
|
||||
## Language
|
||||
|
||||
**Evo_Bot**:
|
||||
The bot itself — handles movement and firing discipline. Guns are pluggable modules.
|
||||
_Avoid_: robot, tank
|
||||
|
||||
**Guess Factor (GF)**:
|
||||
A value from -1 to +1 representing where on the maximum escape angle arc the enemy is. 0 = directly ahead, -1 = full left dodge, +1 = full right dodge. The gun's prediction target.
|
||||
_Avoid_: aim offset, dodge index
|
||||
|
||||
**Max Escape Angle (MEA)**:
|
||||
The widest angle the enemy can reach before a bullet arrives, computed from distance and bullet speed. Guess factor is multiplied by MEA to get the aim offset.
|
||||
|
||||
**Lateral Velocity**:
|
||||
Enemy speed projected perpendicular to the line between you and them. The primary signal for guess factor prediction.
|
||||
_Avoid_: tangential speed, sideways velocity
|
||||
|
||||
**Sliding Window**:
|
||||
The last N ticks (default 30) of enemy state fed as input to the network. Each tick contains lateral velocity, heading delta, and wall distance ahead.
|
||||
_Avoid_: observation buffer, input history
|
||||
|
||||
**Replay Tape**:
|
||||
Rolling buffer of recorded enemy states (~2000 ticks). The evolution thread evaluates gun fitness against this tape.
|
||||
_Avoid_: experience buffer, replay buffer
|
||||
|
||||
**Virtual Gun**:
|
||||
A gun that runs in parallel without actually firing. It tracks where it would have aimed and whether a simulated bullet would have hit. Used to compare gun variants.
|
||||
|
||||
**Virtual Bullet**:
|
||||
A simulated bullet fired by a virtual gun. Never actually sent to the game engine.
|
||||
|
||||
**TOPO_Gun**:
|
||||
Fixed-topology ANN gun evolved by GA. Network shape is predetermined (e.g., 91-8-1), only weights are evolved.
|
||||
_Avoid_: static gun, fixed gun
|
||||
|
||||
**NEAT_Gun**:
|
||||
Variable-topology ANN gun where evolution can add/remove neurons and connections (NEAT algorithm). Deferred — only built if TOPO_Gun hits a ceiling.
|
||||
|
||||
**Population**:
|
||||
The set of candidate networks (default 64-200) being evolved. Each member is a complete set of ANN weights.
|
||||
|
||||
**Champion**:
|
||||
The best-performing network in the current GA population. The champion's weights are pushed to the inference side when it beats the current best.
|
||||
_Avoid_: best, winner, elite
|
||||
|
||||
**Fitness**:
|
||||
Hit count when a network's aim predictions are evaluated as virtual bullets against sampled ticks from the replay tape.
|
||||
_Avoid_: score, reward
|
||||
|
||||
**Cold Start**:
|
||||
The first-ever battle with no saved weights. The gun does not fire until the evolution thread produces its first champion. Subsequent battles load persisted weights.
|
||||
|
||||
**Weight Persistence**:
|
||||
Saving evolved weights to disk. Load order: per-opponent file, then global fallback, then random initialization.
|
||||
_Avoid_: model saving, checkpointing
|
||||
@@ -0,0 +1,3 @@
|
||||
nimble.develop
|
||||
nimble.paths
|
||||
nimbledeps
|
||||
Executable
BIN
Binary file not shown.
@@ -1,11 +1,11 @@
|
||||
{
|
||||
"name": "OscillatorBot",
|
||||
"name": "GotoTest",
|
||||
"version": "0.1.0",
|
||||
"authors": ["Davide Cappellini"],
|
||||
"description": "Predictable zigzag sparring partner for gun testing",
|
||||
"description": "Throwaway goto(x,y) controller prototype",
|
||||
"homepage": "",
|
||||
"countryCodes": ["IT"],
|
||||
"gameTypes": ["classic", "1v1"],
|
||||
"gameTypes": ["classic", "melee", "1v1"],
|
||||
"platform": "Nim",
|
||||
"programmingLang": "Nim"
|
||||
}
|
||||
@@ -0,0 +1,95 @@
|
||||
## GotoTest — throwaway prototype to validate goto(x,y) + aimTo(x,y) controllers.
|
||||
## Diamond pattern with wall-smash north waypoint to test stuck recovery.
|
||||
## Tank Royale: Y=0 is south, Y increases northward. Arena 800x600.
|
||||
|
||||
import std/[strformat, os]
|
||||
import tankroyale_botapi
|
||||
|
||||
const botJsonPath = currentSourcePath().parentDir / "GotoTest.json"
|
||||
|
||||
# Diamond waypoints. (400,600) is AT the north wall — physically impossible, tests wall recovery.
|
||||
const waypoints = [
|
||||
(400.0, 300.0), # center
|
||||
(400.0, 600.0), # north wall — AT wall, unreachable, tests wall-stuck recovery
|
||||
( 36.0, 300.0), # west wall (left)
|
||||
(400.0, 36.0), # south wall (bottom)
|
||||
(764.0, 300.0), # east wall (right)
|
||||
(400.0, 300.0), # center
|
||||
]
|
||||
|
||||
type GotoBot = ref object of Bot
|
||||
waypointIdx: int
|
||||
oldX, oldY: float # position 3 ticks ago for stuck detection
|
||||
tickCount: int # ticks since oldX/oldY was last updated
|
||||
stuckTicks: int # remaining ticks of reverse override
|
||||
|
||||
|
||||
method onRoundStarted*(bot: GotoBot, e: RoundStartedEvent) =
|
||||
setAdjustGunForBodyTurn(true)
|
||||
bot.waypointIdx = 0
|
||||
bot.oldX = 0.0; bot.oldY = 0.0
|
||||
bot.tickCount = 0; bot.stuckTicks = 0
|
||||
|
||||
method run(bot: GotoBot) =
|
||||
while isRunning():
|
||||
let bx = getX()
|
||||
let by = getY()
|
||||
let (tx, ty) = waypoints[bot.waypointIdx]
|
||||
let dist = distanceTo(bx, by, tx, ty)
|
||||
|
||||
if dist < 40.0:
|
||||
bot.waypointIdx = (bot.waypointIdx + 1) mod waypoints.len
|
||||
|
||||
# goto controller — explicit forward/reverse proportional steering
|
||||
# ponytail: no speed-taper on heading error; add when overshooting observed
|
||||
let rawBearing = normalizeRelativeAngle(directionTo(tx, ty) - getDirection())
|
||||
let (dirSign, effBearing) =
|
||||
if abs(rawBearing) > 90.0:
|
||||
(-1.0, normalizeRelativeAngle(rawBearing + 180.0))
|
||||
else:
|
||||
(1.0, rawBearing)
|
||||
let maxTurn = calcMaxTurnRate(getSpeed())
|
||||
setTurnRate(effBearing.clamp(-maxTurn, maxTurn))
|
||||
|
||||
# stuck detector — ponytail: 3-tick sample, upgrade to wall-nav if needed
|
||||
inc bot.tickCount
|
||||
if bot.tickCount >= 3:
|
||||
let moved = distanceTo(bot.oldX, bot.oldY, bx, by)
|
||||
if moved < 1.0 and dist >= 30.0:
|
||||
bot.stuckTicks = 6
|
||||
echo &"STUCK at ({bx:.1f},{by:.1f}) moved={moved:.2f} — reversing"
|
||||
bot.oldX = bx; bot.oldY = by
|
||||
bot.tickCount = 0
|
||||
|
||||
if bot.stuckTicks > 0:
|
||||
# reverse current speed direction to unstick
|
||||
let targetSpd = if getSpeed() >= 0.0: -8.0 else: 8.0
|
||||
setTargetSpeed(targetSpd)
|
||||
dec bot.stuckTicks
|
||||
else:
|
||||
let rawSpeed = getNewTargetSpeed(MAX_SPEED, abs(getSpeed()), dist)
|
||||
setTargetSpeed(dirSign * rawSpeed)
|
||||
|
||||
# aimTo arena center — radarBearingTo reuses getX/getY/getGunDirection internally
|
||||
let cx = getArenaWidth().float / 2.0
|
||||
let cy = getArenaHeight().float / 2.0
|
||||
let gunTurnNeeded = normalizeRelativeAngle(directionTo(cx, cy) - getGunDirection())
|
||||
setGunTurnRate(gunTurnNeeded.clamp(-MAX_GUN_TURN_RATE, MAX_GUN_TURN_RATE))
|
||||
|
||||
echo &"pos=({bx:.1f},{by:.1f}) target=({tx:.1f},{ty:.1f}) dist={dist:.1f} spd={getSpeed():.2f} stuck={bot.stuckTicks}"
|
||||
|
||||
# Debug graphics: waypoint circles + line to current target
|
||||
for i, (wx, wy) in waypoints:
|
||||
if i == bot.waypointIdx:
|
||||
setStrokeColor(RED)
|
||||
else:
|
||||
setStrokeColor(GRAY)
|
||||
drawCircle(wx, wy, 10.0)
|
||||
setStrokeColor(YELLOW)
|
||||
drawLine(bx, by, tx, ty)
|
||||
|
||||
go()
|
||||
|
||||
when isMainModule:
|
||||
var bot = GotoBot()
|
||||
start(bot, botJsonPath)
|
||||
@@ -0,0 +1,10 @@
|
||||
# Package
|
||||
version = "0.1.0"
|
||||
author = "Davide Cappellini"
|
||||
description = "GotoTest — throwaway goto(x,y) controller prototype"
|
||||
license = "MIT"
|
||||
bin = @["GotoTest"]
|
||||
|
||||
# Dependencies
|
||||
requires "nim >= 2.0.0"
|
||||
requires "tankroyale_botapi >= 1.0.0"
|
||||
@@ -0,0 +1,4 @@
|
||||
# begin Nimble config (version 2)
|
||||
when withDir(thisDir(), system.fileExists("nimble.paths")):
|
||||
include "nimble.paths"
|
||||
# end Nimble config
|
||||
@@ -1,48 +0,0 @@
|
||||
# OscillatorBot — predictable zigzag sparring partner for GA gun testing.
|
||||
# Reverses direction + turn every PERIOD ticks. Head-on targeting only.
|
||||
|
||||
import std/[math, os]
|
||||
import tankroyale_botapi
|
||||
|
||||
const botJsonPath = currentSourcePath().parentDir / "OscillatorBot.json"
|
||||
|
||||
const
|
||||
SPEED = 7.0 # forward/backward speed
|
||||
PERIOD = 25 # ticks between direction reversals
|
||||
|
||||
type OscillatorBot = ref object of Bot
|
||||
tickCount: int
|
||||
moveSign: float # +1 forward, -1 backward
|
||||
turnSign: float # +1 right, -1 left
|
||||
|
||||
method onRoundStarted*(bot: OscillatorBot, e: RoundStartedEvent) =
|
||||
setAdjustGunForBodyTurn(true)
|
||||
setAdjustRadarForBodyTurn(true)
|
||||
setAdjustRadarForGunTurn(true)
|
||||
bot.tickCount = 0
|
||||
bot.moveSign = 1.0
|
||||
bot.turnSign = 1.0
|
||||
|
||||
method onScannedBot*(bot: OscillatorBot, e: ScannedBotEvent) =
|
||||
# Head-on targeting: aim gun directly at enemy, fire medium power
|
||||
let bearing = directionTo(getX(), getY(), e.x, e.y)
|
||||
let gunDelta = normalizeRelativeAngle(bearing - getGunDirection())
|
||||
setGunTurnRate(gunDelta.clamp(-MAX_GUN_TURN_RATE, MAX_GUN_TURN_RATE))
|
||||
if abs(gunDelta) < 10.0 and getGunHeat() <= 0.0:
|
||||
discard setFire(2.0)
|
||||
|
||||
method run*(bot: OscillatorBot) =
|
||||
while isRunning():
|
||||
inc bot.tickCount
|
||||
if bot.tickCount mod PERIOD == 0:
|
||||
bot.moveSign *= -1.0
|
||||
bot.turnSign *= -1.0
|
||||
|
||||
setTargetSpeed(bot.moveSign * SPEED)
|
||||
setTurnRate(bot.turnSign * 4.0)
|
||||
setRadarTurnRate(45.0) # spin radar to keep scanning
|
||||
go()
|
||||
|
||||
when isMainModule:
|
||||
var bot = OscillatorBot(moveSign: 1.0, turnSign: 1.0)
|
||||
start(bot, botJsonPath)
|
||||
@@ -1,3 +0,0 @@
|
||||
#!/bin/sh
|
||||
cd "$(dirname "$0")"
|
||||
exec ./OscillatorBot 2>> /tmp/oscillatorbot_stderr.log
|
||||
@@ -1 +0,0 @@
|
||||
--path:"../libs"
|
||||
@@ -0,0 +1,3 @@
|
||||
nimble.develop
|
||||
nimble.paths
|
||||
nimbledeps
|
||||
Executable
BIN
Binary file not shown.
@@ -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 network
|
||||
import actions
|
||||
import training
|
||||
import weights
|
||||
import ./enemy_tracker
|
||||
import ./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,12 @@
|
||||
# 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"
|
||||
# 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"
|
||||
Executable
+7
@@ -0,0 +1,7 @@
|
||||
#!/bin/sh
|
||||
# PPO_Bot — PPO-trained RL bot (compiled native binary)
|
||||
# OPENBLAS_NUM_THREADS=1: prevent OpenBLAS from spawning worker threads,
|
||||
# which deadlock when called from within a multi-threaded Nim bot process.
|
||||
export OPENBLAS_NUM_THREADS=1
|
||||
cd -- "$(dirname -- "$0")"
|
||||
exec "./PPO_Bot"
|
||||
@@ -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,13 @@
|
||||
# Static-link OpenBLAS for portable deployment
|
||||
# ponytail: adjust path per machine, or use pkg-config
|
||||
switch("passL", "-L/nix/store/v07svn2y92bvzjl51aj7c9ca1cwg7rw7-openblas-0.3.32/lib -lopenblas")
|
||||
switch("threads", "on")
|
||||
# begin Nimble config (version 2)
|
||||
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")
|
||||
@@ -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,7 @@
|
||||
{ pkgs ? import <nixpkgs> {} }:
|
||||
|
||||
pkgs.mkShell {
|
||||
buildInputs = with pkgs; [
|
||||
openblas
|
||||
];
|
||||
}
|
||||
@@ -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))
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,69 @@
|
||||
## Assert-based tests for actions.nim.
|
||||
## Run: nim c -r tests/test_actions.nim
|
||||
|
||||
import arraymancer
|
||||
import std/math
|
||||
import "../actions"
|
||||
|
||||
template check(cond: bool, msg: string) =
|
||||
if not cond:
|
||||
quit("FAIL: " & msg, 1)
|
||||
|
||||
let arenaW = 1200.0
|
||||
let arenaH = 800.0
|
||||
|
||||
# Build a 6-element zero tensor and a helper to set individual values
|
||||
proc makeRaw(vals: array[6, float32]): Tensor[float32] =
|
||||
result = zeros[float32](6)
|
||||
for i in 0 ..< 6: result[i] = vals[i]
|
||||
|
||||
# --- 6-dim input produces a valid BotActions ---
|
||||
block basicDecode:
|
||||
let raw = makeRaw([0.0'f32, 0.0, 0.0, 0.0, 0.0, 0.0])
|
||||
let acts = mapActions(raw, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
||||
# sigmoid(0)*arenaW = 0.5*1200 = 600, sigmoid(0)*arenaH = 0.5*800 = 400
|
||||
check abs(acts.gotoX - 600.0) < 1e-6, "gotoX = sigmoid(0)*arenaW"
|
||||
check abs(acts.gotoY - 400.0) < 1e-6, "gotoY = sigmoid(0)*arenaH"
|
||||
check abs(acts.aimToX - 600.0) < 1e-6, "aimToX = sigmoid(0)*arenaW"
|
||||
check abs(acts.aimToY - 400.0) < 1e-6, "aimToY = sigmoid(0)*arenaH"
|
||||
|
||||
# --- Coordinates bounded to arena size ---
|
||||
block coordBounds:
|
||||
# Large positive raw → sigmoid ≈ 1 → close to arenaW/arenaH
|
||||
let rawHigh = makeRaw([100.0'f32, 100.0, 100.0, 100.0, 0.0, 0.0])
|
||||
let actsHigh = mapActions(rawHigh, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
||||
check actsHigh.gotoX <= arenaW + 1e-9, "gotoX <= arenaWidth"
|
||||
check actsHigh.gotoY <= arenaH + 1e-9, "gotoY <= arenaHeight"
|
||||
check actsHigh.gotoX >= 0.0, "gotoX >= 0"
|
||||
# Large negative raw → sigmoid ≈ 0 → close to 0
|
||||
let rawLow = makeRaw([-100.0'f32, -100.0, -100.0, -100.0, 0.0, 0.0])
|
||||
let actsLow = mapActions(rawLow, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
||||
check actsLow.gotoX >= -1e-9, "gotoX >= 0 (low raw)"
|
||||
check actsLow.gotoY >= -1e-9, "gotoY >= 0 (low raw)"
|
||||
|
||||
# --- Fire triggers correctly ---
|
||||
block fireTrigger:
|
||||
# tanh(positive) >= 0 → fire when gunHeat = 0
|
||||
let rawFire = makeRaw([0.0'f32, 0.0, 0.0, 0.0, 1.0, 0.0])
|
||||
let actsFire = mapActions(rawFire, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
||||
check actsFire.shouldFire, "positive tanh → should fire when gun cool"
|
||||
|
||||
# tanh(negative) < 0 → no fire
|
||||
let rawNoFire = makeRaw([0.0'f32, 0.0, 0.0, 0.0, -1.0, 0.0])
|
||||
let actsNoFire = mapActions(rawNoFire, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
||||
check not actsNoFire.shouldFire, "negative tanh → no fire"
|
||||
|
||||
# gunHeat > 0 → no fire even with positive decision
|
||||
let actsHot = mapActions(rawFire, 0.5, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
||||
check not actsHot.shouldFire, "positive tanh but gun hot → no fire"
|
||||
|
||||
# --- Fire power in [0.1, 3.0] ---
|
||||
block firePowerRange:
|
||||
let rawMin = makeRaw([0.0'f32, 0.0, 0.0, 0.0, 0.0, -100.0])
|
||||
let rawMax = makeRaw([0.0'f32, 0.0, 0.0, 0.0, 0.0, 100.0])
|
||||
let actsMin = mapActions(rawMin, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
||||
let actsMax = mapActions(rawMax, 0.0, arenaW, arenaH, 600.0, 400.0, 0.0, 0.0, 0.0, 600.0, 400.0)
|
||||
check actsMin.firePower >= 0.1 - 1e-6, "firePower >= 0.1"
|
||||
check actsMax.firePower <= 3.0 + 1e-6, "firePower <= 3.0"
|
||||
|
||||
echo "test_actions: all passed"
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,54 @@
|
||||
## Assert-based tests for controllers.nim.
|
||||
## Run: nim c -r tests/test_controllers.nim
|
||||
|
||||
import std/math
|
||||
import "../controllers"
|
||||
|
||||
template check(cond: bool, msg: string) =
|
||||
if not cond:
|
||||
quit("FAIL: " & msg, 1)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# gotoTick tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
block forwardMovement:
|
||||
# Bot at (0,0) facing East (geometric 0°), target due east at (100,0) → bearing=0 → forward
|
||||
let (spd2, turn2) = gotoTick(100.0, 0.0, 0.0, 0.0, 0.0, 0.0)
|
||||
check spd2 > 0.0, "forward: targetSpeed should be positive"
|
||||
check abs(turn2) < 1e-9, "forward: no turn needed when already aimed"
|
||||
|
||||
block reverseMovement:
|
||||
# Bot at (100,0) facing East (direction=0), target at (0,0) — directly behind
|
||||
# bearing = normalizeRelativeAngle(180 - 0) = 180 → |bearing|>90 → reverse
|
||||
let (spd, _) = gotoTick(0.0, 0.0, 100.0, 0.0, 0.0, 0.0)
|
||||
check spd < 0.0, "reverse: targetSpeed should be negative when target is behind"
|
||||
|
||||
block turnRateClamping:
|
||||
# Bot at (0,0) facing North (game north = geometric 90°, so direction=90 in geometric)
|
||||
# Target at (100,0) = East. bearing = normalizeRelativeAngle(0 - 90) = -90 → still ≤90
|
||||
# Use large perpendicular target so bearing is 89°, and high speed → small maxTurn
|
||||
# At speed=8, calcMaxTurnRate = 10 - 0.75*8 = 4°
|
||||
# Bot facing East (0°), target at angle 89° bearing (just under 90)
|
||||
let (_, turn) = gotoTick(100.0 * cos(89.0 * PI / 180.0), 100.0 * sin(89.0 * PI / 180.0), 0.0, 0.0, 0.0, 8.0)
|
||||
check abs(turn) <= 4.0 + 1e-9, "turn rate clamped to calcMaxTurnRate at speed=8 (max 4°)"
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# aimToTick tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
block gunShortestArc:
|
||||
# Gun facing East (0°), target due north (geometric 90°) → turn left +90° but clamped to 20
|
||||
let rate = aimToTick(0.0, 100.0, 0.0, 0.0, 0.0)
|
||||
check rate > 0.0, "gun shortest arc: should turn toward target"
|
||||
# Gun facing East (0°), target due south (geometric 270° → normalised -90°)
|
||||
let rate2 = aimToTick(0.0, -100.0, 0.0, 0.0, 0.0)
|
||||
check rate2 < 0.0, "gun shortest arc: should turn other way for target behind"
|
||||
|
||||
block gunTurnRateClamping:
|
||||
# 180° away → clamped to ±20
|
||||
let rate = aimToTick(-100.0, 0.0, 0.0, 0.0, 0.0) # target west, gun east
|
||||
check abs(rate) <= 20.0 + 1e-9, "gun turn rate clamped to ±MAX_GUN_TURN_RATE"
|
||||
check abs(abs(rate) - 20.0) < 1e-9, "gun turn rate at max when 180° away"
|
||||
|
||||
echo "test_controllers: all passed"
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,58 @@
|
||||
## test_network.nim — assert-based tests for network.nim and actions.nim.
|
||||
## Run: nim c --threads:on tests/test_network.nim && ./tests/test_network
|
||||
|
||||
import arraymancer
|
||||
import std/[math, strformat]
|
||||
import ../network
|
||||
import ../actions
|
||||
|
||||
func isFiniteF(x: float32): bool = classify(x) notin {fcNan, fcInf, fcNegInf}
|
||||
func isNaNF(x: float32): bool = classify(x) == fcNan
|
||||
|
||||
when isMainModule:
|
||||
# ---- MLP forward shape ----
|
||||
let mlp = initMLP(STATE_DIM, 64, ACTION_DIM)
|
||||
let inp = zeros[float32](STATE_DIM)
|
||||
let mlpOut = mlp.forward(inp)
|
||||
assert mlpOut.shape[0] == ACTION_DIM, &"MLP output shape wrong: {mlpOut.shape}"
|
||||
|
||||
# ---- ActorCritic actorForward ----
|
||||
let ac = initActorCritic()
|
||||
let state = zeros[float32](STATE_DIM)
|
||||
let (acts, logP) = ac.actorForward(state)
|
||||
assert acts.shape[0] == ACTION_DIM, &"actorForward actions shape wrong: {acts.shape}"
|
||||
assert not isNaNF(logP), "logProb is NaN"
|
||||
assert isFiniteF(logP), &"logProb not finite: {logP}"
|
||||
|
||||
# ---- criticForward ----
|
||||
let v = ac.criticForward(state)
|
||||
assert not isNaNF(v), "critic value is NaN"
|
||||
assert isFiniteF(v), &"critic value not finite: {v}"
|
||||
|
||||
# ---- logStd floor: collapsing logStd should not break actorForward ----
|
||||
var ac2 = initActorCritic()
|
||||
for i in 0..<ACTION_DIM: ac2.logStd[i] = -10.0'f32
|
||||
let (acts2, logP2) = ac2.actorForward(state)
|
||||
assert acts2.shape[0] == ACTION_DIM, "acts2 shape wrong after logStd=-10"
|
||||
assert isFiniteF(logP2), &"logP2 not finite with floored logStd: {logP2}"
|
||||
|
||||
# ---- action mapping ranges ----
|
||||
let raw = randomNormalTensor[float32](ACTION_DIM)
|
||||
let speed = 4.0'f32
|
||||
# arena 800×600, bot at centre, heading north, gun north
|
||||
let botActs = mapActions(raw, 0.0'f32, 800.0, 600.0, 400.0, 300.0, 0.0, speed.float, 0.0, 400.0, 300.0) # gunHeat=0 → fire allowed
|
||||
|
||||
assert botActs.targetSpeed >= -8.0'f32 and botActs.targetSpeed <= 8.0'f32,
|
||||
&"targetSpeed out of range: {botActs.targetSpeed}"
|
||||
|
||||
assert botActs.gunTurnRate >= -20.0'f32 and botActs.gunTurnRate <= 20.0'f32,
|
||||
&"gunTurnRate out of range: {botActs.gunTurnRate}"
|
||||
|
||||
assert botActs.firePower >= 0.1'f32 and botActs.firePower <= 3.0'f32,
|
||||
&"firePower out of range: {botActs.firePower}"
|
||||
|
||||
# shouldFire=false when gunHeat > 0
|
||||
let noFire = mapActions(raw, 1.0'f32, 800.0, 600.0, 400.0, 300.0, 0.0, speed.float, 0.0, 400.0, 300.0)
|
||||
assert not noFire.shouldFire, "shouldFire should be false when gunHeat > 0"
|
||||
|
||||
echo "All tests passed"
|
||||
@@ -0,0 +1,42 @@
|
||||
## Regression test: radar lock must hold on a stationary target.
|
||||
## Run: nim c -r tests/test_radar_lock.nim
|
||||
|
||||
import std/[math, strformat]
|
||||
import "../enemy_tracker"
|
||||
|
||||
template check(cond: bool, msg: string) =
|
||||
if not cond:
|
||||
quit("FAIL: " & msg, 1)
|
||||
|
||||
# Bot at arena center; stationary enemy due north (same x, higher y).
|
||||
# Tank Royale: y increases northward.
|
||||
# Radar uses math convention (0°=east, CCW+); north = 90° in that system.
|
||||
|
||||
let botX = 400.0
|
||||
let botY = 300.0
|
||||
let enemyX = 400.0 # same x → dx = 0
|
||||
let enemyY = 500.0 # north of bot → dy > 0
|
||||
let trueBearing = 90.0 # north in math convention (0=east, CCW+)
|
||||
|
||||
# Radar starts pointing at the enemy (radarDirection = 90°, due north in math convention).
|
||||
var radarDir = 90.0
|
||||
|
||||
var tracker = initEnemyTracker()
|
||||
# Prime with contact at the known position
|
||||
tracker.update(enemyX, enemyY, 0.0, 0.0, 100.0)
|
||||
|
||||
echo "Tick | radarDir | trueBearing | bearingErr"
|
||||
for tick in 1 .. 20:
|
||||
let rate = tracker.getRadarTurnRate(botX, botY, 0.0, radarDir)
|
||||
radarDir = (radarDir + rate + 360.0) mod 360.0
|
||||
|
||||
# Simulate a successful scan every tick (enemy is stationary)
|
||||
tracker.update(enemyX, enemyY, 0.0, 0.0, 100.0)
|
||||
|
||||
# Bearing error: signed difference, wrapped to [-180, 180]
|
||||
let err = ((radarDir - trueBearing) + 540.0) mod 360.0 - 180.0
|
||||
echo &" {tick:2d} | {radarDir:8.3f}° | {trueBearing:8.3f}° | {err:+.3f}°"
|
||||
|
||||
check abs(err) <= 15.0, &"tick {tick}: radar {radarDir:.1f}° drifted > 15° from target {trueBearing:.1f}°"
|
||||
|
||||
echo "All radar lock tests passed"
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,187 @@
|
||||
## Assert-based tests for EnemyTracker and StateVector.
|
||||
## Run: nim c -r tests/test_state.nim
|
||||
|
||||
import std/[math, strformat]
|
||||
import arraymancer
|
||||
|
||||
# Import from parent dir
|
||||
import "../enemy_tracker"
|
||||
import "../state_vector"
|
||||
|
||||
template check(cond: bool, msg: string) =
|
||||
if not cond:
|
||||
quit("FAIL: " & msg, 1)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# EnemyTracker tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
block testBasicUpdate:
|
||||
var t = initEnemyTracker()
|
||||
t.update(200.0, 300.0, 90.0, 5.0, 80.0)
|
||||
check t.hasContact, "hasContact after update"
|
||||
check t.current.x == 200.0, "x after update"
|
||||
check t.current.y == 300.0, "y after update"
|
||||
check t.current.direction == 90.0, "direction after update"
|
||||
check t.current.speed == 5.0, "speed after update"
|
||||
check t.current.energy == 80.0, "energy after update"
|
||||
check t.current.ticksSinceLastScan == 0, "ticksSinceLastScan reset"
|
||||
|
||||
block testFireDetection:
|
||||
var t = initEnemyTracker()
|
||||
# First update sets prevEnergy
|
||||
t.update(100.0, 100.0, 0.0, 0.0, 100.0)
|
||||
# Second update: energy drop of 3.0 → enemy fired power 3.0
|
||||
t.update(100.0, 100.0, 0.0, 0.0, 97.0)
|
||||
check t.current.hasFired, "hasFired when energy drops by 3.0"
|
||||
check abs(t.current.lastFirePower - 3.0) < 0.001, "lastFirePower == 3.0"
|
||||
|
||||
block testNoFireOnSmallDrop:
|
||||
var t = initEnemyTracker()
|
||||
t.update(100.0, 100.0, 0.0, 0.0, 100.0)
|
||||
# Drop of 0.05 — below MIN_FIRE_POWER threshold
|
||||
t.update(100.0, 100.0, 0.0, 0.0, 99.95)
|
||||
check not t.current.hasFired, "no fire on small energy drop"
|
||||
|
||||
block testDeadReckoning:
|
||||
var t = initEnemyTracker()
|
||||
# direction=0° in Tank Royale means north (y increases)
|
||||
t.update(100.0, 100.0, 0.0, 5.0, 100.0)
|
||||
t.deadReckon()
|
||||
# x unchanged (sin 0° = 0), y increases by speed (cos 0° = 1)
|
||||
check abs(t.current.x - 100.0) < 0.001, "dead reckon: x unchanged for dir=0"
|
||||
check abs(t.current.y - 105.0) < 0.001, "dead reckon: y += speed for dir=0"
|
||||
check t.current.ticksSinceLastScan == 1, "ticksSinceLastScan incremented"
|
||||
|
||||
block testDeadReckonEast:
|
||||
var t = initEnemyTracker()
|
||||
# direction=90° → east (sin 90° = 1, cos 90° = 0)
|
||||
t.update(100.0, 100.0, 90.0, 5.0, 100.0)
|
||||
t.deadReckon()
|
||||
check abs(t.current.x - 105.0) < 0.001, "dead reckon east: x += speed"
|
||||
check abs(t.current.y - 100.0) < 0.001, "dead reckon east: y unchanged"
|
||||
|
||||
block testHistoryWindow:
|
||||
var t = initEnemyTracker()
|
||||
# Feed 6 updates — history should hold last 5
|
||||
for i in 1 .. 6:
|
||||
t.update(float64(i) * 10.0, float64(i) * 20.0, 0.0, float64(i), 100.0)
|
||||
check t.historyCount == 5, "historyCount capped at 5"
|
||||
# history[0] should be the second-to-last scan (i=5)
|
||||
check abs(t.history[0].x - 50.0) < 0.001, "history[0].x == 50 (i=5)"
|
||||
check abs(t.history[4].x - 10.0) < 0.001, "history[4].x == 10 (i=1)"
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# StateVector tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
block testStateVectorLength:
|
||||
var t = initEnemyTracker()
|
||||
t.update(400.0, 300.0, 45.0, 3.0, 80.0)
|
||||
let bot = BotStateData(
|
||||
x: 200.0, y: 200.0, direction: 90.0, speed: 4.0, energy: 50.0,
|
||||
gunDirection: 180.0, gunHeat: 0.5,
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
let sv = buildStateVector(bot, t)
|
||||
check sv.shape == [57], "state vector has 57 elements"
|
||||
|
||||
block testStateVectorRange:
|
||||
var t = initEnemyTracker()
|
||||
t.update(400.0, 300.0, 180.0, 8.0, 100.0)
|
||||
let bot = BotStateData(
|
||||
x: 800.0, y: 600.0, direction: 360.0, speed: 8.0, energy: 100.0,
|
||||
gunDirection: 360.0, gunHeat: 1.8,
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
let sv = buildStateVector(bot, t)
|
||||
for i in 0 ..< 57:
|
||||
check sv[i] >= -2.0f32 and sv[i] <= 2.0f32,
|
||||
&"sv[{i}]={sv[i]} out of [-2,2] range"
|
||||
|
||||
block testWallDistances:
|
||||
# Bot at (100, 200) in 800×600 arena
|
||||
# wallMax = max(800, 600) = 800
|
||||
# top = (600 - 200) / 800 = 400/800 = 0.5
|
||||
# bottom = 200 / 800 = 0.25
|
||||
# left = 100 / 800 = 0.125
|
||||
# right = (800 - 100) / 800 = 700/800 = 0.875
|
||||
var t = initEnemyTracker()
|
||||
let bot = BotStateData(
|
||||
x: 100.0, y: 200.0, direction: 0.0, speed: 0.0, energy: 100.0,
|
||||
gunDirection: 0.0, gunHeat: 0.0,
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
let sv = buildStateVector(bot, t)
|
||||
check abs(sv[18] - 0.5f32) < 0.001f32, "top wall = 0.5"
|
||||
check abs(sv[19] - 0.25f32) < 0.001f32, "bottom wall = 0.25"
|
||||
check abs(sv[20] - 0.125f32) < 0.001f32, "left wall = 0.125"
|
||||
check abs(sv[21] - 0.875f32) < 0.001f32, "right wall = 0.875"
|
||||
|
||||
block testRelativeBearing:
|
||||
# Game convention: north=0°, CW. arctan2(dx,dy) used.
|
||||
# Bot at (0,0) dir=0°. Enemy at (0,100) → due north → absDir=0°.
|
||||
# relBearing = (0 - 0 + 540) mod 360 - 180 = 0°. normalized = 0/180 = 0.0
|
||||
var t = initEnemyTracker()
|
||||
t.update(0.0, 100.0, 0.0, 0.0, 100.0)
|
||||
let bot = BotStateData(
|
||||
x: 0.0, y: 0.0, direction: 0.0, speed: 0.0, energy: 100.0,
|
||||
gunDirection: 0.0, gunHeat: 0.0,
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
let sv = buildStateVector(bot, t)
|
||||
check abs(sv[16] - 0.0f32) < 0.01f32, "relative bearing = 0.0 (due north), got " & $sv[16]
|
||||
|
||||
block testRelativeBearingEast:
|
||||
# Enemy at (100,0) → due east → absDir=90°.
|
||||
# relBearing = (90 - 0 + 540) mod 360 - 180 = 90°. normalized = 90/180 = 0.5
|
||||
var t = initEnemyTracker()
|
||||
t.update(100.0, 0.0, 0.0, 0.0, 100.0)
|
||||
let bot = BotStateData(
|
||||
x: 0.0, y: 0.0, direction: 0.0, speed: 0.0, energy: 100.0,
|
||||
gunDirection: 0.0, gunHeat: 0.0,
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
let sv = buildStateVector(bot, t)
|
||||
check abs(sv[16] - 0.5f32) < 0.01f32, "relative bearing east = 0.5, got " & $sv[16]
|
||||
|
||||
block testHistoryPaddedWhenEmpty:
|
||||
var t = initEnemyTracker()
|
||||
let bot = BotStateData(
|
||||
x: 400.0, y: 300.0, direction: 0.0, speed: 0.0, energy: 100.0,
|
||||
gunDirection: 0.0, gunHeat: 0.0,
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
let sv = buildStateVector(bot, t)
|
||||
for i in 22 ..< 42:
|
||||
check sv[i] == 0.0f32, &"history slot {i} should be 0 when no contact"
|
||||
# indices 42-43 (goto inputs) default to 0 when not provided
|
||||
check sv[42] == 0.0f32, "sv[42] (remainingGotoDistance) should default to 0"
|
||||
check sv[43] == 0.0f32, "sv[43] (remainingGunAngle) should default to 0"
|
||||
# indices 44-55 (bullet slots) default to 0 when no bullets provided
|
||||
for i in 44 ..< 56:
|
||||
check sv[i] == 0.0f32, &"bullet slot {i} should be 0 when no bullets"
|
||||
# index 56 (staleness): no contact so ticksSinceLastScan=0 → 0/30 = 0
|
||||
check sv[56] == 0.0f32, "sv[56] staleness should be 0 when no contact"
|
||||
|
||||
block testBulletSlots:
|
||||
var t = initEnemyTracker()
|
||||
t.update(400.0, 300.0, 0.0, 0.0, 100.0) # enemy at (400,300)
|
||||
let bot = BotStateData(
|
||||
x: 200.0, y: 200.0, direction: 0.0, speed: 0.0, energy: 100.0,
|
||||
gunDirection: 0.0, gunHeat: 0.0,
|
||||
arenaWidth: 800.0, arenaHeight: 600.0,
|
||||
)
|
||||
# Bullet at (300,250), power 1.0 → speed = 20-3 = 17
|
||||
# relX = 300-200 = 100, relY = 250-200 = 50 (relative to bot, not enemy)
|
||||
# dist to bot = sqrt(100^2+50^2) ≈ 111.8, ticks ≈ 111.8/17 ≈ 6.58
|
||||
let b = BulletData(x: 300.0, y: 250.0, power: 1.0)
|
||||
let sv = buildStateVector(bot, t, 0.0, 0.0, [b], 1)
|
||||
check abs(sv[44] - (100.0/800.0).float32) < 0.001f32, "bullet relX"
|
||||
check abs(sv[45] - (50.0/600.0).float32) < 0.001f32, "bullet relY"
|
||||
check abs(sv[46] - (17.0/20.0).float32) < 0.001f32, "bullet speed norm"
|
||||
# second slot should be zero-padded
|
||||
for i in 48 ..< 57:
|
||||
check sv[i] == 0.0f32, &"unused bullet slot {i} should be 0"
|
||||
|
||||
echo "All tests passed"
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,175 @@
|
||||
## 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
|
||||
# + 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.01'f32) < 1e-6'f32, "computeTickReward(-2, -10) == 8.01, got " & $r
|
||||
|
||||
# ── computeRoundReward ────────────────────────────────────────────────────────
|
||||
|
||||
block testRoundReward:
|
||||
let r = computeRoundReward(350.0'f32)
|
||||
check abs(r - 7.0'f32) < 1e-6'f32, "computeRoundReward(350) == 7.0, got " & $r
|
||||
# bounded: long-battle cumulative scores must saturate, not blow the value scale
|
||||
check abs(computeRoundReward(89299.0'f32) - 8.0'f32) < 1e-6'f32,
|
||||
"computeRoundReward(89299) == 8.0 (capped), got " & $computeRoundReward(89299.0'f32)
|
||||
|
||||
# ── TrajectoryBuffer ──────────────────────────────────────────────────────────
|
||||
|
||||
block testBuffer:
|
||||
var buf = initTrajectoryBuffer()
|
||||
check buf.len == 0, "empty buffer len == 0"
|
||||
|
||||
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)
|
||||
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](STATE_DIM)
|
||||
let a = randomNormalTensor[float32](ACTION_DIM)
|
||||
let lp = ac.computeLogProb(s, a)
|
||||
let v = ac.criticForward(s)
|
||||
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)
|
||||
|
||||
# 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"
|
||||
|
||||
# ── 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"
|
||||
|
||||
# ── logStd ceiling: raw param must never drift above the collection clamp ─────
|
||||
# Regression for the train/collection std mismatch: logStd starting above the
|
||||
# ceiling (as in the warm-start snapshot) must be clamped back to [floor, ceiling]
|
||||
# by the first Adam step, so recomputed logP matches the acting policy's std.
|
||||
|
||||
block testLogStdCeilingClamp:
|
||||
randomize(45)
|
||||
var ac = initActorCritic()
|
||||
ac.logStd = newTensor[float32](ACTION_DIM).map(proc(v: float32): float32 = 3.0'f32)
|
||||
var buf = initTrajectoryBuffer()
|
||||
for _ in 0..<16:
|
||||
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.1'f32, value: 0.5'f32))
|
||||
var adam: ACAdamStates
|
||||
discard ppoUpdate(ac, buf, lastValue = 0.0'f32, adamStates = adam,
|
||||
epochs = 1, miniBatchSize = 16)
|
||||
for v in ac.logStd:
|
||||
check v <= logStdCeiling + 1e-6'f32, "logStd above ceiling after ppoUpdate"
|
||||
check v >= logStdFloor - 1e-6'f32, "logStd below floor after ppoUpdate"
|
||||
|
||||
echo "All tests passed"
|
||||
Executable
BIN
Binary file not shown.
@@ -0,0 +1,105 @@
|
||||
## test_weights.nim — assert-based tests for weights.nim
|
||||
## Run: nim c tests/test_weights.nim && ./tests/test_weights
|
||||
|
||||
import std/[os, math]
|
||||
import arraymancer
|
||||
import "../network"
|
||||
import "../training"
|
||||
import "../weights"
|
||||
|
||||
template check(cond: bool, msg: string) =
|
||||
if not cond:
|
||||
quit("FAIL: " & msg, 1)
|
||||
|
||||
const tmpBase = "/tmp/test_weights_nim"
|
||||
|
||||
# ── saveWeights / loadWeights roundtrip ───────────────────────────────────────
|
||||
|
||||
block testRoundtrip:
|
||||
let dir = tmpBase & "_roundtrip"
|
||||
removeDir(dir)
|
||||
|
||||
let ac1 = initActorCritic()
|
||||
saveWeights(ac1, dir)
|
||||
|
||||
var ac2 = initActorCritic()
|
||||
loadWeights(ac2, dir)
|
||||
|
||||
# Verify a sample of tensors
|
||||
template tensorEq(a, b: Tensor[float32]) =
|
||||
check a.shape == b.shape, "shape mismatch"
|
||||
let diff = abs(a - b)
|
||||
var maxDiff = 0.0'f32
|
||||
for v in diff: maxDiff = max(maxDiff, v)
|
||||
check maxDiff < 1e-6'f32, "tensor values differ by " & $maxDiff
|
||||
|
||||
tensorEq(ac1.actor.w1, ac2.actor.w1)
|
||||
tensorEq(ac1.actor.b1, ac2.actor.b1)
|
||||
tensorEq(ac1.actor.w3, ac2.actor.w3)
|
||||
tensorEq(ac1.critic.w1, ac2.critic.w1)
|
||||
tensorEq(ac1.critic.b3, ac2.critic.b3)
|
||||
tensorEq(ac1.logStd, ac2.logStd)
|
||||
|
||||
removeDir(dir)
|
||||
|
||||
# ── saveWeightsAtomic ─────────────────────────────────────────────────────────
|
||||
|
||||
block testAtomic:
|
||||
let dir = tmpBase & "_atomic"
|
||||
removeDir(dir)
|
||||
|
||||
let ac = initActorCritic()
|
||||
saveWeightsAtomic(ac, dir)
|
||||
|
||||
check dirExists(dir), "targetDir should exist after atomic save"
|
||||
for f in ["actor_w1.npy", "critic_w1.npy", "log_std.npy"]:
|
||||
check fileExists(dir / f), "missing file: " & f
|
||||
|
||||
removeDir(dir)
|
||||
|
||||
# ── loadBestAvailable ─────────────────────────────────────────────────────────
|
||||
|
||||
block testLoadBest:
|
||||
let root = tmpBase & "_loadbest"
|
||||
removeDir(root)
|
||||
createDir(root)
|
||||
|
||||
let ac0 = initActorCritic()
|
||||
|
||||
# Try with no weights — should return false
|
||||
var acEmpty = initActorCritic()
|
||||
var adamEmpty: ACAdamStates
|
||||
check not loadBestAvailable(acEmpty, adamEmpty, root).loaded, "should return false with no weights"
|
||||
|
||||
# Save to latest/; should load
|
||||
saveWeights(ac0, root / "latest")
|
||||
var ac1 = initActorCritic()
|
||||
var adam1: ACAdamStates
|
||||
check loadBestAvailable(ac1, adam1, root).loaded, "should load from latest/"
|
||||
|
||||
# Remove latest/, save to checkpoint_1/ — should fall back
|
||||
removeDir(root / "latest")
|
||||
saveWeights(ac0, root / "checkpoint_1")
|
||||
var ac2 = initActorCritic()
|
||||
var adam2: ACAdamStates
|
||||
check loadBestAvailable(ac2, adam2, root).loaded, "should load from checkpoint_1/"
|
||||
|
||||
removeDir(root)
|
||||
|
||||
# ── cleanStaleTempDirs ────────────────────────────────────────────────────────
|
||||
|
||||
block testClean:
|
||||
let root = tmpBase & "_clean"
|
||||
removeDir(root)
|
||||
createDir(root)
|
||||
|
||||
let stale = root / "latest_tmp_12345"
|
||||
createDir(stale)
|
||||
check dirExists(stale), "stale dir should exist before clean"
|
||||
|
||||
cleanStaleTempDirs(root)
|
||||
check not dirExists(stale), "stale dir should be gone after clean"
|
||||
|
||||
removeDir(root)
|
||||
|
||||
echo "All tests passed"
|
||||
Executable
BIN
Binary file not shown.
@@ -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)
|
||||
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 @@
|
||||
324564
|
||||
324564
|
||||
324564
|
||||
324564
|
||||
324564
|
||||
324564
|
||||
324564
|
||||
324564
|
||||
324564
|
||||
324564
|
||||
324564
|
||||
324564
|
||||
324564
|
||||
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 @@
|
||||
0
|
||||
@@ -0,0 +1,11 @@
|
||||
{
|
||||
"name": "SAC_LSTM_Bot",
|
||||
"version": "0.1.0",
|
||||
"authors": ["Davide Cappellini"],
|
||||
"description": "SAC+LSTM-trained Tank Royale bot (#37) — training/eval launch config",
|
||||
"homepage": "",
|
||||
"countryCodes": ["IT"],
|
||||
"gameTypes": ["classic", "melee", "1v1"],
|
||||
"platform": "Nim",
|
||||
"programmingLang": "Nim"
|
||||
}
|
||||
@@ -0,0 +1,13 @@
|
||||
# Package
|
||||
version = "0.1.0"
|
||||
author = "Davide Cappellini"
|
||||
description = "SAC+LSTM-trained Tank Royale bot"
|
||||
license = "MIT"
|
||||
srcDir = "src"
|
||||
bin = @["SAC_LSTM_Bot"]
|
||||
|
||||
# Dependencies
|
||||
requires "nim >= 2.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"
|
||||
Executable
+4
@@ -0,0 +1,4 @@
|
||||
#!/bin/sh
|
||||
# Launch config for tools/training_runner/RunTraining.java (#49): the runner
|
||||
# executes <json-basename>.sh inside the bot dir (sample-bots convention).
|
||||
exec "$(dirname "$0")/SAC_LSTM_Bot"
|
||||
@@ -0,0 +1,18 @@
|
||||
# Static-link OpenBLAS for portable deployment
|
||||
# ponytail: adjust path per machine, or use pkg-config
|
||||
switch("passL", "-L/nix/store/v07svn2y92bvzjl51aj7c9ca1cwg7rw7-openblas-0.3.32/lib -lopenblas")
|
||||
# libzip for weight checkpoint zip files
|
||||
# ponytail: nix store path; adjust per machine, or use pkg-config
|
||||
switch("passL", "-L/nix/store/wqvz31s598bvj3zb747943xhl38hjc6h-libzip-1.11.4/lib -lzip")
|
||||
switch("threads", "on")
|
||||
# Submodules import each other as SAC_LSTM_Bot/<mod>; make that resolvable for
|
||||
# the binary build too (tests already add ../src via tests/config.nims).
|
||||
switch("path", thisDir() & "/src")
|
||||
# begin Nimble config (version 2)
|
||||
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.
|
||||
# Must come AFTER the nimble.paths include: later --path wins the import search.
|
||||
switch("path", thisDir() & "/../libs/tankroyale_botapi")
|
||||
switch("path", thisDir() & "/../libs/radar_lock")
|
||||
@@ -0,0 +1,743 @@
|
||||
<svg xmlns="http://www.w3.org/2000/svg" width="1400" height="2000" viewBox="0 0 1400 2000" font-family="sans-serif">
|
||||
<rect width="1400" height="2000" fill="white"/>
|
||||
<text x="700" y="32" text-anchor="middle" font-size="21" font-weight="bold">SAC-LSTM campaign dashboard - live run (current only)</text>
|
||||
<text x="700" y="56" text-anchor="middle" font-size="12" fill="#555">generated 2026-08-23 10:46:52 - auto-reloads every 60 s (open this file in Chrome)</text>
|
||||
<text x="70" y="100" font-size="15" font-weight="bold">Test matches - win % vs opponents</text>
|
||||
<text x="70" y="115" font-size="11" fill="#555">raw dots = single test matches, thick = rolling-mean-10</text>
|
||||
<line x1="70" y1="720.0" x2="697" y2="720.0" stroke="#dddddd"/>
|
||||
<line x1="70" y1="600.0" x2="697" y2="600.0" stroke="#dddddd"/>
|
||||
<line x1="70" y1="480.0" x2="697" y2="480.0" stroke="#dddddd"/>
|
||||
<line x1="70" y1="360.0" x2="697" y2="360.0" stroke="#dddddd"/>
|
||||
<line x1="70" y1="240.0" x2="697" y2="240.0" stroke="#dddddd"/>
|
||||
<line x1="70" y1="120.0" x2="697" y2="120.0" stroke="#dddddd"/>
|
||||
<line x1="70" y1="720" x2="697" y2="720" stroke="black"/>
|
||||
<line x1="70" y1="720" x2="70" y2="120" stroke="black"/>
|
||||
<line x1="128.2" y1="720" x2="128.2" y2="724" stroke="black"/>
|
||||
<text x="128.2" y="737" text-anchor="middle" font-size="11">10</text>
|
||||
<line x1="192.8" y1="720" x2="192.8" y2="724" stroke="black"/>
|
||||
<text x="192.8" y="737" text-anchor="middle" font-size="11">20</text>
|
||||
<line x1="257.5" y1="720" x2="257.5" y2="724" stroke="black"/>
|
||||
<text x="257.5" y="737" text-anchor="middle" font-size="11">30</text>
|
||||
<line x1="322.1" y1="720" x2="322.1" y2="724" stroke="black"/>
|
||||
<text x="322.1" y="737" text-anchor="middle" font-size="11">40</text>
|
||||
<line x1="386.7" y1="720" x2="386.7" y2="724" stroke="black"/>
|
||||
<text x="386.7" y="737" text-anchor="middle" font-size="11">50</text>
|
||||
<line x1="451.4" y1="720" x2="451.4" y2="724" stroke="black"/>
|
||||
<text x="451.4" y="737" text-anchor="middle" font-size="11">60</text>
|
||||
<line x1="516.0" y1="720" x2="516.0" y2="724" stroke="black"/>
|
||||
<text x="516.0" y="737" text-anchor="middle" font-size="11">70</text>
|
||||
<line x1="580.6" y1="720" x2="580.6" y2="724" stroke="black"/>
|
||||
<text x="580.6" y="737" text-anchor="middle" font-size="11">80</text>
|
||||
<line x1="645.3" y1="720" x2="645.3" y2="724" stroke="black"/>
|
||||
<text x="645.3" y="737" text-anchor="middle" font-size="11">90</text>
|
||||
<line x1="66" y1="720.0" x2="70" y2="720.0" stroke="black"/>
|
||||
<text x="63" y="724.0" text-anchor="end" font-size="11">0</text>
|
||||
<line x1="66" y1="600.0" x2="70" y2="600.0" stroke="black"/>
|
||||
<text x="63" y="604.0" text-anchor="end" font-size="11">20</text>
|
||||
<line x1="66" y1="480.0" x2="70" y2="480.0" stroke="black"/>
|
||||
<text x="63" y="484.0" text-anchor="end" font-size="11">40</text>
|
||||
<line x1="66" y1="360.0" x2="70" y2="360.0" stroke="black"/>
|
||||
<text x="63" y="364.0" text-anchor="end" font-size="11">60</text>
|
||||
<line x1="66" y1="240.0" x2="70" y2="240.0" stroke="black"/>
|
||||
<text x="63" y="244.0" text-anchor="end" font-size="11">80</text>
|
||||
<line x1="66" y1="120.0" x2="70" y2="120.0" stroke="black"/>
|
||||
<text x="63" y="124.0" text-anchor="end" font-size="11">100</text>
|
||||
<text x="383" y="753" text-anchor="middle" font-size="12">test match number (each opponent)</text>
|
||||
<text x="16" y="420" text-anchor="middle" font-size="12" transform="rotate(-90 16 420)">win rate (%)</text>
|
||||
<circle cx="70.0" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="76.5" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="82.9" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="89.4" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="95.9" cy="540.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="102.3" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="108.8" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="115.2" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="121.7" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="128.2" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="134.6" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="141.1" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="147.6" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="154.0" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="160.5" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="167.0" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="173.4" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="179.9" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="186.4" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="192.8" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="199.3" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="205.7" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="212.2" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="218.7" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="225.1" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="231.6" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="238.1" cy="660.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="244.5" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="251.0" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="257.5" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="263.9" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="270.4" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="276.8" cy="660.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="283.3" cy="660.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="289.8" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="296.2" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="302.7" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="309.2" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="315.6" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="322.1" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="328.6" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="335.0" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="341.5" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="347.9" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="354.4" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="360.9" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="367.3" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="373.8" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="380.3" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="386.7" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="393.2" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="399.7" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="406.1" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="412.6" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="419.1" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="425.5" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="432.0" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="438.4" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="444.9" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="451.4" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="457.8" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="464.3" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="470.8" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="477.2" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="483.7" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="490.2" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="496.6" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="503.1" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="509.5" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="516.0" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="522.5" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="528.9" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="535.4" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="541.9" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="548.3" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="554.8" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="561.3" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="567.7" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="574.2" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="580.6" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="587.1" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="593.6" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="600.0" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="606.5" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="613.0" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="619.4" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="625.9" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="632.4" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="638.8" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="645.3" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="651.8" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="658.2" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="664.7" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="671.1" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="677.6" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="684.1" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="690.5" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<circle cx="697.0" cy="720.0" r="2" fill="#d62728" opacity="0.25"/>
|
||||
<polyline points="70.0,720.0 76.5,720.0 82.9,720.0 89.4,720.0 95.9,684.0 102.3,690.0 108.8,694.3 115.2,697.5 121.7,700.0 128.2,702.0 134.6,702.0 141.1,702.0 147.6,702.0 154.0,702.0 160.5,720.0 167.0,720.0 173.4,720.0 179.9,720.0 186.4,720.0 192.8,720.0 199.3,720.0 205.7,720.0 212.2,720.0 218.7,720.0 225.1,720.0 231.6,720.0 238.1,714.0 244.5,714.0 251.0,714.0 257.5,714.0 263.9,714.0 270.4,714.0 276.8,708.0 283.3,702.0 289.8,702.0 296.2,702.0 302.7,708.0 309.2,708.0 315.6,708.0 322.1,708.0 328.6,708.0 335.0,708.0 341.5,714.0 347.9,720.0 354.4,720.0 360.9,720.0 367.3,720.0 373.8,720.0 380.3,720.0 386.7,720.0 393.2,720.0 399.7,720.0 406.1,720.0 412.6,720.0 419.1,720.0 425.5,720.0 432.0,720.0 438.4,720.0 444.9,720.0 451.4,720.0 457.8,720.0 464.3,720.0 470.8,720.0 477.2,720.0 483.7,720.0 490.2,720.0 496.6,720.0 503.1,720.0 509.5,720.0 516.0,720.0 522.5,720.0 528.9,720.0 535.4,720.0 541.9,720.0 548.3,720.0 554.8,720.0 561.3,720.0 567.7,720.0 574.2,720.0 580.6,720.0 587.1,720.0 593.6,720.0 600.0,720.0 606.5,720.0 613.0,720.0 619.4,720.0 625.9,720.0 632.4,720.0 638.8,720.0 645.3,720.0 651.8,720.0 658.2,720.0 664.7,720.0 671.1,720.0 677.6,720.0 684.1,720.0 690.5,720.0 697.0,720.0" fill="none" stroke="#d62728" stroke-width="3.5" opacity="1.0"/>
|
||||
<circle cx="70.0" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="76.5" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="82.9" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="89.4" cy="180.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="95.9" cy="480.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="102.3" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="108.8" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="115.2" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="121.7" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="128.2" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="134.6" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="141.1" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="147.6" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="154.0" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="160.5" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="167.0" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="173.4" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="179.9" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="186.4" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="192.8" cy="360.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="199.3" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="205.7" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="212.2" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="218.7" cy="120.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="225.1" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="231.6" cy="600.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="238.1" cy="600.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="244.5" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="251.0" cy="120.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="257.5" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="263.9" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="270.4" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="276.8" cy="540.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="283.3" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="289.8" cy="600.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="296.2" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="302.7" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="309.2" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="315.6" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="322.1" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="328.6" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="335.0" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="341.5" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="347.9" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="354.4" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="360.9" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="367.3" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="373.8" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="380.3" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="386.7" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="393.2" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="399.7" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="406.1" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="412.6" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="419.1" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="425.5" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="432.0" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="438.4" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="444.9" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="451.4" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="457.8" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="464.3" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="470.8" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="477.2" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="483.7" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="490.2" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="496.6" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="503.1" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="509.5" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="516.0" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="522.5" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="528.9" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="535.4" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="541.9" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="548.3" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="554.8" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="561.3" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="567.7" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="574.2" cy="540.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="580.6" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="587.1" cy="600.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="593.6" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="600.0" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="606.5" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="613.0" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="619.4" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="625.9" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="632.4" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="638.8" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="645.3" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="651.8" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="658.2" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="664.7" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="671.1" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="677.6" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="684.1" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="690.5" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<circle cx="697.0" cy="720.0" r="2" fill="#1f77b4" opacity="0.25"/>
|
||||
<polyline points="70.0,720.0 76.5,720.0 82.9,720.0 89.4,585.0 95.9,564.0 102.3,590.0 108.8,608.6 115.2,622.5 121.7,633.3 128.2,642.0 134.6,642.0 141.1,642.0 147.6,642.0 154.0,696.0 160.5,720.0 167.0,720.0 173.4,720.0 179.9,720.0 186.4,720.0 192.8,684.0 199.3,684.0 205.7,684.0 212.2,684.0 218.7,624.0 225.1,624.0 231.6,612.0 238.1,600.0 244.5,600.0 251.0,540.0 257.5,576.0 263.9,576.0 270.4,576.0 276.8,558.0 283.3,618.0 289.8,606.0 296.2,618.0 302.7,630.0 309.2,630.0 315.6,690.0 322.1,690.0 328.6,690.0 335.0,690.0 341.5,708.0 347.9,708.0 354.4,720.0 360.9,720.0 367.3,720.0 373.8,720.0 380.3,720.0 386.7,720.0 393.2,720.0 399.7,720.0 406.1,720.0 412.6,720.0 419.1,720.0 425.5,720.0 432.0,720.0 438.4,720.0 444.9,720.0 451.4,720.0 457.8,720.0 464.3,720.0 470.8,720.0 477.2,720.0 483.7,720.0 490.2,720.0 496.6,720.0 503.1,720.0 509.5,720.0 516.0,720.0 522.5,720.0 528.9,720.0 535.4,720.0 541.9,720.0 548.3,720.0 554.8,720.0 561.3,720.0 567.7,720.0 574.2,702.0 580.6,702.0 587.1,690.0 593.6,690.0 600.0,690.0 606.5,690.0 613.0,690.0 619.4,690.0 625.9,690.0 632.4,690.0 638.8,708.0 645.3,708.0 651.8,720.0 658.2,720.0 664.7,720.0 671.1,720.0 677.6,720.0 684.1,720.0 690.5,720.0 697.0,720.0" fill="none" stroke="#1f77b4" stroke-width="3.5" opacity="1.0"/>
|
||||
<circle cx="70.0" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="76.5" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="82.9" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="89.4" cy="600.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="95.9" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="102.3" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="108.8" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="115.2" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="121.7" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="128.2" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="134.6" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="141.1" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="147.6" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="154.0" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="160.5" cy="660.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="167.0" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="173.4" cy="240.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="179.9" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="186.4" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="192.8" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="199.3" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="205.7" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="212.2" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="218.7" cy="660.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="225.1" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="231.6" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="238.1" cy="360.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="244.5" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="251.0" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="257.5" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="263.9" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="270.4" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="276.8" cy="660.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="283.3" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="289.8" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="296.2" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="302.7" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="309.2" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="315.6" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="322.1" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="328.6" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="335.0" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="341.5" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="347.9" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="354.4" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="360.9" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="367.3" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="373.8" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="380.3" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="386.7" cy="660.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="393.2" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="399.7" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="406.1" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="412.6" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="419.1" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="425.5" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="432.0" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="438.4" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="444.9" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="451.4" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="457.8" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="464.3" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="470.8" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="477.2" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="483.7" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="490.2" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="496.6" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="503.1" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="509.5" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="516.0" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="522.5" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="528.9" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="535.4" cy="660.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="541.9" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="548.3" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="554.8" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="561.3" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="567.7" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="574.2" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="580.6" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="587.1" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="593.6" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="600.0" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="606.5" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="613.0" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="619.4" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="625.9" cy="660.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="632.4" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="638.8" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="645.3" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="651.8" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="658.2" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="664.7" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="671.1" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="677.6" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="684.1" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="690.5" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<circle cx="697.0" cy="720.0" r="2" fill="#2ca02c" opacity="0.25"/>
|
||||
<polyline points="70.0,720.0 76.5,720.0 82.9,720.0 89.4,690.0 95.9,696.0 102.3,700.0 108.8,702.9 115.2,705.0 121.7,706.7 128.2,708.0 134.6,708.0 141.1,708.0 147.6,708.0 154.0,720.0 160.5,714.0 167.0,714.0 173.4,666.0 179.9,666.0 186.4,666.0 192.8,666.0 199.3,666.0 205.7,666.0 212.2,666.0 218.7,660.0 225.1,666.0 231.6,666.0 238.1,678.0 244.5,678.0 251.0,678.0 257.5,678.0 263.9,678.0 270.4,678.0 276.8,672.0 283.3,678.0 289.8,678.0 296.2,678.0 302.7,714.0 309.2,714.0 315.6,714.0 322.1,714.0 328.6,714.0 335.0,714.0 341.5,720.0 347.9,720.0 354.4,720.0 360.9,720.0 367.3,720.0 373.8,720.0 380.3,720.0 386.7,714.0 393.2,714.0 399.7,714.0 406.1,714.0 412.6,714.0 419.1,714.0 425.5,714.0 432.0,714.0 438.4,714.0 444.9,714.0 451.4,720.0 457.8,720.0 464.3,720.0 470.8,720.0 477.2,720.0 483.7,720.0 490.2,720.0 496.6,720.0 503.1,720.0 509.5,720.0 516.0,720.0 522.5,720.0 528.9,720.0 535.4,714.0 541.9,714.0 548.3,714.0 554.8,714.0 561.3,714.0 567.7,714.0 574.2,714.0 580.6,714.0 587.1,714.0 593.6,714.0 600.0,720.0 606.5,720.0 613.0,720.0 619.4,720.0 625.9,714.0 632.4,714.0 638.8,714.0 645.3,714.0 651.8,714.0 658.2,714.0 664.7,714.0 671.1,714.0 677.6,714.0 684.1,714.0 690.5,720.0 697.0,720.0" fill="none" stroke="#2ca02c" stroke-width="3.5" opacity="1.0"/>
|
||||
<line x1="82" y1="772" x2="110" y2="772" stroke="#d62728" stroke-width="3"/>
|
||||
<text x="116" y="776" font-size="12">Corners - 98 evals</text>
|
||||
<line x1="82" y1="790" x2="110" y2="790" stroke="#1f77b4" stroke-width="3"/>
|
||||
<text x="116" y="794" font-size="12">Crazy - 98 evals</text>
|
||||
<line x1="82" y1="808" x2="110" y2="808" stroke="#2ca02c" stroke-width="3"/>
|
||||
<text x="116" y="812" font-size="12">Target - 98 evals</text>
|
||||
<text x="747" y="100" font-size="15" font-weight="bold">Real fights - win % per opponent</text>
|
||||
<text x="747" y="115" font-size="11" fill="#555">training_log.jsonl only - learning in REAL battles, not tests</text>
|
||||
<line x1="747" y1="720.0" x2="1375" y2="720.0" stroke="#dddddd"/>
|
||||
<line x1="747" y1="600.0" x2="1375" y2="600.0" stroke="#dddddd"/>
|
||||
<line x1="747" y1="480.0" x2="1375" y2="480.0" stroke="#dddddd"/>
|
||||
<line x1="747" y1="360.0" x2="1375" y2="360.0" stroke="#dddddd"/>
|
||||
<line x1="747" y1="240.0" x2="1375" y2="240.0" stroke="#dddddd"/>
|
||||
<line x1="747" y1="120.0" x2="1375" y2="120.0" stroke="#dddddd"/>
|
||||
<line x1="747" y1="720" x2="1375" y2="720" stroke="black"/>
|
||||
<line x1="747" y1="720" x2="747" y2="120" stroke="black"/>
|
||||
<line x1="840.9" y1="720" x2="840.9" y2="724" stroke="black"/>
|
||||
<text x="840.9" y="737" text-anchor="middle" font-size="11">300</text>
|
||||
<line x1="935.2" y1="720" x2="935.2" y2="724" stroke="black"/>
|
||||
<text x="935.2" y="737" text-anchor="middle" font-size="11">600</text>
|
||||
<line x1="1029.4" y1="720" x2="1029.4" y2="724" stroke="black"/>
|
||||
<text x="1029.4" y="737" text-anchor="middle" font-size="11">900</text>
|
||||
<line x1="1123.7" y1="720" x2="1123.7" y2="724" stroke="black"/>
|
||||
<text x="1123.7" y="737" text-anchor="middle" font-size="11">1200</text>
|
||||
<line x1="1217.9" y1="720" x2="1217.9" y2="724" stroke="black"/>
|
||||
<text x="1217.9" y="737" text-anchor="middle" font-size="11">1500</text>
|
||||
<line x1="1312.2" y1="720" x2="1312.2" y2="724" stroke="black"/>
|
||||
<text x="1312.2" y="737" text-anchor="middle" font-size="11">1800</text>
|
||||
<line x1="743" y1="720.0" x2="747" y2="720.0" stroke="black"/>
|
||||
<text x="740" y="724.0" text-anchor="end" font-size="11">0</text>
|
||||
<line x1="743" y1="600.0" x2="747" y2="600.0" stroke="black"/>
|
||||
<text x="740" y="604.0" text-anchor="end" font-size="11">20</text>
|
||||
<line x1="743" y1="480.0" x2="747" y2="480.0" stroke="black"/>
|
||||
<text x="740" y="484.0" text-anchor="end" font-size="11">40</text>
|
||||
<line x1="743" y1="360.0" x2="747" y2="360.0" stroke="black"/>
|
||||
<text x="740" y="364.0" text-anchor="end" font-size="11">60</text>
|
||||
<line x1="743" y1="240.0" x2="747" y2="240.0" stroke="black"/>
|
||||
<text x="740" y="244.0" text-anchor="end" font-size="11">80</text>
|
||||
<line x1="743" y1="120.0" x2="747" y2="120.0" stroke="black"/>
|
||||
<text x="740" y="124.0" text-anchor="end" font-size="11">100</text>
|
||||
<text x="1061" y="753" text-anchor="middle" font-size="12">game number (100-game buckets)</text>
|
||||
<text x="16" y="420" text-anchor="middle" font-size="12" transform="rotate(-90 16 420)">win %</text>
|
||||
<polyline points="762.4,720.0 793.8,720.0 825.2,720.0 856.6,720.0 888.1,720.0 919.5,720.0 950.9,720.0 982.3,720.0 1013.7,720.0 1045.1,720.0 1076.6,720.0 1108.0,720.0 1139.4,720.0 1170.8,720.0 1202.2,720.0 1233.6,720.0 1265.0,720.0 1296.5,720.0 1327.9,700.0 1359.3,720.0" fill="none" stroke="#1f77b4" stroke-width="3.5" opacity="1.0"/>
|
||||
<polyline points="825.2,720.0 856.6,720.0 888.1,720.0 919.5,720.0 950.9,720.0 982.3,720.0 1013.7,720.0 1045.1,720.0" fill="none" stroke="#ff7f0e" stroke-width="3.5" opacity="1.0"/>
|
||||
<polyline points="1108.0,720.0 1139.4,720.0 1170.8,720.0" fill="none" stroke="#ff7f0e" stroke-width="3.5" opacity="1.0"/>
|
||||
<polyline points="1233.6,720.0 1265.0,720.0 1296.5,720.0 1327.9,720.0 1359.3,720.0" fill="none" stroke="#ff7f0e" stroke-width="3.5" opacity="1.0"/>
|
||||
<polyline points="762.4,720.0 793.8,720.0" fill="none" stroke="#2ca02c" stroke-width="3.5" opacity="1.0"/>
|
||||
<polyline points="888.1,720.0 919.5,720.0 950.9,720.0 982.3,720.0 1013.7,720.0 1045.1,720.0" fill="none" stroke="#2ca02c" stroke-width="3.5" opacity="1.0"/>
|
||||
<polyline points="1202.2,690.0 1233.6,720.0" fill="none" stroke="#2ca02c" stroke-width="3.5" opacity="1.0"/>
|
||||
<polyline points="1296.5,720.0 1327.9,720.0 1359.3,720.0" fill="none" stroke="#2ca02c" stroke-width="3.5" opacity="1.0"/>
|
||||
<polyline points="762.4,720.0 793.8,720.0 825.2,720.0 856.6,720.0 888.1,720.0 919.5,720.0 950.9,660.0 982.3,720.0 1013.7,720.0" fill="none" stroke="#9467bd" stroke-width="3.5" opacity="1.0"/>
|
||||
<polyline points="1139.4,720.0 1170.8,720.0" fill="none" stroke="#9467bd" stroke-width="3.5" opacity="1.0"/>
|
||||
<polyline points="762.4,720.0 793.8,720.0 825.2,720.0 856.6,720.0 888.1,720.0 919.5,720.0 950.9,720.0 982.3,720.0 1013.7,720.0 1045.1,720.0 1076.6,720.0 1108.0,720.0 1139.4,720.0 1170.8,720.0 1202.2,720.0 1233.6,720.0 1265.0,720.0 1296.5,720.0 1327.9,720.0 1359.3,720.0" fill="none" stroke="#d62728" stroke-width="3.5" opacity="1.0"/>
|
||||
<line x1="759" y1="772" x2="787" y2="772" stroke="#1f77b4" stroke-width="3"/>
|
||||
<text x="793" y="776" font-size="12">Crazy (470 games)</text>
|
||||
<line x1="759" y1="790" x2="787" y2="790" stroke="#ff7f0e" stroke-width="3"/>
|
||||
<text x="793" y="794" font-size="12">RamFire (450 games)</text>
|
||||
<line x1="759" y1="808" x2="787" y2="808" stroke="#2ca02c" stroke-width="3"/>
|
||||
<text x="793" y="812" font-size="12">Target (220 games)</text>
|
||||
<line x1="759" y1="826" x2="787" y2="826" stroke="#9467bd" stroke-width="3"/>
|
||||
<text x="793" y="830" font-size="12">SacTwin (190 games)</text>
|
||||
<line x1="759" y1="844" x2="787" y2="844" stroke="#d62728" stroke-width="3"/>
|
||||
<text x="793" y="848" font-size="12">Corners (670 games)</text>
|
||||
<text x="70" y="940" font-size="15" font-weight="bold">Training losses (log scale)</text>
|
||||
<text x="70" y="955" font-size="11" fill="#555">training_metrics.jsonl - big early spikes are normal</text>
|
||||
<line x1="70" y1="1428.0" x2="697" y2="1428.0" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1401.9" x2="697" y2="1401.9" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1375.8" x2="697" y2="1375.8" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1349.7" x2="697" y2="1349.7" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1323.6" x2="697" y2="1323.6" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1297.4" x2="697" y2="1297.4" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1271.3" x2="697" y2="1271.3" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1245.2" x2="697" y2="1245.2" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1219.1" x2="697" y2="1219.1" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1193.0" x2="697" y2="1193.0" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1166.9" x2="697" y2="1166.9" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1140.8" x2="697" y2="1140.8" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1114.7" x2="697" y2="1114.7" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1088.6" x2="697" y2="1088.6" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1062.4" x2="697" y2="1062.4" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1036.3" x2="697" y2="1036.3" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1010.2" x2="697" y2="1010.2" stroke="#dddddd"/>
|
||||
<line x1="70" y1="984.1" x2="697" y2="984.1" stroke="#dddddd"/>
|
||||
<line x1="70" y1="958.0" x2="697" y2="958.0" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1428" x2="697" y2="1428" stroke="black"/>
|
||||
<line x1="70" y1="1428" x2="70" y2="958" stroke="black"/>
|
||||
<line x1="70.0" y1="1428" x2="70.0" y2="1432" stroke="black"/>
|
||||
<text x="70.0" y="1445" text-anchor="middle" font-size="11">1</text>
|
||||
<line x1="226.8" y1="1428" x2="226.8" y2="1432" stroke="black"/>
|
||||
<text x="226.8" y="1445" text-anchor="middle" font-size="11">51</text>
|
||||
<line x1="383.5" y1="1428" x2="383.5" y2="1432" stroke="black"/>
|
||||
<text x="383.5" y="1445" text-anchor="middle" font-size="11">101</text>
|
||||
<line x1="540.2" y1="1428" x2="540.2" y2="1432" stroke="black"/>
|
||||
<text x="540.2" y="1445" text-anchor="middle" font-size="11">151</text>
|
||||
<line x1="697.0" y1="1428" x2="697.0" y2="1432" stroke="black"/>
|
||||
<text x="697.0" y="1445" text-anchor="middle" font-size="11">201</text>
|
||||
<line x1="66" y1="1428.0" x2="70" y2="1428.0" stroke="black"/>
|
||||
<text x="63" y="1432.0" text-anchor="end" font-size="11">0.1</text>
|
||||
<line x1="66" y1="1310.5" x2="70" y2="1310.5" stroke="black"/>
|
||||
<text x="63" y="1314.5" text-anchor="end" font-size="11">3.16e+03</text>
|
||||
<line x1="66" y1="1193.0" x2="70" y2="1193.0" stroke="black"/>
|
||||
<text x="63" y="1197.0" text-anchor="end" font-size="11">1e+08</text>
|
||||
<line x1="66" y1="1075.5" x2="70" y2="1075.5" stroke="black"/>
|
||||
<text x="63" y="1079.5" text-anchor="end" font-size="11">3.16e+12</text>
|
||||
<line x1="66" y1="958.0" x2="70" y2="958.0" stroke="black"/>
|
||||
<text x="63" y="962.0" text-anchor="end" font-size="11">1e+17</text>
|
||||
<text x="383" y="1461" text-anchor="middle" font-size="12">metric line number</text>
|
||||
<text x="16" y="1193" text-anchor="middle" font-size="12" transform="rotate(-90 16 1193)">loss (log)</text>
|
||||
<polyline points="70.0,1366.3 73.1,1368.7 76.3,1365.5 79.4,1360.9 82.5,1355.0 85.7,1351.6 88.8,1353.9 91.9,1357.7 95.1,979.1 98.2,1370.9 101.3,1393.2 104.5,1379.4 107.6,1366.7 110.8,1357.6 113.9,1351.5 117.0,1349.9 120.2,1348.1 123.3,1347.8 126.4,1345.5 129.6,1346.3 132.7,1348.6 135.8,1350.0 139.0,1103.1 142.1,1359.7 145.2,1364.0 148.4,1366.7 151.5,1064.5 154.6,1382.5 157.8,1382.4 160.9,1385.5 164.1,1383.9 167.2,1386.3 170.3,1390.4 173.5,1392.9 176.6,1392.1 179.7,1392.5 182.9,1404.7 186.0,1403.6 189.1,1397.8 192.3,1409.3 195.4,1415.2 198.5,979.1 201.7,979.1 204.8,1391.7 207.9,1385.3 211.1,1375.4 214.2,1372.1 217.3,979.1 220.5,1364.5 223.6,1359.7 226.8,1101.7 229.9,1357.2 233.0,1354.6 236.2,1363.0 239.3,1377.5 242.4,1390.2 245.6,1375.3 248.7,979.1 251.8,979.1 255.0,1353.1 258.1,1350.3 261.2,1352.5 264.4,1352.6 267.5,979.1 270.6,1352.6 273.8,1350.5 276.9,1354.6 280.0,1359.8 283.2,1355.3 286.3,1359.3 289.4,1354.0 292.6,1362.1 295.7,1365.6 298.9,1367.3 302.0,1374.4 305.1,1397.8 308.3,1104.4 311.4,1375.4 314.5,1261.1 317.7,1156.2 320.8,1076.6 323.9,1059.3 327.1,1058.7 330.2,1058.0 333.3,1057.7 336.5,1057.4 339.6,1057.2 342.7,1057.4 345.9,1056.5 349.0,1056.0 352.1,1055.7 355.3,1055.3 358.4,1054.6 361.6,1054.1 364.7,979.2 367.8,979.2 371.0,979.2 374.1,979.3 377.2,1052.4 380.4,1051.9 383.5,1051.6 386.6,1051.2 389.8,1050.7 392.9,1050.2 396.0,1049.7 399.2,1049.3 402.3,1048.8 405.4,1048.3 408.6,1047.8 411.7,1046.9 414.9,1046.4 418.0,1045.9 421.1,1045.5 424.3,1045.0 427.4,1044.6 430.5,1044.1 433.7,1043.6 436.8,1043.1 439.9,979.4 443.1,1042.3 446.2,1042.0 449.3,1041.6 452.5,1041.3 455.6,1040.9 458.7,1040.2 461.9,1039.9 465.0,1039.6 468.1,1039.2 471.3,1038.9 474.4,1038.5 477.6,1038.9 480.7,1037.9 483.8,1037.4 487.0,979.4 490.1,1036.9 493.2,1036.5 496.4,1036.0 499.5,979.4 502.6,1035.4 505.8,1035.1 508.9,1034.7 512.0,1034.4 515.2,1034.0 518.3,1033.7 521.4,1033.2 524.6,1032.8 527.7,979.5 530.8,1032.0 534.0,979.5 537.1,1031.3 540.2,1030.9 543.4,1030.5 546.5,1030.2 549.7,1029.9 552.8,1029.6 555.9,1029.4 559.1,1029.1 562.2,1028.7 565.3,1028.3 568.5,1028.0 571.6,1028.1 574.7,1027.3 577.9,1026.9 581.0,979.6 584.1,1026.2 587.3,1025.6 590.4,1025.3 593.5,1025.0 596.7,979.6 599.8,979.6 603.0,1023.8 606.1,1023.2 609.2,1022.9 612.4,1022.6 615.5,1022.3 618.6,1022.0 621.8,1021.7 624.9,1021.4 628.0,1021.1 631.2,1020.8 634.3,1020.5 637.4,1020.2 640.6,1019.9 643.7,1019.3 646.8,1019.1 650.0,1018.8 653.1,1018.6 656.2,1018.5 659.4,1018.0 662.5,1017.5 665.6,1017.3 668.8,1017.0 671.9,1016.7 675.1,1016.5 678.2,1016.2 681.3,1016.0 684.5,1015.8 687.6,1015.5 690.7,1015.3 693.9,979.8 697.0,1014.8" fill="none" stroke="#1f77b4" stroke-width="1.8" opacity="1.0"/>
|
||||
<polyline points="70.0,1388.6 73.1,1384.3 76.3,1382.4 79.4,1380.1 82.5,1377.2 85.7,1375.8 88.8,1376.7 91.9,1378.2 95.1,1381.7 98.2,1384.2 101.3,1390.4 104.5,1394.2 107.6,1384.3 110.8,1379.2 113.9,1376.0 117.0,1373.9 120.2,1373.1 123.3,1372.6 126.4,1371.7 129.6,1371.0 132.7,1371.2 135.8,1371.3 139.0,1371.8 142.1,1372.0 145.2,1373.0 148.4,1373.4 151.5,1375.0 154.6,1375.2 157.8,1377.9 160.9,1379.3 164.1,1380.5 167.2,1381.0 170.3,1380.8 173.5,1380.4 176.6,1382.1 179.7,1382.3 182.9,1380.9 186.0,1380.5 189.1,1382.1 192.3,1379.9 195.4,1379.2 198.5,1378.0 201.7,1375.1 204.8,1372.7 207.9,1371.7 211.1,1370.2 214.2,1369.3 217.3,1367.8 220.5,1366.8 223.6,1365.7 226.8,1364.8 229.9,1364.0 233.0,1363.3 236.2,1363.7 239.3,1363.9 242.4,1364.5 245.6,1365.6 248.7,1367.1 251.8,1369.1 255.0,1370.0 258.1,1371.0 261.2,1371.8 264.4,1371.9 267.5,1371.2 270.6,1370.3 273.8,1371.1 276.9,1370.3 280.0,1369.3 283.2,1370.3 286.3,1370.0 289.4,1371.0 292.6,1369.7 295.7,1369.3 298.9,1369.5 302.0,1369.1 305.1,1367.5 308.3,1365.0 311.4,1364.7 314.5,1332.0 317.7,1278.9 320.8,1239.1 323.9,1230.5 327.1,1230.2 330.2,1229.8 333.3,1229.7 336.5,1229.5 339.6,1229.4 342.7,1229.2 345.9,1229.1 349.0,1228.8 352.1,1228.7 355.3,1228.5 358.4,1228.1 361.6,1227.9 364.7,1227.8 367.8,1227.6 371.0,1227.5 374.1,1227.3 377.2,1227.0 380.4,1226.8 383.5,1226.6 386.6,1226.4 389.8,1226.2 392.9,1225.9 396.0,1225.7 399.2,1225.5 402.3,1225.2 405.4,1225.0 408.6,1224.7 411.7,1224.3 414.9,1224.0 418.0,1223.8 421.1,1223.6 424.3,1223.3 427.4,1223.1 430.5,1222.9 433.7,1222.6 436.8,1222.4 439.9,1222.2 443.1,1222.0 446.2,1221.8 449.3,1221.7 452.5,1221.5 455.6,1221.3 458.7,1220.9 461.9,1220.8 465.0,1220.6 468.1,1220.5 471.3,1220.3 474.4,1220.1 477.6,1219.9 480.7,1219.8 483.8,1219.5 487.0,1219.4 490.1,1219.3 493.2,1219.1 496.4,1218.8 499.5,1218.6 502.6,1218.5 505.8,1218.4 508.9,1218.2 512.0,1218.0 515.2,1217.8 518.3,1217.7 521.4,1217.5 524.6,1217.3 527.7,1217.0 530.8,1216.8 534.0,1216.7 537.1,1216.5 540.2,1216.3 543.4,1216.1 546.5,1215.9 549.7,1215.8 552.8,1215.6 555.9,1215.5 559.1,1215.4 562.2,1215.2 565.3,1215.0 568.5,1214.8 571.6,1214.6 574.7,1214.5 577.9,1214.3 581.0,1214.1 584.1,1213.8 587.3,1213.6 590.4,1213.5 593.5,1213.3 596.7,1213.1 599.8,1213.0 603.0,1212.8 606.1,1212.4 609.2,1212.3 612.4,1212.1 615.5,1212.0 618.6,1211.8 621.8,1211.7 624.9,1211.5 628.0,1211.4 631.2,1211.2 634.3,1211.1 637.4,1210.9 640.6,1210.8 643.7,1210.5 646.8,1210.4 650.0,1210.2 653.1,1210.1 656.2,1210.0 659.4,1209.8 662.5,1209.6 665.6,1209.5 668.8,1209.3 671.9,1209.2 675.1,1209.1 678.2,1208.9 681.3,1208.8 684.5,1208.7 687.6,1208.6 690.7,1208.5 693.9,1208.3 697.0,1208.2" fill="none" stroke="#ff7f0e" stroke-width="1.8" opacity="1.0"/>
|
||||
<line x1="82" y1="1480" x2="110" y2="1480" stroke="#1f77b4" stroke-width="3"/>
|
||||
<text x="116" y="1484" font-size="12">critic_loss</text>
|
||||
<line x1="82" y1="1498" x2="110" y2="1498" stroke="#ff7f0e" stroke-width="3"/>
|
||||
<text x="116" y="1502" font-size="12">|actor_loss|</text>
|
||||
<text x="747" y="940" font-size="15" font-weight="bold">Alpha temperature</text>
|
||||
<text x="747" y="955" font-size="11" fill="#555">training_metrics.jsonl - high = exploring, low = exploiting</text>
|
||||
<line x1="747" y1="1428.0" x2="1375" y2="1428.0" stroke="#dddddd"/>
|
||||
<line x1="747" y1="1310.5" x2="1375" y2="1310.5" stroke="#dddddd"/>
|
||||
<line x1="747" y1="1193.0" x2="1375" y2="1193.0" stroke="#dddddd"/>
|
||||
<line x1="747" y1="1075.5" x2="1375" y2="1075.5" stroke="#dddddd"/>
|
||||
<line x1="747" y1="958.0" x2="1375" y2="958.0" stroke="#dddddd"/>
|
||||
<line x1="747" y1="1428" x2="1375" y2="1428" stroke="black"/>
|
||||
<line x1="747" y1="1428" x2="747" y2="958" stroke="black"/>
|
||||
<line x1="747.0" y1="1428" x2="747.0" y2="1432" stroke="black"/>
|
||||
<text x="747.0" y="1445" text-anchor="middle" font-size="11">1</text>
|
||||
<line x1="904.0" y1="1428" x2="904.0" y2="1432" stroke="black"/>
|
||||
<text x="904.0" y="1445" text-anchor="middle" font-size="11">51</text>
|
||||
<line x1="1061.0" y1="1428" x2="1061.0" y2="1432" stroke="black"/>
|
||||
<text x="1061.0" y="1445" text-anchor="middle" font-size="11">101</text>
|
||||
<line x1="1218.0" y1="1428" x2="1218.0" y2="1432" stroke="black"/>
|
||||
<text x="1218.0" y="1445" text-anchor="middle" font-size="11">151</text>
|
||||
<line x1="1375.0" y1="1428" x2="1375.0" y2="1432" stroke="black"/>
|
||||
<text x="1375.0" y="1445" text-anchor="middle" font-size="11">201</text>
|
||||
<line x1="743" y1="1428.0" x2="747" y2="1428.0" stroke="black"/>
|
||||
<text x="740" y="1432.0" text-anchor="end" font-size="11">0</text>
|
||||
<line x1="743" y1="1310.5" x2="747" y2="1310.5" stroke="black"/>
|
||||
<text x="740" y="1314.5" text-anchor="end" font-size="11">0.252</text>
|
||||
<line x1="743" y1="1193.0" x2="747" y2="1193.0" stroke="black"/>
|
||||
<text x="740" y="1197.0" text-anchor="end" font-size="11">0.504</text>
|
||||
<line x1="743" y1="1075.5" x2="747" y2="1075.5" stroke="black"/>
|
||||
<text x="740" y="1079.5" text-anchor="end" font-size="11">0.756</text>
|
||||
<line x1="743" y1="958.0" x2="747" y2="958.0" stroke="black"/>
|
||||
<text x="740" y="962.0" text-anchor="end" font-size="11">1.01</text>
|
||||
<text x="1061" y="1461" text-anchor="middle" font-size="12">metric line number</text>
|
||||
<text x="16" y="1193" text-anchor="middle" font-size="12" transform="rotate(-90 16 1193)">alpha</text>
|
||||
<polyline points="747.0,961.7 750.1,962.3 753.3,962.6 756.4,963.0 759.6,963.6 762.7,964.5 765.8,964.8 769.0,965.1 772.1,965.6 775.3,965.9 778.4,966.5 781.5,967.4 784.7,967.6 787.8,967.0 791.0,966.5 794.1,965.9 797.2,965.4 800.4,964.8 803.5,964.1 806.7,963.6 809.8,963.0 812.9,962.5 816.1,961.9 819.2,961.3 822.4,960.8 825.5,960.5 828.6,959.6 831.8,959.3 834.9,959.8 838.1,960.1 841.2,960.5 844.3,960.8 847.5,961.1 850.6,961.5 853.8,962.0 856.9,962.5 860.0,963.3 863.2,963.9 866.3,964.4 869.5,965.0 872.6,965.5 875.7,966.0 878.9,966.3 882.0,965.9 885.2,965.5 888.3,965.2 891.4,964.9 894.6,964.4 897.7,964.0 900.9,963.5 904.0,963.0 907.1,962.4 910.3,962.0 913.4,961.6 916.6,960.6 919.7,960.1 922.8,959.5 926.0,958.5 929.1,958.0 932.3,958.4 935.4,958.8 938.5,959.4 941.7,960.0 944.8,960.5 948.0,961.5 951.1,962.1 954.2,962.6 957.4,963.2 960.5,963.7 963.7,964.0 966.8,964.6 969.9,965.1 973.1,965.7 976.2,965.9 979.4,965.9 982.5,965.4 985.6,964.3 988.8,963.7 991.9,963.5 995.1,964.1 998.2,964.6 1001.3,964.9 1004.5,965.6 1007.6,966.4 1010.8,966.7 1013.9,967.0 1017.0,967.2 1020.2,967.8 1023.3,968.1 1026.5,968.6 1029.6,969.2 1032.7,969.7 1035.9,970.4 1039.0,971.0 1042.2,971.2 1045.3,971.5 1048.4,971.8 1051.6,972.3 1054.7,972.9 1057.9,973.4 1061.0,973.7 1064.1,974.2 1067.3,974.8 1070.4,975.3 1073.6,975.9 1076.7,976.4 1079.8,977.0 1083.0,977.5 1086.1,978.0 1089.3,979.0 1092.4,979.5 1095.5,980.1 1098.7,980.5 1101.8,981.0 1105.0,981.4 1108.1,981.9 1111.2,982.5 1114.4,983.0 1117.5,983.5 1120.7,983.9 1123.8,984.3 1126.9,984.7 1130.1,985.3 1133.2,985.8 1136.4,986.7 1139.5,987.3 1142.6,987.7 1145.8,988.2 1148.9,988.7 1152.1,989.2 1155.2,989.8 1158.3,990.2 1161.5,990.9 1164.6,991.3 1167.8,991.7 1170.9,992.3 1174.0,993.0 1177.2,993.6 1180.3,994.0 1183.5,994.3 1186.6,994.9 1189.7,995.4 1192.9,995.9 1196.0,996.4 1199.2,996.9 1202.3,997.5 1205.4,998.0 1208.6,998.5 1211.7,999.0 1214.9,999.5 1218.0,1000.0 1221.1,1000.5 1224.3,1000.9 1227.4,1001.3 1230.6,1001.7 1233.7,1002.1 1236.8,1002.5 1240.0,1003.0 1243.1,1003.5 1246.3,1004.0 1249.4,1004.5 1252.5,1005.0 1255.7,1005.5 1258.8,1006.0 1262.0,1006.9 1265.1,1007.4 1268.2,1007.9 1271.4,1008.4 1274.5,1008.9 1277.7,1009.4 1280.8,1010.1 1283.9,1011.1 1287.1,1011.6 1290.2,1012.1 1293.4,1012.6 1296.5,1013.1 1299.6,1013.6 1302.8,1014.1 1305.9,1014.5 1309.1,1015.0 1312.2,1015.7 1315.3,1016.2 1318.5,1016.7 1321.6,1017.6 1324.8,1018.0 1327.9,1018.5 1331.0,1019.0 1334.2,1019.5 1337.3,1020.0 1340.5,1020.8 1343.6,1021.3 1346.7,1021.8 1349.9,1022.3 1353.0,1022.8 1356.2,1023.3 1359.3,1023.6 1362.4,1024.1 1365.6,1024.6 1368.7,1025.1 1371.9,1025.6 1375.0,1026.0" fill="none" stroke="#9467bd" stroke-width="1.8" opacity="1.0"/>
|
||||
<text x="70" y="1600" font-size="15" font-weight="bold">Throughput - games per hour</text>
|
||||
<text x="70" y="1615" font-size="11" fill="#555">method: training_metrics.jsonl 'epoch' deltas (training_log.jsonl has no timestamps); 1 row = one 10-game chunk</text>
|
||||
<line x1="70" y1="1780.0" x2="1375" y2="1780.0" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1739.5" x2="1375" y2="1739.5" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1699.0" x2="1375" y2="1699.0" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1658.5" x2="1375" y2="1658.5" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1618.0" x2="1375" y2="1618.0" stroke="#dddddd"/>
|
||||
<line x1="70" y1="1780" x2="1375" y2="1780" stroke="black"/>
|
||||
<line x1="70" y1="1780" x2="70" y2="1618" stroke="black"/>
|
||||
<line x1="194.6" y1="1780" x2="194.6" y2="1784" stroke="black"/>
|
||||
<text x="194.6" y="1797" text-anchor="middle" font-size="11">20</text>
|
||||
<line x1="325.8" y1="1780" x2="325.8" y2="1784" stroke="black"/>
|
||||
<text x="325.8" y="1797" text-anchor="middle" font-size="11">40</text>
|
||||
<line x1="456.9" y1="1780" x2="456.9" y2="1784" stroke="black"/>
|
||||
<text x="456.9" y="1797" text-anchor="middle" font-size="11">60</text>
|
||||
<line x1="588.1" y1="1780" x2="588.1" y2="1784" stroke="black"/>
|
||||
<text x="588.1" y="1797" text-anchor="middle" font-size="11">80</text>
|
||||
<line x1="719.2" y1="1780" x2="719.2" y2="1784" stroke="black"/>
|
||||
<text x="719.2" y="1797" text-anchor="middle" font-size="11">100</text>
|
||||
<line x1="850.4" y1="1780" x2="850.4" y2="1784" stroke="black"/>
|
||||
<text x="850.4" y="1797" text-anchor="middle" font-size="11">120</text>
|
||||
<line x1="981.5" y1="1780" x2="981.5" y2="1784" stroke="black"/>
|
||||
<text x="981.5" y="1797" text-anchor="middle" font-size="11">140</text>
|
||||
<line x1="1112.7" y1="1780" x2="1112.7" y2="1784" stroke="black"/>
|
||||
<text x="1112.7" y="1797" text-anchor="middle" font-size="11">160</text>
|
||||
<line x1="1243.8" y1="1780" x2="1243.8" y2="1784" stroke="black"/>
|
||||
<text x="1243.8" y="1797" text-anchor="middle" font-size="11">180</text>
|
||||
<line x1="1375.0" y1="1780" x2="1375.0" y2="1784" stroke="black"/>
|
||||
<text x="1375.0" y="1797" text-anchor="middle" font-size="11">200</text>
|
||||
<line x1="66" y1="1780.0" x2="70" y2="1780.0" stroke="black"/>
|
||||
<text x="63" y="1784.0" text-anchor="end" font-size="11">0</text>
|
||||
<line x1="66" y1="1739.5" x2="70" y2="1739.5" stroke="black"/>
|
||||
<text x="63" y="1743.5" text-anchor="end" font-size="11">1688</text>
|
||||
<line x1="66" y1="1699.0" x2="70" y2="1699.0" stroke="black"/>
|
||||
<text x="63" y="1703.0" text-anchor="end" font-size="11">3375</text>
|
||||
<line x1="66" y1="1658.5" x2="70" y2="1658.5" stroke="black"/>
|
||||
<text x="63" y="1662.5" text-anchor="end" font-size="11">5063</text>
|
||||
<line x1="66" y1="1618.0" x2="70" y2="1618.0" stroke="black"/>
|
||||
<text x="63" y="1622.0" text-anchor="end" font-size="11">6750</text>
|
||||
<text x="722" y="1813" text-anchor="middle" font-size="12">chunk interval (10-game chunks)</text>
|
||||
<text x="16" y="1699" text-anchor="middle" font-size="12" transform="rotate(-90 16 1699)">games / hour</text>
|
||||
<circle cx="70.0" cy="1698.0" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="76.6" cy="1774.9" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="83.1" cy="1658.1" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="89.7" cy="1776.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="96.2" cy="1762.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="102.8" cy="1776.0" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="109.3" cy="1655.7" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="115.9" cy="1775.9" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="122.5" cy="1654.3" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="129.0" cy="1776.0" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="135.6" cy="1762.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="142.1" cy="1773.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="148.7" cy="1689.3" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="155.3" cy="1772.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="161.8" cy="1698.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="168.4" cy="1771.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="174.9" cy="1693.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="181.5" cy="1772.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="188.0" cy="1696.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="194.6" cy="1771.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="201.2" cy="1692.0" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="207.7" cy="1771.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="214.3" cy="1689.3" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="220.8" cy="1771.1" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="227.4" cy="1653.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="233.9" cy="1773.9" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="240.5" cy="1687.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="247.1" cy="1772.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="253.6" cy="1653.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="260.2" cy="1775.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="266.7" cy="1650.9" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="273.3" cy="1773.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="279.8" cy="1682.0" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="286.4" cy="1774.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="293.0" cy="1687.1" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="299.5" cy="1774.0" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="306.1" cy="1691.6" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="312.6" cy="1774.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="319.2" cy="1690.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="325.8" cy="1775.7" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="332.3" cy="1690.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="338.9" cy="1776.1" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="345.4" cy="1659.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="352.0" cy="1772.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="358.5" cy="1645.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="365.1" cy="1773.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="371.7" cy="1680.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="378.2" cy="1772.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="384.8" cy="1666.9" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="391.3" cy="1775.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="397.9" cy="1690.0" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="404.4" cy="1772.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="411.0" cy="1679.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="417.6" cy="1775.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="424.1" cy="1686.0" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="430.7" cy="1775.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="437.2" cy="1762.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="443.8" cy="1775.9" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="450.4" cy="1694.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="456.9" cy="1775.1" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="463.5" cy="1680.9" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="470.0" cy="1775.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="476.6" cy="1681.3" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="483.1" cy="1775.9" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="489.7" cy="1691.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="496.3" cy="1772.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="502.8" cy="1683.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="509.4" cy="1774.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="515.9" cy="1642.7" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="522.5" cy="1775.3" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="529.0" cy="1683.7" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="535.6" cy="1774.6" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="542.2" cy="1639.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="548.7" cy="1775.6" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="555.3" cy="1696.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="561.8" cy="1775.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="568.4" cy="1691.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="574.9" cy="1773.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="581.5" cy="1696.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="588.1" cy="1771.7" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="594.6" cy="1632.7" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="601.2" cy="1757.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="607.7" cy="1771.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="614.3" cy="1641.1" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="620.9" cy="1770.0" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="627.4" cy="1650.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="634.0" cy="1771.1" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="640.5" cy="1639.0" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="647.1" cy="1770.9" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="653.6" cy="1689.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="660.2" cy="1770.7" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="666.8" cy="1705.7" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="673.3" cy="1771.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="679.9" cy="1652.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="686.4" cy="1771.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="693.0" cy="1660.7" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="699.5" cy="1771.9" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="706.1" cy="1681.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="712.7" cy="1770.7" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="719.2" cy="1643.7" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="725.8" cy="1770.3" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="732.3" cy="1690.1" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="738.9" cy="1771.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="745.5" cy="1684.3" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="752.0" cy="1769.3" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="758.6" cy="1688.1" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="765.1" cy="1769.9" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="771.7" cy="1695.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="778.2" cy="1773.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="784.8" cy="1694.9" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="791.4" cy="1770.6" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="797.9" cy="1673.3" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="804.5" cy="1770.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="811.0" cy="1676.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="817.6" cy="1771.0" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="824.1" cy="1691.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="830.7" cy="1769.9" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="837.3" cy="1692.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="843.8" cy="1770.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="850.4" cy="1670.9" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="856.9" cy="1770.6" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="863.5" cy="1698.0" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="870.1" cy="1770.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="876.6" cy="1762.9" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="883.2" cy="1770.1" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="889.7" cy="1688.9" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="896.3" cy="1770.1" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="902.8" cy="1690.1" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="909.4" cy="1771.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="916.0" cy="1693.3" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="922.5" cy="1769.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="929.1" cy="1760.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="935.6" cy="1769.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="942.2" cy="1682.6" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="948.7" cy="1770.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="955.3" cy="1761.1" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="961.9" cy="1770.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="968.4" cy="1680.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="975.0" cy="1769.6" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="981.5" cy="1681.1" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="988.1" cy="1769.9" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="994.6" cy="1692.7" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1001.2" cy="1771.3" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1007.8" cy="1685.1" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1014.3" cy="1771.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1020.9" cy="1690.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1027.4" cy="1772.0" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1034.0" cy="1696.3" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1040.6" cy="1770.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1047.1" cy="1682.1" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1053.7" cy="1771.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1060.2" cy="1673.0" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1066.8" cy="1770.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1073.3" cy="1676.0" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1079.9" cy="1769.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1086.5" cy="1668.3" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1093.0" cy="1771.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1099.6" cy="1694.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1106.1" cy="1771.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1112.7" cy="1695.6" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1119.2" cy="1771.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1125.8" cy="1689.7" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1132.4" cy="1771.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1138.9" cy="1761.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1145.5" cy="1771.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1152.0" cy="1688.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1158.6" cy="1771.7" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1165.2" cy="1684.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1171.7" cy="1772.3" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1178.3" cy="1702.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1184.8" cy="1774.6" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1191.4" cy="1702.0" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1197.9" cy="1772.3" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1204.5" cy="1688.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1211.1" cy="1772.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1217.6" cy="1696.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1224.2" cy="1772.1" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1230.7" cy="1695.0" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1237.3" cy="1771.9" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1243.8" cy="1707.2" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1250.4" cy="1772.1" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1257.0" cy="1699.3" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1263.5" cy="1774.1" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1270.1" cy="1672.7" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1276.6" cy="1771.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1283.2" cy="1689.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1289.7" cy="1771.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1296.3" cy="1696.0" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1302.9" cy="1774.3" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1309.4" cy="1689.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1316.0" cy="1772.3" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1322.5" cy="1687.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1329.1" cy="1771.8" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1335.7" cy="1690.1" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1342.2" cy="1772.0" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1348.8" cy="1691.0" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1355.3" cy="1772.5" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1361.9" cy="1701.1" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1368.4" cy="1770.4" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<circle cx="1375.0" cy="1682.9" r="1.6" fill="#7f7f7f" opacity="0.25"/>
|
||||
<polyline points="96.2,1730.7 161.8,1740.2 227.4,1724.1 293.0,1727.4 358.5,1721.3 424.1,1738.6 489.7,1725.4 555.3,1727.8 620.9,1709.3 686.4,1720.0 752.0,1730.7 817.6,1725.8 883.2,1738.6 948.7,1741.6 1014.3,1730.2 1079.9,1726.2 1145.5,1738.5 1211.1,1735.3 1276.6,1731.1 1342.2,1731.1" fill="none" stroke="#2ca02c" stroke-width="3.5" opacity="1.0"/>
|
||||
<rect x="0" y="1820" width="1400" height="180" fill="#f2f2f2"/>
|
||||
<text x="16" y="1841" font-size="15" font-weight="bold">How to read</text>
|
||||
<text x="1384" y="1840" text-anchor="end" font-size="11" fill="#666">Regenerate anytime: python3 tools/plot_progress.py</text>
|
||||
<text x="16" y="1863" font-size="13">Test matches: dots are single fights, thick line shows trend.</text>
|
||||
<text x="16" y="1880" font-size="13">Real battles only. Rising lines mean the bot improves.</text>
|
||||
<text x="16" y="1897" font-size="13">Loss spikes are normal early; endless growth is bad.</text>
|
||||
<text x="16" y="1914" font-size="13">Alpha high means experimenting; falling too fast freezes habits.</text>
|
||||
<text x="16" y="1931" font-size="13">Throughput flat is healthy; dips mean something slowed.</text>
|
||||
<text x="16" y="1948" font-size="13">This file reloads itself in Chrome every sixty seconds.</text>
|
||||
<text x="16" y="1965" font-size="13">Regenerate anytime with tools/watch_dashboard.sh or the python command.</text>
|
||||
<script type="text/javascript"><![CDATA[ setTimeout(function(){ location.reload(); }, 60000); ]]></script>
|
||||
</svg>
|
||||
|
After Width: | Height: | Size: 63 KiB |
@@ -0,0 +1,332 @@
|
||||
## The story so far, in simple words
|
||||
|
||||
This project trains a robot tank. It plays many fights against other tanks.
|
||||
After each fight it changes itself a little. It keeps the changes that helped it win.
|
||||
|
||||
Night 1 (run 1) finished without problems. It ran for 14 hours alone. It never crashed.
|
||||
It beat an old copy of itself most of the time. It beat Crazy about half the time.
|
||||
Three things went badly. First, learning was not stable. Good skill appeared, then disappeared again.
|
||||
Second, the saved "best" version came from one lucky perfect score. It was not really its best.
|
||||
Third, the bot learned to hide and survive. It almost never shot back.
|
||||
|
||||
The human approved five fixes. All five were put into the code.
|
||||
Run 2 used these fixes. Its error numbers grew far too big. Learning broke.
|
||||
We made one speed number smaller. This number sets how fast one part learns.
|
||||
Then we dropped the broken progress and started clean. This is run 3. It is running now.
|
||||
|
||||
Next we watch run 3. One of three doors will open.
|
||||
Door 1: it stays steady. We let it run to the end.
|
||||
Door 2: the numbers grow too big again. We turn the next speed number down.
|
||||
Door 3: it stays steady but still fights badly. We teach aiming as a separate, direct lesson.
|
||||
|
||||
Updated: 2026-08-23 — this section is refreshed at every major step.
|
||||
|
||||
## Small dictionary
|
||||
|
||||
- **training**: the time when the bot plays fights and changes itself to improve. It learns only during training.
|
||||
- **battle**: one group of fights against one opponent. The bot restarts between groups.
|
||||
- **round**: one single fight. Win it by destroying the enemy tank or outliving it.
|
||||
- **chunk**: one work block: a battle of up to 10 rounds, then some learning from it.
|
||||
- **eval (test match)**: a test match. The bot does not learn during these. We use them only to measure.
|
||||
- **win rate**: how many test matches were won, as a percent. 8 wins in 10 matches = 80%.
|
||||
- **checkpoint**: a saved copy of the bot's brain (a zip file). Written every few learning steps.
|
||||
- **"best" checkpoint**: the saved copy we currently call best. Run 1 picked one from a lucky score, hence the quotes.
|
||||
- **replay buffer**: the bot's memory of past moments: what it saw, did, and received. Learning picks old moments from it.
|
||||
- **loss (critic/actor)**: a number saying how wrong the bot's inner guesses are. Lower usually means better. Losses growing huge mean trouble.
|
||||
- **alpha**: a dial setting how much the bot tries new moves instead of repeating known good ones.
|
||||
- **MA / composite score**: MA is the average of the last few win rates; it smooths luck. Composite is the average of MAs across all test opponents.
|
||||
- **twin (SacTwin)**: a frozen copy of our own bot, used as a practice partner. Beating it proves real improvement.
|
||||
- **lever**: one numbered change we prepared, waiting for approval. There are levers 1 to 5.
|
||||
- **watchman**: a helper who checks the running training at set times and stops it if something breaks.
|
||||
|
||||
---
|
||||
|
||||
# Campaign Notebook — campaign-v1 (SAC_LSTM_Bot)
|
||||
|
||||
Overnight training campaign on branch `research/goto-controller`.
|
||||
Companion tickets: config+launch = **#56**, morning verdict = **#57**, map = **#53**.
|
||||
Monitoring contract: observability inventory in **#55 comment 461** (9 signals, thresholds, four-way discrimination).
|
||||
|
||||
## Locked config (campaign-v1)
|
||||
|
||||
Launch command (tmux session `sac_campaign`, stdout teed to `campaign_stdout.log`):
|
||||
|
||||
```bash
|
||||
SAC_OPPONENTS='Corners:3,Crazy:2,RamFire:1,Target:1,SacTwin:1' \
|
||||
SAC_EVAL_OPPONENT=Corners \
|
||||
SAC_TOTAL_ROUNDS=25000 \
|
||||
SAC_CHUNK_SIZE=10 \
|
||||
SAC_EVAL_INTERVAL=2 \
|
||||
SAC_EVAL_ROUNDS=10 \
|
||||
SAC_MAX_CRASHES=5 \
|
||||
SACLSTM_HIDDEN_SIZE=256 \
|
||||
SACLSTM_BATCH_SIZE=16 \
|
||||
SACLSTM_UTD_RATIO=1 \
|
||||
SACLSTM_SAVE_INTERVAL=5 \
|
||||
./sac_train.sh 2>&1 | tee -a campaign_stdout.log
|
||||
```
|
||||
|
||||
| Knob | Value | Source / rationale |
|
||||
|------|-------|--------------------|
|
||||
| `SAC_OPPONENTS` | `Corners:3,Crazy:2,RamFire:1,Target:1,SacTwin:1` | Working recommendation kept. Twin pinned at **1 not 2**: #54 showed mirror battles end early ⇒ fewer transitions per chunk; weight 2 would starve the replay buffer. |
|
||||
| `SAC_EVAL_OPPONENT` | `Corners` | Pinned explicitly (= first pool entry default, #54 note) — removes reorder footgun. |
|
||||
| `SAC_TOTAL_ROUNDS` | `25000` | Sized so the **wall-clock ceiling binds first**: measured throughput ~3400 rounds/h early (drops as battles lengthen) ⇒ 2000 would have exhausted in ~1 h. |
|
||||
| `SAC_CHUNK_SIZE` | `10` | Harness default. |
|
||||
| `SAC_EVAL_INTERVAL` | `2` | Harness default — eval every ~20 rounds. |
|
||||
| `SAC_EVAL_ROUNDS` | `10` | Harness default. |
|
||||
| `SAC_MAX_CRASHES` | `5` | Harness self-abort; monitor intervenes earlier at ≥3 consecutive crashes (#55). |
|
||||
| `SACLSTM_HIDDEN_SIZE` | `256` | Module default (`network.nim`); real capacity vs #49 smoke's 32; under `MaxHidden`=512 cap. |
|
||||
| `SACLSTM_BATCH_SIZE` | `16` | Module default (`integration.nim`). |
|
||||
| `SACLSTM_UTD_RATIO` | `1` | Module default. |
|
||||
| `SACLSTM_SAVE_INTERVAL` | `5` | **Deviation** from default 500 and from #54's "10–20": at hidden 256 a gradient step takes ~1 s and a chunk process fits only ~10–18 steps (see incident below) — interval must sit **inside the per-process step budget**. 5 ⇒ checkpoint every ~5–10 s of active training; IO trivial (10.5 MB zip, atomic replace). |
|
||||
| *(not pinned)* | module defaults | `LR_ACTOR/LR_CRITIC/LR_ALPHA=3e-4`, `GAMMA=0.99`, `TAU=0.005`, `TARGET_ENTROPY=-4.0`, `BUFFER_CAPACITY=500000`, `BURN_IN=8`, `TRAIN_WINDOW=16`. |
|
||||
|
||||
**Budget**: generous wall-clock **ceiling, not a deadline** — originally **T+12 h** from launch 00:21:55 CEST 2026-08-22 ⇒ 12:22 CEST (epoch 1787394115); **extended 2026-08-22 ~07:35 by orchestrator decision on human mandate** ("no deadlines — let the 25000-round budget complete", ~14:35 projected) ⇒ ceiling now **16:30 CEST 2026-08-22 (epoch 1787409000)**, enforced by a hard user-systemd net unit `sac-ceiling-net` (sleeps to the epoch, then kills the tmux session and any straggler harness processes). While HEALTHY per #55 discrimination rules the run continues; a monitor kills the tmux session at the ceiling or on an intervene threshold.
|
||||
|
||||
**Fresh start**: pre-campaign `weights/` held #49-smoke 32-hidden checkpoints, incompatible with hidden=256. Archived to `weights_smoke49_backup/`; campaign baseline re-established by probe battles vs Corners (random-init hidden-256 checkpoint, `best_score.txt` reset then re-raised to 20 by a genuine eval). Twin regenerated via `./make_twin.sh` from that baseline (md5 `61521cff…` verified seed).
|
||||
|
||||
## Phase log
|
||||
|
||||
Every entry below starts with a plain-language first sentence. Technical detail follows for those who want it.
|
||||
|
||||
- **2026-08-21 23:22** — Claim posted on #56 (comment 471). Config locked, notebook committed (`2f49cb2`).
|
||||
- **2026-08-21 23:27** — Smoke weights archived; release build; bootstrap + probe battles vs Corners established a hidden-256 baseline checkpoint (`sac_latest.zip`, 10.5 MB) and twin seed.
|
||||
- **2026-08-21 23:32** — Twin regenerated (md5-verified). **Launch attempt 1** (SAVE_INTERVAL=20, TOTAL_ROUNDS=2000): ran 16+ chunks, evals every 2 chunks — but **zero checkpoints persisted** (see incident). Killed 23:48.
|
||||
- **2026-08-21 23:52–00:10** — Diagnosis (see incident): interval=1 fired, interval=2/20 never; instrumentation + /proc thread forensics ⇒ per-process step budget ~10–18 at ~1 s/step; save check ran only between drain-burst passes.
|
||||
- **2026-08-22 00:12** — Fix: save check moved inside the gradient-step loop, committed `2653671`. Validated: interval=5 save fired ~12 s into a battle.
|
||||
- **2026-08-22 00:21:55** — **Launch (final)**: tmux `sac_campaign`, config above. First campaign save on disk at t+54 s; eval #1 on cadence.
|
||||
- **2026-08-22 00:30** — HEALTHY checklist passed (see below).
|
||||
- **2026-08-22 03:45** — Watch shift 1 (00:34–03:35): liveness flawless (10/10 HEALTHY, zero banners, zip ≤30 s). Learning signal: steady-state eval vs Corners 0–5% with two isolated 10/10 spikes (~01:15) → capability emerged, then lost. Eval-regression intervene threshold fired per #55; intervention DEFERRED to Campaign verdict (#57) — rationale: n=2 evidence, no loss metrics, buffer-loss on restart, run completes ~08:05 anyway. Milestone issue: #58 "Campaign-v1 watch: eval-regression threshold fired — intervention deferred to verdict".
|
||||
- **2026-08-22 ~07:15** — Morning audit: policy demonstrably learning off-benchmark (SacTwin 73→100%, Crazy 26→56%) while Corners eval stays ~0–9% with 3 transient 10/10s; wall-clock ceiling extended to let the 25k complete (~14:35 projected); instability-vs-plateau question left to the curve.
|
||||
- **2026-08-22 ~07:50 — score:60 anatomy**: Tank Royale survival(50)+last-survivor(10) awarded when opponent dies while we survive; exactly-60 ⇒ zero damage dealt by us that round (opponent self-destructed via wasted shots + 0.1/turn inactivity drain). 3,878 rounds (19%); modal vs SacTwin; vs Corners 986 damageless outlives vs 242 true wins; combined with 43% of rounds being score:0, texture = survivor-not-fighter against walls. RL rewards are event-driven (rewards module), so behavioral evidence, not reward poisoning. Feeds #57 levers: aggression shaping / specialist-vs-generalist.
|
||||
- **2026-08-22 ~08:05 — `.part` debris forensics + sweep**: 99 `*.zip.tmp.*.part` (483 MB) are NOT weights.nim debris (that proc uses fixed `.tmp` + finally-cleanup, working); naming matches an external write-temp→rename copier killed mid-write, bursts correlating with kill events; possible culprit: a folder-sync client fighting a file that changes every ~20 s (**human asked to confirm**). Swept with `-mmin +10` age guard (protects in-flight writes); 483 MB freed, real zips untouched — note 2 fresh `.part` reappeared minutes later, copier still active.
|
||||
- **2026-08-22 ~08:00 — ceiling defused**: the 12:22 'self-kill' was notebook prose instructing watchmen — never an OS mechanism; rewritten to 16:30 CEST AND armed a real systemd --user net `sac-ceiling-net` firing 16:30:00 (epoch 1787409000); training uninterrupted (counter +99/147 s verified); annotated #56 comment 487.
|
||||
- **2026-08-22 ~14:28 — CAMPAIGN ENDED NATURALLY**: banner `>>> training complete: 2500 chunks` after **14 h 07 m** (00:21:55 → ~14:28 CEST); round_counter **38912**; **zero crash banners across the whole run**; the `sac-ceiling-net` backstop never fired — cancelled unneeded. Budget note: the harness loop is **chunk-based** — `SAC_TOTAL_ROUNDS=25000` ÷ `CHUNK_SIZE=10` ⇒ **2500 chunk battles** of ≤10 rounds each; the "25k-rounds" label was a misnomer (the counter also accrues 1287×10 eval rounds and rerun chunks). Final eval vs Corners: 10%.
|
||||
- **2026-08-22 ~14:50 — VERDICT posted (→ #57)**: **RETUNE BEFORE SCALING.** Ops layer PROVEN (14 h autonomous, zero crashes, self-healing restarts, natural completion — the harness scales); learning REAL BUT NARROW (within-opponent gains genuine — SacTwin 90.7%, Crazy 26→48% — but specialist-not-generalist, walls untouched; 12 eval spikes ≥8/10 incl. 5×10/10, none retained); benchmark pathology: Corners-only deterministic eval + single-max best gating froze `sac_best.zip` at 01:01:57 on a fluke 10/10. Five code-level levers staged **awaiting human sign-off**; v2 NOT launched. Full rationale: Results below + #57 resolution comment.
|
||||
- **2026-08-22 ~14:50 — HYGIENE (campaign over, no live writers — safe)**: `sac-ceiling-net` stopped + `reset-failed` (backstop obsolete); `.part` corpse sweep **48 → 0** (no age guard needed — nothing writes anymore); future-run guard added to `sac_train.sh`: startup `rm -f "$WEIGHTS_DIR"/sac_latest.zip.tmp.*.part "$WEIGHTS_DIR"/sac_latest.zip.tmp` so a SIGKILLed run's libzip modify-path corpses can't accumulate again.
|
||||
|
||||
## Decision-issue index
|
||||
|
||||
| Issue | What it decided |
|
||||
|-------|-----------------|
|
||||
| #37–#48 | Bot built: skeleton, state, actions, rewards, LSTM network, weights, SAC+LSTM training, integration. |
|
||||
| #49 | Training harness + smoke run (toy hyperparams: hidden 32). |
|
||||
| #54 | Mirror-twin sparring partner; SAVE_INTERVAL persistence rule; eval opponent = first pool entry. |
|
||||
| #55 | 9-signal observability inventory; CRASHED/STALLED/SLOW-LEARNER/HEALTHY discriminators; monitor thresholds. |
|
||||
| #56 | This campaign: locked config + launch + the save-check fix (`2653671`). |
|
||||
| #57 | Morning verdict — consumes this notebook + logs. |
|
||||
|
||||
## Incidents & checks
|
||||
|
||||
### Incident 1 — zero checkpoint persistence at production sizes (launch blockers, fixed)
|
||||
|
||||
**Symptom**: campaign ran 26+ chunk processes across two attempts without a single `sac_latest.zip` update, while rounds/evals flowed normally. #54's rule ("keep `SACLSTM_SAVE_INTERVAL` well below per-chunk gradient-step counts, 10–20 fired in smokes") silently broke at hidden 256.
|
||||
|
||||
**Diagnosis chain** (all reproducible):
|
||||
1. Interval=1 saved within seconds; interval=2 and 20 never saved — through the *same* harness ⇒ not env propagation.
|
||||
2. Temporary step instrumentation (bot stderr via a one-line `SAC_LSTM_Bot.sh` redirect — the vendored runner swallows bot stderr, #55 gap S9-adjacent): steps cost **~1.06 s each**; a drain burst queued 53 steps; logging stopped mid-pass while rounds kept completing.
|
||||
3. `/proc/<pid>/task` sampling: training thread alive and RUNNING (~13 s CPU per ~40 s process) — not deadlocked, just slow ⇒ **per-process step budget ≈ 10–18 steps**.
|
||||
4. The save check lived *between* drain-burst passes; with bursts queueing minutes of steps, `stepCount` never reached `nextSave` before process teardown. Smoke runs masked this: hidden 32 steps were sub-millisecond, so hundreds of steps fit per chunk.
|
||||
|
||||
**Fix** (commit `2653671`): save check relocated **inside** the step loop (checked every gradient step; `packFull`+`trySend` unchanged). Validated: interval=5 save fires ~12 s into a battle; campaign save fired 54 s after launch.
|
||||
|
||||
**Config consequences**: `SACLSTM_SAVE_INTERVAL=5` (inside the per-process budget; #54's 10–20 was derived at smoke speeds). `SAC_TOTAL_ROUNDS=25000` (throughput measured ~3400 rounds/h, so 2000 was a 1-hour budget, not an overnight one). OMP_NUM_THREADS=1 tested and **not** needed (hang was step-budget exhaustion, not OpenMP).
|
||||
|
||||
### Launch health check (t+9 min, 00:30:16) — **HEALTHY** per #55 checklist
|
||||
|
||||
| Signal | Reading | Verdict |
|
||||
|--------|---------|---------|
|
||||
| S1 harness stdout | teed to `campaign_stdout.log`; **0** `crash`/`aborted` banners | ✓ |
|
||||
| S4 round_counter | 1890, +460 in 9 min (~51 rounds/min) | ✓ advancing |
|
||||
| S6 sac_latest.zip mtime | **6 s old**; first save at t+54 s | ✓ fresh |
|
||||
| S2 training_log.jsonl | 1492 lines, growing; last ticks=668, plausible | ✓ |
|
||||
| S3 eval_log.jsonl | age 3 s (atomic replace); eval every 2 chunks | ✓ on cadence |
|
||||
| S5 best_score | 20 (from a genuine campaign-1 eval; non-decreasing) | ✓ |
|
||||
| Sampling | 112 chunks: Corners 41 / Crazy 25 / RamFire 15 / SacTwin 16 / Target 15 ≈ weights 3:2:1:1:1 | ✓ plausible |
|
||||
| Disk | 418 GB free | ✓ |
|
||||
|
||||
Early evals 0% vs Corners — expected for a near-random policy minutes in; SLOW-LEARNER watch rule (flat ≥5 evals = watch) applies, never intervene.
|
||||
|
||||
## Check-in procedure (for monitor sessions)
|
||||
|
||||
```bash
|
||||
tmux capture-pane -p -t sac_campaign | tail -5 # S1: banners, crashes
|
||||
cat ~/Projects/SirRoboGarage/SAC_LSTM_Bot/weights/round_counter.txt
|
||||
stat -c '%Y' ~/Projects/SirRoboGarage/SAC_LSTM_Bot/weights/sac_latest.zip # age <~600s = training alive
|
||||
tail -1 ~/Projects/SirRoboGarage/SAC_LSTM_Bot/training_log.jsonl
|
||||
tail -3 ~/Projects/SirRoboGarage/SAC_LSTM_Bot/campaign_stdout.log # eval results / new best
|
||||
cat ~/Projects/SirRoboGarage/SAC_LSTM_Bot/weights/best_score.txt
|
||||
```
|
||||
|
||||
Intervene per #55 thresholds: ≥3 consecutive `crash #N` banners; ΔS4=0 over ≥15 min; zip mtime >10 min stale while S4 advances (STALLED); disk <1 GB. At the **ceiling (16:30 CEST Aug 22, epoch 1787409000 — extended from 12:22 per human mandate)**: `tmux kill-session -t sac_campaign` if still running (the hard net unit `sac-ceiling-net` fires at the same epoch regardless of monitors) — final state is in `weights/`, logs, and this notebook.
|
||||
|
||||
## Results (campaign-v1 — filled by #57)
|
||||
|
||||
**Run**: 2500/2500 chunks · **14 h 07 m** autonomous (00:21:55 → ~14:28 CEST 2026-08-22) · round_counter 38912 · **zero crashes** · graceful banner `>>> training complete: 2500 chunks` · systemd net never fired. Budget was chunk-based (see phase log) — the "25k-rounds" label was a misnomer.
|
||||
|
||||
**Per-opponent training win rates** (run-3 slice of `training_log.jsonl`):
|
||||
|
||||
| Opponent | Win rate | Record | Note |
|
||||
|----------|----------|--------|------|
|
||||
| SacTwin | **90.7%** | 2693/2970 | vs frozen past-self — genuine self-play gain |
|
||||
| Crazy | 48.3% | 2965/6140 | doubled from 26% early-run |
|
||||
| Target | 9.6% | — | static, barely moved |
|
||||
| Corners | 7.2% | — | walls untouched |
|
||||
| RamFire | **0%** | 0/2920 | mirrors the PPO-era ladder — ram-class needs dedicated pressure |
|
||||
|
||||
**Eval vs Corners (pinned benchmark)**: 1287 evals · overall mean **7.7%** · histogram headline: `0/10 = 867 (67%)`, spikes ≥8/10 = **12** (incl. **5× perfect 10/10**) · final eval 10%. Stdout log carries no timestamps; timing reconstructed from file mtimes.
|
||||
|
||||
**Best-zip paradox**: `sac_best.zip` frozen since **01:01:57** — a single lucky 10/10 at ~round 3.5k wrote `best_score=100`, and no later eval could outrank a perfect score (even genuine ~50%-winrate stretches elsewhere). Best checkpoint = lottery ticket, decoupled from the steady-state policy (which sat at 0–10% vs Corners).
|
||||
|
||||
**Verdict: RETUNE BEFORE SCALING** — full rationale in #57 resolution comment. Five code-level levers staged for human sign-off: (1) eval rotation across pool + moving-average best gating; (2) reward shaping toward damage/aggression incl. anti-ram signal; (3) training-loss/step metrics logged from the training thread (#55 gap #1); (4) gate `sendTrainingMsg` off in eval mode (#55 gap #3); (5) optional stability knobs (lower LR / entropy coeff) once loss curves exist.
|
||||
|
||||
---
|
||||
|
||||
# Campaign Notebook — campaign-v2 (SAC_LSTM_Bot)
|
||||
|
||||
## Locked config (campaign-v2)
|
||||
|
||||
Launched verbatim from orchestrator mandate (umbrella ticket #59, all five levers approved & implemented in `a07e530` + `6fc01eb`):
|
||||
|
||||
```bash
|
||||
cd /home/davide/Projects/SirRoboGarage/SAC_LSTM_Bot && \
|
||||
SAC_OPPONENTS='Corners:3,Crazy:2,RamFire:2,Target:1,SacTwin:1' \
|
||||
SAC_EVAL_OPPONENTS='Corners,Crazy,Target' \
|
||||
SAC_EVAL_INTERVAL=2 SAC_EVAL_ROUNDS=10 \
|
||||
SAC_TOTAL_ROUNDS=25000 SAC_CHUNK_SIZE=10 SAC_MAX_CRASHES=5 \
|
||||
SACLSTM_HIDDEN_SIZE=256 SACLSTM_BATCH_SIZE=16 SACLSTM_SAVE_INTERVAL=5 \
|
||||
./sac_train.sh 2>&1 | tee -a campaign_v2_stdout.log
|
||||
```
|
||||
|
||||
| Knob | Value | Rationale |
|
||||
|------|-------|-----------|
|
||||
| `SAC_OPPONENTS` | Corners:3, Crazy:2, **RamFire:2**, Target:1, SacTwin:1 | RamFire bumped 1→2 vs v1: anti-ram shaping (lever 2, `6fc01eb`) needs exposure to fire; without samples there is no gradient signal against ram-class |
|
||||
| `SAC_EVAL_OPPONENTS` | Corners,Crazy,Target | Lever-1 rotation set (default); composite = mean of per-opponent MA-5 win rates |
|
||||
| `SAC_EVAL_INTERVAL/ROUNDS` | 2 / 10 | Unchanged from v1 cadence |
|
||||
| `SAC_TOTAL_ROUNDS/CHUNK_SIZE/MAX_CRASHES` | 25000 / 10 / 5 | Same budget semantics as v1 (chunk-based) |
|
||||
| `SACLSTM_HIDDEN_SIZE/BATCH_SIZE/SAVE_INTERVAL` | 256 / 16 / 5 | Architecture + throughput knobs carried over; LR/entropy defaults untouched = **lever-5 conditional posture** (activation decided by loss-curve evidence, not upfront) |
|
||||
|
||||
## Fresh start & archive
|
||||
|
||||
v1 state archived intact (notebook references preserved) into `SAC_LSTM_Bot/weights_v1_archive/`: `sac_latest.zip`, `sac_best.zip`, `best_score.txt`, `round_counter.txt`, `training_log.jsonl`, `eval_log.jsonl`, `campaign_stdout.log`, plus the lever-3 smoke leftover `training_metrics.jsonl` (from `src/SAC_LSTM_Bot/`, moved so v2 loss curves start clean for lever-5 reading). Main `weights/` verified empty afterwards ⇒ main bot takes the genuine random-init path (`loadOrInitFull` → `randomFull()`).
|
||||
|
||||
**Twin reseed**: `make_twin.sh` requires a seed zip, but the fresh-start baseline has none. Generated a fresh random-init checkpoint (hidden=256, alpha=1.0 matching `logAlpha=0`) via a throwaway Nim script against `network.nim`/`weights.nim`, seeded it as `sac_best.zip` transiently, ran `./make_twin.sh`, removed the transient copy. Verified twin dir got byte-identical fresh zips (`cmp` OK; NOT v1 zips) + `round_counter.txt=0`. No script changes needed.
|
||||
|
||||
## Safety net
|
||||
|
||||
`systemd-run --user --unit=sac-ceiling-net-v2` armed at launch: sleeps 72000 s then `tmux kill-session -t sac_campaign_v2; sleep 5; pkill -f sac_train.sh`. Unit active at 19:59:40 CEST 2026-08-22, fires **15:59:40 CEST 2026-08-23** (epoch 1787493580).
|
||||
|
||||
## Launch & health evidence (first ~45 min)
|
||||
|
||||
Launched 20:00:19 CEST 2026-08-22 (epoch 1787421619), tmux session `sac_campaign_v2`.
|
||||
|
||||
| Check | Evidence |
|
||||
|-------|----------|
|
||||
| Round counter advances | `round_counter.txt` 95→100→465 across polls; RunTraining `Counter check passed: N == N` every chunk |
|
||||
| Metrics JSONL with scalars | `training_metrics.jsonl` growing (20 lines @ t+45m): full `{epoch, steps, buffer_size, drained, grad_steps, critic_loss, actor_loss, alpha_loss, alpha}` per line |
|
||||
| Eval rotation cycles ≥2 | All 3 opponents EVERY cycle: Corners→Crazy→Target ×4+ cycles in `eval_log.jsonl` (10 games each per cycle) |
|
||||
| MA files written | `weights/ma_history_{Corners,Crazy,Target}.txt` created at first cycle, appended since |
|
||||
| Best-gate on composite only | First write exactly when composite 0.0000 > −1 (missing-file default); later 0% cycles correctly did NOT rewrite (strict improvement enforced) |
|
||||
| No transitions during eval windows | Lever-4 gate active by construction (`sendTrainingMsg` drops all msgs under `SACLSTM_EVAL_MODE=1`, unit-tested); metrics epochs cluster at chunk boundaries |
|
||||
| Zero crash banners | `grep -c 'crash #'` = 0 through 20 chunks |
|
||||
| Sampling distribution plausible | 20 chunks: Corners 10, RamFire 5, Crazy 3, SacTwin 2, Target 0 — within small-n noise of weights (3/2/2/1/1)/9; RamFire already sampled (exposure goal met) |
|
||||
|
||||
Twin liveness: own `round_counter.txt` advancing, own `sac_latest.zip` updating during SacTwin chunks, own metrics file separate from the main bot's.
|
||||
|
||||
### Watch items (not blockers)
|
||||
|
||||
1. **Training-throughput signature**: metrics lines consistently show `steps=1, buffer_size=24 (= burnIn 8 + trainWindow 16, i.e. exact canSample threshold), drained=1` — one gradient step per pass at threshold-crossing moments rather than large drain bursts. Mechanism unexplained by static code read (per-tick sends should yield bigger bursts); v1 learned to its score-60 state under the same integration code without instrumentation, so learning is not obviously broken — but effective grad-steps/hour is THE number to check at first review. This is precisely what lever-3 instrumentation exists to surface.
|
||||
2. **Early critic-loss spikes**: two `1.56e16` outliers (t+6:16, t+12:04) amid otherwise sane values (~8–35) — same class as the random-init artifact flagged in #59 phase-2 notes (~1e12 there); expect decay. If persistent past early chunks, feeds the lever-5 decision.
|
||||
3. **`.part` corpses**: three `sac_latest.zip.tmp.*.part` files accumulated mid-run (libzip interrupted-write artifact, #57 forensics); harmless — atomic renames keep the main zips valid, startup sweep clears them next restart.
|
||||
|
||||
## ~21:30 — v2 attempt-1 checkpoint finding
|
||||
|
||||
Attempt-1's rolling zip was dead ~40 min at birth: `SACLSTM_SAVE_INTERVAL` counts **gradient steps**, and each training pass inside the short battle processes carries exactly **1 step** (the `steps=1, drained=1` signature of watch item 1) ⇒ interval-5 never reached its trigger within a process lifetime — no `sac_latest.zip` roll ever fired despite healthy training. Same class as v1's 23:32 incident, resurfacing through a different seam (per-process step budget vs per-pass step count).
|
||||
|
||||
**Fix (env-level, no rebuild)**: `SACLSTM_SAVE_INTERVAL=1`, restart 21:09:30 CEST. Saves verified flowing: 21:09 / 21:13 / 21:16 (and still flowing at audit time — `sac_latest.zip` mtime 21:22:28, `sac_best.zip` 21:12:50). The locked-config table above keeps the original attempt-1 launch line for the record; live relaunch differs only in this knob.
|
||||
|
||||
## ~21:30 — twin-freeze contract corrected
|
||||
|
||||
`make_twin.sh` twin launcher now exports `SACLSTM_EVAL_MODE=1` (commit `19f34ab`). Without lever-4's gate on the twin side, the **v1 SacTwin had been TRAINING throughout**, not frozen — every mirror battle updated the twin's own weights. Consequence: v1's headline "**90.7% vs twin**" is retroactively an **arms-race win rate** (both policies co-evolving), not a fixed-benchmark score, and the #57 verdict phrasing implying a frozen sparring partner is corrected by this entry. No numbers change; the story does.
|
||||
|
||||
## ~21:30 — net extended + pacing decision
|
||||
|
||||
Slowdown attribution complete: **89% of wall clock = eval rotation by design** (`SAC_EVAL_INTERVAL=2` × 3 opponents × 10 rounds each); trainer exonerated at **~500 ticks/s**. Decision recorded in issue **#60**.
|
||||
|
||||
Ceiling net re-armed to match the extended budget: fires **Thu 2026-08-27 07:59:40 CEST (epoch 1787810380)** — supersedes the 2026-08-23 15:59:40 fire noted under Safety net. Attempt-1 artifacts preserved in `/tmp/v2_attempt1_backup/`.
|
||||
|
||||
# Campaign v2 — attempt-3 (2026-08-23)
|
||||
|
||||
## ~06:45 — LEVER-5 TRIGGERED (watchman shift 2, 02:11–06:22 CEST)
|
||||
|
||||
Evidence-gated activation of the lever-5 conditional posture (`SACLSTM_LR_CRITIC`), human signed off. Trigger evidence:
|
||||
|
||||
| Signal | Observation |
|
||||
|--------|-------------|
|
||||
| actor_loss | \|34M\| → \|106M\| monotone (+15M/h), **no plateau** |
|
||||
| critic_loss | spikes >1e12 in **100%** of NEW metric records; campaign share 80.2% |
|
||||
| Signature | exact `1.5625e16` recurring (= float32 saturation neighborhood) |
|
||||
| alpha | decayed 0.951 → 0.711 (entropy collapse under runaway Q scale) |
|
||||
| best_score.txt | frozen 00:32 (=83.3333) through 5h50m of composite oscillation 3.3–43.3 |
|
||||
| Per-opponent MAs | whipsawing (Crazy 30→100→90→0→60→90→10→0) |
|
||||
|
||||
## DECISION — LR_CRITIC=1e-4, single knob, fresh start
|
||||
|
||||
- **Adjudication**: primary suspect is critic/Q value-scale growth. Actor loss inherits Q magnitude through the policy gradient, so the actor explosion is downstream; alpha decay is a *symptom* (entropy temperature chasing a blown value scale), not a cause. Therefore ONE knob moves: `SACLSTM_LR_CRITIC=1e-4` (critic learns slower → Q estimates stay estimable); LR_ACTOR/LR_ALPHA stay at default 3e-4, TARGET_ENTROPY −4. Multi-knob changes would confound attempt-3's attribution.
|
||||
- **Fresh start over warm start**: attempt-2's weights carry saturated Q representations; warm-starting them under a new LR would inherit the pathology we're trying to escape (contamination risk). Attempt-2 archived intact, nothing discarded.
|
||||
|
||||
## MA-wipe root cause FOUND & FIXED (`167bcc4`)
|
||||
|
||||
Watchman anomaly explained: `ma_history_*.txt` held exactly one leading-space value per file (e.g. `[ 10 ]`, `[ 20 ]`, `[ 0 ]` in the attempt-2 backup) instead of the designed 5-value history, degrading the composite best-gate to last-cycle mean — which fully accounts for the composite whipsaw 3.3–43.3.
|
||||
|
||||
Root cause: in `eval_checkpoint`, the append line was
|
||||
|
||||
```bash
|
||||
printf '%s\n' "$(cat "$(ma_hist_file "$opp")")" "$wr" | tail -n 5 | tr '\n' ' ' > "$(ma_hist_file "$opp")"
|
||||
```
|
||||
|
||||
Bash sets up the `>` redirect **before** running the command substitution, so `$(cat f)` always read the freshly truncated file ⇒ every cycle wiped history to `" $wr "`. Reproduced standalone on bash 5.3 before touching the script; no other writer exists (grep). Fix hoists the read into its own statement and word-splits it so `tail -n 5` keeps exactly the last MA_WINDOW values. Verified live in attempt-3: two eval cycles produced `[0 0 ]` per opponent (previously impossible).
|
||||
|
||||
## Attempt-3 launch (2026-08-23 07:01:28 CEST, epoch 1787461288)
|
||||
|
||||
- Attempt-2 stopped cleanly (tmux kill; no stray procs); artifacts (76 items, 336 MB incl. `.part` corpses + MA/best/latest/counter) → `/tmp/v2_attempt2_backup/`; final round_counter **8452**; `weights/` verified empty.
|
||||
- Twin regenerated via documented transient-seed method: throwaway Nim script (`initSACTrainer(35,4)` + `saveWeights`, hidden=256 env, alpha=1.0) seeded as transient `weights/sac_best.zip` → `./make_twin.sh` → twin zips byte-identical to seed (`cmp` OK), twin counter=0 → transient removed.
|
||||
- Launch line = locked config verbatim plus the two live deltas:
|
||||
|
||||
```bash
|
||||
SAC_OPPONENTS='Corners:3,Crazy:2,RamFire:2,Target:1,SacTwin:1' \
|
||||
SAC_EVAL_OPPONENTS='Corners,Crazy,Target' \
|
||||
SAC_EVAL_INTERVAL=2 SAC_EVAL_ROUNDS=10 \
|
||||
SAC_TOTAL_ROUNDS=25000 SAC_CHUNK_SIZE=10 SAC_MAX_CRASHES=5 \
|
||||
SACLSTM_HIDDEN_SIZE=256 SACLSTM_BATCH_SIZE=16 SACLSTM_SAVE_INTERVAL=1 \
|
||||
SACLSTM_LR_CRITIC=1e-4 \
|
||||
./sac_train.sh 2>&1 | tee -a campaign_v2_stdout.log
|
||||
```
|
||||
|
||||
- Safety net: existing transient unit `sac-ceiling-net-v2.service` re-verified active; fire epoch start+384587 s = **1787810380 = Thu 2026-08-27 07:59:40 CEST exactly** (delta 0 vs mandate) — kept, not duplicated.
|
||||
|
||||
### Health baseline (t+15 m)
|
||||
|
||||
| Check | Evidence |
|
||||
|-------|----------|
|
||||
| Counter | 33 (t+8m) → **125** (t+15m) |
|
||||
| Metrics JSONL | flowing (6 lines); baseline line: `critic_loss=23.09, actor_loss=-3.228, alpha_loss=0.0, alpha=0.99970` (random-init magnitudes — attempt-3's divergence curve starts here); t+15m line: critic 84.5, actor −9.99, alpha 0.9937 |
|
||||
| Saves | `SAVE_INTERVAL=1`: `sac_latest.zip` written from first grad-step onward |
|
||||
| Eval rotation | full cycle landed: Corners/Crazy/Target ×10 games each in `eval_log.jsonl` |
|
||||
| MA files | `[0 0 ]` per opponent after two cycles — fix confirmed in vivo |
|
||||
| Best gate | `best_score.txt=0.0000` written on first composite (0 > −1 default) |
|
||||
| Crash banners | 0 |
|
||||
|
||||
|
||||
|
||||
### Progress graphs
|
||||
|
||||
One live dashboard: `docs/campaign_dashboard.svg` (current run only — test wins, real-fight wins, losses, alpha, throughput; auto-reloads every 60 s when open in Chrome). Keep it fresh with `tools/watch_dashboard.sh` (regenerates every 60 s), or one-shot `python3 tools/plot_progress.py` (pure stdlib; paths overridable via argv, `--selftest` for sanity check).
|
||||
|
||||
## ~10:47 — the dashboard's axes were upside-down since creation
|
||||
|
||||
The progress graphs have been lying since they were made: 0% was drawn at the TOP of every panel and the newest games appeared on the LEFT. The cause is a one-line formula bug in `tools/plot_progress.py`: `map_fn` interpolated as `p1 - t*(p1-p0)` instead of `p0 + t*(p1-p0)`, so all five panels plotted `100 - value` on y and reversed time on x. Tick labels were computed by separate (correct) code, which is why the numbers on the axes never matched the ink.
|
||||
|
||||
Fix + guard: formula corrected; every call site audited (panels 1/2/5 use `map_fn` for both axes and are fixed by the same line; panels 3/4 already used a correct local x-lambda; no other consumer of `map_fn` exists in the repo). `--selftest` now renders a known rising series through the full build path and fails loudly unless higher value = smaller SVG y and newer data = further right — proven to catch this exact bug when the old formula is re-injected. Dashboard regenerated from live logs.
|
||||
|
||||
Honest status while reading the now-correct charts: run 3 is only hours old and winning ~0% — recent evals are 0/10 vs Corners, Crazy and Target alike, and real-fight buckets sit at 0–1 wins per 100 games. Expected for a fresh brain. Night-1's gains were real but were intentionally reset by the stability restart that began attempt-3; the curve starts from zero again here.
|
||||
Executable
+61
@@ -0,0 +1,61 @@
|
||||
#!/usr/bin/env bash
|
||||
# make_twin.sh — #54 mirror-twin sparring partner generator.
|
||||
#
|
||||
# Builds a self-contained `SacTwin` bot dir inside the sample-bots archive so
|
||||
# RunTraining.java resolves it like any sample bot ($SAMPLE_BOTS_DIR/<name>).
|
||||
# The twin is the SAME binary as SAC_LSTM_Bot but with:
|
||||
# - distinct identity (SacTwin.json; SACLSTM_BOT_JSON overrides the baked-in
|
||||
# src json — loadBotInfo gives json total precedence over env, #49 lesson)
|
||||
# - its OWN weights dir, seeded from a FROZEN copy of weights/sac_best.zip
|
||||
# (no checkpoint write races with the main bot)
|
||||
# - its OWN round_counter.txt (main liveness guard untouched)
|
||||
#
|
||||
# Re-running resets the twin to the frozen baseline (reproducible opponent).
|
||||
#
|
||||
# Usage: ./make_twin.sh [target-archive-dir]
|
||||
set -euo pipefail
|
||||
|
||||
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)"
|
||||
TARGET="${1:-${SAMPLE_BOTS_DIR:-/home/davide/Projects/tank-royale/sample-bots/java/build/archive}}/SacTwin"
|
||||
BIN="$SCRIPT_DIR/SAC_LSTM_Bot" # nimble build -d:release output
|
||||
SEED="$SCRIPT_DIR/weights/sac_best.zip" # frozen baseline
|
||||
|
||||
[ -x "$BIN" ] || { echo ">>> $BIN missing — run 'nimble build -d:release' first"; exit 1; }
|
||||
[ -f "$SEED" ] || { echo ">>> $SEED missing — need at least one eval'd checkpoint"; exit 1; }
|
||||
|
||||
mkdir -p "$TARGET/weights"
|
||||
cp "$BIN" "$TARGET/SAC_LSTM_Bot"
|
||||
|
||||
cat > "$TARGET/SacTwin.json" <<EOF
|
||||
{
|
||||
"name": "SacTwin",
|
||||
"version": "0.1.0",
|
||||
"authors": ["Davide Cappellini"],
|
||||
"description": "Mirror-twin sparring partner of SAC_LSTM_Bot (#54), generated by make_twin.sh — own weights dir, frozen seed",
|
||||
"homepage": "",
|
||||
"countryCodes": ["IT"],
|
||||
"gameTypes": ["classic", "melee", "1v1"],
|
||||
"platform": "Nim",
|
||||
"programmingLang": "Nim"
|
||||
}
|
||||
EOF
|
||||
|
||||
cat > "$TARGET/SacTwin.sh" <<'EOF'
|
||||
#!/bin/sh
|
||||
# Twin launcher (#54): identity + weights fully decoupled from the main bot.
|
||||
DIR="$(cd "$(dirname "$0")" && pwd)"
|
||||
export SACLSTM_BOT_JSON="$DIR/SacTwin.json"
|
||||
export SACLSTM_WEIGHTS_PATH="$DIR/weights/sac_latest.zip"
|
||||
# Freeze contract (#54): the twin is a FROZEN sparring partner, not a
|
||||
# co-learner. Eval-mode gate (lever 4) suppresses all sendTrainingMsg traffic,
|
||||
# so the twin never trains — not even in-RAM within a battle. Without this,
|
||||
# any main-bot checkpoint-interval change would let twin drift accumulate.
|
||||
export SACLSTM_EVAL_MODE=1
|
||||
exec "$DIR/SAC_LSTM_Bot"
|
||||
EOF
|
||||
chmod +x "$TARGET/SacTwin.sh" "$TARGET/SAC_LSTM_Bot"
|
||||
|
||||
cp "$SEED" "$TARGET/weights/sac_latest.zip"
|
||||
cp "$SEED" "$TARGET/weights/sac_best.zip"
|
||||
echo 0 > "$TARGET/weights/round_counter.txt"
|
||||
echo ">>> twin ready: $TARGET (seeded from $(stat -c %y "$SEED" | cut -d. -f1) snapshot of sac_best.zip)"
|
||||
Executable
+181
@@ -0,0 +1,181 @@
|
||||
#!/usr/bin/env bash
|
||||
# sac_train.sh — #49 training orchestration for SAC_LSTM_Bot.
|
||||
#
|
||||
# Drives chunked self-play via tools/training_runner/RunTraining.java (which
|
||||
# owns server lifecycle, opponent connection and dead-bot liveness detection
|
||||
# through weights/round_counter.txt), samples opponents by weight per chunk,
|
||||
# runs deterministic evaluation (SACLSTM_EVAL_MODE=1) every N chunks, and keeps
|
||||
# the best checkpoint (weights/sac_best.zip) by a moving-average composite over
|
||||
# the eval opponent set (campaign v2 lever 1, #59).
|
||||
#
|
||||
# Config (env vars):
|
||||
# SAC_OPPONENTS "Name:weight,Name:weight,..." (default below)
|
||||
# SAC_TOTAL_ROUNDS total training-round budget (default 100)
|
||||
# SAC_CHUNK_SIZE rounds per RunTraining battle (default 10)
|
||||
# SAC_EVAL_INTERVAL eval every N chunks (default 2)
|
||||
# SAC_EVAL_ROUNDS rounds per evaluation battle (default 10)
|
||||
# SAC_EVAL_OPPONENTS comma-separated eval set (default Corners,Crazy,Target)
|
||||
# — each cycle evaluates EVERY one; results all land in
|
||||
# eval_log.jsonl (lines carry "opponent":"Name")
|
||||
# SAC_MAX_CRASHES consecutive crashes before abort (default 5)
|
||||
# SAC_LOG_FILE / SAC_EVAL_LOG_FILE (JSON-lines logs)
|
||||
# SACLSTM_* passed through to the bot (UTD_RATIO, BATCH_SIZE, ...)
|
||||
set -uo pipefail
|
||||
|
||||
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)"
|
||||
REPO_ROOT="$(cd "$SCRIPT_DIR/.." && pwd)"
|
||||
RUNNER_DIR="$REPO_ROOT/tools/training_runner"
|
||||
JAR="${TANK_ROYALE_JAR:-/home/davide/Projects/tank-royale/runner/examples/lib/robocode-tankroyale-runner.jar}"
|
||||
|
||||
export PPO_BOT_DIR="$SCRIPT_DIR" # runner launches THIS bot dir
|
||||
export BOT_NAME="${BOT_NAME:-SAC_LSTM_Bot}" # RunTraining result matching
|
||||
export SAMPLE_BOTS_DIR="${SAMPLE_BOTS_DIR:-/home/davide/Projects/tank-royale/sample-bots/java/build/archive}"
|
||||
# Liveness contract (#49): RunTraining.java watches $BOT_DIR/weights/round_counter.txt
|
||||
# and integration.bumpRoundCounter() writes it next to the weights — so the bot's
|
||||
# weights path is pinned here, NOT env-overridable.
|
||||
export SACLSTM_WEIGHTS_PATH="$SCRIPT_DIR/weights/sac_latest.zip"
|
||||
WEIGHTS_DIR="$(dirname "$SACLSTM_WEIGHTS_PATH")"
|
||||
|
||||
OPPONENTS="${SAC_OPPONENTS:-Corners:3,Crazy:2,RamFire:1,Target:1}"
|
||||
TOTAL_ROUNDS="${SAC_TOTAL_ROUNDS:-100}"
|
||||
CHUNK_SIZE="${SAC_CHUNK_SIZE:-10}"
|
||||
EVAL_INTERVAL="${SAC_EVAL_INTERVAL:-2}"
|
||||
EVAL_ROUNDS="${SAC_EVAL_ROUNDS:-10}"
|
||||
EVAL_OPPONENTS="${SAC_EVAL_OPPONENTS:-Corners,Crazy,Target}"
|
||||
MA_WINDOW=5 # lever 1 (#59): per-opponent moving average over last N evals
|
||||
MAX_CRASHES="${SAC_MAX_CRASHES:-5}"
|
||||
LOG_FILE="${SAC_LOG_FILE:-$SCRIPT_DIR/training_log.jsonl}"
|
||||
EVAL_LOG_FILE="${SAC_EVAL_LOG_FILE:-$SCRIPT_DIR/eval_log.jsonl}"
|
||||
CLASSES_DIR="/tmp/opencode/sac_train_classes"
|
||||
|
||||
echo "=== SAC_LSTM_Bot training harness ==="
|
||||
echo "Opponents: $OPPONENTS | budget: $TOTAL_ROUNDS rounds in chunks of $CHUNK_SIZE"
|
||||
echo "Eval: every $EVAL_INTERVAL chunks, $EVAL_ROUNDS rounds vs [$EVAL_OPPONENTS], MA-$MA_WINDOW composite best-gating"
|
||||
echo "Weights: $SACLSTM_WEIGHTS_PATH"
|
||||
|
||||
# ── compile bot + java runner ─────────────────────────────────────────────────
|
||||
(cd "$SCRIPT_DIR" && nimble build -d:release) || { echo ">>> bot build failed"; exit 1; }
|
||||
mkdir -p "$WEIGHTS_DIR" "$CLASSES_DIR"
|
||||
# Startup sweep: atomic saves leave sac_latest.zip.tmp.*.part corpses behind if
|
||||
# a run is SIGKILLed (libzip modify-path, see notebook forensics) — clear them.
|
||||
rm -f "$WEIGHTS_DIR"/sac_latest.zip.tmp.*.part "$WEIGHTS_DIR"/sac_latest.zip.tmp
|
||||
javac -cp "$JAR" -d "$CLASSES_DIR" "$RUNNER_DIR/RunTraining.java" || { echo ">>> javac failed"; exit 1; }
|
||||
|
||||
# ── weighted opponent pick over "Name:w,Name:w" ───────────────────────────────
|
||||
pick_opponent() {
|
||||
local total=0 p name w r
|
||||
local pairs
|
||||
IFS=',' read -ra pairs <<< "$OPPONENTS"
|
||||
for p in "${pairs[@]}"; do total=$(( total + ${p##*:} )); done
|
||||
r=$(( RANDOM % total ))
|
||||
for p in "${pairs[@]}"; do
|
||||
name="${p%%:*}"; w="${p##*:}"
|
||||
if (( r < w )); then echo "$name"; return; fi
|
||||
r=$(( r - w ))
|
||||
done
|
||||
echo "${pairs[0]%%:*}"
|
||||
}
|
||||
|
||||
run_battle() { # $1=opponent $2=rounds $3=log file
|
||||
PPOB_LOG_FILE="$3" java -cp "$CLASSES_DIR:$JAR" RunTraining "$1" "$2"
|
||||
}
|
||||
|
||||
# ── Lever 1 (#59): eval rotation + MA best-gating ─────────────────────────────
|
||||
# best_score.txt FORMAT CHANGE: it used to store the single-opponent integer
|
||||
# win rate (%); that semantics is retired. It now stores the COMPOSITE score —
|
||||
# the mean over SAC_EVAL_OPPONENTS of each opponent's moving average (last
|
||||
# MA_WINDOW eval win rates, %). sac_best.zip is rewritten only when the
|
||||
# composite strictly improves.
|
||||
ma_hist_file() { echo "$WEIGHTS_DIR/ma_history_$1.txt"; }
|
||||
|
||||
composite_of() { # reads one "w w w ..." history line per opponent on stdin
|
||||
awk -v W="$MA_WINDOW" '
|
||||
NF > 0 { n=NF; k=(n>W)?W:n; s=0; for(j=n-k+1;j<=n;j++) s+=$j; tot+=s/k; c++ }
|
||||
END { if (c>0) printf "%.4f", tot/c; else print "-1" }'
|
||||
}
|
||||
|
||||
eval_checkpoint() {
|
||||
# ponytail: opponent names are split by whitespace — fine for Tank Royale bot
|
||||
# names (no spaces); switch to a mapfile IFS=',\n' read if that ever changes.
|
||||
local opps=(${EVAL_OPPONENTS//,/ })
|
||||
local tmp="$EVAL_LOG_FILE.tmp" otmp opp wins rounds wr composite best
|
||||
: > "$tmp"
|
||||
for opp in "${opps[@]}"; do
|
||||
otmp="$EVAL_LOG_FILE.$opp.tmp"
|
||||
: > "$otmp"
|
||||
echo ">>> [eval] $EVAL_ROUNDS deterministic rounds vs $opp"
|
||||
if ! SACLSTM_EVAL_MODE=1 run_battle "$opp" "$EVAL_ROUNDS" "$otmp"; then
|
||||
rm -f "$otmp" "$tmp"
|
||||
echo ">>> [eval] crashed vs $opp — keeping previous best"
|
||||
return 0
|
||||
fi
|
||||
wins=$(grep -c '"win":true' "$otmp" || true)
|
||||
rounds=$(grep -c '"type":"game"' "$otmp" || true)
|
||||
if (( rounds == 0 )); then
|
||||
rm -f "$otmp" "$tmp"
|
||||
echo ">>> [eval] no results vs $opp — keeping previous best"
|
||||
return 0
|
||||
fi
|
||||
wr=$(( 100 * wins / rounds ))
|
||||
echo ">>> [eval] win rate: $wins/$rounds ($wr%) vs $opp"
|
||||
cat "$otmp" >> "$tmp"; rm -f "$otmp"
|
||||
# Per-opponent history: append this cycle's win rate, keep last MA_WINDOW
|
||||
# values on one line. Read FIRST, separately: `$(cat f)` inside a command
|
||||
# redirected `> f` executes against the already-truncated file (bash sets
|
||||
# up the redirect before running the substitution) — every cycle wiped the
|
||||
# history back to a single leading-space value (#60). $hist is UNQUOTED on
|
||||
# purpose: word-splitting turns the stored line into one value per line so
|
||||
# tail keeps the last MA_WINDOW values.
|
||||
local hist
|
||||
hist="$(cat "$(ma_hist_file "$opp")" 2>/dev/null)"
|
||||
printf '%s\n' $hist "$wr" \
|
||||
| tail -n "$MA_WINDOW" | tr '\n' ' ' > "$(ma_hist_file "$opp")"
|
||||
done
|
||||
mv "$tmp" "$EVAL_LOG_FILE"
|
||||
composite=$(for opp in "${opps[@]}"; do cat "$(ma_hist_file "$opp")"; echo; done | composite_of)
|
||||
# ponytail: best-score state is a plain file next to the checkpoint; survives
|
||||
# harness restarts, no lock needed (single harness instance assumed).
|
||||
best=$(cat "$WEIGHTS_DIR/best_score.txt" 2>/dev/null)
|
||||
[ -z "$best" ] && best=-1
|
||||
if awk -v a="$composite" -v b="$best" 'BEGIN{exit !(a+0 > b+0)}' \
|
||||
&& [ -f "$SACLSTM_WEIGHTS_PATH" ]; then
|
||||
echo "$composite" > "$WEIGHTS_DIR/best_score.txt"
|
||||
cp "$SACLSTM_WEIGHTS_PATH" "$WEIGHTS_DIR/sac_best.zip"
|
||||
echo ">>> [eval] new best composite ($composite) -> sac_best.zip"
|
||||
fi
|
||||
}
|
||||
|
||||
NUM_CHUNKS=$(( (TOTAL_ROUNDS + CHUNK_SIZE - 1) / CHUNK_SIZE ))
|
||||
fails=0
|
||||
chunk=1
|
||||
# while, not `for chunk in $(seq ...)`: a crash on the FINAL chunk must rerun
|
||||
# it (#54 — seq list is exhausted by then, so ((chunk--));continue fell through
|
||||
# and the harness exited 0 with the budget incomplete).
|
||||
while (( chunk <= NUM_CHUNKS )); do
|
||||
ROUNDS=$CHUNK_SIZE
|
||||
(( TOTAL_ROUNDS - (chunk - 1) * CHUNK_SIZE < CHUNK_SIZE )) && \
|
||||
ROUNDS=$(( TOTAL_ROUNDS - (chunk - 1) * CHUNK_SIZE ))
|
||||
OPP=$(pick_opponent)
|
||||
echo "=== Chunk $chunk/$NUM_CHUNKS: $ROUNDS rounds vs $OPP ==="
|
||||
if ! run_battle "$OPP" "$ROUNDS" "$LOG_FILE"; then
|
||||
fails=$(( fails + 1 ))
|
||||
if (( fails >= MAX_CRASHES )); then
|
||||
echo ">>> aborted: $fails consecutive crashes (bot process dying?)"
|
||||
exit 1
|
||||
fi
|
||||
# Crash recovery: RunTraining's liveness detection exited; the bot reloads
|
||||
# its latest checkpoint on restart, so just rerun this chunk.
|
||||
echo ">>> crash #$fails — restarting chunk from latest checkpoint"
|
||||
(( chunk-- )); continue
|
||||
fi
|
||||
fails=0
|
||||
(( chunk % EVAL_INTERVAL == 0 )) && eval_checkpoint
|
||||
((chunk += 1))
|
||||
done
|
||||
|
||||
echo ">>> training complete: $NUM_CHUNKS chunks. Logs:"
|
||||
echo " training: $LOG_FILE"
|
||||
echo " eval: $EVAL_LOG_FILE"
|
||||
[ -f "$WEIGHTS_DIR/sac_best.zip" ] && \
|
||||
echo " best: $WEIGHTS_DIR/sac_best.zip (composite $(cat "$WEIGHTS_DIR/best_score.txt"))"
|
||||
exit 0
|
||||
@@ -0,0 +1,11 @@
|
||||
{
|
||||
"name": "SAC_LSTM_Bot",
|
||||
"version": "0.1.0",
|
||||
"authors": ["Davide Cappellini"],
|
||||
"description": "SAC+LSTM Tank Royale bot — self-reported identity; MUST match the name in ../SAC_LSTM_Bot.json (booter identity) or the training runner never sees this bot join",
|
||||
"homepage": "",
|
||||
"countryCodes": ["IT"],
|
||||
"gameTypes": ["classic", "melee", "1v1"],
|
||||
"platform": "Nim",
|
||||
"programmingLang": "Nim"
|
||||
}
|
||||
@@ -0,0 +1,349 @@
|
||||
## SAC_LSTM_Bot — Recurrent SAC-v2 bot (Gitea #48).
|
||||
##
|
||||
## Thread layout (decisions Q1–Q14, see integration.nim for plumbing):
|
||||
## bot thread — this file's run(): inference only, <2ms/tick.
|
||||
## training thread — permanent background SAC updates (integration.nim).
|
||||
## I/O thread — atomic weight saves (integration.nim).
|
||||
## No Arraymancer tensor ever crosses a thread boundary: the bot object and
|
||||
## channels carry plain scalars / fixed arrays / plain seqs only.
|
||||
|
||||
import std/[os, math, random, algorithm]
|
||||
import arraymancer except Linear
|
||||
import tankroyale_botapi
|
||||
import radar_lock
|
||||
import SAC_LSTM_Bot/state
|
||||
import SAC_LSTM_Bot/network
|
||||
import SAC_LSTM_Bot/actions
|
||||
import SAC_LSTM_Bot/rewards
|
||||
import SAC_LSTM_Bot/integration
|
||||
|
||||
# Identity json: baked-in src json by default; SACLSTM_BOT_JSON lets a mirror
|
||||
# twin (#54) boot the same binary under its own name (loadBotInfo gives the
|
||||
# json total precedence over env, so the twin must point at its own file).
|
||||
let botJsonPath = getEnv("SACLSTM_BOT_JSON",
|
||||
currentSourcePath().parentDir / "SAC_LSTM_Bot.json")
|
||||
|
||||
# ── Colors (Recurrent Royalty palette) ───────────────────────────────────────
|
||||
|
||||
const
|
||||
ColBody = fromHex("#7B2FBE")
|
||||
ColTurret = fromHex("#FFD700")
|
||||
ColGun = fromHex("#4A0E6B")
|
||||
ColRadar = fromHex("#FFD700")
|
||||
ColScan = fromHex("#FFB000")
|
||||
ColBullet = fromHex("#FFC125")
|
||||
ColTracks = fromHex("#2C2C34")
|
||||
|
||||
proc applyColors() =
|
||||
setBodyColor(ColBody)
|
||||
setTurretColor(ColTurret)
|
||||
setGunColor(ColGun)
|
||||
setRadarColor(ColRadar)
|
||||
setScanColor(ColScan)
|
||||
setBulletColor(ColBullet)
|
||||
setTracksColor(ColTracks)
|
||||
|
||||
# ── Enemy bullet tracking ────────────────────────────────────────────────────
|
||||
# PPO_Bot-proven pattern: dead-reckoned fixed buffer. getBulletStates() is NOT
|
||||
# used from the bot thread — its seq refcount is shared with the main thread.
|
||||
|
||||
type InFlightBullet = object
|
||||
x, y, vx, vy, power: float64
|
||||
|
||||
const MaxBotBullets = 4
|
||||
const MaxHidden = 512 # ponytail: cap for the plain-array LSTM persistence; raise if SACLSTM_HIDDEN_SIZE > 512
|
||||
|
||||
# ── Bot type — PLAIN DATA ONLY on the shared object (no tensors/heap seqs:
|
||||
# each round runs a fresh bot thread; heap blocks owned by the previous
|
||||
# round's thread must not be freed from another thread) ─────────────────────
|
||||
|
||||
type SacBot = ref object of Bot
|
||||
enemyBearing: float # last known absolute bearing to enemy
|
||||
battleId: int # main thread bumps in onGameStarted; bot thread compares
|
||||
seenBattle: int # bot-thread copy for battle-change detection
|
||||
newBattleSent: bool # first scan of THIS battle emits NewBattle
|
||||
hasContact: bool
|
||||
enemy: EnemyData
|
||||
ticksSinceScan: int
|
||||
# per-step reward accumulators (consumed by the next tick's transition)
|
||||
dmgDealt, dmgTaken, wastedPower: float64
|
||||
wallHits, hits, ramTaken: int
|
||||
# pending transition (episode spans the whole battle; round end is NOT a boundary)
|
||||
hasLastTrans: bool
|
||||
lastState: array[STATE_DIM, float32]
|
||||
lastAction: array[ACTION_DIM, float32]
|
||||
rn: RewardNormalizer # Welford running stats, persists across battles
|
||||
bullets: array[MaxBotBullets, InFlightBullet]
|
||||
bulletCount: int
|
||||
hArr, cArr: array[MaxHidden, float32] # LSTM state across rounds; zeros at battle start
|
||||
|
||||
# ── Plain-array <-> tensor helpers (bot thread only) ──────────────────────────
|
||||
|
||||
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]
|
||||
|
||||
proc hiddenToTensor(arr: array[MaxHidden, float32]; n: int): Tensor[float32] =
|
||||
result = newTensor[float32](n)
|
||||
for i in 0 ..< n: result[i] = arr[i]
|
||||
|
||||
proc tensorToHidden(t: Tensor[float32]; arr: var array[MaxHidden, float32]) =
|
||||
for i in 0 ..< t.shape[0]: arr[i] = t[i]
|
||||
|
||||
# ── Reward ────────────────────────────────────────────────────────────────────
|
||||
|
||||
proc takeReward(bot: SacBot; win = false, loss = false): float32 =
|
||||
## Consume accumulated step events -> Welford-normalized reward (#44).
|
||||
## Lever 2 (#59): pass enemy distance (frac of arena diagonal) for the
|
||||
## anti-charge term; sentinel 2.0 (> ChargeDistFrac) when no contact.
|
||||
var distFrac = 2.0
|
||||
if bot.hasContact:
|
||||
let diag = hypot(getArenaWidth().float64, getArenaHeight().float64)
|
||||
distFrac = hypot(bot.enemy.x - getX(), bot.enemy.y - getY()) / diag
|
||||
let raw = computeReward(
|
||||
damageInflicted = bot.dmgDealt,
|
||||
damageReceived = bot.dmgTaken,
|
||||
wallHitTicks = bot.wallHits,
|
||||
wastedShotPower = bot.wastedPower,
|
||||
hitCount = bot.hits,
|
||||
ramTakenCount = bot.ramTaken,
|
||||
enemyDistFrac = distFrac,
|
||||
win = win, loss = loss)
|
||||
# Lever-2 (#59) observability: env-gated one-liner for smoke/calibration
|
||||
# greps — proves hit/ram/charge terms fire and shows raw magnitudes. File
|
||||
# (not stderr): the battle runner swallows bot process streams.
|
||||
# ponytail: grows unbounded if left on; keep off outside smokes.
|
||||
if getEnv("SACLSTM_REWARD_DEBUG") == "1" and
|
||||
(bot.hits > 0 or bot.ramTaken > 0 or (distFrac < ChargeDistFrac and bot.dmgDealt <= 0.0)):
|
||||
try:
|
||||
let f = open(getWeightsPath().parentDir.parentDir / "reward_debug.log", fmAppend)
|
||||
f.writeLine("raw=" & $raw & " hits=" & $bot.hits & " ram=" & $bot.ramTaken &
|
||||
" distFrac=" & $distFrac)
|
||||
f.close()
|
||||
except CatchableError:
|
||||
discard
|
||||
bot.dmgDealt = 0; bot.dmgTaken = 0; bot.wastedPower = 0; bot.wallHits = 0
|
||||
bot.hits = 0; bot.ramTaken = 0
|
||||
let norm = rewards.normalize(bot.rn, raw)
|
||||
rewards.update(bot.rn, raw)
|
||||
norm.float32
|
||||
|
||||
# ── Event handlers ────────────────────────────────────────────────────────────
|
||||
# onGameStarted/onRoundStarted fire on the MAIN thread (bot thread not yet
|
||||
# started or already joined) — plain-field writes only, no tensors here.
|
||||
|
||||
method onGameStarted*(bot: SacBot, e: GameStartedEventForBot) =
|
||||
inc bot.battleId # bot thread zeroes LSTM + per-battle flags at next tick (Q4/Q12)
|
||||
|
||||
method onRoundStarted*(bot: SacBot, e: RoundStartedEvent) =
|
||||
setAdjustRadarForBodyTurn(true)
|
||||
setAdjustRadarForGunTurn(true)
|
||||
radar_lock.init()
|
||||
applyColors()
|
||||
# Per-round reset ONLY. NOT the LSTM hidden state (persists across rounds, Q4);
|
||||
# NOT hasLastTrans (the pending transition spans the round boundary — episode
|
||||
# ends at battle end only).
|
||||
bot.hasContact = false
|
||||
bot.enemy = EnemyData()
|
||||
bot.ticksSinceScan = 0
|
||||
bot.bulletCount = 0
|
||||
|
||||
method onScannedBot*(bot: SacBot, e: ScannedBotEvent) =
|
||||
bot.enemyBearing = directionTo(getX(), getY(), e.x, e.y)
|
||||
# Same-tick radar lock: apply turn rate immediately so it takes effect this tick.
|
||||
setRadarTurnRate(radar_lock.doRadar(getRadarDirection(), bot.enemyBearing))
|
||||
# Fire detection: energy drop in [0.1, 3.0] between scans (PPO_Bot heuristic).
|
||||
let prevE = if bot.hasContact: bot.enemy.energy else: e.energy
|
||||
let drop = prevE - e.energy
|
||||
bot.enemy.hasFired = bot.hasContact and drop >= 0.1 and drop <= 3.0
|
||||
if bot.enemy.hasFired:
|
||||
bot.enemy.lastFirePower = drop
|
||||
# Keep previous-scan deltas before overwriting (state.nim derives accel/turn rate).
|
||||
bot.enemy.prevSpeed = bot.enemy.speed
|
||||
bot.enemy.prevDirection = bot.enemy.direction
|
||||
bot.enemy.hasPrevScan = bot.hasContact
|
||||
bot.enemy.x = e.x
|
||||
bot.enemy.y = e.y
|
||||
bot.enemy.direction = e.direction
|
||||
bot.enemy.speed = e.speed
|
||||
bot.enemy.energy = e.energy
|
||||
bot.hasContact = true
|
||||
bot.ticksSinceScan = 0
|
||||
# Q12b/Q14+#49: one NewBattle per battle, numeric scannedBotId. Identity is
|
||||
# keyed on getBotName(id) training-side (integration.opponentKey), numeric
|
||||
# fallback in the pre-BotListUpdate window.
|
||||
if not bot.newBattleSent:
|
||||
bot.newBattleSent = true
|
||||
discard sendTrainingMsg(TrainingMsg(kind: tmkNewBattle, enemyId: e.scannedBotId))
|
||||
|
||||
method onBulletHit*(bot: SacBot, e: BulletHitBotEvent) =
|
||||
bot.dmgDealt += e.damage
|
||||
inc bot.hits
|
||||
|
||||
method onHitByBullet*(bot: SacBot, e: HitByBulletEvent) =
|
||||
bot.dmgTaken += e.damage
|
||||
|
||||
# Lever 2 (#59): anti-ram — every BotHitBotEvent receipt means a bot-bot
|
||||
# collision happened and we ate RAM_DAMAGE (server deals 0.6 to both parties;
|
||||
# only the hitter gets notified). Flat per-event penalty; being rammed without
|
||||
# hitting back stays event-invisible.
|
||||
# ponytail: enemy-initiated rams undetected — add energy-residual detection if
|
||||
# v2 battle data shows ram-heavy losses.
|
||||
method onHitBot*(bot: SacBot, e: BotHitBotEvent) =
|
||||
inc bot.ramTaken
|
||||
|
||||
method onHitWall*(bot: SacBot, e: BotHitWallEvent) =
|
||||
inc bot.wallHits
|
||||
|
||||
method onBulletHitWall*(bot: SacBot, e: BulletHitWallEvent) =
|
||||
if e.bullet.ownerId == getMyId():
|
||||
bot.wastedPower += e.bullet.power
|
||||
|
||||
method onGameAborted*(bot: SacBot) =
|
||||
# Mid-round abort: drop the pending transition rather than leak it into the
|
||||
# next battle's data.
|
||||
bot.hasLastTrans = false
|
||||
bot.dmgDealt = 0; bot.dmgTaken = 0; bot.wastedPower = 0; bot.wallHits = 0
|
||||
bot.hits = 0; bot.ramTaken = 0
|
||||
|
||||
method onRoundEnded*(bot: SacBot, e: RoundEndedEventForBot) =
|
||||
## Harness liveness signal (#49): RunTraining.java watches round_counter.txt
|
||||
## and aborts the battle if it freezes (dead bot process). Main thread.
|
||||
bumpRoundCounter()
|
||||
|
||||
method onGameEnded*(bot: SacBot, e: GameEndedEventForBot) =
|
||||
## Battle end -> terminal transition with done=true. Main thread; the API has
|
||||
## joined the bot thread before this fires, so these plain fields are quiescent.
|
||||
if bot.hasLastTrans:
|
||||
let r = bot.takeReward(win = e.results.rank == 1, loss = e.results.rank != 1)
|
||||
discard sendTrainingMsg(TrainingMsg(kind: tmkTransition,
|
||||
state: bot.lastState, action: bot.lastAction,
|
||||
reward: r, nextState: bot.lastState, done: true))
|
||||
bot.hasLastTrans = false
|
||||
|
||||
# ── Run loop (bot thread) ─────────────────────────────────────────────────────
|
||||
|
||||
method run(bot: SacBot) =
|
||||
randomize()
|
||||
# Per-thread locals: born and freed on THIS thread, every round. Nothing
|
||||
# heap-owned survives the round boundary except the bot object's plain fields.
|
||||
var actor: ActorNet
|
||||
var actorReady = false
|
||||
var myVersion = 0
|
||||
var myHidden = 0
|
||||
var flat: seq[float32]
|
||||
|
||||
while isRunning():
|
||||
# Battle boundary (Q4/Q12): zero LSTM persistence + per-battle flags.
|
||||
if bot.battleId != bot.seenBattle:
|
||||
bot.seenBattle = bot.battleId
|
||||
bot.newBattleSent = false
|
||||
zeroMem(addr bot.hArr, sizeof(bot.hArr))
|
||||
zeroMem(addr bot.cArr, sizeof(bot.cArr))
|
||||
|
||||
# Weight sync (Q7/Q11): always-latest; rebuild this thread's tensors on change.
|
||||
if pullWeights(myVersion, myHidden, flat):
|
||||
if myHidden > MaxHidden:
|
||||
# hArr/cArr are fixed-capacity; a bigger SACLSTM_HIDDEN_SIZE would
|
||||
# heap-overflow them in tensorToHidden. Loud misconfig beats corruption.
|
||||
raise newException(ValueError, "SACLSTM_HIDDEN_SIZE=" & $myHidden &
|
||||
" exceeds MaxHidden=" & $MaxHidden & " (bot-side LSTM persistence cap)")
|
||||
var cur = 0
|
||||
actor = actorFromFlat(flat, cur, myHidden)
|
||||
actorReady = true
|
||||
|
||||
# Spawn an enemy bullet when a fresh scan shows they fired; then advance and
|
||||
# prune the tracked bullets (positions feed state slots 22–33).
|
||||
if bot.hasContact and bot.enemy.hasFired and bot.bulletCount < MaxBotBullets:
|
||||
let p = bot.enemy.lastFirePower
|
||||
let spd = 20.0 - 3.0 * p
|
||||
let ang = arctan2(getY() - bot.enemy.y, getX() - bot.enemy.x)
|
||||
bot.bullets[bot.bulletCount] = InFlightBullet(x: bot.enemy.x, y: bot.enemy.y,
|
||||
vx: spd * cos(ang), vy: spd * sin(ang), power: p)
|
||||
inc bot.bulletCount
|
||||
let aW = float64(getArenaWidth())
|
||||
let aH = float64(getArenaHeight())
|
||||
var alive = 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[alive] = InFlightBullet(x: nx, y: ny, vx: b.vx, vy: b.vy, power: b.power)
|
||||
inc alive
|
||||
bot.bulletCount = alive
|
||||
|
||||
# State build (35-dim, this thread's tensor).
|
||||
inc bot.ticksSinceScan
|
||||
var gs: GameState
|
||||
gs.x = getX()
|
||||
gs.y = getY()
|
||||
gs.direction = getDirection()
|
||||
gs.speed = getSpeed()
|
||||
gs.energy = getEnergy()
|
||||
gs.gunDirection = getGunDirection()
|
||||
gs.gunHeat = getGunHeat()
|
||||
gs.arenaWidth = aW
|
||||
gs.arenaHeight = aH
|
||||
gs.hasContact = bot.hasContact
|
||||
gs.enemy = bot.enemy
|
||||
gs.ticksSinceLastScan = bot.ticksSinceScan
|
||||
var bd: array[MaxBotBullets, BulletData]
|
||||
for i in 0 ..< bot.bulletCount:
|
||||
bd[i] = BulletData(x: bot.bullets[i].x, y: bot.bullets[i].y, power: bot.bullets[i].power)
|
||||
if bot.bulletCount > 1: # closest threats fill slots 0-2
|
||||
bd.toOpenArray(0, bot.bulletCount - 1).sort(proc(a, b: BulletData): int =
|
||||
cmp(hypot(a.x - gs.x, a.y - gs.y), hypot(b.x - gs.x, b.y - gs.y)))
|
||||
gs.bulletCount = min(bot.bulletCount, 3)
|
||||
for i in 0 ..< gs.bulletCount:
|
||||
gs.bullets[i] = bd[i]
|
||||
# Consume the fired pulse AFTER the state saw it (one shot -> one bullet).
|
||||
bot.enemy.hasFired = false
|
||||
|
||||
let stateT = buildState(gs)
|
||||
|
||||
# Finalize the PREVIOUS transition: reward from events since the last tick,
|
||||
# nextState is this tick's observation (PPO_Bot alignment).
|
||||
if actorReady and bot.hasLastTrans:
|
||||
discard sendTrainingMsg(TrainingMsg(kind: tmkTransition,
|
||||
state: bot.lastState, action: bot.lastAction,
|
||||
reward: bot.takeReward(), nextState: stateToArr(stateT), done: false))
|
||||
bot.lastState = stateToArr(stateT)
|
||||
|
||||
if not actorReady:
|
||||
setRadarTurnRate(45.0) # no weights yet (defensive; main pre-inits) — just sweep
|
||||
go()
|
||||
continue
|
||||
|
||||
# Inference: hidden state lives as plain arrays on the bot object (persists
|
||||
# across rounds); tensors are rebuilt per tick on this thread.
|
||||
let h = hiddenToTensor(bot.hArr, myHidden)
|
||||
let c = hiddenToTensor(bot.cArr, myHidden)
|
||||
let fwd = actor.actorForward(stateT, (h: h, c: c))
|
||||
tensorToHidden(fwd.lstm.h, bot.hArr)
|
||||
tensorToHidden(fwd.lstm.c, bot.cArr)
|
||||
|
||||
bot.lastAction = actionToArr(fwd.actions)
|
||||
bot.hasLastTrans = true
|
||||
|
||||
# Actions -> intents (go() snapshots them at send time).
|
||||
let mapped = mapActions(fwd.actions, getSpeed(), getGunHeat())
|
||||
setTurnRate(mapped.turnRate)
|
||||
setTargetSpeed(getSpeed() + mapped.acceleration) # actions.nim contract
|
||||
setGunTurnRate(mapped.gunTurnRate)
|
||||
if mapped.firePower > 0.0:
|
||||
discard setFire(mapped.firePower)
|
||||
if not bot.hasContact:
|
||||
setRadarTurnRate(45.0) # sweep until first lock (onScannedBot overrides same-tick)
|
||||
|
||||
go()
|
||||
|
||||
# ── Entry point ───────────────────────────────────────────────────────────────
|
||||
|
||||
when isMainModule:
|
||||
initIntegration() # spawn training + I/O threads, seed weight snapshot
|
||||
var bot = SacBot()
|
||||
start(bot, botJsonPath) # blocks until server disconnect
|
||||
shutdownIntegration() # Shutdown msg -> final save -> joins
|
||||
@@ -0,0 +1,36 @@
|
||||
## actions.nim — map raw network output (4 tanh values) to bot intent fields.
|
||||
##
|
||||
## Note on acceleration vs targetSpeed:
|
||||
## TankRoyale uses setTargetSpeed(), not setAcceleration().
|
||||
## The mapped `acceleration` field is a delta; callers must compute:
|
||||
## newTargetSpeed = clamp(currentSpeed + acceleration, -8.0, 8.0)
|
||||
## and call setTargetSpeed(newTargetSpeed).
|
||||
|
||||
import arraymancer
|
||||
|
||||
const ACTION_DIM* = 4
|
||||
|
||||
type
|
||||
MappedActions* = object
|
||||
turnRate*: float ## degrees/tick, speed-aware; [-10, 10] at speed 0
|
||||
acceleration*: float ## delta speed in [-2, +1]; caller adds to currentSpeed
|
||||
gunTurnRate*: float ## degrees/tick in [-20, 20]
|
||||
firePower*: float ## 0 = don't fire; (0.1, 3.0] = fire with this power
|
||||
|
||||
proc mapActions*(networkOutput: Tensor[float32],
|
||||
currentSpeed: float,
|
||||
gunHeat: float): MappedActions =
|
||||
## networkOutput: [4] tensor of tanh values in [-1, 1].
|
||||
let a0 = networkOutput[0].float
|
||||
let a1 = networkOutput[1].float
|
||||
let a2 = networkOutput[2].float
|
||||
let a3 = networkOutput[3].float
|
||||
|
||||
result.turnRate = a0 * (10.0 - 0.75 * abs(currentSpeed))
|
||||
# asymmetric accel: [-1,1] -> [-2, +1] via (value * 1.5 - 0.5)
|
||||
result.acceleration = a1 * 1.5 - 0.5
|
||||
result.gunTurnRate = a2 * 20.0
|
||||
if a3 > 0.0 and gunHeat <= 0.0:
|
||||
result.firePower = a3 * 2.9 + 0.1
|
||||
else:
|
||||
result.firePower = 0.0
|
||||
@@ -0,0 +1,463 @@
|
||||
## integration.nim — #48 thread plumbing for SAC_LSTM_Bot.
|
||||
##
|
||||
## Threads added by the bot (on top of the bot API's main/bot/sender):
|
||||
## training thread — permanent, drain-then-train loop, owns SACTrainer +
|
||||
## ReplayBuffer. Q1/Q2/Q10/Q12.
|
||||
## I/O thread — cap-1 channel of weight snapshots, atomic zip saves. Q5.
|
||||
##
|
||||
## Cross-thread payloads are plain arrays/seqs ONLY. No Arraymancer tensor ever
|
||||
## crosses a thread boundary: Tensor is a ref type and ORC refcounts are
|
||||
## non-atomic — sharing them across threads is SIGSEGV territory (PPO_Bot,
|
||||
## empirically confirmed). Each thread builds its own tensors from plain data.
|
||||
## Decisions Q1–Q14: Gitea #48.
|
||||
|
||||
import arraymancer except Linear
|
||||
import std/[locks, os, math, random, strutils, times]
|
||||
import tankroyale_botapi # getBotName (#49 name-based opponent identity)
|
||||
import SAC_LSTM_Bot/network
|
||||
import SAC_LSTM_Bot/state # STATE_DIM
|
||||
import SAC_LSTM_Bot/actions # ACTION_DIM
|
||||
import SAC_LSTM_Bot/replay_buffer
|
||||
import SAC_LSTM_Bot/training
|
||||
import SAC_LSTM_Bot/weights
|
||||
|
||||
# ── Config ────────────────────────────────────────────────────────────────────
|
||||
|
||||
proc getUtdRatio*(): int =
|
||||
## Gradient steps per drained transition (Q2). 200 ticks -> 200 steps at 1.
|
||||
parseInt(getEnv("SACLSTM_UTD_RATIO", "1"))
|
||||
|
||||
proc getBatchSize*(): int =
|
||||
# ponytail: 16 is an unprofiled guess sized so UTD=1 keeps up with 30tps;
|
||||
# env knob is the calibration point if steps/sec falls short.
|
||||
parseInt(getEnv("SACLSTM_BATCH_SIZE", "16"))
|
||||
|
||||
proc getSaveInterval*(): int =
|
||||
parseInt(getEnv("SACLSTM_SAVE_INTERVAL", "500"))
|
||||
|
||||
proc getWeightsPath*(): string =
|
||||
getEnv("SACLSTM_WEIGHTS_PATH",
|
||||
currentSourcePath().parentDir / "weights" / "sac_latest.zip")
|
||||
|
||||
proc opponentKey*(enemyId: int): string =
|
||||
## Q14 follow-up (#49): name-based opponent identity from the v1.0.1
|
||||
## BotListUpdate table; numeric-id fallback for the window before the first
|
||||
## update arrives (getBotName still returns "" then). Called on the training
|
||||
## thread — the lookup is lock-guarded in the API, no cross-thread refs.
|
||||
let name = getBotName(enemyId)
|
||||
if name.len > 0: name else: $enemyId
|
||||
|
||||
proc bumpRoundCounter*() =
|
||||
## Liveness signal for tools/training_runner/RunTraining.java (#49): one
|
||||
## increment per round end; the runner aborts when it freezes (dead bot).
|
||||
# ponytail: non-atomic read-modify-write; single writer (main thread) and the
|
||||
# runner re-polls every 500ms with multi-round tolerance, torn reads self-heal.
|
||||
let p = getWeightsPath().parentDir / "round_counter.txt"
|
||||
var n = 0
|
||||
try:
|
||||
n = parseInt(readFile(p).strip())
|
||||
except CatchableError:
|
||||
discard # absent/garbage -> start at 1
|
||||
try:
|
||||
createDir(p.parentDir)
|
||||
writeFile(p, $(n + 1))
|
||||
except CatchableError:
|
||||
discard # counter is best-effort liveness; never kill an event handler
|
||||
|
||||
# ── Channel message (plain data only) ────────────────────────────────────────
|
||||
|
||||
type
|
||||
TrainingMsgKind* = enum tmkTransition, tmkNewBattle, tmkShutdown
|
||||
|
||||
TrainingMsg* = object
|
||||
case kind*: TrainingMsgKind
|
||||
of tmkTransition:
|
||||
state*: array[STATE_DIM, float32] # Q13 plain arrays, no tensors
|
||||
action*: array[ACTION_DIM, float32]
|
||||
reward*: float32
|
||||
nextState*: array[STATE_DIM, float32]
|
||||
done*: bool # true only at battle end
|
||||
of tmkNewBattle:
|
||||
enemyId*: int # Q14 numeric scannedBotId
|
||||
of tmkShutdown:
|
||||
discard
|
||||
|
||||
proc arrToTensor*[N: static int](arr: array[N, float32]): Tensor[float32] =
|
||||
result = newTensor[float32](N)
|
||||
for i in 0 ..< N: result[i] = arr[i]
|
||||
|
||||
# ── Flat weight snapshots ─────────────────────────────────────────────────────
|
||||
# Layout mirrors network.nim's fixed architecture (fc1 -> hiddenDim-wide LSTM
|
||||
# with [4h, 2h] combined weights, 128-wide fc2, 4-out heads). The asserts catch
|
||||
# layout drift if network.nim shapes ever change.
|
||||
|
||||
proc actorSize*(h: int): int = 8*h*h + 168*h + 1160
|
||||
proc criticSize*(h: int): int = 8*h*h + 172*h + 257
|
||||
|
||||
proc putT(t: Tensor[float32]; dst: var seq[float32]; c: var int) =
|
||||
for v in t:
|
||||
dst[c] = v
|
||||
inc c
|
||||
|
||||
proc takeT(src: seq[float32]; c: var int; rows, cols: int): Tensor[float32] =
|
||||
# seq slice copies, then toTensor copies again: result owns its memory —
|
||||
# never a view into src (src may be a cross-thread buffer).
|
||||
let n = rows * cols
|
||||
result = src[c ..< c + n].toTensor().reshape(rows, cols)
|
||||
c += n
|
||||
|
||||
proc takeV(src: seq[float32]; c: var int; n: int): Tensor[float32] =
|
||||
## Rank-1 vector (biases) — reshape(n) keeps rank 1.
|
||||
result = src[c ..< c + n].toTensor().reshape(n)
|
||||
c += n
|
||||
|
||||
proc packActor*(a: ActorNet; dst: var seq[float32]; c: var int) =
|
||||
putT(a.fc1.w, dst, c); putT(a.fc1.b, dst, c)
|
||||
putT(a.lstm.wCombined, dst, c); putT(a.lstm.bCombined, dst, c)
|
||||
putT(a.fc2.w, dst, c); putT(a.fc2.b, dst, c)
|
||||
putT(a.muHead.w, dst, c); putT(a.muHead.b, dst, c)
|
||||
putT(a.logStdHead.w, dst, c); putT(a.logStdHead.b, dst, c)
|
||||
|
||||
proc packCritic*(net: CriticNet; dst: var seq[float32]; c: var int) =
|
||||
putT(net.fc1.w, dst, c); putT(net.fc1.b, dst, c)
|
||||
putT(net.lstm.wCombined, dst, c); putT(net.lstm.bCombined, dst, c)
|
||||
putT(net.fc2.w, dst, c); putT(net.fc2.b, dst, c)
|
||||
putT(net.fc3.w, dst, c); putT(net.fc3.b, dst, c)
|
||||
|
||||
proc actorFromFlat*(src: seq[float32]; c: var int; h: int): ActorNet =
|
||||
result.fc1.w = takeT(src, c, h, 35)
|
||||
result.fc1.b = takeV(src, c, h)
|
||||
result.lstm.wCombined = takeT(src, c, 4*h, 2*h)
|
||||
result.lstm.bCombined = takeV(src, c, 4*h)
|
||||
result.fc2.w = takeT(src, c, 128, h)
|
||||
result.fc2.b = takeV(src, c, 128)
|
||||
result.muHead.w = takeT(src, c, 4, 128)
|
||||
result.muHead.b = takeV(src, c, 4)
|
||||
result.logStdHead.w = takeT(src, c, 4, 128)
|
||||
result.logStdHead.b = takeV(src, c, 4)
|
||||
result.lstm.hiddenDim = result.lstm.bCombined.size div 4 # same as weights.nim loadLSTMCell
|
||||
result.hiddenDim = h
|
||||
assert c == actorSize(h), "actor flat layout drift"
|
||||
|
||||
proc criticFromFlat*(src: seq[float32]; c: var int; h: int): CriticNet =
|
||||
result.fc1.w = takeT(src, c, h, 39)
|
||||
result.fc1.b = takeV(src, c, h)
|
||||
result.lstm.wCombined = takeT(src, c, 4*h, 2*h)
|
||||
result.lstm.bCombined = takeV(src, c, 4*h)
|
||||
result.fc2.w = takeT(src, c, 128, h)
|
||||
result.fc2.b = takeV(src, c, 128)
|
||||
result.fc3.w = takeT(src, c, 1, 128)
|
||||
result.fc3.b = takeV(src, c, 1)
|
||||
result.lstm.hiddenDim = result.lstm.bCombined.size div 4
|
||||
result.hiddenDim = h
|
||||
|
||||
type
|
||||
FullSnap* = object
|
||||
hiddenDim*: int
|
||||
data*: seq[float32] # actor | critic1 | critic2 | targetCritic1 | targetCritic2 | alpha
|
||||
|
||||
proc packFull*(t: SACTrainer): FullSnap =
|
||||
let h = t.actor.hiddenDim
|
||||
result.hiddenDim = h
|
||||
result.data = newSeq[float32](actorSize(h) + 4 * criticSize(h) + 1)
|
||||
var c = 0
|
||||
packActor(t.actor, result.data, c)
|
||||
packCritic(t.critic1, result.data, c)
|
||||
packCritic(t.critic2, result.data, c)
|
||||
packCritic(t.targetCritic1, result.data, c)
|
||||
packCritic(t.targetCritic2, result.data, c)
|
||||
assert abs(t.alpha().float64 - exp(t.logAlpha.float64)) < 1e-6
|
||||
result.data[c] = t.alpha()
|
||||
inc c
|
||||
assert c == result.data.len, "full snapshot layout drift"
|
||||
|
||||
proc unpackFull*(fs: FullSnap):
|
||||
tuple[a: ActorNet, c1, c2, t1, t2: CriticNet, alpha: float32] =
|
||||
var c = 0
|
||||
result.a = actorFromFlat(fs.data, c, fs.hiddenDim)
|
||||
result.c1 = criticFromFlat(fs.data, c, fs.hiddenDim)
|
||||
result.c2 = criticFromFlat(fs.data, c, fs.hiddenDim)
|
||||
result.t1 = criticFromFlat(fs.data, c, fs.hiddenDim)
|
||||
result.t2 = criticFromFlat(fs.data, c, fs.hiddenDim)
|
||||
result.alpha = fs.data[c]
|
||||
|
||||
# ── Shared weight snapshot (training thread writes, bot thread copies out) ────
|
||||
# Q7/Q11: Lock + always-latest semantics. `data` is allocated ONCE and written
|
||||
# IN PLACE under gWeightLock — it is never reassigned, so the shared heap block
|
||||
# never sees cross-thread refcount traffic. Readers copy element-wise out.
|
||||
|
||||
type
|
||||
WeightSnapshot* = object
|
||||
hiddenDim*: int
|
||||
version*: int
|
||||
data*: seq[float32]
|
||||
|
||||
var gWeightLock: Lock
|
||||
var gSharedSnap: WeightSnapshot
|
||||
var gTrainChan: Channel[TrainingMsg]
|
||||
var gSaveChan: Channel[FullSnap]
|
||||
var gTrainingThread: Thread[void]
|
||||
var gIoThread: Thread[void]
|
||||
var gInitialFull: FullSnap # built on main before spawn; read-once after (happens-before)
|
||||
var gWeightsPath: string
|
||||
|
||||
proc pullWeights*(myVersion: var int; hidden: var int;
|
||||
flat: var seq[float32]): bool =
|
||||
## Copy the latest actor snapshot out under the lock. Returns true when a new
|
||||
## version arrived (caller rebuilds its tensors on ITS OWN thread).
|
||||
withLock(gWeightLock):
|
||||
if gSharedSnap.version == myVersion:
|
||||
return false
|
||||
if flat.len != gSharedSnap.data.len:
|
||||
flat = newSeq[float32](gSharedSnap.data.len) # caller-thread-owned buffer
|
||||
for i in 0 ..< flat.len:
|
||||
flat[i] = gSharedSnap.data[i]
|
||||
myVersion = gSharedSnap.version
|
||||
hidden = gSharedSnap.hiddenDim
|
||||
true
|
||||
|
||||
proc evalModeActive*(): bool {.inline.} =
|
||||
## Lever 4 (#59): the harness's deterministic eval battles already run the bot
|
||||
## with SACLSTM_EVAL_MODE=1 (sac_train.sh eval_checkpoint, mechanism from #49).
|
||||
## While set, eval ticks must NOT feed the trainer — transitions would pollute
|
||||
## the replay buffer with eval-only data and trigger gradient updates.
|
||||
getEnv("SACLSTM_EVAL_MODE") == "1"
|
||||
|
||||
proc sendTrainingMsg*(msg: TrainingMsg): bool {.inline.} =
|
||||
## Bot-side enqueue (cap-256, drops on overflow per Q10). Thread-safe.
|
||||
## Lever 4 (#59): fully suppressed in eval mode — NewBattle drops too, so an
|
||||
## eval battle can neither add transitions nor clear/retarget the buffer.
|
||||
if evalModeActive(): return false
|
||||
gTrainChan.trySend(msg)
|
||||
|
||||
# ── Training state (testable without threads) ─────────────────────────────────
|
||||
|
||||
type
|
||||
TrainState* = object
|
||||
trainer*: SACTrainer
|
||||
buf*: ReplayBuffer
|
||||
lastEnemyKey*: string # opponent identity key (#49 name-based, Q14)
|
||||
stepCount*: int
|
||||
nextSave*: int
|
||||
|
||||
proc trainerFromFull*(initial: FullSnap): SACTrainer =
|
||||
let (a, c1, c2, t1, t2, alpha) = unpackFull(initial)
|
||||
result.actor = a
|
||||
result.critic1 = c1
|
||||
result.critic2 = c2
|
||||
result.targetCritic1 = t1
|
||||
result.targetCritic2 = t2
|
||||
assert alpha > 0.0'f32, "checkpoint alpha must be positive"
|
||||
result.logAlpha = ln(alpha)
|
||||
result.targetEntropy = getTargetEntropy()
|
||||
result.tau = getTau()
|
||||
result.lrActor = getLrActor()
|
||||
result.lrCritic = getLrCritic()
|
||||
result.lrAlpha = getLrAlpha()
|
||||
result.gamma = getGamma()
|
||||
result.adam = initSACAdamStates(result.actor, result.critic1, result.critic2)
|
||||
# ponytail: adam momentum not carried in the flat snapshot — optimizer restarts
|
||||
# fresh each process; switch to saveCheckpoint/loadCheckpoint end-to-end when
|
||||
# resume quality matters (#49 harness owns checkpoint management).
|
||||
|
||||
proc initTrainState*(initial: FullSnap): TrainState =
|
||||
result.trainer = trainerFromFull(initial)
|
||||
result.buf = newReplayBuffer(getBufferCapacity(), STATE_DIM, ACTION_DIM)
|
||||
result.lastEnemyKey = ""
|
||||
result.nextSave = getSaveInterval()
|
||||
|
||||
# ── Training-loss metrics (campaign v2 lever 3, #59) ──────────────────────────
|
||||
|
||||
proc metricsFilePath*(): string =
|
||||
## Sits next to the weights dir's parent: SAC_LSTM_Bot/training_metrics.jsonl
|
||||
## under the #49 harness (weights live in SAC_LSTM_Bot/weights/).
|
||||
getWeightsPath().parentDir.parentDir / "training_metrics.jsonl"
|
||||
|
||||
proc metricsLine*(epoch: float64; stepCount, bufferLen, drained, gradSteps: int;
|
||||
m: SACMetrics): string =
|
||||
## One JSONL line with exactly the scalars SACTrainer.sacUpdate exposes
|
||||
## (#59 lever 3 — SACMetrics was already returned, no trainer change needed):
|
||||
## losses/alpha averaged over this pass's gradient steps, buffer size from
|
||||
## replay_buffer.len, cumulative step count and drained transition count.
|
||||
"{\"epoch\":" & $epoch &
|
||||
",\"steps\":" & $stepCount &
|
||||
",\"buffer_size\":" & $bufferLen &
|
||||
",\"drained\":" & $drained &
|
||||
",\"grad_steps\":" & $gradSteps &
|
||||
",\"critic_loss\":" & $m.criticLoss &
|
||||
",\"actor_loss\":" & $m.actorLoss &
|
||||
",\"alpha_loss\":" & $m.alphaLoss &
|
||||
",\"alpha\":" & $m.alpha & "}"
|
||||
|
||||
proc appendMetricsLine(st: TrainState; drained, gradSteps: int; m: SACMetrics) =
|
||||
## Lever 3 (#59): one append per trainPass (never per gradient step). Open,
|
||||
## write, close — cheap and crash-tolerant; a metrics failure never kills
|
||||
## training.
|
||||
try:
|
||||
let f = open(metricsFilePath(), fmAppend)
|
||||
f.writeLine(metricsLine(epochTime(), st.stepCount, st.buf.len,
|
||||
drained, gradSteps, m))
|
||||
f.close()
|
||||
except CatchableError:
|
||||
discard
|
||||
|
||||
proc handleTrainingMsg*(st: var TrainState; msg: TrainingMsg): bool =
|
||||
## Process one message. Returns false for Shutdown (caller stops).
|
||||
## Tensors are born HERE from the message's plain arrays — training thread only.
|
||||
case msg.kind
|
||||
of tmkTransition:
|
||||
st.buf.add(Transition(
|
||||
state: msg.state.arrToTensor,
|
||||
action: msg.action.arrToTensor,
|
||||
reward: msg.reward,
|
||||
nextState: msg.nextState.arrToTensor,
|
||||
done: msg.done))
|
||||
of tmkNewBattle:
|
||||
# Q12a: keep the buffer if the opponent is unchanged, clear otherwise.
|
||||
# Identity keyed on NAME (#49); numeric-id fallback pre-BotListUpdate.
|
||||
let key = opponentKey(msg.enemyId)
|
||||
if key != st.lastEnemyKey:
|
||||
st.buf.clear()
|
||||
st.lastEnemyKey = key
|
||||
of tmkShutdown:
|
||||
return false
|
||||
true
|
||||
|
||||
proc trainPass*(st: var TrainState; drained: int) =
|
||||
## UTD gradient steps for the transitions drained this pass, then publish the
|
||||
## latest actor to the shared snapshot and request periodic disk saves.
|
||||
if drained <= 0 or not st.buf.canSample:
|
||||
return
|
||||
let steps = drained * getUtdRatio() # Q2
|
||||
var gradSteps = 0
|
||||
var sumCritic, sumActor, sumAlphaLoss, sumAlpha = 0.0'f32
|
||||
for i in 1 .. steps:
|
||||
let seqs = st.buf.sampleSequences(getBatchSize())
|
||||
if seqs.len == 0:
|
||||
break
|
||||
let m = sacUpdate(st.trainer, seqs)
|
||||
sumCritic += m.criticLoss; sumActor += m.actorLoss
|
||||
sumAlphaLoss += m.alphaLoss; sumAlpha += m.alpha
|
||||
inc gradSteps
|
||||
inc st.stepCount
|
||||
# Save check INSIDE the step loop (#56 launch finding): at production sizes
|
||||
# (hidden 256 ⇒ ~1 s/step) a drain burst queues minutes of steps; checking
|
||||
# only between passes meant the process died mid-loop before stepCount ever
|
||||
# reached nextSave — zero checkpoints persisted for the whole campaign.
|
||||
# Mid-loop checks + SAVE_INTERVAL≤20 (#54) keep saves ~20 s apart.
|
||||
if st.stepCount >= st.nextSave:
|
||||
st.nextSave += getSaveInterval()
|
||||
var full = packFull(st.trainer)
|
||||
discard gSaveChan.trySend(move(full)) # cap-1: drop if I/O thread is busy (Q5)
|
||||
if gradSteps > 0:
|
||||
# Lever 3 (#59): one metrics line per pass, losses averaged over its steps.
|
||||
appendMetricsLine(st, drained, gradSteps, SACMetrics(
|
||||
criticLoss: sumCritic / gradSteps.float32,
|
||||
actorLoss: sumActor / gradSteps.float32,
|
||||
alphaLoss: sumAlphaLoss / gradSteps.float32,
|
||||
alpha: sumAlpha / gradSteps.float32))
|
||||
# Publish latest actor (Q7): in-place write under the lock, bump version.
|
||||
withLock(gWeightLock):
|
||||
assert gSharedSnap.hiddenDim == st.trainer.actor.hiddenDim,
|
||||
"snapshot/trainer hidden size mismatch"
|
||||
var c = 0
|
||||
packActor(st.trainer.actor, gSharedSnap.data, c)
|
||||
inc gSharedSnap.version
|
||||
|
||||
# ── Threads ───────────────────────────────────────────────────────────────────
|
||||
|
||||
proc trainingThreadEntry() {.thread.} =
|
||||
{.cast(gcsafe).}:
|
||||
randomize()
|
||||
var st = initTrainState(gInitialFull)
|
||||
var running = true
|
||||
while running:
|
||||
let first = gTrainChan.recv() # block until traffic (no busy spin)
|
||||
var drained = 0
|
||||
var msg = first
|
||||
while running:
|
||||
if msg.kind == tmkTransition:
|
||||
inc drained
|
||||
if not handleTrainingMsg(st, msg):
|
||||
running = false # Shutdown
|
||||
break
|
||||
let (more, nxt) = gTrainChan.tryRecv()
|
||||
if not more:
|
||||
break # drained — now train (Q10)
|
||||
msg = nxt
|
||||
if not running:
|
||||
var full = packFull(st.trainer) # final save request, then exit
|
||||
discard gSaveChan.trySend(move(full))
|
||||
break
|
||||
trainPass(st, drained)
|
||||
|
||||
proc ioThreadEntry() {.thread.} =
|
||||
{.cast(gcsafe).}:
|
||||
while true:
|
||||
let fs = gSaveChan.recv() # blocks; exits via empty-data sentinel
|
||||
if fs.data.len == 0:
|
||||
break
|
||||
let (a, c1, c2, t1, t2, alpha) = unpackFull(fs)
|
||||
saveWeights(gWeightsPath, a, c1, c2, t1, t2, alpha) # atomic zip (weights.nim)
|
||||
|
||||
# ── Lifecycle ─────────────────────────────────────────────────────────────────
|
||||
|
||||
proc randomFull(): FullSnap =
|
||||
let t = initSACTrainer(STATE_DIM, ACTION_DIM) # random nets, born + freed here
|
||||
packFull(t)
|
||||
|
||||
proc loadOrInitFull(): FullSnap =
|
||||
let path = getWeightsPath()
|
||||
if fileExists(path):
|
||||
try:
|
||||
let cp = loadCheckpoint(path)
|
||||
var t = initSACTrainer(STATE_DIM, ACTION_DIM) # env hyperparams
|
||||
t.actor = cp.actor
|
||||
t.critic1 = cp.critic1
|
||||
t.critic2 = cp.critic2
|
||||
t.targetCritic1 = cp.targetCritic1
|
||||
t.targetCritic2 = cp.targetCritic2
|
||||
assert cp.alpha > 0.0'f32, "checkpoint alpha must be positive"
|
||||
t.logAlpha = ln(cp.alpha)
|
||||
result = packFull(t)
|
||||
except Exception as e:
|
||||
stderr.writeLine "[sac] checkpoint load failed (" & e.msg & ") — random init"
|
||||
result = randomFull()
|
||||
else:
|
||||
result = randomFull()
|
||||
|
||||
proc initIntegration*() =
|
||||
## Open channels, build the initial weight snapshot, spawn both threads.
|
||||
## Call once from the main module before start().
|
||||
gWeightsPath = getWeightsPath()
|
||||
gInitialFull = loadOrInitFull()
|
||||
gSharedSnap.hiddenDim = gInitialFull.hiddenDim
|
||||
gSharedSnap.data = newSeq[float32](actorSize(gInitialFull.hiddenDim))
|
||||
for i in 0 ..< gSharedSnap.data.len: # element-wise: no refcount traffic
|
||||
gSharedSnap.data[i] = gInitialFull.data[i]
|
||||
gSharedSnap.version = 1
|
||||
initLock(gWeightLock)
|
||||
gTrainChan.open(256) # Q10 cap-256
|
||||
gSaveChan.open(1) # Q5 cap-1
|
||||
# Lever 4 (#59): one-time visibility for the suppression gate (see
|
||||
# sendTrainingMsg) — the eval bot trains nothing by design.
|
||||
if evalModeActive():
|
||||
stderr.writeLine "[sac] SACLSTM_EVAL_MODE=1 — training input suppressed (lever 4, #59)"
|
||||
createThread(gTrainingThread, trainingThreadEntry)
|
||||
createThread(gIoThread, ioThreadEntry)
|
||||
|
||||
proc shutdownIntegration*() =
|
||||
## Stop both threads cleanly. Called after the bot disconnects (start returned).
|
||||
while gTrainChan.tryRecv().dataAvailable:
|
||||
discard # drop pending transitions — process exiting
|
||||
discard gTrainChan.trySend(TrainingMsg(kind: tmkShutdown))
|
||||
joinThread(gTrainingThread) # training requests one final save
|
||||
# Nim's recv() blocks even on closed channels, so the I/O thread exits via an
|
||||
# empty-data sentinel. Retry while it is still busy saving: every failed
|
||||
# trySend means a save is in flight and will be consumed, so this terminates.
|
||||
var stop = FullSnap(hiddenDim: -1)
|
||||
while not gSaveChan.trySend(move(stop)):
|
||||
sleep(50)
|
||||
stop = FullSnap(hiddenDim: -1)
|
||||
joinThread(gIoThread)
|
||||
gTrainChan.close()
|
||||
@@ -0,0 +1,161 @@
|
||||
## network.nim — LSTM-based Actor and dual Critic for SAC-v2.
|
||||
## No autograd; inference only. Manual LSTM cell from scratch.
|
||||
|
||||
import arraymancer
|
||||
import std/[math, random, os, strutils]
|
||||
|
||||
# ── Configuration ─────────────────────────────────────────────────────────────
|
||||
|
||||
proc getHiddenSize*(): int =
|
||||
let s = getEnv("SACLSTM_HIDDEN_SIZE", "256")
|
||||
result = parseInt(s)
|
||||
|
||||
proc isEvalMode*(): bool =
|
||||
getEnv("SACLSTM_EVAL_MODE", "0") == "1"
|
||||
|
||||
# ── Types ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
type
|
||||
Linear* = object
|
||||
w*, b*: Tensor[float32] # w: [out, in], b: [out]
|
||||
|
||||
LSTMCell* = object
|
||||
## Combined weight matrix Wi|Wf|Wg|Wo stacked: [4*hidden, input+hidden]
|
||||
## Combined bias stacked: [4*hidden]
|
||||
wCombined*: Tensor[float32]
|
||||
bCombined*: Tensor[float32]
|
||||
hiddenDim*: int
|
||||
|
||||
ActorNet* = object
|
||||
fc1*: Linear
|
||||
lstm*: LSTMCell
|
||||
fc2*: Linear
|
||||
muHead*: Linear
|
||||
logStdHead*: Linear
|
||||
hiddenDim*: int
|
||||
|
||||
CriticNet* = object
|
||||
fc1*: Linear
|
||||
lstm*: LSTMCell
|
||||
fc2*: Linear
|
||||
fc3*: Linear
|
||||
hiddenDim*: int
|
||||
|
||||
LSTMState* = tuple[h, c: Tensor[float32]] # each [hiddenDim]
|
||||
|
||||
# ── Init helpers ──────────────────────────────────────────────────────────────
|
||||
|
||||
proc initLinear*(inDim, outDim: int; scale: float32): Linear =
|
||||
result.w = randomNormalTensor[float32]([outDim, inDim]) *. scale
|
||||
result.b = zeros[float32](outDim)
|
||||
|
||||
proc initLinearHe*(inDim, outDim: int): Linear =
|
||||
initLinear(inDim, outDim, sqrt(2.0'f32 / inDim.float32))
|
||||
|
||||
proc initLinearOut*(inDim, outDim: int): Linear =
|
||||
initLinear(inDim, outDim, sqrt(1.0'f32 / inDim.float32))
|
||||
|
||||
proc initLSTMCell*(inputDim, hiddenDim: int): LSTMCell =
|
||||
result.hiddenDim = hiddenDim
|
||||
let fanIn = (inputDim + hiddenDim).float32
|
||||
let scale = sqrt(1.0'f32 / fanIn)
|
||||
result.wCombined = randomNormalTensor[float32]([4 * hiddenDim, inputDim + hiddenDim]) *. scale
|
||||
result.bCombined = zeros[float32](4 * hiddenDim)
|
||||
|
||||
proc zeroState*(hiddenDim: int): LSTMState =
|
||||
result = (h: zeros[float32](hiddenDim), c: zeros[float32](hiddenDim))
|
||||
|
||||
proc initActorNet*(stateDim: int): ActorNet =
|
||||
let h = getHiddenSize()
|
||||
result.hiddenDim = h
|
||||
result.fc1 = initLinearHe(stateDim, h)
|
||||
result.lstm = initLSTMCell(h, h)
|
||||
result.fc2 = initLinearHe(h, 128)
|
||||
result.muHead = initLinearOut(128, 4)
|
||||
result.logStdHead = initLinearOut(128, 4)
|
||||
|
||||
proc initCriticNet*(stateDim, actionDim: int): CriticNet =
|
||||
let h = getHiddenSize()
|
||||
result.hiddenDim = h
|
||||
result.fc1 = initLinearHe(stateDim + actionDim, h)
|
||||
result.lstm = initLSTMCell(h, h)
|
||||
result.fc2 = initLinearHe(h, 128)
|
||||
result.fc3 = initLinearOut(128, 1)
|
||||
|
||||
# ── Forward helpers ───────────────────────────────────────────────────────────
|
||||
|
||||
proc linear*(l: Linear; x: Tensor[float32]): Tensor[float32] =
|
||||
l.w * x + l.b
|
||||
|
||||
proc relu*(x: Tensor[float32]): Tensor[float32] =
|
||||
x.map(proc(v: float32): float32 = max(0.0'f32, v))
|
||||
|
||||
proc sigmoid*(x: Tensor[float32]): Tensor[float32] =
|
||||
x.map(proc(v: float32): float32 = 1.0'f32 / (1.0'f32 + exp(-v)))
|
||||
|
||||
proc tanhT*(x: Tensor[float32]): Tensor[float32] =
|
||||
x.map(proc(v: float32): float32 = tanh(v))
|
||||
|
||||
proc lstmStep*(cell: LSTMCell; x, h, c: Tensor[float32]): LSTMState =
|
||||
## x: [inputDim], h/c: [hiddenDim] → h', c': [hiddenDim]
|
||||
let xh = concat(x, h, axis = 0) # [inputDim + hiddenDim]
|
||||
let gates = cell.wCombined * xh + cell.bCombined # [4*hidden]
|
||||
let hd = cell.hiddenDim
|
||||
let iGate = sigmoid(gates[0 ..< hd])
|
||||
let fGate = sigmoid(gates[hd ..< 2*hd])
|
||||
let gGate = tanhT(gates[2*hd ..< 3*hd])
|
||||
let oGate = sigmoid(gates[3*hd ..< 4*hd])
|
||||
let cPrime = fGate *. c + iGate *. gGate
|
||||
let hPrime = oGate *. tanhT(cPrime)
|
||||
result = (h: hPrime, c: cPrime)
|
||||
|
||||
# ── Actor forward ─────────────────────────────────────────────────────────────
|
||||
|
||||
const
|
||||
LOG_STD_MIN = -5.0'f32
|
||||
LOG_STD_MAX = 2.0'f32
|
||||
LOG_PROB_EPS = 1e-6'f32
|
||||
|
||||
proc actorForward*(net: ActorNet; state: Tensor[float32]; lstm: LSTMState;
|
||||
deterministic = false):
|
||||
tuple[actions: Tensor[float32]; logProb: float32; lstm: LSTMState] =
|
||||
## state: [stateDim], lstm: (h,c) each [hiddenDim]
|
||||
## Returns actions [4], scalar logProb, updated (h',c').
|
||||
let h1 = relu(net.fc1.linear(state))
|
||||
let lstmOut = lstmStep(net.lstm, h1, lstm.h, lstm.c)
|
||||
let h2 = relu(net.fc2.linear(lstmOut.h))
|
||||
let mu = net.muHead.linear(h2)
|
||||
let logStdRaw = net.logStdHead.linear(h2)
|
||||
let logStd = logStdRaw.map(proc(v: float32): float32 = clamp(v, LOG_STD_MIN, LOG_STD_MAX))
|
||||
|
||||
if deterministic or isEvalMode():
|
||||
let actions = tanhT(mu)
|
||||
return (actions: actions, logProb: 0.0'f32, lstm: lstmOut)
|
||||
|
||||
# Reparameterization: z = mu + std * eps, action = tanh(z)
|
||||
let std = logStd.map(proc(v: float32): float32 = exp(v))
|
||||
var actions = newTensor[float32](4)
|
||||
var logProb = 0.0'f32
|
||||
let twoPiLog = 0.5'f32 * ln(2.0'f32 * PI.float32)
|
||||
for i in 0 ..< 4:
|
||||
let eps = gauss(0.0'f64, 1.0'f64).float32
|
||||
let z = mu[i] + std[i] * eps
|
||||
actions[i] = tanh(z)
|
||||
# log N(z | mu, std) - log(1 - tanh²(z) + eps)
|
||||
let diff = (z - mu[i]) / std[i]
|
||||
let logNorm = -0.5'f32 * diff * diff - ln(std[i]) - twoPiLog
|
||||
let tanhCorr = ln(1.0'f32 - actions[i] * actions[i] + LOG_PROB_EPS)
|
||||
logProb += logNorm - tanhCorr
|
||||
|
||||
result = (actions: actions, logProb: logProb, lstm: lstmOut)
|
||||
|
||||
# ── Critic forward ────────────────────────────────────────────────────────────
|
||||
|
||||
proc criticForward*(net: CriticNet; stateAction: Tensor[float32]; lstm: LSTMState):
|
||||
tuple[q: float32; lstm: LSTMState] =
|
||||
## stateAction: [stateDim + actionDim], lstm: (h,c) each [hiddenDim]
|
||||
let h1 = relu(net.fc1.linear(stateAction))
|
||||
let lstmOut = lstmStep(net.lstm, h1, lstm.h, lstm.c)
|
||||
let h2 = relu(net.fc2.linear(lstmOut.h))
|
||||
let q = net.fc3.linear(h2)
|
||||
result = (q: q[0], lstm: lstmOut)
|
||||
@@ -0,0 +1,136 @@
|
||||
## replay_buffer.nim — sequential ring buffer for off-policy SAC+LSTM training.
|
||||
##
|
||||
## Stores transitions and samples contiguous sequences for recurrent training.
|
||||
## Sequences NEVER cross battle boundaries (done=true).
|
||||
##
|
||||
## Config env vars:
|
||||
## SACLSTM_BUFFER_CAPACITY (default: 500_000)
|
||||
## SACLSTM_BURN_IN (default: 8)
|
||||
## SACLSTM_TRAIN_WINDOW (default: 16)
|
||||
|
||||
import arraymancer
|
||||
import std/[os, strutils, random]
|
||||
|
||||
# ── Config ────────────────────────────────────────────────────────────────────
|
||||
|
||||
proc getBufferCapacity*(): int =
|
||||
parseInt(getEnv("SACLSTM_BUFFER_CAPACITY", "500000"))
|
||||
|
||||
proc getBurnIn*(): int =
|
||||
parseInt(getEnv("SACLSTM_BURN_IN", "8"))
|
||||
|
||||
proc getTrainWindow*(): int =
|
||||
parseInt(getEnv("SACLSTM_TRAIN_WINDOW", "16"))
|
||||
|
||||
# ── Types ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
type
|
||||
Transition* = object
|
||||
state*: Tensor[float32] # [stateDim]
|
||||
action*: Tensor[float32] # [actionDim]
|
||||
reward*: float32
|
||||
nextState*: Tensor[float32] # [stateDim]
|
||||
done*: bool # true = battle end
|
||||
|
||||
Sequence* = object
|
||||
burnIn*: seq[Transition] # first burnIn steps (for LSTM warm-up)
|
||||
train*: seq[Transition] # next trainWindow steps (for gradient computation)
|
||||
|
||||
ReplayBuffer* = object
|
||||
## Ring buffer. `head` is the next write position. `count` tracks fill level.
|
||||
transitions: seq[Transition]
|
||||
capacity: int
|
||||
stateDim: int
|
||||
actionDim: int
|
||||
head: int # next write index
|
||||
count: int # number of valid transitions stored
|
||||
burnIn: int
|
||||
trainWindow: int
|
||||
|
||||
# ── Construction ──────────────────────────────────────────────────────────────
|
||||
|
||||
proc newReplayBuffer*(capacity, stateDim, actionDim: int;
|
||||
burnIn = getBurnIn();
|
||||
trainWindow = getTrainWindow()): ReplayBuffer =
|
||||
result.capacity = capacity
|
||||
result.stateDim = stateDim
|
||||
result.actionDim = actionDim
|
||||
result.burnIn = burnIn
|
||||
result.trainWindow = trainWindow
|
||||
result.head = 0
|
||||
result.count = 0
|
||||
result.transitions = newSeq[Transition](capacity)
|
||||
|
||||
# ── Core operations ───────────────────────────────────────────────────────────
|
||||
|
||||
proc add*(buf: var ReplayBuffer; t: Transition) =
|
||||
buf.transitions[buf.head] = t
|
||||
buf.head = (buf.head + 1) mod buf.capacity
|
||||
if buf.count < buf.capacity:
|
||||
inc buf.count
|
||||
|
||||
proc len*(buf: ReplayBuffer): int = buf.count
|
||||
|
||||
proc clear*(buf: var ReplayBuffer) =
|
||||
## Drop all transitions (#48: opponent changed across battles).
|
||||
## Old slots keep stale tensors until the ring overwrites them.
|
||||
buf.head = 0
|
||||
buf.count = 0
|
||||
|
||||
proc canSample*(buf: ReplayBuffer): bool =
|
||||
buf.count >= buf.burnIn + buf.trainWindow
|
||||
|
||||
# ── Sampling ──────────────────────────────────────────────────────────────────
|
||||
|
||||
proc sampleSequences*(buf: ReplayBuffer; batchSize: int): seq[Sequence] =
|
||||
## Sample `batchSize` contiguous sequences of length burnIn+trainWindow.
|
||||
## Sequences never cross a done=true boundary and never wrap the ring buffer.
|
||||
##
|
||||
## Returns fewer than batchSize sequences if not enough valid starts exist.
|
||||
## Returns empty seq if canSample is false.
|
||||
if not buf.canSample: return @[]
|
||||
|
||||
let seqLen = buf.burnIn + buf.trainWindow
|
||||
let oldest = if buf.count < buf.capacity: 0
|
||||
else: buf.head # oldest valid index when full
|
||||
|
||||
# Build valid starting indices.
|
||||
# ponytail: O(count) scan per sample call; upgrade to an indexed set of
|
||||
# boundary positions if count reaches hundreds of thousands and profiling shows
|
||||
# this is a bottleneck.
|
||||
var validStarts: seq[int]
|
||||
for i in 0 ..< buf.count - seqLen + 1:
|
||||
# Absolute ring-buffer index for the i-th oldest transition
|
||||
let startIdx = (oldest + i) mod buf.capacity
|
||||
# Check: the sequence [startIdx .. startIdx+seqLen-2] must not contain done=true
|
||||
# (a done at position k means the battle ended there; the next transition is
|
||||
# from a new battle, so the sequence would cross a boundary).
|
||||
# Also, the sequence must not wrap around the ring buffer.
|
||||
let endIdx = startIdx + seqLen - 1 # exclusive of wrap check
|
||||
if endIdx >= buf.capacity:
|
||||
# Sequence wraps the ring buffer — invalid starting point.
|
||||
continue
|
||||
var crosses = false
|
||||
for j in 0 ..< seqLen - 1:
|
||||
if buf.transitions[startIdx + j].done:
|
||||
crosses = true
|
||||
break
|
||||
if not crosses:
|
||||
validStarts.add(startIdx)
|
||||
|
||||
if validStarts.len == 0: return @[]
|
||||
|
||||
result = newSeq[Sequence](min(batchSize, validStarts.len))
|
||||
# Sample with replacement if batchSize > validStarts.len, else sample without.
|
||||
# ponytail: sampling with replacement for simplicity; shuffle+take for
|
||||
# without-replacement if the caller needs it.
|
||||
for i in 0 ..< result.len:
|
||||
let startIdx = validStarts[rand(validStarts.len - 1)]
|
||||
var s: Sequence
|
||||
s.burnIn = newSeq[Transition](buf.burnIn)
|
||||
s.train = newSeq[Transition](buf.trainWindow)
|
||||
for j in 0 ..< buf.burnIn:
|
||||
s.burnIn[j] = buf.transitions[startIdx + j]
|
||||
for j in 0 ..< buf.trainWindow:
|
||||
s.train[j] = buf.transitions[startIdx + buf.burnIn + j]
|
||||
result[i] = s
|
||||
@@ -0,0 +1,84 @@
|
||||
## rewards.nim — Raw reward computation + running mean/variance normalizer.
|
||||
## Welford online algorithm; safe cold-start (0 or 1 samples).
|
||||
|
||||
import std/math
|
||||
|
||||
# ── Lever-2 shaping constants (#59, campaign v2) — TUNABLE ────────────────────
|
||||
# Scale discipline: commensurate with existing magnitudes (dealt p=1 was +4,
|
||||
# wall tick -5/tick, win +20). Death/loss and win terms stay dominant; these
|
||||
# only re-rank mid-band behaviors (fight vs outlive vs get-rammed).
|
||||
|
||||
const
|
||||
# Multiplier on the bullet-damage-dealt term: p=1 hit +4 -> +5. Low-power
|
||||
# spam stays unprofitable (6*0.1-2 = -1.4 < 0 even after x1.25).
|
||||
AggressionMult* = 1.25 # ponytail: TUNABLE — raise toward 1.5 if v2 bot still passivity-leaning
|
||||
# Flat per landed shot on top of damage: discrete accuracy signal.
|
||||
HitBonus* = 0.5 # ponytail: TUNABLE — keep < 6p-2 at min viable power (~0.34)
|
||||
# Per bot-bot collision (BotHitBotEvent): server deals RAM_DAMAGE=0.6 to
|
||||
# both parties but only notifies the hitter — each receipt = damage taken.
|
||||
RamTakenPenalty* = 3.0 # ponytail: TUNABLE — vs p=0.8 bullet received (-6.8)
|
||||
# Enemy-charging deterrent: distance/diagonal below this => escalating
|
||||
# negative (max at zero distance), suppressed while we deal damage that step.
|
||||
ChargeDistFrac* = 0.12 # ponytail: TUNABLE — ~120u of 800x600 diag (1000)
|
||||
ChargePenalty* = 2.0 # ponytail: TUNABLE — per-tick ceiling, milder than wall (-5/tick)
|
||||
|
||||
# ── Raw reward ────────────────────────────────────────────────────────────────
|
||||
|
||||
proc computeReward*(
|
||||
damageInflicted: float64 = 0.0, # fire power p of own shot that hit
|
||||
damageReceived: float64 = 0.0, # fire power p_e of enemy shot that hit
|
||||
wallHitTicks: int = 0, # ticks in wall contact this step
|
||||
wastedShotPower: float64 = 0.0, # fire power of shot that missed/hit wall
|
||||
hitCount: int = 0, # own bullets that hit the enemy this step (#59)
|
||||
ramTakenCount: int = 0, # collisions where we were the victim (#59)
|
||||
enemyDistFrac: float64 = 2.0, # enemy dist / arena diag; >ChargeDistFrac when no contact (#59)
|
||||
win: bool = false,
|
||||
loss: bool = false
|
||||
): float64 =
|
||||
## Returns the raw (un-normalized) reward for one decision step.
|
||||
## Damage formula: 4p + 2(p-1) = 6p - 2 (matches Tank Royale bullet rules).
|
||||
let p = damageInflicted
|
||||
let pe = damageReceived
|
||||
if p > 0.0:
|
||||
result += AggressionMult * (6.0 * p - 2.0)
|
||||
if hitCount > 0: result += HitBonus * hitCount.float64
|
||||
if pe > 0.0: result -= 6.0 * pe - 2.0
|
||||
result -= RamTakenPenalty * ramTakenCount.float64
|
||||
if enemyDistFrac < ChargeDistFrac and p <= 0.0:
|
||||
result -= ChargePenalty * (1.0 - enemyDistFrac / ChargeDistFrac)
|
||||
result -= 5.0 * wallHitTicks.float64
|
||||
if wastedShotPower > 0.0: result -= 0.1 * wastedShotPower
|
||||
if win: result += 20.0
|
||||
if loss: result -= 10.0
|
||||
|
||||
# ── Running normalizer (Welford) ──────────────────────────────────────────────
|
||||
|
||||
const NormEps = 1e-8
|
||||
|
||||
type
|
||||
RewardNormalizer* = object
|
||||
n*: int # samples seen
|
||||
mean*: float64
|
||||
m2*: float64 # sum of squared deviations (Welford M2)
|
||||
|
||||
proc update*(rn: var RewardNormalizer; r: float64) =
|
||||
rn.n += 1
|
||||
let delta = r - rn.mean
|
||||
rn.mean += delta / rn.n.float64
|
||||
let delta2 = r - rn.mean
|
||||
rn.m2 += delta * delta2
|
||||
|
||||
proc normalize*(rn: RewardNormalizer; r: float64): float64 =
|
||||
## Returns (r - mean) / (std + eps) once statistics are meaningful
|
||||
## (n >= 4 and spread well above zero). Before that, returns the RAW
|
||||
## reward unchanged — Welford M2 collapses to exactly 0 when early raw
|
||||
## rewards are identical, and dividing by the 1e-8 floor then z-scores
|
||||
## the first differing reward to ~1e8, poisoning TD targets.
|
||||
if rn.n < 4: return r
|
||||
let variance = rn.m2 / rn.n.float64 # ponytail: population var; switch to n-1 if bias matters
|
||||
let stddev = sqrt(variance)
|
||||
if stddev <= 1e-3 * (abs(rn.mean) + 1.0): return r
|
||||
# ponytail: warm-up pass-through ceiling — raw rewards bypass normalization
|
||||
# until stats are meaningful; upgrade = persist Welford state in checkpoint
|
||||
# if warm-up noise ever hurts learning.
|
||||
result = (r - rn.mean) / (stddev + NormEps)
|
||||
@@ -0,0 +1,117 @@
|
||||
## State vector module — produces a 35-dimensional normalized tensor for SAC+LSTM policy.
|
||||
## No bot API imports; takes plain data structs populated from game events.
|
||||
## The LSTM handles temporal context, so no explicit history window here.
|
||||
|
||||
import std/math
|
||||
import arraymancer
|
||||
|
||||
const STATE_DIM* = 35
|
||||
|
||||
type
|
||||
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
|
||||
|
||||
EnemyData* = object
|
||||
## Current enemy state, from the most recent onScannedBot event.
|
||||
x*, y*: float64
|
||||
direction*: float64
|
||||
speed*: float64
|
||||
energy*: float64
|
||||
hasFired*: bool
|
||||
lastFirePower*: float64
|
||||
prevSpeed*: float64 # speed from the previous scan (for acceleration)
|
||||
prevDirection*: float64 # direction from the previous scan (for turn rate)
|
||||
hasPrevScan*: bool # true once we have at least two scans
|
||||
|
||||
GameState* = object
|
||||
## Accumulates data from bot events. Populate fields before calling buildState.
|
||||
# Own bot
|
||||
x*, y*: float64
|
||||
direction*: float64
|
||||
speed*: float64
|
||||
energy*: float64
|
||||
gunDirection*: float64
|
||||
gunHeat*: float64
|
||||
arenaWidth*, arenaHeight*: float64
|
||||
# Enemy
|
||||
hasContact*: bool
|
||||
enemy*: EnemyData
|
||||
ticksSinceLastScan*: int
|
||||
# Bullets in flight (up to 3 tracked)
|
||||
bullets*: array[3, BulletData]
|
||||
bulletCount*: int
|
||||
|
||||
proc buildState*(gs: GameState): Tensor[float32] =
|
||||
## Build the 35-float normalized state tensor.
|
||||
##
|
||||
## Layout:
|
||||
## [0-6] own bot: x/aW, y/aH, dir/360, speed/8, energy/100, gunDir/360, gunHeat/1.8
|
||||
## [7-13] enemy: x/aW, y/aH, dir/360, speed/8, energy/100, hasFired, lastFirePower/3
|
||||
## [14-17] derived: enemyAccel/8, enemyTurnRate/180, relBearing/180, distance/diag
|
||||
## [18-21] walls: top, bottom, left, right — each / max(aW,aH)
|
||||
## [22-33] bullets: up to 3 × (relX/aW, relY/aH, speed/20, ticksToImpact clamped to 1)
|
||||
## [34] scan staleness: ticksSinceLastScan/30 clamped to 1
|
||||
result = zeros[float32](STATE_DIM)
|
||||
|
||||
let aW = gs.arenaWidth
|
||||
let aH = gs.arenaHeight
|
||||
let diag = sqrt(aW * aW + aH * aH)
|
||||
let wMax = max(aW, aH)
|
||||
|
||||
# --- Own bot (0-6) ---
|
||||
result[0] = float32(gs.x / aW)
|
||||
result[1] = float32(gs.y / aH)
|
||||
result[2] = float32(gs.direction / 360.0)
|
||||
result[3] = float32(gs.speed / 8.0)
|
||||
result[4] = float32(gs.energy / 100.0)
|
||||
result[5] = float32(gs.gunDirection / 360.0)
|
||||
result[6] = float32(gs.gunHeat / 1.8)
|
||||
|
||||
# --- Enemy current (7-13) ---
|
||||
if gs.hasContact:
|
||||
result[7] = float32(gs.enemy.x / aW)
|
||||
result[8] = float32(gs.enemy.y / aH)
|
||||
result[9] = float32(gs.enemy.direction / 360.0)
|
||||
result[10] = float32(gs.enemy.speed / 8.0)
|
||||
result[11] = float32(gs.enemy.energy / 100.0)
|
||||
result[12] = float32(if gs.enemy.hasFired: 1.0 else: 0.0)
|
||||
result[13] = float32(gs.enemy.lastFirePower / 3.0)
|
||||
|
||||
# --- Derived (14-17) ---
|
||||
if gs.hasContact:
|
||||
if gs.enemy.hasPrevScan:
|
||||
result[14] = float32((gs.enemy.speed - gs.enemy.prevSpeed) / 8.0)
|
||||
let dDir = ((gs.enemy.direction - gs.enemy.prevDirection) + 540.0) mod 360.0 - 180.0
|
||||
result[15] = float32(dDir / 180.0)
|
||||
let dx = gs.enemy.x - gs.x
|
||||
let dy = gs.enemy.y - gs.y
|
||||
let absDir = (180.0 * arctan2(dx, dy) / PI + 360.0) mod 360.0
|
||||
let relBearing = ((absDir - gs.direction) + 540.0) mod 360.0 - 180.0
|
||||
result[16] = float32(relBearing / 180.0)
|
||||
result[17] = float32(sqrt(dx * dx + dy * dy) / diag)
|
||||
|
||||
# --- Wall distances (18-21): top, bottom, left, right ---
|
||||
result[18] = float32((aH - gs.y) / wMax)
|
||||
result[19] = float32(gs.y / wMax)
|
||||
result[20] = float32(gs.x / wMax)
|
||||
result[21] = float32((aW - gs.x) / wMax)
|
||||
|
||||
# --- Bullet tracking (22-33): up to 3 bullets × 4 floats ---
|
||||
# Per slot: relX/aW, relY/aH, speed/20, ticksToImpact/diag (clamped to 1)
|
||||
for i in 0 ..< min(gs.bulletCount, 3):
|
||||
let b = gs.bullets[i]
|
||||
let bSpd = 20.0 - 3.0 * b.power
|
||||
let bdx = b.x - gs.x
|
||||
let bdy = b.y - gs.y
|
||||
let bdist = sqrt(bdx * bdx + bdy * bdy)
|
||||
let ticks = if bSpd > 0.0: min(bdist / bSpd / diag, 1.0) else: 0.0
|
||||
let base = 22 + i * 4
|
||||
result[base + 0] = float32(bdx / aW)
|
||||
result[base + 1] = float32(bdy / aH)
|
||||
result[base + 2] = float32(bSpd / 20.0)
|
||||
result[base + 3] = float32(ticks)
|
||||
|
||||
# --- Scan staleness (34) ---
|
||||
result[34] = float32(min(gs.ticksSinceLastScan.float64 / 30.0, 1.0))
|
||||
@@ -0,0 +1,611 @@
|
||||
## training.nim — SAC-v2 update for the LSTM Actor + twin Critic.
|
||||
## Manual backprop; no autograd. Uses Arraymancer tensors throughout.
|
||||
##
|
||||
## Config env vars:
|
||||
## SACLSTM_LR_ACTOR (default: 3e-4)
|
||||
## SACLSTM_LR_CRITIC (default: 3e-4)
|
||||
## SACLSTM_LR_ALPHA (default: 3e-4)
|
||||
## SACLSTM_GAMMA (default: 0.99)
|
||||
## SACLSTM_TAU (default: 0.005)
|
||||
## SACLSTM_TARGET_ENTROPY (default: -4.0)
|
||||
|
||||
import arraymancer except Linear
|
||||
import std/[math, os, strutils]
|
||||
import SAC_LSTM_Bot/network
|
||||
import SAC_LSTM_Bot/weights
|
||||
import SAC_LSTM_Bot/replay_buffer
|
||||
|
||||
# ── Config ────────────────────────────────────────────────────────────────────
|
||||
|
||||
proc getLrActor*(): float32 = parseFloat(getEnv("SACLSTM_LR_ACTOR", "3e-4")).float32
|
||||
proc getLrCritic*(): float32 = parseFloat(getEnv("SACLSTM_LR_CRITIC", "3e-4")).float32
|
||||
proc getLrAlpha*(): float32 = parseFloat(getEnv("SACLSTM_LR_ALPHA", "3e-4")).float32
|
||||
proc getGamma*(): float32 = parseFloat(getEnv("SACLSTM_GAMMA", "0.99")).float32
|
||||
proc getTau*(): float32 = parseFloat(getEnv("SACLSTM_TAU", "0.005")).float32
|
||||
proc getTargetEntropy*(): float32 =
|
||||
parseFloat(getEnv("SACLSTM_TARGET_ENTROPY", "-4.0")).float32
|
||||
|
||||
# ── SACTrainer ────────────────────────────────────────────────────────────────
|
||||
|
||||
type
|
||||
SACTrainer* = object
|
||||
actor*: ActorNet
|
||||
critic1*: CriticNet
|
||||
critic2*: CriticNet
|
||||
targetCritic1*: CriticNet
|
||||
targetCritic2*: CriticNet
|
||||
logAlpha*: float32 ## log of entropy temperature; alpha = exp(logAlpha)
|
||||
targetEntropy*: float32
|
||||
tau*: float32
|
||||
lrActor*: float32
|
||||
lrCritic*: float32
|
||||
lrAlpha*: float32
|
||||
gamma*: float32
|
||||
adam*: SACAdamStates
|
||||
|
||||
SACMetrics* = object
|
||||
criticLoss*: float32
|
||||
actorLoss*: float32
|
||||
alphaLoss*: float32
|
||||
alpha*: float32
|
||||
|
||||
proc initSACTrainer*(stateDim, actionDim: int): SACTrainer =
|
||||
result.actor = initActorNet(stateDim)
|
||||
result.critic1 = initCriticNet(stateDim, actionDim)
|
||||
result.critic2 = initCriticNet(stateDim, actionDim)
|
||||
result.targetCritic1 = result.critic1
|
||||
result.targetCritic2 = result.critic2
|
||||
result.logAlpha = 0.0'f32
|
||||
result.targetEntropy = getTargetEntropy()
|
||||
result.tau = getTau()
|
||||
result.lrActor = getLrActor()
|
||||
result.lrCritic = getLrCritic()
|
||||
result.lrAlpha = getLrAlpha()
|
||||
result.gamma = getGamma()
|
||||
result.adam = initSACAdamStates(result.actor, result.critic1, result.critic2)
|
||||
|
||||
proc alpha*(t: SACTrainer): float32 = exp(t.logAlpha)
|
||||
|
||||
# ── Adam steps ────────────────────────────────────────────────────────────────
|
||||
|
||||
proc adamStepScalar(param: var float32; grad: float32;
|
||||
state: var AdamVar; lr: float32) =
|
||||
## Scalar Adam for logAlpha (state.m/v are shape [1] tensors).
|
||||
inc state.t
|
||||
let b1 = 0.9'f32; let b2 = 0.999'f32; let eps = 1e-8'f32
|
||||
state.m[0] = b1 * state.m[0] + (1.0'f32 - b1) * grad
|
||||
state.v[0] = b2 * state.v[0] + (1.0'f32 - b2) * grad * grad
|
||||
let mHat = state.m[0] / (1.0'f32 - b1 ^ state.t.float32)
|
||||
let vHat = state.v[0] / (1.0'f32 - b2 ^ state.t.float32)
|
||||
param -= lr * mHat / (sqrt(vHat) + eps)
|
||||
|
||||
proc adamStepTensor(param: var Tensor[float32]; grad: Tensor[float32];
|
||||
state: var AdamVar; lr: float32) =
|
||||
## Tensor Adam (same pattern as PPO_Bot/training.nim adamStep).
|
||||
inc state.t
|
||||
let b1 = 0.9'f32; let b2 = 0.999'f32; let eps = 1e-8'f32
|
||||
state.m = b1 *. state.m + (1.0'f32 - b1) *. grad
|
||||
state.v = b2 *. state.v + (1.0'f32 - b2) *. (grad *. grad)
|
||||
let mHat = state.m /. (1.0'f32 - b1 ^ state.t.float32)
|
||||
let vHat = state.v /. (1.0'f32 - b2 ^ state.t.float32)
|
||||
param -= lr *. mHat /. vHat.map(proc(x: float32): float32 = sqrt(x) + eps)
|
||||
|
||||
# ── Gradient clipping ─────────────────────────────────────────────────────────
|
||||
|
||||
proc globalNorm(grads: seq[Tensor[float32]]): float32 =
|
||||
var sumSq = 0.0'f32
|
||||
for g in grads:
|
||||
for v in g: sumSq += v * v
|
||||
sqrt(sumSq)
|
||||
|
||||
proc clipGrads(grads: var seq[Tensor[float32]]; maxNorm: float32) =
|
||||
let norm = globalNorm(grads)
|
||||
if norm > maxNorm and norm == norm:
|
||||
let scale = maxNorm / norm
|
||||
for g in grads.mitems: g = g *. scale
|
||||
|
||||
# ── Forward caches (for backprop) ─────────────────────────────────────────────
|
||||
|
||||
type
|
||||
LinearFwd = object
|
||||
inp, pre, act: Tensor[float32] # input, pre-relu, post-relu (or linear)
|
||||
|
||||
proc linearReluFwd(l: Linear; x: Tensor[float32]): LinearFwd =
|
||||
result.inp = x
|
||||
result.pre = l.w * x + l.b
|
||||
result.act = relu(result.pre)
|
||||
|
||||
proc linearFwd(l: Linear; x: Tensor[float32]): LinearFwd =
|
||||
result.inp = x
|
||||
result.pre = l.w * x + l.b
|
||||
result.act = result.pre # no nonlinearity
|
||||
|
||||
type
|
||||
LSTMFwdCache = object
|
||||
xh, gatesPre: Tensor[float32] # [inputDim+hd], [4*hd]
|
||||
iGate, fGate, gGate, oGate: Tensor[float32] # [hd] each
|
||||
cPrev, cPrime, hPrime: Tensor[float32] # [hd] each
|
||||
|
||||
proc lstmStepCached(cell: LSTMCell; x, h, c: Tensor[float32]): LSTMFwdCache =
|
||||
result.cPrev = c
|
||||
result.xh = concat(x, h, axis = 0)
|
||||
result.gatesPre = cell.wCombined * result.xh + cell.bCombined
|
||||
let hd = cell.hiddenDim
|
||||
result.iGate = sigmoid(result.gatesPre[0 ..< hd])
|
||||
result.fGate = sigmoid(result.gatesPre[hd ..< 2*hd])
|
||||
result.gGate = tanhT(result.gatesPre[2*hd ..< 3*hd])
|
||||
result.oGate = sigmoid(result.gatesPre[3*hd ..< 4*hd])
|
||||
result.cPrime = result.fGate *. c + result.iGate *. result.gGate
|
||||
result.hPrime = result.oGate *. tanhT(result.cPrime)
|
||||
|
||||
# ── Backward helpers ──────────────────────────────────────────────────────────
|
||||
|
||||
proc reluGrad(pre, dAct: Tensor[float32]): Tensor[float32] =
|
||||
result = newTensor[float32](dAct.shape)
|
||||
for i in 0 ..< dAct.shape[0]:
|
||||
result[i] = if pre[i] > 0.0'f32: dAct[i] else: 0.0'f32
|
||||
|
||||
## Linear layer backward: returns (dx, dw, db) given upstream grad dAct.
|
||||
## If hasRelu, applies relu' gate before computing gradients.
|
||||
proc linearBack(w: Tensor[float32]; fwd: LinearFwd;
|
||||
dAct: Tensor[float32]; hasRelu: bool):
|
||||
tuple[dx, dw, db: Tensor[float32]] =
|
||||
let dPre = if hasRelu: reluGrad(fwd.pre, dAct) else: dAct
|
||||
result.dw = dPre.unsqueeze(1) * fwd.inp.unsqueeze(0) # [out, in]
|
||||
result.db = dPre
|
||||
result.dx = w.transpose * dPre # [in]
|
||||
|
||||
## LSTM single-step backward. dHPrime: [hd], dCPrime: [hd] (use zeros for truncated BPTT).
|
||||
## Returns (dwCombined, dbCombined, dxh).
|
||||
proc lstmBack(cell: LSTMCell; cache: LSTMFwdCache;
|
||||
dHPrime, dCPrime: Tensor[float32]):
|
||||
tuple[dwCombined, dbCombined, dxh: Tensor[float32]] =
|
||||
let hd = cell.hiddenDim
|
||||
let tanhCPrime = tanhT(cache.cPrime)
|
||||
|
||||
# Output gate
|
||||
let dOGate_post = dHPrime *. tanhCPrime
|
||||
# Cell state: gradient from h' and from downstream dCPrime
|
||||
let dCPrimeTotal = dHPrime *. cache.oGate *.
|
||||
(ones[float32](hd) - tanhCPrime *. tanhCPrime) + dCPrime
|
||||
|
||||
# Gate post-activation gradients
|
||||
let dFGate_post = dCPrimeTotal *. cache.cPrev
|
||||
let dIGate_post = dCPrimeTotal *. cache.gGate
|
||||
let dGGate_post = dCPrimeTotal *. cache.iGate
|
||||
|
||||
# Gate pre-activation gradients (sigmoid', tanh')
|
||||
let dIPre = dIGate_post *. cache.iGate *. (ones[float32](hd) - cache.iGate)
|
||||
let dFPre = dFGate_post *. cache.fGate *. (ones[float32](hd) - cache.fGate)
|
||||
let dGPre = dGGate_post *. (ones[float32](hd) - cache.gGate *. cache.gGate)
|
||||
let dOPre = dOGate_post *. cache.oGate *. (ones[float32](hd) - cache.oGate)
|
||||
|
||||
# Concatenated gate gradient [4*hd]
|
||||
let dGatesPre = concat(dIPre, dFPre, dGPre, dOPre, axis = 0)
|
||||
|
||||
result.dwCombined = dGatesPre.unsqueeze(1) * cache.xh.unsqueeze(0)
|
||||
result.dbCombined = dGatesPre
|
||||
result.dxh = cell.wCombined.transpose * dGatesPre
|
||||
|
||||
# ── Squashed-Gaussian log-prob and its gradients ──────────────────────────────
|
||||
|
||||
const
|
||||
LOG_PROB_EPS = 1e-6'f32
|
||||
LOG_STD_MIN = -5.0'f32
|
||||
LOG_STD_MAX = 2.0'f32
|
||||
|
||||
## Given stored mu, clamped logStd, and sampled action = tanh(z), recover
|
||||
## log π(a|s) and gradients w.r.t. mu and logStd.
|
||||
proc squashedLogProb(mu, logStd, action: Tensor[float32]):
|
||||
tuple[logProb: float32;
|
||||
dLogProbDMu, dLogProbDLogStd: Tensor[float32]] =
|
||||
let std = logStd.map(proc(v: float32): float32 = exp(v))
|
||||
let z = mu # deterministic reparam: action = tanh(mu), so z ≡ mu, diff ≡ 0
|
||||
result.dLogProbDMu = newTensor[float32](4)
|
||||
result.dLogProbDLogStd = newTensor[float32](4)
|
||||
let twoPiLog = 0.5'f32 * ln(2.0'f32 * PI.float32)
|
||||
var lp = 0.0'f32
|
||||
for i in 0 ..< 4:
|
||||
let diff = (z[i] - mu[i]) / std[i]
|
||||
let logNorm = -0.5'f32 * diff * diff - ln(std[i]) - twoPiLog
|
||||
let tanhCorr = ln(1.0'f32 - action[i] * action[i] + LOG_PROB_EPS)
|
||||
lp += logNorm - tanhCorr
|
||||
result.dLogProbDMu[i] = diff / std[i] # (z-mu)/std²
|
||||
result.dLogProbDLogStd[i] = diff * diff - 1.0'f32 # d logN / d logStd
|
||||
result.logProb = lp
|
||||
|
||||
# ── Critic forward with activation cache ──────────────────────────────────────
|
||||
|
||||
type
|
||||
CriticFwdCache = object
|
||||
fc1: LinearFwd
|
||||
lstm: LSTMFwdCache
|
||||
fc2: LinearFwd
|
||||
fc3: LinearFwd
|
||||
q: float32
|
||||
|
||||
proc criticFwdCached(net: CriticNet; stateAction: Tensor[float32];
|
||||
h, c: Tensor[float32]): CriticFwdCache =
|
||||
result.fc1 = linearReluFwd(net.fc1, stateAction)
|
||||
result.lstm = lstmStepCached(net.lstm, result.fc1.act, h, c)
|
||||
result.fc2 = linearReluFwd(net.fc2, result.lstm.hPrime)
|
||||
result.fc3 = linearFwd(net.fc3, result.fc2.act)
|
||||
result.q = result.fc3.act[0]
|
||||
|
||||
# ── Critic backward ───────────────────────────────────────────────────────────
|
||||
|
||||
type
|
||||
CriticGrads = object
|
||||
dFc1W, dFc1B: Tensor[float32]
|
||||
dFc2W, dFc2B: Tensor[float32]
|
||||
dFc3W, dFc3B: Tensor[float32]
|
||||
dLstmW, dLstmB: Tensor[float32]
|
||||
dInput: Tensor[float32] ## grad w.r.t. stateAction input
|
||||
|
||||
proc criticBack(net: CriticNet; cache: CriticFwdCache; dQ: float32): CriticGrads =
|
||||
let dFc3Act = [dQ].toTensor()
|
||||
let fc3b = linearBack(net.fc3.w, cache.fc3, dFc3Act, hasRelu = false)
|
||||
result.dFc3W = fc3b.dw; result.dFc3B = fc3b.db
|
||||
|
||||
let fc2b = linearBack(net.fc2.w, cache.fc2, fc3b.dx, hasRelu = true)
|
||||
result.dFc2W = fc2b.dw; result.dFc2B = fc2b.db
|
||||
|
||||
let zeros_hd = zeros[float32](net.hiddenDim)
|
||||
let lstmb = lstmBack(net.lstm, cache.lstm, fc2b.dx, zeros_hd)
|
||||
result.dLstmW = lstmb.dwCombined; result.dLstmB = lstmb.dbCombined
|
||||
|
||||
# xh = [fc1.act | h_prev], dx is the x-part (fc1 output dim = hiddenDim)
|
||||
let dLstmX = lstmb.dxh[0 ..< cache.fc1.act.shape[0]]
|
||||
let fc1b = linearBack(net.fc1.w, cache.fc1, dLstmX, hasRelu = true)
|
||||
result.dFc1W = fc1b.dw; result.dFc1B = fc1b.db
|
||||
result.dInput = fc1b.dx # [stateDim + actionDim]
|
||||
|
||||
# ── Actor forward with activation cache ───────────────────────────────────────
|
||||
|
||||
type
|
||||
ActorFwdCache = object
|
||||
fc1: LinearFwd
|
||||
lstm: LSTMFwdCache
|
||||
fc2: LinearFwd
|
||||
muHead: LinearFwd
|
||||
lsHead: LinearFwd ## logStd head
|
||||
mu: Tensor[float32] ## [4]
|
||||
logStd: Tensor[float32] ## [4] clamped
|
||||
action: Tensor[float32] ## [4] tanh(mu) — deterministic for gradient
|
||||
|
||||
proc actorFwdCached(net: ActorNet; state: Tensor[float32];
|
||||
h, c: Tensor[float32]): ActorFwdCache =
|
||||
result.fc1 = linearReluFwd(net.fc1, state)
|
||||
result.lstm = lstmStepCached(net.lstm, result.fc1.act, h, c)
|
||||
result.fc2 = linearReluFwd(net.fc2, result.lstm.hPrime)
|
||||
result.muHead = linearFwd(net.muHead, result.fc2.act)
|
||||
result.lsHead = linearFwd(net.logStdHead, result.fc2.act)
|
||||
result.mu = result.muHead.act
|
||||
result.logStd = result.lsHead.act.map(
|
||||
proc(v: float32): float32 = clamp(v, LOG_STD_MIN, LOG_STD_MAX))
|
||||
# Use tanh(mu) as the action for gradient computation (reparameterization).
|
||||
# ponytail: deterministic here; add stochastic sample if off-policy bias matters.
|
||||
result.action = tanhT(result.mu)
|
||||
|
||||
# ── Actor backward ────────────────────────────────────────────────────────────
|
||||
|
||||
type
|
||||
ActorGrads = object
|
||||
dFc1W, dFc1B: Tensor[float32]
|
||||
dFc2W, dFc2B: Tensor[float32]
|
||||
dMuW, dMuB: Tensor[float32]
|
||||
dLogStdW, dLogStdB: Tensor[float32]
|
||||
dLstmW, dLstmB: Tensor[float32]
|
||||
|
||||
proc actorBack(net: ActorNet; cache: ActorFwdCache;
|
||||
dMu, dLogStd: Tensor[float32]): ActorGrads =
|
||||
let muBack = linearBack(net.muHead.w, cache.muHead, dMu, hasRelu = false)
|
||||
result.dMuW = muBack.dw; result.dMuB = muBack.db
|
||||
|
||||
let lsBack = linearBack(net.logStdHead.w, cache.lsHead, dLogStd, hasRelu = false)
|
||||
result.dLogStdW = lsBack.dw; result.dLogStdB = lsBack.db
|
||||
|
||||
# fc2 gets grads from both output heads
|
||||
let fc2b = linearBack(net.fc2.w, cache.fc2, muBack.dx + lsBack.dx, hasRelu = true)
|
||||
result.dFc2W = fc2b.dw; result.dFc2B = fc2b.db
|
||||
|
||||
let zeros_hd = zeros[float32](net.hiddenDim)
|
||||
let lstmb = lstmBack(net.lstm, cache.lstm, fc2b.dx, zeros_hd)
|
||||
result.dLstmW = lstmb.dwCombined; result.dLstmB = lstmb.dbCombined
|
||||
|
||||
let dLstmX = lstmb.dxh[0 ..< cache.fc1.act.shape[0]]
|
||||
let fc1b = linearBack(net.fc1.w, cache.fc1, dLstmX, hasRelu = true)
|
||||
result.dFc1W = fc1b.dw; result.dFc1B = fc1b.db
|
||||
|
||||
# ── Adam application ──────────────────────────────────────────────────────────
|
||||
|
||||
proc applyActorAdam(net: var ActorNet; g: ActorGrads;
|
||||
adam: var ActorAdam; lr: float32) =
|
||||
adamStepTensor(net.fc1.w, g.dFc1W, adam.fc1.w, lr)
|
||||
adamStepTensor(net.fc1.b, g.dFc1B, adam.fc1.b, lr)
|
||||
adamStepTensor(net.fc2.w, g.dFc2W, adam.fc2.w, lr)
|
||||
adamStepTensor(net.fc2.b, g.dFc2B, adam.fc2.b, lr)
|
||||
adamStepTensor(net.muHead.w, g.dMuW, adam.muHead.w, lr)
|
||||
adamStepTensor(net.muHead.b, g.dMuB, adam.muHead.b, lr)
|
||||
adamStepTensor(net.logStdHead.w, g.dLogStdW, adam.logStdHead.w, lr)
|
||||
adamStepTensor(net.logStdHead.b, g.dLogStdB, adam.logStdHead.b, lr)
|
||||
adamStepTensor(net.lstm.wCombined, g.dLstmW, adam.lstm.wCombined, lr)
|
||||
adamStepTensor(net.lstm.bCombined, g.dLstmB, adam.lstm.bCombined, lr)
|
||||
|
||||
proc applyCriticAdam(net: var CriticNet; g: CriticGrads;
|
||||
adam: var CriticAdam; lr: float32) =
|
||||
adamStepTensor(net.fc1.w, g.dFc1W, adam.fc1.w, lr)
|
||||
adamStepTensor(net.fc1.b, g.dFc1B, adam.fc1.b, lr)
|
||||
adamStepTensor(net.fc2.w, g.dFc2W, adam.fc2.w, lr)
|
||||
adamStepTensor(net.fc2.b, g.dFc2B, adam.fc2.b, lr)
|
||||
adamStepTensor(net.fc3.w, g.dFc3W, adam.fc3.w, lr)
|
||||
adamStepTensor(net.fc3.b, g.dFc3B, adam.fc3.b, lr)
|
||||
adamStepTensor(net.lstm.wCombined, g.dLstmW, adam.lstm.wCombined, lr)
|
||||
adamStepTensor(net.lstm.bCombined, g.dLstmB, adam.lstm.bCombined, lr)
|
||||
|
||||
# ── Soft target update ────────────────────────────────────────────────────────
|
||||
|
||||
proc softUpdateLinear(target: var Linear; src: Linear; tau: float32) =
|
||||
target.w = tau *. src.w + (1.0'f32 - tau) *. target.w
|
||||
target.b = tau *. src.b + (1.0'f32 - tau) *. target.b
|
||||
|
||||
proc softUpdateLSTM(target: var LSTMCell; src: LSTMCell; tau: float32) =
|
||||
target.wCombined = tau *. src.wCombined + (1.0'f32 - tau) *. target.wCombined
|
||||
target.bCombined = tau *. src.bCombined + (1.0'f32 - tau) *. target.bCombined
|
||||
|
||||
proc softUpdateCritic(target: var CriticNet; src: CriticNet; tau: float32) =
|
||||
softUpdateLinear(target.fc1, src.fc1, tau)
|
||||
softUpdateLSTM(target.lstm, src.lstm, tau)
|
||||
softUpdateLinear(target.fc2, src.fc2, tau)
|
||||
softUpdateLinear(target.fc3, src.fc3, tau)
|
||||
|
||||
# ── Gradient accumulators ─────────────────────────────────────────────────────
|
||||
|
||||
proc zeroCriticGrads(net: CriticNet): CriticGrads =
|
||||
result.dFc1W = zeros[float32](net.fc1.w.shape)
|
||||
result.dFc1B = zeros[float32](net.fc1.b.shape)
|
||||
result.dFc2W = zeros[float32](net.fc2.w.shape)
|
||||
result.dFc2B = zeros[float32](net.fc2.b.shape)
|
||||
result.dFc3W = zeros[float32](net.fc3.w.shape)
|
||||
result.dFc3B = zeros[float32](net.fc3.b.shape)
|
||||
result.dLstmW = zeros[float32](net.lstm.wCombined.shape)
|
||||
result.dLstmB = zeros[float32](net.lstm.bCombined.shape)
|
||||
result.dInput = zeros[float32](net.fc1.w.shape[1]) # [stateDim+actionDim]
|
||||
|
||||
proc zeroActorGrads(net: ActorNet): ActorGrads =
|
||||
result.dFc1W = zeros[float32](net.fc1.w.shape)
|
||||
result.dFc1B = zeros[float32](net.fc1.b.shape)
|
||||
result.dFc2W = zeros[float32](net.fc2.w.shape)
|
||||
result.dFc2B = zeros[float32](net.fc2.b.shape)
|
||||
result.dMuW = zeros[float32](net.muHead.w.shape)
|
||||
result.dMuB = zeros[float32](net.muHead.b.shape)
|
||||
result.dLogStdW = zeros[float32](net.logStdHead.w.shape)
|
||||
result.dLogStdB = zeros[float32](net.logStdHead.b.shape)
|
||||
result.dLstmW = zeros[float32](net.lstm.wCombined.shape)
|
||||
result.dLstmB = zeros[float32](net.lstm.bCombined.shape)
|
||||
|
||||
proc addCriticGrads(a: var CriticGrads; b: CriticGrads) =
|
||||
a.dFc1W += b.dFc1W; a.dFc1B += b.dFc1B
|
||||
a.dFc2W += b.dFc2W; a.dFc2B += b.dFc2B
|
||||
a.dFc3W += b.dFc3W; a.dFc3B += b.dFc3B
|
||||
a.dLstmW += b.dLstmW; a.dLstmB += b.dLstmB
|
||||
# dInput not accumulated (not used for parameter update)
|
||||
|
||||
proc addActorGrads(a: var ActorGrads; b: ActorGrads) =
|
||||
a.dFc1W += b.dFc1W; a.dFc1B += b.dFc1B
|
||||
a.dFc2W += b.dFc2W; a.dFc2B += b.dFc2B
|
||||
a.dMuW += b.dMuW; a.dMuB += b.dMuB
|
||||
a.dLogStdW += b.dLogStdW; a.dLogStdB += b.dLogStdB
|
||||
a.dLstmW += b.dLstmW; a.dLstmB += b.dLstmB
|
||||
|
||||
proc scaleCriticGrads(g: var CriticGrads; s: float32) =
|
||||
g.dFc1W = g.dFc1W *. s; g.dFc1B = g.dFc1B *. s
|
||||
g.dFc2W = g.dFc2W *. s; g.dFc2B = g.dFc2B *. s
|
||||
g.dFc3W = g.dFc3W *. s; g.dFc3B = g.dFc3B *. s
|
||||
g.dLstmW = g.dLstmW *. s; g.dLstmB = g.dLstmB *. s
|
||||
|
||||
proc scaleActorGrads(g: var ActorGrads; s: float32) =
|
||||
g.dFc1W = g.dFc1W *. s; g.dFc1B = g.dFc1B *. s
|
||||
g.dFc2W = g.dFc2W *. s; g.dFc2B = g.dFc2B *. s
|
||||
g.dMuW = g.dMuW *. s; g.dMuB = g.dMuB *. s
|
||||
g.dLogStdW = g.dLogStdW *. s; g.dLogStdB = g.dLogStdB *. s
|
||||
g.dLstmW = g.dLstmW *. s; g.dLstmB = g.dLstmB *. s
|
||||
|
||||
proc criticGradsAsSeq(g: CriticGrads): seq[Tensor[float32]] =
|
||||
@[g.dFc1W, g.dFc1B, g.dFc2W, g.dFc2B, g.dFc3W, g.dFc3B, g.dLstmW, g.dLstmB]
|
||||
|
||||
proc applyClipToCritic(g: var CriticGrads; maxNorm: float32) =
|
||||
var gs = criticGradsAsSeq(g)
|
||||
clipGrads(gs, maxNorm)
|
||||
g.dFc1W = gs[0]; g.dFc1B = gs[1]
|
||||
g.dFc2W = gs[2]; g.dFc2B = gs[3]
|
||||
g.dFc3W = gs[4]; g.dFc3B = gs[5]
|
||||
g.dLstmW = gs[6]; g.dLstmB = gs[7]
|
||||
|
||||
proc applyClipToActor(g: var ActorGrads; maxNorm: float32) =
|
||||
var gs = @[g.dFc1W, g.dFc1B, g.dFc2W, g.dFc2B,
|
||||
g.dMuW, g.dMuB, g.dLogStdW, g.dLogStdB, g.dLstmW, g.dLstmB]
|
||||
clipGrads(gs, maxNorm)
|
||||
g.dFc1W = gs[0]; g.dFc1B = gs[1]
|
||||
g.dFc2W = gs[2]; g.dFc2B = gs[3]
|
||||
g.dMuW = gs[4]; g.dMuB = gs[5]
|
||||
g.dLogStdW = gs[6]; g.dLogStdB = gs[7]
|
||||
g.dLstmW = gs[8]; g.dLstmB = gs[9]
|
||||
|
||||
# ── SAC update ────────────────────────────────────────────────────────────────
|
||||
|
||||
proc sacUpdate*(trainer: var SACTrainer; sequences: seq[Sequence]): SACMetrics =
|
||||
## One SAC-v2 update given a batch of sequences. No-op if empty.
|
||||
if sequences.len == 0: return
|
||||
|
||||
let N = sequences.len.float32
|
||||
let alph = trainer.alpha()
|
||||
let gamma = trainer.gamma
|
||||
|
||||
var totalCriticLoss = 0.0'f32
|
||||
var totalActorLoss = 0.0'f32
|
||||
var totalAlphaLoss = 0.0'f32
|
||||
|
||||
var accC1Grads = zeroCriticGrads(trainer.critic1)
|
||||
var accC2Grads = zeroCriticGrads(trainer.critic2)
|
||||
var accAGrads = zeroActorGrads(trainer.actor)
|
||||
var dLogAlpha = 0.0'f32
|
||||
|
||||
for sq in sequences:
|
||||
# ── 1. Burn-in: warm up hidden states, no gradient ──────────────────────
|
||||
var actorH = zeros[float32](trainer.actor.hiddenDim)
|
||||
var actorC = zeros[float32](trainer.actor.hiddenDim)
|
||||
var c1H = zeros[float32](trainer.critic1.hiddenDim)
|
||||
var c1C = zeros[float32](trainer.critic1.hiddenDim)
|
||||
var c2H = zeros[float32](trainer.critic2.hiddenDim)
|
||||
var c2C = zeros[float32](trainer.critic2.hiddenDim)
|
||||
var tc1H = zeros[float32](trainer.targetCritic1.hiddenDim)
|
||||
var tc1C = zeros[float32](trainer.targetCritic1.hiddenDim)
|
||||
var tc2H = zeros[float32](trainer.targetCritic2.hiddenDim)
|
||||
var tc2C = zeros[float32](trainer.targetCritic2.hiddenDim)
|
||||
|
||||
for tr in sq.burnIn:
|
||||
let sa = concat(tr.state, tr.action, axis = 0)
|
||||
let af = lstmStepCached(trainer.actor.lstm,
|
||||
relu(trainer.actor.fc1.linear(tr.state)), actorH, actorC)
|
||||
actorH = af.hPrime; actorC = af.cPrime
|
||||
let c1f = lstmStepCached(trainer.critic1.lstm,
|
||||
relu(trainer.critic1.fc1.linear(sa)), c1H, c1C)
|
||||
c1H = c1f.hPrime; c1C = c1f.cPrime
|
||||
let c2f = lstmStepCached(trainer.critic2.lstm,
|
||||
relu(trainer.critic2.fc1.linear(sa)), c2H, c2C)
|
||||
c2H = c2f.hPrime; c2C = c2f.cPrime
|
||||
let tc1f = lstmStepCached(trainer.targetCritic1.lstm,
|
||||
relu(trainer.targetCritic1.fc1.linear(sa)), tc1H, tc1C)
|
||||
tc1H = tc1f.hPrime; tc1C = tc1f.cPrime
|
||||
let tc2f = lstmStepCached(trainer.targetCritic2.lstm,
|
||||
relu(trainer.targetCritic2.fc1.linear(sa)), tc2H, tc2C)
|
||||
tc2H = tc2f.hPrime; tc2C = tc2f.cPrime
|
||||
|
||||
# ── 2–4. Training window ─────────────────────────────────────────────────
|
||||
let T = sq.train.len.float32
|
||||
|
||||
var seqC1Grads = zeroCriticGrads(trainer.critic1)
|
||||
var seqC2Grads = zeroCriticGrads(trainer.critic2)
|
||||
var seqAGrads = zeroActorGrads(trainer.actor)
|
||||
var seqDLogAlpha = 0.0'f32
|
||||
|
||||
for tr in sq.train:
|
||||
let s = tr.state
|
||||
let a = tr.action
|
||||
let r = tr.reward
|
||||
let sn = tr.nextState
|
||||
let d = if tr.done: 0.0'f32 else: 1.0'f32
|
||||
let sa = concat(s, a, axis = 0)
|
||||
|
||||
# ── 2. Critic update ─────────────────────────────────────────────────
|
||||
|
||||
let c1Cache = criticFwdCached(trainer.critic1, sa, c1H, c1C)
|
||||
let c2Cache = criticFwdCached(trainer.critic2, sa, c2H, c2C)
|
||||
|
||||
# ── 3. Actor update (run first to get actorFwd on s before advancing h/c) ─
|
||||
|
||||
let actorFwd = actorFwdCached(trainer.actor, s, actorH, actorC)
|
||||
let aCurr = actorFwd.action
|
||||
let lpResult = squashedLogProb(actorFwd.mu, actorFwd.logStd, aCurr)
|
||||
let logProbA = lpResult.logProb
|
||||
|
||||
# Advance actor hidden state from s → sn before computing actorNxt
|
||||
actorH = actorFwd.lstm.hPrime; actorC = actorFwd.lstm.cPrime
|
||||
|
||||
# Next-state action from current actor (uses h/c advanced through s)
|
||||
let actorNxt = actorFwdCached(trainer.actor, sn, actorH, actorC)
|
||||
let aN = actorNxt.action
|
||||
let lpN = squashedLogProb(actorNxt.mu, actorNxt.logStd, aN).logProb
|
||||
let saN = concat(sn, aN, axis = 0)
|
||||
|
||||
# Target Q
|
||||
let tc1Cache = criticFwdCached(trainer.targetCritic1, saN, tc1H, tc1C)
|
||||
let tc2Cache = criticFwdCached(trainer.targetCritic2, saN, tc2H, tc2C)
|
||||
let minQTarg = min(tc1Cache.q, tc2Cache.q)
|
||||
|
||||
# Bellman target
|
||||
let y = r + gamma * d * (minQTarg - alph * lpN)
|
||||
let errQ1 = c1Cache.q - y
|
||||
let errQ2 = c2Cache.q - y
|
||||
totalCriticLoss += 0.5'f32 * (errQ1 * errQ1 + errQ2 * errQ2)
|
||||
|
||||
# MSE gradient: d_loss/d_q = (q - y) [scaling applied at accumulation]
|
||||
addCriticGrads(seqC1Grads, criticBack(trainer.critic1, c1Cache, errQ1))
|
||||
addCriticGrads(seqC2Grads, criticBack(trainer.critic2, c2Cache, errQ2))
|
||||
|
||||
# Advance critic hidden states
|
||||
c1H = c1Cache.lstm.hPrime; c1C = c1Cache.lstm.cPrime
|
||||
c2H = c2Cache.lstm.hPrime; c2C = c2Cache.lstm.cPrime
|
||||
tc1H = tc1Cache.lstm.hPrime; tc1C = tc1Cache.lstm.cPrime
|
||||
tc2H = tc2Cache.lstm.hPrime; tc2C = tc2Cache.lstm.cPrime
|
||||
|
||||
# ── 3 (cont). Actor gradient via critic ─────────────────────────────────
|
||||
|
||||
# Q-values for current policy action (critics used as frozen estimators)
|
||||
let saCurr = concat(s, aCurr, axis = 0)
|
||||
let qA1Cache = criticFwdCached(trainer.critic1, saCurr, c1H, c1C)
|
||||
let qA2Cache = criticFwdCached(trainer.critic2, saCurr, c2H, c2C)
|
||||
let q1Val = qA1Cache.q
|
||||
let q2Val = qA2Cache.q
|
||||
let qA1Back = criticBack(trainer.critic1, qA1Cache, -1.0'f32)
|
||||
let qA2Back = criticBack(trainer.critic2, qA2Cache, -1.0'f32)
|
||||
totalActorLoss += alph * logProbA - min(q1Val, q2Val)
|
||||
|
||||
# Gradient of -minQ w.r.t. action = dInput[stateDim ..< stateDim+actionDim]
|
||||
# from the critic whose Q was smaller.
|
||||
let minQBack = if q1Val <= q2Val: qA1Back else: qA2Back
|
||||
let stateDim = s.shape[0]
|
||||
let actionDim = aCurr.shape[0]
|
||||
let dQdA = minQBack.dInput[stateDim ..< stateDim + actionDim]
|
||||
|
||||
# Chain through tanh: d(tanh(mu))/d(mu) = 1 - action²
|
||||
let dTanh = aCurr.map(proc(a: float32): float32 = 1.0'f32 - a * a)
|
||||
|
||||
# Total gradient w.r.t. mu: (alpha * dLogP/dMu + dQ/dA) * dTanh/dMu
|
||||
let dMu = (alph *. lpResult.dLogProbDMu + dQdA) *. dTanh
|
||||
let dLogStd = alph *. lpResult.dLogProbDLogStd
|
||||
|
||||
addActorGrads(seqAGrads, actorBack(trainer.actor, actorFwd, dMu, dLogStd))
|
||||
|
||||
# ── 4. Alpha update ──────────────────────────────────────────────────
|
||||
# Loss = -log_alpha * stop_grad(logProb + targetEntropy)
|
||||
# d_loss/d_log_alpha = -(logProb + targetEntropy)
|
||||
totalAlphaLoss += -trainer.logAlpha * (logProbA + trainer.targetEntropy)
|
||||
seqDLogAlpha += -(logProbA + trainer.targetEntropy)
|
||||
|
||||
# Average sequence grads over T steps, accumulate over batch
|
||||
scaleCriticGrads(seqC1Grads, 1.0'f32 / T)
|
||||
scaleCriticGrads(seqC2Grads, 1.0'f32 / T)
|
||||
scaleActorGrads(seqAGrads, 1.0'f32 / T)
|
||||
addCriticGrads(accC1Grads, seqC1Grads)
|
||||
addCriticGrads(accC2Grads, seqC2Grads)
|
||||
addActorGrads(accAGrads, seqAGrads)
|
||||
dLogAlpha += seqDLogAlpha / T
|
||||
|
||||
# Average over batch
|
||||
scaleCriticGrads(accC1Grads, 1.0'f32 / N)
|
||||
scaleCriticGrads(accC2Grads, 1.0'f32 / N)
|
||||
scaleActorGrads(accAGrads, 1.0'f32 / N)
|
||||
dLogAlpha /= N
|
||||
|
||||
# Gradient clipping (max_norm = 1.0)
|
||||
applyClipToCritic(accC1Grads, 1.0'f32)
|
||||
applyClipToCritic(accC2Grads, 1.0'f32)
|
||||
applyClipToActor(accAGrads, 1.0'f32)
|
||||
|
||||
# Apply Adam updates
|
||||
applyCriticAdam(trainer.critic1, accC1Grads, trainer.adam.critic1, trainer.lrCritic)
|
||||
applyCriticAdam(trainer.critic2, accC2Grads, trainer.adam.critic2, trainer.lrCritic)
|
||||
applyActorAdam(trainer.actor, accAGrads, trainer.adam.actor, trainer.lrActor)
|
||||
adamStepScalar(trainer.logAlpha, dLogAlpha, trainer.adam.alpha, trainer.lrAlpha)
|
||||
|
||||
# ── 5. Soft target update ──────────────────────────────────────────────────
|
||||
softUpdateCritic(trainer.targetCritic1, trainer.critic1, trainer.tau)
|
||||
softUpdateCritic(trainer.targetCritic2, trainer.critic2, trainer.tau)
|
||||
|
||||
let totalSteps = N * sequences[0].train.len.float32
|
||||
result.criticLoss = totalCriticLoss / totalSteps
|
||||
result.actorLoss = totalActorLoss / totalSteps
|
||||
result.alphaLoss = totalAlphaLoss / totalSteps
|
||||
result.alpha = trainer.alpha()
|
||||
@@ -0,0 +1,352 @@
|
||||
## weights.nim — save/load all SAC-LSTM network tensors as .npy inside a .zip.
|
||||
##
|
||||
## Strategy: write_npy writes to paths; zip/zipfiles.addFile reads from paths.
|
||||
## So we write each tensor to a temp .npy, add it to the zip, then delete temps.
|
||||
## Load reverses: extract each entry to a temp .npy, read_npy, delete.
|
||||
## Atomic save: build the zip in a temp path, then rename over the target.
|
||||
|
||||
import arraymancer except Linear
|
||||
import zip/zipfiles
|
||||
import std/[os, times, strutils]
|
||||
import SAC_LSTM_Bot/network
|
||||
|
||||
# ── Adam state types (used by training.nim) ───────────────────────────────────
|
||||
|
||||
type
|
||||
AdamVar* = object
|
||||
m*, v*: Tensor[float32]
|
||||
t*: int
|
||||
|
||||
## Adam states for one Linear layer (w and b).
|
||||
LinearAdam* = object
|
||||
w*, b*: AdamVar
|
||||
|
||||
## Adam states for one LSTMCell (wCombined and bCombined).
|
||||
LSTMCellAdam* = object
|
||||
wCombined*, bCombined*: AdamVar
|
||||
|
||||
## Adam states for one ActorNet.
|
||||
ActorAdam* = object
|
||||
fc1*, fc2*, muHead*, logStdHead*: LinearAdam
|
||||
lstm*: LSTMCellAdam
|
||||
|
||||
## Adam states for one CriticNet.
|
||||
CriticAdam* = object
|
||||
fc1*, fc2*, fc3*: LinearAdam
|
||||
lstm*: LSTMCellAdam
|
||||
|
||||
SACAdamStates* = object
|
||||
actor*: ActorAdam
|
||||
critic1*: CriticAdam
|
||||
critic2*: CriticAdam
|
||||
alpha*: AdamVar # scalar, shape [1]
|
||||
initialized*: bool
|
||||
|
||||
# ── Init helpers ──────────────────────────────────────────────────────────────
|
||||
|
||||
proc initAdamVar(t: Tensor[float32]): AdamVar =
|
||||
AdamVar(m: zeros[float32](t.shape), v: zeros[float32](t.shape), t: 0)
|
||||
|
||||
proc initLinearAdam*(l: Linear): LinearAdam =
|
||||
LinearAdam(w: initAdamVar(l.w), b: initAdamVar(l.b))
|
||||
|
||||
proc initLSTMCellAdam*(c: LSTMCell): LSTMCellAdam =
|
||||
LSTMCellAdam(
|
||||
wCombined: initAdamVar(c.wCombined),
|
||||
bCombined: initAdamVar(c.bCombined))
|
||||
|
||||
proc initActorAdam*(a: ActorNet): ActorAdam =
|
||||
ActorAdam(
|
||||
fc1: initLinearAdam(a.fc1),
|
||||
fc2: initLinearAdam(a.fc2),
|
||||
muHead: initLinearAdam(a.muHead),
|
||||
logStdHead: initLinearAdam(a.logStdHead),
|
||||
lstm: initLSTMCellAdam(a.lstm))
|
||||
|
||||
proc initCriticAdam*(c: CriticNet): CriticAdam =
|
||||
CriticAdam(
|
||||
fc1: initLinearAdam(c.fc1),
|
||||
fc2: initLinearAdam(c.fc2),
|
||||
fc3: initLinearAdam(c.fc3),
|
||||
lstm: initLSTMCellAdam(c.lstm))
|
||||
|
||||
proc initSACAdamStates*(actor: ActorNet; critic1, critic2: CriticNet): SACAdamStates =
|
||||
result.actor = initActorAdam(actor)
|
||||
result.critic1 = initCriticAdam(critic1)
|
||||
result.critic2 = initCriticAdam(critic2)
|
||||
result.alpha = initAdamVar(ones[float32](1))
|
||||
result.initialized = true
|
||||
|
||||
# ── Internal: temp dir per save ───────────────────────────────────────────────
|
||||
|
||||
proc tmpDir(): string =
|
||||
getTempDir() / ("sacw_" & $int(epochTime() * 1000))
|
||||
|
||||
# ── Save helpers ──────────────────────────────────────────────────────────────
|
||||
|
||||
template addT(z: var ZipArchive; name: string; t: Tensor[float32]; tmp: string) =
|
||||
## Write tensor to a temp file, add to zip, delete temp file.
|
||||
let p = tmp / name
|
||||
t.write_npy(p)
|
||||
z.addFile(name, p)
|
||||
|
||||
proc addLinear(z: var ZipArchive; prefix: string; l: Linear; tmp: string) =
|
||||
addT(z, prefix & "_w.npy", l.w, tmp)
|
||||
addT(z, prefix & "_b.npy", l.b, tmp)
|
||||
|
||||
proc addLSTMCell(z: var ZipArchive; prefix: string; c: LSTMCell; tmp: string) =
|
||||
addT(z, prefix & "_wc.npy", c.wCombined, tmp)
|
||||
addT(z, prefix & "_bc.npy", c.bCombined, tmp)
|
||||
|
||||
proc addActorNet(z: var ZipArchive; prefix: string; a: ActorNet; tmp: string) =
|
||||
addLinear(z, prefix & "_fc1", a.fc1, tmp)
|
||||
addLSTMCell(z, prefix & "_lstm", a.lstm, tmp)
|
||||
addLinear(z, prefix & "_fc2", a.fc2, tmp)
|
||||
addLinear(z, prefix & "_mu", a.muHead, tmp)
|
||||
addLinear(z, prefix & "_logstd", a.logStdHead, tmp)
|
||||
|
||||
proc addCriticNet(z: var ZipArchive; prefix: string; c: CriticNet; tmp: string) =
|
||||
addLinear(z, prefix & "_fc1", c.fc1, tmp)
|
||||
addLSTMCell(z, prefix & "_lstm", c.lstm, tmp)
|
||||
addLinear(z, prefix & "_fc2", c.fc2, tmp)
|
||||
addLinear(z, prefix & "_fc3", c.fc3, tmp)
|
||||
|
||||
proc addAdamVar(z: var ZipArchive; prefix: string; v: AdamVar; tmp: string) =
|
||||
addT(z, prefix & "_m.npy", v.m, tmp)
|
||||
addT(z, prefix & "_v.npy", v.v, tmp)
|
||||
|
||||
proc addLinearAdam(z: var ZipArchive; prefix: string; la: LinearAdam; tmp: string) =
|
||||
addAdamVar(z, prefix & "_w", la.w, tmp)
|
||||
addAdamVar(z, prefix & "_b", la.b, tmp)
|
||||
|
||||
proc addLSTMCellAdam(z: var ZipArchive; prefix: string; la: LSTMCellAdam; tmp: string) =
|
||||
addAdamVar(z, prefix & "_wc", la.wCombined, tmp)
|
||||
addAdamVar(z, prefix & "_bc", la.bCombined, tmp)
|
||||
|
||||
proc addActorAdam(z: var ZipArchive; prefix: string; a: ActorAdam; tmp: string) =
|
||||
addLinearAdam(z, prefix & "_fc1", a.fc1, tmp)
|
||||
addLSTMCellAdam(z, prefix & "_lstm", a.lstm, tmp)
|
||||
addLinearAdam(z, prefix & "_fc2", a.fc2, tmp)
|
||||
addLinearAdam(z, prefix & "_mu", a.muHead, tmp)
|
||||
addLinearAdam(z, prefix & "_logstd", a.logStdHead, tmp)
|
||||
|
||||
proc addCriticAdam(z: var ZipArchive; prefix: string; c: CriticAdam; tmp: string) =
|
||||
addLinearAdam(z, prefix & "_fc1", c.fc1, tmp)
|
||||
addLSTMCellAdam(z, prefix & "_lstm", c.lstm, tmp)
|
||||
addLinearAdam(z, prefix & "_fc2", c.fc2, tmp)
|
||||
addLinearAdam(z, prefix & "_fc3", c.fc3, tmp)
|
||||
|
||||
# ── Public API ────────────────────────────────────────────────────────────────
|
||||
|
||||
proc saveWeights*(path: string;
|
||||
actor: ActorNet;
|
||||
critic1, critic2: CriticNet;
|
||||
targetCritic1, targetCritic2: CriticNet;
|
||||
alpha: float32) =
|
||||
## Save network tensors (no Adam states) to `path` (.zip).
|
||||
## Atomic: writes to a temp path first, then renames.
|
||||
let tmp = tmpDir()
|
||||
createDir(tmp)
|
||||
let tmpZip = path & ".tmp"
|
||||
try:
|
||||
var z: ZipArchive
|
||||
if not z.open(tmpZip, fmWrite):
|
||||
raise newException(IOError, "cannot create zip: " & tmpZip)
|
||||
addActorNet(z, "actor", actor, tmp)
|
||||
addCriticNet(z, "c1", critic1, tmp)
|
||||
addCriticNet(z, "c2", critic2, tmp)
|
||||
addCriticNet(z, "tc1", targetCritic1,tmp)
|
||||
addCriticNet(z, "tc2", targetCritic2,tmp)
|
||||
# alpha: store as a 1-element tensor
|
||||
addT(z, "alpha.npy", [alpha].toTensor.asType(float32), tmp)
|
||||
z.close()
|
||||
createDir(path.parentDir)
|
||||
moveFile(tmpZip, path)
|
||||
finally:
|
||||
removeDir(tmp)
|
||||
if fileExists(tmpZip): removeFile(tmpZip)
|
||||
|
||||
proc saveCheckpoint*(path: string;
|
||||
actor: ActorNet;
|
||||
critic1, critic2: CriticNet;
|
||||
targetCritic1, targetCritic2: CriticNet;
|
||||
alpha: float32;
|
||||
adam: SACAdamStates) =
|
||||
## Save networks + Adam states to `path` (.zip). Atomic.
|
||||
let tmp = tmpDir()
|
||||
createDir(tmp)
|
||||
let tmpZip = path & ".tmp"
|
||||
try:
|
||||
var z: ZipArchive
|
||||
if not z.open(tmpZip, fmWrite):
|
||||
raise newException(IOError, "cannot create zip: " & tmpZip)
|
||||
addActorNet(z, "actor", actor, tmp)
|
||||
addCriticNet(z, "c1", critic1, tmp)
|
||||
addCriticNet(z, "c2", critic2, tmp)
|
||||
addCriticNet(z, "tc1", targetCritic1,tmp)
|
||||
addCriticNet(z, "tc2", targetCritic2,tmp)
|
||||
addT(z, "alpha.npy", [alpha].toTensor.asType(float32), tmp)
|
||||
if adam.initialized:
|
||||
addActorAdam(z, "adam_actor", adam.actor, tmp)
|
||||
addCriticAdam(z, "adam_c1", adam.critic1, tmp)
|
||||
addCriticAdam(z, "adam_c2", adam.critic2, tmp)
|
||||
addAdamVar(z, "adam_alpha", adam.alpha, tmp)
|
||||
# t counters (all in lockstep; store as text)
|
||||
writeFile(tmp / "adam_t.txt",
|
||||
$adam.actor.fc1.w.t & "\n" &
|
||||
$adam.critic1.fc1.w.t & "\n" &
|
||||
$adam.critic2.fc1.w.t & "\n" &
|
||||
$adam.alpha.t)
|
||||
z.addFile("adam_t.txt", tmp / "adam_t.txt")
|
||||
z.close()
|
||||
createDir(path.parentDir)
|
||||
moveFile(tmpZip, path)
|
||||
finally:
|
||||
removeDir(tmp)
|
||||
if fileExists(tmpZip): removeFile(tmpZip)
|
||||
|
||||
# ── Load helpers ──────────────────────────────────────────────────────────────
|
||||
|
||||
template loadT(name: string; tmp: string): Tensor[float32] =
|
||||
read_npy[float32](tmp / name)
|
||||
|
||||
proc loadLinear(z: var ZipArchive; prefix, tmp: string): Linear =
|
||||
z.extractFile(prefix & "_w.npy", tmp / (prefix & "_w.npy"))
|
||||
z.extractFile(prefix & "_b.npy", tmp / (prefix & "_b.npy"))
|
||||
result.w = read_npy[float32](tmp / (prefix & "_w.npy"))
|
||||
result.b = read_npy[float32](tmp / (prefix & "_b.npy"))
|
||||
|
||||
proc loadLSTMCell(z: var ZipArchive; prefix, tmp: string): LSTMCell =
|
||||
z.extractFile(prefix & "_wc.npy", tmp / (prefix & "_wc.npy"))
|
||||
z.extractFile(prefix & "_bc.npy", tmp / (prefix & "_bc.npy"))
|
||||
result.wCombined = read_npy[float32](tmp / (prefix & "_wc.npy"))
|
||||
result.bCombined = read_npy[float32](tmp / (prefix & "_bc.npy"))
|
||||
result.hiddenDim = result.bCombined.shape[0] div 4
|
||||
|
||||
proc loadActorNet(z: var ZipArchive; prefix, tmp: string): ActorNet =
|
||||
result.fc1 = loadLinear(z, prefix & "_fc1", tmp)
|
||||
result.lstm = loadLSTMCell(z, prefix & "_lstm", tmp)
|
||||
result.fc2 = loadLinear(z, prefix & "_fc2", tmp)
|
||||
result.muHead = loadLinear(z, prefix & "_mu", tmp)
|
||||
result.logStdHead = loadLinear(z, prefix & "_logstd", tmp)
|
||||
result.hiddenDim = result.lstm.hiddenDim
|
||||
|
||||
proc loadCriticNet(z: var ZipArchive; prefix, tmp: string): CriticNet =
|
||||
result.fc1 = loadLinear(z, prefix & "_fc1", tmp)
|
||||
result.lstm = loadLSTMCell(z, prefix & "_lstm", tmp)
|
||||
result.fc2 = loadLinear(z, prefix & "_fc2", tmp)
|
||||
result.fc3 = loadLinear(z, prefix & "_fc3", tmp)
|
||||
result.hiddenDim = result.lstm.hiddenDim
|
||||
|
||||
proc loadAdamVarFromZip(z: var ZipArchive; prefix, tmp: string): AdamVar =
|
||||
z.extractFile(prefix & "_m.npy", tmp / (prefix & "_m.npy"))
|
||||
z.extractFile(prefix & "_v.npy", tmp / (prefix & "_v.npy"))
|
||||
result.m = read_npy[float32](tmp / (prefix & "_m.npy"))
|
||||
result.v = read_npy[float32](tmp / (prefix & "_v.npy"))
|
||||
|
||||
proc loadLinearAdam(z: var ZipArchive; prefix, tmp: string): LinearAdam =
|
||||
result.w = loadAdamVarFromZip(z, prefix & "_w", tmp)
|
||||
result.b = loadAdamVarFromZip(z, prefix & "_b", tmp)
|
||||
|
||||
proc loadLSTMCellAdam(z: var ZipArchive; prefix, tmp: string): LSTMCellAdam =
|
||||
result.wCombined = loadAdamVarFromZip(z, prefix & "_wc", tmp)
|
||||
result.bCombined = loadAdamVarFromZip(z, prefix & "_bc", tmp)
|
||||
|
||||
proc loadActorAdam(z: var ZipArchive; prefix, tmp: string): ActorAdam =
|
||||
result.fc1 = loadLinearAdam(z, prefix & "_fc1", tmp)
|
||||
result.lstm = loadLSTMCellAdam(z, prefix & "_lstm", tmp)
|
||||
result.fc2 = loadLinearAdam(z, prefix & "_fc2", tmp)
|
||||
result.muHead = loadLinearAdam(z, prefix & "_mu", tmp)
|
||||
result.logStdHead = loadLinearAdam(z, prefix & "_logstd", tmp)
|
||||
|
||||
proc loadCriticAdam(z: var ZipArchive; prefix, tmp: string): CriticAdam =
|
||||
result.fc1 = loadLinearAdam(z, prefix & "_fc1", tmp)
|
||||
result.lstm = loadLSTMCellAdam(z, prefix & "_lstm", tmp)
|
||||
result.fc2 = loadLinearAdam(z, prefix & "_fc2", tmp)
|
||||
result.fc3 = loadLinearAdam(z, prefix & "_fc3", tmp)
|
||||
|
||||
type
|
||||
WeightCheckpoint* = object
|
||||
actor*: ActorNet
|
||||
critic1*: CriticNet
|
||||
critic2*: CriticNet
|
||||
targetCritic1*: CriticNet
|
||||
targetCritic2*: CriticNet
|
||||
alpha*: float32
|
||||
adam*: SACAdamStates ## initialized=false if not present in zip
|
||||
|
||||
proc loadCheckpoint*(path: string): WeightCheckpoint =
|
||||
## Load all tensors from `path` (.zip). Raises IOError if file not found.
|
||||
## Adam states loaded only if present; result.adam.initialized reflects this.
|
||||
if not fileExists(path):
|
||||
raise newException(IOError, "checkpoint not found: " & path)
|
||||
let tmp = tmpDir()
|
||||
createDir(tmp)
|
||||
try:
|
||||
var z: ZipArchive
|
||||
if not z.open(path, fmRead):
|
||||
raise newException(IOError, "cannot open zip: " & path)
|
||||
|
||||
result.actor = loadActorNet(z, "actor", tmp)
|
||||
result.critic1 = loadCriticNet(z, "c1", tmp)
|
||||
result.critic2 = loadCriticNet(z, "c2", tmp)
|
||||
result.targetCritic1 = loadCriticNet(z, "tc1", tmp)
|
||||
result.targetCritic2 = loadCriticNet(z, "tc2", tmp)
|
||||
|
||||
z.extractFile("alpha.npy", tmp / "alpha.npy")
|
||||
let alphaTensor = read_npy[float32](tmp / "alpha.npy")
|
||||
result.alpha = alphaTensor[0]
|
||||
|
||||
# Adam states — optional
|
||||
var hasAdam = false
|
||||
for f in z.walkFiles:
|
||||
if f.startsWith("adam_"):
|
||||
hasAdam = true
|
||||
break
|
||||
if hasAdam:
|
||||
result.adam.actor = loadActorAdam(z, "adam_actor", tmp)
|
||||
result.adam.critic1 = loadCriticAdam(z, "adam_c1", tmp)
|
||||
result.adam.critic2 = loadCriticAdam(z, "adam_c2", tmp)
|
||||
result.adam.alpha = loadAdamVarFromZip(z, "adam_alpha", tmp)
|
||||
# t counters
|
||||
z.extractFile("adam_t.txt", tmp / "adam_t.txt")
|
||||
let ts = readFile(tmp / "adam_t.txt").strip().splitLines()
|
||||
if ts.len >= 4:
|
||||
let tActor = parseInt(ts[0])
|
||||
let tCritic1 = parseInt(ts[1])
|
||||
let tCritic2 = parseInt(ts[2])
|
||||
let tAlpha = parseInt(ts[3])
|
||||
# propagate t to all Adam vars
|
||||
template setT(v: var AdamVar; tval: int) = v.t = tval
|
||||
setT(result.adam.actor.fc1.w, tActor)
|
||||
setT(result.adam.actor.fc1.b, tActor)
|
||||
setT(result.adam.actor.lstm.wCombined,tActor)
|
||||
setT(result.adam.actor.lstm.bCombined,tActor)
|
||||
setT(result.adam.actor.fc2.w, tActor)
|
||||
setT(result.adam.actor.fc2.b, tActor)
|
||||
setT(result.adam.actor.muHead.w, tActor)
|
||||
setT(result.adam.actor.muHead.b, tActor)
|
||||
setT(result.adam.actor.logStdHead.w, tActor)
|
||||
setT(result.adam.actor.logStdHead.b, tActor)
|
||||
setT(result.adam.critic1.fc1.w, tCritic1)
|
||||
setT(result.adam.critic1.fc1.b, tCritic1)
|
||||
setT(result.adam.critic1.lstm.wCombined,tCritic1)
|
||||
setT(result.adam.critic1.lstm.bCombined,tCritic1)
|
||||
setT(result.adam.critic1.fc2.w, tCritic1)
|
||||
setT(result.adam.critic1.fc2.b, tCritic1)
|
||||
setT(result.adam.critic1.fc3.w, tCritic1)
|
||||
setT(result.adam.critic1.fc3.b, tCritic1)
|
||||
setT(result.adam.critic2.fc1.w, tCritic2)
|
||||
setT(result.adam.critic2.fc1.b, tCritic2)
|
||||
setT(result.adam.critic2.lstm.wCombined,tCritic2)
|
||||
setT(result.adam.critic2.lstm.bCombined,tCritic2)
|
||||
setT(result.adam.critic2.fc2.w, tCritic2)
|
||||
setT(result.adam.critic2.fc2.b, tCritic2)
|
||||
setT(result.adam.critic2.fc3.w, tCritic2)
|
||||
setT(result.adam.critic2.fc3.b, tCritic2)
|
||||
setT(result.adam.alpha, tAlpha)
|
||||
result.adam.initialized = true
|
||||
|
||||
z.close()
|
||||
finally:
|
||||
removeDir(tmp)
|
||||
@@ -0,0 +1,2 @@
|
||||
switch("path", "../src")
|
||||
switch("path", "../../libs")
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user