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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions tests/test_cloud_cold_latency_case.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,7 +55,9 @@ def test_cloud_cold_latency_case_accepts_label_filter():
case = CloudColdLatencyCase(label_percentage=0.9)

assert case.label_percentage == 0.9
assert case.filter_rate == pytest.approx(0.1)
assert case.filters.type == FilterOp.StrEqual
assert case.filters.filter_rate == pytest.approx(0.1)


def test_cloud_cold_latency_case_rejects_two_filter_types():
Expand Down
12 changes: 12 additions & 0 deletions tests/test_cloud_payload_case.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
from vectordb_bench.backend.clients.api import EmptyDBCaseConfig
from vectordb_bench.backend.data_source import DatasetSource
from vectordb_bench.backend.dataset import DatasetWithSizeType
from vectordb_bench.backend.filter import FilterOp
from vectordb_bench.backend.payload import PayloadProfile
from vectordb_bench.backend.runner.mp_runner import MultiProcessingSearchRunner
from vectordb_bench.backend.runner.serial_runner import SerialSearchRunner
Expand Down Expand Up @@ -101,6 +102,17 @@ def test_case_runner_reuse_key_distinguishes_scalar_label_schema_requirement():
assert hash(ids_only_runner) != hash(scalar_label_runner)


def test_cloud_payload_label_filter_sets_filter_rate():
case = CloudPayloadSearchCase(
dataset_with_size_type=DatasetWithSizeType.CohereMedium.value,
label_percentage=0.2,
)

assert case.filter_rate == pytest.approx(0.8)
assert case.filters.type == FilterOp.StrEqual
assert case.filters.filter_rate == pytest.approx(0.8)


