Skip to content

Commit 3ad0141

Browse files
committed
feat(example): more 2d distributions and better graph comparision
1 parent dc0fd34 commit 3ad0141

5 files changed

Lines changed: 239 additions & 116 deletions

File tree

‎examples/graph_2d/dataset.py‎

Lines changed: 66 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,19 @@
77
__all__ = ["MAX_POINTS", "MIN_POINTS", "PRESETS", "make_points"]
88

99
#: Synthetic 2D distributions the graph can be built on.
10-
PRESETS: tuple[str, ...] = ("blobs", "moons", "circles", "spiral", "grid", "uniform")
10+
PRESETS: tuple[str, ...] = (
11+
"blobs",
12+
"moons",
13+
"circles",
14+
"spiral",
15+
"grid",
16+
"uniform",
17+
"s-curve",
18+
"annulus",
19+
"overlap",
20+
"twins",
21+
"segments",
22+
)
1123

1224
#: Smallest and largest cloud the example accepts. The upper bound is not a memory limit but a runtime
1325
#: one: the theoretical reference graphs in `theory.py` compare the DEG against the Delaunay graph, the
@@ -70,13 +82,66 @@ def _uniform(num_points: int, rng: np.random.Generator) -> tuple[np.ndarray, np.
7082
return rng.uniform(0.0, 10.0, size=(num_points, 2)), np.zeros(num_points, dtype=np.int64)
7183

7284

85+
def _s_curve(num_points: int, rng: np.random.Generator) -> tuple[np.ndarray, np.ndarray]:
86+
"""A sine-wave manifold with thickness — a narrow curved strip to navigate along."""
87+
t = rng.uniform(0.0, 3.0 * np.pi, size=num_points)
88+
points = np.column_stack([t, np.sin(t) * 3.0])
89+
points = points + rng.standard_normal((num_points, 2)) * 0.18
90+
group = (np.sin(t) > 0.0).astype(np.int64)
91+
return points, group
92+
93+
94+
def _annulus(num_points: int, rng: np.random.Generator) -> tuple[np.ndarray, np.ndarray]:
95+
"""A filled disc whose density falls off toward the rim, grouped by radius band."""
96+
bands = 3
97+
theta = rng.uniform(0.0, 2.0 * np.pi, size=num_points)
98+
radius = rng.uniform(0.0, 1.0, size=num_points) ** 1.6 * 4.6
99+
group = np.minimum((radius / 4.6 * bands).astype(np.int64), bands - 1)
100+
return np.column_stack([radius * np.cos(theta), radius * np.sin(theta)]), group
101+
102+
103+
def _overlap(num_points: int, rng: np.random.Generator) -> tuple[np.ndarray, np.ndarray]:
104+
"""A dense blob laid over a uniform fill — heterogeneous density that shares space."""
105+
group = (rng.random(num_points) < 0.5).astype(np.int64)
106+
dense = rng.standard_normal((num_points, 2)) * 0.7 + np.array([1.5, 1.0])
107+
flat = rng.uniform(-5.0, 5.0, size=(num_points, 2))
108+
return np.where(group[:, None] == 1, dense, flat), group
109+
110+
111+
def _twins(num_points: int, rng: np.random.Generator) -> tuple[np.ndarray, np.ndarray]:
112+
"""Two spiral arms wound in opposite senses — strongly non-linearly separable."""
113+
arms = 2
114+
group = np.arange(num_points) % arms
115+
radius = np.sqrt(rng.uniform(0.0, 1.0, size=num_points)) * 4.5
116+
sense = np.where(group == 0, 1.0, -1.0)
117+
theta = radius * 1.5 * sense + group * np.pi
118+
points = np.column_stack([radius * np.cos(theta), radius * np.sin(theta)])
119+
return points + rng.standard_normal((num_points, 2)) * 0.07, group
120+
121+
122+
def _segments(num_points: int, rng: np.random.Generator) -> tuple[np.ndarray, np.ndarray]:
123+
"""Points strung along a handful of random line segments — sparse, filamentary structure."""
124+
lines = int(np.clip(num_points // 120, 3, 8))
125+
start = rng.uniform(-5.0, 5.0, size=(lines, 2))
126+
end = rng.uniform(-5.0, 5.0, size=(lines, 2))
127+
group = np.arange(num_points) % lines
128+
t = rng.uniform(0.0, 1.0, size=num_points)
129+
points = start[group] * (1.0 - t[:, None]) + end[group] * t[:, None]
130+
return points + rng.standard_normal((num_points, 2)) * 0.12, group
131+
132+
73133
_GENERATORS = {
74134
"blobs": _blobs,
75135
"moons": _moons,
76136
"circles": _circles,
77137
"spiral": _spiral,
78138
"grid": _grid,
79139
"uniform": _uniform,
140+
"s-curve": _s_curve,
141+
"annulus": _annulus,
142+
"overlap": _overlap,
143+
"twins": _twins,
144+
"segments": _segments,
80145
}
81146

82147

‎examples/graph_2d/deg_graph.py‎

Lines changed: 2 additions & 70 deletions
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,7 @@
1111
from deglib.distances import FloatSpace, Metric
1212
from deglib.optimization import mips_l2_transform_query
1313
from pysearch import epsilon_search
14-
from theory import TheoryReport, compare, dissimilarities
14+
from theory import TheoryReport, compare, dissimilarities, knng_edges, nsw_edges
1515

1616
__all__ = [
1717
"DEFAULT_K",
@@ -240,74 +240,6 @@ def deg_hit(self) -> bool:
240240
return bool(self.exact in self.deg_indices)
241241

242242

243-
def _nearest_neighbours(points: np.ndarray, k: int, metric: Metric) -> np.ndarray:
244-
"""
245-
The `[n, m]` matrix of each vertex's `m = min(k, n - 1)` nearest neighbour indices under `metric`.
246-
247-
The knng reads its links from this one computation. The neighbour set is read from the same
248-
dissimilarity the reference graphs use (`theory.dissimilarities`, whose infinite diagonal keeps a
249-
vertex out of its own neighbour set), so an inner-product cloud links by inner product and an L2
250-
cloud by Euclidean distance. The indices are distinct by construction, so a vertex's out-degree is
251-
exactly `m`.
252-
"""
253-
points = np.asarray(points)
254-
num_points = int(points.shape[0])
255-
neighbours = min(int(k), max(num_points - 1, 0))
256-
if num_points < 2 or neighbours <= 0:
257-
return np.zeros((num_points, 0), dtype=np.int64)
258-
distances = dissimilarities(points, metric)
259-
return np.argsort(distances, axis=1)[:, :neighbours].astype(np.int64)
260-
261-
262-
def knng_edges(points: np.ndarray, k: int, metric: Metric = Metric.FP32_L2) -> np.ndarray:
263-
"""
264-
The directed k-nearest-neighbour graph over `points`, as a directed `[e, 2]` int32 edge list.
265-
266-
Each vertex points at its `k` nearest neighbours under `metric`, read from the shared neighbour
267-
computation, and the link is kept in the direction it was chosen — the edge list holds `(v, neighbour)`
268-
for every selection and is never symmetrized. So the out-degree is exactly `k` per vertex (fewer only
269-
when `k` exceeds the number of other vertices), while an in-degree can exceed `k` when a vertex is a
270-
popular neighbour. A vertex can only walk the neighbours it chose itself.
271-
"""
272-
points = np.asarray(points)
273-
num_points = int(points.shape[0])
274-
nearest = _nearest_neighbours(points, k, metric)
275-
neighbours = int(nearest.shape[1])
276-
if neighbours <= 0:
277-
return np.zeros((0, 2), dtype=np.int32)
278-
279-
sources = np.repeat(np.arange(num_points, dtype=np.int32), neighbours)
280-
targets = nearest.reshape(-1).astype(np.int32)
281-
return np.column_stack((sources, targets)).astype(np.int32).reshape(-1, 2)
282-
283-
284-
def nsw_edges(points: np.ndarray, k: int, metric: Metric = Metric.FP32_L2) -> np.ndarray:
285-
"""
286-
The navigable small world over `points`, built incrementally as an undirected `[e, 2]` int32 edge list.
287-
288-
Vertices are inserted in index order, and each new vertex links — undirected — to its `k` nearest among
289-
the vertices already in the graph, the same dissimilarity the reference graphs read. Because a later
290-
vertex can add a link to an earlier one, an earlier vertex ends up with more than `k` neighbours: the
291-
degree is the count of vertices that chose it plus the `k` it chose itself, so the maximum degree is
292-
unbounded. The first vertex has no earlier vertex to link and the first few hold fewer than `k` links,
293-
simply because too few vertices precede them.
294-
"""
295-
points = np.asarray(points)
296-
num_points = int(points.shape[0])
297-
if num_points < 2:
298-
return np.zeros((0, 2), dtype=np.int32)
299-
300-
distances = dissimilarities(points, metric)
301-
kept: list[tuple[int, int]] = []
302-
for i in range(1, num_points):
303-
earlier = distances[i, :i]
304-
chosen = min(int(k), i)
305-
nearest = np.argpartition(earlier, chosen - 1)[:chosen]
306-
nearest = nearest[np.argsort(earlier[nearest], kind="stable")]
307-
kept.extend((i, int(j)) for j in nearest)
308-
return np.asarray(kept, dtype=np.int32).reshape(-1, 2)
309-
310-
311243
def build_model(
312244
points: np.ndarray,
313245
groups: np.ndarray,
@@ -382,7 +314,7 @@ def build_model(
382314
metric=metric,
383315
mips=mips,
384316
graph=graph,
385-
theory=compare(points, edges, metric) if theory else None,
317+
theory=compare(points, edges, metric, k) if theory else None,
386318
)
387319

388320

‎examples/graph_2d/test_graph_2d.py‎

Lines changed: 54 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -387,8 +387,8 @@ def test_describe_covers_the_built_graph(model) -> None:
387387

388388
assert "vertices" in text and "build" in text and "4" in text
389389
overlap = [line for line in lines if line.startswith("∩")]
390-
assert [line.split()[1] for line in overlap] == ["delaunay", "gabriel", "rng", "mst", "mrng"], (
391-
"one line per reference"
390+
assert [line.split()[1] for line in overlap] == ["knng", "nsw", "delaunay", "gabriel", "rng", "mst", "mrng", "nsg"], (
391+
"one line per reference, in the selector's order"
392392
)
393393
for line, item in zip(overlap, model.theory.graphs):
394394
assert f"{item.edges:6d}" in line and f"{item.shared:7d}" in line, "the panel quotes the report"
@@ -1052,9 +1052,11 @@ def test_the_k_field_follows_the_view_it_lands_on(viewer: GraphViewer) -> None:
10521052

10531053
class _SpinStub:
10541054
state = "normal"
1055+
increment = 1
10551056

1056-
def configure(self, state: str) -> None:
1057+
def configure(self, state: str, increment: int = 1) -> None:
10571058
self.state = state
1059+
self.increment = increment
10581060

10591061
viewer._k_spin = _SpinStub()
10601062
try:
@@ -1073,6 +1075,31 @@ def configure(self, state: str) -> None:
10731075
viewer._k_spin = None
10741076

10751077

1078+
def test_the_k_field_steps_by_two_only_for_the_deg(viewer: GraphViewer) -> None:
1079+
"""The DEG builds at an even degree, so its k arrows step by 2; the knng and NSW read k raw and step by 1."""
1080+
1081+
class _SpinStub:
1082+
state = "normal"
1083+
increment = 1
1084+
1085+
def configure(self, state: str, increment: int = 1) -> None:
1086+
self.state = state
1087+
self.increment = increment
1088+
1089+
viewer._k_spin = _SpinStub()
1090+
try:
1091+
viewer._on_view(DEG_VIEW)
1092+
assert viewer._k_spin.increment == 2, "the DEG snaps k to even, so a step of 1 would leave the down arrow inert"
1093+
1094+
viewer._on_view(KNNG_VIEW)
1095+
assert viewer._k_spin.increment == 1, "the knng reads the raw neighbour count, so it steps by 1"
1096+
1097+
viewer._on_view(NSW_VIEW)
1098+
assert viewer._k_spin.increment == 1, "the NSW reads the raw k too"
1099+
finally:
1100+
viewer._k_spin = None
1101+
1102+
10761103
def test_viewer_picks_the_vertex_under_the_cursor(viewer: GraphViewer) -> None:
10771104
viewer._fig.canvas.draw()
10781105
target = 42
@@ -1112,9 +1139,11 @@ def test_the_mrng_view_is_offered_undirected_and_timed(viewer: GraphViewer) -> N
11121139

11131140
class _SpinStub:
11141141
state = "normal"
1142+
increment = 1
11151143

1116-
def configure(self, state: str) -> None:
1144+
def configure(self, state: str, increment: int = 1) -> None:
11171145
self.state = state
1146+
self.increment = increment
11181147

11191148
viewer._k_spin = _SpinStub()
11201149
try:
@@ -1138,19 +1167,22 @@ def configure(self, state: str) -> None:
11381167
assert drawn <= symmetrised, "and that graph is a subgraph of the symmetrised NSG graph"
11391168

11401169

1141-
def test_the_nsg_view_is_offered_directed_but_never_compared(viewer: GraphViewer) -> None:
1170+
def test_the_nsg_view_is_offered_directed_and_scored_in_the_report(viewer: GraphViewer) -> None:
11421171
"""
11431172
The NSG is a viewer-built view like the knng and the NSW: offered under every metric, drawn directed,
1144-
reads no k — yet it holds no report entry, so it is never scored against the DEG.
1173+
reads no k. It carries a report entry, so the panel scores it against the DEG — yet it stays out of
1174+
the colour key, so its own view states its edge count and time, never a DEG share.
11451175
"""
11461176
assert NSG_VIEW in viewer.views, "the NSG is offered in the selector"
1147-
assert NSG_VIEW not in viewer.model.theory.names, "but it is not a reference graph in the comparison"
1177+
assert NSG_VIEW in viewer.model.theory.names, "the NSG is scored against the DEG in the report"
11481178

11491179
class _SpinStub:
11501180
state = "normal"
1181+
increment = 1
11511182

1152-
def configure(self, state: str) -> None:
1183+
def configure(self, state: str, increment: int = 1) -> None:
11531184
self.state = state
1185+
self.increment = increment
11541186

11551187
viewer._k_spin = _SpinStub()
11561188
try:
@@ -1162,7 +1194,7 @@ def configure(self, state: str) -> None:
11621194
assert viewer._view_directed(), "the NSG is drawn directed"
11631195
panel = viewer._panel_text.get_text()
11641196
assert panel.startswith("nsg"), "the column heads with the NSG's name"
1165-
assert "of DEG" not in panel, "the NSG states no DEG share — it is not compared"
1197+
assert "of DEG" not in panel, "the NSG view states its own edge count and time, not a DEG share"
11661198

11671199
drawn = {(int(u), int(v)) for u, v in viewer._display_edges()}
11681200
library = {
@@ -1499,7 +1531,7 @@ def test_compare_counts_the_deg_edges_each_graph_shares() -> None:
14991531

15001532
undirected = {tuple(sorted(edge)) for edge in scene.edges.reshape(-1, 2)}
15011533
assert report.deg_edges == len(undirected)
1502-
assert tuple(item.name for item in report.graphs) == ("delaunay", "gabriel", "rng", "mst", "mrng")
1534+
assert tuple(item.name for item in report.graphs) == ("knng", "nsw", "delaunay", "gabriel", "rng", "mst", "mrng", "nsg")
15031535
assert all(0 < item.shared <= min(report.deg_edges, item.edges) for item in report.graphs)
15041536

15051537
by_name = {item.name: item for item in report.graphs}
@@ -1631,8 +1663,8 @@ def test_an_inner_product_report_carries_only_the_graphs_it_can_decide() -> None
16311663
euclidean = build_scene("blobs", 120, DEFAULT_K, 5, Metric.FP32_L2, threads=1).theory
16321664
inner = build_scene("blobs", 120, DEFAULT_K, 5, Metric.FP32_InnerProduct, threads=1).theory
16331665

1634-
assert euclidean.names == ("delaunay", "gabriel", "rng", "mst", "mrng")
1635-
assert inner.names == ("rng", "mst", "mrng")
1666+
assert euclidean.names == ("knng", "nsw", "delaunay", "gabriel", "rng", "mst", "mrng", "nsg")
1667+
assert inner.names == ("knng", "nsw", "rng", "mst", "mrng", "nsg")
16361668

16371669

16381670
def test_the_colour_key_names_only_what_the_report_holds() -> None:
@@ -1831,7 +1863,7 @@ def test_viewer_reports_the_theory_overlap_on_every_build(viewer: GraphViewer) -
18311863

18321864
assert viewer.model.theory is not None, "a fresh cloud is compared again"
18331865
names = tuple(item.name for item in viewer.model.theory.graphs)
1834-
assert len(names) == 5, "every reference graph is listed"
1866+
assert len(names) == 8, "every graph the selector offers is listed, the knng and NSW among them"
18351867
assert all(f"∩ {name}" in viewer._panel_text.get_text() for name in names), "the reseeded panel lists them all"
18361868
assert all(f"∩ {name}" in text for name in names), "and so did the panel before the reseed"
18371869
assert all(item.shared > 0 for item in viewer.model.theory.graphs)
@@ -1925,13 +1957,13 @@ def test_an_overlay_switch_survives_a_rebuild(viewer: GraphViewer) -> None:
19251957

19261958
viewer._on_metric("IP")
19271959

1928-
assert viewer.overlays == {"mst", "mrng"}, "the graphs that were on stay on, the hidden one stays off"
1960+
assert viewer.overlays == {"mst", "mrng", "nsg", "knng", "nsw"}, "the graphs that were on stay on, the hidden one stays off"
19291961
assert viewer._legend_names() == ("DEG only", "rng", "mrng", "mst"), "and the key names exactly those"
19301962

19311963
viewer._on_metric("L2")
19321964

19331965
assert "rng" not in viewer.overlays, "the graph is still hidden when the metric brings it back"
1934-
assert viewer.overlays == {"delaunay", "gabriel", "mst", "mrng"}
1966+
assert viewer.overlays == {"delaunay", "gabriel", "mst", "mrng", "nsg", "knng", "nsw"}
19351967

19361968

19371969
def test_a_graph_new_to_the_report_defaults_on(viewer: GraphViewer) -> None:
@@ -2117,16 +2149,20 @@ def test_viewer_draws_the_graph_the_selector_picks(viewer: GraphViewer) -> None:
21172149
DEG_VIEW,
21182150
KNNG_VIEW,
21192151
NSW_VIEW,
2120-
*(overlap.name for overlap in viewer.model.theory.graphs),
2152+
"delaunay",
2153+
"gabriel",
2154+
"rng",
2155+
"mst",
2156+
"mrng",
21212157
NSG_VIEW,
21222158
NONE_VIEW,
2123-
), "the DEG, the knng, the NSW, the reference graphs, the NSG and the empty view are the selector's entries"
2159+
), "the selector lists every graph in GRAPH_ORDER's order, the knng, NSW and NSG among them"
21242160

21252161
viewer._fig.canvas.draw()
21262162
assert len(viewer._edges.get_segments()) == viewer.model.edges.shape[0], "the DEG view draws the graph itself"
21272163
assert viewer._legend.get_visible(), "and has to explain what its edge colours mean"
21282164

2129-
for name in viewer.model.theory.names:
2165+
for name in (n for n in OVERLAP_ORDER if n in viewer.model.theory.names):
21302166
viewer._on_view(name)
21312167
viewer._fig.canvas.draw()
21322168

0 commit comments

Comments
 (0)