feat(PPO_Bot): weight persistence + background training (#17)
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,101 @@
|
||||
## test_weights.nim — assert-based tests for weights.nim
|
||||
## Run: nim c tests/test_weights.nim && ./tests/test_weights
|
||||
|
||||
import std/[os, math]
|
||||
import arraymancer
|
||||
import "../network"
|
||||
import "../weights"
|
||||
|
||||
template check(cond: bool, msg: string) =
|
||||
if not cond:
|
||||
quit("FAIL: " & msg, 1)
|
||||
|
||||
const tmpBase = "/tmp/test_weights_nim"
|
||||
|
||||
# ── saveWeights / loadWeights roundtrip ───────────────────────────────────────
|
||||
|
||||
block testRoundtrip:
|
||||
let dir = tmpBase & "_roundtrip"
|
||||
removeDir(dir)
|
||||
|
||||
let ac1 = initActorCritic()
|
||||
saveWeights(ac1, dir)
|
||||
|
||||
var ac2 = initActorCritic()
|
||||
loadWeights(ac2, dir)
|
||||
|
||||
# Verify a sample of tensors
|
||||
template tensorEq(a, b: Tensor[float32]) =
|
||||
check a.shape == b.shape, "shape mismatch"
|
||||
let diff = abs(a - b)
|
||||
var maxDiff = 0.0'f32
|
||||
for v in diff: maxDiff = max(maxDiff, v)
|
||||
check maxDiff < 1e-6'f32, "tensor values differ by " & $maxDiff
|
||||
|
||||
tensorEq(ac1.actor.w1, ac2.actor.w1)
|
||||
tensorEq(ac1.actor.b1, ac2.actor.b1)
|
||||
tensorEq(ac1.actor.w3, ac2.actor.w3)
|
||||
tensorEq(ac1.critic.w1, ac2.critic.w1)
|
||||
tensorEq(ac1.critic.b3, ac2.critic.b3)
|
||||
tensorEq(ac1.logStd, ac2.logStd)
|
||||
|
||||
removeDir(dir)
|
||||
|
||||
# ── saveWeightsAtomic ─────────────────────────────────────────────────────────
|
||||
|
||||
block testAtomic:
|
||||
let dir = tmpBase & "_atomic"
|
||||
removeDir(dir)
|
||||
|
||||
let ac = initActorCritic()
|
||||
saveWeightsAtomic(ac, dir)
|
||||
|
||||
check dirExists(dir), "targetDir should exist after atomic save"
|
||||
for f in ["actor_w1.npy", "critic_w1.npy", "log_std.npy"]:
|
||||
check fileExists(dir / f), "missing file: " & f
|
||||
|
||||
removeDir(dir)
|
||||
|
||||
# ── loadBestAvailable ─────────────────────────────────────────────────────────
|
||||
|
||||
block testLoadBest:
|
||||
let root = tmpBase & "_loadbest"
|
||||
removeDir(root)
|
||||
createDir(root)
|
||||
|
||||
let ac0 = initActorCritic()
|
||||
|
||||
# Try with no weights — should return false
|
||||
var acEmpty = initActorCritic()
|
||||
check not loadBestAvailable(acEmpty, root), "should return false with no weights"
|
||||
|
||||
# Save to latest/; should load
|
||||
saveWeights(ac0, root / "latest")
|
||||
var ac1 = initActorCritic()
|
||||
check loadBestAvailable(ac1, root), "should load from latest/"
|
||||
|
||||
# Remove latest/, save to checkpoint_1/ — should fall back
|
||||
removeDir(root / "latest")
|
||||
saveWeights(ac0, root / "checkpoint_1")
|
||||
var ac2 = initActorCritic()
|
||||
check loadBestAvailable(ac2, root), "should load from checkpoint_1/"
|
||||
|
||||
removeDir(root)
|
||||
|
||||
# ── cleanStaleTempDirs ────────────────────────────────────────────────────────
|
||||
|
||||
block testClean:
|
||||
let root = tmpBase & "_clean"
|
||||
removeDir(root)
|
||||
createDir(root)
|
||||
|
||||
let stale = root / "latest_tmp_12345"
|
||||
createDir(stale)
|
||||
check dirExists(stale), "stale dir should exist before clean"
|
||||
|
||||
cleanStaleTempDirs(root)
|
||||
check not dirExists(stale), "stale dir should be gone after clean"
|
||||
|
||||
removeDir(root)
|
||||
|
||||
echo "All tests passed"
|
||||
Reference in New Issue
Block a user