Compare commits
29 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 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 | |||
| 4ee0d8272c | |||
| add3e34926 |
@@ -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"
|
||||
}
|
||||
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"
|
||||
@@ -1,7 +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")
|
||||
# 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"
|
||||
|
||||
@@ -0,0 +1,275 @@
|
||||
# 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
|
||||
|
||||
- **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 |
|
||||
|
||||
|
||||
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
|
||||
@@ -1,8 +1,8 @@
|
||||
{
|
||||
"name": "Recurrent Royalty",
|
||||
"name": "SAC_LSTM_Bot",
|
||||
"version": "0.1.0",
|
||||
"authors": ["Davide Cappellini"],
|
||||
"description": "SAC+LSTM Tank Royale bot — skeleton with radar lock",
|
||||
"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"],
|
||||
|
||||
@@ -1,13 +1,30 @@
|
||||
## SAC_LSTM_Bot — skeleton: radar lock + "Recurrent Royalty" color scheme.
|
||||
## No RL yet. Connects, sets colors, locks radar onto enemy.
|
||||
## 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
|
||||
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
|
||||
|
||||
const botJsonPath = currentSourcePath().parentDir / "SAC_LSTM_Bot.json"
|
||||
# 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")
|
||||
@@ -26,35 +43,307 @@ proc applyColors() =
|
||||
setBulletColor(ColBullet)
|
||||
setTracksColor(ColTracks)
|
||||
|
||||
# ── Bot type ──────────────────────────────────────────────────────────────────
|
||||
# ── 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
|
||||
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))
|
||||
|
||||
# ── Run loop ──────────────────────────────────────────────────────────────────
|
||||
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():
|
||||
# Spin radar when no enemy is visible (full sweep).
|
||||
if bot.enemyBearing == 0.0:
|
||||
setRadarTurnRate(45.0)
|
||||
# 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)
|
||||
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
|
||||
@@ -3,6 +3,25 @@
|
||||
|
||||
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*(
|
||||
@@ -10,6 +29,9 @@ proc computeReward*(
|
||||
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 =
|
||||
@@ -17,8 +39,13 @@ proc computeReward*(
|
||||
## Damage formula: 4p + 2(p-1) = 6p - 2 (matches Tank Royale bullet rules).
|
||||
let p = damageInflicted
|
||||
let pe = damageReceived
|
||||
if p > 0.0: result += 6.0 * p - 2.0
|
||||
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
|
||||
|
||||
@@ -0,0 +1,614 @@
|
||||
## 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))
|
||||
# Recover pre-tanh z ≈ arctanh(action)
|
||||
let z = action.map(proc(a: float32): float32 =
|
||||
let ac = clamp(a, -1.0'f32 + 1e-6'f32, 1.0'f32 - 1e-6'f32)
|
||||
0.5'f32 * ln((1.0'f32 + ac) / (1.0'f32 - ac)))
|
||||
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,97 @@
|
||||
import unittest
|
||||
import arraymancer
|
||||
import std/math
|
||||
import SAC_LSTM_Bot/actions
|
||||
|
||||
proc makeOutput(a0, a1, a2, a3: float): Tensor[float32] =
|
||||
result = newTensor[float32](4)
|
||||
result[0] = a0.float32
|
||||
result[1] = a1.float32
|
||||
result[2] = a2.float32
|
||||
result[3] = a3.float32
|
||||
|
||||
suite "mapActions":
|
||||
|
||||
test "ACTION_DIM is 4":
|
||||
check ACTION_DIM == 4
|
||||
|
||||
# Speed-aware turn rate
|
||||
test "turn: output +1 at speed 0 -> +10 degrees":
|
||||
let m = mapActions(makeOutput(1.0, 0.0, 0.0, -1.0), 0.0, 0.0)
|
||||
check abs(m.turnRate - 10.0) < 1e-6
|
||||
|
||||
test "turn: output -1 at speed 0 -> -10 degrees":
|
||||
let m = mapActions(makeOutput(-1.0, 0.0, 0.0, -1.0), 0.0, 0.0)
|
||||
check abs(m.turnRate - (-10.0)) < 1e-6
|
||||
|
||||
test "turn: output +1 at speed 8 -> +4 degrees":
|
||||
let m = mapActions(makeOutput(1.0, 0.0, 0.0, -1.0), 8.0, 0.0)
|
||||
check abs(m.turnRate - 4.0) < 1e-6
|
||||
|
||||
test "turn: output -1 at speed 8 -> -4 degrees":
|
||||
let m = mapActions(makeOutput(-1.0, 0.0, 0.0, -1.0), 8.0, 0.0)
|
||||
check abs(m.turnRate - (-4.0)) < 1e-6
|
||||
|
||||
test "turn: output 0 -> 0 regardless of speed":
|
||||
let m = mapActions(makeOutput(0.0, 0.0, 0.0, -1.0), 5.0, 0.0)
|
||||
check abs(m.turnRate) < 1e-6
|
||||
|
||||
# Acceleration range
|
||||
test "accel: output -1 -> -2.0":
|
||||
let m = mapActions(makeOutput(0.0, -1.0, 0.0, -1.0), 0.0, 0.0)
|
||||
check abs(m.acceleration - (-2.0)) < 1e-6
|
||||
|
||||
test "accel: output +1 -> +1.0":
|
||||
let m = mapActions(makeOutput(0.0, 1.0, 0.0, -1.0), 0.0, 0.0)
|
||||
check abs(m.acceleration - 1.0) < 1e-6
|
||||
|
||||
test "accel: output 0 -> midpoint -0.5":
|
||||
# value*1.5 - 0.5 at value=0 -> -0.5 (correct midpoint between -2 and +1)
|
||||
let m = mapActions(makeOutput(0.0, 0.0, 0.0, -1.0), 0.0, 0.0)
|
||||
check abs(m.acceleration - (-0.5)) < 1e-6
|
||||
|
||||
# Gun turn rate
|
||||
test "gun turn: output +1 -> +20 degrees":
|
||||
let m = mapActions(makeOutput(0.0, 0.0, 1.0, -1.0), 0.0, 0.0)
|
||||
check abs(m.gunTurnRate - 20.0) < 1e-6
|
||||
|
||||
test "gun turn: output -1 -> -20 degrees":
|
||||
let m = mapActions(makeOutput(0.0, 0.0, -1.0, -1.0), 0.0, 0.0)
|
||||
check abs(m.gunTurnRate - (-20.0)) < 1e-6
|
||||
|
||||
test "gun turn: output 0 -> 0":
|
||||
let m = mapActions(makeOutput(0.0, 0.0, 0.0, -1.0), 0.0, 0.0)
|
||||
check abs(m.gunTurnRate) < 1e-6
|
||||
|
||||
# Fire threshold
|
||||
test "fire: output -0.5 -> no fire (firePower == 0)":
|
||||
let m = mapActions(makeOutput(0.0, 0.0, 0.0, -0.5), 0.0, 0.0)
|
||||
check m.firePower == 0.0
|
||||
|
||||
test "fire: output 0.0 -> no fire (boundary, not positive)":
|
||||
let m = mapActions(makeOutput(0.0, 0.0, 0.0, 0.0), 0.0, 0.0)
|
||||
check m.firePower == 0.0
|
||||
|
||||
test "fire: output +0.5 -> fire with correct power":
|
||||
let m = mapActions(makeOutput(0.0, 0.0, 0.0, 0.5), 0.0, 0.0)
|
||||
# 0.5 * 2.9 + 0.1 = 1.55
|
||||
check abs(m.firePower - 1.55) < 1e-5
|
||||
|
||||
test "fire: output +1.0 -> fire power near 3.0":
|
||||
let m = mapActions(makeOutput(0.0, 0.0, 0.0, 1.0), 0.0, 0.0)
|
||||
check abs(m.firePower - 3.0) < 1e-5
|
||||
|
||||
test "fire: output +1.0 but gunHeat > 0 -> no fire":
|
||||
let m = mapActions(makeOutput(0.0, 0.0, 0.0, 1.0), 0.0, 1.5)
|
||||
check m.firePower == 0.0
|
||||
|
||||
test "fire: output +0.001 (just above 0) -> fires with power near 0.1":
|
||||
let m = mapActions(makeOutput(0.0, 0.0, 0.0, 0.001), 0.0, 0.0)
|
||||
check m.firePower > 0.0
|
||||
check m.firePower < 0.2
|
||||
|
||||
# Speed-aware turn with negative speed (reverse)
|
||||
test "turn: speed -8 (reversing) -> same magnitude as speed +8":
|
||||
let fwd = mapActions(makeOutput(1.0, 0.0, 0.0, -1.0), 8.0, 0.0)
|
||||
let rev = mapActions(makeOutput(1.0, 0.0, 0.0, -1.0), -8.0, 0.0)
|
||||
check abs(fwd.turnRate - rev.turnRate) < 1e-6
|
||||
@@ -0,0 +1,150 @@
|
||||
## Tests for integration.nim (#48) — assert-based, no framework.
|
||||
## Covers: TrainingMsg channel round-trip (plain arrays through a channel),
|
||||
## drain-then-train NewBattle/Shutdown handling, trainPass safety below canSample.
|
||||
|
||||
import arraymancer except Linear
|
||||
import std/[locks, json, os]
|
||||
import tankroyale_botapi # updateBotNames: seed the vendored id->name table
|
||||
import SAC_LSTM_Bot/integration
|
||||
import SAC_LSTM_Bot/state # STATE_DIM
|
||||
import SAC_LSTM_Bot/actions # ACTION_DIM
|
||||
import SAC_LSTM_Bot/training # initSACTrainer
|
||||
import SAC_LSTM_Bot/replay_buffer
|
||||
|
||||
# ── 1. TrainingMsg round-trips through a channel with arrays intact ──────────
|
||||
block:
|
||||
var ch: Channel[TrainingMsg]
|
||||
ch.open(4)
|
||||
var msg = TrainingMsg(kind: tmkTransition)
|
||||
for i in 0 ..< STATE_DIM:
|
||||
msg.state[i] = float32(i) * 0.5'f32
|
||||
msg.nextState[i] = float32(i) * 2.0'f32
|
||||
for i in 0 ..< ACTION_DIM:
|
||||
msg.action[i] = float32(i) - 2.0'f32
|
||||
msg.reward = -1.25'f32
|
||||
msg.done = true
|
||||
assert ch.trySend(msg)
|
||||
assert ch.trySend(TrainingMsg(kind: tmkNewBattle, enemyId: 4242))
|
||||
assert ch.trySend(TrainingMsg(kind: tmkShutdown))
|
||||
|
||||
let r1 = ch.recv()
|
||||
assert r1.kind == tmkTransition, "first msg is a transition"
|
||||
for i in 0 ..< STATE_DIM:
|
||||
assert r1.state[i] == float32(i) * 0.5'f32, "state round-trip at " & $i
|
||||
assert r1.nextState[i] == float32(i) * 2.0'f32, "nextState round-trip at " & $i
|
||||
for i in 0 ..< ACTION_DIM:
|
||||
assert r1.action[i] == float32(i) - 2.0'f32, "action round-trip at " & $i
|
||||
assert r1.reward == -1.25'f32 and r1.done
|
||||
|
||||
let r2 = ch.recv()
|
||||
assert r2.kind == tmkNewBattle and r2.enemyId == 4242
|
||||
let r3 = ch.recv()
|
||||
assert r3.kind == tmkShutdown
|
||||
|
||||
ch.close()
|
||||
# closed + empty -> tryRecv reports no data (Nim 2.2: recv would block forever)
|
||||
let (ok4, _) = ch.tryRecv()
|
||||
assert not ok4, "closed channel must report dataAvailable=false"
|
||||
echo "PASS TrainingMsg channel round-trip"
|
||||
|
||||
# ── 2. NewBattle clears only when the opponent NAME changes; Shutdown stops ───
|
||||
block:
|
||||
# Seed the API's id->name table (v1.0.1) for name-based identity (#49).
|
||||
updateBotNames(parseJson(
|
||||
"""{"bots":[{"id":7,"name":"Corners"},{"id":8,"name":"Crazy"}]}"""))
|
||||
assert opponentKey(7) == "Corners", "known id resolves to name"
|
||||
assert opponentKey(99) == "99", "unknown id falls back to numeric string"
|
||||
|
||||
var st: TrainState
|
||||
st.trainer = initSACTrainer(STATE_DIM, ACTION_DIM)
|
||||
st.buf = newReplayBuffer(64, STATE_DIM, ACTION_DIM, burnIn = 2, trainWindow = 3)
|
||||
|
||||
assert handleTrainingMsg(st, TrainingMsg(kind: tmkNewBattle, enemyId: 7))
|
||||
assert st.lastEnemyKey == "Corners"
|
||||
|
||||
for i in 0 ..< 5:
|
||||
var m = TrainingMsg(kind: tmkTransition)
|
||||
m.reward = float32(i)
|
||||
assert handleTrainingMsg(st, m)
|
||||
assert st.buf.len == 5, "transitions stored"
|
||||
|
||||
# Same opponent -> buffer kept (Q12a).
|
||||
assert handleTrainingMsg(st, TrainingMsg(kind: tmkNewBattle, enemyId: 7))
|
||||
assert st.buf.len == 5, "same opponent must NOT clear"
|
||||
|
||||
# Opponent changed -> clear.
|
||||
assert handleTrainingMsg(st, TrainingMsg(kind: tmkNewBattle, enemyId: 8))
|
||||
assert st.buf.len == 0, "opponent change must clear"
|
||||
assert st.lastEnemyKey == "Crazy"
|
||||
|
||||
# Nameless window (pre-BotListUpdate) is a distinct key.
|
||||
assert handleTrainingMsg(st, TrainingMsg(kind: tmkNewBattle, enemyId: 99))
|
||||
assert st.lastEnemyKey == "99" and st.buf.len == 0
|
||||
|
||||
# Shutdown stops the caller's loop.
|
||||
assert not handleTrainingMsg(st, TrainingMsg(kind: tmkShutdown))
|
||||
echo "PASS name-keyed NewBattle clear / Shutdown"
|
||||
|
||||
# ── 2b. bumpRoundCounter increments per call, cold-starts at 1 ────────────────
|
||||
block:
|
||||
let tmp = getTempDir() / "sac_test_weights_" & $getCurrentProcessId()
|
||||
putEnv("SACLSTM_WEIGHTS_PATH", tmp / "sac_latest.zip")
|
||||
bumpRoundCounter()
|
||||
bumpRoundCounter()
|
||||
assert readFile(tmp / "round_counter.txt") == "2", "counter increments per round"
|
||||
echo "PASS bumpRoundCounter"
|
||||
|
||||
# ── 3. trainPass is a safe no-op below canSample (no steps, no publish) ───────
|
||||
block:
|
||||
var st: TrainState
|
||||
st.trainer = initSACTrainer(STATE_DIM, ACTION_DIM)
|
||||
st.buf = newReplayBuffer(64, STATE_DIM, ACTION_DIM, burnIn = 2, trainWindow = 3)
|
||||
for i in 0 ..< 4:
|
||||
assert handleTrainingMsg(st, TrainingMsg(kind: tmkTransition))
|
||||
trainPass(st, 4) # 4 < burnIn+trainWindow = 5
|
||||
assert st.stepCount == 0, "no gradient steps below canSample"
|
||||
echo "PASS trainPass no-op below canSample"
|
||||
|
||||
# ── 4. Flat snapshot layout: pack/unpack round-trips weights exactly ──────────
|
||||
block:
|
||||
let t0 = initSACTrainer(STATE_DIM, ACTION_DIM)
|
||||
let fs = packFull(t0)
|
||||
assert fs.data.len == actorSize(t0.actor.hiddenDim) + 4 * criticSize(t0.actor.hiddenDim) + 1
|
||||
let (a, c1, c2, tc1, tc2, alpha) = unpackFull(fs)
|
||||
assert a.hiddenDim == t0.actor.hiddenDim
|
||||
assert c1.fc3.b.shape[0] == 1
|
||||
let fw = a.muHead.w.flatten()
|
||||
let fw0 = t0.actor.muHead.w.flatten()
|
||||
for i in 0 ..< fw.size:
|
||||
assert fw[i] == fw0[i], "actor mu weights round-trip"
|
||||
for i in 0 ..< c2.lstm.bCombined.size:
|
||||
assert c2.lstm.bCombined[i] == t0.critic2.lstm.bCombined[i], "critic lstm bias round-trip"
|
||||
assert alpha == t0.alpha()
|
||||
discard tc1
|
||||
echo "PASS flat snapshot pack/unpack round-trip"
|
||||
|
||||
# ── 5. Lever 3 (#59): metricsLine emits exactly the exposed trainer scalars ───
|
||||
block:
|
||||
let line = metricsLine(1787394115.123, 42, 500, 20, 20,
|
||||
SACMetrics(criticLoss: 0.5'f32, actorLoss: -1.5'f32,
|
||||
alphaLoss: 0.25'f32, alpha: 2.0'f32))
|
||||
let j = parseJson(line) # throws on malformed JSONL
|
||||
assert j["steps"].getInt() == 42 and j["buffer_size"].getInt() == 500
|
||||
assert j["drained"].getInt() == 20 and j["grad_steps"].getInt() == 20
|
||||
assert abs(j["critic_loss"].getFloat() - 0.5) < 1e-3
|
||||
assert abs(j["actor_loss"].getFloat() + 1.5) < 1e-3
|
||||
assert abs(j["alpha_loss"].getFloat() - 0.25) < 1e-3
|
||||
assert abs(j["alpha"].getFloat() - 2.0) < 1e-3
|
||||
assert j["epoch"].getFloat() > 1e9
|
||||
echo "PASS metricsLine JSONL scalars"
|
||||
|
||||
# ── 6. Lever 4 (#59): SACLSTM_EVAL_MODE=1 suppresses training input ───────────
|
||||
block:
|
||||
putEnv("SACLSTM_EVAL_MODE", "1")
|
||||
# Gate fires before any channel traffic: false = dropped, nothing enqueued.
|
||||
assert not sendTrainingMsg(TrainingMsg(kind: tmkTransition)),
|
||||
"eval mode must drop transitions"
|
||||
assert not sendTrainingMsg(TrainingMsg(kind: tmkNewBattle, enemyId: 7)),
|
||||
"eval mode must drop NewBattle (no buffer clears from eval)"
|
||||
delEnv("SACLSTM_EVAL_MODE")
|
||||
echo "PASS eval-mode training-input suppression"
|
||||
@@ -0,0 +1,99 @@
|
||||
import unittest
|
||||
import arraymancer
|
||||
import std/[math, os]
|
||||
import SAC_LSTM_Bot/network
|
||||
|
||||
const
|
||||
STATE_DIM = 20
|
||||
ACTION_DIM = 4
|
||||
|
||||
suite "ActorNet":
|
||||
setup:
|
||||
let actor = initActorNet(STATE_DIM)
|
||||
let state = randomNormalTensor[float32](STATE_DIM)
|
||||
let ls0 = zeroState(actor.hiddenDim)
|
||||
|
||||
test "forward output shapes":
|
||||
let (actions, lp, ls1) = actorForward(actor, state, ls0)
|
||||
check actions.shape == [4]
|
||||
check ls1.h.shape == [actor.hiddenDim]
|
||||
check ls1.c.shape == [actor.hiddenDim]
|
||||
check not lp.isNaN
|
||||
|
||||
test "actions clamped in (-1, 1) — tanh output":
|
||||
let (actions, _, _) = actorForward(actor, state, ls0)
|
||||
for i in 0 ..< 4:
|
||||
check actions[i] > -1.0'f32
|
||||
check actions[i] < 1.0'f32
|
||||
|
||||
test "hidden state propagates (h/c change after step)":
|
||||
let (_, _, ls1) = actorForward(actor, state, ls0)
|
||||
# h' should differ from zero init for non-trivial input
|
||||
var hChanged = false
|
||||
for i in 0 ..< actor.hiddenDim:
|
||||
if abs(ls1.h[i] - ls0.h[i]) > 1e-7'f32:
|
||||
hChanged = true
|
||||
break
|
||||
check hChanged
|
||||
|
||||
test "deterministic mode: same input → same output":
|
||||
let (a1, _, _) = actorForward(actor, state, ls0, deterministic = true)
|
||||
let (a2, _, _) = actorForward(actor, state, ls0, deterministic = true)
|
||||
for i in 0 ..< 4:
|
||||
check abs(a1[i] - a2[i]) < 1e-7'f32
|
||||
|
||||
test "stochastic mode: outputs may differ (sampling)":
|
||||
# Run many times; at least one pair should differ
|
||||
let (a1, _, _) = actorForward(actor, state, ls0)
|
||||
let (a2, _, _) = actorForward(actor, state, ls0)
|
||||
var anyDiff = false
|
||||
for i in 0 ..< 4:
|
||||
if abs(a1[i] - a2[i]) > 1e-7'f32:
|
||||
anyDiff = true
|
||||
break
|
||||
check anyDiff
|
||||
|
||||
suite "CriticNet":
|
||||
setup:
|
||||
let critic = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let state = randomNormalTensor[float32](STATE_DIM)
|
||||
let actions = randomNormalTensor[float32](ACTION_DIM)
|
||||
let stateAct = concat(state, actions, axis = 0)
|
||||
let ls0 = zeroState(critic.hiddenDim)
|
||||
|
||||
test "forward output shape":
|
||||
let (q, ls1) = criticForward(critic, stateAct, ls0)
|
||||
check ls1.h.shape == [critic.hiddenDim]
|
||||
check ls1.c.shape == [critic.hiddenDim]
|
||||
# q is scalar float — just ensure it doesn't NaN
|
||||
check not q.isNaN
|
||||
|
||||
test "hidden state propagates":
|
||||
let (_, ls1) = criticForward(critic, stateAct, ls0)
|
||||
var hChanged = false
|
||||
for i in 0 ..< critic.hiddenDim:
|
||||
if abs(ls1.h[i] - ls0.h[i]) > 1e-7'f32:
|
||||
hChanged = true
|
||||
break
|
||||
check hChanged
|
||||
|
||||
suite "Dual critics":
|
||||
test "two independent critics produce different Q values":
|
||||
let c1 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let c2 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let sa = randomNormalTensor[float32](STATE_DIM + ACTION_DIM)
|
||||
let ls = zeroState(c1.hiddenDim)
|
||||
let (q1, _) = criticForward(c1, sa, ls)
|
||||
let (q2, _) = criticForward(c2, sa, ls)
|
||||
check abs(q1 - q2) > 1e-7'f32
|
||||
|
||||
suite "Hidden size config":
|
||||
test "hidden size 128 works":
|
||||
putEnv("SACLSTM_HIDDEN_SIZE", "128")
|
||||
let actor = initActorNet(STATE_DIM)
|
||||
let state = randomNormalTensor[float32](STATE_DIM)
|
||||
let ls0 = zeroState(128)
|
||||
let (actions, _, ls1) = actorForward(actor, state, ls0)
|
||||
check actions.shape == [4]
|
||||
check ls1.h.shape == [128]
|
||||
putEnv("SACLSTM_HIDDEN_SIZE", "256")
|
||||
@@ -0,0 +1,156 @@
|
||||
## Tests for replay_buffer.nim — assert-based, no framework.
|
||||
|
||||
import arraymancer
|
||||
import std/sequtils
|
||||
import SAC_LSTM_Bot/replay_buffer
|
||||
import SAC_LSTM_Bot/state # STATE_DIM
|
||||
|
||||
const
|
||||
S_DIM = STATE_DIM # 35
|
||||
A_DIM = 4
|
||||
|
||||
proc makeTrans(reward: float32; done: bool): Transition =
|
||||
Transition(
|
||||
state: zeros[float32](S_DIM),
|
||||
action: zeros[float32](A_DIM),
|
||||
reward: reward,
|
||||
nextState: zeros[float32](S_DIM),
|
||||
done: done
|
||||
)
|
||||
|
||||
# ── 1. len and canSample on empty buffer ──────────────────────────────────────
|
||||
block:
|
||||
var buf = newReplayBuffer(100, S_DIM, A_DIM, burnIn = 8, trainWindow = 16)
|
||||
assert buf.len == 0, "empty len"
|
||||
assert not buf.canSample, "empty canSample"
|
||||
echo "PASS empty buffer"
|
||||
|
||||
# ── 2. len grows, canSample becomes true ──────────────────────────────────────
|
||||
block:
|
||||
var buf = newReplayBuffer(100, S_DIM, A_DIM, burnIn = 2, trainWindow = 3)
|
||||
for i in 0 ..< 4:
|
||||
buf.add(makeTrans(float32(i), false))
|
||||
assert buf.len == 4
|
||||
assert not buf.canSample, "needs 5 to sample"
|
||||
buf.add(makeTrans(99, false))
|
||||
assert buf.len == 5
|
||||
assert buf.canSample, "5 transitions, seqLen=5 → canSample"
|
||||
echo "PASS len + canSample"
|
||||
|
||||
# ── 3. Ring buffer wraps at capacity ─────────────────────────────────────────
|
||||
block:
|
||||
var buf = newReplayBuffer(10, S_DIM, A_DIM, burnIn = 2, trainWindow = 3)
|
||||
for i in 0 ..< 15:
|
||||
buf.add(makeTrans(float32(i), false))
|
||||
assert buf.len == 10, "wraps at capacity, len stays 10"
|
||||
echo "PASS wrap at capacity"
|
||||
|
||||
# ── 4. Sampled sequences have correct length ──────────────────────────────────
|
||||
block:
|
||||
var buf = newReplayBuffer(200, S_DIM, A_DIM, burnIn = 4, trainWindow = 6)
|
||||
for i in 0 ..< 50:
|
||||
buf.add(makeTrans(float32(i), false))
|
||||
let seqs = buf.sampleSequences(8)
|
||||
assert seqs.len > 0, "should have valid starts"
|
||||
for s in seqs:
|
||||
assert s.burnIn.len == 4, "burnIn len"
|
||||
assert s.train.len == 6, "train len"
|
||||
echo "PASS sequence length"
|
||||
|
||||
# ── 5. Burn-in / train split is correct ───────────────────────────────────────
|
||||
block:
|
||||
# Fill with distinct rewards so we can identify positions
|
||||
var buf = newReplayBuffer(200, S_DIM, A_DIM, burnIn = 3, trainWindow = 4)
|
||||
for i in 0 ..< 30:
|
||||
buf.add(makeTrans(float32(i), false))
|
||||
let seqs = buf.sampleSequences(1)
|
||||
assert seqs.len == 1
|
||||
let s = seqs[0]
|
||||
# The 4th reward of the sequence (index 3) must match s.train[0].reward
|
||||
# and s.burnIn[2].reward must be s.burnIn[2].reward (just check no overlap)
|
||||
let allRewards = s.burnIn.mapIt(it.reward) & s.train.mapIt(it.reward)
|
||||
# consecutive integer rewards → each must be strictly increasing by 1
|
||||
var ok = true
|
||||
for i in 1 ..< allRewards.len:
|
||||
if allRewards[i] != allRewards[i-1] + 1.0f32:
|
||||
ok = false
|
||||
break
|
||||
assert ok, "burn-in and train must form a contiguous sequence"
|
||||
echo "PASS burn-in/train split"
|
||||
|
||||
# ── 6. Sequences never cross a done=true boundary ─────────────────────────────
|
||||
block:
|
||||
var buf = newReplayBuffer(200, S_DIM, A_DIM, burnIn = 2, trainWindow = 3)
|
||||
# Episode 1: transitions 0..4 (done at index 4)
|
||||
for i in 0 ..< 4:
|
||||
buf.add(makeTrans(float32(i), false))
|
||||
buf.add(makeTrans(99, true)) # battle end at index 4
|
||||
# Episode 2: transitions 5..14
|
||||
for i in 5 ..< 15:
|
||||
buf.add(makeTrans(float32(i), false))
|
||||
|
||||
let seqs = buf.sampleSequences(20)
|
||||
for s in seqs:
|
||||
# No transition in burnIn (except the last) or train (except the last)
|
||||
# may have done=true, since that would mean the next step crosses a boundary.
|
||||
for i in 0 ..< s.burnIn.len - 1:
|
||||
assert not s.burnIn[i].done, "done in middle of burnIn"
|
||||
for i in 0 ..< s.train.len - 1:
|
||||
assert not s.train[i].done, "done in middle of train"
|
||||
# The join between burnIn and train must not cross a done=true
|
||||
if s.burnIn.len > 0:
|
||||
assert not s.burnIn[^1].done, "done at end of burnIn crosses boundary to train"
|
||||
echo "PASS no cross-boundary sequences"
|
||||
|
||||
# ── 7. canSample false when buffer too small ───────────────────────────────────
|
||||
block:
|
||||
var buf = newReplayBuffer(100, S_DIM, A_DIM, burnIn = 8, trainWindow = 16)
|
||||
for i in 0 ..< 23: # seqLen = 24; 23 < 24
|
||||
buf.add(makeTrans(float32(i), false))
|
||||
assert not buf.canSample, "23 < 24 seqLen"
|
||||
buf.add(makeTrans(23, false))
|
||||
assert buf.canSample, "24 == seqLen"
|
||||
echo "PASS canSample threshold"
|
||||
|
||||
# ── 8. Empty buffer doesn't crash on sample ────────────────────────────────────
|
||||
block:
|
||||
var buf = newReplayBuffer(100, S_DIM, A_DIM, burnIn = 8, trainWindow = 16)
|
||||
let seqs = buf.sampleSequences(4)
|
||||
assert seqs.len == 0, "empty buffer → empty result"
|
||||
echo "PASS empty sample"
|
||||
|
||||
# ── 9. Single episode (no done except at very end) ────────────────────────────
|
||||
block:
|
||||
var buf = newReplayBuffer(200, S_DIM, A_DIM, burnIn = 3, trainWindow = 5)
|
||||
for i in 0 ..< 19:
|
||||
buf.add(makeTrans(float32(i), false))
|
||||
buf.add(makeTrans(19, true)) # last
|
||||
let seqs = buf.sampleSequences(5)
|
||||
assert seqs.len > 0, "should find valid starts"
|
||||
for s in seqs:
|
||||
assert s.burnIn.len == 3
|
||||
assert s.train.len == 5
|
||||
echo "PASS single episode"
|
||||
|
||||
# ── 10. Multiple short episodes: all boundaries respected ─────────────────────
|
||||
block:
|
||||
var buf = newReplayBuffer(300, S_DIM, A_DIM, burnIn = 2, trainWindow = 3)
|
||||
# 10 episodes of 5 transitions each (done at end of each episode)
|
||||
var reward = 0'f32
|
||||
for ep in 0 ..< 10:
|
||||
for i in 0 ..< 4:
|
||||
buf.add(makeTrans(reward, false))
|
||||
reward += 1
|
||||
buf.add(makeTrans(reward, true)) # battle end
|
||||
reward += 1
|
||||
|
||||
let seqs = buf.sampleSequences(30)
|
||||
assert seqs.len > 0
|
||||
for s in seqs:
|
||||
# No done in the middle of any sequence
|
||||
let all = s.burnIn & s.train
|
||||
for i in 0 ..< all.len - 1:
|
||||
assert not all[i].done, "boundary crossed in multi-episode test"
|
||||
echo "PASS multiple episodes"
|
||||
|
||||
echo "ALL TESTS PASSED"
|
||||
@@ -11,17 +11,46 @@ template check(cond: bool, msg: string) =
|
||||
# ── computeReward ─────────────────────────────────────────────────────────────
|
||||
|
||||
block damageInflicted:
|
||||
# p=1: 6*1 - 2 = 4
|
||||
check abs(computeReward(damageInflicted = 1.0) - 4.0) < 1e-9, "p=1 damage = +4"
|
||||
# p=3: 6*3 - 2 = 16
|
||||
check abs(computeReward(damageInflicted = 3.0) - 16.0) < 1e-9, "p=3 damage = +16"
|
||||
# p=1: 1.25 * (6*1 - 2) = 5 (lever 2 aggression mult, #59)
|
||||
check abs(computeReward(damageInflicted = 1.0) - 5.0) < 1e-9, "p=1 damage = +5"
|
||||
# p=3: 1.25 * (6*3 - 2) = 20
|
||||
check abs(computeReward(damageInflicted = 3.0) - 20.0) < 1e-9, "p=3 damage = +20"
|
||||
# low-power spam stays unprofitable: 1.25*(6*0.1-2) < 0
|
||||
check computeReward(damageInflicted = 0.1) < 0.0, "p=0.1 spam still negative"
|
||||
|
||||
block damageReceived:
|
||||
# p_e=1: -(6*1 - 2) = -4
|
||||
# p_e=1: -(6*1 - 2) = -4 (unchanged by lever 2)
|
||||
check abs(computeReward(damageReceived = 1.0) - (-4.0)) < 1e-9, "p_e=1 received = -4"
|
||||
# p_e=3: -(6*3 - 2) = -16
|
||||
check abs(computeReward(damageReceived = 3.0) - (-16.0)) < 1e-9, "p_e=3 received = -16"
|
||||
|
||||
block hitBonus:
|
||||
# flat +0.5 per landed shot: p=1 hit -> 5.0 + 0.5
|
||||
check abs(computeReward(damageInflicted = 1.0, hitCount = 1) - 5.5) < 1e-9,
|
||||
"p=1 hit = +5.5"
|
||||
# two hits in one step: 1.25*(6*2-2) + 2*0.5 = 12.5 + 1.0 = 13.5
|
||||
check abs(computeReward(damageInflicted = 2.0, hitCount = 2) - 13.5) < 1e-9,
|
||||
"two hits = +13.5"
|
||||
|
||||
block ramTaken:
|
||||
# flat per victim collision (#59)
|
||||
check abs(computeReward(ramTakenCount = 1) - (-3.0)) < 1e-9, "ram taken x1 = -3"
|
||||
check abs(computeReward(ramTakenCount = 2) - (-6.0)) < 1e-9, "ram taken x2 = -6"
|
||||
|
||||
block chargeDeterrent:
|
||||
# zero-damage case at half threshold depth: -2 * (1 - 0.06/0.12) = -1
|
||||
let rHalf = computeReward(enemyDistFrac = 0.06)
|
||||
check abs(rHalf - (-1.0)) < 1e-9, "charge at frac 0.06 = -1"
|
||||
# at zero distance: full ceiling
|
||||
check abs(computeReward(enemyDistFrac = 0.0) - (-2.0)) < 1e-9, "charge at frac 0 = -2"
|
||||
# at/beyond threshold and no-contact sentinel: no penalty
|
||||
check abs(computeReward(enemyDistFrac = 0.12)) < 1e-9, "at threshold = 0"
|
||||
check abs(computeReward(enemyDistFrac = 0.5)) < 1e-9, "beyond threshold = 0"
|
||||
check abs(computeReward(enemyDistFrac = 2.0)) < 1e-9, "no-contact sentinel = 0"
|
||||
# suppressed while dealing damage that step (fighting back at close range is fine)
|
||||
let rFight = computeReward(damageInflicted = 1.0, enemyDistFrac = 0.06)
|
||||
check abs(rFight - 5.0) < 1e-9, "dealing damage cancels charge penalty"
|
||||
|
||||
block wallHit:
|
||||
check abs(computeReward(wallHitTicks = 1) - (-5.0)) < 1e-9, "1 wall tick = -5"
|
||||
|
||||
@@ -30,6 +59,7 @@ block wastedShot:
|
||||
check abs(computeReward(wastedShotPower = 2.0) - (-0.2)) < 1e-9, "missed p=2 = -0.2"
|
||||
|
||||
block winLoss:
|
||||
# terminal terms stay dominant over shaping (#59 scale discipline)
|
||||
check abs(computeReward(win = true) - 20.0) < 1e-9, "win = +20"
|
||||
check abs(computeReward(loss = true) - (-10.0)) < 1e-9, "loss = -10"
|
||||
|
||||
|
||||
@@ -0,0 +1,159 @@
|
||||
## test_training.nim — stdlib unittest for SAC training module.
|
||||
import unittest
|
||||
import arraymancer
|
||||
import std/[math, random, os]
|
||||
import SAC_LSTM_Bot/network
|
||||
import SAC_LSTM_Bot/weights
|
||||
import SAC_LSTM_Bot/replay_buffer
|
||||
import SAC_LSTM_Bot/training
|
||||
|
||||
const
|
||||
STATE_DIM = 10
|
||||
ACTION_DIM = 4
|
||||
HIDDEN = 16 # small for speed; override via env not needed in tests
|
||||
|
||||
proc makeTrainer(): SACTrainer =
|
||||
## Small trainer for tests — override hidden size via env before calling.
|
||||
putEnv("SACLSTM_HIDDEN_SIZE", "16")
|
||||
initSACTrainer(STATE_DIM, ACTION_DIM)
|
||||
|
||||
proc makeBuffer(): ReplayBuffer =
|
||||
newReplayBuffer(capacity = 2000, stateDim = STATE_DIM, actionDim = ACTION_DIM,
|
||||
burnIn = 4, trainWindow = 8)
|
||||
|
||||
proc randState(): Tensor[float32] =
|
||||
randomNormalTensor[float32](STATE_DIM)
|
||||
|
||||
proc randAction(): Tensor[float32] =
|
||||
randomTensor[float32](ACTION_DIM, 1.0'f32) *. 2.0'f32 -. 1.0'f32 # uniform [-1,1]
|
||||
|
||||
proc fillBuffer(buf: var ReplayBuffer; n: int; donePeriod = 20) =
|
||||
for i in 0 ..< n:
|
||||
let t = Transition(
|
||||
state: randState(),
|
||||
action: randAction(),
|
||||
reward: rand(-1.0'f32 .. 1.0'f32),
|
||||
nextState: randState(),
|
||||
done: (i mod donePeriod == donePeriod - 1))
|
||||
buf.add(t)
|
||||
|
||||
proc isFinite(x: float32): bool =
|
||||
not (x != x) and x < Inf and x > -Inf # not NaN and not Inf
|
||||
|
||||
suite "SACTrainer — basic update":
|
||||
|
||||
setup:
|
||||
randomize(42)
|
||||
var trainer = makeTrainer()
|
||||
var buf = makeBuffer()
|
||||
fillBuffer(buf, 1000)
|
||||
let seqs = buf.sampleSequences(4)
|
||||
|
||||
test "sampleSequences returns non-empty batch":
|
||||
check seqs.len > 0
|
||||
|
||||
test "sacUpdate returns finite losses":
|
||||
let m = sacUpdate(trainer, seqs)
|
||||
check isFinite(m.criticLoss)
|
||||
check isFinite(m.actorLoss)
|
||||
check isFinite(m.alphaLoss)
|
||||
|
||||
test "alpha stays positive after update":
|
||||
var t2 = trainer
|
||||
discard sacUpdate(t2, seqs)
|
||||
check t2.alpha() > 0.0'f32
|
||||
|
||||
test "critic1 weights change after update":
|
||||
let w_before = trainer.critic1.fc3.w.clone()
|
||||
discard sacUpdate(trainer, seqs)
|
||||
let w_after = trainer.critic1.fc3.w
|
||||
var changed = false
|
||||
for i in 0 ..< w_before.size:
|
||||
if abs(w_before.unsafe_raw_offset[i] - w_after.unsafe_raw_offset[i]) > 1e-10'f32:
|
||||
changed = true
|
||||
break
|
||||
check changed
|
||||
|
||||
test "critic2 weights change after update":
|
||||
let w_before = trainer.critic2.fc3.w.clone()
|
||||
discard sacUpdate(trainer, seqs)
|
||||
let w_after = trainer.critic2.fc3.w
|
||||
var changed = false
|
||||
for i in 0 ..< w_before.size:
|
||||
if abs(w_before.unsafe_raw_offset[i] - w_after.unsafe_raw_offset[i]) > 1e-10'f32:
|
||||
changed = true
|
||||
break
|
||||
check changed
|
||||
|
||||
test "actor weights change after update":
|
||||
let w_before = trainer.actor.fc1.w.clone()
|
||||
discard sacUpdate(trainer, seqs)
|
||||
let w_after = trainer.actor.fc1.w
|
||||
var changed = false
|
||||
for i in 0 ..< w_before.size:
|
||||
if abs(w_before.unsafe_raw_offset[i] - w_after.unsafe_raw_offset[i]) > 1e-10'f32:
|
||||
changed = true
|
||||
break
|
||||
check changed
|
||||
|
||||
test "soft target update: target moves toward critic":
|
||||
## After update, target fc3.w should be closer to critic1.fc3.w than before.
|
||||
let targetBefore = trainer.targetCritic1.fc3.w.clone()
|
||||
let criticW = trainer.critic1.fc3.w.clone()
|
||||
discard sacUpdate(trainer, seqs)
|
||||
let targetAfter = trainer.targetCritic1.fc3.w
|
||||
|
||||
# Distance before: ||targetBefore - criticW||
|
||||
var distBefore = 0.0'f32
|
||||
for i in 0 ..< targetBefore.size:
|
||||
let d = targetBefore.unsafe_raw_offset[i] - criticW.unsafe_raw_offset[i]
|
||||
distBefore += d * d
|
||||
|
||||
# Distance after: ||targetAfter - criticW_new|| (critic changed too, use original for ref)
|
||||
var distAfter = 0.0'f32
|
||||
for i in 0 ..< targetAfter.size:
|
||||
let d = targetAfter.unsafe_raw_offset[i] - criticW.unsafe_raw_offset[i]
|
||||
distAfter += d * d
|
||||
|
||||
# Target moved toward critic (distance decreased).
|
||||
# With tau=0.005, it moves a tiny bit — just check direction.
|
||||
check distAfter <= distBefore + 1e-3'f32 # soft bound; critic also moves
|
||||
|
||||
test "empty sequences is a no-op":
|
||||
let m = sacUpdate(trainer, @[])
|
||||
check m.criticLoss == 0.0'f32
|
||||
check m.actorLoss == 0.0'f32
|
||||
check m.alpha == 0.0'f32
|
||||
|
||||
suite "SACTrainer — done=true terminal transitions":
|
||||
|
||||
test "update with done=true transitions produces finite losses":
|
||||
randomize(7)
|
||||
putEnv("SACLSTM_HIDDEN_SIZE", "16")
|
||||
var trainer = makeTrainer()
|
||||
var buf = makeBuffer()
|
||||
# Fill with short episodes: done every 12 steps (burnIn=4, trainWindow=8 → seqLen=12)
|
||||
fillBuffer(buf, 800, donePeriod = 12)
|
||||
let seqs = buf.sampleSequences(2)
|
||||
if seqs.len > 0:
|
||||
let m = sacUpdate(trainer, seqs)
|
||||
check isFinite(m.criticLoss)
|
||||
check isFinite(m.actorLoss)
|
||||
check isFinite(m.alphaLoss)
|
||||
check m.alpha > 0.0'f32
|
||||
|
||||
suite "SACTrainer — multiple updates":
|
||||
|
||||
test "three sequential updates all stay finite":
|
||||
randomize(99)
|
||||
putEnv("SACLSTM_HIDDEN_SIZE", "16")
|
||||
var trainer = makeTrainer()
|
||||
var buf = makeBuffer()
|
||||
fillBuffer(buf, 1000)
|
||||
for _ in 1..3:
|
||||
let seqs = buf.sampleSequences(4)
|
||||
if seqs.len > 0:
|
||||
let m = sacUpdate(trainer, seqs)
|
||||
check isFinite(m.criticLoss)
|
||||
check isFinite(m.actorLoss)
|
||||
check m.alpha > 0.0'f32
|
||||
@@ -0,0 +1,141 @@
|
||||
import unittest
|
||||
import arraymancer
|
||||
import zip/zipfiles
|
||||
import std/[os, math, strutils]
|
||||
import SAC_LSTM_Bot/network
|
||||
import SAC_LSTM_Bot/weights
|
||||
|
||||
const
|
||||
STATE_DIM = 20
|
||||
ACTION_DIM = 4
|
||||
|
||||
proc tensorsEqual(a, b: Tensor[float32]; tol: float32 = 1e-6'f32): bool =
|
||||
if a.shape != b.shape: return false
|
||||
for i in 0 ..< a.size:
|
||||
if abs(a.unsafe_raw_offset[i] - b.unsafe_raw_offset[i]) > tol: return false
|
||||
true
|
||||
|
||||
suite "saveWeights / loadCheckpoint":
|
||||
|
||||
setup:
|
||||
let actor = initActorNet(STATE_DIM)
|
||||
let c1 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let c2 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let tc1 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let tc2 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let alpha = 0.2'f32
|
||||
let zipPath = getTempDir() / "test_weights_latest.zip"
|
||||
|
||||
teardown:
|
||||
if fileExists(zipPath): removeFile(zipPath)
|
||||
|
||||
test "creates a valid zip file":
|
||||
saveWeights(zipPath, actor, c1, c2, tc1, tc2, alpha)
|
||||
check fileExists(zipPath)
|
||||
var z: ZipArchive
|
||||
check z.open(zipPath, fmRead)
|
||||
var count = 0
|
||||
for f in z.walkFiles: inc count
|
||||
z.close()
|
||||
check count > 0
|
||||
|
||||
test "all entries are .npy files (+ alpha.npy)":
|
||||
saveWeights(zipPath, actor, c1, c2, tc1, tc2, alpha)
|
||||
var z: ZipArchive
|
||||
discard z.open(zipPath, fmRead)
|
||||
var allNpy = true
|
||||
for f in z.walkFiles:
|
||||
if not f.endsWith(".npy"): allNpy = false
|
||||
z.close()
|
||||
check allNpy
|
||||
|
||||
test "round-trip: actor weights preserved":
|
||||
saveWeights(zipPath, actor, c1, c2, tc1, tc2, alpha)
|
||||
let ck = loadCheckpoint(zipPath)
|
||||
check tensorsEqual(actor.fc1.w, ck.actor.fc1.w)
|
||||
check tensorsEqual(actor.fc1.b, ck.actor.fc1.b)
|
||||
check tensorsEqual(actor.lstm.wCombined, ck.actor.lstm.wCombined)
|
||||
check tensorsEqual(actor.lstm.bCombined, ck.actor.lstm.bCombined)
|
||||
check tensorsEqual(actor.fc2.w, ck.actor.fc2.w)
|
||||
check tensorsEqual(actor.muHead.w, ck.actor.muHead.w)
|
||||
check tensorsEqual(actor.logStdHead.w, ck.actor.logStdHead.w)
|
||||
|
||||
test "round-trip: critic weights preserved":
|
||||
saveWeights(zipPath, actor, c1, c2, tc1, tc2, alpha)
|
||||
let ck = loadCheckpoint(zipPath)
|
||||
check tensorsEqual(c1.fc1.w, ck.critic1.fc1.w)
|
||||
check tensorsEqual(c1.lstm.wCombined, ck.critic1.lstm.wCombined)
|
||||
check tensorsEqual(c1.fc3.w, ck.critic1.fc3.w)
|
||||
check tensorsEqual(tc1.fc1.w, ck.targetCritic1.fc1.w)
|
||||
check tensorsEqual(tc2.fc1.w, ck.targetCritic2.fc1.w)
|
||||
|
||||
test "round-trip: alpha preserved":
|
||||
saveWeights(zipPath, actor, c1, c2, tc1, tc2, alpha)
|
||||
let ck = loadCheckpoint(zipPath)
|
||||
check abs(ck.alpha - alpha) < 1e-6'f32
|
||||
|
||||
test "round-trip: hiddenDim reconstructed":
|
||||
saveWeights(zipPath, actor, c1, c2, tc1, tc2, alpha)
|
||||
let ck = loadCheckpoint(zipPath)
|
||||
check ck.actor.hiddenDim == actor.hiddenDim
|
||||
check ck.critic1.hiddenDim == c1.hiddenDim
|
||||
|
||||
test "adam not present → initialized=false":
|
||||
saveWeights(zipPath, actor, c1, c2, tc1, tc2, alpha)
|
||||
let ck = loadCheckpoint(zipPath)
|
||||
check not ck.adam.initialized
|
||||
|
||||
suite "saveCheckpoint (with Adam)":
|
||||
|
||||
setup:
|
||||
let actor = initActorNet(STATE_DIM)
|
||||
let c1 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let c2 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let tc1 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let tc2 = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
let alpha = 0.1'f32
|
||||
var adam = initSACAdamStates(actor, c1, c2)
|
||||
# put some non-zero values in Adam state
|
||||
adam.actor.fc1.w.m[0, 0] = 0.5'f32
|
||||
adam.actor.fc1.w.t = 42
|
||||
adam.critic1.fc1.w.t = 7
|
||||
adam.alpha.t = 99
|
||||
let zipPath = getTempDir() / "test_checkpoint.zip"
|
||||
|
||||
teardown:
|
||||
if fileExists(zipPath): removeFile(zipPath)
|
||||
|
||||
test "round-trip: Adam m tensor":
|
||||
saveCheckpoint(zipPath, actor, c1, c2, tc1, tc2, alpha, adam)
|
||||
let ck = loadCheckpoint(zipPath)
|
||||
check ck.adam.initialized
|
||||
check abs(ck.adam.actor.fc1.w.m[0, 0] - 0.5'f32) < 1e-6'f32
|
||||
|
||||
test "round-trip: Adam t counters":
|
||||
saveCheckpoint(zipPath, actor, c1, c2, tc1, tc2, alpha, adam)
|
||||
let ck = loadCheckpoint(zipPath)
|
||||
check ck.adam.actor.fc1.w.t == 42
|
||||
check ck.adam.critic1.fc1.w.t == 7
|
||||
check ck.adam.alpha.t == 99
|
||||
|
||||
suite "Atomic save":
|
||||
|
||||
test "temp file is cleaned up after successful save":
|
||||
let zipPath = getTempDir() / "test_atomic.zip"
|
||||
let tmpZip = zipPath & ".tmp"
|
||||
let actor = initActorNet(STATE_DIM)
|
||||
let c = initCriticNet(STATE_DIM, ACTION_DIM)
|
||||
saveWeights(zipPath, actor, c, c, c, c, 0.2'f32)
|
||||
check fileExists(zipPath)
|
||||
check not fileExists(tmpZip)
|
||||
removeFile(zipPath)
|
||||
|
||||
suite "Error handling":
|
||||
|
||||
test "loadCheckpoint missing file → IOError":
|
||||
var raised = false
|
||||
try:
|
||||
discard loadCheckpoint(getTempDir() / "nonexistent_xxxxxx.zip")
|
||||
except IOError:
|
||||
raised = true
|
||||
check raised
|
||||
@@ -209,6 +209,8 @@ proc runReceiveLoop*(ws: SyncWebSocket; info: BotInfo; secret: string; serverUrl
|
||||
of "SkippedTurnEvent":
|
||||
let e = node.to(SkippedTurnEvent)
|
||||
gBot.onSkippedTurn(e)
|
||||
of "BotListUpdate":
|
||||
updateBotNames(node)
|
||||
else:
|
||||
discard # unknown message type — ignore
|
||||
except Exception as e:
|
||||
|
||||
@@ -13,7 +13,7 @@
|
||||
## tickChan main → bot (true = new tick ready; false = stop)
|
||||
## intentChan bot → sender (JSON string to send to server; "" = stop)
|
||||
|
||||
import std/[json, locks, math, os, posix, syncio]
|
||||
import std/[json, locks, math, os, posix, syncio, tables]
|
||||
import ./constants
|
||||
import ./schemas
|
||||
import ./color
|
||||
@@ -68,6 +68,7 @@ var gGameSetup {.guard: gLock.}: GameSetup
|
||||
var gTeammateIds{.guard: gLock.}: seq[int]
|
||||
var gVariant {.guard: gLock.}: string
|
||||
var gServerVersion {.guard: gLock.}: string
|
||||
var gBotNames {.guard: gLock.}: Table[int, string]
|
||||
# Events handed main -> bot thread. Channel move, NOT a locked shared seq:
|
||||
# the old locked seq[BotEvent] copied GC'd payloads (strings/teamMessages)
|
||||
# across threads on every tick -> refcount churn under ORC + --threads:on ->
|
||||
@@ -166,6 +167,28 @@ proc getTurnRemaining*(): float = gTurnRemaining
|
||||
proc getGunTurnRemaining*(): float = gGunTurnRemaining
|
||||
proc getRadarTurnRemaining*(): float = gRadarTurnRemaining
|
||||
|
||||
proc getBotName*(id: int): string =
|
||||
## Lookup bot name by id from the last BotListUpdate. Returns "" if unknown.
|
||||
withLock(gLock): result = gBotNames.getOrDefault(id, "")
|
||||
|
||||
proc updateBotNames*(node: JsonNode) =
|
||||
## Update the id → name table from a BotListUpdate message (full replacement).
|
||||
withLock(gLock):
|
||||
gBotNames.clear()
|
||||
if node.isNil: return
|
||||
let botsNode = node{"bots"}
|
||||
if botsNode.isNil or botsNode.kind != JArray: return
|
||||
for b in botsNode:
|
||||
let name = b{"name"}.getStr("")
|
||||
if name.len == 0: continue
|
||||
var id = -1
|
||||
if not b{"id"}.isNil:
|
||||
id = b{"id"}.getInt(-1)
|
||||
if id == -1 and not b{"botId"}.isNil:
|
||||
id = b{"botId"}.getInt(-1)
|
||||
if id == -1: continue
|
||||
gBotNames[id] = name
|
||||
|
||||
proc getMaxSpeed*(): float = gMaxSpeed
|
||||
proc getMaxTurnRate*(): float = gMaxTurnRate
|
||||
proc getMaxGunTurnRate*(): float = gMaxGunTurnRate
|
||||
@@ -862,6 +885,8 @@ proc initGlobals*() =
|
||||
gEventChan.open(8)
|
||||
initLock(gLock)
|
||||
gEventQueue = initEventQueue()
|
||||
withLock(gLock):
|
||||
gBotNames = initTable[int, string]()
|
||||
# Debug log is opt-in: it is written every tick from two threads, so leaving
|
||||
# it on by default is a disk hog and an I/O stall source. Enable with
|
||||
# PPOB_DEBUG_LOG=1 to debug; starts fresh (truncated) each run.
|
||||
|
||||
@@ -18,9 +18,10 @@ import java.util.logging.Logger;
|
||||
* to the same log file via PPOB_LOG_FILE — the shell wrapper stitches them.
|
||||
*
|
||||
* Usage (env vars):
|
||||
* PPO_BOT_DIR — path to PPO_Bot dir
|
||||
* PPO_BOT_DIR — path to the bot dir (any Tank Royale bot)
|
||||
* SAMPLE_BOTS_DIR — path to sample bots archive
|
||||
* PPOB_LOG_FILE — path to training_log.jsonl (appended)
|
||||
* BOT_NAME — bot name to match in round results (default: PPO_Bot)
|
||||
* TRAINING_OPPONENT — opponent bot name (default: Target)
|
||||
* TRAINING_ROUNDS — number of rounds to run (CLI arg or env var)
|
||||
*
|
||||
@@ -39,8 +40,9 @@ public class RunTraining {
|
||||
: System.getenv().getOrDefault("TRAINING_OPPONENT", "Target");
|
||||
int totalRounds = args.length > 1 ? Integer.parseInt(args[1])
|
||||
: Integer.parseInt(System.getenv().getOrDefault("TRAINING_ROUNDS", "100"));
|
||||
String botName = System.getenv().getOrDefault("BOT_NAME", "PPO_Bot");
|
||||
|
||||
System.out.printf("Training: PPO_Bot vs %s for %d rounds%n", opponent, totalRounds);
|
||||
System.out.printf("Training: %s vs %s for %d rounds%n", botName, opponent, totalRounds);
|
||||
System.out.printf("Log: %s%n", logFile);
|
||||
|
||||
// Dead-bot guard: the runner keeps listing a crashed PPO_Bot in the
|
||||
@@ -81,7 +83,7 @@ public class RunTraining {
|
||||
boolean win = false;
|
||||
boolean found = false;
|
||||
for (var r : event.getResults()) {
|
||||
if (r.getName().equals("PPO_Bot")) {
|
||||
if (r.getName().equals(botName)) {
|
||||
found = true;
|
||||
totalScore = r.getTotalScore();
|
||||
win = r.getRank() == 1;
|
||||
@@ -99,9 +101,9 @@ public class RunTraining {
|
||||
frozenRounds[0]++;
|
||||
long frozenMs = System.currentTimeMillis() - lastAdvanceMs[0];
|
||||
if (frozenRounds[0] >= 10 && frozenMs >= 10_000) {
|
||||
System.err.printf("PPO_Bot round_counter frozen at %d for %d "
|
||||
System.err.printf("%s round_counter frozen at %d for %d "
|
||||
+ "harness rounds / %.0fs — process dead, aborting for restart%n",
|
||||
ctr, frozenRounds[0], frozenMs / 1000.0);
|
||||
botName, ctr, frozenRounds[0], frozenMs / 1000.0);
|
||||
System.exit(1);
|
||||
}
|
||||
} else {
|
||||
@@ -110,7 +112,7 @@ public class RunTraining {
|
||||
lastAdvanceMs[0] = System.currentTimeMillis();
|
||||
}
|
||||
if (!found) {
|
||||
System.err.println("PPO_Bot missing from round " + round
|
||||
System.err.println(botName + " missing from round " + round
|
||||
+ " results — process died, aborting battle for restart");
|
||||
System.exit(1);
|
||||
}
|
||||
@@ -147,10 +149,10 @@ public class RunTraining {
|
||||
endCounter = readCounter(counterPath);
|
||||
}
|
||||
if (endCounter < expectedEnd) {
|
||||
System.err.printf("PPO_Bot round_counter %d < expected %d (start+%d) at battle "
|
||||
System.err.printf("%s round_counter %d < expected %d (start+%d) at battle "
|
||||
+ "end (waited 60s) — %d rounds never trained (corpse?), aborting for "
|
||||
+ "restart%n",
|
||||
endCounter, expectedEnd, totalRounds, expectedEnd - endCounter);
|
||||
botName, endCounter, expectedEnd, totalRounds, expectedEnd - endCounter);
|
||||
System.exit(1);
|
||||
}
|
||||
System.out.printf("Counter check passed: %d == expected %d%n", endCounter, expectedEnd);
|
||||
|
||||
Reference in New Issue
Block a user