From ea49e4e403d9613c5daacdf4f5186134276e181d Mon Sep 17 00:00:00 2001 From: Alex Bondarev Date: Fri, 25 Sep 2026 14:39:58 -0500 Subject: [PATCH] feat(milvus): add native BF16 HNSW support Add IndexType.HNSW_BF16, HNSWBF16Config, and a MilvusHNSWBF16 CLI command. The client creates BFLOAT16_VECTOR fields and converts float32 insert batches and query vectors to round-to-nearest-even BF16 byte payloads accepted by pymilvus. The conversion uses NumPy bit operations, avoiding an additional ml_dtypes runtime dependency. Milvus still receives index_type "HNSW"; BF16 is selected by the collection field's data type, so FP32, FTS, and GPU configurations keep their existing behavior. Tests cover the rounding and special-value encoding, BFLOAT16_VECTOR schema selection against the FP32 default, insert and batch-search conversion, the CLI command wiring, and HNSW_BF16 case-config registration. Co-authored-by: RJ Silk --- tests/test_milvus.py | 154 +++++++++++++++++- tests/test_milvus_zilliz_cli.py | 35 ++++ vectordb_bench/backend/clients/api.py | 1 + vectordb_bench/backend/clients/milvus/cli.py | 20 +++ .../backend/clients/milvus/config.py | 12 ++ .../backend/clients/milvus/milvus.py | 31 +++- vectordb_bench/cli/vectordbbench.py | 3 +- 7 files changed, 248 insertions(+), 8 deletions(-) diff --git a/tests/test_milvus.py b/tests/test_milvus.py index 83fabf957..77aba7121 100644 --- a/tests/test_milvus.py +++ b/tests/test_milvus.py @@ -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 @@ -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=" 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 = {} diff --git a/tests/test_milvus_zilliz_cli.py b/tests/test_milvus_zilliz_cli.py index 6db46e9c7..c4d08e829 100644 --- a/tests/test_milvus_zilliz_cli.py +++ b/tests/test_milvus_zilliz_cli.py @@ -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 @@ -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)], diff --git a/vectordb_bench/backend/clients/api.py b/vectordb_bench/backend/clients/api.py index e1c74e9df..862929a5d 100644 --- a/vectordb_bench/backend/clients/api.py +++ b/vectordb_bench/backend/clients/api.py @@ -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" diff --git a/vectordb_bench/backend/clients/milvus/cli.py b/vectordb_bench/backend/clients/milvus/cli.py index ec11015c8..e57925103 100644 --- a/vectordb_bench/backend/clients/milvus/cli.py +++ b/vectordb_bench/backend/clients/milvus/cli.py @@ -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, diff --git a/vectordb_bench/backend/clients/milvus/config.py b/vectordb_bench/backend/clients/milvus/config.py index caa9a82e2..9629003ab 100644 --- a/vectordb_bench/backend/clients/milvus/config.py +++ b/vectordb_bench/backend/clients/milvus/config.py @@ -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 @@ -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, diff --git a/vectordb_bench/backend/clients/milvus/milvus.py b/vectordb_bench/backend/clients/milvus/milvus.py index dcb7e9a5d..bb9de9ab4 100644 --- a/vectordb_bench/backend/clients/milvus/milvus.py +++ b/vectordb_bench/backend/clients/milvus/milvus.py @@ -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__) @@ -21,6 +22,19 @@ 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="> np.uint32(16)) & np.uint32(1)) + bf16 = ((bits + rounding_bias) >> np.uint32(16)).astype(" np.uint32(0x7F800000) + bf16[nan_mask] = ((bits[nan_mask] >> np.uint32(16)) | np.uint32(0x40)).astype(" bool: @@ -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"), @@ -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( @@ -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] @@ -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, diff --git a/vectordb_bench/cli/vectordbbench.py b/vectordb_bench/cli/vectordbbench.py index 1bbc462ef..49af7afd0 100644 --- a/vectordb_bench/cli/vectordbbench.py +++ b/vectordb_bench/cli/vectordbbench.py @@ -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 @@ -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)