def test_cli_propagates_cloud_payload_dataset_selection():
custom_case = get_custom_case_config(
{
Expand Down
49 changes: 49 additions & 0 deletions tests/test_custom_dataset_filter.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,49 @@
import pytest

from vectordb_bench.backend.cases import CaseType, PerformanceCustomDataset
from vectordb_bench.backend.filter import LabelFilter
from vectordb_bench.models import CaseConfig


def _custom_case_kwargs(**overrides):
kwargs = {
"name": "custom-ds",
"description": "",
"load_timeout": 1.0,
"optimize_timeout": 1.0,
"dataset_config": {
"name": "custom-ds",
"dir": "/tmp/custom-ds",
"size": 1000,
"dim": 128,
"metric_type": "L2",
"file_count": 1,
},
}
kwargs.update(overrides)
return kwargs


def test_custom_dataset_filter_sets_filter_rate_from_label_percentage():
case = PerformanceCustomDataset(**_custom_case_kwargs(use_filter=True, label_percentage=0.01))
assert case.filter_rate == pytest.approx(0.99)
assert isinstance(case.filters, LabelFilter)
assert case.filters.filter_rate == pytest.approx(0.99)


def test_custom_dataset_filter_matches_label_filter_formula():
case = PerformanceCustomDataset(**_custom_case_kwargs(use_filter=True, label_percentage=0.2))
assert case.filter_rate == pytest.approx(LabelFilter(label_percentage=0.2).filter_rate)


def test_custom_dataset_without_filter_leaves_filter_rate_unset():
case = PerformanceCustomDataset(**_custom_case_kwargs(use_filter=False))
assert case.filter_rate is None


def test_case_config_builds_custom_dataset_with_filter_rate():
case = CaseConfig(
case_id=CaseType.PerformanceCustomDataset,
custom_case=_custom_case_kwargs(use_filter=True, label_percentage=0.5),
).case
assert case.filter_rate == pytest.approx(0.5)
41 changes: 41 additions & 0 deletions tests/test_filter_charts.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,41 @@
from unittest.mock import MagicMock, patch

import pytest

from vectordb_bench.frontend.components.int_filter import charts as int_charts
from vectordb_bench.frontend.components.label_filter import charts as label_charts

NONE_FILTER_DATA = [
{
"filter_rate": None,
"qps": 10.0,
"recall": 0.9,
"db_name": "A",
"dataset_name": "custom",
},
{
"filter_rate": 0.99,
"qps": 20.0,
"recall": 0.8,
"db_name": "B",
"dataset_name": "custom",
},
]


@pytest.mark.parametrize("charts", [label_charts, int_charts])
def test_get_range_coerces_none_filter_rate(charts):
data = [{"filter_rate": None}, {"filter_rate": 0.5}]
xrange = charts.getRange("filter_rate", data, [0.05, 0.1])
assert xrange[0] <= 0
assert xrange[1] >= 0.5


@pytest.mark.parametrize("charts", [label_charts, int_charts])
def test_draw_chart_does_not_crash_when_filter_rate_is_none(charts):
st = MagicMock()
with patch.object(charts, "px") as px:
px.line.return_value = MagicMock()
charts.drawChart(st, list(NONE_FILTER_DATA), "qps")
px.line.assert_called_once()
st.plotly_chart.assert_called_once()
2 changes: 2 additions & 0 deletions tests/test_multitenant_case.py
Original file line number Diff line number Diff line change
Expand Up @@ -65,7 +65,9 @@ def test_case_config_constructs_multitenant_case():

assert isinstance(case, CloudMultiTenantSearchCase)
assert case.payload_profile == PayloadProfile.SCALAR_LABEL
assert case.filter_rate == pytest.approx(0.99)
assert case.filters.type == FilterOp.StrEqual
assert case.filters.filter_rate == pytest.approx(0.99)
assert case.tenant_labels() == ["tenant_0000", "tenant_0001", "tenant_0002", "tenant_0003", "tenant_0004"]


Expand Down
12 changes: 12 additions & 0 deletions vectordb_bench/backend/cases.py
Original file line number Diff line number Diff line change
Expand Up @@ -439,6 +439,11 @@ def __init__(
gt_neighbors_field=dataset_config.gt_col_name,
scalar_labels_file=f"{dataset_config.scalar_labels_name}.parquet",
)
filter_rate = (
LabelFilter(label_percentage=label_percentage).filter_rate
if (use_filter and label_percentage is not None)
else None
)
super().__init__(
name=name,
description=description,
Expand All @@ -448,6 +453,7 @@ def __init__(
dataset=DatasetManager(data=dataset),
use_filter=use_filter,
label_percentage=label_percentage,
filter_rate=filter_rate,
)

@property
Expand Down Expand Up @@ -670,6 +676,8 @@ def __init__(
"Cloud leaderboard search envelope case with explicit response payload profile. "
f"Payload profile: {payload_profile.value}; dataset: {dataset_name}."
)
if label_percentage is not None:
filter_rate = LabelFilter(label_percentage=label_percentage).filter_rate
super().__init__(
name=name,
description=description,
Expand Down Expand Up @@ -739,6 +747,8 @@ def __init__(
"Cloud leaderboard cold/warm serial latency case with explicit response payload profile. "
f"Payload profile: {payload_profile.value}; dataset: {dataset_name}; query count: {query_count}."
)
if label_percentage is not None:
filter_rate = LabelFilter(label_percentage=label_percentage).filter_rate
super().__init__(
name=name,
description=description,
Expand Down Expand Up @@ -840,6 +850,8 @@ def __init__(
raise ValueError(msg)

dataset = dataset_with_size_type.get_manager()
if label_percentage is not None:
filter_rate = LabelFilter(label_percentage=label_percentage).filter_rate
super().__init__(
name=f"Cloud Multi-Tenant Search - {dataset_with_size_type.value}, {tenant_count} tenants",
description=(
Expand Down
8 changes: 5 additions & 3 deletions vectordb_bench/frontend/components/int_filter/charts.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,8 +21,10 @@ def drawChartByMetric(st, data, metrics=("qps", "recall"), **kwargs):


def getRange(metric, data, padding_multipliers):
minV = min([d.get(metric, 0) for d in data])
maxV = max([d.get(metric, 0) for d in data])
# dict.get(metric, 0) still returns None when the key is present with a None value.
values = [d.get(metric) or 0 for d in data]
minV = min(values)
maxV = max(values)
padding = maxV - minV
rangeV = [
minV - padding * padding_multipliers[0],
Expand All @@ -39,7 +41,7 @@ def drawChart(st, data: list[object], metric):
y = metric
yrange = getRange(y, data, [0.2, 0.1])

data.sort(key=lambda a: a[x])
data.sort(key=lambda a: a.get(x) or 0)

fig = px.line(
data,
Expand Down
8 changes: 5 additions & 3 deletions vectordb_bench/frontend/components/label_filter/charts.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,8 +21,10 @@ def drawChartByMetric(st, data, metrics=("qps", "recall"), **kwargs):


def getRange(metric, data, padding_multipliers):
minV = min([d.get(metric, 0) for d in data])
maxV = max([d.get(metric, 0) for d in data])
# dict.get(metric, 0) still returns None when the key is present with a None value.
values = [d.get(metric) or 0 for d in data]

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Two small notes on this coercion (identical in int_filter/charts.py):

  1. A one-line rationale comment would help — d.get(metric) or 0 is easy to "simplify" back to d.get(metric, 0), which reintroduces the crash for present-but-None values.

  2. For the legacy filter_rate=None rows this is meant to keep rendering: plotly converts None x to NaN, so the point is dropped from the line (not drawn at 0), while the axis still extends to 0. The run is invisible rather than misplaced. Once the producers are fixed (see comment on cases.py:442) this path is mostly unreachable; consider dropping unknown rows before the axis computation instead of padding the range with 0.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I added a comment that dict.get(metric, 0) still returns None when the key is present. I left the coerce in place so leftover result files with filter_rate=None still do not TypeError on min or sort. Dropping those rows would change the axis for mixed leftover files, which I left out of this follow-up.

minV = min(values)
maxV = max(values)
padding = maxV - minV
rangeV = [
minV - padding * padding_multipliers[0],
Expand All @@ -39,7 +41,7 @@ def drawChart(st, data: list[object], metric):
y = metric
yrange = getRange(y, data, [0.2, 0.1])

data.sort(key=lambda a: a[x])
data.sort(key=lambda a: a.get(x) or 0)

fig = px.line(
data,
Expand Down
Loading