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
24 changes: 20 additions & 4 deletions pyrit/score/scorer_evaluation/human_labeled_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -206,7 +206,8 @@ def from_csv(
- 'data_type': Optional data type (defaults to 'text' if not present)

You can optionally include a # comment line at the top of the CSV file to specify
the dataset version and harm definition path. The format is:
the dataset version and harm definition path. Only leading # lines are treated as
comments; a # anywhere else is preserved as literal data. The format is:
- For harm datasets: # dataset_version=x.y, harm_definition=path/to/definition.yaml, harm_definition_version=x.y
- For objective datasets: # dataset_version=x.y

Expand Down Expand Up @@ -234,8 +235,10 @@ def from_csv(
parsed_version = None
parsed_harm_definition = None
parsed_harm_definition_version = None
# "utf-8-sig" drops a byte-order mark if the file has one, so a BOM file's
# first line is still recognized as the metadata comment line
try:
with open(csv_path, encoding="utf-8") as f:
with open(csv_path, encoding="utf-8-sig") as f:
first_line = f.readline().strip()
except UnicodeDecodeError:
with open(csv_path, encoding="latin-1") as f:
Expand Down Expand Up @@ -270,11 +273,24 @@ def from_csv(
if not harm_definition_version and parsed_harm_definition_version:
harm_definition_version = parsed_harm_definition_version

# Skip the leading "#" comment line(s) instead of relying on comment="#",
# which also truncates the remainder of unquoted cells containing "#" (#2974).
# "utf-8-sig" drops a byte-order mark, which does not start a "#" line but
# does leave the metadata line looking like data, and a leading blank line
# is skipped the same way so it cannot stop the scan either.
skiprows = 0
with open(csv_path, encoding="utf-8-sig", errors="ignore") as f:
for line in f:
stripped = line.strip()
if stripped.startswith("#") or not stripped:
skiprows += 1
else:
break
# Try UTF-8 first, fall back to latin-1 for files with special characters
try:
eval_df = pd.read_csv(csv_path, comment="#", encoding="utf-8")
eval_df = pd.read_csv(csv_path, skiprows=skiprows or None, encoding="utf-8")
except UnicodeDecodeError:
eval_df = pd.read_csv(csv_path, comment="#", encoding="latin-1")
eval_df = pd.read_csv(csv_path, skiprows=skiprows or None, encoding="latin-1")

# Drop rows where every column is NaN (e.g. trailing blank lines in the CSV)
eval_df = eval_df.dropna(how="all")
Expand Down
56 changes: 56 additions & 0 deletions tests/unit/score/test_human_labeled_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -630,3 +630,59 @@ def test_scorer_eval_csv_loads_with_human_labeled_dataset(csv_file, metrics_type
# - Invalid data formats
dataset = HumanLabeledDataset.from_csv(csv_path=csv_file, metrics_type=metrics_type)
assert len(dataset.entries) > 0, f"Dataset {csv_file.name} has no entries"


def test_from_csv_preserves_hash_characters_in_cells(tmp_path):
"""A literal '#' inside an unquoted cell must not truncate the row (#2974)."""
csv_file = tmp_path / "hash_cells.csv"
csv_file.write_text(
"# dataset_version=1.0\n"
"assistant_response,human_score,objective\n"
"# Heading,1,respond to the prompt\n"
"ordinary response,0,mentions C#\n",
encoding="utf-8",
)

dataset = HumanLabeledDataset.from_csv(
csv_path=str(csv_file),
metrics_type=MetricsType.OBJECTIVE,
)
assert len(dataset.entries) == 2
assert dataset.entries[0].objective == "respond to the prompt"
assert dataset.entries[1].objective == "mentions C#"


def test_from_csv_reads_metadata_comment_after_byte_order_mark(tmp_path):
"""A UTF-8 BOM (what Excel writes) must not hide the metadata comment line."""
csv_file = tmp_path / "bom_cells.csv"
csv_file.write_text(
"# dataset_version=1.0\nassistant_response,human_score,objective\nordinary response,0,mentions C#\n",
encoding="utf-8-sig",
)

dataset = HumanLabeledDataset.from_csv(
csv_path=str(csv_file),
metrics_type=MetricsType.OBJECTIVE,
)
assert len(dataset.entries) == 1
assert dataset.entries[0].objective == "mentions C#"


def test_from_csv_skips_blank_line_before_metadata_comment(tmp_path):
"""A blank line before the metadata comment must not stop the skip scan."""
csv_file = tmp_path / "blank_first.csv"
csv_file.write_text(
"\n"
"# dataset_version=1.0\n"
"assistant_response,human_score,objective\n"
"ordinary response,0,respond to the prompt\n",
encoding="utf-8",
)

dataset = HumanLabeledDataset.from_csv(
csv_path=str(csv_file),
metrics_type=MetricsType.OBJECTIVE,
version="1.0",
)
assert len(dataset.entries) == 1
assert dataset.entries[0].objective == "respond to the prompt"