import json, math, os, random import numpy as np import torch from torch import nn SEED = 2016 np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED) def tents(diagram, grid): d = np.asarray(diagram, dtype=float).reshape(-1, 2) b, death = d[:, 0:1], d[:, 1:2] return np.maximum(0.0, np.minimum(grid[None, :] - b, death - grid[None, :])) def landscape(diagram, grid, K): v = tents(diagram, grid) if len(v) == 0: return np.zeros((K, len(grid))) v = np.sort(v, axis=0)[::-1] out = np.zeros((K, len(grid))) out[:min(K, len(v))] = v[:K] return out def lp_grid(a, b, p=2, dt=0.002): # Common grid over the supports; sufficiently fine for this sanity check. lo = min([x[0] for x in a] + [x[0] for x in b]) - .01 hi = max([x[1] for x in a] + [x[1] for x in b]) + .01 g = np.arange(lo, hi + dt/2, dt) return (np.sum(np.abs(landscape(a, g, max(len(a),len(b))) - landscape(b, g, max(len(a),len(b))))**p) * dt) ** (1/p) def tent_distance(x, y, dt=0.001): lo = min(x[0], y[0]) - .01; hi = max(x[1], y[1]) + .01 g = np.arange(lo, hi + dt/2, dt) return np.sqrt(np.sum(np.abs(tents([x], g)[0] - tents([y], g)[0])**2) * dt) def matching_bound(a, b, dt=0.001): # Equal-cardinality matched-pair W_2^triangle upper bound (the chosen matching). return math.sqrt(sum(tent_distance(x, y, dt)**2 for x, y in zip(a, b))) class PersistenceLandscape(nn.Module): def __init__(self, t_min=0., t_max=1., n_grid=32, K=3): super().__init__() self.K, self.n_grid = K, n_grid self.register_buffer('grid', torch.linspace(t_min, t_max, n_grid)) def forward(self, diagrams, mask=None): # diagrams [B,N,2], mask [B,N], padded rows are ignored. b = diagrams[..., 0:1]; d = diagrams[..., 1:2] t = self.grid.view(1, 1, -1) v = torch.relu(torch.minimum(t-b, d-t)) if mask is not None: v = v.masked_fill(~mask[..., None], -1e9) vals, _ = torch.topk(v, k=min(self.K, v.shape[1]), dim=1) if vals.shape[1] < self.K: vals = torch.cat([vals, torch.zeros(*vals.shape[:1], self.K-vals.shape[1], vals.shape[2], device=vals.device)], 1) return vals.clamp_min(0.) def make_data(n, N=6, noise=0., seed=0): rng=np.random.default_rng(seed); ds=[]; ys=[] for i in range(n): y=int(i % 2); center=.34 if y==0 else .66 births=np.clip(rng.normal(center, .09, N), .03, .82) pers=np.clip(rng.normal(.22, .05, N), .04, .38) deaths=np.minimum(births+pers, .97) ds.append(np.stack([births,deaths],1)); ys.append(y) ds=np.asarray(ds) if noise: ds=np.stack([np.clip(ds[:,:,0]+rng.normal(0,noise,ds.shape[:2]),0,.9), np.clip(ds[:,:,1]+rng.normal(0,noise,ds.shape[:2]),.05,1)],2); ds[:,:,1]=np.maximum(ds[:,:,1],ds[:,:,0]+.01) return ds.astype('float32'), np.asarray(ys, dtype='int64') def train_eval(kind, train, ytr, test, yte, epochs=80): device='cuda' if torch.cuda.is_available() else 'cpu' try: if kind=='landscape': layer=PersistenceLandscape(.0,1.,32,3); Xtr=layer(torch.tensor(train)).flatten(1); Xte=layer(torch.tensor(test)).flatten(1) else: Xtr=torch.tensor(train).flatten(1); Xte=torch.tensor(test).flatten(1) model=nn.Sequential(nn.Linear(Xtr.shape[1],32),nn.ReLU(),nn.Linear(32,2)).to(device) Xtr,Xte=Xtr.to(device),Xte.to(device); yt=torch.tensor(ytr,device=device); yv=torch.tensor(yte,device=device) opt=torch.optim.Adam(model.parameters(),lr=.01) for _ in range(epochs): opt.zero_grad(); loss=nn.functional.cross_entropy(model(Xtr),yt); loss.backward(); opt.step() with torch.no_grad(): acc=(model(Xte).argmax(1)==yv).float().mean().item() return acc except Exception: # CUDA OOM or driver issues: rerun on CPU. torch.cuda.empty_cache() if torch.cuda.is_available() else None old=torch.cuda.is_available; torch.cuda.is_available=lambda:False try: return train_eval(kind,train,ytr,test,yte,epochs) finally: torch.cuda.is_available=old def main(): # Mechanism verification: predictions are (P1) one point saturates the bound, # (P2) every multi-point ratio is <=1, and (P3) larger noise cannot amplify # landscape perturbation beyond the chosen matching bound. grid=np.linspace(0,1,2001); rows=[] a=[(.15,.55)]; for shift in [.002,.01,.03,.08]: b=[(a[0][0]+shift,a[0][1]+shift)] lhs=np.linalg.norm(landscape(a,grid,3)-landscape(b,grid,3))*math.sqrt(1/2000) rhs=tent_distance(a[0],b[0]); rows.append({'single_shift':shift,'landscape_norm':lhs,'bound':rhs,'ratio':lhs/rhs}) rng=np.random.default_rng(4); ratios=[] for _ in range(100): A=[]; B=[] for j in range(5): x=float(rng.uniform(.05,.75)); p=float(rng.uniform(.08,.3)); A.append((x,x+p)) dx=float(rng.normal(0,.04)); dp=float(rng.normal(0,.03)); z=max(.01,min(.9,x+dx)); B.append((z,min(.99,z+max(.02,p+dp)))) ratios.append(lp_grid(A,B)/matching_bound(A,B)) # P3: piecewise-linear tents integrated on a uniform grid should converge # to the continuous L2 norm as dt shrinks (first-order rectangle estimate). ca=[(.13,.57),(.31,.76),(.61,.91)]; cb=[(.16,.55),(.28,.79),(.65,.88)] ref=lp_grid(ca, cb, dt=0.0001) discretization=[] for dt in [.02,.01,.005,.0025,.00125]: val=lp_grid(ca, cb, dt=dt) discretization.append({'dt':dt,'grid_norm':val,'abs_error':abs(val-ref),'error_over_dt':abs(val-ref)/dt}) # Autograd prediction: away from max/tent kinks, the layer has finite gradients. q=torch.tensor([[[.18,.52],[.43,.81]]],dtype=torch.float32,requires_grad=True) qout=PersistenceLandscape(.0,1.,64,2)(q) qout.sum().backward() grad_ok=bool(torch.isfinite(q.grad).all() and q.grad.abs().sum()>0) prediction={'single_point_equality':rows,'multi_point_max_ratio':float(max(ratios)), 'multi_point_mean_ratio':float(np.mean(ratios)),'predicted_bound':1.0, 'grid_convergence':{'reference_dt':0.0001,'sweep':discretization, 'prediction':'error decreases as dt decreases'}, 'autograd_finite_nonzero':grad_ok} tr,ytr=make_data(160,seed=10); te,yte=make_data(80,seed=11) clean={'padded_birth_death':train_eval('raw',tr,ytr,te,yte),'landscape':train_eval('landscape',tr,ytr,te,yte)} noisy=[] for sigma in [.01,.03,.07]: nt,ny=make_data(80,noise=sigma,seed=20); noisy.append({'sigma':sigma,'padded_birth_death':train_eval('raw',tr,ytr,nt,ny),'landscape':train_eval('landscape',tr,ytr,nt,ny)}) result={'prediction_checks':prediction,'clean_accuracy':clean,'noise_accuracy':noisy} with open('results.json','w') as f: json.dump(result,f,indent=2) print(json.dumps(result,indent=2)) if __name__=='__main__': main()