# Cell B. Sweep the predictor scale yourself: theory curve vs SGD runs.
import numpy as np

N, p_tgt, beta, steepness = 10, 1.0, 1.0, 1.2
rng = np.random.default_rng(0)
s0 = 1.0 / np.arange(1, N + 1) ** steepness
R = np.linalg.qr(rng.normal(size=(N, N)))[0]
C = R @ np.diag(s0) @ R.T
C = C / np.sqrt(np.outer(np.diag(C), np.diag(C)))   # unit-diagonal data covariance
s, V = np.linalg.eigh(C)                             # its spectrum is what the theory needs
L = np.linalg.cholesky(C)

def rank_budget(alpha, p):
    mu = p_tgt * s / (alpha * (p * s + 1 - p))
    return int((mu > 1).sum())

def sgd_count(alpha, p, seed=0, batch=256, dt=0.02, t1=2.0, t2=6.0):
    """Train the linear JEPA by SGD and count the directions that grow."""
    rng = np.random.default_rng(seed + 1)
    W = 1e-2 * rng.normal(size=(N, N)); Wa = W.copy(); Wp = alpha * np.eye(N)
    for i in range(int(t2 / dt) + 1):
        x = rng.normal(size=(batch, N)) @ L.T
        x1 = x * (rng.random((batch, N)) < p)
        m2 = rng.random((batch, N)) < p_tgt
        err = (x1 @ W.T) @ Wp.T - m2 * (x @ Wa.T)
        W -= dt * Wp.T @ err.T @ x1 / batch
        Wa += dt * beta * (W - Wa)
        if abs(i * dt - t1) < dt / 2:
            G1 = np.linalg.norm(W @ V, axis=0)
    return int((np.linalg.norm(W @ V, axis=0) > G1).sum())

alphas = np.logspace(-0.6, 0.4, 21)
series = {f"theory, p = {p}": [rank_budget(a, p) for a in alphas] for p in (0.3, 0.7, 0.95)}
plot(np.log10(alphas), series, xlabel="log10 predictor scale", ylabel="rank budget")

print("p = 0.7      alpha   theory   SGD")
for a in alphas[::4]:
    print(f"          {a:6.2f}   {rank_budget(a, 0.7):4d}   {sgd_count(a, 0.7):4d}")
