Kernel-Prompted Random Transformer / stage2_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json, sys, time
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6
  7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  8from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report
  9
 10SEED = 2052
 11N_SUPPORT = 64
 12DH = 64
 13EPOCHS = 20
 14# Union of step sizes is shared by baseline and idea configs.
 15LR_GRID = [1e-3, 3e-3, 1e-2]
 16SIGMA_GRID = [0.4, 0.6, 0.8]
 17records = {'baseline': {}, 'idea': {}}
 18
 19class FrozenAttentionPrompt(nn.Module):
 20    """One frozen single-head attention layer, with trainable prompt only."""
 21    def __init__(self, d_in=32, dh=DH, n_prompt=N_SUPPORT, seed=0):
 22        super().__init__()
 23        r = np.random.RandomState(seed)
 24        q = r.normal(0, 1/np.sqrt(dh), (dh, dh)).astype('float32')
 25        k = r.normal(0, 1/np.sqrt(dh), (dh, dh)).astype('float32')
 26        v = r.normal(0, 1/np.sqrt(dh), (dh, 1)).astype('float32')
 27        self.register_buffer('Q', torch.from_numpy(q))
 28        self.register_buffer('K', torch.from_numpy(k))
 29        self.register_buffer('V', torch.from_numpy(v))
 30        self.prompt = nn.Parameter(torch.zeros(n_prompt, dh))
 31        nn.init.normal_(self.prompt, std=0.05)
 32        self.d_in, self.dh = d_in, dh
 33
 34    def embed(self, x):
 35        z = torch.zeros(x.shape[0], self.dh, device=x.device, dtype=x.dtype)
 36        z[:, :self.d_in] = x
 37        z[:, self.d_in] = 1.
 38        return z
 39
 40    def forward(self, x):
 41        e = self.embed(x)
 42        q = e @ self.Q
 43        keys = self.prompt @ self.K
 44        logits = q @ keys.T / np.sqrt(self.dh)
 45        w = torch.softmax(logits, dim=1)
 46        vals = self.prompt @ self.V
 47        return w @ vals
 48
 49def set_seed(seed):
 50    np.random.seed(seed); torch.manual_seed(seed)
 51    if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
 52
 53def dataset(seed):
 54    d = get_dataset('sequence', seed, n_train=400, n_test=400)
 55    # Bench train_model expects torch tensors; get_dataset already supplies them.
 56    return d
 57
 58def baseline_fn(cfg):
 59    def run(seed):
 60        set_seed(seed)
 61        ds = dataset(seed)
 62        model = FrozenAttentionPrompt(seed=seed)
 63        model, metric, _ = train_model(model, ds, epochs=EPOCHS, lr=cfg['lr'],
 64                                       batch=128, weight_decay=0.0, log=lambda *_: None)
 65        if model is None: return float('nan')
 66        with torch.no_grad():
 67            device = next(model.parameters()).device
 68            x = ds['xte'].to(device); pred = model(x)
 69            # observed trained-model attention entropy and prompt norm
 70            e = model.embed(x[:32]); logits = (e @ model.Q) @ (model.prompt @ model.K).T / np.sqrt(DH)
 71            records['baseline'][int(seed)] = {'prompt_norm': float(model.prompt.norm()),
 72                'attention_entropy': float((-(torch.softmax(logits,1)*torch.log_softmax(logits,1)).sum(1)).mean())}
 73        return float(metric)
 74    return run
 75
 76def analytic_prompt(xs, ys, sigma, seed):
 77    r = np.random.RandomState(seed)
 78    q = r.normal(0, 1/np.sqrt(DH), (DH, DH))
 79    k = r.normal(0, 1/np.sqrt(DH), (DH, DH))
 80    v = r.normal(0, 1/np.sqrt(DH), (DH, 1))
 81    A = (q.T @ k)[:33] / np.sqrt(DH)
 82    B = np.vstack([A, v.T])
 83    c = np.column_stack([xs / sigma**2, -np.sum(xs*xs, axis=1)/(2*sigma**2), ys])
 84    P = np.linalg.lstsq(B, c.T, rcond=None)[0].T
 85    return P, q, k, v
 86
 87def idea_fn(cfg):
 88    def run(seed):
 89        ds = dataset(seed)
 90        xs = ds['xtr'][:N_SUPPORT].numpy().astype('float64')
 91        ys = ds['ytr'][:N_SUPPORT].numpy().astype('float64').reshape(-1,1)
 92        P,q,k,v = analytic_prompt(xs, ys, cfg['sigma'], seed)
 93        xt = ds['xte'].numpy().astype('float64')
 94        e = np.zeros((len(xt), DH)); e[:,:32] = xt; e[:,32] = 1
 95        logits = e @ ((q.T @ k) @ P.T) / np.sqrt(DH)
 96        logits -= logits.max(1, keepdims=True)
 97        w = np.exp(logits); w /= w.sum(1, keepdims=True)
 98        pred = w @ (P @ v)
 99        metric = float(np.mean((pred[:,0] - ds['yte'].numpy())**2))
