Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
154 changes: 151 additions & 3 deletions tests/test_milvus.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,14 +8,25 @@
from types import SimpleNamespace
from unittest.mock import MagicMock, call

import numpy as np
import pytest
from pydantic import SecretStr
from pymilvus import DataType

from vectordb_bench.backend.cases import CaseType
from vectordb_bench.backend.clients import DB
from vectordb_bench.backend.clients.api import IndexType
from vectordb_bench.backend.clients.milvus.config import MilvusConfig, MilvusFtsConfig
from vectordb_bench.backend.clients.milvus.milvus import MILVUS_FORCE_MERGE_TARGET_SIZE_MB, Milvus
from vectordb_bench.backend.clients.api import IndexType, MetricType
from vectordb_bench.backend.clients.milvus.config import (
HNSWBF16Config,
HNSWConfig,
MilvusConfig,
MilvusFtsConfig,
)
from vectordb_bench.backend.clients.milvus.milvus import (
MILVUS_FORCE_MERGE_TARGET_SIZE_MB,
Milvus,
_float32_to_bf16_bytes,
)
from vectordb_bench.backend.payload import PayloadProfile
from vectordb_bench.backend.runner.mp_runner import MultiProcessingSearchRunner
from vectordb_bench.interface import BenchMarkRunner
Expand All @@ -24,6 +35,143 @@
log = logging.getLogger(__name__)


def test_float32_to_bf16_bytes_uses_round_to_nearest_even() -> None:
float32_bits = np.array(
[
[0x3F807FFF, 0x3F808000, 0x3F808001, 0x3F818000],
[0x7F800000, 0xFF800000, 0x7FC00001, 0xFFC00001],
],
dtype=np.uint32,
)

encoded = _float32_to_bf16_bytes(float32_bits.view(np.float32))
actual = np.frombuffer(b"".join(encoded), dtype="<u2").reshape(float32_bits.shape)

expected = np.array(
[
[0x3F80, 0x3F80, 0x3F81, 0x3F82],
[0x7F80, 0xFF80, 0x7FC0, 0xFFC0],
],
dtype=np.uint16,
)
np.testing.assert_array_equal(actual, expected)


def test_hnsw_bf16_config_requests_plain_hnsw_index_type() -> None:
config = HNSWBF16Config(M=30, efConstruction=360, ef=100, metric_type=MetricType.COSINE)

assert config.index == IndexType.HNSW_BF16
# Milvus has no HNSW_BF16 index type; BF16 is selected by the field's data type.
assert config.index_param() == {
"metric_type": "COSINE",
"index_type": "HNSW",
"params": {"M": 30, "efConstruction": 360},
}
assert config.search_param() == {"metric_type": "COSINE", "params": {"ef": 100}}


def test_hnsw_bf16_index_type_resolves_to_bf16_case_config() -> None:
assert DB.Milvus.case_config_cls(IndexType.HNSW_BF16) is HNSWBF16Config


def _vector_field_schema_call(
monkeypatch: pytest.MonkeyPatch,
db_case_config: HNSWConfig,
dim: int = 4,
) -> call:
client = MagicMock()
client.has_collection.return_value = False
client_cls = MagicMock(return_value=client)
schema = MagicMock()
client_cls.create_schema.return_value = schema
client_cls.prepare_index_params.return_value = MagicMock()
monkeypatch.setattr("vectordb_bench.backend.clients.milvus.milvus.MilvusClient", client_cls)

Milvus(dim=dim, db_config={"uri": "http://example.invalid"}, db_case_config=db_case_config)

return next(c for c in schema.add_field.call_args_list if c.args[0] == "vector")


def test_milvus_bf16_case_config_creates_bfloat16_vector_field(monkeypatch: pytest.MonkeyPatch) -> None:
field_call = _vector_field_schema_call(
monkeypatch,
HNSWBF16Config(M=8, efConstruction=64, ef=32, metric_type=MetricType.COSINE),
)

assert field_call.args[1] == DataType.BFLOAT16_VECTOR
assert field_call.kwargs["dim"] == 4


