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
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
import asyncio
import logging
from enum import Enum
from typing import Any
from uuid import uuid4

from typing_extensions import override
Expand Down Expand Up @@ -130,7 +131,7 @@ class _AegisContentSafetyDataset(_RemoteDatasetLoader):
HF_DATASET_NAME: str = "nvidia/Aegis-AI-Content-Safety-Dataset-2.0"
harm_categories: list[str] = [c.value.lower() for c in AegisHarmCategory]
modalities: tuple[Modality, ...] = (Modality.TEXT,)
size: str = "huge" # 19093 annotated human-LLM interactions across all splits after filtering
size: str = "huge" # 13246 unique unsafe prompts from 19093 annotated rows across all splits
tags: frozenset[str] = frozenset({"default", "safety"})

def __init__(
Expand Down Expand Up @@ -231,7 +232,9 @@ async def _fetch_dataset_async(self, *, cache: bool = True) -> SeedDataset:
],
}

seed_prompts: list[SeedUnion] = []
# A prompt can appear in several rows (one per labeled response). Keep one seed per prompt, with the
# first row's metadata and the categories of every kept row.
merged: dict[str, tuple[dict[str, Any], list[str]]] = {}

for split_name in hf_dataset:
for example in hf_dataset[split_name]:
Expand All @@ -251,36 +254,41 @@ async def _fetch_dataset_async(self, *, cache: bool = True) -> SeedDataset:
if violated_categories
else []
)
standardized_categories = self._standardize_harm_categories(
prompt_harm_categories,
alias_overrides=alias_overrides,
)

# Filter by harm_categories if specified
if self._selected_category_values is not None and not any(
cat in self._selected_category_values for cat in prompt_harm_categories
):
continue

