# Moteur vectorise du toy a decoder appris (E 1D, Dec lineaire 1->19, CE)
C = 19
PYI = PA + PB # classes entieres 0..18 (indices one-hot)
YOH = np.zeros((N_PAIRS, C))
YOH[np.arange(N_PAIRS), PYI] = 1.0
def batch_cls(lrd_a, wd_a, seed_a, lr_rep=1e-3, steps=12000, frac=0.82,
thr=0.90, gap=1000, log=20):
"""N runs (lr_dec, wd, seed) en un seul appel numpy vectorise.
Retourne phases Table 1 + accuracies finales + temps de croisement."""
R = len(lrd_a)
E = np.stack([np.random.default_rng(int(s) * 31 + k).normal(0, .5, P)
for k, s in enumerate(seed_a)])
W = np.zeros((R, C)); b = np.zeros((R, C))
mE = np.zeros_like(E); vE = np.zeros_like(E)
mW = np.zeros_like(W); vW = np.zeros_like(W)
mb = np.zeros_like(b); vb = np.zeros_like(b)
lrd = np.asarray(lrd_a, dtype=float)[:, None]
wd = np.asarray(wd_a, dtype=float)[:, None]
lrp = np.full((R, 1), lr_rep)
trm = np.zeros((R, N_PAIRS), dtype=bool)
for r, s in enumerate(seed_a):
perm = np.random.default_rng(int(s) * 7 + 3).permutation(N_PAIRS)
trm[r, perm[:int(round(frac * N_PAIRS))]] = True
ntr = trm.sum(1)[:, None, None]
t_tr = np.full(R, np.nan); t_va = np.full(R, np.nan)
a_tr = np.zeros(R); a_va = np.zeros(R)
for step in range(1, steps + 1):
S = (E[:, PA] + E[:, PB])[:, :, None]
logits = S * W[:, None, :] + b[:, None, :]
Z = logits - logits.max(-1, keepdims=True)
ex = np.exp(Z)
p = ex / ex.sum(-1, keepdims=True)
G = (p - YOH[None]) / ntr * trm[:, :, None]
gW = (G * S).sum(1)
gb = G.sum(1)
gS = (G * W[:, None, :]).sum(-1)
gE = np.zeros_like(E)
np.add.at(gE.T, PA, gS.T)
np.add.at(gE.T, PB, gS.T)
bc1, bc2 = 1 - .9 ** step, 1 - .999 ** step
mE = .9 * mE + .1 * gE; vE = .999 * vE + .001 * gE ** 2
E -= lrp * (mE / bc1) / (np.sqrt(vE / bc2) + 1e-8)
mW = .9 * mW + .1 * gW; vW = .999 * vW + .001 * gW ** 2
W -= lrd * (mW / bc1) / (np.sqrt(vW / bc2) + 1e-8) + lrd * wd * W
mb = .9 * mb + .1 * gb; vb = .999 * vb + .001 * gb ** 2
b -= lrd * (mb / bc1) / (np.sqrt(vb / bc2) + 1e-8) + lrd * wd * b
if step % log == 0 or step == steps:
ok = logits.argmax(-1) == PYI[None, :]
a_tr = (ok * trm).sum(1) / trm.sum(1)
a_va = (ok * ~trm).sum(1) / np.maximum((~trm).sum(1), 1)
h = (a_tr >= thr) & np.isnan(t_tr); t_tr[h] = step
h = (a_va >= thr) & np.isnan(t_va); t_va[h] = step
ph = np.where(np.isnan(t_tr) | (a_tr < thr), "CONFU",
np.where(np.isnan(t_va) | (a_va < thr), "MEMOR",
np.where(t_va - t_tr >= gap, "GROKK", "COMPR")))
return ph, a_tr, a_va, t_tr, t_va
# Carte : lr_dec x fraction (wd = 0), 4 graines par cellule
LRD_GRID = [1e-3, 3e-3, 1e-2, 3e-2, 1e-1]
FR_GRID = [0.35, 0.50, 0.60, 0.70, 0.82]
t0 = time.time()
phase_map = {}
va_map = {}
for f in FR_GRID:
la, sa = [], []
for l in LRD_GRID:
for s in SEEDS:
la.append(l); sa.append(s)
ph, _, a_va, _, _ = batch_cls(np.array(la), np.zeros(len(la)), np.array(sa), frac=f)
phase_map[f] = [max(set(ph[j * 4:(j + 1) * 4]), key=ph[j * 4:(j + 1) * 4].tolist().count)
for j in range(len(LRD_GRID))]
va_map[f] = [a_va[j * 4:(j + 1) * 4].mean() for j in range(len(LRD_GRID))]
print(f"carte {len(FR_GRID) * len(LRD_GRID) * 4} runs en {time.time() - t0:.0f}s")
LABELS = {"COMPR": 0, "GROKK": 1, "MEMOR": 2, "CONFU": 3}
COLORS = ["#2e7d32", "#f9a825", "#c62828", "#424242"]
NAMES_FR = {"COMPR": "comprehension", "GROKK": "grokking", "MEMOR": "memorization",
"CONFU": "confusion"}
fig, axes = plt.subplots(1, 2, figsize=(13, 4.2))
grid = np.array([[LABELS[phase_map[f][j]] for j in range(len(LRD_GRID))]
for f in FR_GRID])
axes[0].imshow(grid, cmap=plt.matplotlib.colors.ListedColormap(COLORS),
aspect="auto", vmin=-0.5, vmax=3.5)
axes[0].set_xticks(range(len(LRD_GRID)), [f"{x:.0e}" for x in LRD_GRID])
axes[0].set_yticks(range(len(FR_GRID)), [f"{f:.2f}" for f in FR_GRID])
axes[0].set_xlabel("taux d'apprentissage du decoder")
axes[0].set_ylabel("fraction de donnees")
axes[0].set_title("Phases (Table 1, vote majoritaire 4 graines)")
for i in range(len(FR_GRID)):
for j in range(len(LRD_GRID)):
axes[0].text(j, i, phase_map[FR_GRID[i]][j][:2], ha="center",
va="center", color="white", fontsize=9, fontweight="bold")
handles = [plt.Rectangle((0, 0), 1, 1, color=COLORS[i]) for i in range(4)]
axes[0].legend(handles, [NAMES_FR[k] for k in LABELS], fontsize=8, loc="lower left")
vam = np.array([[va_map[f][j] for j in range(len(LRD_GRID))] for f in FR_GRID])
im = axes[1].imshow(vam, cmap="viridis", vmin=0, vmax=1)
axes[1].set_xticks(range(len(LRD_GRID)), [f"{x:.0e}" for x in LRD_GRID])
axes[1].set_yticks(range(len(FR_GRID)), [f"{f:.2f}" for f in FR_GRID])
axes[1].set_xlabel("taux d'apprentissage du decoder")
axes[1].set_title("accuracy de validation finale (moyenne 4 graines)")
plt.colorbar(im, ax=axes[1])
plt.tight_layout(); plt.show()
print("Phase dominante par cellule (lrd x frac) :")
print("frac\\lrd " + " ".join(f"{x:>7.0e}" for x in LRD_GRID))
for f in FR_GRID:
print(f"{f:>7.2f} " + " ".join(f"{x:>7s}" for x in phase_map[f]))
print()
print("accuracy de validation moyenne :")
print("frac\\lrd " + " ".join(f"{x:>7.0e}" for x in LRD_GRID))
for f in FR_GRID:
print(f"{f:>7.2f} " + " ".join(f"{v:>7.2f}" for v in va_map[f]))