415d4e3738
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>
160 lines
5.2 KiB
Nim
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
|