mirror of
https://github.com/lancedb/lancedb.git
synced 2026-09-03 20:18:54 +00:00
Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 41d9c13fd3 |
@@ -162,6 +162,15 @@ def connect(
|
||||
... },
|
||||
... )
|
||||
|
||||
Azure managed identity authentication requires only the storage account name.
|
||||
LanceDB acquires and refreshes managed identity tokens automatically. For a
|
||||
user-assigned identity, also set ``azure_storage_client_id``:
|
||||
|
||||
>>> db = lancedb.connect( # doctest: +SKIP
|
||||
... "az://my-container/lancedb",
|
||||
... storage_options={"account_name": "my-storage-account"},
|
||||
... )
|
||||
|
||||
For tests and temporary data, use an in-memory database:
|
||||
|
||||
>>> db = lancedb.connect("memory://")
|
||||
@@ -455,6 +464,11 @@ async def connect_async(
|
||||
... db = await lancedb.connect_async("s3://my-bucket/lancedb",
|
||||
... storage_options={
|
||||
... "aws_access_key_id": "***"})
|
||||
... # Azure managed identity tokens are acquired and refreshed automatically
|
||||
... db = await lancedb.connect_async(
|
||||
... "az://my-container/lancedb",
|
||||
... storage_options={"account_name": "my-storage-account"},
|
||||
... )
|
||||
... # For tests and temporary data, use an in-memory database
|
||||
... db = await lancedb.connect_async("memory://")
|
||||
... # Connect to LanceDB cloud
|
||||
|
||||
@@ -261,7 +261,6 @@ 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,8 @@ class HnswPq:
|
||||
m: int = 20
|
||||
ef_construction: int = 300
|
||||
target_partition_size: Optional[int] = None
|
||||
# Reserved for future accelerator support. create_index() currently raises if set.
|
||||
# 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.
|
||||
accelerator: Optional[str] = None
|
||||
|
||||
|
||||
@@ -421,7 +422,8 @@ class HnswSq:
|
||||
m: int = 20
|
||||
ef_construction: int = 300
|
||||
target_partition_size: Optional[int] = None
|
||||
# Reserved for future accelerator support. create_index() currently raises if set.
|
||||
# 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.
|
||||
accelerator: Optional[str] = None
|
||||
|
||||
|
||||
@@ -616,7 +618,8 @@ class IvfFlat:
|
||||
max_iterations: int = 50
|
||||
sample_rate: int = 256
|
||||
target_partition_size: Optional[int] = None
|
||||
# Reserved for future accelerator support. create_index() currently raises if set.
|
||||
# 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.
|
||||
accelerator: Optional[str] = None
|
||||
|
||||
|
||||
@@ -648,7 +651,8 @@ class IvfSq:
|
||||
max_iterations: int = 50
|
||||
sample_rate: int = 256
|
||||
target_partition_size: Optional[int] = None
|
||||
# Reserved for future accelerator support. create_index() currently raises if set.
|
||||
# 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.
|
||||
accelerator: Optional[str] = None
|
||||
|
||||
|
||||
@@ -780,7 +784,7 @@ class IvfPq:
|
||||
max_iterations: int = 50
|
||||
sample_rate: int = 256
|
||||
target_partition_size: Optional[int] = None
|
||||
# Name of the accelerator ("cuda" or "mps") to use for IVF training. When set,
|
||||
# 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.
|
||||
accelerator: Optional[str] = None
|
||||
|
||||
@@ -836,7 +840,8 @@ class IvfRq:
|
||||
max_iterations: int = 50
|
||||
sample_rate: int = 256
|
||||
target_partition_size: Optional[int] = None
|
||||
# Reserved for future accelerator support. create_index() currently raises if set.
|
||||
# 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.
|
||||
accelerator: Optional[str] = None
|
||||
|
||||
|
||||
|
||||
@@ -68,14 +68,6 @@ 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,
|
||||
@@ -465,8 +457,6 @@ 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,
|
||||
@@ -494,6 +484,12 @@ 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."
|
||||
@@ -561,8 +557,6 @@ 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,45 +214,6 @@ 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
|
||||
@@ -2776,17 +2737,20 @@ class LanceTable(Table):
|
||||
)
|
||||
|
||||
# Handle accelerator through pylance
|
||||
accelerated_options = _pylance_accelerated_index_options(
|
||||
config, accelerator=accelerator, index_type=index_type
|
||||
)
|
||||
if accelerated_options is not None:
|
||||
if accelerator 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,
|
||||
name=name,
|
||||
train=train,
|
||||
**accelerated_options,
|
||||
num_bits=num_bits,
|
||||
m=m,
|
||||
ef_construction=ef_construction,
|
||||
target_partition_size=target_partition_size,
|
||||
)
|
||||
self.checkout_latest()
|
||||
return
|
||||
@@ -2794,21 +2758,39 @@ class LanceTable(Table):
|
||||
# New API: metric is the column name
|
||||
column = metric
|
||||
|
||||
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
|
||||
# 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
|
||||
|
||||
return LOOP.run(
|
||||
self._table.create_index(
|
||||
@@ -2836,11 +2818,6 @@ 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(
|
||||
@@ -4853,7 +4830,6 @@ class AsyncTable:
|
||||
config: Optional[
|
||||
Union[
|
||||
IvfFlat,
|
||||
IvfSq,
|
||||
IvfPq,
|
||||
IvfRq,
|
||||
HnswPq,
|
||||
@@ -4924,23 +4900,6 @@ 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,
|
||||
@@ -4967,7 +4926,6 @@ class AsyncTable:
|
||||
config: Optional[
|
||||
Union[
|
||||
IvfFlat,
|
||||
IvfSq,
|
||||
IvfPq,
|
||||
IvfRq,
|
||||
HnswPq,
|
||||
@@ -4990,11 +4948,6 @@ 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,
|
||||
|
||||
@@ -102,6 +102,10 @@ def fs_from_uri(uri: str) -> Tuple[pa_fs.FileSystem, str]:
|
||||
az_blob_fs = adlfs.AzureBlobFileSystem(
|
||||
account_name=os.environ.get("AZURE_STORAGE_ACCOUNT_NAME"),
|
||||
account_key=os.environ.get("AZURE_STORAGE_ACCOUNT_KEY"),
|
||||
# Without an explicit key, authenticate with DefaultAzureCredential
|
||||
# instead of attempting anonymous access. In particular, this enables
|
||||
# managed identity authentication on Azure hosts.
|
||||
anon=False,
|
||||
)
|
||||
|
||||
fs = pa_fs.PyFileSystem(pa_fs.FSSpecHandler(az_blob_fs))
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
|
||||
import lancedb
|
||||
@@ -12,11 +13,19 @@ import pytest
|
||||
# AWS_PROFILE=default TEST_S3_BASE_URL=s3://my_bucket/dataset pytest tests/test_io.py
|
||||
#
|
||||
# Azure:
|
||||
# You need to setup Azure credentials an a base path to run this test. Example
|
||||
# You need to set up Azure credentials and a base path to run this test. Examples:
|
||||
#
|
||||
# Account key:
|
||||
# export AZURE_STORAGE_ACCOUNT_NAME="<account>"
|
||||
# export AZURE_STORAGE_ACCOUNT_KEY="<key>"
|
||||
# export REMOTE_BASE_URL=az://my_blob/dataset
|
||||
# pytest tests/test_io.py
|
||||
#
|
||||
# Managed identity (system-assigned or user-assigned):
|
||||
# export AZURE_STORAGE_ACCOUNT_NAME="<account>"
|
||||
# export AZURE_STORAGE_CLIENT_ID="<client-id>" # user-assigned identity only
|
||||
# export REMOTE_BASE_URL=az://my_blob/dataset
|
||||
# pytest tests/test_io.py
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True, scope="module")
|
||||
@@ -58,3 +67,28 @@ def test_remote_io():
|
||||
assert len(db) == 1
|
||||
|
||||
assert db.open_table("test").name == db["test"].name
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
(os.environ.get("REMOTE_BASE_URL") is None),
|
||||
reason="please setup remote base url",
|
||||
)
|
||||
def test_remote_io_async():
|
||||
async def run():
|
||||
db = await lancedb.connect_async(os.environ["REMOTE_BASE_URL"])
|
||||
table = await db.create_table(
|
||||
"test_async",
|
||||
data=[
|
||||
{"vector": [3.1, 4.1], "item": "foo"},
|
||||
{"vector": [5.9, 26.5], "item": "bar"},
|
||||
],
|
||||
)
|
||||
|
||||
assert await table.count_rows() == 2
|
||||
assert (await db.open_table("test_async")).name == "test_async"
|
||||
assert "test_async" in await db.table_names()
|
||||
|
||||
await db.drop_table("test_async")
|
||||
assert "test_async" not in await db.table_names()
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
@@ -875,25 +875,6 @@ 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
|
||||
|
||||
@@ -10,21 +10,11 @@ from concurrent.futures import ThreadPoolExecutor
|
||||
from datetime import date, datetime, timedelta
|
||||
from time import sleep
|
||||
from typing import List
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from unittest.mock import patch
|
||||
|
||||
import lancedb
|
||||
from lancedb.dependencies import _PANDAS_AVAILABLE
|
||||
from lancedb.index import (
|
||||
BTree,
|
||||
FTS,
|
||||
HnswFlat,
|
||||
HnswPq,
|
||||
HnswSq,
|
||||
IvfFlat,
|
||||
IvfPq,
|
||||
IvfRq,
|
||||
IvfSq,
|
||||
)
|
||||
from lancedb.index import BTree, FTS, HnswFlat, HnswPq, HnswSq, IvfPq
|
||||
import numpy as np
|
||||
import polars as pl
|
||||
import pyarrow as pa
|
||||
@@ -35,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 AsyncTable, LanceTable
|
||||
from lancedb.table import LanceTable
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
@@ -1422,174 +1412,6 @@ 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(
|
||||
|
||||
@@ -26,7 +26,14 @@ import pandas as pd
|
||||
import polars as pl
|
||||
import pytest
|
||||
import lancedb
|
||||
from lancedb.util import flatten_columns, get_uri_scheme, join_uri, value_to_sql
|
||||
from lancedb import util
|
||||
from lancedb.util import (
|
||||
flatten_columns,
|
||||
fs_from_uri,
|
||||
get_uri_scheme,
|
||||
join_uri,
|
||||
value_to_sql,
|
||||
)
|
||||
from utils import exception_output
|
||||
|
||||
|
||||
@@ -78,6 +85,33 @@ def test_normalize_uri():
|
||||
assert parsed_scheme == expected_scheme
|
||||
|
||||
|
||||
def test_fs_from_uri_azure_uses_default_credential(monkeypatch):
|
||||
azure_options = {}
|
||||
filesystem = object()
|
||||
|
||||
class MockAdlfs:
|
||||
@staticmethod
|
||||
def AzureBlobFileSystem(**kwargs):
|
||||
azure_options.update(kwargs)
|
||||
return object()
|
||||
|
||||
monkeypatch.setenv("AZURE_STORAGE_ACCOUNT_NAME", "account")
|
||||
monkeypatch.delenv("AZURE_STORAGE_ACCOUNT_KEY", raising=False)
|
||||
monkeypatch.setattr(util, "adlfs", MockAdlfs)
|
||||
monkeypatch.setattr(util.pa_fs, "FSSpecHandler", lambda _: object())
|
||||
monkeypatch.setattr(util.pa_fs, "PyFileSystem", lambda _: filesystem)
|
||||
|
||||
actual_filesystem, path = fs_from_uri("az://container/database")
|
||||
|
||||
assert actual_filesystem is filesystem
|
||||
assert path == "container/database"
|
||||
assert azure_options == {
|
||||
"account_name": "account",
|
||||
"account_key": None,
|
||||
"anon": False,
|
||||
}
|
||||
|
||||
|
||||
def test_join_uri_remote():
|
||||
schemes = ["s3", "az", "gs"]
|
||||
for scheme in schemes:
|
||||
|
||||
@@ -622,10 +622,6 @@ 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