import sys, json, math, random from itertools import permutations import numpy as np import torch from torch import nn import torch.nn.functional as F sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import sweep_baseline, evaluate, make_report from latent_mixture_transport_local import get_dataset TRACK = 'latent_mixture_transport' MODEL = 'mlp_tiny' SEEDS = tuple(range(8)) SWEEP_SEEDS = (0, 1, 2, 3) K, ZDIM = 4, 2 def seed_all(s): random.seed(s); np.random.seed(s); torch.manual_seed(s) if torch.cuda.is_available(): torch.cuda.manual_seed_all(s) def get_device(): try: if torch.cuda.is_available(): torch.zeros(1, device='cuda') return torch.device('cuda') except Exception: pass return torch.device('cpu') def cov_from(raw, floor=0.08): L = torch.tril(raw) d = torch.diagonal(L, dim1=1, dim2=2) L = L - torch.diag_embed(d) + torch.diag_embed(F.softplus(d) + 0.15) return L @ L.transpose(1, 2) + floor * torch.eye(ZDIM, device=raw.device) def sym_penalty(w, mu, cov, eps=0.7, lc=0.4, lw=0.8): out = torch.zeros((), device=mu.device) ident = tuple(range(K)) for p in permutations(range(K)): if p == ident: continue q = torch.tensor(p, device=mu.device) dist = torch.linalg.vector_norm(mu - mu[q], dim=1) dist = dist + lc * torch.linalg.matrix_norm(cov - cov[q], dim=(1, 2)) + lw * torch.abs(w - w[q]) out = out + F.softplus(eps - dist).sum() return out / (K * (K - 1)) def mixture_nll(z, logits, mu, raw): cov = cov_from(raw) inv = torch.linalg.inv(cov) diff = z[:, None, :] - mu[None, :, :] q = torch.einsum('nkd,kde,nke->nk', diff, inv, diff) ld = torch.logdet(cov) lp = F.log_softmax(logits, 0)[None, :] - 0.5 * (q + ld[None, :] + ZDIM * math.log(2 * math.pi)) return -torch.logsumexp(lp, 1).mean() class MixtureAE(nn.Module): def __init__(self, inp, out): super().__init__() self.enc = nn.Sequential(nn.Linear(inp, 64), nn.Tanh(), nn.Linear(64, ZDIM)) self.dec = nn.Sequential(nn.Linear(ZDIM, 64), nn.Tanh(), nn.Linear(64, out)) self.logits = nn.Parameter(torch.zeros(K)) self.mu = nn.Parameter(torch.randn(K, ZDIM) * 0.7) self.raw = nn.Parameter(torch.randn(K, ZDIM, ZDIM) * 0.05) def forward(self, x): z = self.enc(x) return self.dec(z), z def signatures(self): return F.softmax(self.logits, 0), self.mu, cov_from(self.raw) def run_one(seed, cfg, idea): seed_all(seed) ds = get_dataset(seed, 400, 200) dev = get_device() xtr = torch.as_tensor(ds['xtr'], dtype=torch.float32, device=dev) ytr = torch.as_tensor(ds['ytr'], dtype=torch.float32, device=dev) xte = torch.as_tensor(ds['xte'], dtype=torch.float32, device=dev) yte = torch.as_tensor(ds['yte'], dtype=torch.float32, device=dev) net = MixtureAE(xtr.shape[1], ytr.shape[1]).to(dev) opt = torch.optim.Adam(net.parameters(), lr=cfg['lr'], weight_decay=cfg.get('weight_decay', 0.0)) bs = 64 for _ in range(cfg['epochs']): order = torch.randperm(len(xtr), device=dev) for ix in order.split(bs): pred, z = net(xtr[ix]) w, mu, cov = net.signatures() loss = F.mse_loss(pred, ytr[ix]) + 0.03 * mixture_nll(z, net.logits, mu, net.raw) if idea: loss = loss + cfg['eta'] * sym_penalty(w, mu, cov, cfg['eps'], cfg['lc'], cfg['lw']) opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(net.parameters(), 5.0); opt.step() with torch.no_grad(): pred, z = net(xte) metric = F.mse_loss(pred, yte).item() w, mu, cov = net.signatures() sig = {'mean_pair_min': float(torch.pdist(mu).min().item()), 'weight_std': float(w.std().item()), 'mixture_nll': float(mixture_nll(net.enc(xte), net.logits, mu, net.raw).item())} return metric, sig def train_fn(cfg, idea): def f(seed): return run_one(seed, cfg, idea)[0] return f def eval_with_sig(cfg, idea): vals=[]; sigs=[] for s in SEEDS: v, sig = run_one(s, cfg, idea); vals.append(v); sigs.append(sig) return {'mean': float(np.mean(vals)), 'std': float(np.std(vals)), 'per_seed': vals, 'signatures': sigs, 'cfg': cfg} def main(): # Union parity: every lr is evaluated by both methods; baseline also sweeps its weight decay. grid = [] for lr in [0.001, 0.003, 0.006]: for wd in [0.0, 1e-4]: grid.append({'lr':lr, 'weight_decay':wd, 'epochs':12}) base = sweep_baseline(lambda c: train_fn(c, False), grid, seeds=SWEEP_SEEDS) idea_cfgs = [dict(base['best_cfg'], eta=e, eps=.7, lc=.4, lw=.8) for e in [.03, .08, .15]] # Include all lr/step sizes tried by the idea in the baseline sweep (already present in grid). idea_runs = [eval_with_sig(c, True) for c in idea_cfgs] idea = min(idea_runs, key=lambda r:r['mean']) rep = make_report(TRACK, MODEL, base, idea, {'mechanism_signature': { 'prediction': 'signature separation penalty increases minimum component-mean separation and lowers residual component similarity', 'baseline_mean_pair_min': float(np.mean([x['mean_pair_min'] for x in eval_with_sig(base['best_cfg'], False)['signatures']])), 'idea_mean_pair_min': float(np.mean([x['mean_pair_min'] for x in idea['signatures']])), 'predicted_direction': 'idea_mean_pair_min > baseline_mean_pair_min', 'confirmed': bool(np.mean([x['mean_pair_min'] for x in idea['signatures']]) > np.mean([x['mean_pair_min'] for x in eval_with_sig(base['best_cfg'], False)['signatures']])) }, 'track_rationale': 'latent_mixture_transport directly contains Gaussian-mixture component structure; built-in tracks do not.'}, ) with open('bench_report.json','w') as f: json.dump(rep,f,indent=2) print(json.dumps(rep, indent=2)) if __name__ == '__main__': main()