File size: 4,119 Bytes
de10ad1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
"""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"]))