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>
This commit is contained in:
@@ -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
|
||||
Reference in New Issue
Block a user