mirror of
https://github.com/lancedb/lancedb.git
synced 2026-08-31 18:48:25 +00:00
Compare commits
4 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| bc3837c4fe | |||
| 2a4f4f338b | |||
| 7357d63e87 | |||
| 624a75edf7 |
@@ -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(
|
||||
|
||||
@@ -707,6 +707,9 @@ class LanceDBConnection(DBConnection):
|
||||
self._namespace_client_properties = namespace_client_properties
|
||||
if _inner is not None:
|
||||
self._conn = _inner
|
||||
# Native-derived wrappers resolve this in their async reconstruction
|
||||
# path so construction never synchronously re-enters LOOP.
|
||||
self._read_consistency_interval = read_consistency_interval
|
||||
self._cached_namespace_client = None
|
||||
return
|
||||
|
||||
@@ -756,11 +759,14 @@ class LanceDBConnection(DBConnection):
|
||||
# storage_options. Also, this class really shouldn't be holding any state
|
||||
# beyond _conn.
|
||||
self._conn = AsyncConnection(LOOP.run(do_connect()))
|
||||
# Keep property access synchronous so debugger introspection cannot wait on
|
||||
# the background loop while that thread is suspended at a breakpoint.
|
||||
self._read_consistency_interval = read_consistency_interval
|
||||
self._cached_namespace_client: Optional[LanceNamespace] = None
|
||||
|
||||
@property
|
||||
def read_consistency_interval(self) -> Optional[timedelta]:
|
||||
return LOOP.run(self._conn.get_read_consistency_interval())
|
||||
return self._read_consistency_interval
|
||||
|
||||
@property
|
||||
def session(self) -> Optional[Session]:
|
||||
@@ -771,8 +777,16 @@ class LanceDBConnection(DBConnection):
|
||||
return self._conn.uri
|
||||
|
||||
@classmethod
|
||||
def from_inner(cls, inner: LanceDbConnection):
|
||||
return cls(None, _inner=inner)
|
||||
def from_inner(
|
||||
cls,
|
||||
inner: LanceDbConnection,
|
||||
read_consistency_interval: Optional[timedelta],
|
||||
):
|
||||
return cls(
|
||||
None,
|
||||
read_consistency_interval=read_consistency_interval,
|
||||
_inner=inner,
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(uri={self._conn.uri!r})"
|
||||
|
||||
@@ -314,8 +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,
|
||||
# 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 (e.g. "cuda") 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 (e.g. "cuda") 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 (e.g. "cuda") 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
|
||||
|
||||
|
||||
@@ -784,7 +780,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,8 +836,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,
|
||||
# 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
|
||||
|
||||
|
||||
|
||||
@@ -226,7 +226,7 @@ class PermutationBuilder:
|
||||
|
||||
async def do_execute():
|
||||
inner_tbl = await self._async.execute()
|
||||
return LanceTable.from_inner(inner_tbl)
|
||||
return await LanceTable.from_inner(inner_tbl)
|
||||
|
||||
return LOOP.run(do_execute())
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -214,6 +214,45 @@ IndexConfigType = Union[
|
||||
# Known distance metrics for legacy API detection
|
||||
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(
|
||||
data, schema: Optional[pa.Schema] = None
|
||||
@@ -2182,11 +2221,15 @@ class LanceTable(Table):
|
||||
return self.name
|
||||
|
||||
@classmethod
|
||||
def from_inner(cls, tbl: LanceDBTable):
|
||||
from .db import LanceDBConnection
|
||||
async def from_inner(cls, tbl: LanceDBTable):
|
||||
from .db import AsyncConnection, LanceDBConnection
|
||||
|
||||
async_tbl = AsyncTable(tbl)
|
||||
conn = LanceDBConnection.from_inner(tbl.database())
|
||||
inner_conn = tbl.database()
|
||||
read_consistency_interval = await AsyncConnection(
|
||||
inner_conn
|
||||
).get_read_consistency_interval()
|
||||
conn = LanceDBConnection.from_inner(inner_conn, read_consistency_interval)
|
||||
return cls(
|
||||
conn,
|
||||
async_tbl.name,
|
||||
@@ -2733,20 +2776,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
|
||||
@@ -2754,39 +2794,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(
|
||||
@@ -2814,6 +2836,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(
|
||||
@@ -4826,6 +4853,7 @@ class AsyncTable:
|
||||
config: Optional[
|
||||
Union[
|
||||
IvfFlat,
|
||||
IvfSq,
|
||||
IvfPq,
|
||||
IvfRq,
|
||||
HnswPq,
|
||||
@@ -4896,6 +4924,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,
|
||||
@@ -4922,6 +4967,7 @@ class AsyncTable:
|
||||
config: Optional[
|
||||
Union[
|
||||
IvfFlat,
|
||||
IvfSq,
|
||||
IvfPq,
|
||||
IvfRq,
|
||||
HnswPq,
|
||||
@@ -4944,6 +4990,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,
|
||||
|
||||
@@ -77,6 +77,23 @@ def test_sync_repr_does_not_use_background_loop(tmp_path, monkeypatch):
|
||||
assert repr(table) == f"LanceTable(name='test', _conn={db!r})"
|
||||
|
||||
|
||||
def test_read_consistency_interval_does_not_use_background_loop(tmp_path, monkeypatch):
|
||||
from lancedb.background_loop import LOOP
|
||||
from lancedb.db import LanceDBConnection
|
||||
|
||||
consistency_interval = timedelta(seconds=5)
|
||||
db = lancedb.connect(tmp_path, read_consistency_interval=consistency_interval)
|
||||
db_from_inner = LanceDBConnection.from_inner(db._inner, consistency_interval)
|
||||
|
||||
def fail_run(*args, **kwargs):
|
||||
raise AssertionError("properties should not use the Python background loop")
|
||||
|
||||
monkeypatch.setattr(LOOP, "run", fail_run)
|
||||
|
||||
assert db.read_consistency_interval == consistency_interval
|
||||
assert db_from_inner.read_consistency_interval == consistency_interval
|
||||
|
||||
|
||||
def test_ingest_pd(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
|
||||
|
||||
@@ -6,6 +6,7 @@ import math
|
||||
import pytest
|
||||
|
||||
from lancedb import DBConnection, Table, connect
|
||||
from lancedb.background_loop import LOOP
|
||||
from lancedb.permutation import Permutation, Permutations, permutation_builder
|
||||
|
||||
|
||||
@@ -31,6 +32,25 @@ def test_split_random_ratios(mem_db):
|
||||
assert 65 <= split_1_count <= 75 # ~70% ± tolerance
|
||||
|
||||
|
||||
def test_execute_does_not_reenter_background_loop(tmp_path, monkeypatch):
|
||||
import threading
|
||||
|
||||
db = connect(tmp_path)
|
||||
tbl = db.create_table("test_table", pa.table({"x": range(10)}))
|
||||
original_run = LOOP.run
|
||||
|
||||
def fail_on_reentry(future):
|
||||
assert threading.current_thread() is not LOOP.thread
|
||||
return original_run(future)
|
||||
|
||||
monkeypatch.setattr(LOOP, "run", fail_on_reentry)
|
||||
|
||||
permutation_tbl = permutation_builder(tbl).execute()
|
||||
|
||||
assert permutation_tbl.count_rows() == 10
|
||||
assert permutation_tbl._conn.read_consistency_interval is None
|
||||
|
||||
|
||||
def test_split_random_counts(mem_db):
|
||||
"""Test random splitting with absolute counts."""
|
||||
tbl = mem_db.create_table(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -6,14 +6,25 @@ import os
|
||||
import sys
|
||||
import threading
|
||||
import warnings
|
||||
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
|
||||
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
|
||||
@@ -24,7 +35,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
|
||||
|
||||
|
||||
@@ -1411,6 +1422,174 @@ 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.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")
|
||||
def test_create_index_method(mock_create_index, mem_db: DBConnection):
|
||||
table = mem_db.create_table(
|
||||
@@ -2124,6 +2303,27 @@ def test_delete(mem_db: DBConnection):
|
||||
assert table.to_arrow()["id"].to_pylist() == [1]
|
||||
|
||||
|
||||
def test_concurrent_deletes_are_thread_safe(mem_db: DBConnection):
|
||||
num_workers = 8
|
||||
table = mem_db.create_table(
|
||||
"my_table", data=[{"id": row_id} for row_id in range(num_workers)]
|
||||
)
|
||||
barrier = threading.Barrier(num_workers)
|
||||
|
||||
def delete(row_id: int):
|
||||
barrier.wait()
|
||||
return table.delete(f"id = {row_id}")
|
||||
|
||||
with ThreadPoolExecutor(max_workers=num_workers) as pool:
|
||||
results = list(pool.map(delete, range(num_workers)))
|
||||
|
||||
assert all(result.num_deleted_rows == 1 for result in results)
|
||||
assert sorted(result.version for result in results) == list(
|
||||
range(2, num_workers + 2)
|
||||
)
|
||||
assert table.count_rows() == 0
|
||||
|
||||
|
||||
def test_delete_expr(mem_db: DBConnection):
|
||||
table = mem_db.create_table(
|
||||
"my_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();
|
||||
@@ -745,6 +749,9 @@ impl Table {
|
||||
|
||||
#[allow(private_interfaces)]
|
||||
pub fn delete(self_: PyRef<'_, Self>, condition: PredicateArg) -> PyResult<Bound<'_, PyAny>> {
|
||||
// Do not hold the Python borrow across the await. The cloned Rust table
|
||||
// handle is thread-safe and allows deletes on the same Python table to
|
||||
// run concurrently without PyO3 reporting "Already borrowed".
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let result = match &condition {
|
||||
|
||||
Reference in New Issue
Block a user