From 302f569b5036826395ecefb67e43864aa655d4ec Mon Sep 17 00:00:00 2001 From: jamesgao-jpg Date: Tue, 14 Jul 2026 00:08:39 +0800 Subject: [PATCH 1/9] feat: add insert batch and streaming rate controls Replace environment and CloudInsert-specific batch sizing with a task-level insert batch size shared by CLI, REST, frontend, datasets, and runners. Keep streaming insert rate case-specific and validate its relationship to batching. Signed-off-by: jamesgao-jpg --- .env.example | 1 - README.md | 10 +- docs/release/2026-05-cloud-leaderboard.md | 2 +- tests/test_case_runner_reuse.py | 12 + tests/test_cloud_insert_case.py | 45 ++-- tests/test_concurrent_runner.py | 16 +- tests/test_frontend_run_settings.py | 72 ++++++ tests/test_insert_control_cli.py | 131 +++++++++++ tests/test_insert_control_contract.py | 148 +++++++++++++ tests/test_insert_control_runners.py | 205 ++++++++++++++++++ tests/test_rate_runner.py | 160 +++++++------- vectordb_bench/__init__.py | 4 +- vectordb_bench/backend/cases.py | 28 +-- vectordb_bench/backend/clients/api.py | 3 +- vectordb_bench/backend/clients/doris/cli.py | 2 +- .../backend/clients/memorydb/cli.py | 12 +- .../backend/clients/memorydb/config.py | 7 +- .../backend/clients/memorydb/memorydb.py | 4 +- vectordb_bench/backend/dataset.py | 13 +- .../backend/runner/concurrent_runner.py | 6 +- vectordb_bench/backend/runner/mp_runner.py | 1 - vectordb_bench/backend/runner/rate_runner.py | 15 +- .../backend/runner/read_write_runner.py | 6 +- .../backend/runner/serial_runner.py | 19 +- vectordb_bench/backend/task_runner.py | 9 +- vectordb_bench/cli/cli.py | 39 +++- .../components/run_test/generateTasks.py | 9 +- .../components/run_test/runSettings.py | 59 +++++ .../frontend/config/dbCaseConfigs.py | 5 +- vectordb_bench/frontend/pages/run_test.py | 8 +- vectordb_bench/models.py | 49 ++++- vectordb_bench/restful/app.py | 2 + 32 files changed, 908 insertions(+), 194 deletions(-) create mode 100644 tests/test_frontend_run_settings.py create mode 100644 tests/test_insert_control_cli.py create mode 100644 tests/test_insert_control_contract.py create mode 100644 tests/test_insert_control_runners.py create mode 100644 vectordb_bench/frontend/components/run_test/runSettings.py diff --git a/.env.example b/.env.example index e495ea999..d0f5abbe3 100644 --- a/.env.example +++ b/.env.example @@ -2,7 +2,6 @@ LOG_LEVEL=INFO LOG_FILE="logs/vectordb_bench.log" # TIMEZONE= -# NUM_PER_BATCH= # DEFAULT_DATASET_URL= DATASET_LOCAL_DIR="/tmp/vectordb_bench/dataset" diff --git a/README.md b/README.md index 7851ef81b..f0795701b 100644 --- a/README.md +++ b/README.md @@ -477,7 +477,7 @@ pip install 'vectordb-bench[hologres]' 'psycopg[binary]' pgvector Execute tests for the index types: HGraph. ```shell -NUM_PER_BATCH=10000 vectordbbench hologreshgraph --host Hologres_Endpoint --port 80 \ +vectordbbench hologreshgraph --host Hologres_Endpoint --port 80 --insert-batch-size 10000 \ --user ACCESS_ID --password ACCESS_KEY --database DATABASE_NAME \ --m 64 --ef-construction 400 --case-type Performance768D10M \ --index-type HGraph --ef-search 400 --k 10 --num-concurrency 1,60,70,75,80,90,95,100,105,110,115,120,125,130 \ @@ -531,12 +531,12 @@ To list the options for zvec, execute vectordbbench zvec --help Doris supports ann index with type hnsw from version 4.0.x ```shell -NUM_PER_BATCH=1000000 vectordbbench doris --http-port=8030 --port=9030 --db-name=vector_test --case-type=Performance768D1M --stream-load-rows-per-batch=500000 +vectordbbench doris --http-port=8030 --port=9030 --db-name=vector_test --case-type=Performance768D1M --insert-batch-size=1000000 --stream-load-rows-per-batch=500000 ``` Using flag `--session-var`, if you want to test doris with some customized session variables. For example: ```shell -NUM_PER_BATCH=1000000 vectordbbench doris --http-port=8030 --port=9030 --db-name=vector_test --case-type=Performance768D1M --stream-load-rows-per-batch=500000 --session-var enable_profile=True +vectordbbench doris --http-port=8030 --port=9030 --db-name=vector_test --case-type=Performance768D1M --insert-batch-size=1000000 --stream-load-rows-per-batch=500000 --session-var enable_profile=True ``` Mote options: @@ -557,8 +557,8 @@ Mote options: --session-var TEXT Session variable key=value applied to each SQL session (repeatable) --stream-load-rows-per-batch INTEGER - Rows per single stream load request; default - uses NUM_PER_BATCH + Rows per Doris stream-load request; when + omitted, the Doris client default is used --no-index Create table without ANN index ``` diff --git a/docs/release/2026-05-cloud-leaderboard.md b/docs/release/2026-05-cloud-leaderboard.md index c9dc76a7b..252b6cd3b 100644 --- a/docs/release/2026-05-cloud-leaderboard.md +++ b/docs/release/2026-05-cloud-leaderboard.md @@ -65,7 +65,7 @@ vectordbbench zillizautoindex \ --uri "$ZILLIZ_URI" \ --token "$ZILLIZ_TOKEN" \ --collection-name cloud_insert_laion100m_bs10k \ - --cloud-insert-batch-size 10000 \ + --insert-batch-size 10000 \ --load-concurrency 16 \ --skip-search-serial \ --skip-search-concurrent \ diff --git a/tests/test_case_runner_reuse.py b/tests/test_case_runner_reuse.py index 55dfbda0b..ce06e5ba0 100644 --- a/tests/test_case_runner_reuse.py +++ b/tests/test_case_runner_reuse.py @@ -1,5 +1,6 @@ from pydantic import SecretStr +from vectordb_bench import config from vectordb_bench.backend.clients import DB from vectordb_bench.backend.clients.api import EmptyDBCaseConfig, MetricType from vectordb_bench.backend.clients.doris.config import DorisCaseConfig, DorisConfig @@ -12,6 +13,8 @@ from vectordb_bench.metric import Metric from vectordb_bench.models import CaseConfig, CaseType, TaskConfig, TaskStage, TestResult +DEFAULT_INSERT_BATCH_SIZE = config.DEFAULT_INSERT_BATCH_SIZE + def make_runner( *, @@ -21,6 +24,7 @@ def make_runner( db_config=None, db_case_config=None, stages: list[TaskStage] | None = None, + insert_batch_size: int = DEFAULT_INSERT_BATCH_SIZE, ) -> CaseRunner: if db_config is None: if db == DB.TurboPuffer: @@ -45,6 +49,7 @@ def make_runner( db_case_config=db_case_config, case_config=CaseConfig(case_id=case_id, custom_case=custom_case or {}), stages=stages or [TaskStage.DROP_OLD, TaskStage.LOAD, TaskStage.SEARCH_SERIAL], + insert_batch_size=insert_batch_size, ) return CaseRunner( run_id="run-id", @@ -111,6 +116,13 @@ def test_reuse_key_preserves_safe_payload_reuse(): assert hash(ids_only) == hash(vector) +def test_reuse_key_distinguishes_insert_batch_size(): + assert_not_reusable( + make_runner(insert_batch_size=100), + make_runner(insert_batch_size=200), + ) + + def test_reuse_key_distinguishes_physical_db_targets(): assert_not_reusable( make_runner(db_config=TurboPufferConfig(api_key="key", region="aws-us-east-1", namespace="namespace_a")), diff --git a/tests/test_cloud_insert_case.py b/tests/test_cloud_insert_case.py index 28e70f7bd..9fc8ba1a6 100644 --- a/tests/test_cloud_insert_case.py +++ b/tests/test_cloud_insert_case.py @@ -69,13 +69,13 @@ def iter_batches(self, batch_size): def test_cloud_insert_case_defaults_to_laion_100m(): - case = CloudInsertCase(batch_size=1000) + case = CloudInsertCase() assert case.case_id == CaseType.CloudInsertCase assert case.label == CaseLabel.CloudInsert assert case.dataset.data.name == "LAION" assert case.dataset.data.size == 100_000_000 - assert case.batch_size == 1000 + assert not hasattr(case, "batch_size") assert case.duration is None assert case.readiness_timeout is None @@ -84,14 +84,13 @@ def test_case_config_builds_cloud_insert_case_from_custom_case(): case = CaseConfig( case_id=CaseType.CloudInsertCase, custom_case={ - "batch_size": 5000, "duration": 1800, "dataset_with_size_type": DatasetWithSizeType.CohereMedium.value, }, ).case assert isinstance(case, CloudInsertCase) - assert case.batch_size == 5000 + assert not hasattr(case, "batch_size") assert case.duration == 1800 assert case.dataset.data.name == "Cohere" assert case.dataset.data.size == 1_000_000 @@ -101,13 +100,12 @@ def test_case_config_builds_cloud_insert_case_from_laion_100m_dataset_option(): case = CaseConfig( case_id=CaseType.CloudInsertCase, custom_case={ - "batch_size": 10_000, "dataset_with_size_type": "Large LAION (768dim, 100M)", }, ).case assert isinstance(case, CloudInsertCase) - assert case.batch_size == 10_000 + assert not hasattr(case, "batch_size") assert case.dataset_with_size_type == DatasetWithSizeType.LAIONLarge assert case.dataset.data.name == "LAION" assert case.dataset.data.size == 100_000_000 @@ -122,7 +120,6 @@ def test_laion_100m_dataset_option_uses_100m_timeouts(): def test_cli_builds_cloud_insert_custom_case_config(): params = { "case_type": "CloudInsertCase", - "cloud_insert_batch_size": 10_000, "cloud_insert_duration": 1800, "cloud_insert_readiness_timeout": 7200, "cloud_insert_readiness_poll_interval": 10, @@ -130,7 +127,6 @@ def test_cli_builds_cloud_insert_custom_case_config(): } assert get_custom_case_config(params) == { - "batch_size": 10_000, "duration": 1800, "readiness_timeout": 7200, "readiness_poll_interval": 10, @@ -142,7 +138,6 @@ def test_cli_builds_cloud_insert_custom_case_config_with_laion_100m_dataset(): cfg = get_custom_case_config( { "case_type": "CloudInsertCase", - "cloud_insert_batch_size": 10_000, "cloud_insert_duration": None, "cloud_insert_readiness_timeout": None, "cloud_insert_readiness_poll_interval": None, @@ -151,7 +146,6 @@ def test_cli_builds_cloud_insert_custom_case_config_with_laion_100m_dataset(): ) assert cfg == { - "batch_size": 10_000, "duration": None, "dataset_with_size_type": DatasetWithSizeType.LAIONLarge.value, } @@ -166,7 +160,6 @@ def test_cli_builds_cloud_insert_custom_case_config_with_default_dataset(): cfg = get_custom_case_config( { "case_type": "CloudInsertCase", - "cloud_insert_batch_size": 10_000, "cloud_insert_duration": None, "cloud_insert_readiness_timeout": None, "cloud_insert_readiness_poll_interval": None, @@ -175,7 +168,6 @@ def test_cli_builds_cloud_insert_custom_case_config_with_default_dataset(): ) assert cfg == { - "batch_size": 10_000, "duration": None, "dataset_with_size_type": DatasetWithSizeType.CohereMedium.value, } @@ -243,11 +235,11 @@ def test_assembler_schedules_cloud_insert_case(): case_config=CaseConfig( case_id=CaseType.CloudInsertCase, custom_case={ - "batch_size": 1000, "dataset_with_size_type": DatasetWithSizeType.CohereMedium.value, }, ), stages=[TaskStage.DROP_OLD, TaskStage.LOAD], + insert_batch_size=1000, ) runner = Assembler.assemble_all("run-id", "task-label", [task], DatasetSource.S3) @@ -311,10 +303,11 @@ def test_cloud_insert_result_file_uses_insert_only_metrics(tmp_path: Path): db_case_config=EmptyDBCaseConfig(), case_config=CaseConfig( case_id=CaseType.CloudInsertCase, - custom_case={"batch_size": 1000, "duration": None}, + custom_case={"duration": None}, ), stages=[TaskStage.DROP_OLD, TaskStage.LOAD], load_concurrency=0, + insert_batch_size=1000, ), metrics=Metric( inserted_count=100_000_000, @@ -343,14 +336,16 @@ def test_cloud_insert_result_file_uses_insert_only_metrics(tmp_path: Path): } assert written["results"][0]["task_config"]["db_config"]["api_key"] == "**********" assert written["results"][0]["task_config"]["db_config"]["index_name"] == "laion100m" + assert written["results"][0]["task_config"]["insert_batch_size"] == 1000 assert written["results"][0]["task_config"]["case_config"] == { "case_id": 600, - "custom_case": {"batch_size": 1000, "duration": None}, + "custom_case": {"duration": None}, } read_back = TestResult.read_file(result_file) assert read_back.results[0].task_config.case_config.case_id == CaseType.CloudInsertCase - assert read_back.results[0].task_config.case_config.custom_case == {"batch_size": 1000, "duration": None} + assert read_back.results[0].task_config.case_config.custom_case == {"duration": None} + assert read_back.results[0].task_config.insert_batch_size == 1000 collected = ResultCollector.collect(tmp_path) assert len(collected) == 1 @@ -423,7 +418,7 @@ def write(self, **kwargs): def test_milvus_insert_readiness_uses_entity_count_and_index_progress(): db = Milvus.__new__(Milvus) db.collection_name = "c" - db._vector_index_name = "vector_idx" + db._main_index_name = "vector_idx" db.client = type( "Client", (), @@ -697,9 +692,9 @@ def poll_insert_readiness(self, expected_count): db = DB() monkeypatch.setattr("vectordb_bench.backend.task_runner.time.sleep", lambda _: None) - case = CloudInsertCase(batch_size=2) + case = CloudInsertCase() case.dataset = Dataset() - config = type("Config", (), {"load_concurrency": 1})() + config = type("Config", (), {"load_concurrency": 1, "insert_batch_size": 2})() runner = CaseRunner.construct(ca=case, db=db, config=config) metric = runner._run_cloud_insert_case() @@ -752,9 +747,13 @@ def fail_on_sleep(_seconds): monkeypatch.setattr("vectordb_bench.backend.task_runner.ConcurrentInsertRunner", FakeConcurrentInsertRunner) monkeypatch.setattr("vectordb_bench.backend.task_runner.time.sleep", fail_on_sleep) - case = CloudInsertCase(batch_size=1, readiness_timeout=0, readiness_poll_interval=0) + case = CloudInsertCase(readiness_timeout=0, readiness_poll_interval=0) case.dataset = Dataset() - runner = CaseRunner.construct(ca=case, db=DB(), config=type("Config", (), {"load_concurrency": 1})()) + runner = CaseRunner.construct( + ca=case, + db=DB(), + config=type("Config", (), {"load_concurrency": 1, "insert_batch_size": 1})(), + ) with pytest.raises(TimeoutError, match="fully_searchable.*last_status.*stalled"): runner._run_cloud_insert_case() @@ -798,9 +797,9 @@ def poll_insert_readiness(self, expected_count): return {"fully_searchable": True, "fully_indexed": True, "additional_parameters": {}} monkeypatch.setattr("vectordb_bench.backend.task_runner.ConcurrentInsertRunner", FakeConcurrentInsertRunner) - case = CloudInsertCase(batch_size=1000, duration=60) + case = CloudInsertCase(duration=60) case.dataset = Dataset() - config = type("Config", (), {"load_concurrency": 7})() + config = type("Config", (), {"load_concurrency": 7, "insert_batch_size": 1000})() runner = CaseRunner.construct(ca=case, db=DB(), config=config) metric = runner._run_cloud_insert_case() diff --git a/tests/test_concurrent_runner.py b/tests/test_concurrent_runner.py index c9e5d9267..8ad2a06b1 100644 --- a/tests/test_concurrent_runner.py +++ b/tests/test_concurrent_runner.py @@ -4,8 +4,8 @@ - Correctness tests (threading & async backends) - Parameterized benchmark: serial vs concurrent across (batch_size, workers) matrix -NUM_PER_BATCH is set via os.environ before each run. Since runners execute -task() in a spawn subprocess that re-imports config, the env var takes effect. +Batch size is passed directly to each runner so subprocess execution uses the +same explicit benchmark value. Requires: - Milvus running at localhost:19530 @@ -21,7 +21,6 @@ from __future__ import annotations import logging -import os import time from vectordb_bench.backend.clients import DB @@ -55,10 +54,6 @@ def prepare_dataset(): return dataset -def set_batch_size(batch_size: int) -> None: - os.environ["NUM_PER_BATCH"] = str(batch_size) - - def timed_run(runner: SerialInsertRunner | ConcurrentInsertRunner) -> tuple[int, float]: start = time.perf_counter() count = runner.run() @@ -100,23 +95,23 @@ def test_concurrent_insert_async(): def run_serial(batch_size: int) -> tuple[int, float]: - set_batch_size(batch_size) runner = SerialInsertRunner( db=get_milvus_db(f"bench_serial_b{batch_size}"), dataset=prepare_dataset(), normalize=False, + batch_size=batch_size, ) return timed_run(runner) def run_concurrent(batch_size: int, workers: int) -> tuple[int, float]: - set_batch_size(batch_size) runner = ConcurrentInsertRunner( db=get_milvus_db(f"bench_conc_b{batch_size}_w{workers}"), dataset=prepare_dataset(), normalize=False, max_workers=workers, backend=ExecutorBackend.THREADING, + batch_size=batch_size, ) return timed_run(runner) @@ -151,9 +146,6 @@ def bench_matrix(): print(f" {dur_s / dur_c:>11.2f}x", end="") print() - # restore default - set_batch_size(100) - if __name__ == "__main__": bench_matrix() diff --git a/tests/test_frontend_run_settings.py b/tests/test_frontend_run_settings.py new file mode 100644 index 000000000..4c421aa92 --- /dev/null +++ b/tests/test_frontend_run_settings.py @@ -0,0 +1,72 @@ +from collections import defaultdict + +import pytest + +from vectordb_bench.backend.cases import CaseType +from vectordb_bench.backend.clients import DB +from vectordb_bench.frontend.components.run_test import generateTasks +from vectordb_bench.frontend.components.run_test.runSettings import ( + DEFAULT_STREAMING_INSERT_RATE, + validate_streaming_insert_rates, +) +from vectordb_bench.models import CaseConfig + + +def streaming_case(insert_rate: int | None = None) -> CaseConfig: + custom_case = {} if insert_rate is None else {"insert_rate": insert_rate} + return CaseConfig(case_id=CaseType.StreamingPerformanceCase, custom_case=custom_case) + + +@pytest.mark.parametrize( + ("insert_rate", "batch_size", "expected_message"), + [ + (400, 500, "must be greater than or equal to"), + (750, 500, "must be divisible by"), + ], +) +def test_validate_streaming_insert_rates_rejects_invalid_rate( + insert_rate: int, + batch_size: int, + expected_message: str, +): + is_valid, errors = validate_streaming_insert_rates([streaming_case(insert_rate)], batch_size) + + assert not is_valid + assert len(errors) == 1 + assert expected_message in errors[0] + + +def test_validate_streaming_insert_rates_checks_each_streaming_case_and_uses_default(): + cases = [ + CaseConfig(case_id=CaseType.Performance768D1M), + streaming_case(), + CaseConfig(case_id=CaseType.StreamingCustomDataset, custom_case={"insert_rate": 1_000}), + ] + + is_valid, errors = validate_streaming_insert_rates(cases, DEFAULT_STREAMING_INSERT_RATE) + + assert is_valid + assert errors == [] + + +def test_generate_tasks_passes_batch_size_to_task_config(monkeypatch: pytest.MonkeyPatch): + captured_task_configs: list[dict[str, object]] = [] + + class CapturedTaskConfig: + def __init__(self, **kwargs: object): + captured_task_configs.append(kwargs) + + monkeypatch.setattr(generateTasks, "TaskConfig", CapturedTaskConfig) + case = CaseConfig(case_id=CaseType.Performance768D1M) + all_case_configs = defaultdict(lambda: defaultdict(dict)) + + tasks = generateTasks.generate_tasks( + [DB.Test], + {DB.Test: DB.Test.config_cls()}, + [case], + all_case_configs, + batch_size=250, + ) + + assert len(tasks) == 1 + assert captured_task_configs[0]["insert_batch_size"] == 250 diff --git a/tests/test_insert_control_cli.py b/tests/test_insert_control_cli.py new file mode 100644 index 000000000..df9ce7c6a --- /dev/null +++ b/tests/test_insert_control_cli.py @@ -0,0 +1,131 @@ +import pytest +from click.testing import CliRunner + +from vectordb_bench.backend.clients.memorydb import cli as memorydb_cli +from vectordb_bench.backend.clients.memorydb.config import MemoryDBHNSWConfig +from vectordb_bench.backend.clients.test import cli as test_cli +from vectordb_bench.cli import cli as core_cli + + +@pytest.mark.parametrize( + ("args", "expected_batch_size"), + [ + (["--dry-run"], 100), + (["--dry-run", "--insert-batch-size", "250"], 250), + ], +) +def test_common_insert_batch_size_is_forwarded_to_task_config( + monkeypatch: pytest.MonkeyPatch, + args: list[str], + expected_batch_size: int, +) -> None: + captured = {} + + class FakeTaskConfig: + def __init__(self, **kwargs): + captured.update(kwargs) + + monkeypatch.setattr(core_cli, "TaskConfig", FakeTaskConfig) + + result = CliRunner().invoke(test_cli.Test, args) + + assert result.exit_code == 0, result.output + assert captured["insert_batch_size"] == expected_batch_size + + +@pytest.mark.parametrize("option", ["--insert-batch-size", "--streaming-insert-rate"]) +def test_positive_insert_controls_reject_zero(option: str) -> None: + result = CliRunner().invoke(test_cli.Test, ["--dry-run", option, "0"]) + + assert result.exit_code == 2 + assert "x>=1" in result.output + + +@pytest.mark.parametrize( + "case_type", + ["StreamingPerformanceCase", "StreamingCustomDataset"], +) +def test_streaming_insert_rate_only_maps_to_streaming_cases(case_type: str) -> None: + streaming = core_cli.get_custom_case_config( + { + "case_type": case_type, + "dataset_with_size_type": None, + "streaming_insert_rate": 750, + }, + ) + non_streaming = core_cli.get_custom_case_config( + { + "case_type": "Performance1536D50K", + "dataset_with_size_type": None, + "streaming_insert_rate": 750, + }, + ) + + assert streaming == {"insert_rate": 750} + assert non_streaming == {} + + +def test_cloud_insert_no_longer_has_a_custom_batch_mapping() -> None: + custom_case = core_cli.get_custom_case_config( + { + "case_type": "CloudInsertCase", + "dataset_with_size_type": None, + "cloud_insert_duration": None, + "cloud_insert_readiness_timeout": None, + "cloud_insert_readiness_poll_interval": None, + }, + ) + + assert "batch_size" not in custom_case + + +def test_memorydb_cli_keeps_task_and_pipeline_batch_sizes_distinct(monkeypatch: pytest.MonkeyPatch) -> None: + captured = {} + + def fake_run(**kwargs): + captured.update(kwargs) + + monkeypatch.setattr(memorydb_cli, "run", fake_run) + + result = CliRunner().invoke( + memorydb_cli.MemoryDB, + [ + "--host", + "localhost", + "--dry-run", + "--insert-batch-size", + "200", + "--memorydb-pipeline-batch-size", + "8", + ], + ) + + assert result.exit_code == 0, result.output + assert captured["insert_batch_size"] == 200 + assert captured["db_case_config"].pipeline_batch_size == 8 + + +def test_memorydb_config_accepts_legacy_insert_batch_size() -> None: + config = MemoryDBHNSWConfig.model_validate({"insert_batch_size": 12}) + + assert config.pipeline_batch_size == 12 + assert config.model_dump()["pipeline_batch_size"] == 12 + assert "insert_batch_size" not in config.model_dump() + + +def test_memorydb_config_prefers_canonical_pipeline_batch_size() -> None: + config = MemoryDBHNSWConfig.model_validate( + { + "pipeline_batch_size": 8, + "insert_batch_size": 12, + }, + ) + + assert config.pipeline_batch_size == 8 + + +def test_removed_cloud_insert_batch_option_is_rejected() -> None: + result = CliRunner().invoke(test_cli.Test, ["--dry-run", "--cloud-insert-batch-size", "5000"]) + + assert result.exit_code == 2 + assert "No such option '--cloud-insert-batch-size'" in result.output diff --git a/tests/test_insert_control_contract.py b/tests/test_insert_control_contract.py new file mode 100644 index 000000000..c4ce5f8b8 --- /dev/null +++ b/tests/test_insert_control_contract.py @@ -0,0 +1,148 @@ +import importlib +import json +from pathlib import Path +from typing import Any + +import pytest +from pydantic import ValidationError + +from vectordb_bench import config +from vectordb_bench.backend.cases import CaseType +from vectordb_bench.backend.clients import DB, EmptyDBCaseConfig +from vectordb_bench.backend.clients.test.config import TestConfig +from vectordb_bench.metric import Metric +from vectordb_bench.models import CaseConfig, CaseResult, TaskConfig, TestResult + + +def make_task( + case_id: CaseType = CaseType.Performance768D1M, + custom_case: dict | None = None, + **overrides: Any, +) -> TaskConfig: + values = { + "db": DB.Test, + "db_config": TestConfig(), + "db_case_config": EmptyDBCaseConfig(), + "case_config": CaseConfig(case_id=case_id, custom_case=custom_case), + } + values.update(overrides) + return TaskConfig(**values) + + +def write_legacy_result(tmp_path: Path, task: TaskConfig, metrics: Metric) -> Path: + result = TestResult( + run_id="legacy-run", + task_label="legacy", + results=[CaseResult(task_config=task, metrics=metrics)], + ) + raw = result.model_dump(mode="json", serialize_as_any=True) + raw["results"][0]["task_config"].pop("insert_batch_size") + result_path = tmp_path / "legacy.json" + result_path.write_text(json.dumps(raw)) + return result_path + + +def test_insert_control_defaults_and_serialization(): + task = make_task() + + assert config.DEFAULT_INSERT_BATCH_SIZE == 100 + assert config.DEFAULT_STREAMING_INSERT_RATE == 500 + assert task.insert_batch_size == 100 + assert task.model_dump(mode="json")["insert_batch_size"] == 100 + assert "insert_rate" not in TaskConfig.model_fields + + +@pytest.mark.parametrize("insert_batch_size", [0, -1]) +def test_insert_batch_size_must_be_positive(insert_batch_size: int): + with pytest.raises(ValidationError, match="greater than 0"): + make_task(insert_batch_size=insert_batch_size) + + +@pytest.mark.parametrize( + "case_id", + [CaseType.StreamingPerformanceCase, CaseType.StreamingCustomDataset], +) +def test_streaming_rate_must_cover_and_divide_batch(case_id: CaseType): + valid = make_task(case_id, {"insert_rate": 1000}, insert_batch_size=250) + assert valid.case_config.custom_case["insert_rate"] == 1000 + + with pytest.raises(ValidationError, match="greater than or equal"): + make_task(case_id, {"insert_rate": 200}, insert_batch_size=250) + with pytest.raises(ValidationError, match="divisible"): + make_task(case_id, {"insert_rate": 550}, insert_batch_size=100) + + +def test_streaming_rate_uses_stable_default(): + assert make_task(CaseType.StreamingPerformanceCase, insert_batch_size=250).insert_batch_size == 250 + with pytest.raises(ValidationError, match="divisible"): + make_task(CaseType.StreamingPerformanceCase, insert_batch_size=300) + + +def test_read_file_migrates_cloud_insert_batch_size(tmp_path: Path): + path = write_legacy_result( + tmp_path, + make_task( + CaseType.CloudInsertCase, + {"batch_size": 250}, + insert_batch_size=100, + ), + Metric(additional_parameters={"num_per_batch": 100}), + ) + + result = TestResult.read_file(path) + + assert result.results[0].task_config.insert_batch_size == 250 + + +def test_read_file_migrates_metrics_batch_size(tmp_path: Path): + path = write_legacy_result( + tmp_path, + make_task( + CaseType.StreamingPerformanceCase, + {"insert_rate": 500}, + insert_batch_size=100, + ), + Metric(additional_parameters={"num_per_batch": 250}), + ) + + result = TestResult.read_file(path) + + assert result.results[0].task_config.insert_batch_size == 250 + + +def test_rest_run_accepts_insert_batch_size(monkeypatch: pytest.MonkeyPatch): + pytest.importorskip("flask") + restful_app = importlib.import_module("vectordb_bench.restful.app") + + captured: dict[str, Any] = {} + monkeypatch.setattr(restful_app.benchmark_runner, "has_running", lambda: False) + monkeypatch.setattr(restful_app.benchmark_runner, "set_download_address", lambda _value: None) + monkeypatch.setattr( + restful_app.benchmark_runner, + "run", + lambda tasks, task_label: captured.update(tasks=tasks, task_label=task_label), + ) + + response = restful_app.app.test_client().post( + "/run", + json={ + "task_label": "contract", + "tasks": [ + { + "db": DB.Test.value, + "db_config": {}, + "db_case_config": {}, + "case_config": { + "case_id": CaseType.StreamingPerformanceCase.value, + "custom_case": {"insert_rate": 1000}, + }, + "stages": [], + "insert_batch_size": 250, + } + ], + }, + ) + + assert response.get_json()["code"] == 0 + assert captured["task_label"] == "contract" + assert captured["tasks"][0].insert_batch_size == 250 diff --git a/tests/test_insert_control_runners.py b/tests/test_insert_control_runners.py new file mode 100644 index 000000000..db69e0fa0 --- /dev/null +++ b/tests/test_insert_control_runners.py @@ -0,0 +1,205 @@ +from contextlib import contextmanager +from types import SimpleNamespace +from typing import Any + +import pytest + +from vectordb_bench import config +from vectordb_bench.backend import task_runner as task_runner_module +from vectordb_bench.backend.cases import CaseLabel, StreamingPerformanceCase +from vectordb_bench.backend.dataset import DataSetIterator, FtsDocumentIterator +from vectordb_bench.backend.filter import non_filter +from vectordb_bench.backend.runner.concurrent_runner import ConcurrentInsertRunner +from vectordb_bench.backend.runner.serial_runner import SerialInsertRunner +from vectordb_bench.backend.task_runner import CaseRunner +from vectordb_bench.backend.workload import WorkloadKind + +DEFAULT_INSERT_BATCH_SIZE = config.DEFAULT_INSERT_BATCH_SIZE + + +class FakeDB: + name = "FakeDB" + thread_safe = True + + @contextmanager + def init(self): + yield + + def need_normalize_cosine(self): + return False + + +def make_case_runner( + case: Any, + *, + batch_size: int = 17, + load_concurrency: int = 3, + db: Any | None = None, +) -> CaseRunner: + task_config = SimpleNamespace( + insert_batch_size=batch_size, + load_concurrency=load_concurrency, + case_config=SimpleNamespace(k=10), + ) + return CaseRunner.model_construct(ca=case, config=task_config, db=db or FakeDB()) + + +def test_streaming_case_preserves_requested_insert_rate(): + case = StreamingPerformanceCase(insert_rate=550) + + assert case.insert_rate == 550 + assert "550 rows/s" in case.name + + +def test_direct_callers_use_stable_batch_default(): + dataset = SimpleNamespace(train_files=[]) + concurrent_dataset = SimpleNamespace(data=SimpleNamespace()) + + assert DataSetIterator(dataset)._batch_size == DEFAULT_INSERT_BATCH_SIZE + assert FtsDocumentIterator(SimpleNamespace())._batch_size == DEFAULT_INSERT_BATCH_SIZE + assert ConcurrentInsertRunner(FakeDB(), concurrent_dataset, normalize=False).batch_size == DEFAULT_INSERT_BATCH_SIZE + assert SerialInsertRunner(FakeDB(), dataset, normalize=False).batch_size == DEFAULT_INSERT_BATCH_SIZE + + +def test_serial_insert_runner_groups_rows_by_explicit_batch_size(): + class InsertDB(FakeDB): + def __init__(self): + self.metadata_batches = [] + + def insert_embeddings( + self, + embeddings: list[Any], + metadata: list[Any], + ) -> tuple[int, None]: + self.metadata_batches.append(metadata) + return len(metadata), None + + db = InsertDB() + runner = SerialInsertRunner(db, SimpleNamespace(), normalize=False, batch_size=2) + + inserted = runner.endless_insert_data( + all_embeddings=[[0.1], [0.2], [0.3], [0.4], [0.5]], + all_metadata=[0, 1, 2, 3, 4], + ) + + assert inserted == 5 + assert db.metadata_batches == [[0, 1], [2, 3], [4]] + + +@pytest.mark.parametrize( + ("label", "workload_kind"), + [ + (CaseLabel.Performance, WorkloadKind.VECTOR), + (CaseLabel.FullTextSearchPerformance, WorkloadKind.FULL_TEXT), + ], +) +def test_performance_load_propagates_task_batch( + monkeypatch: pytest.MonkeyPatch, + label: CaseLabel, + workload_kind: WorkloadKind, +): + created: dict[str, Any] = {} + + class FakeConcurrentInsertRunner: + def __init__(self, *args, **kwargs): + created.update(kwargs) + + def run(self): + return 9, 1.25 + + case = SimpleNamespace( + label=label, + is_multitenant=False, + dataset=SimpleNamespace(data=SimpleNamespace(metric_type="L2")), + filters=non_filter, + load_timeout=30, + with_scalar_labels=False, + ) + monkeypatch.setattr(task_runner_module, "ConcurrentInsertRunner", FakeConcurrentInsertRunner) + + result = make_case_runner(case)._load_train_data() + + assert result == (9, 1.25) + assert created["batch_size"] == 17 + assert created["workload_kind"] == workload_kind + + +def test_capacity_load_propagates_task_batch(monkeypatch: pytest.MonkeyPatch): + created: dict[str, Any] = {} + + class FakeSerialInsertRunner: + def __init__(self, *args, **kwargs): + created.update(kwargs) + + def run_endlessness(self): + return 123 + + case = SimpleNamespace( + label=CaseLabel.Load, + dataset=SimpleNamespace(data=SimpleNamespace(metric_type="L2")), + filters=non_filter, + load_timeout=30, + ) + monkeypatch.setattr(task_runner_module, "SerialInsertRunner", FakeSerialInsertRunner) + + metric = make_case_runner(case)._run_capacity_case() + + assert metric.max_load_count == 123 + assert created["batch_size"] == 17 + + +def test_cloud_insert_propagates_task_batch(monkeypatch: pytest.MonkeyPatch): + created: dict[str, Any] = {} + + class FakeConcurrentInsertRunner: + def __init__(self, *args, **kwargs): + created.update(kwargs) + + def task(self): + return 3 + + class ReadinessDB(FakeDB): + def poll_insert_readiness(self, expected_count: int) -> dict[str, Any]: + assert expected_count == 3 + return {"fully_searchable": True, "fully_indexed": True, "additional_parameters": {}} + + case = SimpleNamespace( + label=CaseLabel.CloudInsert, + is_multitenant=False, + dataset=SimpleNamespace(data=SimpleNamespace(metric_type="L2")), + filters=non_filter, + duration=60, + readiness_timeout=None, + readiness_poll_interval=0, + ) + monkeypatch.setattr(task_runner_module, "ConcurrentInsertRunner", FakeConcurrentInsertRunner) + + metric = make_case_runner(case, db=ReadinessDB())._run_cloud_insert_case() + + assert metric.inserted_count == 3 + assert created["batch_size"] == 17 + assert created["duration"] == 60 + + +def test_streaming_runner_propagates_task_batch(monkeypatch: pytest.MonkeyPatch): + created: dict[str, Any] = {} + + class FakeReadWriteRunner: + def __init__(self, **kwargs): + created.update(kwargs) + + case = SimpleNamespace( + label=CaseLabel.Streaming, + dataset=SimpleNamespace(data=SimpleNamespace(metric_type="L2")), + insert_rate=34, + search_stages=[0.5], + optimize_after_write=False, + read_dur_after_write=10, + concurrencies=[1], + ) + monkeypatch.setattr(task_runner_module, "ReadWriteRunner", FakeReadWriteRunner) + + make_case_runner(case)._init_read_write_runner() + + assert created["insert_rate"] == 34 + assert created["batch_size"] == 17 diff --git a/tests/test_rate_runner.py b/tests/test_rate_runner.py index df92b0dd7..a02b9da36 100644 --- a/tests/test_rate_runner.py +++ b/tests/test_rate_runner.py @@ -1,88 +1,90 @@ -from typing import Iterable -import argparse -from vectordb_bench.backend.dataset import Dataset, DatasetSource +from types import SimpleNamespace + +import pytest + +from vectordb_bench import config +from vectordb_bench.backend.runner import read_write_runner as read_write_runner_module +from vectordb_bench.backend.runner.mp_runner import MultiProcessingSearchRunner from vectordb_bench.backend.runner.rate_runner import RatedMultiThreadingInsertRunner from vectordb_bench.backend.runner.read_write_runner import ReadWriteRunner -from vectordb_bench.backend.clients import DB, VectorDB -from vectordb_bench.backend.clients.milvus.config import FLATConfig -from vectordb_bench.backend.clients.zilliz_cloud.config import AutoIndexConfig -import logging +DEFAULT_INSERT_BATCH_SIZE = config.DEFAULT_INSERT_BATCH_SIZE + + +class FakeDB: + name = "FakeDB" + -log = logging.getLogger("vectordb_bench") -log.setLevel(logging.DEBUG) +def test_rate_runner_uses_explicit_batch_size(): + runner = RatedMultiThreadingInsertRunner( + rate=30, + db=FakeDB(), + dataset_iter=iter(()), + batch_size=5, + ) + + assert runner.insert_rate == 30 + assert runner.batch_size == 5 + assert runner.batch_rate == 6 -def get_rate_runner(db): - cohere = Dataset.COHERE.manager(100_000) - prepared = cohere.prepare(DatasetSource.AliyunOSS) - assert prepared + +def test_rate_runner_direct_caller_uses_stable_batch_default(): runner = RatedMultiThreadingInsertRunner( - rate = 10, - db = db, - dataset = cohere, + rate=DEFAULT_INSERT_BATCH_SIZE, + db=FakeDB(), + dataset_iter=iter(()), ) - return runner - -def test_rate_runner(db, insert_rate): - runner = get_rate_runner(db) - - _, t = runner.run_with_rate() - log.info(f"insert run done, time={t}") - -def test_read_write_runner(db, insert_rate, conc: list, search_stage: Iterable[float], read_dur_after_write: int, local: bool=False): - cohere = Dataset.COHERE.manager(1_000_000) - if local is True: - source = DatasetSource.AliyunOSS - else: - source = DatasetSource.S3 - prepared = cohere.prepare(source) - assert prepared - - rw_runner = ReadWriteRunner( - db=db, - dataset=cohere, - insert_rate=insert_rate, - search_stage=search_stage, - read_dur_after_write=read_dur_after_write, - concurrencies=conc + assert runner.batch_size == DEFAULT_INSERT_BATCH_SIZE + assert runner.batch_rate == 1 + + +@pytest.mark.parametrize( + ("rate", "batch_size", "message"), + [ + (0, 10, "insert rate must be greater than 0"), + (-10, 10, "insert rate must be greater than 0"), + (10, 0, "insert batch size must be greater than 0"), + (10, -1, "insert batch size must be greater than 0"), + (10, 4, "insert rate 10 must be divisible by insert batch size 4"), + ], +) +def test_rate_runner_rejects_invalid_rate_batch_combinations(rate, batch_size, message): + with pytest.raises(ValueError, match=message): + RatedMultiThreadingInsertRunner( + rate=rate, + db=FakeDB(), + dataset_iter=iter(()), + batch_size=batch_size, + ) + + +def test_read_write_runner_requests_task_batch_from_dataset(monkeypatch): + requested_batch_sizes = [] + + class Dataset: + data = SimpleNamespace(size=100) + test_data = [] + gt_data = [] + + def iter_batches(self, batch_size): + requested_batch_sizes.append(batch_size) + return iter(()) + + class FakeSerialSearchRunner: + def __init__(self, **kwargs): + pass + + monkeypatch.setattr(MultiProcessingSearchRunner, "__init__", lambda self, **kwargs: None) + monkeypatch.setattr(read_write_runner_module, "SerialSearchRunner", FakeSerialSearchRunner) + + runner = ReadWriteRunner( + db=FakeDB(), + dataset=Dataset(), + insert_rate=30, + batch_size=5, ) - rw_runner.run_read_write() - - -def get_db(db: str, config: dict) -> VectorDB: - if db == DB.Milvus.name: - return DB.Milvus.init_cls(dim=768, db_config=config, db_case_config=FLATConfig(metric_type="COSINE"), drop_old=True) - elif db == DB.ZillizCloud.name: - return DB.ZillizCloud.init_cls(dim=768, db_config=config, db_case_config=AutoIndexConfig(metric_type="COSINE"), drop_old=True) - else: - raise ValueError(f"unknown db: {db}") - - -if __name__ == "__main__": - parser = argparse.ArgumentParser() - parser.add_argument("-r", "--insert_rate", type=int, default="1000", help="insert entity row count per seconds, cps") - parser.add_argument("-d", "--db", type=str, default=DB.Milvus.name, help="db name") - parser.add_argument("-t", "--duration", type=int, default=300, help="stage search duration in seconds") - parser.add_argument("--use_s3", action='store_true', help="whether to use S3 dataset") - - flags = parser.parse_args() - - # TODO read uri, user, password from .env - config = { - "uri": "http://localhost:19530", - "user": "", - "password": "", - } - - conc = (1, 15, 50) - search_stage = (0.5, 0.6, 0.7, 0.8, 0.9) - - db = get_db(flags.db, config) - test_read_write_runner( - db=db, - insert_rate=flags.insert_rate, - conc=conc, - search_stage=search_stage, - read_dur_after_write=flags.duration, - local=flags.use_s3) + + assert requested_batch_sizes == [5] + assert runner.batch_size == 5 + assert runner.batch_rate == 6 diff --git a/vectordb_bench/__init__.py b/vectordb_bench/__init__.py index 1491630e0..6bc44f290 100644 --- a/vectordb_bench/__init__.py +++ b/vectordb_bench/__init__.py @@ -13,13 +13,15 @@ class config: ALIYUN_OSS_URL = "assets.zilliz.com.cn/benchmark/" AWS_S3_URL = "assets.zilliz.com/benchmark/" + DEFAULT_INSERT_BATCH_SIZE = 100 + DEFAULT_STREAMING_INSERT_RATE = 500 + LOG_LEVEL = env.str("LOG_LEVEL", "INFO") LOG_FILE = env.str("LOG_FILE", "logs/vectordb_bench.log") DEFAULT_DATASET_URL = env.str("DEFAULT_DATASET_URL", AWS_S3_URL) DATASET_SOURCE = env.str("DATASET_SOURCE", "S3") # Options "S3", "AliyunOSS", or "IR_DATASETS" DATASET_LOCAL_DIR = env.path("DATASET_LOCAL_DIR", "/tmp/vectordb_bench/dataset") - NUM_PER_BATCH = env.int("NUM_PER_BATCH", 100) LOAD_CONCURRENCY = env.int("LOAD_CONCURRENCY", 0) # 0 = cpu_count TIME_PER_BATCH = 1 # 1s. for streaming insertion. MAX_INSERT_RETRY = 5 diff --git a/vectordb_bench/backend/cases.py b/vectordb_bench/backend/cases.py index 24eaf768e..f007aa06b 100644 --- a/vectordb_bench/backend/cases.py +++ b/vectordb_bench/backend/cases.py @@ -475,17 +475,6 @@ def __init__( concurrencies: list[int] | str = (5, 10), **kwargs, ): - num_per_batch = config.NUM_PER_BATCH - if insert_rate % config.NUM_PER_BATCH != 0: - _insert_rate = max( - num_per_batch, - insert_rate // num_per_batch * num_per_batch, - ) - log.warning( - f"[streaming_case init] insert_rate(={insert_rate}) should be " - f"divisible by NUM_PER_BATCH={num_per_batch}), reset to {_insert_rate}", - ) - insert_rate = _insert_rate if not isinstance(dataset_with_size_type, DatasetWithSizeType): dataset_with_size_type = DatasetWithSizeType(dataset_with_size_type) dataset = dataset_with_size_type.get_manager() @@ -535,18 +524,6 @@ def __init__( read_dur_after_write: int = 30, **kwargs, ): - num_per_batch = config.NUM_PER_BATCH - if insert_rate % config.NUM_PER_BATCH != 0: - _insert_rate = max( - num_per_batch, - insert_rate // num_per_batch * num_per_batch, - ) - log.warning( - f"[streaming_case init] insert_rate(={insert_rate}) should be " - f"divisible by NUM_PER_BATCH={num_per_batch}), reset to {_insert_rate}", - ) - insert_rate = _insert_rate - dataset_config = CustomDatasetConfig(**dataset_config) dataset = CustomDataset( name=dataset_config.name, @@ -767,7 +744,6 @@ def filters(self) -> Filter: class CloudInsertCase(Case): case_id: CaseType = CaseType.CloudInsertCase label: CaseLabel = CaseLabel.CloudInsert - batch_size: int duration: float | None = None readiness_timeout: float | None = config.CLOUD_INSERT_READINESS_TIMEOUT readiness_poll_interval: float = config.CLOUD_INSERT_READINESS_POLL_INTERVAL @@ -775,7 +751,6 @@ class CloudInsertCase(Case): def __init__( self, - batch_size: int, duration: float | None = None, readiness_timeout: float | None = config.CLOUD_INSERT_READINESS_TIMEOUT, readiness_poll_interval: float = config.CLOUD_INSERT_READINESS_POLL_INTERVAL, @@ -790,10 +765,9 @@ def __init__( else dataset_with_size_type.get_manager() ) super().__init__( - name=f"Cloud Insert - batch {batch_size}", + name="Cloud Insert", description="Cloud leaderboard insert-only case with readiness polling.", dataset=dataset, - batch_size=batch_size, duration=duration, readiness_timeout=readiness_timeout, readiness_poll_interval=readiness_poll_interval, diff --git a/vectordb_bench/backend/clients/api.py b/vectordb_bench/backend/clients/api.py index bac90b7be..f6e6dfbc2 100644 --- a/vectordb_bench/backend/clients/api.py +++ b/vectordb_bench/backend/clients/api.py @@ -331,8 +331,7 @@ def insert_embeddings( tenant_labels_data: list[str] | None = None, **kwargs, ) -> tuple[int, Exception]: - """Insert the embeddings to the vector database. The default number of embeddings for - each insert_embeddings is 5000. + """Insert one task-configured batch of embeddings into the vector database. Args: embeddings(list[list[float]]): list of embedding to add to the vector database. diff --git a/vectordb_bench/backend/clients/doris/cli.py b/vectordb_bench/backend/clients/doris/cli.py index 8153b412d..9e796f937 100644 --- a/vectordb_bench/backend/clients/doris/cli.py +++ b/vectordb_bench/backend/clients/doris/cli.py @@ -144,7 +144,7 @@ class DorisTypedDict(CommonTypedDict, HNSWBaseTypedDict): "--stream-load-rows-per-batch", type=int, required=False, - help="Rows per single stream load request; default uses NUM_PER_BATCH", + help="Rows per Doris stream-load request; when omitted, the Doris client default is used", ), ] no_index: Annotated[ diff --git a/vectordb_bench/backend/clients/memorydb/cli.py b/vectordb_bench/backend/clients/memorydb/cli.py index 568eec2a3..8e43f0843 100644 --- a/vectordb_bench/backend/clients/memorydb/cli.py +++ b/vectordb_bench/backend/clients/memorydb/cli.py @@ -48,13 +48,15 @@ class MemoryDBTypedDict(TypedDict): ), ), ] - insert_batch_size: Annotated[ + pipeline_batch_size: Annotated[ int, click.option( - "--insert-batch-size", - type=int, + "--memorydb-pipeline-batch-size", + "pipeline_batch_size", + type=click.IntRange(min=1), default=10, - help="Batch size for inserting data. Adjust this as needed, but don't make it too big", + show_default=True, + help="Commands buffered in each MemoryDB pipeline execution", ), ] @@ -82,7 +84,7 @@ def MemoryDB(**parameters: Unpack[MemoryDBHNSWTypedDict]): M=parameters["m"], ef_construction=parameters["ef_construction"], ef_runtime=parameters["ef_runtime"], - insert_batch_size=parameters["insert_batch_size"], + pipeline_batch_size=parameters["pipeline_batch_size"], ), **parameters, ) diff --git a/vectordb_bench/backend/clients/memorydb/config.py b/vectordb_bench/backend/clients/memorydb/config.py index 2c40ff546..befa3c23e 100644 --- a/vectordb_bench/backend/clients/memorydb/config.py +++ b/vectordb_bench/backend/clients/memorydb/config.py @@ -1,4 +1,4 @@ -from pydantic import BaseModel, SecretStr +from pydantic import AliasChoices, BaseModel, Field, PositiveInt, SecretStr from ..api import DBCaseConfig, DBConfig, IndexType, MetricType @@ -24,7 +24,10 @@ def to_dict(self) -> dict: class MemoryDBIndexConfig(BaseModel, DBCaseConfig): metric_type: MetricType | None = None - insert_batch_size: int | None = None + pipeline_batch_size: PositiveInt = Field( + default=10, + validation_alias=AliasChoices("pipeline_batch_size", "insert_batch_size"), + ) def parse_metric(self) -> str: if self.metric_type == MetricType.L2: diff --git a/vectordb_bench/backend/clients/memorydb/memorydb.py b/vectordb_bench/backend/clients/memorydb/memorydb.py index 7e7a8650b..d816d0fd4 100644 --- a/vectordb_bench/backend/clients/memorydb/memorydb.py +++ b/vectordb_bench/backend/clients/memorydb/memorydb.py @@ -32,7 +32,7 @@ def __init__( self.case_config = db_case_config self.collection_name = INDEX_NAME self.target_nodes = RedisCluster.RANDOM if not self.db_config["cmd"] else None - self.insert_batch_size = db_case_config.insert_batch_size + self.pipeline_batch_size = db_case_config.pipeline_batch_size self.dbsize = kwargs.get("num_rows") # Create a MemoryDB connection, if db has password configured, add it to the connection here and in init(): @@ -190,7 +190,7 @@ def insert_embeddings( }, ) # Execute the pipe so we don't keep too much in memory at once - if (i + 1) % self.insert_batch_size == 0: + if (i + 1) % self.pipeline_batch_size == 0: pipe.execute() pipe.execute() diff --git a/vectordb_bench/backend/dataset.py b/vectordb_bench/backend/dataset.py index 176c08e17..9f465577f 100644 --- a/vectordb_bench/backend/dataset.py +++ b/vectordb_bench/backend/dataset.py @@ -31,6 +31,7 @@ from .filter import Filter, FilterOp, non_filter log = logging.getLogger(__name__) +DEFAULT_INSERT_BATCH_SIZE = config.DEFAULT_INSERT_BATCH_SIZE class SizeLabel(NamedTuple): @@ -416,7 +417,10 @@ def _read_file(self, file_name: str) -> pl.DataFrame: class DataSetIterator: - def __init__(self, dataset: DatasetManager, batch_size: int = config.NUM_PER_BATCH): + def __init__(self, dataset: DatasetManager, batch_size: int = DEFAULT_INSERT_BATCH_SIZE): + if batch_size <= 0: + msg = f"insert batch size must be greater than 0, got {batch_size}" + raise ValueError(msg) self._ds = dataset self._batch_size = batch_size self._idx = 0 # file number @@ -1080,7 +1084,7 @@ def prepare( log.info(f"FTS dataset preparation completed: {self.data.full_name}") return True - def iter_batches(self, batch_size: int = config.NUM_PER_BATCH): + def iter_batches(self, batch_size: int = DEFAULT_INSERT_BATCH_SIZE): """Return an iterator for streaming FTS document batches.""" return FtsDocumentIterator(self, batch_size=batch_size) @@ -1107,7 +1111,10 @@ class FtsDocumentIterator: processing of large datasets. """ - def __init__(self, dataset: FtsDatasetManager, batch_size: int = config.NUM_PER_BATCH): + def __init__(self, dataset: FtsDatasetManager, batch_size: int = DEFAULT_INSERT_BATCH_SIZE): + if batch_size <= 0: + msg = f"insert batch size must be greater than 0, got {batch_size}" + raise ValueError(msg) self._ds = dataset self._batch_size = batch_size self._finished = False diff --git a/vectordb_bench/backend/runner/concurrent_runner.py b/vectordb_bench/backend/runner/concurrent_runner.py index e41af22ee..8b659e76e 100644 --- a/vectordb_bench/backend/runner/concurrent_runner.py +++ b/vectordb_bench/backend/runner/concurrent_runner.py @@ -33,6 +33,7 @@ from .executor import TaskExecutor log = logging.getLogger(__name__) +DEFAULT_INSERT_BATCH_SIZE = config.DEFAULT_INSERT_BATCH_SIZE class ExecutorBackend(StrEnum): @@ -65,12 +66,15 @@ def __init__( timeout: float | None = None, max_workers: int | None = None, backend: ExecutorBackend = ExecutorBackend.THREADING, - batch_size: int = config.NUM_PER_BATCH, + batch_size: int = DEFAULT_INSERT_BATCH_SIZE, duration: float | None = None, with_scalar_labels: bool = False, tenant_case=None, # noqa: ANN001 workload_kind: WorkloadKind = WorkloadKind.VECTOR, ): + if batch_size <= 0: + msg = f"insert batch size must be greater than 0, got {batch_size}" + raise ValueError(msg) self.timeout = timeout if isinstance(timeout, int | float) else None self.dataset: DatasetManager | FtsDatasetManager = dataset self.db = db diff --git a/vectordb_bench/backend/runner/mp_runner.py b/vectordb_bench/backend/runner/mp_runner.py index bf81d2a7e..b5a7fe17d 100644 --- a/vectordb_bench/backend/runner/mp_runner.py +++ b/vectordb_bench/backend/runner/mp_runner.py @@ -19,7 +19,6 @@ from ...models import ConcurrencySlotTimeoutError from ..clients import api -NUM_PER_BATCH = config.NUM_PER_BATCH log = logging.getLogger(__name__) # HDR Histogram constants diff --git a/vectordb_bench/backend/runner/rate_runner.py b/vectordb_bench/backend/runner/rate_runner.py index 91d0bb3ee..4b1f1a7b4 100644 --- a/vectordb_bench/backend/runner/rate_runner.py +++ b/vectordb_bench/backend/runner/rate_runner.py @@ -13,6 +13,7 @@ from .util import get_data log = logging.getLogger(__name__) +DEFAULT_INSERT_BATCH_SIZE = config.DEFAULT_INSERT_BATCH_SIZE class RatedMultiThreadingInsertRunner: @@ -23,13 +24,25 @@ def __init__( dataset_iter: DataSetIterator, normalize: bool = False, timeout: float | None = None, + batch_size: int = DEFAULT_INSERT_BATCH_SIZE, ): + if batch_size <= 0: + msg = f"insert batch size must be greater than 0, got {batch_size}" + raise ValueError(msg) + if rate <= 0: + msg = f"insert rate must be greater than 0, got {rate}" + raise ValueError(msg) + if rate % batch_size != 0: + msg = f"insert rate {rate} must be divisible by insert batch size {batch_size}" + raise ValueError(msg) + self.timeout = timeout if isinstance(timeout, int | float) else None self.dataset = dataset_iter self.db = db self.normalize = normalize self.insert_rate = rate - self.batch_rate = rate // config.NUM_PER_BATCH + self.batch_size = batch_size + self.batch_rate = rate // batch_size self.executing_futures = [] self.sig_idx = 0 diff --git a/vectordb_bench/backend/runner/read_write_runner.py b/vectordb_bench/backend/runner/read_write_runner.py index d3d1df2fa..8293128fe 100644 --- a/vectordb_bench/backend/runner/read_write_runner.py +++ b/vectordb_bench/backend/runner/read_write_runner.py @@ -8,6 +8,7 @@ import numpy as np +from vectordb_bench import config from vectordb_bench.backend.clients import api from vectordb_bench.backend.dataset import DatasetManager from vectordb_bench.backend.filter import Filter, non_filter @@ -19,6 +20,7 @@ from .serial_runner import SerialSearchRunner log = logging.getLogger(__name__) +DEFAULT_INSERT_BATCH_SIZE = config.DEFAULT_INSERT_BATCH_SIZE class ReadWriteRunner(MultiProcessingSearchRunner, RatedMultiThreadingInsertRunner): @@ -41,6 +43,7 @@ def __init__( optimize_after_write: bool = True, read_dur_after_write: int = 300, # seconds, search duration when insertion is done timeout: float | None = None, + batch_size: int = DEFAULT_INSERT_BATCH_SIZE, ): self.insert_rate = insert_rate self.data_volume = dataset.data.size @@ -75,8 +78,9 @@ def __init__( self, rate=insert_rate, db=db, - dataset_iter=iter(dataset), + dataset_iter=dataset.iter_batches(batch_size), normalize=normalize, + batch_size=batch_size, ) self.serial_search_runner = SerialSearchRunner( db=db, diff --git a/vectordb_bench/backend/runner/serial_runner.py b/vectordb_bench/backend/runner/serial_runner.py index f671d381a..3aba25c1c 100644 --- a/vectordb_bench/backend/runner/serial_runner.py +++ b/vectordb_bench/backend/runner/serial_runner.py @@ -19,8 +19,8 @@ from .. import utils from ..clients import api -NUM_PER_BATCH = config.NUM_PER_BATCH LOAD_MAX_TRY_COUNT = config.LOAD_MAX_TRY_COUNT +DEFAULT_INSERT_BATCH_SIZE = config.DEFAULT_INSERT_BATCH_SIZE log = logging.getLogger(__name__) @@ -36,29 +36,34 @@ def __init__( normalize: bool, filters: Filter = non_filter, timeout: float | None = None, + batch_size: int = DEFAULT_INSERT_BATCH_SIZE, ): + if batch_size <= 0: + msg = f"insert batch size must be greater than 0, got {batch_size}" + raise ValueError(msg) self.timeout = timeout if isinstance(timeout, int | float) else None self.dataset = dataset self.db = db self.normalize = normalize self.filters = filters + self.batch_size = batch_size def endless_insert_data(self, all_embeddings: list, all_metadata: list, left_id: int = 0) -> int: with self.db.init(): # unique id for endlessness insertion all_metadata = [i + left_id for i in all_metadata] - num_batches = math.ceil(len(all_embeddings) / NUM_PER_BATCH) + num_batches = math.ceil(len(all_embeddings) / self.batch_size) log.info( f"({mp.current_process().name:16}) Start inserting {len(all_embeddings)} " - f"embeddings in batch {NUM_PER_BATCH}" + f"embeddings in batch {self.batch_size}" ) count = 0 for batch_id in range(num_batches): retry_count = 0 already_insert_count = 0 - metadata = all_metadata[batch_id * NUM_PER_BATCH : (batch_id + 1) * NUM_PER_BATCH] - embeddings = all_embeddings[batch_id * NUM_PER_BATCH : (batch_id + 1) * NUM_PER_BATCH] + metadata = all_metadata[batch_id * self.batch_size : (batch_id + 1) * self.batch_size] + embeddings = all_embeddings[batch_id * self.batch_size : (batch_id + 1) * self.batch_size] log.debug( f"({mp.current_process().name:16}) batch [{batch_id:3}/{num_batches}], " @@ -88,7 +93,7 @@ def endless_insert_data(self, all_embeddings: list, all_metadata: list, left_id: count += already_insert_count log.info( f"({mp.current_process().name:16}) Finish inserting {len(all_embeddings)} embeddings in " - f"batch {NUM_PER_BATCH}" + f"batch {self.batch_size}" ) return count @@ -96,7 +101,7 @@ def run_endlessness(self) -> int: """run forever util DB raises exception or crash""" # datasets for load tests are quite small, can fit into memory # only 1 file - data_df = next(iter(self.dataset)) + data_df = next(self.dataset.iter_batches(self.batch_size)) all_embeddings, all_metadata = ( np.stack(data_df[self.dataset.data.train_vector_field]).tolist(), data_df[self.dataset.data.train_id_field].tolist(), diff --git a/vectordb_bench/backend/task_runner.py b/vectordb_bench/backend/task_runner.py index d66d6f454..d2fd4a149 100644 --- a/vectordb_bench/backend/task_runner.py +++ b/vectordb_bench/backend/task_runner.py @@ -8,7 +8,6 @@ import numpy as np -from .. import config from ..base import BaseModel from ..metric import Metric from ..models import PerformanceTimeoutError, TaskConfig, TaskStage @@ -82,6 +81,7 @@ def load_reuse_key(self) -> tuple | None: self._db_case_config_hash_key(), self._collection_name_hash_key(), self._dataset_hash_key(), + self.config.insert_batch_size, self.ca.with_scalar_labels, self.ca.is_multitenant, self._multitenant_routing_hash_key(), @@ -298,6 +298,7 @@ def _run_capacity_case(self) -> Metric: self.normalize, self.ca.filters, self.ca.load_timeout, + batch_size=self.config.insert_batch_size, ) count = runner.run_endlessness() except Exception as e: @@ -335,7 +336,7 @@ def _run_perf_case(self, drop_old: bool = True) -> Metric: m.load_duration = round(load_dur + build_dur, 4) m.additional_parameters.update( { - "num_per_batch": config.NUM_PER_BATCH, + "insert_batch_size": self.config.insert_batch_size, "load_concurrency": self.config.load_concurrency, } ) @@ -408,7 +409,7 @@ def _run_cloud_insert_case(self) -> Metric: self.normalize, self.ca.filters, max_workers=self.config.load_concurrency or None, - batch_size=self.ca.batch_size, + batch_size=self.config.insert_batch_size, duration=self.ca.duration, **runner_kwargs, ) @@ -511,6 +512,7 @@ def _load_train_data(self): self.ca.filters, self.ca.load_timeout, max_workers=self.config.load_concurrency or None, + batch_size=self.config.insert_batch_size, with_scalar_labels=self.ca.with_scalar_labels, workload_kind=self.workload_kind, **runner_kwargs, @@ -699,6 +701,7 @@ def _init_read_write_runner(self): concurrencies=ca.concurrencies, k=self.config.case_config.k, normalize=self.normalize, + batch_size=self.config.insert_batch_size, ) def stop(self): diff --git a/vectordb_bench/cli/cli.py b/vectordb_bench/cli/cli.py index abcb91607..fbc8cd0a2 100644 --- a/vectordb_bench/cli/cli.py +++ b/vectordb_bench/cli/cli.py @@ -233,6 +233,13 @@ def get_custom_case_config(parameters: dict) -> dict: "dataset_with_size_type": dataset_with_size_type, "label_percentage": parameters["label_percentage"], } + elif parameters["case_type"] in { + "StreamingPerformanceCase", + "StreamingCustomDataset", + }: + custom_case_config = { + "insert_rate": parameters["streaming_insert_rate"], + } elif parameters["case_type"] == "CloudPayloadSearchCase": custom_case_config = { "payload_profile": parameters["payload_profile"], @@ -252,7 +259,6 @@ def get_custom_case_config(parameters: dict) -> dict: copy_if_not_none(custom_case_config, parameters, "cloud_label_percentage", "label_percentage") elif parameters["case_type"] == "CloudInsertCase": custom_case_config = { - "batch_size": parameters["cloud_insert_batch_size"], "duration": parameters["cloud_insert_duration"], "dataset_with_size_type": dataset_with_size_type, } @@ -384,6 +390,26 @@ class CommonTypedDict(TypedDict): help="Number of concurrent workers for data loading in performance cases (0 = cpu_count)", ), ] + insert_batch_size: Annotated[ + int, + click.option( + "--insert-batch-size", + type=click.IntRange(min=1), + default=config.DEFAULT_INSERT_BATCH_SIZE, + show_default=True, + help="Rows or documents in each logical VDBBench insert batch; backends may split it further", + ), + ] + streaming_insert_rate: Annotated[ + int, + click.option( + "--streaming-insert-rate", + type=click.IntRange(min=1), + default=config.DEFAULT_STREAMING_INSERT_RATE, + show_default=True, + help="Rows inserted per second for StreamingPerformanceCase", + ), + ] search_serial: Annotated[ bool, click.option( @@ -706,16 +732,6 @@ class CommonTypedDict(TypedDict): help="Number of serial queries per cold/warm pass for CloudColdLatencyCase", ), ] - cloud_insert_batch_size: Annotated[ - int, - click.option( - "--cloud-insert-batch-size", - type=int, - default=5000, - show_default=True, - help="Insert batch size for CloudInsertCase", - ), - ] cloud_insert_duration: Annotated[ float | None, click.option( @@ -967,6 +983,7 @@ def run( parameters["search_concurrent"], ), load_concurrency=parameters["load_concurrency"], + insert_batch_size=parameters["insert_batch_size"], ) task_label = parameters["task_label"] diff --git a/vectordb_bench/frontend/components/run_test/generateTasks.py b/vectordb_bench/frontend/components/run_test/generateTasks.py index 5d848bb94..e78726381 100644 --- a/vectordb_bench/frontend/components/run_test/generateTasks.py +++ b/vectordb_bench/frontend/components/run_test/generateTasks.py @@ -3,7 +3,13 @@ from vectordb_bench.models import CaseConfig, CaseConfigParamType, TaskConfig -def generate_tasks(activedDbList: list[DB], dbConfigs, activedCaseList: list[CaseConfig], allCaseConfigs): +def generate_tasks( + activedDbList: list[DB], + dbConfigs, + activedCaseList: list[CaseConfig], + allCaseConfigs, + batch_size: int, +): tasks = [] for db in activedDbList: for case in activedCaseList: @@ -35,6 +41,7 @@ def generate_tasks(activedDbList: list[DB], dbConfigs, activedCaseList: list[Cas db_config=dbConfigs[db], case_config=case, db_case_config=db_case_config, + insert_batch_size=batch_size, ) tasks.append(task) diff --git a/vectordb_bench/frontend/components/run_test/runSettings.py b/vectordb_bench/frontend/components/run_test/runSettings.py new file mode 100644 index 000000000..ea459af5e --- /dev/null +++ b/vectordb_bench/frontend/components/run_test/runSettings.py @@ -0,0 +1,59 @@ +from vectordb_bench import config +from vectordb_bench.backend.cases import CaseType +from vectordb_bench.models import CaseConfig + +DEFAULT_INSERT_BATCH_SIZE = config.DEFAULT_INSERT_BATCH_SIZE +DEFAULT_STREAMING_INSERT_RATE = config.DEFAULT_STREAMING_INSERT_RATE +MAX_STREAMLIT_INT = (1 << 53) - 1 +STREAMING_CASE_TYPES = { + CaseType.StreamingPerformanceCase, + CaseType.StreamingCustomDataset, +} + + +def validate_streaming_insert_rates( + activedCaseList: list[CaseConfig], + batch_size: int, +) -> tuple[bool, list[str]]: + errors = [] + for case_config in activedCaseList: + if case_config.case_id not in STREAMING_CASE_TYPES: + continue + + custom_case = case_config.custom_case or {} + insert_rate = custom_case.get("insert_rate", DEFAULT_STREAMING_INSERT_RATE) + case_name = case_config.case_id.name + if insert_rate < batch_size: + errors.append( + f"{case_name}: Streaming Insert Rate ({insert_rate}) must be greater than or equal to " + f"Insert Batch Size ({batch_size})." + ) + elif insert_rate % batch_size != 0: + errors.append( + f"{case_name}: Streaming Insert Rate ({insert_rate}) must be divisible by " + f"Insert Batch Size ({batch_size})." + ) + + return len(errors) == 0, errors + + +def runSettings(container, activedCaseList: list[CaseConfig]) -> tuple[int, bool]: + container.markdown( + "
", + unsafe_allow_html=True, + ) + container.subheader("Run Settings") + batch_size = container.number_input( + "Insert Batch Size", + min_value=1, + max_value=MAX_STREAMLIT_INT, + value=DEFAULT_INSERT_BATCH_SIZE, + step=100, + help="Rows or documents in each logical VDBBench insert batch. Backends may split it further.", + ) + + is_valid, errors = validate_streaming_insert_rates(activedCaseList, batch_size) + for error in errors: + container.error(error) + + return batch_size, is_valid diff --git a/vectordb_bench/frontend/config/dbCaseConfigs.py b/vectordb_bench/frontend/config/dbCaseConfigs.py index daacae26a..2a74a415e 100644 --- a/vectordb_bench/frontend/config/dbCaseConfigs.py +++ b/vectordb_bench/frontend/config/dbCaseConfigs.py @@ -240,9 +240,10 @@ def generate_custom_streaming_case() -> CaseConfig: ), ConfigInput( label=CaseConfigParamType.insert_rate, + displayLabel="Streaming Insert Rate", inputType=InputType.Number, - inputConfig=dict(step=100, min=100, max=4_000, value=200), - inputHelp="fixed insertion rate (rows/s), must be divisible by 100", + inputConfig=dict(step=100, min=100, max=MAX_STREAMLIT_INT, value=500), + inputHelp="Fixed streaming insertion rate (rows/s); must be at least and divisible by Insert Batch Size.", ), ConfigInput( label=CaseConfigParamType.search_stages, diff --git a/vectordb_bench/frontend/pages/run_test.py b/vectordb_bench/frontend/pages/run_test.py index 64115ff17..6a995f48e 100644 --- a/vectordb_bench/frontend/pages/run_test.py +++ b/vectordb_bench/frontend/pages/run_test.py @@ -5,6 +5,7 @@ from vectordb_bench.frontend.components.run_test.generateTasks import generate_tasks from vectordb_bench.frontend.components.run_test.hideSidebar import hideSidebar from vectordb_bench.frontend.components.run_test.initStyle import initStyle +from vectordb_bench.frontend.components.run_test.runSettings import runSettings from vectordb_bench.frontend.components.run_test.submitTask import submitTask from vectordb_bench.frontend.components.check_results.nav import NavToResults, NavToPages from vectordb_bench.frontend.components.check_results.headerIcon import drawHeaderIcon @@ -46,8 +47,13 @@ def main(): caseSelectorContainer = st.container() activedCaseList, allCaseConfigs = caseSelector(caseSelectorContainer, activedDbList) + # run settings + runSettingsContainer = st.container() + batch_size, areRunSettingsValid = runSettings(runSettingsContainer, activedCaseList) + isAllValid = isAllValid and areRunSettingsValid + # generate tasks - tasks = generate_tasks(activedDbList, dbConfigs, activedCaseList, allCaseConfigs) if isAllValid else [] + tasks = generate_tasks(activedDbList, dbConfigs, activedCaseList, allCaseConfigs, batch_size) if isAllValid else [] # submit submitContainer = st.container() diff --git a/vectordb_bench/models.py b/vectordb_bench/models.py index d427d4f58..1c4924ed1 100644 --- a/vectordb_bench/models.py +++ b/vectordb_bench/models.py @@ -6,6 +6,7 @@ from typing import Any, ClassVar, Self import ujson +from pydantic import PositiveInt, model_validator from vectordb_bench.backend.cases import type2case from vectordb_bench.backend.dataset import DatasetWithSizeMap @@ -267,6 +268,38 @@ class TaskConfig(BaseModel): case_config: CaseConfig stages: list[TaskStage] = ALL_TASK_STAGES load_concurrency: int = config.LOAD_CONCURRENCY + insert_batch_size: PositiveInt = config.DEFAULT_INSERT_BATCH_SIZE + + @model_validator(mode="after") + def validate_streaming_insert_rate(self) -> Self: + streaming_case_types = { + CaseType.StreamingPerformanceCase, + CaseType.StreamingCustomDataset, + } + if self.case_config.case_id not in streaming_case_types: + return self + + custom_case = self.case_config.custom_case or {} + insert_rate = custom_case.get("insert_rate", config.DEFAULT_STREAMING_INSERT_RATE) + if not isinstance(insert_rate, int) or isinstance(insert_rate, bool) or insert_rate <= 0: + raise ValueError("streaming insert_rate must be a positive integer") + + rate_is_too_low = insert_rate < self.insert_batch_size + rate_is_divisible = insert_rate % self.insert_batch_size == 0 + + if rate_is_too_low: + msg = ( + f"streaming insert_rate ({insert_rate}) must be greater than or equal to " + f"insert_batch_size ({self.insert_batch_size})" + ) + raise ValueError(msg) + if not rate_is_divisible: + msg = ( + f"streaming insert_rate ({insert_rate}) must be divisible by " + f"insert_batch_size ({self.insert_batch_size})" + ) + raise ValueError(msg) + return self @property def db_name(self): @@ -430,7 +463,22 @@ def read_file(cls, full_path: pathlib.Path, trans_unit: bool = False) -> Self: for case_result in test_result["results"]: task_config = case_result.get("task_config") case_config = task_config.get("case_config") + metrics = case_result.get("metrics") db = DB(task_config.get("db")) + if "insert_batch_size" not in task_config: + insert_batch_size = None + if CaseType(case_config.get("case_id")) == CaseType.CloudInsertCase: + custom_case = case_config.get("custom_case") or {} + insert_batch_size = custom_case.get("batch_size") + if insert_batch_size is None and metrics: + additional_parameters = metrics.get("additional_parameters") or {} + insert_batch_size = additional_parameters.get( + "insert_batch_size", + additional_parameters.get("num_per_batch"), + ) + if insert_batch_size is not None: + task_config["insert_batch_size"] = insert_batch_size + task_config["db_config"] = db.config_cls(**task_config["db_config"]) # Safely instantiate DBCaseConfig (fallback to EmptyDBCaseConfig on None) raw_case_cfg = task_config.get("db_case_config") or {} @@ -447,7 +495,6 @@ def read_file(cls, full_path: pathlib.Path, trans_unit: bool = False) -> Self: task_config["case_config"] = cls.get_case_config(case_config=case_config) case_result["task_config"] = task_config - metrics = case_result.get("metrics") if ( metrics and CaseType(case_config.get("case_id")) == CaseType.CloudColdLatencyCase diff --git a/vectordb_bench/restful/app.py b/vectordb_bench/restful/app.py index ad0336501..a7de6cd2f 100644 --- a/vectordb_bench/restful/app.py +++ b/vectordb_bench/restful/app.py @@ -1,5 +1,6 @@ from flask import Flask, jsonify, request +from vectordb_bench import config from vectordb_bench.backend.clients import DB from vectordb_bench.interface import benchmark_runner from vectordb_bench.models import ALL_TASK_STAGES, CaseConfig, TaskConfig, TaskStage @@ -82,6 +83,7 @@ def run(): case_config=case_config, db_case_config=db_case_config, stages=stages, + insert_batch_size=task.get("insert_batch_size", config.DEFAULT_INSERT_BATCH_SIZE), ) task_configs.append(task_config) except Exception as e: From 6d26682326bccd05f97f17e5d8e322a5452f874b Mon Sep 17 00:00:00 2001 From: jamesgao-jpg Date: Wed, 12 Aug 2026 09:52:38 +0000 Subject: [PATCH 2/9] fix(milvus): honor configured insert batch size Signed-off-by: jamesgao-jpg --- tests/test_milvus.py | 49 ++++++++++- .../backend/clients/milvus/milvus.py | 85 ++++++++----------- 2 files changed, 82 insertions(+), 52 deletions(-) diff --git a/tests/test_milvus.py b/tests/test_milvus.py index ef131ef87..10731ef76 100644 --- a/tests/test_milvus.py +++ b/tests/test_milvus.py @@ -443,7 +443,6 @@ def insert(collection_name, batch_data): db = object.__new__(Milvus) db.client = SimpleNamespace(insert=insert) db.collection_name = "test_collection" - db.batch_size = 100 db._primary_field = "pk" db._scalar_id_field = "id" db._vector_field = "vector" @@ -465,3 +464,51 @@ def insert(collection_name, batch_data): {"pk": 1, "id": 1, "vector": [0.1, 0.2], "labels": "tenant_0001", "scalar_label": "label_a"}, {"pk": 2, "id": 2, "vector": [0.3, 0.4], "labels": "tenant_0002", "scalar_label": "label_b"}, ] + + +def test_milvus_vector_insert_uses_one_client_call_for_runner_batch() -> 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 + + count, err = db.insert_embeddings( + embeddings=[[float(i)] for i in range(5)], + metadata=list(range(5)), + ) + + assert count == 5 + assert err is None + client.insert.assert_called_once() + assert len(client.insert.call_args.args[1]) == 5 + + +def test_milvus_fts_insert_uses_one_client_call_for_runner_batch() -> None: + client = MagicMock() + client.insert.side_effect = lambda _collection, rows: {"insert_count": len(rows)} + + db = object.__new__(Milvus) + db.client = client + db.name = "Milvus" + db.collection_name = "test_collection" + db._is_fts = True + db._primary_field = "doc_id" + db._text_field = "text" + db._filter_id_field = "filter_id" + db.with_scalar_labels = False + + count, err = db.insert_documents( + texts=[f"document {i}" for i in range(5)], + doc_ids=[f"d{i}" for i in range(5)], + filter_ids=list(range(5)), + ) + + assert count == 5 + assert err is None + client.insert.assert_called_once() diff --git a/vectordb_bench/backend/clients/milvus/milvus.py b/vectordb_bench/backend/clients/milvus/milvus.py index 73c3c242e..71ab8f66c 100644 --- a/vectordb_bench/backend/clients/milvus/milvus.py +++ b/vectordb_bench/backend/clients/milvus/milvus.py @@ -16,8 +16,6 @@ log = logging.getLogger(__name__) -MILVUS_LOAD_REQS_SIZE = 1.5 * 1024 * 1024 -MILVUS_FTS_BATCH_SIZE = 1000 MILVUS_FORCE_MERGE_TARGET_SIZE_MB = ((1 << 63) - 1) // (1024**2) MILVUS_FORCE_MERGE_MAX_ATTEMPTS = 10 MILVUS_FORCE_MERGE_RETRY_INTERVAL_SECONDS = 30 @@ -46,7 +44,6 @@ def __init__( # noqa: PLR0915 drop_old: bool = False, name: str = "Milvus", with_scalar_labels: bool = False, - fts_batch_size: int | None = None, **kwargs, ): """Initialize wrapper around the milvus vector database.""" @@ -63,7 +60,6 @@ def __init__( # noqa: PLR0915 self._is_fts = isinstance(self.case_config, MilvusFtsConfig) if self._is_fts: - self.batch_size = fts_batch_size or MILVUS_FTS_BATCH_SIZE self._primary_field = "doc_id" self._text_field = "text" self._filter_id_field = "filter_id" @@ -75,7 +71,6 @@ def __init__( # noqa: PLR0915 self._sort_index_name = self._doc_id_sort_index_name self._sort_index_field = self._primary_field else: - self.batch_size = int(MILVUS_LOAD_REQS_SIZE / (dim * 4)) self._primary_field = "pk" self._scalar_id_field = "id" self._vector_field = "vector" @@ -403,31 +398,29 @@ def insert_embeddings( tenant_labels_data: list[str] | None = None, **kwargs, ) -> tuple[int, Exception]: - """Insert embeddings into Milvus. should call self.init() first""" + """Insert one runner-provided batch of embeddings into Milvus.""" assert self.client is not None assert len(embeddings) == len(metadata) - insert_count = 0 + + rows = [] + for i in range(len(embeddings)): + row = { + self._primary_field: metadata[i], + self._scalar_id_field: metadata[i], + self._vector_field: embeddings[i], + } + if tenant_labels_data is not None: + row[self._multitenant_partition_key_field] = tenant_labels_data[i] + if self.with_scalar_labels: + row[self._scalar_payload_label_field] = labels_data[i] + rows.append(row) + try: - for batch_start_offset in range(0, len(embeddings), self.batch_size): - batch_end_offset = min(batch_start_offset + self.batch_size, len(embeddings)) - batch_data = [] - for i in range(batch_start_offset, batch_end_offset): - row = { - self._primary_field: metadata[i], - self._scalar_id_field: metadata[i], - self._vector_field: embeddings[i], - } - if tenant_labels_data is not None: - row[self._multitenant_partition_key_field] = tenant_labels_data[i] - if self.with_scalar_labels: - row[self._scalar_payload_label_field] = labels_data[i] - batch_data.append(row) - res = self.client.insert(self.collection_name, batch_data) - insert_count += res["insert_count"] + res = self.client.insert(self.collection_name, rows) except MilvusException as e: log.info(f"Failed to insert data: {e}") - return insert_count, e - return insert_count, None + return 0, e + return res["insert_count"], None def insert_documents( self, @@ -435,7 +428,7 @@ def insert_documents( doc_ids: list[str], **kwargs, ) -> tuple[int, Exception | None]: - """Insert documents into a Milvus BM25 full-text collection.""" + """Insert one runner-provided batch into a Milvus BM25 collection.""" if not self._is_fts: msg = "insert_documents is only valid in FTS mode" raise RuntimeError(msg) @@ -446,40 +439,30 @@ def insert_documents( msg = f"Mismatch between texts ({len(docs)}) and doc_ids ({len(doc_ids)}) lengths" raise ValueError(msg) - batch_size = kwargs.get("batch_size", self.batch_size) labels_data = kwargs.get("labels_data") filter_ids = kwargs.get("filter_ids") if filter_ids is not None and len(filter_ids) != len(docs): msg = f"Mismatch between texts ({len(docs)}) and filter_ids ({len(filter_ids)}) lengths" raise ValueError(msg) - insert_count = 0 + rows = [] + for i, doc in enumerate(docs): + row = { + self._primary_field: str(doc_ids[i]), + self._text_field: doc, + } + if filter_ids is not None: + row[self._filter_id_field] = int(filter_ids[i]) + if self.with_scalar_labels: + row[self._scalar_label_field] = labels_data[i] if labels_data is not None else "" + rows.append(row) + try: - for batch_start_offset in range(0, len(docs), batch_size): - batch_end_offset = min(batch_start_offset + batch_size, len(docs)) - rows = [] - for i in range(batch_start_offset, batch_end_offset): - row = { - self._primary_field: str(doc_ids[i]), - self._text_field: docs[i], - } - if filter_ids is not None: - row[self._filter_id_field] = int(filter_ids[i]) - if self.with_scalar_labels: - row[self._scalar_label_field] = labels_data[i] if labels_data is not None else "" - rows.append(row) - - res = self.client.insert(self.collection_name, rows) - insert_count += res["insert_count"] - if batch_start_offset // batch_size % 10 == 0: - log.debug( - f"{self.name} batch insert progress: {batch_end_offset}/{len(docs)} " - f"({batch_end_offset / len(docs) * 100:.1f}%)" - ) + res = self.client.insert(self.collection_name, rows) except MilvusException as e: log.info(f"{self.name} insert error: {e}") - return insert_count, e - return insert_count, None + return 0, e + return res["insert_count"], None def prepare_filter(self, filters: Filter): if self._is_fts: From a1e46f3188320fe653e8f2310727ce69b6ee011e Mon Sep 17 00:00:00 2001 From: jamesgao-jpg Date: Wed, 12 Aug 2026 16:15:00 +0000 Subject: [PATCH 3/9] fix(opensearch): separate health request timeout Signed-off-by: jamesgao-jpg --- tests/test_oss_opensearch_fts.py | 1 + vectordb_bench/backend/clients/oss_opensearch/oss_opensearch.py | 2 ++ 2 files changed, 3 insertions(+) diff --git a/tests/test_oss_opensearch_fts.py b/tests/test_oss_opensearch_fts.py index 25b5b295c..d728aecaf 100644 --- a/tests/test_oss_opensearch_fts.py +++ b/tests/test_oss_opensearch_fts.py @@ -138,6 +138,7 @@ def health(self, **kwargs): "index": "idx", "wait_for_status": "green", "timeout": "30m", + "request_timeout": 1860, } diff --git a/vectordb_bench/backend/clients/oss_opensearch/oss_opensearch.py b/vectordb_bench/backend/clients/oss_opensearch/oss_opensearch.py index d26514497..12617a640 100644 --- a/vectordb_bench/backend/clients/oss_opensearch/oss_opensearch.py +++ b/vectordb_bench/backend/clients/oss_opensearch/oss_opensearch.py @@ -19,6 +19,7 @@ WAITING_FOR_REFRESH_SEC: Final[int] = 30 WAITING_FOR_FORCE_MERGE_SEC: Final[int] = 30 REPLICA_HEALTH_TIMEOUT: Final[str] = "30m" +REPLICA_HEALTH_REQUEST_TIMEOUT_SEC: Final[int] = 31 * 60 # Central registry for version-dependent OpenSearch index settings. # Add new rules here to automatically support future versions. @@ -813,6 +814,7 @@ def _wait_till_green(self): index=self.index_name, wait_for_status="green", timeout=REPLICA_HEALTH_TIMEOUT, + request_timeout=REPLICA_HEALTH_REQUEST_TIMEOUT_SEC, ) health = response.get("status", "unknown") if response.get("timed_out") or health != "green": From 5ad7b7d0cd2a0723136771c212b93988fd50bec4 Mon Sep 17 00:00:00 2001 From: Ubuntu Date: Thu, 13 Aug 2026 04:12:09 +0000 Subject: [PATCH 4/9] fix(opensearch): propagate vector metric type Signed-off-by: jamesgao-jpg --- tests/test_oss_opensearch_fts.py | 30 +++++++++++++++++++ .../backend/clients/oss_opensearch/cli.py | 1 + 2 files changed, 31 insertions(+) diff --git a/tests/test_oss_opensearch_fts.py b/tests/test_oss_opensearch_fts.py index d728aecaf..d8d8f91a6 100644 --- a/tests/test_oss_opensearch_fts.py +++ b/tests/test_oss_opensearch_fts.py @@ -46,6 +46,36 @@ def test_oss_opensearch_fts_config_defaults(): assert config.search_param() == {} +def test_oss_opensearch_cli_propagates_metric_type(monkeypatch): + captured = {} + + def fake_run(**kwargs): + captured.update(kwargs) + + monkeypatch.setattr(oss_opensearch_cli, "run", fake_run) + result = CliRunner().invoke( + oss_opensearch_cli.OSSOpenSearch, + [ + "--host", + "localhost", + "--metric-type", + "cosine", + "--m", + "16", + "--ef-construction", + "256", + "--ef-search", + "100", + "--case-type", + "Performance768D1M", + "--dry-run", + ], + ) + + assert result.exit_code == 0, result.output + assert captured["db_case_config"].metric_type_name == "cosine" + + def test_oss_opensearch_fts_config_supports_bm25_similarity(): config = OSSOpenSearchFtsConfig(bm25_k1=1.2, bm25_b=0.75) elastic_config = ElasticCloudFtsConfig(bm25_k1=1.2, bm25_b=0.75) diff --git a/vectordb_bench/backend/clients/oss_opensearch/cli.py b/vectordb_bench/backend/clients/oss_opensearch/cli.py index 03da18fdb..5c577d886 100644 --- a/vectordb_bench/backend/clients/oss_opensearch/cli.py +++ b/vectordb_bench/backend/clients/oss_opensearch/cli.py @@ -169,6 +169,7 @@ def OSSOpenSearch(**parameters: Unpack[OSSOpenSearchHNSWTypedDict]): flush_threshold_size=parameters["flush_threshold_size"], index_thread_qty_during_force_merge=parameters["index_thread_qty_during_force_merge"], cb_threshold=parameters["cb_threshold"], + metric_type_name=parameters["metric_type"], efConstruction=parameters["ef_construction"], efSearch=parameters["ef_search"], M=parameters["m"], From 1b8f842cf4241a38945f1aff701120b074b9bd2d Mon Sep 17 00:00:00 2001 From: jamesgao-jpg Date: Thu, 13 Aug 2026 12:41:45 +0000 Subject: [PATCH 5/9] fix(fts): bypass filter preparation for unfiltered loads Signed-off-by: jamesgao-jpg --- tests/test_fts_dataset.py | 32 +++++++++++++++++++++++++++++-- vectordb_bench/backend/dataset.py | 31 ++++++++++++++++++++++-------- 2 files changed, 53 insertions(+), 10 deletions(-) diff --git a/tests/test_fts_dataset.py b/tests/test_fts_dataset.py index a5eaf9791..1faca47b7 100644 --- a/tests/test_fts_dataset.py +++ b/tests/test_fts_dataset.py @@ -15,7 +15,7 @@ MSMarcoTranslator, SizeLabel, ) -from vectordb_bench.backend.filter import NewIntFilter +from vectordb_bench.backend.filter import NewIntFilter, non_filter @dataclass @@ -162,6 +162,7 @@ def test_fts_iterator_preserves_qrel_docs_before_filler(): manager._ir_dataset = FakeDataset() manager.required_doc_ids = {"d4"} manager.selected_doc_ids = manager._build_selected_doc_ids() + manager.filter_stats = {"filter_type": "NumGE"} assert manager.selected_doc_ids == {"d1", "d2", "d4"} @@ -174,6 +175,31 @@ def test_fts_iterator_preserves_qrel_docs_before_filler(): assert [filter_id for _, filter_id in docs] == [2, 0, 1] +def test_fts_unfiltered_iterator_omits_filter_ids(): + manager = make_tiny_msmarco_manager() + manager._ir_dataset = FakeDataset() + manager.selected_doc_ids = {"d1", "d2", "d3"} + + docs = [doc for batch in manager for doc in batch] + + assert [doc.filter_id for doc in docs] == [None, None, None] + + +def test_fts_prepare_non_filter_skips_selected_document_build(monkeypatch: pytest.MonkeyPatch): + manager = make_tiny_msmarco_manager(size=4) + monkeypatch.setattr(manager._translator, "load", FakeDataset) + monkeypatch.setattr( + manager, + "_build_selected_doc_ids", + lambda: pytest.fail("unfiltered FTS must not build selected document IDs"), + ) + + assert manager.prepare(source=None, filters=non_filter) + assert manager.selected_doc_ids is None + assert manager.recall_queries_data == manager.queries_data + assert manager.recall_gt_data == manager.gt_data + + def test_fts_qrel_filter_ids_match_sparse_emitted_documents_when_one_is_skipped(): class UnassignableDocument: doc_id = "d3" @@ -200,6 +226,7 @@ def iter_documents(self, dataset): # noqa: ARG002 manager._translator = SparseTranslator() manager.required_doc_ids = {"d1", "d5"} manager.selected_doc_ids = {"d1", "d3", "d5"} + manager.filter_stats = {"filter_type": "NumGE"} qrel_filter_ids = manager._build_qrel_filter_ids() emitted_filter_ids = {doc.doc_id: doc.filter_id for batch in manager for doc in batch} @@ -278,9 +305,10 @@ def test_fts_cap_rejects_required_qrel_docs_missing_from_corpus(): def test_fts_prepare_propagates_missing_required_qrel_docs(monkeypatch: pytest.MonkeyPatch): manager = FtsDatasetManager(data=MSMarcoFts(size=100_000)) monkeypatch.setattr(manager._translator, "load", FakeDatasetWithMissingQrel) + filters = NewIntFilter(filter_rate=0.5, int_field="filter_id", int_value=50_000) with pytest.raises(ValueError, match="missing from corpus"): - manager.prepare(source=None) + manager.prepare(source=None, filters=filters) def test_fts_dataset_size_registry(): diff --git a/vectordb_bench/backend/dataset.py b/vectordb_bench/backend/dataset.py index 9f465577f..e57647aa6 100644 --- a/vectordb_bench/backend/dataset.py +++ b/vectordb_bench/backend/dataset.py @@ -864,13 +864,13 @@ def _build_selected_doc_ids(self) -> set[str]: return selected_doc_ids - def _iter_selected_documents_with_filter_ids(self) -> Iterator[FtsDocument]: + def _iter_selected_documents_with_filter_ids(self, include_filter_ids: bool = False) -> Iterator[FtsDocument]: """Yield selected documents with the exact filter IDs used for insertion and qrels.""" if self._ir_dataset is None: msg = "ir_datasets dataset not loaded. Call prepare() first." raise RuntimeError(msg) - permutation = FtsFilterIdPermutation.for_size(self.data.size) + permutation = FtsFilterIdPermutation.for_size(self.data.size) if include_filter_ids else None documents = iter(self._translator.iter_documents(self._ir_dataset)) emitted_count = 0 while emitted_count < self.data.size: @@ -879,7 +879,8 @@ def _iter_selected_documents_with_filter_ids(self) -> Iterator[FtsDocument]: doc.doc_id = str(doc.doc_id) if self.selected_doc_ids is not None and doc.doc_id not in self.selected_doc_ids: continue - doc.filter_id = permutation.map(emitted_count) + if permutation is not None: + doc.filter_id = permutation.map(emitted_count) except StopIteration: break except Exception as e: @@ -897,7 +898,7 @@ def _build_qrel_filter_ids(self) -> dict[str, int]: qrel_doc_ids = set(self.required_doc_ids) qrel_filter_ids: dict[str, int] = {} - for doc in self._iter_selected_documents_with_filter_ids(): + for doc in self._iter_selected_documents_with_filter_ids(include_filter_ids=True): doc_id = doc.doc_id if doc_id in qrel_doc_ids: qrel_filter_ids[doc_id] = doc.filter_id @@ -1050,7 +1051,8 @@ def prepare( raise ValueError(msg) # noqa: TRY301 self.required_doc_ids = {doc_id for qrels in self.gt_data for doc_id in qrels} - self.selected_doc_ids = self._build_selected_doc_ids() + is_filtered = filters is not None and filters.type != FilterOp.NonFilter + self.selected_doc_ids = self._build_selected_doc_ids() if is_filtered else None self.recall_queries_data, self.recall_gt_data = self._apply_filters_to_qrels( self.queries_data, self.gt_data, @@ -1061,7 +1063,7 @@ def prepare( "selected %s corpus docs including %s qrel docs", len(self.gt_data), len(self.recall_gt_data), - len(self.selected_doc_ids), + len(self.selected_doc_ids) if self.selected_doc_ids is not None else self.data.size, len(self.required_doc_ids), ) else: @@ -1119,6 +1121,7 @@ def __init__(self, dataset: FtsDatasetManager, batch_size: int = DEFAULT_INSERT_ self._batch_size = batch_size self._finished = False self._docs_iter = None + self._doc_count = 0 def __iter__(self): return self @@ -1136,12 +1139,24 @@ def __next__(self) -> list[FtsDocument]: raise StopIteration if self._docs_iter is None: - self._docs_iter = self._ds._iter_selected_documents_with_filter_ids() + if self._ds.filter_stats: + self._docs_iter = self._ds._iter_selected_documents_with_filter_ids(include_filter_ids=True) + else: + self._docs_iter = self._ds._translator.iter_documents(self._ds._ir_dataset) batch = [] while len(batch) < self._batch_size: + if self._doc_count >= self._ds.data.size: + self._finished = True + if batch: + return batch + raise StopIteration try: - batch.append(next(self._docs_iter)) + doc = next(self._docs_iter) + if not self._ds.filter_stats: + doc.doc_id = str(doc.doc_id) + batch.append(doc) + self._doc_count += 1 except StopIteration: self._finished = True if batch: From 410b686b098e3b8515f9abdde797202d77dbec26 Mon Sep 17 00:00:00 2001 From: jamesgao-jpg Date: Thu, 13 Aug 2026 13:28:06 +0000 Subject: [PATCH 6/9] fix(fts): omit filter schema for unfiltered Milvus Signed-off-by: jamesgao-jpg --- tests/test_fts_filter_runner.py | 31 ++++++++++++++++++- tests/test_milvus.py | 27 ++++++++++++++++ .../backend/clients/milvus/milvus.py | 6 ++-- vectordb_bench/backend/task_runner.py | 3 ++ 4 files changed, 64 insertions(+), 3 deletions(-) diff --git a/tests/test_fts_filter_runner.py b/tests/test_fts_filter_runner.py index fe7179334..50803ae19 100644 --- a/tests/test_fts_filter_runner.py +++ b/tests/test_fts_filter_runner.py @@ -3,7 +3,7 @@ from vectordb_bench.backend.cases import CaseLabel from vectordb_bench.backend.data_source import DatasetSource from vectordb_bench.backend.dataset import FtsDocument, FtsQuery -from vectordb_bench.backend.filter import non_filter +from vectordb_bench.backend.filter import NewIntFilter, non_filter from vectordb_bench.backend.payload import PayloadProfile from vectordb_bench.backend.runner.concurrent_runner import ConcurrentInsertRunner from vectordb_bench.backend.task_runner import CaseRunner @@ -59,6 +59,35 @@ class Case: assert init_calls == [False] +def test_fts_init_db_passes_filter_schema_flag(): + captured = [] + + class Db: + def __init__(self, **kwargs): + captured.append(kwargs["fts_filter_enabled"]) + + class Config: + db = type("DbConfig", (), {"init_cls": Db})() + db_config = type("Config", (), {"to_dict": lambda self: {}})() + db_case_config = object() + + for filters in (non_filter, NewIntFilter(filter_rate=0.5, int_field="filter_id", int_value=50)): + case = type( + "Case", + (), + { + "label": CaseLabel.FullTextSearchPerformance, + "dataset": type("Dataset", (), {"data": object()})(), + "filters": filters, + "is_multitenant": False, + "with_scalar_labels": False, + }, + )() + CaseRunner.construct(ca=case, config=Config()).init_db() + + assert captured == [False, True] + + def test_fts_perf_metric_includes_dataset_filter_stats(): class Dataset: filter_stats = { diff --git a/tests/test_milvus.py b/tests/test_milvus.py index 10731ef76..2d6163276 100644 --- a/tests/test_milvus.py +++ b/tests/test_milvus.py @@ -512,3 +512,30 @@ def test_milvus_fts_insert_uses_one_client_call_for_runner_batch() -> None: assert count == 5 assert err is None client.insert.assert_called_once() + + +def test_milvus_fts_filter_index_is_conditional(monkeypatch: pytest.MonkeyPatch) -> None: + db = object.__new__(Milvus) + db._is_fts = True + db._sparse_field = "sparse_vector" + db._main_index_name = "sparse_vector_idx" + db._sort_index_field = "doc_id" + db._sort_index_name = "doc_id_sort_idx" + db._filter_id_field = "filter_id" + db._filter_id_sort_index_name = "filter_id_sort_idx" + db.with_scalar_labels = False + db.case_config = SimpleNamespace(sparse_index_param=lambda: {}) + + for enabled, expected_fields in ( + (False, {"sparse_vector", "doc_id"}), + (True, {"sparse_vector", "doc_id", "filter_id"}), + ): + params = MagicMock() + monkeypatch.setattr( + "vectordb_bench.backend.clients.milvus.milvus.MilvusClient.prepare_index_params", + lambda: params, + ) + db._fts_filter_enabled = enabled + db._build_index_params() + fields = {call.kwargs["field_name"] for call in params.add_index.call_args_list} + assert fields == expected_fields diff --git a/vectordb_bench/backend/clients/milvus/milvus.py b/vectordb_bench/backend/clients/milvus/milvus.py index 71ab8f66c..42a949f4e 100644 --- a/vectordb_bench/backend/clients/milvus/milvus.py +++ b/vectordb_bench/backend/clients/milvus/milvus.py @@ -58,6 +58,7 @@ def __init__( # noqa: PLR0915 self._multitenant_partition_key_field = self._scalar_label_field self._scalar_labels_index_name = "labels_idx" self._is_fts = isinstance(self.case_config, MilvusFtsConfig) + self._fts_filter_enabled = bool(kwargs.get("fts_filter_enabled", self._is_fts)) if self._is_fts: self._primary_field = "doc_id" @@ -107,7 +108,8 @@ def __init__( # noqa: PLR0915 else self.case_config.index_param().get("analyzer_params", {"type": "english"}) ) schema.add_field(self._primary_field, DataType.VARCHAR, max_length=512, is_primary=True) - schema.add_field(self._filter_id_field, DataType.INT64) + if self._fts_filter_enabled: + schema.add_field(self._filter_id_field, DataType.INT64) schema.add_field( self._text_field, DataType.VARCHAR, @@ -200,7 +202,7 @@ def _build_index_params(self): index_name=self._sort_index_name, index_type="STL_SORT", ) - if self._is_fts: + if self._is_fts and self._fts_filter_enabled: index_params.add_index( field_name=self._filter_id_field, index_name=self._filter_id_sort_index_name, diff --git a/vectordb_bench/backend/task_runner.py b/vectordb_bench/backend/task_runner.py index d2fd4a149..9c69f7078 100644 --- a/vectordb_bench/backend/task_runner.py +++ b/vectordb_bench/backend/task_runner.py @@ -15,6 +15,7 @@ from .cases import Case, CaseLabel, StreamingPerformanceCase from .clients import DB, MetricType, api from .data_source import DatasetSource +from .filter import FilterOp from .runner import ( ColdWarmSearchRunner, ConcurrentInsertRunner, @@ -200,6 +201,8 @@ def init_db(self, drop_old: bool = True) -> None: extra_db_kwargs["collection_name"] = collection_name if self.ca.is_multitenant: extra_db_kwargs["multitenant_tenant_labels"] = self.ca.tenant_labels() + if self.is_fts: + extra_db_kwargs["fts_filter_enabled"] = self.ca.filters.type != FilterOp.NonFilter self.db = db_cls( dim=getattr(self.ca.dataset.data, "dim", 0), From 8d297c88423e0d90b87e9b8b21f1c4046f89bc4c Mon Sep 17 00:00:00 2001 From: jamesgao-jpg Date: Mon, 17 Aug 2026 09:01:59 +0000 Subject: [PATCH 7/9] fix(aws_opensearch): honor unified insert_batch_size for serverless bulk Merging #843 reintroduced config.NUM_PER_BATCH, which the insert-batch-size control feature removed. Route the AOSS serverless bulk chunk size through the task-level insert_batch_size instead, keeping the default of 100 and the non-positive guard. Signed-off-by: jamesgao-jpg --- README.md | 6 +++--- tests/test_aws_opensearch.py | 15 ++++++--------- .../clients/aws_opensearch/aws_opensearch.py | 5 +++-- vectordb_bench/backend/task_runner.py | 2 ++ 4 files changed, 14 insertions(+), 14 deletions(-) diff --git a/README.md b/README.md index 53dfa93f5..b08458c79 100644 --- a/README.md +++ b/README.md @@ -287,7 +287,7 @@ OpenSearch Serverless (AOSS) is a serverless deployment option for Amazon OpenSe **Example: Run performance test on OpenSearch Serverless** ```shell -NUM_PER_BATCH=100 vectordbbench awsopensearch --db-label aoss \ +vectordbbench awsopensearch --db-label aoss --insert-batch-size 100 \ --serverless --aws-region us-east-1 \ --host .aoss.us-east-1.on.aws --port 443 \ --case-type Performance768D1M \ @@ -303,13 +303,13 @@ OpenSearch Serverless-specific options: |--------|-------------| | `--serverless` | Enable OpenSearch Serverless mode (uses AWS SigV4 auth) | | `--aws-region` | AWS region for the AOSS collection (default: `us-east-1`) | -| `NUM_PER_BATCH` | Number of vectors per Serverless bulk request (default: `100`) | +| `--insert-batch-size` | Number of vectors per Serverless bulk request (default: `100`) | > **Notes:** > - `--user` and `--password` are not needed for Serverless mode > - `--engine` is accepted but ignored internally (AOSS manages the engine) > - `--force-merge-enabled`, `--refresh-interval`, `--flush-threshold-size`, and `--cb-threshold` are ignored for Serverless -> - Keep `NUM_PER_BATCH` small enough for the Serverless bulk API request limits +> - Keep `--insert-batch-size` small enough for the Serverless bulk API request limits ### Run Elastic Cloud from command line diff --git a/tests/test_aws_opensearch.py b/tests/test_aws_opensearch.py index 2e2930338..4a2bafa1a 100644 --- a/tests/test_aws_opensearch.py +++ b/tests/test_aws_opensearch.py @@ -2,21 +2,19 @@ import pytest -from vectordb_bench import config from vectordb_bench.backend.clients.aws_opensearch.aws_opensearch import AWSOpenSearch -def test_serverless_insert_uses_configured_batch_size(monkeypatch) -> None: +def test_serverless_insert_uses_configured_batch_size() -> None: bulk_requests = [] - def bulk(*, body): + def bulk(*, body: list) -> None: bulk_requests.append(body) - monkeypatch.setattr(config, "NUM_PER_BATCH", 2) - db = object.__new__(AWSOpenSearch) db.client = SimpleNamespace(bulk=bulk) db._is_serverless = True + db._insert_batch_size = 2 db.index_name = "test-index" db.vector_col_name = "embedding" db.with_scalar_labels = False @@ -33,13 +31,12 @@ def bulk(*, body): @pytest.mark.parametrize("batch_size", [0, -1]) -def test_serverless_insert_rejects_non_positive_batch_size(monkeypatch, batch_size: int) -> None: - monkeypatch.setattr(config, "NUM_PER_BATCH", batch_size) - +def test_serverless_insert_rejects_non_positive_batch_size(batch_size: int) -> None: db = object.__new__(AWSOpenSearch) db._is_serverless = True + db._insert_batch_size = batch_size - with pytest.raises(ValueError, match="NUM_PER_BATCH must be greater than 0"): + with pytest.raises(ValueError, match="insert_batch_size must be greater than 0"): db._insert_with_single_client( embeddings=[[0.1]], metadata=[1], diff --git a/vectordb_bench/backend/clients/aws_opensearch/aws_opensearch.py b/vectordb_bench/backend/clients/aws_opensearch/aws_opensearch.py index eb0195b88..f6774b1b3 100644 --- a/vectordb_bench/backend/clients/aws_opensearch/aws_opensearch.py +++ b/vectordb_bench/backend/clients/aws_opensearch/aws_opensearch.py @@ -46,6 +46,7 @@ def __init__( self.label_col_name = label_col_name self.vector_col_name = vector_col_name self.with_scalar_labels = with_scalar_labels + self._insert_batch_size = kwargs.get("insert_batch_size", config.DEFAULT_INSERT_BATCH_SIZE) log.info(f"AWS_OpenSearch client config: {self.db_config}") log.info(f"AWS_OpenSearch db case config : {self.case_config}") @@ -257,9 +258,9 @@ def _insert_with_single_client( labels_data: list[str] | None = None, ) -> tuple[int, Exception]: embeddings_list = list(embeddings) - batch_size = config.NUM_PER_BATCH if self._is_serverless else len(embeddings_list) + batch_size = self._insert_batch_size if self._is_serverless else len(embeddings_list) if self._is_serverless and batch_size <= 0: - raise ValueError("NUM_PER_BATCH must be greater than 0 for OpenSearch Serverless") + raise ValueError("insert_batch_size must be greater than 0 for OpenSearch Serverless") total_inserted = 0 for i in range(0, len(embeddings_list), batch_size): diff --git a/vectordb_bench/backend/task_runner.py b/vectordb_bench/backend/task_runner.py index 9c69f7078..4b9314e9d 100644 --- a/vectordb_bench/backend/task_runner.py +++ b/vectordb_bench/backend/task_runner.py @@ -203,6 +203,8 @@ def init_db(self, drop_old: bool = True) -> None: extra_db_kwargs["multitenant_tenant_labels"] = self.ca.tenant_labels() if self.is_fts: extra_db_kwargs["fts_filter_enabled"] = self.ca.filters.type != FilterOp.NonFilter + if self.config.db is DB.AWSOpenSearch: + extra_db_kwargs["insert_batch_size"] = self.config.insert_batch_size self.db = db_cls( dim=getattr(self.ca.dataset.data, "dim", 0), From 64771cda9d90e6be4f9f57f4cfa48cb3281ee5ed Mon Sep 17 00:00:00 2001 From: jamesgao-jpg Date: Mon, 24 Aug 2026 03:59:19 +0000 Subject: [PATCH 8/9] fix(fts): prepare qrel-preserving capped corpus Materialize capped qrel-preserving documents during dataset preparation and stream the prepared corpus for filtered and unfiltered loads. This keeps semantic ground truth retrievable without scanning the full raw corpus during timed insertion. Signed-off-by: jamesgao-jpg --- tests/test_fts_dataset.py | 31 +++++++--- vectordb_bench/backend/dataset.py | 96 ++++++++++++++++++++++++------- 2 files changed, 100 insertions(+), 27 deletions(-) diff --git a/tests/test_fts_dataset.py b/tests/test_fts_dataset.py index 1faca47b7..8806d7fbb 100644 --- a/tests/test_fts_dataset.py +++ b/tests/test_fts_dataset.py @@ -71,6 +71,12 @@ def __init__(self): self.qrels = [Qrel("q1", "d3", 1), Qrel("q2", "d4", 2)] +class FakeDatasetWithLateQrel(FakeDataset): + def __init__(self): + super().__init__() + self.qrels = [Qrel("q1", "d4", 1)] + + def make_tiny_msmarco_manager(size: int = 3) -> FtsDatasetManager: small_label = MSMarcoFts._size_label[100_000] @@ -185,17 +191,21 @@ def test_fts_unfiltered_iterator_omits_filter_ids(): assert [doc.filter_id for doc in docs] == [None, None, None] -def test_fts_prepare_non_filter_skips_selected_document_build(monkeypatch: pytest.MonkeyPatch): - manager = make_tiny_msmarco_manager(size=4) - monkeypatch.setattr(manager._translator, "load", FakeDataset) +def test_fts_prepare_non_filter_materializes_qrel_preserving_corpus(monkeypatch: pytest.MonkeyPatch): + manager = make_tiny_msmarco_manager(size=3) + monkeypatch.setattr(manager._translator, "load", FakeDatasetWithLateQrel) + + assert manager.prepare(source=None, filters=non_filter) monkeypatch.setattr( - manager, - "_build_selected_doc_ids", - lambda: pytest.fail("unfiltered FTS must not build selected document IDs"), + manager._translator, + "iter_documents", + lambda dataset: pytest.fail("timed insertion must use the prepared corpus"), ) + docs = [doc for batch in manager for doc in batch] - assert manager.prepare(source=None, filters=non_filter) assert manager.selected_doc_ids is None + assert [doc.doc_id for doc in docs] == ["d1", "d2", "d4"] + assert [doc.filter_id for doc in docs] == [None, None, None] assert manager.recall_queries_data == manager.queries_data assert manager.recall_gt_data == manager.gt_data @@ -241,6 +251,12 @@ def test_fts_prepare_integer_filter_derives_filtered_qrels(monkeypatch: pytest.M filters = NewIntFilter(filter_rate=0.5, int_field="filter_id", int_value=2) assert manager.prepare(source=None, filters=filters) + monkeypatch.setattr( + manager._translator, + "iter_documents", + lambda dataset: pytest.fail("timed insertion must use the prepared corpus"), + ) + emitted_filter_ids = {doc.doc_id: doc.filter_id for batch in manager for doc in batch} assert [query.query_id for query in manager.queries_data] == ["q1", "q2"] assert manager.gt_data == [{"d3": 1}, {"d1": 2}] @@ -249,6 +265,7 @@ def test_fts_prepare_integer_filter_derives_filtered_qrels(monkeypatch: pytest.M assert manager.recall_skipped is False assert manager.recall_skip_reason is None assert manager.qrel_filter_ids == {"d1": 3, "d3": 1} + assert emitted_filter_ids == {"d1": 3, "d2": 2, "d3": 1, "d4": 0} assert manager.filter_stats == { "filter_type": "NumGE", "filter_field": "filter_id", diff --git a/vectordb_bench/backend/dataset.py b/vectordb_bench/backend/dataset.py index e57647aa6..f797f5bed 100644 --- a/vectordb_bench/backend/dataset.py +++ b/vectordb_bench/backend/dataset.py @@ -4,9 +4,11 @@ >>> Dataset.Cohere.get(100_000) """ +import json import logging import math import pathlib +import tempfile import types import typing from abc import ABC, abstractmethod @@ -794,6 +796,8 @@ class FtsDatasetManager(BaseModel): qrel_filter_ids: dict[str, int] = PydanticField(default_factory=dict) filter_stats: dict[str, int | float | str] = PydanticField(default_factory=dict) _ir_dataset: typing.Any = PrivateAttr(default=None) + _prepared_documents_dir: typing.Any = PrivateAttr(default=None) + _prepared_documents_path: pathlib.Path | None = PrivateAttr(default=None) def __init__(self, **data): super().__init__(**data) @@ -864,6 +868,69 @@ def _build_selected_doc_ids(self) -> set[str]: return selected_doc_ids + def _prepare_qrel_preserving_documents(self) -> None: + """Materialize capped documents before timed insertion.""" + if self._prepared_documents_dir is not None: + self._prepared_documents_dir.cleanup() + self._prepared_documents_dir = None + self._prepared_documents_path = None + + if self.data.size == max(self.data._size_label): + return + + required_doc_ids = set(self.required_doc_ids) + self._validate_cap(required_doc_ids=required_doc_ids, target_size=self.data.size) + filler_limit = self.data.size - len(required_doc_ids) + filler_count = 0 + selected_doc_ids: set[str] = set() + found_required_doc_ids: set[str] = set() + prepared_dir = tempfile.TemporaryDirectory(prefix="vdbbench_fts_qrel_v1_") + prepared_path = pathlib.Path(prepared_dir.name, f"{self.data.dir_name}.jsonl") + + try: + with prepared_path.open("w", encoding="utf-8") as output: + for doc in self._translator.iter_documents(self._ir_dataset): + doc_id = str(doc.doc_id) + if doc_id in selected_doc_ids: + continue + if doc_id in required_doc_ids: + found_required_doc_ids.add(doc_id) + elif filler_count < filler_limit: + filler_count += 1 + else: + continue + selected_doc_ids.add(doc_id) + output.write(json.dumps([doc_id, doc.text], ensure_ascii=False) + "\n") + if len(selected_doc_ids) == self.data.size and found_required_doc_ids == required_doc_ids: + break + + missing_doc_ids = required_doc_ids - found_required_doc_ids + if missing_doc_ids: + preview = ", ".join(sorted(missing_doc_ids)[:10]) + msg = ( + f"{self.data.full_name} semantic qrel docs missing from corpus: {preview}" + f"{'...' if len(missing_doc_ids) > 10 else ''}" + ) + raise ValueError(msg) # noqa: TRY301 + if len(selected_doc_ids) != self.data.size: + msg = f"{self.data.full_name} prepared {len(selected_doc_ids)} documents, expected {self.data.size}" + raise ValueError(msg) # noqa: TRY301 + except Exception: + prepared_dir.cleanup() + raise + + self._prepared_documents_dir = prepared_dir + self._prepared_documents_path = prepared_path + + def _iter_prepared_documents(self) -> Iterator[FtsDocument]: + if self._prepared_documents_path is None: + yield from self._translator.iter_documents(self._ir_dataset) + return + with self._prepared_documents_path.open(encoding="utf-8") as prepared: + for line in prepared: + doc_id, text = json.loads(line) + yield FtsDocument(doc_id=doc_id, text=text) + def _iter_selected_documents_with_filter_ids(self, include_filter_ids: bool = False) -> Iterator[FtsDocument]: """Yield selected documents with the exact filter IDs used for insertion and qrels.""" if self._ir_dataset is None: @@ -871,7 +938,7 @@ def _iter_selected_documents_with_filter_ids(self, include_filter_ids: bool = Fa raise RuntimeError(msg) permutation = FtsFilterIdPermutation.for_size(self.data.size) if include_filter_ids else None - documents = iter(self._translator.iter_documents(self._ir_dataset)) + documents = iter(self._iter_prepared_documents()) emitted_count = 0 while emitted_count < self.data.size: try: @@ -892,8 +959,8 @@ def _iter_selected_documents_with_filter_ids(self, include_filter_ids: bool = Fa def _build_qrel_filter_ids(self) -> dict[str, int]: """Map qrel doc IDs to their deterministic permuted FTS filter ID.""" - if self.selected_doc_ids is None: - msg = "selected_doc_ids is required before building FTS filter IDs" + if self._ir_dataset is None: + msg = "ir_datasets dataset not loaded. Call prepare() first." raise RuntimeError(msg) qrel_doc_ids = set(self.required_doc_ids) @@ -1051,8 +1118,8 @@ def prepare( raise ValueError(msg) # noqa: TRY301 self.required_doc_ids = {doc_id for qrels in self.gt_data for doc_id in qrels} - is_filtered = filters is not None and filters.type != FilterOp.NonFilter - self.selected_doc_ids = self._build_selected_doc_ids() if is_filtered else None + self.selected_doc_ids = None + self._prepare_qrel_preserving_documents() self.recall_queries_data, self.recall_gt_data = self._apply_filters_to_qrels( self.queries_data, self.gt_data, @@ -1121,7 +1188,6 @@ def __init__(self, dataset: FtsDatasetManager, batch_size: int = DEFAULT_INSERT_ self._batch_size = batch_size self._finished = False self._docs_iter = None - self._doc_count = 0 def __iter__(self): return self @@ -1139,24 +1205,14 @@ def __next__(self) -> list[FtsDocument]: raise StopIteration if self._docs_iter is None: - if self._ds.filter_stats: - self._docs_iter = self._ds._iter_selected_documents_with_filter_ids(include_filter_ids=True) - else: - self._docs_iter = self._ds._translator.iter_documents(self._ds._ir_dataset) + self._docs_iter = self._ds._iter_selected_documents_with_filter_ids( + include_filter_ids=bool(self._ds.filter_stats), + ) batch = [] while len(batch) < self._batch_size: - if self._doc_count >= self._ds.data.size: - self._finished = True - if batch: - return batch - raise StopIteration try: - doc = next(self._docs_iter) - if not self._ds.filter_stats: - doc.doc_id = str(doc.doc_id) - batch.append(doc) - self._doc_count += 1 + batch.append(next(self._docs_iter)) except StopIteration: self._finished = True if batch: From a22c48a66e1a91b6a37e4a1a8fd064b3ab86bcd8 Mon Sep 17 00:00:00 2001 From: jamesgao-jpg Date: Mon, 24 Aug 2026 07:44:41 +0000 Subject: [PATCH 9/9] fix(frontend): allow small streaming insert rates Lower the streaming insert rate control minimum to one and rely on the existing batch/rate validation for valid combinations. Signed-off-by: jamesgao-jpg --- tests/test_frontend_run_settings.py | 13 ++++++++++++- vectordb_bench/frontend/config/dbCaseConfigs.py | 2 +- 2 files changed, 13 insertions(+), 2 deletions(-) diff --git a/tests/test_frontend_run_settings.py b/tests/test_frontend_run_settings.py index 4c421aa92..8e4b2a1eb 100644 --- a/tests/test_frontend_run_settings.py +++ b/tests/test_frontend_run_settings.py @@ -9,7 +9,8 @@ DEFAULT_STREAMING_INSERT_RATE, validate_streaming_insert_rates, ) -from vectordb_bench.models import CaseConfig +from vectordb_bench.frontend.config.dbCaseConfigs import custom_streaming_config +from vectordb_bench.models import CaseConfig, CaseConfigParamType def streaming_case(insert_rate: int | None = None) -> CaseConfig: @@ -49,6 +50,16 @@ def test_validate_streaming_insert_rates_checks_each_streaming_case_and_uses_def assert errors == [] +def test_streaming_rate_control_accepts_valid_values_below_100(): + rate_input = next(item for item in custom_streaming_config if item.label == CaseConfigParamType.insert_rate) + + is_valid, errors = validate_streaming_insert_rates([streaming_case(50)], batch_size=10) + + assert rate_input.inputConfig["min"] == 1 + assert is_valid + assert errors == [] + + def test_generate_tasks_passes_batch_size_to_task_config(monkeypatch: pytest.MonkeyPatch): captured_task_configs: list[dict[str, object]] = [] diff --git a/vectordb_bench/frontend/config/dbCaseConfigs.py b/vectordb_bench/frontend/config/dbCaseConfigs.py index 2a74a415e..84ed76f8e 100644 --- a/vectordb_bench/frontend/config/dbCaseConfigs.py +++ b/vectordb_bench/frontend/config/dbCaseConfigs.py @@ -242,7 +242,7 @@ def generate_custom_streaming_case() -> CaseConfig: label=CaseConfigParamType.insert_rate, displayLabel="Streaming Insert Rate", inputType=InputType.Number, - inputConfig=dict(step=100, min=100, max=MAX_STREAMLIT_INT, value=500), + inputConfig=dict(step=100, min=1, max=MAX_STREAMLIT_INT, value=500), inputHelp="Fixed streaming insertion rate (rows/s); must be at least and divisible by Insert Batch Size.", ), ConfigInput(