fix(python): honor MPS accelerator in async indexing

This commit is contained in:
Gatefixer
2026-08-06 05:29:08 +00:00
parent 7357d63e87
commit 2a4f4f338b
5 changed files with 212 additions and 51 deletions
+1
View File
@@ -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(
+6 -6
View File
@@ -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
+92 -43
View File
@@ -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,
+109 -2
View File
@@ -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(
+4
View File
@@ -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();