# Cell D. ResidualPred's attention at initialisation, in numpy.
import numpy as np

def attention(Q, K, V, beta_init):
    """softmax(Q K^T / sqrt(d) + beta_init * I) V, the diagonal-bias attention of ResidualPred."""
    n, d = Q.shape
    logits = Q @ K.T / np.sqrt(d) + beta_init * np.eye(n)
    A = np.exp(logits - logits.max(axis=1, keepdims=True))
    A /= A.sum(axis=1, keepdims=True)
    return A @ V, A

rng = np.random.default_rng(0)
n, d = 54, 32          # 54 context tokens, as on CIFAR with keep probability 0.85
Q, K = 0.4 * rng.normal(size=(2, n, d))   # a standard initialisation: logits with std about 0.16
V = rng.normal(size=(n, d))
print("beta_init   min A_ii   output = input?")
for b in (0, 2, 5, np.log(n) + 5, 10):
    out, A = attention(Q, K, V, b)
    print(f"{b:8.2f}   {A.diagonal().min():.4f}     max |out - V| = {np.abs(out - V).max():.3f}")

# In the released code the bias sits on the diagonal of the context block only (the target tokens
# appended after the context attend freely), and it is a learnable parameter, one per head.
n_ctx, n_tgt = 54, 10
Qa, Ka = 0.4 * rng.normal(size=(2, n_ctx + n_tgt, d))
bias = np.zeros((n_ctx + n_tgt, n_ctx + n_tgt)); bias[:n_ctx, :n_ctx] = np.eye(n_ctx)
logits = Qa @ Ka.T / np.sqrt(d) + 10.0 * bias
A = np.exp(logits - logits.max(axis=1, keepdims=True)); A /= A.sum(axis=1, keepdims=True)
print(f"context rows: min A_ii = {A[:n_ctx].diagonal().min():.4f};  target rows: max weight = {A[n_ctx:].max():.3f}")