seed_prompts.append(
SeedPrompt(
value=prompt_value,
data_type="text",
dataset_name=self.dataset_name,
harm_categories=standardized_categories if standardized_categories else None,
source=self.source,
authors=self._AUTHORS,
groups=self._GROUPS,
metadata={
"id": example.get("id"),
"prompt_label": example.get("prompt_label"),
"response_label": example.get("response_label"),
"prompt_label_source": example.get("prompt_label_source"),
"response_label_source": example.get("response_label_source"),
"aegis_violated_categories": ", ".join(prompt_harm_categories),
},
)
_, merged_categories = merged.setdefault(prompt_value, (example, []))
merged_categories.extend(cat for cat in prompt_harm_categories if cat not in merged_categories)

seed_prompts: list[SeedUnion] = []
for prompt_value, (example, prompt_harm_categories) in merged.items():
standardized_categories = self._standardize_harm_categories(
prompt_harm_categories,
alias_overrides=alias_overrides,
)
seed_prompts.append(
SeedPrompt(
value=prompt_value,
data_type="text",
dataset_name=self.dataset_name,
harm_categories=standardized_categories if standardized_categories else None,
source=self.source,
authors=self._AUTHORS,
groups=self._GROUPS,
metadata={
"id": example.get("id"),
"prompt_label": example.get("prompt_label"),
"response_label": example.get("response_label"),
"prompt_label_source": example.get("prompt_label_source"),
"response_label_source": example.get("response_label_source"),
"aegis_violated_categories": ", ".join(prompt_harm_categories),
},
)
)

if not seed_prompts:
raise ValueError("SeedDataset cannot be empty. Check your filter criteria.")
Expand Down
57 changes: 32 additions & 25 deletions pyrit/datasets/seed_datasets/remote/beaver_tails_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -63,7 +63,7 @@ class _BeaverTailsDataset(_RemoteDatasetLoader):

# Metadata
modalities: tuple[Modality, ...] = (Modality.TEXT,)
size: str = "huge" # 166382 annotated prompt-response entries (default config)
size: str = "huge" # 14402 unique unsafe prompts from 166382 prompt-response entries (default config)
tags: frozenset[str] = frozenset({"default", "safety"})

def __init__(
Expand Down Expand Up @@ -130,35 +130,42 @@ async def _fetch_dataset_async(self, *, cache: bool = True) -> SeedDataset:
"Center on Frontiers of Computing Studies, School of Computer Science, Peking University",
]

seed_prompts = []
# BeaverTails has one row per prompt-response pair, so a prompt repeats once per response. Keep one
# seed per prompt and merge the harm labels its responses got.
merged: dict[str, tuple[list[str], dict[str, bool]]] = {}
for item in data:
if self.unsafe_only and item["is_safe"]:
continue

raw_harm_categories = [
part.strip() for k, v in item["category"].items() if v for part in k.split(",") if part.strip()
]
harm_categories = self._standardize_harm_categories(
raw_harm_categories,
alias_overrides=self.HARM_CATEGORY_ALIAS_OVERRIDES,
)

seed_prompts.append(
SeedPrompt(
value=item["prompt"],
data_type="text",
dataset_name=self.dataset_name,
harm_categories=harm_categories,
description=description,
source=source_url,
authors=authors,
groups=groups,
metadata={
"beaver_tails_categories": ",".join(raw_harm_categories),
"beaver_tails_category_flags": json.dumps(item["category"], sort_keys=True),
},
)
raw_harm_categories, category_flags = merged.setdefault(item["prompt"], ([], {}))
for key, flagged in item["category"].items():
category_flags[key] = category_flags.get(key, False) or bool(flagged)
if flagged:
for part in key.split(","):
part = part.strip()
if part and part not in raw_harm_categories:
raw_harm_categories.append(part)

seed_prompts = [
SeedPrompt(
value=prompt,
data_type="text",
dataset_name=self.dataset_name,
harm_categories=self._standardize_harm_categories(
raw_harm_categories,
alias_overrides=self.HARM_CATEGORY_ALIAS_OVERRIDES,
),
description=description,
source=source_url,
authors=authors,
groups=groups,
metadata={
"beaver_tails_categories": ",".join(raw_harm_categories),
"beaver_tails_category_flags": json.dumps(category_flags, sort_keys=True),
},
)
for prompt, (raw_harm_categories, category_flags) in merged.items()
]

logger.info(f"Successfully loaded {len(seed_prompts)} prompts from BeaverTails dataset")

Expand Down
44 changes: 24 additions & 20 deletions pyrit/datasets/seed_datasets/remote/pku_safe_rlhf_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,7 +50,7 @@ class _PKUSafeRLHFDataset(_RemoteDatasetLoader):

# Metadata
modalities: tuple[Modality, ...] = (Modality.TEXT,)
size: str = "huge" # 73907 prompt-response pairs across 19 harm categories
size: str = "huge" # 38641 unique prompts from 73907 prompt-response pairs across 19 harm categories
tags: frozenset[str] = frozenset({"default", "safety"})

def __init__(
Expand Down Expand Up @@ -146,7 +146,9 @@ async def _fetch_dataset_async(self, *, cache: bool = True) -> SeedDataset:
"Violence": [HarmCategory.VIOLENT_CONTENT],
"White-Collar Crime": [HarmCategory.SCAMS, HarmCategory.DECEPTION],
}
seed_prompts: list[SeedPrompt] = []
# PKU-SafeRLHF has one row per response pair, so a prompt can repeat many times. Keep one seed per
# prompt and merge the harm categories of every kept row.
harm_categories_by_prompt: dict[str, set[str]] = {}

for item in data:
is_unsafe = not (item["is_response_0_safe"] and item["is_response_1_safe"])
Expand All @@ -169,26 +171,28 @@ async def _fetch_dataset_async(self, *, cache: bool = True) -> SeedDataset:
if not self.filter_harm_categories or any(
category in self.filter_harm_categories for category in harm_categories
):
standardized_harm_categories = self._standardize_harm_categories(
harm_categories_by_prompt.setdefault(item["prompt"], set()).update(harm_categories)

seed_prompts = [
SeedPrompt(
value=prompt,
data_type="text",
dataset_name=self.dataset_name,
harm_categories=self._standardize_harm_categories(
sorted(harm_categories),
alias_overrides=harm_category_alias_overrides,
)
seed_prompts.append(
SeedPrompt(
value=item["prompt"],
data_type="text",
dataset_name=self.dataset_name,
harm_categories=standardized_harm_categories,
description=(
"This is a Hugging Face dataset that labels a prompt and 2 responses categorizing "
"their helpfulness or harmfulness. Only the 'prompt' column is extracted."
),
source=f"https://huggingface.co/datasets/{self.source}",
authors=self._AUTHORS,
groups=self._GROUPS,
metadata=({"pku_categories": ", ".join(sorted(harm_categories))} if harm_categories else None),
)
)
),
description=(
"This is a Hugging Face dataset that labels a prompt and 2 responses categorizing "
"their helpfulness or harmfulness. Only the 'prompt' column is extracted."
),
source=f"https://huggingface.co/datasets/{self.source}",
authors=self._AUTHORS,
groups=self._GROUPS,
metadata=({"pku_categories": ", ".join(sorted(harm_categories))} if harm_categories else None),
)
for prompt, harm_categories in harm_categories_by_prompt.items()
]

logger.info(f"Successfully loaded {len(seed_prompts)} prompts from PKU-SafeRLHF dataset")

Expand Down
28 changes: 28 additions & 0 deletions tests/unit/datasets/test_aegis_ai_content_safety_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -235,3 +235,31 @@ def test_init_empty_harm_categories_raises():
def test_invalid_harm_category_raises():
with pytest.raises(ValueError, match="Expected AegisHarmCategory"):
_AegisContentSafetyDataset(harm_categories=["Malware"])


async def test_fetch_dataset_merges_repeated_prompts():
"""A prompt labeled in several rows should load as one seed with the first row's metadata."""

def row(row_id: str, categories: str) -> dict[str, str]:
return {"id": row_id, "prompt": "Same prompt", "prompt_label": "unsafe", "violated_categories": categories}

rows = {
"train": [row("1", "Violence"), row("2", "Violence, Harassment")],
"test": [row("3", "Criminal Planning/Confessions")],
}
loader = _AegisContentSafetyDataset()

with patch.object(loader, "_fetch_from_huggingface_async", new_callable=AsyncMock, return_value=rows):
dataset = await loader.fetch_dataset_async()

assert len(dataset.seeds) == 1
assert dataset.seeds[0].metadata["id"] == "1"
assert dataset.seeds[0].metadata["aegis_violated_categories"] == (
"Violence, Harassment, Criminal Planning/Confessions"
)
assert set(dataset.seeds[0].harm_categories) == {
"VIOLENT_CONTENT",
"VIOLENT_THREATS",
"COORDINATION_HARM",
"HARASSMENT",
}
25 changes: 25 additions & 0 deletions tests/unit/datasets/test_beaver_tails_dataset.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT license.

import json
from unittest.mock import AsyncMock, patch

import pytest
Expand Down Expand Up @@ -174,3 +175,27 @@ def test_harm_category_alias_overrides_cover_beaver_tails_leaf_labels(self):
)
== expected
)


