mirror of
https://github.com/lancedb/lancedb.git
synced 2026-09-11 15:52:17 +00:00
Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bc3837c4fe | ||
|
|
2a4f4f338b |
@@ -261,6 +261,7 @@ class Table:
|
|||||||
def name(self) -> str: ...
|
def name(self) -> str: ...
|
||||||
def __repr__(self) -> str: ...
|
def __repr__(self) -> str: ...
|
||||||
def is_open(self) -> bool: ...
|
def is_open(self) -> bool: ...
|
||||||
|
def _is_native(self) -> bool: ...
|
||||||
def close(self) -> None: ...
|
def close(self) -> None: ...
|
||||||
async def schema(self) -> pa.Schema: ...
|
async def schema(self) -> pa.Schema: ...
|
||||||
async def add(
|
async def add(
|
||||||
|
|||||||
@@ -314,8 +314,7 @@ class HnswPq:
|
|||||||
m: int = 20
|
m: int = 20
|
||||||
ef_construction: int = 300
|
ef_construction: int = 300
|
||||||
target_partition_size: Optional[int] = None
|
target_partition_size: Optional[int] = None
|
||||||
# Name of the accelerator (e.g. "cuda") to use for IVF training. When set,
|
# Reserved for future accelerator support. create_index() currently raises if set.
|
||||||
# create_index() dispatches to pylance to build the index on the accelerator.
|
|
||||||
accelerator: Optional[str] = None
|
accelerator: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
@@ -422,8 +421,7 @@ class HnswSq:
|
|||||||
m: int = 20
|
m: int = 20
|
||||||
ef_construction: int = 300
|
ef_construction: int = 300
|
||||||
target_partition_size: Optional[int] = None
|
target_partition_size: Optional[int] = None
|
||||||
# Name of the accelerator (e.g. "cuda") to use for IVF training. When set,
|
# Reserved for future accelerator support. create_index() currently raises if set.
|
||||||
# create_index() dispatches to pylance to build the index on the accelerator.
|
|
||||||
accelerator: Optional[str] = None
|
accelerator: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
@@ -618,8 +616,7 @@ class IvfFlat:
|
|||||||
max_iterations: int = 50
|
max_iterations: int = 50
|
||||||
sample_rate: int = 256
|
sample_rate: int = 256
|
||||||
target_partition_size: Optional[int] = None
|
target_partition_size: Optional[int] = None
|
||||||
# Name of the accelerator (e.g. "cuda") to use for IVF training. When set,
|
# Reserved for future accelerator support. create_index() currently raises if set.
|
||||||
# create_index() dispatches to pylance to build the index on the accelerator.
|
|
||||||
accelerator: Optional[str] = None
|
accelerator: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
@@ -651,8 +648,7 @@ class IvfSq:
|
|||||||
max_iterations: int = 50
|
max_iterations: int = 50
|
||||||
sample_rate: int = 256
|
sample_rate: int = 256
|
||||||
target_partition_size: Optional[int] = None
|
target_partition_size: Optional[int] = None
|
||||||
# Name of the accelerator (e.g. "cuda") to use for IVF training. When set,
|
# Reserved for future accelerator support. create_index() currently raises if set.
|
||||||
# create_index() dispatches to pylance to build the index on the accelerator.
|
|
||||||
accelerator: Optional[str] = None
|
accelerator: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
@@ -784,7 +780,7 @@ class IvfPq:
|
|||||||
max_iterations: int = 50
|
max_iterations: int = 50
|
||||||
sample_rate: int = 256
|
sample_rate: int = 256
|
||||||
target_partition_size: Optional[int] = None
|
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.
|
# create_index() dispatches to pylance to build the index on the accelerator.
|
||||||
accelerator: Optional[str] = None
|
accelerator: Optional[str] = None
|
||||||
|
|
||||||
@@ -840,8 +836,7 @@ class IvfRq:
|
|||||||
max_iterations: int = 50
|
max_iterations: int = 50
|
||||||
sample_rate: int = 256
|
sample_rate: int = 256
|
||||||
target_partition_size: Optional[int] = None
|
target_partition_size: Optional[int] = None
|
||||||
# Name of the accelerator (e.g. "cuda") to use for IVF training. When set,
|
# Reserved for future accelerator support. create_index() currently raises if set.
|
||||||
# create_index() dispatches to pylance to build the index on the accelerator.
|
|
||||||
accelerator: Optional[str] = None
|
accelerator: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -68,6 +68,14 @@ from ..table import AsyncTable, BlobMode, Branches, IndexStatistics, Query, Tabl
|
|||||||
from ..types import BaseTokenizerType
|
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):
|
class RemoteTable(Table):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -457,6 +465,8 @@ class RemoteTable(Table):
|
|||||||
... "l2", vector_column_name="vector"
|
... "l2", vector_column_name="vector"
|
||||||
... )
|
... )
|
||||||
"""
|
"""
|
||||||
|
_reject_index_accelerator(config, accelerator)
|
||||||
|
|
||||||
# Detect whether this is a legacy API call
|
# Detect whether this is a legacy API call
|
||||||
is_legacy = self._is_legacy_create_index_call(
|
is_legacy = self._is_legacy_create_index_call(
|
||||||
metric,
|
metric,
|
||||||
@@ -484,12 +494,6 @@ class RemoteTable(Table):
|
|||||||
|
|
||||||
column = vector_column_name
|
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:
|
if replace is not None:
|
||||||
logging.warning(
|
logging.warning(
|
||||||
"replace is not supported on LanceDB cloud."
|
"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 job may already be complete when returned; callers must not assume
|
||||||
the index exists until :meth:`Job.wait` returns.
|
the index exists until :meth:`Job.wait` returns.
|
||||||
"""
|
"""
|
||||||
|
_reject_index_accelerator(config)
|
||||||
|
|
||||||
return Job(
|
return Job(
|
||||||
LOOP.run(
|
LOOP.run(
|
||||||
self._table.create_index_async(
|
self._table.create_index_async(
|
||||||
|
|||||||
@@ -214,6 +214,45 @@ IndexConfigType = Union[
|
|||||||
# Known distance metrics for legacy API detection
|
# Known distance metrics for legacy API detection
|
||||||
KNOWN_METRICS = {"l2", "cosine", "dot", "hamming"}
|
KNOWN_METRICS = {"l2", "cosine", "dot", "hamming"}
|
||||||
|
|
||||||
|
_PYLANCE_ACCELERATED_INDEX_TYPE = "IVF_PQ"
|
||||||
|
|
||||||
|
|
||||||
|
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_TYPE
|
||||||
|
if isinstance(config, IvfPq)
|
||||||
|
else type(config).__name__
|
||||||
|
)
|
||||||
|
if index_type.upper() != _PYLANCE_ACCELERATED_INDEX_TYPE:
|
||||||
|
raise ValueError(
|
||||||
|
f"Index type {index_type} does not support an accelerator; "
|
||||||
|
f"only {_PYLANCE_ACCELERATED_INDEX_TYPE} supports acceleration"
|
||||||
|
)
|
||||||
|
|
||||||
|
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(
|
def _into_pyarrow_reader(
|
||||||
data, schema: Optional[pa.Schema] = None
|
data, schema: Optional[pa.Schema] = None
|
||||||
@@ -2737,20 +2776,17 @@ class LanceTable(Table):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Handle accelerator through pylance
|
# 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(
|
self.to_lance().create_index(
|
||||||
column=column,
|
column=column,
|
||||||
index_type=index_type,
|
|
||||||
metric=metric,
|
|
||||||
num_partitions=num_partitions,
|
|
||||||
num_sub_vectors=num_sub_vectors,
|
|
||||||
replace=replace,
|
replace=replace,
|
||||||
accelerator=accelerator,
|
|
||||||
index_cache_size=index_cache_size,
|
index_cache_size=index_cache_size,
|
||||||
num_bits=num_bits,
|
name=name,
|
||||||
m=m,
|
train=train,
|
||||||
ef_construction=ef_construction,
|
**accelerated_options,
|
||||||
target_partition_size=target_partition_size,
|
|
||||||
)
|
)
|
||||||
self.checkout_latest()
|
self.checkout_latest()
|
||||||
return
|
return
|
||||||
@@ -2758,39 +2794,21 @@ class LanceTable(Table):
|
|||||||
# New API: metric is the column name
|
# New API: metric is the column name
|
||||||
column = metric
|
column = metric
|
||||||
|
|
||||||
# Check if config has accelerator set and dispatch to pylance
|
accelerated_options = (
|
||||||
if config is not None and hasattr(config, "accelerator"):
|
_pylance_accelerated_index_options(config)
|
||||||
acc = getattr(config, "accelerator", None)
|
if config is not None
|
||||||
if acc is not None:
|
else None
|
||||||
# Dispatch to pylance for GPU acceleration
|
)
|
||||||
index_type_map = {
|
if accelerated_options is not None:
|
||||||
"IvfFlat": "IVF_FLAT",
|
self.to_lance().create_index(
|
||||||
"IvfSq": "IVF_SQ",
|
column=column,
|
||||||
"IvfPq": "IVF_PQ",
|
replace=replace,
|
||||||
"IvfRq": "IVF_RQ",
|
name=name,
|
||||||
"HnswPq": "IVF_HNSW_PQ",
|
train=train,
|
||||||
"HnswSq": "IVF_HNSW_SQ",
|
**accelerated_options,
|
||||||
}
|
)
|
||||||
cfg_type = type(config).__name__
|
self.checkout_latest()
|
||||||
lance_index_type = index_type_map.get(cfg_type, "IVF_PQ")
|
return
|
||||||
|
|
||||||
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
|
|
||||||
|
|
||||||
return LOOP.run(
|
return LOOP.run(
|
||||||
self._table.create_index(
|
self._table.create_index(
|
||||||
@@ -2818,6 +2836,11 @@ class LanceTable(Table):
|
|||||||
The job may already be complete when returned; callers must not assume
|
The job may already be complete when returned; callers must not assume
|
||||||
the index exists until :meth:`Job.wait` returns.
|
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(
|
return Job(
|
||||||
LOOP.run(
|
LOOP.run(
|
||||||
self._table.create_index_async(
|
self._table.create_index_async(
|
||||||
@@ -4830,6 +4853,7 @@ class AsyncTable:
|
|||||||
config: Optional[
|
config: Optional[
|
||||||
Union[
|
Union[
|
||||||
IvfFlat,
|
IvfFlat,
|
||||||
|
IvfSq,
|
||||||
IvfPq,
|
IvfPq,
|
||||||
IvfRq,
|
IvfRq,
|
||||||
HnswPq,
|
HnswPq,
|
||||||
@@ -4900,6 +4924,23 @@ class AsyncTable:
|
|||||||
" BTree, Bitmap, LabelList, Fm, or FTS, but got "
|
" BTree, Bitmap, LabelList, Fm, or FTS, but got "
|
||||||
+ str(type(config))
|
+ 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:
|
try:
|
||||||
await self._inner.create_index(
|
await self._inner.create_index(
|
||||||
column,
|
column,
|
||||||
@@ -4926,6 +4967,7 @@ class AsyncTable:
|
|||||||
config: Optional[
|
config: Optional[
|
||||||
Union[
|
Union[
|
||||||
IvfFlat,
|
IvfFlat,
|
||||||
|
IvfSq,
|
||||||
IvfPq,
|
IvfPq,
|
||||||
IvfRq,
|
IvfRq,
|
||||||
HnswPq,
|
HnswPq,
|
||||||
@@ -4948,6 +4990,11 @@ class AsyncTable:
|
|||||||
be complete when returned; callers must not assume the index exists
|
be complete when returned; callers must not assume the index exists
|
||||||
until :meth:`AsyncJob.wait` resolves.
|
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(
|
job = await self._inner.create_index_async(
|
||||||
column,
|
column,
|
||||||
index=config,
|
index=config,
|
||||||
|
|||||||
@@ -875,6 +875,25 @@ def test_remote_create_index_async_returns_job():
|
|||||||
job.cancel()
|
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():
|
def test_remote_job_wait_raises_on_failure():
|
||||||
from lancedb.exceptions import JobFailedError
|
from lancedb.exceptions import JobFailedError
|
||||||
from lancedb.index import BTree
|
from lancedb.index import BTree
|
||||||
|
|||||||
@@ -10,11 +10,21 @@ from concurrent.futures import ThreadPoolExecutor
|
|||||||
from datetime import date, datetime, timedelta
|
from datetime import date, datetime, timedelta
|
||||||
from time import sleep
|
from time import sleep
|
||||||
from typing import List
|
from typing import List
|
||||||
from unittest.mock import patch
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
import lancedb
|
import lancedb
|
||||||
from lancedb.dependencies import _PANDAS_AVAILABLE
|
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 numpy as np
|
||||||
import polars as pl
|
import polars as pl
|
||||||
import pyarrow as pa
|
import pyarrow as pa
|
||||||
@@ -25,7 +35,7 @@ from lancedb.db import AsyncConnection, DBConnection
|
|||||||
from lancedb.embeddings import EmbeddingFunctionConfig, EmbeddingFunctionRegistry
|
from lancedb.embeddings import EmbeddingFunctionConfig, EmbeddingFunctionRegistry
|
||||||
from lancedb.expr import col, lit
|
from lancedb.expr import col, lit
|
||||||
from lancedb.pydantic import LanceModel, Vector
|
from lancedb.pydantic import LanceModel, Vector
|
||||||
from lancedb.table import LanceTable
|
from lancedb.table import AsyncTable, LanceTable
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
|
|
||||||
|
|
||||||
@@ -1412,6 +1422,174 @@ def test_create_index_async_returns_done_job(mem_db: DBConnection):
|
|||||||
job.cancel()
|
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.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()
|
||||||
|
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_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()
|
||||||
|
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")
|
@patch("lancedb.table.AsyncTable.create_index")
|
||||||
def test_create_index_method(mock_create_index, mem_db: DBConnection):
|
def test_create_index_method(mock_create_index, mem_db: DBConnection):
|
||||||
table = mem_db.create_table(
|
table = mem_db.create_table(
|
||||||
|
|||||||
@@ -622,6 +622,10 @@ impl Table {
|
|||||||
self.inner.is_some()
|
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.
|
/// Closes the table, releasing any resources associated with it.
|
||||||
pub fn close(&mut self) {
|
pub fn close(&mut self) {
|
||||||
self.inner.take();
|
self.inner.take();
|
||||||
|
|||||||
Reference in New Issue
Block a user