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
72 changes: 60 additions & 12 deletions assert_ai/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,11 @@
from assert_ai.core.judge import get_verdict_dimension, infer_judge_status, is_valid_event_flag
from assert_ai.display import label_metric, label_run_status, label_stage, label_stage_status, label_status
from assert_ai.logging_config import configure_logging
from assert_ai.results import compute_dimension_summary, detect_dimensions
from assert_ai.results import (
compute_dimension_summary,
compute_policy_violation_by_permissibility,
detect_dimensions,
)
from assert_ai.stages import STAGE_NAMES

ROOT = Path(__file__).resolve().parent.parent
Expand Down Expand Up @@ -387,7 +391,10 @@ def _reject_ordinal_compare(run_summaries: Iterable[dict[str, Any]], metric: str
)


def _compute_prompt_metrics(rows: list[dict[str, Any]]) -> dict[str, Any] | None:
def _compute_prompt_metrics(
rows: list[dict[str, Any]],
behavior_categories: Iterable[dict[str, Any]] = (),
) -> dict[str, Any] | None:
if not rows:
return None

Expand All @@ -414,27 +421,44 @@ def _compute_prompt_metrics(rows: list[dict[str, Any]]) -> dict[str, Any] | None
),
) or "-"
permissible_rows = [row for row in scored_rows if get_permissible_flag(row, default=False)]
not_permissible_rows = [row for row in scored_rows if not get_permissible_flag(row, default=False)]
permissibility_split = compute_policy_violation_by_permissibility(
scored_rows,
behavior_categories,
)

return {
metrics: dict[str, Any] = {
"total": len(rows),
"scored_total": scored_total,
"judge_failures": judge_failures,
"judge_failure_rate": judge_failures / len(rows) if rows else 0.0,
"policy_violation_rate": _dimension_rate({"dimensions": dimensions}, "policy_violation"),
"overrefusal_rate": _dimension_rate({"dimensions": dimensions}, "overrefusal"),
"permissible_overrefusal_rate": _compute_dimension_summary(permissible_rows, "overrefusal")["rate"],
"not_permissible_policy_violation_rate": _compute_dimension_summary(
not_permissible_rows,
"policy_violation",
)["rate"],
"dimensions": dimensions,
"target": target,
"judge_model": judge_model,
}

if permissibility_split["permissible"] is not None:
permissible = permissibility_split["permissible"]
not_permissible = permissibility_split["not_permissible"]
assert not_permissible is not None
metrics.update(
{
"permissible_policy_violation_rate": permissible["rate"],
"not_permissible_policy_violation_rate": not_permissible["rate"],
"policy_violation_on_permissible": permissible,
"policy_violation_on_not_permissible": not_permissible,
}
)

return metrics


def _compute_scenario_metrics(rows: list[dict[str, Any]]) -> dict[str, Any] | None:
def _compute_scenario_metrics(
rows: list[dict[str, Any]],
behavior_categories: Iterable[dict[str, Any]] = (),
) -> dict[str, Any] | None:
if not rows:
return None

Expand Down Expand Up @@ -470,7 +494,12 @@ def _compute_scenario_metrics(rows: list[dict[str, Any]]) -> dict[str, Any] | No
),
) or "-"

return {
permissibility_split = compute_policy_violation_by_permissibility(
scored_rows,
behavior_categories,
)

metrics: dict[str, Any] = {
"total": len(rows),
"scored_total": scored_total,
"judge_failures": judge_failures,
Expand All @@ -483,10 +512,29 @@ def _compute_scenario_metrics(rows: list[dict[str, Any]]) -> dict[str, Any] | No
"judge_model": judge_model,
}

if permissibility_split["permissible"] is not None:
permissible = permissibility_split["permissible"]
not_permissible = permissibility_split["not_permissible"]
assert not_permissible is not None
metrics.update(
{
"permissible_policy_violation_rate": permissible["rate"],
"not_permissible_policy_violation_rate": not_permissible["rate"],
"policy_violation_on_permissible": permissible,
"policy_violation_on_not_permissible": not_permissible,
}
)

return metrics


def _load_run_summary(run_dir: Path) -> dict[str, Any] | None:
manifest = load_json(run_dir / "manifest.json")
score_rows = load_jsonl(run_dir / "scores.jsonl")
taxonomy = load_json(run_dir.parent / "taxonomy.json")
behavior_categories = (taxonomy or {}).get("behavior_categories")
if not isinstance(behavior_categories, list):
behavior_categories = []
prompt_rows = [row for row in score_rows if not row.get("tester_model")]
scenario_rows = [row for row in score_rows if row.get("tester_model")]

Expand All @@ -507,8 +555,8 @@ def _load_run_summary(run_dir: Path) -> dict[str, Any] | None:
"current_stage": current_stage,
"started_at": (manifest or {}).get("started_at"),
"ended_at": (manifest or {}).get("ended_at"),
"prompt_metrics": _compute_prompt_metrics(prompt_rows),
"scenario_metrics": _compute_scenario_metrics(scenario_rows),
"prompt_metrics": _compute_prompt_metrics(prompt_rows, behavior_categories),
"scenario_metrics": _compute_scenario_metrics(scenario_rows, behavior_categories),
"prompt_rows": prompt_rows,
"scenario_rows": scenario_rows,
}
Expand Down
132 changes: 126 additions & 6 deletions assert_ai/results.py
Original file line number Diff line number Diff line change
Expand Up @@ -176,6 +176,94 @@ def dimension_rate(metrics: dict[str, Any], metric: str) -> float | None:
return float(rate) if isinstance(rate, (int, float)) else None


