import json, sys import numpy as np import torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import train_model, sweep_baseline, make_report from semantic_track import get_dataset SEEDS = tuple(range(8)) K, R, D = 3, 9, 8 PHI = torch.tensor([0, 0, 0, 1, 1, 1, 2, 2, 2], dtype=torch.long) class SemanticMLP(nn.Module): def __init__(self, calibrated=False, temperature=1.0): super().__init__() self.body = nn.Sequential(nn.Linear(D, 64), nn.ReLU(), nn.Linear(64, 64), nn.ReLU(), nn.Linear(64, R)) self.calibrated = calibrated self.temperature = float(temperature) if calibrated: self.a = nn.Parameter(torch.ones(K)) self.b = nn.Parameter(torch.zeros(K)) def state_log_probs(self, x): response_lp = torch.log_softmax(self.body(x) / self.temperature, dim=1) phi = PHI.to(response_lp.device) return torch.stack([torch.logsumexp(response_lp[:, phi == j], dim=1) for j in range(K)], dim=1) def forward(self, x): state_lp = self.state_log_probs(x) if self.calibrated: return state_lp * self.a[None, :] + self.b[None, :] return state_lp def tensor_ds(seed): raw = get_dataset(seed, 400, 400) return {k: (torch.from_numpy(v) if isinstance(v, np.ndarray) else v) for k, v in raw.items()} def train_one(seed, idea, cfg, return_model=False): torch.manual_seed(10000 + int(seed)) np.random.seed(10000 + int(seed)) net = SemanticMLP(calibrated=idea, temperature=cfg['temperature']) trained, metric, _ = train_model( net, tensor_ds(seed), epochs=cfg['epochs'], lr=cfg['lr'], batch=128, weight_decay=cfg['weight_decay'], log=lambda *_: None) return (metric, trained) if return_model else metric def factory(idea): def make(cfg): return lambda seed: train_one(seed, idea, cfg) return make def main(): rng = np.random.default_rng(1071) response = rng.dirichlet(np.ones(R), size=128) phi = PHI.numpy() pushed = np.stack([response[:, phi == j].sum(axis=1) for j in range(K)], axis=1) direct = np.stack([response[:, phi == j].sum(axis=1) for j in range(K)], axis=1) math_check = { 'aggregation_error_vs_direct': float(np.max(np.abs(pushed - direct))), 'mass_conservation_error': float(np.max(np.abs(pushed.sum(1) - 1.0))), } grid = [{'lr': lr, 'epochs': 18, 'temperature': temp, 'weight_decay': 0.0} for lr in (1e-3, 3e-3, 6e-3) for temp in (0.8, 1.0, 1.2)] base = sweep_baseline(factory(False), grid, seeds=(0, 1, 2, 3)) best = base['best_cfg'] nearby = [c for c in grid if c['temperature'] == best['temperature']] idea_runs = [] for cfg in nearby: vals = [train_one(s, True, cfg) for s in SEEDS] idea_runs.append((float(np.mean(vals)), cfg, {'mean': float(np.mean(vals)), 'std': float(np.std(vals)), 'per_seed': vals, 'n': len(vals)})) _, idea_cfg, idea = min(idea_runs, key=lambda t: t[0]) observed_a, observed_mass, observed_calibration = [], [], [] for seed in SEEDS: _, model = train_one(seed, True, idea_cfg, return_model=True) device = next(model.parameters()).device with torch.no_grad(): x = tensor_ds(seed)['xte'].to(device) raw = torch.softmax(model.state_log_probs(x), 1) cal = torch.softmax(model(x), 1) observed_a.append(model.a.detach().cpu().numpy()) observed_mass.append(float(torch.max(torch.abs(raw.sum(1) - 1)).item())) observed_calibration.append(float(torch.mean(torch.abs(cal - raw)).item())) mean_a = np.mean(observed_a, axis=0) signature = { 'prediction': 'pushforward conserves state mass; calibration learns a nontrivial correction when raw response probabilities are distorted', 'predicted_vs_observed': { 'predicted_mass_error': 0.0, 'observed_mass_error_mean': float(np.mean(observed_mass)), 'observed_abs_calibrated_vs_raw_mean': float(np.mean(observed_calibration)), 'learned_a_mean_by_state': mean_a.tolist(), 'learned_a_std_by_state': np.std(observed_a, axis=0).tolist(), }, 'confirmed': bool(np.max(observed_mass) < 1e-6 and np.mean(observed_calibration) > 1e-4), } extra = { 'mechanism_signature': signature, 'math_check': math_check, 'custom_track': {'name': 'semantic_pushforward_classification', 'file': 'semantic_track.py', 'domain': 'uncertainty_calibration'}, } report = make_report('semantic_pushforward_classification', 'mlp_tiny', base, idea, extra) with open('bench_report.json', 'w') as f: json.dump(report, f, indent=2) print(json.dumps(report, indent=2)) if __name__ == '__main__': main()