thefinalboss commited on
Commit
de10ad1
·
verified ·
1 Parent(s): 8d9996a

opt: cumsum/chunked attention kernels, memory-flat CE, block checkpointing, v2 trainer (proven equivalent, 46 tests)

Browse files
Files changed (1) hide show
  1. tests/test_train_loop_smoke.py +112 -0
tests/test_train_loop_smoke.py ADDED
@@ -0,0 +1,112 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """CPU smoke test of the v2 training loop pieces, end to end.
2
+
3
+ Exercises exactly what fast4gpu_boost_v2.py does per step — synthetic int32
4
+ memmap shard -> per-chunk fetch -> tick_chunk_train_ce -> backward -> SGD ->
5
+ scheduled sampling via sample_tokens_chunked — under BOTH attention kernels,
6
+ asserting they produce equivalent losses. No GPU required.
7
+
8
+ Open-heart value: if this passes, swapping the live pod's code cannot change
9
+ the training signal beyond documented float32 rounding.
10
+ """
11
+ import sys
12
+ import tempfile
13
+ from pathlib import Path
14
+
15
+ import numpy as np
16
+ import pytest
17
+ import torch
18
+ import torch.nn.functional as F
19
+
20
+ sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
21
+
22
+ from fractus.continuous_engine import ContinuousThoughtEngine
23
+ from fractus.nn import attention as attn_mod
24
+ from fractus.nn.ce import sample_tokens_chunked
25
+
26
+ VOCAB = 1000
27
+ CFG = dict(vocab_size=VOCAB, d_model=64, n_heads=1, d_head=64,
28
+ n_levels=2, n_oscillators=8, coupling_rank=4,
29
+ n_experts=8, top_k=2, expert_d_ff=64, siren_rank=16,
30
+ n_layers=2)
31
+
32
+
33
+ def _make_shard(tmp: Path, n_tokens: int = 4096, seed: int = 0) -> Path:
34
+ rng = np.random.default_rng(seed)
35
+ p = tmp / "shard_gpu0.npy"
36
+ mm = np.lib.format.open_memmap(p, mode="w+", dtype=np.int32,
37
+ shape=(n_tokens,))
38
+ mm[:] = rng.integers(0, VOCAB, size=n_tokens)
39
+ mm.flush()
40
+ return p
41
+
42
+
43
+ def _fetch(mm, start, count):
44
+ view = np.asarray(mm[start : start + count])
45
+ return torch.from_numpy(view).to(torch.int64)
46
+
47
+
48
+ def _run_loop(shard_path, impl, n_steps=6):
49
+ prev_impl = attn_mod._ACTIVE_IMPL
50
+ attn_mod.set_attention_impl(impl)
51
+ try:
52
+ torch.manual_seed(123)
53
+ eng = ContinuousThoughtEngine(**CFG)
54
+ eng.reset_thought(batch_size=2)
55
+ opt = torch.optim.SGD(eng.parameters(), lr=7e-4, momentum=0.9)
56
+
57
+ mm = np.load(str(shard_path), mmap_mode="r")
58
+ B, SEQ = 2, 16
59
+ step = B * SEQ
60
+ losses = []
61
+ for i, start in enumerate(range(0, n_steps * step, step)):
62
+ block = _fetch(mm, start, step + 1)
63
+ chunk = block[:step].view(B, SEQ).long()
64
+ target = block[1:].view(B, SEQ)
65
+
66
+ ce_tf, lb, h = eng.tick_chunk_train_ce(chunk, target,
67
+ ce_chunk=128,
68
+ return_hidden=True)
69
+ loss = ce_tf + 0.02 * lb
70
+ opt.zero_grad(set_to_none=True)
71
+ loss.backward()
72
+ opt.step()
73
+
74
+ if i % 2 == 0:
75
+ # scheduled-sampling branch: exercise sampling mechanics only
76
+ with torch.no_grad():
77
+ samp = sample_tokens_chunked(
78
+ h.reshape(-1, h.shape[-1]).detach(),
79
+ eng.output_head.weight, temperature=0.9, ce_chunk=128)
80
+ assert samp.shape == (B * SEQ,)
81
+ assert samp.min() >= 0 and samp.max() < VOCAB
82
+
83
+ assert torch.isfinite(loss), f"non-finite loss at step {i}"
84
+ losses.append(float(ce_tf.detach()))
85
+ return losses
86
+ finally:
87
+ attn_mod.set_attention_impl(prev_impl)
88
+
89
+
90
+ def test_v2_loop_runs_and_kernels_agree(tmp_path):
91
+ shard = _make_shard(tmp_path)
92
+ l_cs = _run_loop(shard, "cumsum")
93
+ l_ch = _run_loop(shard, "chunked")
94
+ # same data, same init (same seed): losses must match within fp32 rounding
95
+ for a, b in zip(l_cs, l_ch):
96
+ rel = abs(a - b) / max(abs(a), 1e-9)
97
+ assert rel < 5e-4, f"losses diverge: cumsum={a:.6f} chunked={b:.6f} ({rel:.2e})"
98
+
99
+
100
+ def test_fetch_matches_legacy_whole_shard_load(tmp_path):
101
+ """The v2 per-chunk int32 fetch must give the exact tokens v1 loaded."""
102
+ shard = _make_shard(tmp_path, n_tokens=512, seed=42)
103
+ legacy = torch.from_numpy(np.load(str(shard), mmap_mode="r")).to(torch.int64)
104
+ mm = np.load(str(shard), mmap_mode="r")
105
+ for start in (0, 37, 256):
106
+ got = _fetch(mm, start, 33)
107
+ want = legacy[start : start + 33]
108
+ assert torch.equal(got, want)
109
+
110
+
111
+ if __name__ == "__main__":
112
+ sys.exit(pytest.main([__file__, "-v"]))