def compute_policy_violation_by_permissibility(
rows: Iterable[dict[str, Any]],
behavior_categories: Iterable[dict[str, Any]],
) -> dict[str, dict[str, Any] | None]:
"""Split policy violations by node permissibility, one vote per row.

Each row contributes to a bucket only when at least one behavior in that
bucket was relevant. Its vote is true when any such behavior was violated.
The test case's originating behavior does not affect this calculation.
"""
categories = [
(index, entry)
for index, entry in enumerate(behavior_categories)
if isinstance(entry, dict)
]
if not categories:
return {"permissible": None, "not_permissible": None}

permissible_by_index = {
index: entry.get("permissible") is True
for index, entry in categories
}
permissible_by_name = {
str(entry.get("name") or "").strip(): entry.get("permissible") is True
for _, entry in categories
if str(entry.get("name") or "").strip()
}
flags: dict[bool, list[bool]] = {True: [], False: []}
not_applicable: dict[bool, int] = {True: 0, False: 0}

for row in rows:
verdict = row.get("verdict")
node_judgments = verdict.get("node_judgments") if isinstance(verdict, dict) else None
row_applicable = {True: False, False: False}
row_violated = {True: False, False: False}

if isinstance(node_judgments, list):
for node in node_judgments:
if not isinstance(node, dict):
continue
if "relevant" in node and node.get("relevant") is not True:
continue
violated = node.get("violated")
if not isinstance(violated, bool):
continue

node_index = node.get("node_index")
if (
isinstance(node_index, int)
and not isinstance(node_index, bool)
and node_index in permissible_by_index
):
permissible = permissible_by_index[node_index]
else:
node_name = str(node.get("node_name") or "").strip()
if node_name not in permissible_by_name:
continue
permissible = permissible_by_name[node_name]

row_applicable[permissible] = True
row_violated[permissible] = row_violated[permissible] or violated

for permissible in (True, False):
if row_applicable[permissible]:
flags[permissible].append(row_violated[permissible])
else:
not_applicable[permissible] += 1

def summarize(permissible: bool) -> dict[str, Any]:
values = flags[permissible]
flagged_count = sum(values)
clear_count = len(values) - flagged_count
return {
"rate": flagged_count / len(values) if values else None,
"counts": {0: clear_count, 1: flagged_count},
"count": len(values),
"applicable_count": len(values),
"not_applicable_count": not_applicable[permissible],
"flagged_count": flagged_count,
"clear_count": clear_count,
}

return {
"permissible": summarize(True),
"not_permissible": summarize(False),
}


def _first_str(rows: Iterable[dict[str, Any]], key: str) -> str:
for row in rows:
value = row.get(key)
Expand All @@ -188,6 +276,7 @@ def _compute_test_set_metrics(
rows: list[dict[str, Any]],
*,
include_tester_model: bool = False,
behavior_categories: Iterable[dict[str, Any]] = (),
) -> dict[str, Any] | None:
if not rows:
return None
Expand All @@ -211,26 +300,57 @@ def _compute_test_set_metrics(
"judge_model": _first_str(rows, "judge_model"),
}

permissibility_split = compute_policy_violation_by_permissibility(
scored_rows,
behavior_categories,
)
if permissibility_split["permissible"] is not None:
permissible = permissibility_split["permissible"]
not_permissible = permissibility_split["not_permissible"]
assert not_permissible is not None
metrics.update(
{
"permissible_policy_violation_rate": permissible["rate"],
"not_permissible_policy_violation_rate": not_permissible["rate"],
"policy_violation_on_permissible": permissible,
"policy_violation_on_not_permissible": not_permissible,
}
)

