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