fractus-cte / tests /test_train_loop_smoke.py
thefinalboss's picture
opt: cumsum/chunked attention kernels, memory-flat CE, block checkpointing, v2 trainer (proven equivalent, 46 tests)
de10ad1 verified
Raw History Blame
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"]))