mirror of
https://github.com/lancedb/lancedb.git
synced 2026-08-18 20:18:37 +00:00
fix(python): reject unsupported index accelerators
This commit is contained in:
@@ -314,8 +314,7 @@ class HnswPq:
|
||||
m: int = 20
|
||||
ef_construction: int = 300
|
||||
target_partition_size: Optional[int] = None
|
||||
# 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.
|
||||
# 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 ("cuda" or "mps") 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 ("cuda" or "mps") 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 ("cuda" or "mps") 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
|
||||
|
||||
|
||||
@@ -840,8 +836,7 @@ class IvfRq:
|
||||
max_iterations: int = 50
|
||||
sample_rate: int = 256
|
||||
target_partition_size: Optional[int] = None
|
||||
# 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.
|
||||
# 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,14 +214,7 @@ IndexConfigType = Union[
|
||||
# Known distance metrics for legacy API detection
|
||||
KNOWN_METRICS = {"l2", "cosine", "dot", "hamming"}
|
||||
|
||||
_PYLANCE_ACCELERATED_INDEX_TYPES = {
|
||||
IvfFlat: "IVF_FLAT",
|
||||
IvfSq: "IVF_SQ",
|
||||
IvfPq: "IVF_PQ",
|
||||
IvfRq: "IVF_RQ",
|
||||
HnswPq: "IVF_HNSW_PQ",
|
||||
HnswSq: "IVF_HNSW_SQ",
|
||||
}
|
||||
_PYLANCE_ACCELERATED_INDEX_TYPE = "IVF_PQ"
|
||||
|
||||
|
||||
def _pylance_accelerated_index_options(
|
||||
@@ -237,10 +230,15 @@ def _pylance_accelerated_index_options(
|
||||
return None
|
||||
|
||||
if index_type is None:
|
||||
index_type = _PYLANCE_ACCELERATED_INDEX_TYPES.get(type(config))
|
||||
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 {type(config).__name__} does not support an accelerator"
|
||||
f"Index type {index_type} does not support an accelerator; "
|
||||
f"only {_PYLANCE_ACCELERATED_INDEX_TYPE} supports acceleration"
|
||||
)
|
||||
|
||||
return {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -14,7 +14,17 @@ 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
|
||||
@@ -1455,6 +1465,51 @@ def test_create_index_dispatches_mps_to_pylance(mem_db: DBConnection):
|
||||
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()
|
||||
@@ -1494,6 +1549,22 @@ async def test_async_create_index_dispatches_mps_to_pylance():
|
||||
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()
|
||||
|
||||
Reference in New Issue
Block a user