if include_tester_model:
metrics["tester_model"] = _first_str(rows, "tester_model")

return metrics


def compute_prompt_metrics(rows: list[dict[str, Any]]) -> dict[str, Any] | None:
def compute_prompt_metrics(
rows: list[dict[str, Any]],
behavior_categories: Iterable[dict[str, Any]] = (),
) -> dict[str, Any] | None:
"""Compute prompt-only summary metrics."""
return _compute_test_set_metrics(rows)
return _compute_test_set_metrics(rows, behavior_categories=behavior_categories)


def compute_scenario_metrics(rows: list[dict[str, Any]]) -> dict[str, Any] | None:
def compute_scenario_metrics(
rows: list[dict[str, Any]],
behavior_categories: Iterable[dict[str, Any]] = (),
) -> dict[str, Any] | None:
"""Compute scenario-only summary metrics."""
return _compute_test_set_metrics(rows, include_tester_model=True)
return _compute_test_set_metrics(
rows,
include_tester_model=True,
behavior_categories=behavior_categories,
)


def load_run_summary(run_dir: Path) -> dict[str, Any] | None:
"""Load one run's manifest and score-derived summaries."""
manifest = load_json(run_dir / "manifest.json")
score_rows = load_jsonl(run_dir / "scores.jsonl")
taxonomy = load_json(run_dir.parent / "taxonomy.json")
behavior_categories = (taxonomy or {}).get("behavior_categories")
if not isinstance(behavior_categories, list):
behavior_categories = []
prompt_rows = [row for row in score_rows if not row.get("tester_model")]
scenario_rows = [row for row in score_rows if row.get("tester_model")]

Expand All @@ -251,8 +371,8 @@ def load_run_summary(run_dir: Path) -> dict[str, Any] | None:
"current_stage": current_stage,
"started_at": (manifest or {}).get("started_at"),
"ended_at": (manifest or {}).get("ended_at"),
"prompt_metrics": compute_prompt_metrics(prompt_rows),
"scenario_metrics": compute_scenario_metrics(scenario_rows),
"prompt_metrics": compute_prompt_metrics(prompt_rows, behavior_categories),
"scenario_metrics": compute_scenario_metrics(scenario_rows, behavior_categories),
"prompt_rows": prompt_rows,
"scenario_rows": scenario_rows,
}
Expand Down
26 changes: 13 additions & 13 deletions scripts/export_suite_results.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,7 @@
)
from assert_ai.core.judge import get_verdict_dimension, infer_judge_status, is_not_applicable_dimension, is_valid_event_flag
from assert_ai.core.transcript import Transcript, TranscriptEvent, TranscriptMetadata
from assert_ai.results import compute_dimension_summary, detect_dimensions
from assert_ai.results import compute_dimension_summary, compute_policy_violation_by_permissibility, detect_dimensions

EXPORT_DIR_NAME = "exports"
CSV_FORMAT = "csv"
Expand Down Expand Up @@ -643,16 +643,17 @@ def load_suite_tables(
for key in dimensions_payload:
relevant_dimensions.add(str(key))

permissible_scores = [
row for row in score_rows
if infer_judge_status(row) == "ok"
and _row_permissible(row, permissible_by_name)
ok_score_rows = [
row for row in score_rows if infer_judge_status(row) == "ok"
]
not_permissible_scores = [
row for row in score_rows
if infer_judge_status(row) == "ok"
and not _row_permissible(row, permissible_by_name)
permissible_scores = [
row for row in ok_score_rows
if _row_permissible(row, permissible_by_name)
]
policy_violation_split = compute_policy_violation_by_permissibility(
ok_score_rows,
taxonomy.get("behavior_categories") or [],
)
run_rows.append(
{
"suite_id": suite_id,
Expand All @@ -671,10 +672,9 @@ def load_suite_tables(
"policy_violation_rate": _event_rate(score_rows, "policy_violation"),
"overrefusal_rate": _event_rate(score_rows, "overrefusal"),
"permissible_overrefusal_rate": _event_rate(permissible_scores, "overrefusal"),
"not_permissible_policy_violation_rate": _event_rate(
not_permissible_scores,
"policy_violation",
),
"not_permissible_policy_violation_rate": (
policy_violation_split["not_permissible"] or {}
).get("rate"),
}
)

Expand Down
Loading
Loading