File size: 8,318 Bytes
1da7ac7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
# tests/test_diagnostics.py
"""Plan 6, tâche 2 — diagnostics.py : l'observabilité à ciel ouvert.

expert_routing_snapshot : structure valide sur un tiny CteCore (une clé
`layer_{i}` par couche, 2 experts retenus par layer, poids renormalisés
sommant à 1, ids dans [0, n_experts), triés par gate décroissante) ;
EXACTITUDE contre la formule von Mises réécrite à la main sur les phases
captées par un hook frais (κ_eff = 4.0/2.5 = 1.6 — le PIÈGE température du
doc §MoE batché — top-2 renormalisé) ; reproductibilité malgré la pollution
des états porteurs (reset au départ) ; hooks retirés après usage.

head_readout : floats Python bornés [0, 1) sur CTE (les deux têtes), None
sur FractusCore (pas de têtes confiance/salience — duck-typing honnête).
"""
import pytest

torch = pytest.importorskip(
    "torch", reason="diagnostics nécessite torch (substrat sans torch)"
)

from fractus_vorax.model.cte_core import CteCore, CteCoreConfig
from fractus_vorax.model.diagnostics import expert_routing_snapshot, head_readout


def _tiny(**overrides) -> CteCoreConfig:
    return CteCoreConfig(vocab_size=128, **overrides)


def _tiny_fractus():
    from fractus_vorax.model.fractus_core import FractusCore, FractusCoreConfig

    cfg = FractusCoreConfig(
        vocab_size=128, d_model=64, n_layers=2, n_experts=4, top_k=2,
        rank=16, max_seq_len=64,
    )
    return FractusCore(cfg)


# ---------------------------------------------------------------------------
# expert_routing_snapshot — structure
# ---------------------------------------------------------------------------

def test_snapshot_structure_tiny():
    torch.manual_seed(0)
    cfg = _tiny()
    core = CteCore(cfg)
    snap = expert_routing_snapshot(core, [10, 20, 30])
    assert set(snap.keys()) == {f"layer_{i}" for i in range(cfg.n_layers)}
    for entries in snap.values():
        assert len(entries) == cfg.top_k  # 2 experts retenus par layer
        ids = [e for e, _ in entries]
        weights = [w for _, w in entries]
        assert all(isinstance(e, int) and 0 <= e < cfg.n_experts for e in ids)
        assert all(isinstance(w, float) and w > 0.0 for w in weights)
        assert abs(sum(weights) - 1.0) < 1e-5  # top-k RENORMALISÉ
        # triées par gate décroissante (convention torch.topk)
        assert weights == sorted(weights, reverse=True)


def test_snapshot_layer_count_follows_config():
    """Le snapshot suit le nombre de couches du noyau (tiny 2, ici 3 ; le 1B
    réel en aurait 16 — une clé layer_{i} par bloc)."""
    torch.manual_seed(0)
    cfg = _tiny(n_layers=3)
    core = CteCore(cfg)
    snap = expert_routing_snapshot(core, [1, 2, 3, 4, 5])
    assert len(snap) == 3
    assert set(snap) == {"layer_0", "layer_1", "layer_2"}


# ---------------------------------------------------------------------------
# expert_routing_snapshot — exactitude vs la formule von Mises à la main
# ---------------------------------------------------------------------------

def test_snapshot_matches_von_mises_formula():
    """EXACTITUDE : ids EXACTS et poids ≈ identiques à la formule von Mises
    réécrite à la main dans CE test (doc §MoE batché) sur les phases captées
    au hook — κ_eff = 4.0/2.5 = 1.6 hardcodé ici : le test attraperait un
    noyau qui aurait oublié le piège température."""
    torch.manual_seed(0)
    cfg = _tiny()
    core = CteCore(cfg)
    ids = torch.tensor([[10, 20, 30]])
    snap = expert_routing_snapshot(core, ids)

    captured = []  # (moe, phases réellement consommées par la passe MoE)

    def hook(module, args, output):
        captured.append((module, args[1].detach().clone()))

    handles = [blk.moe.register_forward_hook(hook) for blk in core.blocks]
    try:
        core.reset_states(1)
        core(ids)  # même forward, mêmes états initiaux => mêmes phases
    finally:
        for h in handles:
            h.remove()

    assert len(captured) == cfg.n_layers
    kappa_eff = 4.0 / 2.5  # doc §MoE batché : κ=4.0, temperature=2.5
    for i, (moe, phases) in enumerate(captured):
        theta_bar = torch.atan2(
            torch.sin(phases).sum(dim=-1), torch.cos(phases).sum(dim=-1)
        )[0, -1]
        gates = torch.exp(kappa_eff * torch.cos(theta_bar - moe.expert_phases))
        gates = gates / gates.sum()
        vals, idx = gates.topk(cfg.top_k)
        expected = [
            (int(e), float(w))
            for e, w in zip(idx.tolist(), (vals / vals.sum()).tolist())
        ]
        got = snap[f"layer_{i}"]
        assert [e for e, _ in got] == [e for e, _ in expected]
        assert all(abs(a - b) < 1e-5 for (_, a), (_, b) in zip(got, expected))


