rayzin 0.0.1__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
rayzin/__init__.py ADDED
@@ -0,0 +1,26 @@
1
+ __version__ = "0.0.1"
2
+
3
+ from rayzin.enums import MetricType, SearchBackendType
4
+ from rayzin.pipeline import (
5
+ build_manifest,
6
+ build_manifest_from_cogs,
7
+ build_manifest_from_zarr,
8
+ knn_cog_search,
9
+ knn_zarr_search,
10
+ )
11
+ from rayzin.types import SearchResults
12
+
13
+ __all__ = [
14
+ "__version__",
15
+ # pipeline
16
+ "knn_zarr_search",
17
+ "knn_cog_search",
18
+ # pipeline — manifest build
19
+ "build_manifest",
20
+ "build_manifest_from_zarr",
21
+ "build_manifest_from_cogs",
22
+ # pipeline API types
23
+ "MetricType",
24
+ "SearchBackendType",
25
+ "SearchResults",
26
+ ]
rayzin/enums.py ADDED
@@ -0,0 +1,17 @@
1
+ from enum import StrEnum
2
+
3
+
4
+ class MetricType(StrEnum):
5
+ EUCLIDEAN = "euclidean"
6
+ COSINE = "cosine"
7
+
8
+
9
+ class ReaderType(StrEnum):
10
+ ZARR = "zarr"
11
+ COG = "cog"
12
+
13
+
14
+ class SearchBackendType(StrEnum):
15
+ NUMPY = "numpy"
16
+ FAISS_CPU = "faiss_cpu"
17
+ FAISS_GPU = "faiss_gpu"
@@ -0,0 +1,7 @@
1
+ from rayzin.manifest.filtering import filter_manifest
2
+ from rayzin.manifest.schema import MANIFEST_SCHEMA
3
+
4
+ __all__ = [
5
+ "MANIFEST_SCHEMA",
6
+ "filter_manifest",
7
+ ]
@@ -0,0 +1,166 @@
1
+ from collections.abc import Iterator
2
+ from itertools import product
3
+ from typing import Any
4
+
5
+ import numpy as np
6
+ import pyarrow as pa # type: ignore[import-untyped]
7
+ import zarr
8
+
9
+ from rayzin.manifest.schema import CHUNK_SCHEMA, MANIFEST_SCHEMA, ChunkTable, ManifestTable
10
+ from rayzin.readers.protocol import VectorReader
11
+ from rayzin.readers.zarr_layout import index_axis_names, index_axis_positions
12
+ from rayzin.types import (
13
+ COL_DIM,
14
+ COL_SLICE,
15
+ COL_START,
16
+ COL_STOP,
17
+ COL_URL,
18
+ ChunkRecord,
19
+ ChunkRef,
20
+ DimSlice,
21
+ IndexSlice,
22
+ )
23
+
24
+
25
+ def build_zarr_chunk_table(
26
+ store_url: str,
27
+ *,
28
+ array_name: str,
29
+ store_kwargs: dict[str, Any],
30
+ embedding_dim_name: str = "embedding",
31
+ ) -> ChunkTable:
32
+ return pa.Table.from_pylist(
33
+ list(
34
+ iter_zarr_chunk_slices(
35
+ store_url,
36
+ array_name=array_name,
37
+ store_kwargs=store_kwargs,
38
+ embedding_dim_name=embedding_dim_name,
39
+ )
40
+ ),
41
+ schema=CHUNK_SCHEMA,
42
+ )
43
+
44
+
45
+ def iter_zarr_chunk_slices(
46
+ store_url: str,
47
+ *,
48
+ array_name: str,
49
+ store_kwargs: dict[str, Any],
50
+ embedding_dim_name: str,
51
+ ) -> Iterator[ChunkRef]:
52
+ group = zarr.open_group(
53
+ store=store_url,
54
+ mode="r",
55
+ storage_options=store_kwargs or None,
56
+ )
57
+ array = group[array_name]
58
+ assert isinstance(array, zarr.Array)
59
+
60
+ if array.ndim < 2:
61
+ msg = (
62
+ "Expected at least one index axis plus the embedding axis, "
63
+ f"got {array.ndim} dimensions."
64
+ )
65
+ raise ValueError(msg)
66
+
67
+ axis_positions = index_axis_positions(array, embedding_dim_name=embedding_dim_name)
68
+ axis_names = index_axis_names(array, embedding_dim_name=embedding_dim_name)
69
+ index_shape = tuple(array.shape[index] for index in axis_positions)
70
+ index_chunks = tuple(array.chunks[index] for index in axis_positions)
71
+ start_ranges = [range(0, size, chunk) for size, chunk in zip(index_shape, index_chunks)]
72
+
73
+ for starts in product(*start_ranges):
74
+ yield ChunkRef(
75
+ url=store_url,
76
+ slice=[
77
+ DimSlice(
78
+ dim=axis_name,
79
+ start=start,
80
+ stop=min(start + chunk, size),
81
+ )
82
+ for axis_name, start, chunk, size in zip(
83
+ axis_names,
84
+ starts,
85
+ index_chunks,
86
+ index_shape,
87
+ )
88
+ ],
89
+ )
90
+
91
+
92
+ def compute_chunk_summary_arrow(batch: ChunkTable, reader: VectorReader) -> ManifestTable:
93
+ return _compute_chunk_summary(batch, reader)
94
+
95
+
96
+ def _compute_chunk_summary(batch: ChunkTable, reader: VectorReader) -> ManifestTable:
97
+ rows = _chunk_slice_rows(batch)
98
+ centroids: list[list[float]] = []
99
+ radii: list[float] = []
100
+ counts: list[int] = []
101
+
102
+ for row in rows:
103
+ chunk = _chunk_record_from_slice(row)
104
+ vectors, _shape = reader.read(chunk)
105
+ if vectors.shape[0] == 0:
106
+ msg = f"Chunk {chunk[COL_URL]} produced no vectors."
107
+ raise ValueError(msg)
108
+
109
+ centroid = vectors.mean(axis=0).astype(np.float32)
110
+ distances = np.linalg.norm(vectors - centroid[None, :], axis=1)
111
+ centroids.append(centroid.tolist())
112
+ radii.append(float(distances.max()))
113
+ counts.append(int(vectors.shape[0]))
114
+
115
+ return pa.Table.from_arrays(
116
+ [
117
+ batch.column(COL_URL),
118
+ batch.column(COL_SLICE),
119
+ pa.array(counts, type=pa.int32()),
120
+ pa.array(centroids, type=pa.list_(pa.float32())),
121
+ pa.array(radii, type=pa.float32()),
122
+ ],
123
+ schema=MANIFEST_SCHEMA,
124
+ )
125
+
126
+
127
+ def _chunk_slice_rows(batch: ChunkTable) -> list[ChunkRef]:
128
+ urls = batch.column(COL_URL).to_pylist()
129
+ slices = batch.column(COL_SLICE).to_pylist()
130
+ return [
131
+ ChunkRef(
132
+ url=str(urls[i]),
133
+ slice=_coerce_index_slice(slices[i]),
134
+ )
135
+ for i in range(batch.num_rows)
136
+ ]
137
+
138
+
139
+ def _chunk_record_from_slice(chunk: ChunkRef) -> ChunkRecord:
140
+ return ChunkRecord(
141
+ url=chunk[COL_URL],
142
+ slice=chunk[COL_SLICE],
143
+ count=0,
144
+ centroid=np.empty(0, dtype=np.float32),
145
+ radius=0.0,
146
+ )
147
+
148
+
149
+ def _coerce_index_slice(raw: Any) -> IndexSlice:
150
+ if not isinstance(raw, list):
151
+ msg = f"Expected slice to be a list, got {type(raw).__name__}."
152
+ raise TypeError(msg)
153
+
154
+ parts: IndexSlice = []
155
+ for part in raw:
156
+ if not isinstance(part, dict):
157
+ msg = f"Expected each slice entry to be a dict, got {type(part).__name__}."
158
+ raise TypeError(msg)
159
+ parts.append(
160
+ DimSlice(
161
+ dim=str(part[COL_DIM]),
162
+ start=int(part[COL_START]),
163
+ stop=int(part[COL_STOP]),
164
+ )
165
+ )
166
+ return parts
@@ -0,0 +1,60 @@
1
+ from typing import Any
2
+
3
+ import ray.data
4
+ from ray.data.expressions import Expr
5
+ from shapely.geometry.base import BaseGeometry # type: ignore[import-untyped]
6
+
7
+ from rayzin.manifest.spatial import (
8
+ GeoTransform,
9
+ chunk_from_row,
10
+ chunk_polygon,
11
+ read_zarr_geotransform,
12
+ validate_aoi_geometry,
13
+ )
14
+
15
+
16
+ def filter_manifest(
17
+ dataset: ray.data.Dataset,
18
+ *,
19
+ filter_expr: Expr | None = None,
20
+ aoi: BaseGeometry | None = None,
21
+ store_kwargs: dict[str, Any] | None = None,
22
+ ) -> ray.data.Dataset:
23
+ filtered = dataset
24
+ if isinstance(filter_expr, Expr):
25
+ filtered = filtered.filter(expr=filter_expr)
26
+ if aoi is not None:
27
+ filtered = filtered.filter(
28
+ ChunkIntersectsAOI,
29
+ fn_constructor_kwargs={
30
+ "aoi": aoi,
31
+ "store_kwargs": store_kwargs or {},
32
+ },
33
+ )
34
+ return filtered
35
+
36
+
37
+ class ChunkIntersectsAOI:
38
+ # TODO: we can get away with opening tha zarr store once since we can extract the geotransform
39
+ # and analytically check for intersection given a chunks key.
40
+ # TODO: this should probably be renamed to something zarr specific and we should probable
41
+ # support cog by forcing manifest to be geoparquet and adding geometry up front per cog tile.
42
+ def __init__(
43
+ self,
44
+ aoi: BaseGeometry,
45
+ store_kwargs: dict[str, Any] | None = None,
46
+ ) -> None:
47
+ self._aoi = validate_aoi_geometry(aoi)
48
+ self._store_kwargs = store_kwargs or {}
49
+ self._transform_cache: dict[str, GeoTransform] = {}
50
+
51
+ def __call__(self, row: dict[str, Any]) -> bool:
52
+ chunk = chunk_from_row(row)
53
+ transform = self._transform_cache.get(chunk["url"])
54
+ if transform is None:
55
+ transform = read_zarr_geotransform(
56
+ chunk["url"],
57
+ store_kwargs=self._store_kwargs,
58
+ )
59
+ self._transform_cache[chunk["url"]] = transform
60
+ return bool(chunk_polygon(chunk, transform).intersects(self._aoi))
@@ -0,0 +1,77 @@
1
+ from typing import Annotated, TypeAlias
2
+
3
+ import pyarrow as pa # type: ignore[import-untyped]
4
+
5
+ from rayzin.types import (
6
+ COL_CENTROID,
7
+ COL_CHUNK_ID,
8
+ COL_COUNT,
9
+ COL_DIM,
10
+ COL_DISTANCE,
11
+ COL_LOWER_BOUNDS,
12
+ COL_MIN_LOWER_BOUND,
13
+ COL_OFFSET,
14
+ COL_QUERY_ID,
15
+ COL_RADIUS,
16
+ COL_SLICE,
17
+ COL_START,
18
+ COL_STOP,
19
+ COL_URL,
20
+ )
21
+
22
+ INDEX_SLICE_TYPE: pa.ListType = pa.list_(
23
+ pa.struct(
24
+ [
25
+ pa.field(COL_DIM, pa.string()),
26
+ pa.field(COL_START, pa.int64()),
27
+ pa.field(COL_STOP, pa.int64()),
28
+ ]
29
+ )
30
+ )
31
+
32
+ CHUNK_SCHEMA: pa.Schema = pa.schema(
33
+ [
34
+ pa.field(COL_URL, pa.string()),
35
+ pa.field(COL_SLICE, INDEX_SLICE_TYPE),
36
+ ]
37
+ )
38
+
39
+ MANIFEST_SCHEMA: pa.Schema = pa.schema(
40
+ [
41
+ pa.field(COL_URL, pa.string()),
42
+ pa.field(COL_SLICE, INDEX_SLICE_TYPE),
43
+ pa.field(COL_COUNT, pa.int32()),
44
+ pa.field(COL_CENTROID, pa.list_(pa.float32())),
45
+ pa.field(COL_RADIUS, pa.float32()),
46
+ ]
47
+ )
48
+
49
+ LOWER_BOUND_SCHEMA: pa.Schema = MANIFEST_SCHEMA.append(
50
+ pa.field(COL_LOWER_BOUNDS, pa.list_(pa.float32()))
51
+ ).append(pa.field(COL_MIN_LOWER_BOUND, pa.float32()))
52
+
53
+ SEARCH_RESULT_SCHEMA: pa.Schema = pa.schema(
54
+ [
55
+ pa.field(COL_QUERY_ID, pa.int64()),
56
+ pa.field(COL_CHUNK_ID, pa.string()),
57
+ pa.field(COL_URL, pa.string()),
58
+ pa.field(COL_SLICE, INDEX_SLICE_TYPE),
59
+ pa.field(COL_OFFSET, pa.int64()),
60
+ pa.field(COL_DISTANCE, pa.float32()),
61
+ ]
62
+ )
63
+
64
+ BLOCK_SEARCH_SUMMARY_SCHEMA: pa.Schema = pa.schema(
65
+ [
66
+ pa.field("rows_seen", pa.int64()),
67
+ pa.field("rows_searched", pa.int64()),
68
+ pa.field("query_evaluations", pa.int64()),
69
+ pa.field("results_added", pa.int64()),
70
+ ]
71
+ )
72
+
73
+ ChunkTable: TypeAlias = Annotated[pa.Table, CHUNK_SCHEMA]
74
+ ManifestTable: TypeAlias = Annotated[pa.Table, MANIFEST_SCHEMA]
75
+ LowerBoundTable: TypeAlias = Annotated[pa.Table, LOWER_BOUND_SCHEMA]
76
+ SearchResultTable: TypeAlias = Annotated[pa.Table, SEARCH_RESULT_SCHEMA]
77
+ BlockSearchSummaryTable: TypeAlias = Annotated[pa.Table, BLOCK_SEARCH_SUMMARY_SCHEMA]
@@ -0,0 +1,137 @@
1
+ from collections.abc import Sequence
2
+ from typing import Any
3
+
4
+ import zarr
5
+ from shapely.geometry.base import BaseGeometry # type: ignore[import-untyped]
6
+
7
+ from rayzin.types import (
8
+ COL_DIM,
9
+ COL_SLICE,
10
+ COL_START,
11
+ COL_STOP,
12
+ COL_URL,
13
+ ChunkRef,
14
+ DimSlice,
15
+ )
16
+
17
+ GeoTransform = tuple[float, float, float, float, float, float]
18
+
19
+
20
+ def read_zarr_geotransform(
21
+ store_url: str,
22
+ *,
23
+ store_kwargs: dict[str, Any] | None = None,
24
+ ) -> GeoTransform:
25
+ # TODO: also support spatial zarr conventions.
26
+ group = zarr.open_group(
27
+ store=store_url,
28
+ mode="r",
29
+ storage_options=store_kwargs or None,
30
+ )
31
+ if "spatial_ref" not in group:
32
+ msg = f"Zarr store {store_url} is missing spatial_ref."
33
+ raise ValueError(msg)
34
+ spatial_ref = group["spatial_ref"]
35
+
36
+ raw_transform = spatial_ref.attrs.get("GeoTransform")
37
+ if raw_transform is None:
38
+ msg = f"Zarr store {store_url} is missing spatial_ref.GeoTransform."
39
+ raise ValueError(msg)
40
+
41
+ return parse_geotransform(raw_transform)
42
+
43
+
44
+ def parse_geotransform(raw_transform: Any) -> GeoTransform:
45
+ parts: Sequence[Any]
46
+ if isinstance(raw_transform, str):
47
+ parts = raw_transform.split()
48
+ elif isinstance(raw_transform, Sequence):
49
+ parts = raw_transform
50
+ else:
51
+ msg = f"Unsupported GeoTransform value: {raw_transform!r}"
52
+ raise TypeError(msg)
53
+
54
+ if len(parts) != 6:
55
+ msg = f"Expected 6 GeoTransform coefficients, got {len(parts)}."
56
+ raise ValueError(msg)
57
+
58
+ origin_x = _coerce_float(parts[0])
59
+ pixel_width = _coerce_float(parts[1])
60
+ row_rotation = _coerce_float(parts[2])
61
+ origin_y = _coerce_float(parts[3])
62
+ column_rotation = _coerce_float(parts[4])
63
+ pixel_height = _coerce_float(parts[5])
64
+ return origin_x, pixel_width, row_rotation, origin_y, column_rotation, pixel_height
65
+
66
+
67
+ def chunk_polygon(chunk: ChunkRef, transform: GeoTransform) -> BaseGeometry:
68
+ from shapely.geometry import Polygon # type: ignore[import-untyped]
69
+
70
+ x_start, x_stop = _dim_interval(chunk["slice"], "x")
71
+ y_start, y_stop = _dim_interval(chunk["slice"], "y")
72
+ x0, y0 = pixel_to_world(transform, x_start, y_start)
73
+ x1, y1 = pixel_to_world(transform, x_stop, y_start)
74
+ x2, y2 = pixel_to_world(transform, x_stop, y_stop)
75
+ x3, y3 = pixel_to_world(transform, x_start, y_stop)
76
+ return Polygon([(x0, y0), (x1, y1), (x2, y2), (x3, y3)])
77
+
78
+
79
+ def pixel_to_world(transform: GeoTransform, x: int, y: int) -> tuple[float, float]:
80
+ origin_x, pixel_width, row_rotation, origin_y, column_rotation, pixel_height = transform
81
+ world_x = origin_x + (x * pixel_width) + (y * row_rotation)
82
+ world_y = origin_y + (x * column_rotation) + (y * pixel_height)
83
+ return world_x, world_y
84
+
85
+
86
+ def validate_aoi_geometry(aoi: BaseGeometry) -> BaseGeometry:
87
+ if isinstance(aoi, BaseGeometry):
88
+ if aoi.is_empty:
89
+ msg = "AOI geometry is empty."
90
+ raise ValueError(msg)
91
+ return aoi
92
+
93
+ msg = "AOI must be a shapely BaseGeometry. " f"Got {type(aoi).__name__}."
94
+ raise TypeError(msg)
95
+
96
+
97
+ def chunk_from_row(row: dict[str, Any]) -> ChunkRef:
98
+ return ChunkRef(
99
+ url=str(row[COL_URL]),
100
+ slice=_coerce_slice(row[COL_SLICE]),
101
+ )
102
+
103
+
104
+ def _dim_interval(parts: list[DimSlice], dim: str) -> tuple[int, int]:
105
+ for part in parts:
106
+ if part[COL_DIM] == dim:
107
+ return int(part[COL_START]), int(part[COL_STOP])
108
+ msg = f"Chunk is missing the required {dim!r} dimension."
109
+ raise KeyError(msg)
110
+
111
+
112
+ def _coerce_slice(raw: Any) -> list[DimSlice]:
113
+ if not isinstance(raw, list):
114
+ msg = f"Expected slice to be a list, got {type(raw).__name__}."
115
+ raise TypeError(msg)
116
+
117
+ parts: list[DimSlice] = []
118
+ for part in raw:
119
+ if not isinstance(part, dict):
120
+ msg = f"Expected each slice entry to be a dict, got {type(part).__name__}."
121
+ raise TypeError(msg)
122
+ parts.append(
123
+ DimSlice(
124
+ dim=str(part[COL_DIM]),
125
+ start=_coerce_int(part[COL_START]),
126
+ stop=_coerce_int(part[COL_STOP]),
127
+ )
128
+ )
129
+ return parts
130
+
131
+
132
+ def _coerce_float(value: Any) -> float:
133
+ return float(str(value))
134
+
135
+
136
+ def _coerce_int(value: Any) -> int:
137
+ return int(str(value))
rayzin/metrics.py ADDED
@@ -0,0 +1,156 @@
1
+ from typing import Protocol, runtime_checkable
2
+
3
+ import numpy as np
4
+ import pyarrow as pa # type: ignore[import-untyped]
5
+
6
+ from rayzin.enums import MetricType
7
+ from rayzin.manifest.schema import LOWER_BOUND_SCHEMA, LowerBoundTable, ManifestTable
8
+ from rayzin.types import (
9
+ COL_CENTROID,
10
+ COL_COUNT,
11
+ COL_RADIUS,
12
+ COL_SLICE,
13
+ COL_URL,
14
+ Float32Array,
15
+ )
16
+
17
+
18
+ @runtime_checkable
19
+ class Metric(Protocol):
20
+ def distance(self, a: Float32Array, b: Float32Array) -> float: ...
21
+ def pairwise(self, vectors: Float32Array, query: Float32Array) -> Float32Array: ...
22
+ def lower_bound(self, query: Float32Array, centroid: Float32Array, radius: float) -> float: ...
23
+
24
+
25
+ class EuclideanMetric:
26
+ """Squared L2 distance, matching FAISS IndexFlatL2 semantics."""
27
+
28
+ def distance(self, a: Float32Array, b: Float32Array) -> float:
29
+ delta = np.asarray(a - b, dtype=np.float32)
30
+ return float(np.dot(delta, delta))
31
+
32
+ def pairwise(self, vectors: Float32Array, query: Float32Array) -> Float32Array:
33
+ query_array = np.asarray(query, dtype=np.float32)
34
+ if query_array.ndim == 1:
35
+ deltas = np.asarray(vectors - query_array[None, :], dtype=np.float32)
36
+ return np.asarray(np.sum(deltas * deltas, axis=1, dtype=np.float32), dtype=np.float32)
37
+
38
+ deltas = np.asarray(vectors[None, :, :] - query_array[:, None, :], dtype=np.float32)
39
+ return np.asarray(np.sum(deltas * deltas, axis=2, dtype=np.float32), dtype=np.float32)
40
+
41
+ def lower_bound(self, query: Float32Array, centroid: Float32Array, radius: float) -> float:
42
+ centroid_distance = float(np.linalg.norm(query - centroid))
43
+ return float(max(0.0, centroid_distance - radius) ** 2)
44
+
45
+
46
+ class CosineMetric:
47
+ """Cosine distance implemented as 1 - inner_product(normalized(a), normalized(b))."""
48
+
49
+ def distance(self, a: Float32Array, b: Float32Array) -> float:
50
+ return float(1.0 - np.dot(_normalize(a), _normalize(b)))
51
+
52
+ def pairwise(self, vectors: Float32Array, query: Float32Array) -> Float32Array:
53
+ normalized_vectors = _normalize(vectors)
54
+ normalized_query = _normalize(query)
55
+ if normalized_query.ndim == 1:
56
+ scores = normalized_vectors @ normalized_query
57
+ return np.asarray(1.0 - scores, dtype=np.float32)
58
+
59
+ scores = normalized_query @ normalized_vectors.T
60
+ return np.asarray(1.0 - scores, dtype=np.float32)
61
+
62
+ def lower_bound(self, query: Float32Array, centroid: Float32Array, radius: float) -> float:
63
+ return max(0.0, self.distance(query, centroid) - radius)
64
+
65
+
66
+ EUCLIDEAN = EuclideanMetric()
67
+ COSINE = CosineMetric()
68
+
69
+
70
+ def make_metric(metric_type: MetricType) -> Metric:
71
+ if metric_type == MetricType.EUCLIDEAN:
72
+ return EuclideanMetric()
73
+ if metric_type == MetricType.COSINE:
74
+ return CosineMetric()
75
+ raise ValueError(metric_type)
76
+
77
+
78
+ def add_lower_bounds_fn(
79
+ batch: ManifestTable,
80
+ queries: Float32Array,
81
+ metric_type: str,
82
+ ) -> LowerBoundTable:
83
+ return _add_lower_bounds(batch, queries=queries, metric_type=metric_type)
84
+
85
+
86
+ def _add_lower_bounds(
87
+ batch: ManifestTable,
88
+ queries: Float32Array,
89
+ metric_type: str,
90
+ ) -> LowerBoundTable:
91
+ query_matrix = np.asarray(queries, dtype=np.float32)
92
+ if query_matrix.ndim != 2:
93
+ msg = f"Expected queries to have shape (nq, d), got {query_matrix.shape!r}."
94
+ raise ValueError(msg)
95
+
96
+ centroids = np.asarray(batch.column(COL_CENTROID).to_pylist(), dtype=np.float32)
97
+ radii = np.asarray(batch.column(COL_RADIUS).to_pylist(), dtype=np.float32)
98
+ lower_bounds = _lower_bounds(
99
+ query_matrix,
100
+ centroids,
101
+ radii,
102
+ metric_type=MetricType(metric_type),
103
+ )
104
+ min_lower_bounds = np.asarray(np.min(lower_bounds, axis=1), dtype=np.float32)
105
+
106
+ return pa.Table.from_arrays(
107
+ [
108
+ batch.column(COL_URL),
109
+ batch.column(COL_SLICE),
110
+ batch.column(COL_COUNT),
111
+ batch.column(COL_CENTROID),
112
+ batch.column(COL_RADIUS),
113
+ pa.array(lower_bounds.tolist(), type=pa.list_(pa.float32())),
114
+ pa.array(min_lower_bounds.tolist(), type=pa.float32()),
115
+ ],
116
+ schema=LOWER_BOUND_SCHEMA,
117
+ )
118
+
119
+
120
+ def _lower_bounds(
121
+ queries: Float32Array,
122
+ centroids: Float32Array,
123
+ radii: Float32Array,
124
+ *,
125
+ metric_type: MetricType,
126
+ ) -> Float32Array:
127
+ if metric_type == MetricType.EUCLIDEAN:
128
+ deltas = np.asarray(
129
+ centroids[:, None, :] - queries[None, :, :],
130
+ dtype=np.float32,
131
+ )
132
+ centroid_distances = np.linalg.norm(deltas, axis=2)
133
+ return np.asarray(
134
+ np.square(np.maximum(0.0, centroid_distances - radii[:, None])),
135
+ dtype=np.float32,
136
+ )
137
+
138
+ metric = make_metric(metric_type)
139
+ distances = np.asarray(metric.pairwise(centroids, queries), dtype=np.float32).T
140
+ if distances.ndim != 2:
141
+ msg = f"Expected pairwise distances to have shape (n_rows, nq), got {distances.shape!r}."
142
+ raise ValueError(msg)
143
+ return np.asarray(np.maximum(0.0, distances - radii[:, None]), dtype=np.float32)
144
+
145
+
146
+ def _normalize(vectors: Float32Array) -> Float32Array:
147
+ array = np.asarray(vectors, dtype=np.float32)
148
+ if array.ndim == 1:
149
+ denominator = np.float32(np.linalg.norm(array))
150
+ if denominator < np.finfo(np.float32).eps:
151
+ denominator = np.float32(np.finfo(np.float32).eps)
152
+ return np.asarray(array / denominator, dtype=np.float32)
153
+
154
+ norms = np.linalg.norm(array, axis=1, keepdims=True)
155
+ safe_norms = np.maximum(norms, np.finfo(np.float32).eps)
156
+ return np.asarray(array / safe_norms, dtype=np.float32)