mirror of
https://github.com/lancedb/lancedb.git
synced 2026-09-03 12:08:52 +00:00
Compare commits
2 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| bc3837c4fe | |||
| 2a4f4f338b |
Generated
-1
@@ -5442,7 +5442,6 @@ dependencies = [
|
||||
"pprof 0.14.1",
|
||||
"rand 0.9.5",
|
||||
"random_word",
|
||||
"rayon",
|
||||
"regex",
|
||||
"reqwest 0.12.28",
|
||||
"rstest",
|
||||
|
||||
@@ -60,7 +60,6 @@ moka = { version = "0.12", features = ["future"] }
|
||||
object_store = "0.13.2"
|
||||
pin-project = "1.0.7"
|
||||
rand = "0.9"
|
||||
rayon = "1"
|
||||
snafu = "0.8"
|
||||
url = "2"
|
||||
num-traits = "0.2"
|
||||
|
||||
@@ -261,6 +261,7 @@ class Table:
|
||||
def name(self) -> str: ...
|
||||
def __repr__(self) -> str: ...
|
||||
def is_open(self) -> bool: ...
|
||||
def _is_native(self) -> bool: ...
|
||||
def close(self) -> None: ...
|
||||
async def schema(self) -> pa.Schema: ...
|
||||
async def add(
|
||||
|
||||
@@ -314,8 +314,7 @@ class HnswPq:
|
||||
m: int = 20
|
||||
ef_construction: int = 300
|
||||
target_partition_size: Optional[int] = None
|
||||
# Name of the accelerator (e.g. "cuda") to use for IVF training. When set,
|
||||
# create_index() dispatches to pylance to build the index on the accelerator.
|
||||
# Reserved for future accelerator support. create_index() currently raises if set.
|
||||
accelerator: Optional[str] = None
|
||||
|
||||
|
||||
@@ -422,8 +421,7 @@ class HnswSq:
|
||||
m: int = 20
|
||||
ef_construction: int = 300
|
||||
target_partition_size: Optional[int] = None
|
||||
# Name of the accelerator (e.g. "cuda") to use for IVF training. When set,
|
||||
# create_index() dispatches to pylance to build the index on the accelerator.
|
||||
# Reserved for future accelerator support. create_index() currently raises if set.
|
||||
accelerator: Optional[str] = None
|
||||
|
||||
|
||||
@@ -618,8 +616,7 @@ class IvfFlat:
|
||||
max_iterations: int = 50
|
||||
sample_rate: int = 256
|
||||
target_partition_size: Optional[int] = None
|
||||
# Name of the accelerator (e.g. "cuda") to use for IVF training. When set,
|
||||
# create_index() dispatches to pylance to build the index on the accelerator.
|
||||
# Reserved for future accelerator support. create_index() currently raises if set.
|
||||
accelerator: Optional[str] = None
|
||||
|
||||
|
||||
@@ -651,8 +648,7 @@ class IvfSq:
|
||||
max_iterations: int = 50
|
||||
sample_rate: int = 256
|
||||
target_partition_size: Optional[int] = None
|
||||
# Name of the accelerator (e.g. "cuda") to use for IVF training. When set,
|
||||
# create_index() dispatches to pylance to build the index on the accelerator.
|
||||
# Reserved for future accelerator support. create_index() currently raises if set.
|
||||
accelerator: Optional[str] = None
|
||||
|
||||
|
||||
@@ -769,13 +765,6 @@ class IvfPq:
|
||||
|
||||
The default value is 256.
|
||||
|
||||
seed: int, optional
|
||||
Seed used for deterministic sampling and training. Given identical data in
|
||||
the same row order and identical index parameters, the same seed produces
|
||||
the same IVF centroids and PQ codebook. This option is supported for local
|
||||
CPU index builds; remote and accelerator builds reject it explicitly. If
|
||||
omitted, training remains random.
|
||||
|
||||
target_partition_size: int, default is 8192
|
||||
|
||||
The target size of each partition.
|
||||
@@ -790,18 +779,11 @@ class IvfPq:
|
||||
num_bits: int = 8
|
||||
max_iterations: int = 50
|
||||
sample_rate: int = 256
|
||||
seed: Optional[int] = None
|
||||
target_partition_size: Optional[int] = None
|
||||
# Name of the accelerator (e.g. "cuda") to use for IVF training. When set,
|
||||
# Name of the accelerator ("cuda" or "mps") to use for IVF training. When set,
|
||||
# create_index() dispatches to pylance to build the index on the accelerator.
|
||||
accelerator: Optional[str] = None
|
||||
|
||||
def __post_init__(self):
|
||||
if self.seed is not None and self.accelerator is not None:
|
||||
raise ValueError(
|
||||
"IvfPq seed is not supported with accelerator-based index training"
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class IvfRq:
|
||||
@@ -854,8 +836,7 @@ class IvfRq:
|
||||
max_iterations: int = 50
|
||||
sample_rate: int = 256
|
||||
target_partition_size: Optional[int] = None
|
||||
# Name of the accelerator (e.g. "cuda") to use for IVF training. When set,
|
||||
# create_index() dispatches to pylance to build the index on the accelerator.
|
||||
# Reserved for future accelerator support. create_index() currently raises if set.
|
||||
accelerator: Optional[str] = None
|
||||
|
||||
|
||||
|
||||
@@ -68,6 +68,14 @@ from ..table import AsyncTable, BlobMode, Branches, IndexStatistics, Query, Tabl
|
||||
from ..types import BaseTokenizerType
|
||||
|
||||
|
||||
def _reject_index_accelerator(
|
||||
config: Optional[IndexConfigType] = None,
|
||||
accelerator: Optional[str] = None,
|
||||
) -> None:
|
||||
if accelerator is not None or getattr(config, "accelerator", None) is not None:
|
||||
raise ValueError("Index accelerators are not supported on LanceDB Cloud.")
|
||||
|
||||
|
||||
class RemoteTable(Table):
|
||||
def __init__(
|
||||
self,
|
||||
@@ -457,6 +465,8 @@ class RemoteTable(Table):
|
||||
... "l2", vector_column_name="vector"
|
||||
... )
|
||||
"""
|
||||
_reject_index_accelerator(config, accelerator)
|
||||
|
||||
# Detect whether this is a legacy API call
|
||||
is_legacy = self._is_legacy_create_index_call(
|
||||
metric,
|
||||
@@ -484,12 +494,6 @@ class RemoteTable(Table):
|
||||
|
||||
column = vector_column_name
|
||||
|
||||
if accelerator is not None:
|
||||
logging.warning(
|
||||
"GPU accelerator is not yet supported on LanceDB cloud."
|
||||
"If you have 100M+ vectors to index,"
|
||||
"please contact us at contact@lancedb.com"
|
||||
)
|
||||
if replace is not None:
|
||||
logging.warning(
|
||||
"replace is not supported on LanceDB cloud."
|
||||
@@ -557,6 +561,8 @@ class RemoteTable(Table):
|
||||
The job may already be complete when returned; callers must not assume
|
||||
the index exists until :meth:`Job.wait` returns.
|
||||
"""
|
||||
_reject_index_accelerator(config)
|
||||
|
||||
return Job(
|
||||
LOOP.run(
|
||||
self._table.create_index_async(
|
||||
|
||||
@@ -214,6 +214,45 @@ IndexConfigType = Union[
|
||||
# Known distance metrics for legacy API detection
|
||||
KNOWN_METRICS = {"l2", "cosine", "dot", "hamming"}
|
||||
|
||||
_PYLANCE_ACCELERATED_INDEX_TYPE = "IVF_PQ"
|
||||
|
||||
|
||||
def _pylance_accelerated_index_options(
|
||||
config: IndexConfigType,
|
||||
*,
|
||||
accelerator: Optional[str] = None,
|
||||
index_type: Optional[str] = None,
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""Translate an accelerated vector config into PyLance index options."""
|
||||
if accelerator is None:
|
||||
accelerator = getattr(config, "accelerator", None)
|
||||
if accelerator is None:
|
||||
return None
|
||||
|
||||
if index_type is None:
|
||||
index_type = (
|
||||
_PYLANCE_ACCELERATED_INDEX_TYPE
|
||||
if isinstance(config, IvfPq)
|
||||
else type(config).__name__
|
||||
)
|
||||
if index_type.upper() != _PYLANCE_ACCELERATED_INDEX_TYPE:
|
||||
raise ValueError(
|
||||
f"Index type {index_type} does not support an accelerator; "
|
||||
f"only {_PYLANCE_ACCELERATED_INDEX_TYPE} supports acceleration"
|
||||
)
|
||||
|
||||
return {
|
||||
"index_type": index_type,
|
||||
"metric": getattr(config, "distance_type", "l2"),
|
||||
"num_partitions": getattr(config, "num_partitions", None),
|
||||
"num_sub_vectors": getattr(config, "num_sub_vectors", None),
|
||||
"accelerator": accelerator,
|
||||
"num_bits": getattr(config, "num_bits", 8),
|
||||
"m": getattr(config, "m", 20),
|
||||
"ef_construction": getattr(config, "ef_construction", 300),
|
||||
"target_partition_size": getattr(config, "target_partition_size", None),
|
||||
}
|
||||
|
||||
|
||||
def _into_pyarrow_reader(
|
||||
data, schema: Optional[pa.Schema] = None
|
||||
@@ -2737,20 +2776,17 @@ class LanceTable(Table):
|
||||
)
|
||||
|
||||
# Handle accelerator through pylance
|
||||
if accelerator is not None:
|
||||
accelerated_options = _pylance_accelerated_index_options(
|
||||
config, accelerator=accelerator, index_type=index_type
|
||||
)
|
||||
if accelerated_options is not None:
|
||||
self.to_lance().create_index(
|
||||
column=column,
|
||||
index_type=index_type,
|
||||
metric=metric,
|
||||
num_partitions=num_partitions,
|
||||
num_sub_vectors=num_sub_vectors,
|
||||
replace=replace,
|
||||
accelerator=accelerator,
|
||||
index_cache_size=index_cache_size,
|
||||
num_bits=num_bits,
|
||||
m=m,
|
||||
ef_construction=ef_construction,
|
||||
target_partition_size=target_partition_size,
|
||||
name=name,
|
||||
train=train,
|
||||
**accelerated_options,
|
||||
)
|
||||
self.checkout_latest()
|
||||
return
|
||||
@@ -2758,44 +2794,21 @@ class LanceTable(Table):
|
||||
# New API: metric is the column name
|
||||
column = metric
|
||||
|
||||
# Check if config has accelerator set and dispatch to pylance
|
||||
if config is not None and hasattr(config, "accelerator"):
|
||||
acc = getattr(config, "accelerator", None)
|
||||
if acc is not None:
|
||||
if isinstance(config, IvfPq) and config.seed is not None:
|
||||
raise ValueError(
|
||||
"IvfPq seed is not supported with accelerator-based "
|
||||
"index training"
|
||||
)
|
||||
# Dispatch to pylance for GPU acceleration
|
||||
index_type_map = {
|
||||
"IvfFlat": "IVF_FLAT",
|
||||
"IvfSq": "IVF_SQ",
|
||||
"IvfPq": "IVF_PQ",
|
||||
"IvfRq": "IVF_RQ",
|
||||
"HnswPq": "IVF_HNSW_PQ",
|
||||
"HnswSq": "IVF_HNSW_SQ",
|
||||
}
|
||||
cfg_type = type(config).__name__
|
||||
lance_index_type = index_type_map.get(cfg_type, "IVF_PQ")
|
||||
|
||||
self.to_lance().create_index(
|
||||
column=column,
|
||||
index_type=lance_index_type,
|
||||
metric=getattr(config, "distance_type", "l2"),
|
||||
num_partitions=getattr(config, "num_partitions", None),
|
||||
num_sub_vectors=getattr(config, "num_sub_vectors", None),
|
||||
replace=replace,
|
||||
accelerator=acc,
|
||||
num_bits=getattr(config, "num_bits", 8),
|
||||
m=getattr(config, "m", 20),
|
||||
ef_construction=getattr(config, "ef_construction", 300),
|
||||
target_partition_size=getattr(
|
||||
config, "target_partition_size", None
|
||||
),
|
||||
)
|
||||
self.checkout_latest()
|
||||
return
|
||||
accelerated_options = (
|
||||
_pylance_accelerated_index_options(config)
|
||||
if config is not None
|
||||
else None
|
||||
)
|
||||
if accelerated_options is not None:
|
||||
self.to_lance().create_index(
|
||||
column=column,
|
||||
replace=replace,
|
||||
name=name,
|
||||
train=train,
|
||||
**accelerated_options,
|
||||
)
|
||||
self.checkout_latest()
|
||||
return
|
||||
|
||||
return LOOP.run(
|
||||
self._table.create_index(
|
||||
@@ -2823,6 +2836,11 @@ class LanceTable(Table):
|
||||
The job may already be complete when returned; callers must not assume
|
||||
the index exists until :meth:`Job.wait` returns.
|
||||
"""
|
||||
if _pylance_accelerated_index_options(config) is not None:
|
||||
raise ValueError(
|
||||
"Accelerated index creation does not support create_index_async; "
|
||||
"use create_index instead."
|
||||
)
|
||||
return Job(
|
||||
LOOP.run(
|
||||
self._table.create_index_async(
|
||||
@@ -4835,6 +4853,7 @@ class AsyncTable:
|
||||
config: Optional[
|
||||
Union[
|
||||
IvfFlat,
|
||||
IvfSq,
|
||||
IvfPq,
|
||||
IvfRq,
|
||||
HnswPq,
|
||||
@@ -4905,6 +4924,23 @@ class AsyncTable:
|
||||
" BTree, Bitmap, LabelList, Fm, or FTS, but got "
|
||||
+ str(type(config))
|
||||
)
|
||||
accelerated_options = (
|
||||
_pylance_accelerated_index_options(config) if config is not None else None
|
||||
)
|
||||
if accelerated_options is not None:
|
||||
if not self._inner._is_native():
|
||||
raise ValueError("GPU accelerator is not supported on LanceDB Cloud.")
|
||||
dataset = await self.to_lance()
|
||||
await asyncio.to_thread(
|
||||
dataset.create_index,
|
||||
column=column,
|
||||
replace=True if replace is None else replace,
|
||||
name=name,
|
||||
train=train,
|
||||
**accelerated_options,
|
||||
)
|
||||
await self.checkout_latest()
|
||||
return
|
||||
try:
|
||||
await self._inner.create_index(
|
||||
column,
|
||||
@@ -4931,6 +4967,7 @@ class AsyncTable:
|
||||
config: Optional[
|
||||
Union[
|
||||
IvfFlat,
|
||||
IvfSq,
|
||||
IvfPq,
|
||||
IvfRq,
|
||||
HnswPq,
|
||||
@@ -4953,6 +4990,11 @@ class AsyncTable:
|
||||
be complete when returned; callers must not assume the index exists
|
||||
until :meth:`AsyncJob.wait` resolves.
|
||||
"""
|
||||
if config is not None and _pylance_accelerated_index_options(config):
|
||||
raise ValueError(
|
||||
"Accelerated index creation does not support create_index_async; "
|
||||
"use create_index instead."
|
||||
)
|
||||
job = await self._inner.create_index_async(
|
||||
column,
|
||||
index=config,
|
||||
|
||||
@@ -26,7 +26,7 @@ from lancedb.index import (
|
||||
HnswFlat,
|
||||
FTS,
|
||||
)
|
||||
from lancedb.table import IndexStatistics, LanceTable
|
||||
from lancedb.table import IndexStatistics
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
@@ -375,7 +375,7 @@ async def test_create_vector_index(some_table: AsyncTable):
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_4bit_ivfpq_index(some_table: AsyncTable):
|
||||
# Can create
|
||||
await some_table.create_index("vector", config=IvfPq(num_bits=4, seed=42))
|
||||
await some_table.create_index("vector", config=IvfPq(num_bits=4))
|
||||
# Can recreate if replace=True
|
||||
await some_table.create_index("vector", config=IvfPq(num_bits=4), replace=True)
|
||||
# Can't recreate if replace=False
|
||||
@@ -395,16 +395,6 @@ async def test_create_4bit_ivfpq_index(some_table: AsyncTable):
|
||||
assert stats.num_indices == 1
|
||||
|
||||
|
||||
def test_seeded_ivfpq_rejects_accelerator():
|
||||
with pytest.raises(ValueError, match="seed is not supported with accelerator"):
|
||||
IvfPq(seed=42, accelerator="cuda")
|
||||
|
||||
config = IvfPq(seed=42)
|
||||
config.accelerator = "cuda"
|
||||
with pytest.raises(ValueError, match="seed is not supported with accelerator"):
|
||||
object.__new__(LanceTable).create_index("vector", config=config)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_ivfrq_index(some_table: AsyncTable):
|
||||
await some_table.create_index("vector", config=IvfRq(num_bits=1))
|
||||
|
||||
@@ -875,6 +875,25 @@ def test_remote_create_index_async_returns_job():
|
||||
job.cancel()
|
||||
|
||||
|
||||
def test_remote_create_index_rejects_accelerator():
|
||||
from lancedb.index import IvfPq
|
||||
from lancedb.remote.table import RemoteTable
|
||||
|
||||
inner = MagicMock()
|
||||
inner.name = "test"
|
||||
table = RemoteTable(inner, "dev")
|
||||
|
||||
with pytest.raises(ValueError, match="not supported on LanceDB Cloud"):
|
||||
table.create_index(accelerator="mps")
|
||||
with pytest.raises(ValueError, match="not supported on LanceDB Cloud"):
|
||||
table.create_index("vector", config=IvfPq(accelerator="mps"))
|
||||
with pytest.raises(ValueError, match="not supported on LanceDB Cloud"):
|
||||
table.create_index_async("vector", config=IvfPq(accelerator="mps"))
|
||||
|
||||
inner.create_index.assert_not_called()
|
||||
inner.create_index_async.assert_not_called()
|
||||
|
||||
|
||||
def test_remote_job_wait_raises_on_failure():
|
||||
from lancedb.exceptions import JobFailedError
|
||||
from lancedb.index import BTree
|
||||
|
||||
@@ -10,11 +10,21 @@ from concurrent.futures import ThreadPoolExecutor
|
||||
from datetime import date, datetime, timedelta
|
||||
from time import sleep
|
||||
from typing import List
|
||||
from unittest.mock import patch
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import lancedb
|
||||
from lancedb.dependencies import _PANDAS_AVAILABLE
|
||||
from lancedb.index import BTree, FTS, HnswFlat, HnswPq, HnswSq, IvfPq
|
||||
from lancedb.index import (
|
||||
BTree,
|
||||
FTS,
|
||||
HnswFlat,
|
||||
HnswPq,
|
||||
HnswSq,
|
||||
IvfFlat,
|
||||
IvfPq,
|
||||
IvfRq,
|
||||
IvfSq,
|
||||
)
|
||||
import numpy as np
|
||||
import polars as pl
|
||||
import pyarrow as pa
|
||||
@@ -25,7 +35,7 @@ from lancedb.db import AsyncConnection, DBConnection
|
||||
from lancedb.embeddings import EmbeddingFunctionConfig, EmbeddingFunctionRegistry
|
||||
from lancedb.expr import col, lit
|
||||
from lancedb.pydantic import LanceModel, Vector
|
||||
from lancedb.table import LanceTable
|
||||
from lancedb.table import AsyncTable, LanceTable
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
@@ -1412,6 +1422,174 @@ def test_create_index_async_returns_done_job(mem_db: DBConnection):
|
||||
job.cancel()
|
||||
|
||||
|
||||
def test_create_index_dispatches_mps_to_pylance(mem_db: DBConnection):
|
||||
table = mem_db.create_table(
|
||||
"mps_sync",
|
||||
data=[
|
||||
{"vector": [3.1, 4.1]},
|
||||
{"vector": [5.9, 26.5]},
|
||||
],
|
||||
)
|
||||
dataset = MagicMock()
|
||||
|
||||
with (
|
||||
patch.object(table, "to_lance", return_value=dataset),
|
||||
patch.object(table, "checkout_latest") as checkout_latest,
|
||||
):
|
||||
with pytest.warns(DeprecationWarning, match="create_index"):
|
||||
table.create_index(
|
||||
metric="cosine",
|
||||
num_partitions=4,
|
||||
num_sub_vectors=2,
|
||||
accelerator="mps",
|
||||
replace=False,
|
||||
name="vector_mps",
|
||||
)
|
||||
|
||||
dataset.create_index.assert_called_once_with(
|
||||
column="vector",
|
||||
replace=False,
|
||||
index_cache_size=None,
|
||||
name="vector_mps",
|
||||
train=True,
|
||||
index_type="IVF_PQ",
|
||||
metric="cosine",
|
||||
num_partitions=4,
|
||||
num_sub_vectors=2,
|
||||
accelerator="mps",
|
||||
num_bits=8,
|
||||
m=20,
|
||||
ef_construction=300,
|
||||
target_partition_size=None,
|
||||
)
|
||||
checkout_latest.assert_called_once_with()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"config",
|
||||
[
|
||||
IvfFlat(accelerator="mps"),
|
||||
IvfSq(accelerator="mps"),
|
||||
IvfRq(accelerator="mps"),
|
||||
HnswPq(accelerator="mps"),
|
||||
HnswSq(accelerator="mps"),
|
||||
],
|
||||
)
|
||||
def test_create_index_rejects_unsupported_accelerated_format(
|
||||
mem_db: DBConnection, config
|
||||
):
|
||||
table = mem_db.create_table(
|
||||
"unsupported_accelerator",
|
||||
data=[{"vector": [3.1, 4.1]}, {"vector": [5.9, 26.5]}],
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(table, "to_lance") as to_lance,
|
||||
pytest.raises(ValueError, match="only IVF_PQ supports acceleration"),
|
||||
):
|
||||
table.create_index("vector", config=config)
|
||||
|
||||
to_lance.assert_not_called()
|
||||
|
||||
|
||||
def test_legacy_create_index_rejects_unsupported_accelerated_format(
|
||||
mem_db: DBConnection,
|
||||
):
|
||||
table = mem_db.create_table(
|
||||
"unsupported_legacy_accelerator",
|
||||
data=[{"vector": [3.1, 4.1]}, {"vector": [5.9, 26.5]}],
|
||||
)
|
||||
|
||||
with (
|
||||
pytest.warns(DeprecationWarning, match="create_index"),
|
||||
patch.object(table, "to_lance") as to_lance,
|
||||
pytest.raises(ValueError, match="only IVF_PQ supports acceleration"),
|
||||
):
|
||||
table.create_index(index_type="IVF_FLAT", accelerator="mps")
|
||||
|
||||
to_lance.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_create_index_dispatches_mps_to_pylance():
|
||||
inner = MagicMock()
|
||||
inner._is_native.return_value = True
|
||||
inner.checkout_latest = AsyncMock()
|
||||
table = AsyncTable(inner)
|
||||
dataset = MagicMock()
|
||||
|
||||
with patch.object(table, "to_lance", AsyncMock(return_value=dataset)):
|
||||
await table.create_index(
|
||||
"vector",
|
||||
config=IvfPq(
|
||||
distance_type="cosine",
|
||||
num_partitions=4,
|
||||
num_sub_vectors=2,
|
||||
accelerator="mps",
|
||||
),
|
||||
name="vector_mps",
|
||||
)
|
||||
|
||||
dataset.create_index.assert_called_once_with(
|
||||
column="vector",
|
||||
replace=True,
|
||||
name="vector_mps",
|
||||
train=True,
|
||||
index_type="IVF_PQ",
|
||||
metric="cosine",
|
||||
num_partitions=4,
|
||||
num_sub_vectors=2,
|
||||
accelerator="mps",
|
||||
num_bits=8,
|
||||
m=20,
|
||||
ef_construction=300,
|
||||
target_partition_size=None,
|
||||
)
|
||||
inner.create_index.assert_not_called()
|
||||
inner.checkout_latest.assert_awaited_once_with()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_create_index_rejects_unsupported_accelerated_format():
|
||||
inner = MagicMock()
|
||||
inner._is_native.return_value = True
|
||||
table = AsyncTable(inner)
|
||||
|
||||
with (
|
||||
patch.object(table, "to_lance", AsyncMock()) as to_lance,
|
||||
pytest.raises(ValueError, match="only IVF_PQ supports acceleration"),
|
||||
):
|
||||
await table.create_index("vector", config=IvfFlat(accelerator="mps"))
|
||||
|
||||
to_lance.assert_not_awaited()
|
||||
inner.create_index.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_background_index_rejects_accelerator():
|
||||
inner = MagicMock()
|
||||
inner.create_index_async = AsyncMock()
|
||||
table = AsyncTable(inner)
|
||||
|
||||
with pytest.raises(ValueError, match="Accelerated index creation does not support"):
|
||||
await table.create_index_async("vector", config=IvfPq(accelerator="mps"))
|
||||
|
||||
inner.create_index_async.assert_not_awaited()
|
||||
|
||||
|
||||
def test_background_index_rejects_accelerator(mem_db: DBConnection):
|
||||
table = mem_db.create_table(
|
||||
"mps_background",
|
||||
data=[
|
||||
{"vector": [3.1, 4.1]},
|
||||
{"vector": [5.9, 26.5]},
|
||||
],
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="Accelerated index creation does not support"):
|
||||
table.create_index_async("vector", config=IvfPq(accelerator="mps"))
|
||||
|
||||
|
||||
@patch("lancedb.table.AsyncTable.create_index")
|
||||
def test_create_index_method(mock_create_index, mem_db: DBConnection):
|
||||
table = mem_db.create_table(
|
||||
|
||||
@@ -90,9 +90,6 @@ pub fn extract_index_params(source: &Option<Bound<'_, PyAny>>) -> PyResult<Lance
|
||||
.max_iterations(params.max_iterations)
|
||||
.sample_rate(params.sample_rate)
|
||||
.num_bits(params.num_bits);
|
||||
if let Some(seed) = params.seed {
|
||||
ivf_pq_builder = ivf_pq_builder.seed(seed);
|
||||
}
|
||||
if let Some(num_partitions) = params.num_partitions {
|
||||
ivf_pq_builder = ivf_pq_builder.num_partitions(num_partitions);
|
||||
}
|
||||
@@ -235,7 +232,6 @@ struct IvfPqParams {
|
||||
num_bits: u32,
|
||||
max_iterations: u32,
|
||||
sample_rate: u32,
|
||||
seed: Option<u64>,
|
||||
target_partition_size: Option<u32>,
|
||||
}
|
||||
|
||||
|
||||
@@ -622,6 +622,10 @@ impl Table {
|
||||
self.inner.is_some()
|
||||
}
|
||||
|
||||
pub fn _is_native(&self) -> PyResult<bool> {
|
||||
Ok(self.inner_ref()?.as_native().is_some())
|
||||
}
|
||||
|
||||
/// Closes the table, releasing any resources associated with it.
|
||||
pub fn close(&mut self) {
|
||||
self.inner.take();
|
||||
|
||||
@@ -61,7 +61,6 @@ futures.workspace = true
|
||||
num-traits.workspace = true
|
||||
url.workspace = true
|
||||
rand.workspace = true
|
||||
rayon.workspace = true
|
||||
regex.workspace = true
|
||||
serde = { version = "^1" }
|
||||
serde_json = { version = "1" }
|
||||
|
||||
@@ -274,8 +274,6 @@ pub struct IvfPqIndexBuilder {
|
||||
pub(crate) sample_rate: u32,
|
||||
pub(crate) max_iterations: u32,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub(crate) seed: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub(crate) target_partition_size: Option<u32>,
|
||||
|
||||
// PQ
|
||||
@@ -294,7 +292,6 @@ impl Default for IvfPqIndexBuilder {
|
||||
num_bits: None,
|
||||
sample_rate: 256,
|
||||
max_iterations: 50,
|
||||
seed: None,
|
||||
target_partition_size: None,
|
||||
}
|
||||
}
|
||||
@@ -304,30 +301,6 @@ impl IvfPqIndexBuilder {
|
||||
impl_distance_type_setter!();
|
||||
impl_ivf_params_setter!();
|
||||
impl_pq_params_setter!();
|
||||
|
||||
/// Use a deterministic seed when sampling and training the IVF and PQ models.
|
||||
///
|
||||
/// Given identical data in the same row order and identical index parameters,
|
||||
/// using the same seed produces the same IVF centroids and PQ codebook. This is
|
||||
/// useful when independently-built tables need reproducible approximate-search
|
||||
/// results.
|
||||
///
|
||||
/// Seeded training is supported by native tables. Remote backends reject this
|
||||
/// option unless they can provide the same deterministic training contract.
|
||||
///
|
||||
/// If no seed is provided, index training uses random sampling and initialization.
|
||||
///
|
||||
/// # Examples
|
||||
///
|
||||
/// ```
|
||||
/// use lancedb::index::vector::IvfPqIndexBuilder;
|
||||
///
|
||||
/// let index = IvfPqIndexBuilder::default().seed(42);
|
||||
/// ```
|
||||
pub fn seed(mut self, seed: u64) -> Self {
|
||||
self.seed = Some(seed);
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn suggested_num_sub_vectors(dim: u32) -> u32 {
|
||||
|
||||
@@ -339,12 +339,6 @@ impl<S: HttpSend> RemoteTable<S> {
|
||||
// Auto is special-cased since it needs schema inspection.
|
||||
let (index_type_str, params) = match &index.index {
|
||||
Index::IvfFlat(p) => ("IVF_FLAT", Some(to_json(p)?)),
|
||||
Index::IvfPq(p) if p.seed.is_some() => {
|
||||
return Err(Error::NotSupported {
|
||||
message: "Deterministic IVF PQ training is not supported on remote tables"
|
||||
.to_string(),
|
||||
});
|
||||
}
|
||||
Index::IvfPq(p) => ("IVF_PQ", Some(to_json(p)?)),
|
||||
Index::IvfSq(p) => ("IVF_SQ", Some(to_json(p)?)),
|
||||
Index::IvfHnswSq(p) => ("IVF_HNSW_SQ", Some(to_json(p)?)),
|
||||
@@ -5163,41 +5157,6 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_seeded_ivf_pq_is_rejected_for_remote_tables() {
|
||||
let table = Table::new_with_handler("my_table", |request| match request.url().path() {
|
||||
"/v1/table/my_table/describe/" => {
|
||||
let schema = Schema::new(vec![Field::new(
|
||||
"vector",
|
||||
DataType::FixedSizeList(
|
||||
Arc::new(Field::new("item", DataType::Float32, true)),
|
||||
8,
|
||||
),
|
||||
false,
|
||||
)]);
|
||||
http::Response::builder()
|
||||
.status(200)
|
||||
.body(describe_response(&schema))
|
||||
.unwrap()
|
||||
}
|
||||
path => panic!("Unexpected request for unsupported seeded index: {path}"),
|
||||
});
|
||||
|
||||
let error = table
|
||||
.create_index(
|
||||
&["vector"],
|
||||
Index::IvfPq(IvfPqIndexBuilder::default().seed(42)),
|
||||
)
|
||||
.execute()
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(
|
||||
error,
|
||||
Error::NotSupported { message }
|
||||
if message.contains("not supported on remote tables")
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_create_index_returns_job() {
|
||||
let describe_calls = std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0));
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user