mirror of
https://github.com/lancedb/lancedb.git
synced 2026-08-24 06:58:31 +00:00
fix(python): honor MPS accelerator in async indexing
This commit is contained in:
@@ -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,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
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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();
|
||||
|
||||
Reference in New Issue
Block a user