Boundary-Radial Persistence Loss / boundary_radial.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 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]