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
6 changes: 5 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,13 @@ 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

# A requested axis cannot match metadata that does not declare it.
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_match_when_another_axis_matches(self):
"""An undeclared axis fails even when another filtered axis matches."""
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
Loading