opt: cumsum/chunked attention kernels, memory-flat CE, block checkpointing, v2 trainer (proven equivalent, 46 tests)
Browse files- 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"]))
|