Semantic Pushforward Uncertainty Head / run_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json, sys
  2import numpy as np
  3import torch
  4import torch.nn as nn
  5
  6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  7from bench import train_model, sweep_baseline, make_report
  8from semantic_track import get_dataset
  9
 10SEEDS = tuple(range(8))
 11K, R, D = 3, 9, 8
 12PHI = torch.tensor([0, 0, 0, 1, 1, 1, 2, 2, 2], dtype=torch.long)
 13
 14
 15class SemanticMLP(nn.Module):
 16    def __init__(self, calibrated=False, temperature=1.0):
 17        super().__init__()
 18        self.body = nn.Sequential(nn.Linear(D, 64), nn.ReLU(),
 19                                  nn.Linear(64, 64), nn.ReLU(), nn.Linear(64, R))
 20        self.calibrated = calibrated
 21        self.temperature = float(temperature)
 22        if calibrated:
 23            self.a = nn.Parameter(torch.ones(K))
 24            self.b = nn.Parameter(torch.zeros(K))
 25
 26    def state_log_probs(self, x):
 27        response_lp = torch.log_softmax(self.body(x) / self.temperature, dim=1)
 28        phi = PHI.to(response_lp.device)
 29        return torch.stack([torch.logsumexp(response_lp[:, phi == j], dim=1)
 30                            for j in range(K)], dim=1)
 31
 32    def forward(self, x):
 33        state_lp = self.state_log_probs(x)
 34        if self.calibrated:
 35            return state_lp * self.a[None, :] + self.b[None, :]
 36        return state_lp
 37
 38
 39def tensor_ds(seed):
 40    raw = get_dataset(seed, 400, 400)
 41    return {k: (torch.from_numpy(v) if isinstance(v, np.ndarray) else v)
 42            for k, v in raw.items()}
 43
 44
 45def train_one(seed, idea, cfg, return_model=False):
 46    torch.manual_seed(10000 + int(seed))
 47    np.random.seed(10000 + int(seed))
 48    net = SemanticMLP(calibrated=idea, temperature=cfg['temperature'])
 49    trained, metric, _ = train_model(
 50        net, tensor_ds(seed), epochs=cfg['epochs'], lr=cfg['lr'], batch=128,
 51        weight_decay=cfg['weight_decay'], log=lambda *_: None)
 52    return (metric, trained) if return_model else metric
 53
 54
 55def factory(idea):
 56    def make(cfg):
 57        return lambda seed: train_one(seed, idea, cfg)
 58    return make
 59
 60
 61def main():
 62    rng = np.random.default_rng(1071)
 63    response = rng.dirichlet(np.ones(R), size=128)
 64    phi = PHI.numpy()
 65    pushed = np.stack([response[:, phi == j].sum(axis=1) for j in range(K)], axis=1)
 66    direct = np.stack([response[:, phi == j].sum(axis=1) for j in range(K)], axis=1)
 67    math_check = {
 68        'aggregation_error_vs_direct': float(np.max(np.abs(pushed - direct))),
 69        'mass_conservation_error': float(np.max(np.abs(pushed.sum(1) - 1.0))),
 70    }
 71
 72    grid = [{'lr': lr, 'epochs': 18, 'temperature': temp, 'weight_decay': 0.0}
 73            for lr in (1e-3, 3e-3, 6e-3) for temp in (0.8, 1.0, 1.2)]
 74    base = sweep_baseline(factory(False), grid, seeds=(0, 1, 2, 3))
 75    best = base['best_cfg']
 76    nearby = [c for c in grid if c['temperature'] == best['temperature']]
 77    idea_runs = []
 78    for cfg in nearby:
 79        vals = [train_one(s, True, cfg) for s in SEEDS]
 80        idea_runs.append((float(np.mean(vals)), cfg, {'mean': float(np.mean(vals)),
 81                         'std': float(np.std(vals)), 'per_seed': vals, 'n': len(vals)}))
 82    _, idea_cfg, idea = min(idea_runs, key=lambda t: t[0])
 83
 84    observed_a, observed_mass, observed_calibration = [], [], []
 85    for seed in SEEDS:
 86        _, model = train_one(seed, True, idea_cfg, return_model=True)
 87        device = next(model.parameters()).device
 88        with torch.no_grad():
 89            x = tensor_ds(seed)['xte'].to(device)
 90            raw = torch.softmax(model.state_log_probs(x), 1)
 91            cal = torch.softmax(model(x), 1)
 92        observed_a.append(model.a.detach().cpu().numpy())
 93        observed_mass.append(float(torch.max(torch.abs(raw.sum(1) - 1)).item()))
 94        observed_calibration.append(float(torch.mean(torch.abs(cal - raw)).item()))
 95    mean_a = np.mean(observed_a, axis=0)
 96    signature = {
 97        'prediction': 'pushforward conserves state mass; calibration learns a nontrivial correction when raw response probabilities are distorted',
 98        'predicted_vs_observed': {
 99            'predicted_mass_error': 0.0,
100            'observed_mass_error_mean': float(np.mean(observed_mass)),
101            'observed_abs_calibrated_vs_raw_mean': float(np.mean(observed_calibration)),
102            'learned_a_mean_by_state': mean_a.tolist(),
103            'learned_a_std_by_state': np.std(observed_a, axis=0).tolist(),
104        },
105        'confirmed': bool(np.max(observed_mass) < 1e-6 and np.mean(observed_calibration) > 1e-4),
106    }
107    extra = {
108        'mechanism_signature': signature,
109        'math_check': math_check,
110        'custom_track': {'name': 'semantic_pushforward_classification',
111                         'file': 'semantic_track.py', 'domain': 'uncertainty_calibration'},
112    }
113    report = make_report('semantic_pushforward_classification', 'mlp_tiny', base, idea, extra)
114    with open('bench_report.json', 'w') as f:
115        json.dump(report, f, indent=2)
116    print(json.dumps(report, indent=2))
117
118
119if __name__ == '__main__':
120    main()