100        target = xt @ xs.T / cfg['sigma']**2 - (xs*xs).sum(1)[None,:]/(2*cfg['sigma']**2)
101        # common query term is intentionally omitted, as it cancels in softmax.
102        residual = float(np.sqrt(np.mean((e @ ((q.T@k)@P.T)/np.sqrt(DH)-target)**2)))
103        records['idea'].setdefault(cfg['sigma'], {})[int(seed)] = {
104            'prompt_norm': float(np.linalg.norm(P,axis=1).mean()),
105            'logit_residual': residual,
106            'entropy': float(np.mean(-(w*np.log(np.maximum(w,1e-30))).sum(1))) }
107        return metric
108    return run
109
110def main():
111    # Baseline sweep covers all rates used by idea-side comparison.
112    base_grid = [{'lr': x} for x in LR_GRID]
113    base = sweep_baseline(baseline_fn, base_grid)
114    # Idea sweep has three nearby bandwidth settings; each is evaluated on all 8 seeds.
115    idea_runs = []
116    for sigma, lr in zip(SIGMA_GRID, LR_GRID):
117        r = evaluate(idea_fn({'sigma': sigma, 'lr': lr}))
118        idea_runs.append({'cfg': {'sigma': sigma, 'lr': lr}, 'result': r})
119    best = min(idea_runs, key=lambda z: z['result']['mean'])
120    extra = {'track_choice': 'sequence: multi-token temporal-window correlations require sequence attention',
121             'best_sigma': best['cfg']['sigma'], 'idea_sweep': idea_runs,
122             'trained_baseline_behavior': records['baseline']}
123    sigs = [records['idea'][s] for s in SIGMA_GRID if s in records['idea']]
124    if sigs:
125        means = [float(np.mean([x['prompt_norm'] for x in z.values()])) for z in sigs]
126        ref = means[1]
127        best_resid = float(np.mean([x['logit_residual'] for x in records['idea'][best['cfg']['sigma']].values()]))
128        extra['mechanism_prediction'] = {'prediction': 'prompt norm scales as sigma^-2; affine logit residual is near zero',
129             'observed_prompt_norms': dict(zip(SIGMA_GRID, means)),
130             'predicted_norm_ratio_sigma_0.4_to_0.8': 4.0,
131             'observed_norm_ratio_sigma_0.4_to_0.8': means[0]/means[-1],
132             'observed_best_logit_residual': best_resid,
133             'confirmed': bool(abs(means[0]/means[-1]-4.0) < 0.5 and best_resid < 1e-4)}
134    report = make_report('sequence', 'frozen_random_attention_prompt', base, best['result'], extra)
135    report['baseline_sweep_shared_lr_grid'] = LR_GRID
136    report['cut'] = 'Only sequence track tested; no CIFAR/MNIST transfer and no latency/FLOP study.'
137    Path('bench_report.json').write_text(json.dumps(report, indent=2))
138    print(json.dumps(report, indent=2))
139
140if __name__ == '__main__': main()