def test_milvus_fp32_case_config_still_creates_float_vector_field(monkeypatch: pytest.MonkeyPatch) -> None:
field_call = _vector_field_schema_call(
monkeypatch,
HNSWConfig(M=8, efConstruction=64, ef=32, metric_type=MetricType.COSINE),
)

assert field_call.args[1] == DataType.FLOAT_VECTOR


def test_milvus_fts_case_config_does_not_enable_bf16(monkeypatch: pytest.MonkeyPatch) -> None:
client = MagicMock()
client.has_collection.return_value = False
client_cls = MagicMock(return_value=client)
client_cls.create_schema.return_value = MagicMock()
client_cls.prepare_index_params.return_value = MagicMock()
monkeypatch.setattr("vectordb_bench.backend.clients.milvus.milvus.MilvusClient", client_cls)

# An FTS config has no `index` attribute at all, so the BF16 probe must tolerate its absence.
db = Milvus(dim=4, db_config={"uri": "http://example.invalid"}, db_case_config=MilvusFtsConfig())

assert db._use_bf16 is False


def test_milvus_bf16_insert_converts_runner_batch_to_bf16_bytes() -> None:
client = MagicMock()
client.insert.side_effect = lambda _collection, rows: {"insert_count": len(rows)}

db = object.__new__(Milvus)
db.client = client
db.collection_name = "test_collection"
db._primary_field = "pk"
db._scalar_id_field = "id"
db._vector_field = "vector"
db.with_scalar_labels = False
db._use_bf16 = True

