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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
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(1.0 - 0.2)


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 vectordb_bench/backend/cases.py
Original file line number Diff line number Diff line change
Expand Up @@ -439,6 +439,7 @@ def __init__(
gt_neighbors_field=dataset_config.gt_col_name,
scalar_labels_file=f"{dataset_config.scalar_labels_name}.parquet",
)
filter_rate = (1.0 - label_percentage) if (use_filter and label_percentage is not None) else None
super().__init__(
name=name,
description=description,
Expand All @@ -448,6 +449,7 @@ def __init__(
dataset=DatasetManager(data=dataset),
use_filter=use_filter,
label_percentage=label_percentage,
filter_rate=filter_rate,
)

@property
Expand Down
7 changes: 4 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,9 @@ 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])
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 +40,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
7 changes: 4 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,9 @@ 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])
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 +40,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