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
11 changes: 10 additions & 1 deletion pyrit/datasets/seed_datasets/seed_dataset_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -222,9 +222,18 @@ def _match_single_criterion(
filter_vals = getattr(criterion, field.name)
meta_vals = getattr(metadata, field.name)

if filter_vals is None or meta_vals is None:
if filter_vals is None:
continue

# `meta_vals is None` means the dataset never declared this axis, which
# is not the same as declaring it and matching. Skipping here dropped
# the filtered axis entirely, so every dataset silent on that axis came
# back as a match. It also disagreed with get_all_dataset_names_async,
# which excludes datasets carrying no metadata at all - a dataset with
# partial metadata was treated as better qualified than one with none.
if meta_vals is None:
return False

if strict_match:
if filter_vals - meta_vals:
return False
Expand Down
33 changes: 33 additions & 0 deletions tests/unit/datasets/test_seed_dataset_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -371,6 +371,39 @@ def test_modalities(self):
dataset_filter=SeedDatasetFilter(modalities={"audio"}),
)

def test_undeclared_axis_does_not_match(self):
"""A dataset silent on a filtered axis does not satisfy that axis."""
metadata = SeedDatasetMetadata(tags={"safety"}, size={"small"}, source_type={"local"})
assert not SeedDatasetProvider._match_filter_to_metadata(
metadata=metadata,
dataset_filter=SeedDatasetFilter(modalities={"audio"}),
)
assert not SeedDatasetProvider._match_filter_to_metadata(
metadata=metadata,
dataset_filter=SeedDatasetFilter(harm_categories={"violence"}),
)
# An axis the dataset does declare is still matched normally.
assert SeedDatasetProvider._match_filter_to_metadata(
metadata=metadata,
dataset_filter=SeedDatasetFilter(tags={"safety"}),
)

def test_undeclared_axis_does_not_match_strict(self):
"""strict_match also rejects a dataset silent on the filtered axis."""
metadata = SeedDatasetMetadata(tags={"safety"}, size={"small"}, source_type={"local"})
assert not SeedDatasetProvider._match_filter_to_metadata(
metadata=metadata,
dataset_filter=SeedDatasetFilter(modalities={"audio"}, strict_match=True),
)

def test_undeclared_axis_does_not_mask_a_declared_mismatch(self):
"""Two filtered axes, one declared and mismatching: still no match."""
metadata = SeedDatasetMetadata(modalities={"text"})
assert not SeedDatasetProvider._match_filter_to_metadata(
metadata=metadata,
dataset_filter=SeedDatasetFilter(modalities={"text"}, harm_categories={"violence"}),
)

def test_sources(self):
"""Source filter checks membership."""
metadata = SeedDatasetMetadata(source_type={"remote"})
Expand Down