diff --git a/python/python/lancedb/index.py b/python/python/lancedb/index.py index cd6ee954e..d9bdfb848 100644 --- a/python/python/lancedb/index.py +++ b/python/python/lancedb/index.py @@ -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 diff --git a/python/python/lancedb/remote/table.py b/python/python/lancedb/remote/table.py index acc2f4c9d..8012d2f71 100644 --- a/python/python/lancedb/remote/table.py +++ b/python/python/lancedb/remote/table.py @@ -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( diff --git a/python/python/lancedb/table.py b/python/python/lancedb/table.py index c9b01b816..6735c16fc 100644 --- a/python/python/lancedb/table.py +++ b/python/python/lancedb/table.py @@ -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 { diff --git a/python/python/tests/test_remote_db.py b/python/python/tests/test_remote_db.py index d5d3569d3..98d95c435 100644 --- a/python/python/tests/test_remote_db.py +++ b/python/python/tests/test_remote_db.py @@ -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 diff --git a/python/python/tests/test_table.py b/python/python/tests/test_table.py index 0b19e9864..5176bef57 100644 --- a/python/python/tests/test_table.py +++ b/python/python/tests/test_table.py @@ -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()