diff --git a/python/python/lancedb/_lancedb.pyi b/python/python/lancedb/_lancedb.pyi index 47e727f99..180c78580 100644 --- a/python/python/lancedb/_lancedb.pyi +++ b/python/python/lancedb/_lancedb.pyi @@ -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( diff --git a/python/python/lancedb/index.py b/python/python/lancedb/index.py index aa7846892..cd6ee954e 100644 --- a/python/python/lancedb/index.py +++ b/python/python/lancedb/index.py @@ -314,7 +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, + # 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 @@ -422,7 +422,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, + # 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 @@ -618,7 +618,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, + # 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 @@ -651,7 +651,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, + # 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 @@ -784,7 +784,7 @@ class IvfPq: 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, + # 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 @@ -840,7 +840,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, + # 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 diff --git a/python/python/lancedb/table.py b/python/python/lancedb/table.py index ae36bac7a..c9b01b816 100644 --- a/python/python/lancedb/table.py +++ b/python/python/lancedb/table.py @@ -214,6 +214,47 @@ 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", +} + + +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_TYPES.get(type(config)) + if index_type is None: + raise ValueError( + f"Index type {type(config).__name__} does not support an accelerator" + ) + + 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 +2778,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,39 +2796,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: - # 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( @@ -2818,6 +2838,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( @@ -4830,6 +4855,7 @@ class AsyncTable: config: Optional[ Union[ IvfFlat, + IvfSq, IvfPq, IvfRq, HnswPq, @@ -4900,6 +4926,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, @@ -4926,6 +4969,7 @@ class AsyncTable: config: Optional[ Union[ IvfFlat, + IvfSq, IvfPq, IvfRq, HnswPq, @@ -4948,6 +4992,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, diff --git a/python/python/tests/test_table.py b/python/python/tests/test_table.py index 069527b21..0b19e9864 100644 --- a/python/python/tests/test_table.py +++ b/python/python/tests/test_table.py @@ -10,7 +10,7 @@ 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 @@ -25,7 +25,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 +1412,113 @@ 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.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_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( diff --git a/python/src/table.rs b/python/src/table.rs index 5b5d6596a..99bef4e1d 100644 --- a/python/src/table.rs +++ b/python/src/table.rs @@ -622,6 +622,10 @@ impl Table { self.inner.is_some() } + pub fn _is_native(&self) -> PyResult { + 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();