# Cell C. Add your own anti-collapse mechanism as a term in the dynamics,
# write down its dose on (sigma, gamma), and check the prediction.
import numpy as np

N, p, p_tgt, alpha, beta = 10, 0.7, 1.0, 1.0, 1.0
s = 1.0 / np.arange(1, N + 1) ** 1.2
s = s * N / s.sum()
C = np.diag(s)                                     # unit-diagonal in expectation, eigenbasis = identity
S11 = p * p * C + p * (1 - p) * np.eye(N)          # context auto-covariance under masking
SX1 = p * C                                        # cross-covariance
Wp = alpha * np.eye(N)

eta = 0.15   # <- the dose

def extra(W):
    """Your mechanism, as an extra term in dW/dt. Default: encoder weight decay."""
    return -eta * W

def dose(sigma, gamma):
    """What the mechanism does to one direction. Weight decay adds a floor to the decay."""
    return sigma + eta, gamma

# Prediction from the dose
sigma = alpha**2 * np.diag(S11)
gamma = alpha * p_tgt * np.diag(SX1)
sig2, gam2 = dose(sigma, gamma)
print("escape without the mechanism:", int((gamma / sigma > 1).sum()), "of", N)
print("escape with it, predicted:   ", int((gam2 / sig2 > 1).sum()), "of", N)

# Measurement: integrate the expected flow (Lemma 3.1) with the extra term
rng = np.random.default_rng(0)
W = 1e-2 * rng.normal(size=(N, N)); Wa = W.copy(); dt = 0.01
for i in range(800):
    if i == 200:
        G1 = np.linalg.norm(W, axis=0)
    dW = -Wp.T @ Wp @ W @ S11 + p_tgt * Wp.T @ Wa @ SX1.T + extra(W)
    W, Wa = W + dt * dW, Wa + dt * beta * (W - Wa)
print("escape with it, measured:    ", int((np.linalg.norm(W, axis=0) > G1).sum()), "of", N)
