opt: cumsum/chunked attention kernels, memory-flat CE, block checkpointing, v2 trainer (proven equivalent, 46 tests)
de10ad1 verified Download tests/test_train_loop_smoke.py from thefinalboss/fractus-cte: direct link, hf CLI and curl.
- Browser
- Download file 4.12 kB
-
https://huggingface.co/thefinalboss/fractus-cte/resolve/0906b039a4098c101e45c75f4b65bf8793051c89/tests/test_train_loop_smoke.py
- Command line
-
hf download hf://thefinalboss/fractus-cte@0906b039a4098c101e45c75f4b65bf8793051c89/tests/test_train_loop_smoke.py
-
curl -L -o test_train_loop_smoke.py https://huggingface.co/thefinalboss/fractus-cte/resolve/0906b039a4098c101e45c75f4b65bf8793051c89/tests/test_train_loop_smoke.py
4.12 kB
| """CPU smoke test of the v2 training loop pieces, end to end. | |
| Exercises exactly what fast4gpu_boost_v2.py does per step — synthetic int32 | |
| memmap shard -> per-chunk fetch -> tick_chunk_train_ce -> backward -> SGD -> | |
| scheduled sampling via sample_tokens_chunked — under BOTH attention kernels, | |
| asserting they produce equivalent losses. No GPU required. | |
| Open-heart value: if this passes, swapping the live pod's code cannot change | |
| the training signal beyond documented float32 rounding. | |
| """ | |
| import sys | |
| import tempfile | |
| from pathlib import Path | |
| import numpy as np | |
| import pytest | |
| import torch | |
| import torch.nn.functional as F | |
| sys.path.insert(0, str(Path(__file__).resolve().parents[1])) | |
| from fractus.continuous_engine import ContinuousThoughtEngine | |
| from fractus.nn import attention as attn_mod | |
| from fractus.nn.ce import sample_tokens_chunked | |
| VOCAB = 1000 | |
| CFG = dict(vocab_size=VOCAB, d_model=64, n_heads=1, d_head=64, | |
| n_levels=2, n_oscillators=8, coupling_rank=4, | |
| n_experts=8, top_k=2, expert_d_ff=64, siren_rank=16, | |
| n_layers=2) | |
| def _make_shard(tmp: Path, n_tokens: int = 4096, seed: int = 0) -> Path: | |
| rng = np.random.default_rng(seed) | |
| p = tmp / "shard_gpu0.npy" | |
| mm = np.lib.format.open_memmap(p, mode="w+", dtype=np.int32, | |
| shape=(n_tokens,)) | |
| mm[:] = rng.integers(0, VOCAB, size=n_tokens) | |
| mm.flush() | |
| return p | |
| def _fetch(mm, start, count): | |
| view = np.asarray(mm[start : start + count]) | |
| return torch.from_numpy(view).to(torch.int64) | |
| def _run_loop(shard_path, impl, n_steps=6): | |
| prev_impl = attn_mod._ACTIVE_IMPL | |
| attn_mod.set_attention_impl(impl) | |
| try: | |
| torch.manual_seed(123) | |
| eng = ContinuousThoughtEngine(**CFG) | |
| eng.reset_thought(batch_size=2) | |
| opt = torch.optim.SGD(eng.parameters(), lr=7e-4, momentum=0.9) | |
| mm = np.load(str(shard_path), mmap_mode="r") | |
| B, SEQ = 2, 16 | |
| step = B * SEQ | |
| losses = [] | |
| for i, start in enumerate(range(0, n_steps * step, step)): | |
| block = _fetch(mm, start, step + 1) | |
| chunk = block[:step].view(B, SEQ).long() | |
| target = block[1:].view(B, SEQ) | |
| ce_tf, lb, h = eng.tick_chunk_train_ce(chunk, target, | |
| ce_chunk=128, | |
| return_hidden=True) | |
| loss = ce_tf + 0.02 * lb | |
| opt.zero_grad(set_to_none=True) | |
| loss.backward() | |
| opt.step() | |
| if i % 2 == 0: | |
| # scheduled-sampling branch: exercise sampling mechanics only | |
| with torch.no_grad(): | |
| samp = sample_tokens_chunked( | |
| h.reshape(-1, h.shape[-1]).detach(), | |
| eng.output_head.weight, temperature=0.9, ce_chunk=128) | |
| assert samp.shape == (B * SEQ,) | |
| assert samp.min() >= 0 and samp.max() < VOCAB | |
| assert torch.isfinite(loss), f"non-finite loss at step {i}" | |
| losses.append(float(ce_tf.detach())) | |
| return losses | |
| finally: | |
| attn_mod.set_attention_impl(prev_impl) | |
| def test_v2_loop_runs_and_kernels_agree(tmp_path): | |
| shard = _make_shard(tmp_path) | |
| l_cs = _run_loop(shard, "cumsum") | |
| l_ch = _run_loop(shard, "chunked") | |
| # same data, same init (same seed): losses must match within fp32 rounding | |
| for a, b in zip(l_cs, l_ch): | |
| rel = abs(a - b) / max(abs(a), 1e-9) | |
| assert rel < 5e-4, f"losses diverge: cumsum={a:.6f} chunked={b:.6f} ({rel:.2e})" | |
| def test_fetch_matches_legacy_whole_shard_load(tmp_path): | |
| """The v2 per-chunk int32 fetch must give the exact tokens v1 loaded.""" | |
| shard = _make_shard(tmp_path, n_tokens=512, seed=42) | |
| legacy = torch.from_numpy(np.load(str(shard), mmap_mode="r")).to(torch.int64) | |
| mm = np.load(str(shard), mmap_mode="r") | |
| for start in (0, 37, 256): | |
| got = _fetch(mm, start, 33) | |
| want = legacy[start : start + 33] | |
| assert torch.equal(got, want) | |
| if __name__ == "__main__": | |
| sys.exit(pytest.main([__file__, "-v"])) | |