async def test_fetch_dataset_merges_repeated_prompts():
"""BeaverTails repeats a prompt once per response; each prompt should be one seed with merged labels."""
flags_a = {"animal_abuse": True, "financial_crime,property_crime,theft": False}
flags_b = {"animal_abuse": False, "financial_crime,property_crime,theft": True}
rows = [
{"prompt": "Same prompt", "response": "a", "category": flags_a, "is_safe": False},
{"prompt": "Same prompt", "response": "b", "category": flags_b, "is_safe": False},
{"prompt": "Same prompt", "response": "c", "category": flags_b, "is_safe": True},
{"prompt": "Other prompt", "response": "d", "category": flags_a, "is_safe": False},
]
loader = _BeaverTailsDataset()

with patch.object(loader, "_fetch_from_huggingface_async", new=AsyncMock(return_value=rows)):
dataset = await loader.fetch_dataset_async()

assert [seed.value for seed in dataset.seeds] == ["Same prompt", "Other prompt"]
merged = dataset.seeds[0]
assert merged.metadata["beaver_tails_categories"] == "animal_abuse,financial_crime,property_crime,theft"
assert json.loads(merged.metadata["beaver_tails_category_flags"]) == {
"animal_abuse": True,
"financial_crime,property_crime,theft": True,
}
28 changes: 28 additions & 0 deletions tests/unit/datasets/test_pku_safe_rlhf_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -121,3 +121,31 @@ async def test_fetch_dataset_standardizes_all_native_harm_categories(native_labe

assert len(dataset.seeds) == 1
assert dataset.seeds[0].harm_categories == expected_categories


async def test_fetch_dataset_merges_repeated_prompts():
"""A prompt appears once per response pair; it should load as one seed with every pair's categories."""
rows = [
{
"prompt": "Same prompt",
"is_response_0_safe": False,
"is_response_1_safe": True,
"response_0_harm_category": {"Cybercrime": True, "Violence": False},
"response_1_harm_category": {"Cybercrime": False, "Violence": False},
},
{
"prompt": "Same prompt",
"is_response_0_safe": True,
"is_response_1_safe": False,
"response_0_harm_category": {"Cybercrime": False, "Violence": False},
"response_1_harm_category": {"Cybercrime": False, "Violence": True},
},
]
loader = _PKUSafeRLHFDataset()

with patch.object(loader, "_fetch_from_huggingface_async", new_callable=AsyncMock, return_value=rows):
dataset = await loader.fetch_dataset_async()

assert len(dataset.seeds) == 1
assert dataset.seeds[0].metadata == {"pku_categories": "Cybercrime, Violence"}
assert set(dataset.seeds[0].harm_categories) == {"MALWARE", "COORDINATION_HARM", "VIOLENT_CONTENT"}
Loading