train-kit / test / checkpoint.test.mjs
 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
// Checkpoints: atomic writes, structural validation, interval reloads.
import test from "node:test";
import assert from "node:assert/strict";
import { mkdtempSync, writeFileSync } from "node:fs";
import { tmpdir } from "node:os";
import { join } from "node:path";
import { checkpointPath, candidatePath, loadModel, saveModel, CheckpointStore } from "../dist/checkpoint.js";
import { isMlpModel, isTabularModel } from "../dist/model.js";
import { modelRev } from "../dist/rev.js";
import { mlpInit } from "../dist/forward.js";

const mlp = (game) => ({ v: 1, game, episodes: 3, updatedAt: 1, net: mlpInit([2, 2, 1], () => 0.5) });
const tab = (game) => ({ v: 1, game, episodes: 5, updatedAt: 2, q: { s0: [0, 1] } });

test("save/load roundtrip for both model kinds; guards narrow by game", () => {
  const dir = mkdtempSync(join(tmpdir(), "tk-ck-"));
  saveModel(checkpointPath(dir, "a"), mlp("a"));
  saveModel(checkpointPath(dir, "b"), tab("b"));
  const a = loadModel(checkpointPath(dir, "a"));
  const b = loadModel(checkpointPath(dir, "b"));
  assert.ok(isMlpModel(a, "a"));
  assert.ok(!isMlpModel(a, "z"), "expected-game narrows");
  assert.ok(isTabularModel(b, "b"));
  // the written JSON is the model plus the additive rev the writer computed
  assert.deepEqual(a, { ...mlp("a"), rev: modelRev(mlp("a")) });
  assert.deepEqual(b, { ...tab("b"), rev: modelRev(tab("b")) });
});

test("corrupt, missing, or alien JSON loads as null", () => {
  const dir = mkdtempSync(join(tmpdir(), "tk-ck-"));
  assert.equal(loadModel(join(dir, "nope.json")), null);
  writeFileSync(join(dir, "bad.json"), "{not json");
  assert.equal(loadModel(join(dir, "bad.json")), null);
  writeFileSync(join(dir, "alien.json"), JSON.stringify({ v: 2, game: "x", net: { sizes: [], layers: [] } }));
  assert.equal(loadModel(join(dir, "alien.json")), null, "v must be 1");
  writeFileSync(join(dir, "shape.json"), JSON.stringify({ v: 1, game: "x", episodes: 0, updatedAt: 0 }));
  assert.equal(loadModel(join(dir, "shape.json")), null, "needs net or q");
});

test("candidatePath sits beside the live checkpoint", () => {
  assert.equal(candidatePath("/m", "toy"), "/m/toy.candidate.json");
  assert.equal(checkpointPath("/m", "toy"), "/m/toy.json");
});

test("CheckpointStore loads now and reloads on rewrite", async () => {
  const dir = mkdtempSync(join(tmpdir(), "tk-ck-"));
  const file = checkpointPath(dir, "toy");
  saveModel(file, mlp("toy"));
  const seen = [];
  const store = new CheckpointStore(file, 30, (m) => seen.push(m.episodes));
  assert.equal(store.model?.episodes, 3);
  const next = { ...mlp("toy"), episodes: 9 };
  await new Promise((r) => setTimeout(r, 20)); // let mtime tick past the first stat
  saveModel(file, next);
  await new Promise((r) => setTimeout(r, 120));
  store.close();
  assert.equal(store.model?.episodes, 9, "rewrite picked up on the interval");
});

static mirror of HEAD · about · clone: git clone https://git.ardegazu.ro/train-kit.git