# Cell A. The full linear JEPA, trained by SGD on sampled data with sampled masks.
# Does the count of directions that grow match the count with mu > 1?
import numpy as np

rng = np.random.default_rng(0)
N = 10            # input features = embedding size
p = 0.7           # context keep probability (feature masking)
p_tgt = 1.0       # target keep probability
alpha = 1.0       # predictor scale            <- try 0.5, 1.2, 2.0
beta = 1.0        # EMA speed                  <- changes the clock only
steepness = 1.2   # how uneven the data spectrum is

# Data: unit-diagonal covariance C with a power-law spectrum, randomly rotated
s = 1.0 / np.arange(1, N + 1) ** steepness
R = np.linalg.qr(rng.normal(size=(N, N)))[0]
C = R @ np.diag(s) @ R.T
C = C / np.sqrt(np.outer(np.diag(C), np.diag(C)))
s, V = np.linalg.eigh(C); s, V = s[::-1], V[:, ::-1]
L = np.linalg.cholesky(C)

# Theory: mu_k = p_tgt s_k / (alpha (p s_k + 1 - p)), and the growth rate of each direction
mu = p_tgt * s / (alpha * (p * s + 1 - p))
sigma = alpha**2 * (p * p * s + p * (1 - p))
rate = (-(beta + sigma) + np.sqrt((beta + sigma) ** 2 + 4 * beta * sigma * (mu - 1))) / 2

# Training: encoder W by SGD on the JEPA loss, target W_a by EMA, predictor W_p = alpha I held fixed
Wp = alpha * np.eye(N)
W = 1e-2 * rng.normal(size=(N, N)); Wa = W.copy()
dt, t1, t2, batch = 0.01, 2.0, 8.0, 512
sizes, times = [], []
for i in range(int(t2 / dt) + 1):
    x = rng.normal(size=(batch, N)) @ L.T
    x1 = x * (rng.random((batch, N)) < p)              # context view: masked features
    m2 = rng.random((batch, N)) < p_tgt                # target mask on the embedding
    err = (x1 @ W.T) @ Wp.T - m2 * (x @ Wa.T)          # prediction minus stop-gradient target
    W -= dt * Wp.T @ err.T @ x1 / batch                # gradient step on the context encoder
    Wa += dt * beta * (W - Wa)                         # EMA target
    if i % 10 == 0:
        times.append(i * dt); sizes.append(np.linalg.norm(W @ V, axis=0))
    if abs(i * dt - t1) < dt / 2:
        G1 = np.linalg.norm(W @ V, axis=0)
measured = np.log(np.linalg.norm(W @ V, axis=0) / G1) / (t2 - t1)

print(" k     s_k     mu_k   predicted rate   measured rate")
for k in range(N):
    print(f"{k+1:2d}  {s[k]:6.3f}  {mu[k]:6.2f}      {rate[k]:+.3f}         {measured[k]:+.3f}")
print(f"\npredicted to grow (mu > 1): {int((mu > 1).sum())}   grew in training: {int((measured > 0).sum())}")

sizes = np.array(sizes)
plot(times, {f"direction {k+1}": sizes[:, k] for k in range(N)}, logy=True,
     xlabel="training time", ylabel="size of each direction",
     kinds={f"direction {k+1}": "drive" if mu[k] > 1 else "decay" for k in range(N)})
