Files
SirRoboGarage/SAC_LSTM_Bot/tests/test_training.nim
T
SirStone 415d4e3738 feat(SAC_LSTM_Bot): SAC training module (#47)
Implements sacUpdate with burn-in LSTM warm-up, twin-critic TD update,
actor reparameterization gradient, auto-alpha, and soft target update.
Manual backprop (linear + LSTM single-step, truncated BPTT). 11 new tests
all green; full regression suite (57+ tests) unaffected.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-08-21 00:11:35 +02:00

160 lines
5.2 KiB
Nim

## 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