embeddings = [[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]
count, err = db.insert_embeddings(embeddings=embeddings, metadata=[0, 1, 2])

assert count == 3
assert err is None
rows = client.insert.call_args.args[1]
assert [row["vector"] for row in rows] == _float32_to_bf16_bytes(embeddings)
# Two bytes per dimension, where a float32 payload would use four.
assert {len(row["vector"]) for row in rows} == {4}


def test_milvus_bf16_batch_search_converts_every_query() -> None:
captured = {}

def search(**kwargs):
captured.update(kwargs)
return [[{"pk": 0}], [{"pk": 1}], [{"pk": 2}]]

db = object.__new__(Milvus)
db.client = SimpleNamespace(search=search)
db.collection_name = "test_collection"
db._vector_field = "vector"
db._primary_field = "pk"
db.case_config = SimpleNamespace(search_param=lambda: {"metric_type": "COSINE"})
db.expr = ""
db._use_bf16 = True

queries = [[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]

assert db.search_embeddings(queries, k=1) == [[0], [1], [2]]
# Guards against encoding only the first query of a batch.
assert captured["data"] == _float32_to_bf16_bytes(queries)


def test_milvus_vector_payload_requests_vector_field_and_returns_ids():
captured = {}

Expand Down
35 changes: 35 additions & 0 deletions tests/test_milvus_zilliz_cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,9 @@
from click.testing import CliRunner
from pytest import MonkeyPatch

from vectordb_bench.backend.clients.api import IndexType
from vectordb_bench.backend.clients.milvus import cli as milvus_cli
from vectordb_bench.backend.clients.milvus.config import HNSWBF16Config
from vectordb_bench.backend.clients.zilliz_cloud import cli as zilliz_cli
from vectordb_bench.cli import cli as common_cli

Expand Down Expand Up @@ -57,6 +59,39 @@ def fake_run(**kwargs):
assert captured["db_case_config"].use_partition_key is True


def test_milvus_hnsw_bf16_cli_builds_bf16_case_config(monkeypatch: MonkeyPatch) -> None:
captured = {}

def fake_run(**kwargs):
captured.update(kwargs)

monkeypatch.setattr(milvus_cli, "run", fake_run)

result = CliRunner().invoke(
milvus_cli.MilvusHNSWBF16,
[
"--uri",
"http://localhost:19530",
"--collection-name",
"bench_bf16",
"--m",
"30",
"--ef-construction",
"360",
"--ef-search",
"100",
"--dry-run",
],
)

assert result.exit_code == 0, result.output
case_config = captured["db_case_config"]
assert isinstance(case_config, HNSWBF16Config)
assert case_config.index == IndexType.HNSW_BF16
assert (case_config.M, case_config.efConstruction, case_config.ef) == (30, 360, 100)
assert captured["db_config"].collection_name == "bench_bf16"


@pytest.mark.parametrize(
("level_args", "expected_level"),
[(["--level", "2"], 2), ([], None)],
Expand Down
1 change: 1 addition & 0 deletions vectordb_bench/backend/clients/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@ class IndexType(StrEnum):
HNSW_SQ = "HNSW_SQ"
HNSW_BQ = "HNSW_BQ"
HNSW_PQ = "HNSW_PQ"
HNSW_BF16 = "HNSW_BF16"
HNSW_PRQ = "HNSW_PRQ"
DISKANN = "DISKANN"
STREAMING_DISKANN = "DISKANN"
Expand Down
20 changes: 20 additions & 0 deletions vectordb_bench/backend/clients/milvus/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -156,6 +156,26 @@ def MilvusHNSW(**parameters: Unpack[MilvusHNSWTypedDict]):
)


@cli.command()
@click_parameter_decorators_from_typed_dict(MilvusHNSWTypedDict)
def MilvusHNSWBF16(**parameters: Unpack[MilvusHNSWTypedDict]):
from .config import HNSWBF16Config

run(
db=DBTYPE,
db_config=_build_milvus_config(parameters),
db_case_config=_with_partition_key(
HNSWBF16Config(
M=parameters["m"],
efConstruction=parameters["ef_construction"],
ef=parameters["ef_search"],
),
parameters,
),
**parameters,
)


class MilvusRefineTypedDict(TypedDict):
refine: Annotated[
bool,
Expand Down
12 changes: 12 additions & 0 deletions vectordb_bench/backend/clients/milvus/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -92,6 +92,17 @@ def search_param(self) -> dict:
}


class HNSWBF16Config(HNSWConfig):
index: IndexType = IndexType.HNSW_BF16

def index_param(self) -> dict:
return {
"metric_type": self.parse_metric(),
"index_type": "HNSW",
"params": {"M": self.M, "efConstruction": self.efConstruction},
}


class HNSWSQConfig(HNSWConfig, DBCaseConfig):
index: IndexType = IndexType.HNSW_SQ
sq_type: SQType = SQType.SQ8
Expand Down Expand Up @@ -598,6 +609,7 @@ def search_param(self) -> dict:
IndexType.AUTOINDEX: AutoIndexConfig,
IndexType.FTS: MilvusFtsConfig,
IndexType.HNSW: HNSWConfig,
IndexType.HNSW_BF16: HNSWBF16Config,
IndexType.HNSW_SQ: HNSWSQConfig,
IndexType.HNSW_PQ: HNSWPQConfig,
IndexType.HNSW_PRQ: HNSWPRQConfig,
Expand Down
31 changes: 27 additions & 4 deletions vectordb_bench/backend/clients/milvus/milvus.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,12 +6,13 @@
from contextlib import contextmanager
from typing import Any

import numpy as np
from pymilvus import DataType, Function, FunctionType, MilvusClient, MilvusException

from vectordb_bench.backend.filter import Filter, FilterOp
from vectordb_bench.backend.payload import PayloadProfile

from ..api import VectorDB
from ..api import IndexType, VectorDB
from .config import MilvusFtsConfig, MilvusIndexConfig

log = logging.getLogger(__name__)
Expand All @@ -21,13 +22,27 @@
MILVUS_FORCE_MERGE_RETRY_INTERVAL_SECONDS = 30


def _float32_to_bf16_bytes(vectors: Iterable[list[float]]) -> list[bytes]:
"""Convert float32 vectors to round-to-nearest-even BF16 byte strings."""
# pymilvus accepts raw BF16 bytes; keep this tested bit encoding to avoid an ml_dtypes dependency.
bits = np.asarray(vectors, dtype="<f4").view("<u4")
rounding_bias = np.uint32(0x7FFF) + ((bits >> np.uint32(16)) & np.uint32(1))
bf16 = ((bits + rounding_bias) >> np.uint32(16)).astype("<u2")

# Preserve NaNs instead of allowing payload rounding to overflow to infinity.
nan_mask = (bits & np.uint32(0x7FFFFFFF)) > np.uint32(0x7F800000)
bf16[nan_mask] = ((bits[nan_mask] >> np.uint32(16)) | np.uint32(0x40)).astype("<u2")
return [row.tobytes() for row in bf16]


class Milvus(VectorDB):
supports_batch_search = True
supported_filter_types: list[FilterOp] = [
FilterOp.NonFilter,
FilterOp.NumGE,
FilterOp.StrEqual,
]
_use_bf16: bool = False

@classmethod
def supports_full_text_search(cls) -> bool:
Expand Down Expand Up @@ -89,6 +104,8 @@ def __init__( # noqa: PLR0912, PLR0915
if self.with_scalar_labels:
self._scalar_payload_label_field = "scalar_label"

self._use_bf16 = getattr(self.case_config, "index", None) == IndexType.HNSW_BF16

client = MilvusClient(
uri=self.db_config.get("uri"),
user=self.db_config.get("user"),
Expand Down Expand Up @@ -134,7 +151,11 @@ def __init__( # noqa: PLR0912, PLR0915
else:
schema.add_field(self._primary_field, DataType.INT64, is_primary=True)
schema.add_field(self._scalar_id_field, DataType.INT64)
schema.add_field(self._vector_field, DataType.FLOAT_VECTOR, dim=dim)
schema.add_field(
self._vector_field,
DataType.BFLOAT16_VECTOR if self._use_bf16 else DataType.FLOAT_VECTOR,
dim=dim,
)

if self.multitenant_tenant_labels:
schema.add_field(
Expand Down Expand Up @@ -419,12 +440,13 @@ def insert_embeddings(
assert self.client is not None
assert len(embeddings) == len(metadata)

vectors = _float32_to_bf16_bytes(embeddings) if self._use_bf16 else embeddings
rows = []
for i in range(len(embeddings)):
row = {
self._primary_field: metadata[i],
self._scalar_id_field: metadata[i],
self._vector_field: embeddings[i],
self._vector_field: vectors[i],
}
if tenant_labels_data is not None:
row[self._multitenant_partition_key_field] = tenant_labels_data[i]
Expand Down Expand Up @@ -556,9 +578,10 @@ def search_embeddings(
tenant_expr = f"{tenant_field} == '{tenant}'"
expr = tenant_expr if not expr else f"({expr}) and ({tenant_expr})"

search_data = _float32_to_bf16_bytes(queries) if self._use_bf16 else queries
search_kwargs = {
"collection_name": self.collection_name,
"data": queries,
"data": search_data,
"anns_field": self._vector_field,
"search_params": self.case_config.search_param(),
"limit": k,
Expand Down
3 changes: 2 additions & 1 deletion vectordb_bench/cli/vectordbbench.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,7 @@
from ..backend.clients.lindorm.cli import LindormHNSW, LindormIVFBQ, LindormIVFPQ
from ..backend.clients.mariadb.cli import MariaDBHNSW
from ..backend.clients.memorydb.cli import MemoryDB
from ..backend.clients.milvus.cli import MilvusAutoIndex, MilvusFTS
from ..backend.clients.milvus.cli import MilvusAutoIndex, MilvusFTS, MilvusHNSWBF16
from ..backend.clients.oceanbase.cli import OceanBaseHNSW, OceanBaseIVF
from ..backend.clients.oss_opensearch.cli import OSSOpenSearch
from ..backend.clients.pgdiskann.cli import PgDiskAnn
Expand Down Expand Up @@ -70,6 +70,7 @@
cli.add_command(ZillizAutoIndex)
cli.add_command(MilvusAutoIndex)
cli.add_command(MilvusFTS)
cli.add_command(MilvusHNSWBF16)
cli.add_command(AWSOpenSearch)
cli.add_command(OSSOpenSearch)
cli.add_command(PgVectorScaleDiskAnn)
Expand Down