# ---------------------------------------------------------------------------
# expert_routing_snapshot — robustesse, hygiène, duck-typing
# ---------------------------------------------------------------------------

def test_snapshot_deterministic_across_state_pollution():
    """Le reset au départ rend la lecture reproductible même après que le
    noyau ait avalé d'autres chunks (états porteurs mutés par le forward)."""
    torch.manual_seed(0)
    core = CteCore(_tiny())
    ids = [7, 8, 9]
    first = expert_routing_snapshot(core, ids)
    core(torch.tensor([[40, 50, 60, 61]]))  # pollue attn_S/attn_z/thought_state
    second = expert_routing_snapshot(core, ids)
    assert first == second


def test_snapshot_list_and_tensor_equivalent():
    """Ergonomie : list[int] et tensor (accepté (1, L) et (L,)) disent la
    même chose sur le même noyau."""
    torch.manual_seed(0)
    core = CteCore(_tiny())
    a = expert_routing_snapshot(core, [10, 20, 30])
    b = expert_routing_snapshot(core, torch.tensor([10, 20, 30]))
    c = expert_routing_snapshot(core, torch.tensor([[10, 20, 30]]))
    assert a == b == c


def test_snapshot_rejects_empty_ids():
    torch.manual_seed(0)
    core = CteCore(_tiny())
    with pytest.raises(ValueError):
        expert_routing_snapshot(core, [])


def test_hooks_removed_and_core_still_usable():
    """Hygiène : les hooks partent après l'appel, le noyau fonctionne
    exactement comme avant (aucune trace du passage du diagnosticien)."""
    torch.manual_seed(0)
    core = CteCore(_tiny())
    expert_routing_snapshot(core, [1, 2])
    head_readout(core, [1, 2])
    for blk in core.blocks:
        assert not blk.moe._forward_hooks
    assert not core.output_head._forward_hooks
    logits = core(torch.tensor([[3, 4]]))
    assert tuple(logits.shape) == (1, 2, 128)


def test_snapshot_duck_types_fractus_core():
    """Le routage est lisible sur les DEUX noyaux (mêmes conventions moe :
    forward(h, phases) + _compute_gates) — FractusCore κ=4.0 direct, phases
    par position (on lit la dernière)."""
    torch.manual_seed(0)
    core = _tiny_fractus()
    snap = expert_routing_snapshot(core, [10, 20, 30])
    assert set(snap) == {"layer_0", "layer_1"}
    for entries in snap.values():
        assert len(entries) == 2
        weights = [w for _, w in entries]
        assert abs(sum(weights) - 1.0) < 1e-5


# ---------------------------------------------------------------------------
# head_readout — têtes confiance/salience du CTE
# ---------------------------------------------------------------------------

def test_head_readout_cte_floats_bounded():
    """CTE : les deux têtes se lisent — floats Python dans [0, 1) (sigmoid),
    reproductibles malgré la pollution des états porteurs."""
    torch.manual_seed(0)
    core = CteCore(_tiny())
    out = head_readout(core, [10, 20, 30])
    assert set(out) == {"confidence", "salience"}
    assert isinstance(out["confidence"], float)
    assert isinstance(out["salience"], float)
    assert 0.0 <= out["confidence"] < 1.0  # borné (le plan : « head_readout borné »)
    assert 0.0 <= out["salience"] < 1.0
    core(torch.tensor([[90, 91]]))  # pollue les états porteurs
    again = head_readout(core, [10, 20, 30])
    assert out == again


def test_head_readout_fractus_none():
    """FractusCore n'a PAS de têtes confiance/salience (lm_head seule) :
    lecture honnête → les deux à None, sans lancer le moindre forward."""
    torch.manual_seed(0)
    core = _tiny_fractus()
    assert head_readout(core, [1, 2, 3]) == {"confidence": None, "salience": None}