Regret-aware evidential cost compression / math_experiment.py
Mechanism failed
1import json, time
2import numpy as np
3
4LAYERS, WIDTH = 4, 2
5
6def instance(seed, K=8):
7 r=np.random.default_rng(seed); edges=[]
8 for t in range(LAYERS):
9 src=range(t*WIDTH,(t+1)*WIDTH) if t else [-1]
10 dst=range((t+1)*WIDTH,(t+2)*WIDTH) if t<LAYERS-1 else [-2]
11 edges += [(a,b) for a in src for b in dst]
12 foc=[]
13 for _ in edges:
14 l=np.maximum(.03,r.normal(1.2,.25,K)); u=l+r.uniform(.05,.75,K)
15 foc.append((l,u,r.dirichlet(np.ones(K))))
16 return edges,foc
17
18def routes(edges):
19 lookup={p:i for i,p in enumerate(edges)}; out=[]
20 for q in np.ndindex(*(WIDTH,)*(LAYERS-1)):
21 ns=[-1]+[t*WIDTH+q[t-1] for t in range(1,LAYERS)]+[-2]
22 x=np.zeros(len(edges),dtype=np.int8)
23 for a,b in zip(ns[:-1],ns[1:]): x[lookup[a,b]]=1
24 out.append(x)
25 return np.asarray(out)
26
27def compressed(foc,groups):
28 # Deliberately conservative surrogate: merged block cost is its mass times
29 # max lower endpoint, hence every merge only raises edge cost.
30 vals=[]
31 for e,(l,u,m) in enumerate(foc):
32 vals.append(sum(max(float(np.dot(m[z],l[z])),m[z].sum()*float(np.max(l[z]))) for z in groups[e]))
33 return np.asarray(vals)
34
35def merge_groups(groups,e,i,j):
36 g=[[z[:] for z in ge] for ge in groups]; g[e][i]+=g[e][j]; del g[e][j]; return g
37
38def route(c,X):
39 v=X@c; k=int(np.argmin(v)); return k,float(v[k])
40
41def compress(foc,X,target,mode='regret',lam=.1):
42 groups=[[[i] for i in range(len(foc[e][0]))] for e in range(len(foc))]
43 c=compressed(foc,groups); x0=X[np.argmin(X@c)]
44 while sum(len(z) for ge in groups for z in ge)>target*len(groups):
45 best=None
46 for e,g in enumerate(groups):
47 for i in range(len(g)):
48 for j in range(i+1,len(g)):
49 ng=merge_groups(groups,e,i,j); nc=compressed(foc,ng); d=nc-c
50 if mode=='regret': score=float(d@x0+lam*np.abs(d).sum())
51 elif mode=='random': score=float(np.random.default_rng(e+i*31+j).random())
52 elif mode=='mass':
53 score=-float(foc[e][2][g[i]].sum()+foc[e][2][g[j]].sum())
54 else:
55 l,u,m=foc[e]; ai,bj=g[i],g[j]
56 lo=max(min(l[ai]),min(l[bj])); hi=min(max(u[ai]),max(u[bj]))
57 union=max(max(u[ai]),max(u[bj]))-min(min(l[ai]),min(l[bj]))
58 score=-max(0.,hi-lo)/max(union,1e-12)
59 if best is None or score<best[0]: best=(score,ng,nc)
60 if best is None: break
61 _,groups,c=best
62 return groups,c,x0
63
64def theorem_check():
65 rng=np.random.default_rng(7); violation=0.; ratios=[]
66 for _ in range(8):
67 e,f=instance(int(rng.integers(1e9))); X=routes(e); g,ch,_=compress(f,X,3); c=compressed(f,[[[i] for i in range(8)] for _ in f])
68 delta=X@(ch-c); star=int(np.argmin(X@c)); hat=int(np.argmin(X@ch))
69 regret=float((X[hat]-X[star])@c); bound=float(delta[star])
70 violation=max(violation,regret-bound); ratios.append(regret/(bound+1e-12))
71 return {'max_bound_violation':float(violation),'mean_regret_bound_ratio':float(np.mean(ratios))}
72
73def main():
74 start=time.time(); check=theorem_check(); summary={}
75 for K in (4,8,12):
76 for target in (1,2,3):
77 vals=[]
78 for seed in range(3):
79 e,f=instance(seed+100,K); X=routes(e); c=compressed(f,[[[i] for i in range(K)] for _ in f]); star,_=route(c,X)
80 _,ch,_=compress(f,X,target); hat,_=route(ch,X); d=X@(ch-c)
81 vals.append((float((X[hat]-X[star])@c),int(hat!=star),float(d[star]),float(np.sum(ch-c))))
82 summary[f'K{K}_to{target}']={'mean_regret':float(np.mean([v[0] for v in vals])),'flip_rate':float(np.mean([v[1] for v in vals])),'mean_bound':float(np.mean([v[2] for v in vals])),'mean_inflation':float(np.mean([v[3] for v in vals]))}
83 baselines={}
84 for mode in ('regret','jaccard','random'):
85 vals=[]
86 for seed in range(3):
87 e,f=instance(seed+900,8); X=routes(e); c=compressed(f,[[[i] for i in range(8)] for _ in f]); star,_=route(c,X)
88 _,ch,_=compress(f,X,2,mode); h,_=route(ch,X); vals.append(((X[h]-X[star])@c,h!=star))
89 baselines[mode]={'mean_regret':float(np.mean([v[0] for v in vals])),'flip_rate':float(np.mean([v[1] for v in vals]))}
90 print(json.dumps({'theorem':check,'sweep':summary,'baselines':baselines,'seconds':time.time()-start},indent=2))
91if __name__=='__main__': main()