diff --git a/pyrit/datasets/seed_datasets/remote/aegis_ai_content_safety_dataset.py b/pyrit/datasets/seed_datasets/remote/aegis_ai_content_safety_dataset.py index da85afcf78..1e438cdcd1 100644 --- a/pyrit/datasets/seed_datasets/remote/aegis_ai_content_safety_dataset.py +++ b/pyrit/datasets/seed_datasets/remote/aegis_ai_content_safety_dataset.py @@ -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 @@ -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__( @@ -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]: @@ -251,10 +254,6 @@ 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( @@ -262,25 +261,34 @@ async def _fetch_dataset_async(self, *, cache: bool = True) -> SeedDataset: ): 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.") diff --git a/pyrit/datasets/seed_datasets/remote/beaver_tails_dataset.py b/pyrit/datasets/seed_datasets/remote/beaver_tails_dataset.py index 6c511c4a64..f5a51eef50 100644 --- a/pyrit/datasets/seed_datasets/remote/beaver_tails_dataset.py +++ b/pyrit/datasets/seed_datasets/remote/beaver_tails_dataset.py @@ -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__( @@ -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") diff --git a/pyrit/datasets/seed_datasets/remote/pku_safe_rlhf_dataset.py b/pyrit/datasets/seed_datasets/remote/pku_safe_rlhf_dataset.py index 9a684b6f73..be05ea7441 100644 --- a/pyrit/datasets/seed_datasets/remote/pku_safe_rlhf_dataset.py +++ b/pyrit/datasets/seed_datasets/remote/pku_safe_rlhf_dataset.py @@ -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__( @@ -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"]) @@ -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") diff --git a/tests/unit/datasets/test_aegis_ai_content_safety_dataset.py b/tests/unit/datasets/test_aegis_ai_content_safety_dataset.py index 22fa6011b1..e9f39fbfb7 100644 --- a/tests/unit/datasets/test_aegis_ai_content_safety_dataset.py +++ b/tests/unit/datasets/test_aegis_ai_content_safety_dataset.py @@ -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", + } diff --git a/tests/unit/datasets/test_beaver_tails_dataset.py b/tests/unit/datasets/test_beaver_tails_dataset.py index 32068b4bfc..81ec0b3669 100644 --- a/tests/unit/datasets/test_beaver_tails_dataset.py +++ b/tests/unit/datasets/test_beaver_tails_dataset.py @@ -1,6 +1,7 @@ # Copyright (c) Microsoft Corporation. # Licensed under the MIT license. +import json from unittest.mock import AsyncMock, patch import pytest @@ -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, + } diff --git a/tests/unit/datasets/test_pku_safe_rlhf_dataset.py b/tests/unit/datasets/test_pku_safe_rlhf_dataset.py index cf47dc79ff..0ad8399651 100644 --- a/tests/unit/datasets/test_pku_safe_rlhf_dataset.py +++ b/tests/unit/datasets/test_pku_safe_rlhf_dataset.py @@ -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"}