Boundary-Radial Persistence Loss / boundary_radial.py
Mechanism confirmed, baseline not beaten
1"""MVP boundary-radial persistence loss.
2
3The general routine computes capped H0 persistence on a polygonal boundary
4 graph. Zero-length tie bars are discarded, as standard persistence diagrams
5 discard diagonal points. The toy circle helper uses one canonical interval per
6 connected circular boundary component, making the mechanism sweeps transparent.
7Full XRPH relative/H1 computation and differentiable gradients are not included.
8"""
9from dataclasses import dataclass
10import numpy as np
11from scipy.optimize import linear_sum_assignment
12
13@dataclass(frozen=True)
14class Bar:
15 birth: float
16 death: float
17 weight: float = 1.0
18
19def boundary_graph_bars(points, components, tol=1e-10):
20 points = np.asarray(points, dtype=float)
21 if points.size == 0:
22 return []
23 points = points.reshape((-1, 2))
24 r = np.linalg.norm(points, axis=1)
25 edges = []
26 for comp in components:
27 comp = list(comp)
28 for a, b in zip(comp, comp[1:] + comp[:1]):
29 edges.append((max(r[a], r[b]), a, b))
30 events = [(r[i], 0, i, i) for i in range(len(points))]
31 events += [(f, 1, a, b) for f, a, b in edges]
32 events.sort(key=lambda z: (z[0], z[1]))
33 parent = list(range(len(points))); birth = r.copy()
34 active = np.zeros(len(points), dtype=bool); bars = []
35 def find(x):
36 while parent[x] != x:
37 parent[x] = parent[parent[x]]; x = parent[x]
38 return x
39 for f, typ, a, b in events:
40 if typ == 0:
41 active[a] = True
42 elif active[a] and active[b]:
43 ra, rb = find(a), find(b)
44 if ra != rb:
45 keep, kill = (ra, rb) if birth[ra] <= birth[rb] else (rb, ra)
46 if f - birth[kill] > tol:
47 bars.append(Bar(float(birth[kill]), float(f)))
48 parent[kill] = keep
49 maxf = float(max(r))
50 roots = {find(i) for i in range(len(points)) if active[i]}
51 bars.extend(Bar(float(birth[q]), maxf) for q in roots)
52 return sorted(bars, key=lambda x: (x.birth, x.death))
53
54def match_loss(pred, target, unmatched=1.0):
55 """Hungarian L1 endpoint matching plus unmatched penalty."""
56 n, m = len(pred), len(target)
57 if n == 0 and m == 0: return 0.0, []
58 k = n + m
59 cost = np.full((k, k), float(unmatched))
60 if n and m:
61 cost[:n, :m] = [[abs(p.birth-t.birth)+abs(p.death-t.death)
62 for t in target] for p in pred]
63 cost[n:, m:] = 0.0
64 rows, cols = linear_sum_assignment(cost)
65 return float(cost[rows, cols].sum()), list(zip(rows.tolist(), cols.tolist()))
66
67def make_components(radii, n=24):
68 pts, comps = [], []
69 for rad in radii:
70 start = len(pts)
71 pts.extend([[rad*np.cos(2*np.pi*j/n), rad*np.sin(2*np.pi*j/n)]
72 for j in range(n)])
73 comps.append(list(range(start, start+n)))
74 return np.asarray(pts, dtype=float).reshape((-1, 2)), comps
75
76def bars_from_radii(radii):
77 """Canonical radial boundary-component bars used by the toy experiment."""
78 return [Bar(float(rad), float(rad)) for rad in radii]
79
80def radial_loss(pred_radii, target_radii, unmatched=1.0):
81 return match_loss(bars_from_radii(pred_radii), bars_from_radii(target_radii), unmatched)[0]