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 +26 -0
- rayzin/enums.py +17 -0
- rayzin/manifest/__init__.py +7 -0
- rayzin/manifest/build.py +166 -0
- rayzin/manifest/filtering.py +60 -0
- rayzin/manifest/schema.py +77 -0
- rayzin/manifest/spatial.py +137 -0
- rayzin/metrics.py +156 -0
- rayzin/pipeline.py +234 -0
- rayzin/readers/__init__.py +21 -0
- rayzin/readers/cog_reader.py +17 -0
- rayzin/readers/protocol.py +17 -0
- rayzin/readers/zarr_layout.py +65 -0
- rayzin/readers/zarr_reader.py +85 -0
- rayzin/search/__init__.py +20 -0
- rayzin/search/backends/__init__.py +43 -0
- rayzin/search/backends/faiss.py +292 -0
- rayzin/search/backends/numpy.py +174 -0
- rayzin/search/backends/protocols.py +42 -0
- rayzin/search/block_searcher.py +211 -0
- rayzin/search/heap_actor.py +25 -0
- rayzin/search/results.py +37 -0
- rayzin/types.py +83 -0
- rayzin-0.0.1.dist-info/METADATA +112 -0
- rayzin-0.0.1.dist-info/RECORD +26 -0
- rayzin-0.0.1.dist-info/WHEEL +4 -0
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"
|
rayzin/manifest/build.py
ADDED
|
@@ -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)
|