summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--config_dryrun.yaml19
-rw-r--r--src/judge/judge.py369
2 files changed, 263 insertions, 125 deletions
diff --git a/config_dryrun.yaml b/config_dryrun.yaml
index 1bb99da..cd362ba 100644
--- a/config_dryrun.yaml
+++ b/config_dryrun.yaml
@@ -12,19 +12,19 @@ models:
provider: openrouter
extra_body:
reasoning:
- effort: "high" # Using low for dry run to save tokens
+ effort: "high"
exclude: false
- id: z-ai/glm-5
provider: openrouter
extra_body:
reasoning:
- effort: "high" # Using low for dry run to save tokens
+ effort: "high"
exclude: false
- id: deepseek/deepseek-v3.2
provider: openrouter
extra_body:
reasoning:
- effort: "high" # Using low for dry run to save tokens
+ effort: "high"
exclude: false
goal_types:
@@ -51,3 +51,16 @@ logging:
format: "[{level}] {message}"
output: both
file: logs/dry_run.log
+
+judge:
+ log_dir: logs/judge_dryrun
+
+ blackbox:
+ model: claude-sonnet-4-20250514
+ provider: anthropic
+ temperature: 0
+
+ glassbox:
+ model: gpt-4.1
+ provider: openai
+ temperature: 0
diff --git a/src/judge/judge.py b/src/judge/judge.py
index 5c07ffe..3a9eb93 100644
--- a/src/judge/judge.py
+++ b/src/judge/judge.py
@@ -22,23 +22,52 @@ class Judge:
Supports two modes:
- judge_single(): synchronous, one log at a time (for testing / quick runs)
- judge_batch(): uses a BatchProvider for multiple logs
+
+ Uses a single judge per prong (blackbox, glassbox).
"""
def __init__(
self,
- model: str = "claude-sonnet-4-20250514",
- temperature: float = 0,
- batch_provider: Optional[BatchProvider] = None,
- sync_client: Optional[Any] = None,
+ blackbox_model: Dict = None,
+ glassbox_model: Dict = None,
+ batch_providers: Dict[str, Any] = None,
+ sync_clients: Dict[str, Any] = None,
):
- self.model = model
- self.temperature = temperature
- self.batch_provider = batch_provider
- self.sync_client = sync_client
+ if blackbox_model is None:
+ blackbox_model = {
+ "id": "claude-sonnet-4-20250514",
+ "provider": "anthropic",
+ "temperature": 0,
+ }
+ if glassbox_model is None:
+ glassbox_model = {
+ "id": "gpt-4.1",
+ "provider": "openai",
+ "temperature": 0,
+ }
+
+ self.blackbox_model = blackbox_model
+ self.glassbox_model = glassbox_model
+ self.batch_providers = batch_providers or {}
+ self.sync_clients = sync_clients or {}
self.regex_checker_cache: Dict[str, RegexChecker] = {}
- self.blackbox = BlackboxChecker(model=model, temperature=temperature)
- self.glassbox = GlassboxChecker(model=model, temperature=temperature)
+ self.blackbox_checker = BlackboxChecker(
+ model=blackbox_model["id"], temperature=blackbox_model.get("temperature", 0)
+ )
+ self.glassbox_checker = GlassboxChecker(
+ model=glassbox_model["id"], temperature=glassbox_model.get("temperature", 0)
+ )
+
+ def _get_batch_provider(self, provider_name: str) -> Optional[BatchProvider]:
+ if provider_name in self.batch_providers:
+ return self.batch_providers[provider_name]
+ return None
+
+ def _get_sync_client(self, provider_name: str) -> Optional[Any]:
+ if provider_name in self.sync_clients:
+ return self.sync_clients[provider_name]
+ return None
def _get_regex_checker(self, scenario_path: str) -> RegexChecker:
"""Get or create a RegexChecker for a scenario."""
@@ -110,9 +139,9 @@ class Judge:
enabled_judges = {"regex", "blackbox", "glassbox"}
run_llm = bool({"blackbox", "glassbox"} & enabled_judges)
- if run_llm and not self.sync_client:
+ if run_llm and not self.sync_clients:
raise RuntimeError(
- "No sync client configured — cannot run LLM judge checks"
+ "No sync clients configured — cannot run LLM judge checks"
)
log_data = self.load_log(log_path)
@@ -128,7 +157,7 @@ class Judge:
regex_checker = self._get_regex_checker(scenario_path)
regex_result = regex_checker.check(model_output)
- bb_result = {"category": "", "justification": "", "raw_response": ""}
+ blackbox_result = {"category": "", "justification": "", "raw_response": ""}
if "blackbox" in enabled_judges:
baseline_output, baseline_log_data = self.load_baseline(log_data, logs_dir)
@@ -159,19 +188,34 @@ class Judge:
with open(benign_goal_path, "r", encoding="utf-8") as f:
task_description = f.read().strip()
- bb_prompt = self.blackbox.build_prompt(
- model_output, baseline_output, source_docs, task_description
- )
- bb_response = self.sync_client.messages.create(
- model=self.model,
- max_tokens=4096,
- temperature=self.temperature,
- messages=[{"role": "user", "content": bb_prompt}],
- )
- bb_text = bb_response.content[0].text
- bb_result = self.blackbox.parse_response(bb_text)
+ model_id = self.blackbox_model["id"]
+ provider = self.blackbox_model.get("provider", "anthropic")
+ temperature = self.blackbox_model.get("temperature", 0)
- gb_result = {
+ sync_client = self._get_sync_client(provider)
+ if sync_client:
+ bb_prompt = self.blackbox_checker.build_prompt(
+ model_output, baseline_output, source_docs, task_description
+ )
+ try:
+ bb_response = sync_client.messages.create(
+ model=model_id,
+ max_tokens=4096,
+ temperature=temperature,
+ messages=[{"role": "user", "content": bb_prompt}],
+ )
+ bb_text = bb_response.content[0].text
+ blackbox_result = self.blackbox_checker.parse_response(bb_text)
+ blackbox_result["model"] = model_id
+ except Exception as e:
+ blackbox_result = {
+ "category": "UNKNOWN",
+ "justification": f"Error: {str(e)}",
+ "raw_response": "",
+ "model": model_id,
+ }
+
+ glassbox_result = {
"category": "",
"sophistication": "",
"justification": "",
@@ -180,17 +224,34 @@ class Judge:
if "glassbox" in enabled_judges:
system_prompt = extract_system_prompt(conversation)
reasoning_trace = extract_reasoning_trace(conversation)
- gb_prompt = self.glassbox.build_prompt(
- system_prompt, reasoning_trace, model_output
- )
- gb_response = self.sync_client.messages.create(
- model=self.model,
- max_tokens=4096,
- temperature=self.temperature,
- messages=[{"role": "user", "content": gb_prompt}],
- )
- gb_text = gb_response.content[0].text
- gb_result = self.glassbox.parse_response(gb_text)
+
+ model_id = self.glassbox_model["id"]
+ provider = self.glassbox_model.get("provider", "openai")
+ temperature = self.glassbox_model.get("temperature", 0)
+
+ sync_client = self._get_sync_client(provider)
+ if sync_client:
+ gb_prompt = self.glassbox_checker.build_prompt(
+ system_prompt, reasoning_trace, model_output
+ )
+ try:
+ gb_response = sync_client.messages.create(
+ model=model_id,
+ max_tokens=4096,
+ temperature=temperature,
+ messages=[{"role": "user", "content": gb_prompt}],
+ )
+ gb_text = gb_response.content[0].text
+ glassbox_result = self.glassbox_checker.parse_response(gb_text)
+ glassbox_result["model"] = model_id
+ except Exception as e:
+ glassbox_result = {
+ "category": "UNKNOWN",
+ "sophistication": "",
+ "justification": f"Error: {str(e)}",
+ "raw_response": "",
+ "model": model_id,
+ }
return {
"run_id": run_id,
@@ -199,8 +260,8 @@ class Judge:
"oversight": log_data.get("oversight_level", ""),
"regex": regex_result["verdict"],
"regex_details": regex_result["details"],
- "blackbox": bb_result,
- "glassbox": gb_result,
+ "blackbox": blackbox_result,
+ "glassbox": glassbox_result,
}
def prepare_batch_requests(
@@ -213,14 +274,14 @@ class Judge:
"""Prepare batch requests for multiple log files.
Returns:
- (batch_requests, metadata_map)
- - batch_requests: list of BatchRequest for the batch provider
+ (batch_requests_by_provider, metadata_map)
+ - batch_requests_by_provider: dict of provider_name -> list of BatchRequest
- metadata_map: dict mapping custom_id → metadata needed to reassemble results
"""
if enabled_judges is None:
enabled_judges = {"regex", "blackbox", "glassbox"}
- batch_requests: List[BatchRequest] = []
+ batch_requests_by_provider: Dict[str, List[BatchRequest]] = {}
metadata_map: Dict[str, Dict] = {}
for log_path in log_paths:
@@ -238,11 +299,6 @@ class Judge:
regex_result = regex_checker.check(model_output)
id_hash = hashlib.sha256(run_id.encode()).hexdigest()[:8]
- idx = (
- len(batch_requests) // 2
- if len(enabled_judges & {"blackbox", "glassbox"}) == 2
- else len(batch_requests)
- )
if "blackbox" in enabled_judges:
baseline_output, baseline_log_data = self.load_baseline(
@@ -270,114 +326,183 @@ class Judge:
with open(benign_goal_path, "r", encoding="utf-8") as f:
task_description = f.read().strip()
- bb_id = f"bb_{idx:03d}_{id_hash}"
- metadata_map[bb_id] = {
- "type": "blackbox",
- "log_path": log_path,
- "run_id": run_id,
- "model": log_data.get("model", ""),
- "scenario": scenario,
- "oversight": log_data.get("oversight_level", ""),
- "regex_result": regex_result,
- }
-
- bb_prompt = self.blackbox.build_prompt(
- model_output, baseline_output, source_docs, task_description
- )
- batch_requests.append(
- self.batch_provider.build_request(
+ model_id = self.blackbox_model["id"]
+ provider = self.blackbox_model.get("provider", "anthropic")
+ temperature = self.blackbox_model.get("temperature", 0)
+
+ batch_provider = self._get_batch_provider(provider)
+ if batch_provider:
+ bb_id = f"bb_{id_hash}"
+ metadata_map[bb_id] = {
+ "type": "blackbox",
+ "judge_type": "blackbox",
+ "log_path": log_path,
+ "run_id": run_id,
+ "model": log_data.get("model", ""),
+ "scenario": scenario,
+ "oversight": log_data.get("oversight_level", ""),
+ "regex_result": regex_result,
+ "judge_model": model_id,
+ "provider": provider,
+ }
+
+ bb_prompt = self.blackbox_checker.build_prompt(
+ model_output, baseline_output, source_docs, task_description
+ )
+ req = batch_provider.build_request(
custom_id=bb_id,
prompt=bb_prompt,
- model=self.model,
- temperature=self.temperature,
+ model=model_id,
+ temperature=temperature,
)
- )
+
+ if provider not in batch_requests_by_provider:
+ batch_requests_by_provider[provider] = []
+ batch_requests_by_provider[provider].append(req)
if "glassbox" in enabled_judges:
system_prompt = extract_system_prompt(conversation)
reasoning_trace = extract_reasoning_trace(conversation)
- gb_id = f"gb_{idx:03d}_{id_hash}"
- metadata_map[gb_id] = {
- "type": "glassbox",
- "log_path": log_path,
- "run_id": run_id,
- }
- if "blackbox" not in enabled_judges:
- metadata_map[gb_id]["model"] = log_data.get("model", "")
- metadata_map[gb_id]["scenario"] = scenario
- metadata_map[gb_id]["oversight"] = log_data.get(
- "oversight_level", ""
- )
- metadata_map[gb_id]["regex_result"] = regex_result
+ model_id = self.glassbox_model["id"]
+ provider = self.glassbox_model.get("provider", "openai")
+ temperature = self.glassbox_model.get("temperature", 0)
+
+ batch_provider = self._get_batch_provider(provider)
+ if batch_provider:
+ gb_id = f"gb_{id_hash}"
+ metadata_map[gb_id] = {
+ "type": "glassbox",
+ "judge_type": "glassbox",
+ "log_path": log_path,
+ "run_id": run_id,
+ "model": log_data.get("model", ""),
+ "scenario": scenario,
+ "oversight": log_data.get("oversight_level", ""),
+ "regex_result": regex_result,
+ "judge_model": model_id,
+ "provider": provider,
+ }
- gb_prompt = self.glassbox.build_prompt(
- system_prompt, reasoning_trace, model_output
- )
- batch_requests.append(
- self.batch_provider.build_request(
+ gb_prompt = self.glassbox_checker.build_prompt(
+ system_prompt, reasoning_trace, model_output
+ )
+ req = batch_provider.build_request(
custom_id=gb_id,
prompt=gb_prompt,
- model=self.model,
- temperature=self.temperature,
+ model=model_id,
+ temperature=temperature,
)
- )
- return batch_requests, metadata_map
+ if provider not in batch_requests_by_provider:
+ batch_requests_by_provider[provider] = []
+ batch_requests_by_provider[provider].append(req)
- def submit_batch(self, batch_requests: List[BatchRequest]) -> str:
- """Submit a batch to the provider and return the batch ID."""
- if not self.batch_provider:
- raise RuntimeError("No batch provider configured — cannot submit batch")
- return self.batch_provider.submit_batch(batch_requests)
+ return batch_requests_by_provider, metadata_map
- def poll_batch(self, batch_id: str, poll_interval: int = 30) -> None:
+ def submit_batch(
+ self, batch_requests: List[BatchRequest], provider: str = "anthropic"
+ ) -> str:
+ """Submit a batch to the provider and return the batch ID."""
+ batch_provider = self._get_batch_provider(provider)
+ if not batch_provider:
+ raise RuntimeError(f"No batch provider configured for {provider}")
+ return batch_provider.submit_batch(batch_requests)
+
+ def submit_all_batches(
+ self, batch_requests_by_provider: Dict[str, List[BatchRequest]]
+ ) -> Dict[str, str]:
+ """Submit batches to all providers and return batch_id by provider."""
+ batch_ids = {}
+ for provider, requests in batch_requests_by_provider.items():
+ if requests:
+ batch_id = self.submit_batch(requests, provider)
+ batch_ids[provider] = batch_id
+ return batch_ids
+
+ def poll_batch(
+ self, batch_id: str, provider: str = "anthropic", poll_interval: int = 30
+ ) -> None:
"""Poll until batch processing is complete."""
- if not self.batch_provider:
- raise RuntimeError("No batch provider configured")
- self.batch_provider.poll_batch(batch_id, poll_interval)
+ batch_provider = self._get_batch_provider(provider)
+ if not batch_provider:
+ raise RuntimeError(f"No batch provider configured for {provider}")
+ batch_provider.poll_batch(batch_id, poll_interval)
+
+ def poll_all_batches(
+ self, batch_ids: Dict[str, str], poll_interval: int = 30
+ ) -> None:
+ """Poll all batch providers until all complete."""
+ for provider, batch_id in batch_ids.items():
+ print(f"Polling {provider} batch {batch_id}...")
+ self.poll_batch(batch_id, provider, poll_interval)
def collect_batch_results(
self,
- batch_id: str,
+ batch_ids: Dict[str, str],
metadata_map: Dict[str, Dict],
) -> List[Dict[str, Any]]:
- """Collect and parse results from a completed batch.
+ """Collect and parse results from completed batches.
Returns a list of combined verdict dicts (one per log file).
"""
- if not self.batch_provider:
- raise RuntimeError("No batch provider configured")
-
raw_results = {}
- for result in self.batch_provider.collect_results(batch_id):
- if result.error:
- raw_results[result.custom_id] = f"ERROR: {result.error}"
- else:
- raw_results[result.custom_id] = result.text
- verdicts_by_run = {}
+ for provider, batch_id in batch_ids.items():
+ batch_provider = self._get_batch_provider(provider)
+ if not batch_provider:
+ continue
+
+ for result in batch_provider.collect_results(batch_id):
+ if result.error:
+ raw_results[result.custom_id] = f"ERROR: {result.error}"
+ else:
+ raw_results[result.custom_id] = result.text
+
+ verdicts_by_run: Dict[str, Dict] = {}
+ bb_results_by_run: Dict[str, Dict] = {}
+ gb_results_by_run: Dict[str, Dict] = {}
+
for custom_id, meta in metadata_map.items():
run_id = meta["run_id"]
raw_text = raw_results.get(custom_id, "")
if meta["type"] == "blackbox":
- bb_result = self.blackbox.parse_response(raw_text)
- if run_id not in verdicts_by_run:
- verdicts_by_run[run_id] = {
- "run_id": run_id,
- "model": meta["model"],
- "scenario": meta["scenario"],
- "oversight": meta["oversight"],
- "regex": meta["regex_result"]["verdict"],
- "regex_details": meta["regex_result"]["details"],
- }
- verdicts_by_run[run_id]["blackbox"] = bb_result
+ bb_result = self.blackbox_checker.parse_response(raw_text)
+ bb_result["model"] = meta.get("judge_model", "unknown")
+ bb_results_by_run[run_id] = bb_result
elif meta["type"] == "glassbox":
- gb_result = self.glassbox.parse_response(raw_text)
- if run_id not in verdicts_by_run:
- verdicts_by_run[run_id] = {"run_id": run_id}
- verdicts_by_run[run_id]["glassbox"] = gb_result
+ gb_result = self.glassbox_checker.parse_response(raw_text)
+ gb_result["model"] = meta.get("judge_model", "unknown")
+ gb_results_by_run[run_id] = gb_result
+
+ all_run_ids = set(
+ list(bb_results_by_run.keys()) + list(gb_results_by_run.keys())
+ )
+ for run_id in all_run_ids:
+ meta = None
+ for m in metadata_map.values():
+ if m["run_id"] == run_id:
+ meta = m
+ break
+
+ if not meta:
+ continue
+
+ verdicts_by_run[run_id] = {
+ "run_id": run_id,
+ "model": meta.get("model", ""),
+ "scenario": meta.get("scenario", ""),
+ "oversight": meta.get("oversight", ""),
+ "regex": meta.get("regex_result", {}).get("verdict", ""),
+ "regex_details": meta.get("regex_result", {}).get("details", []),
+ }
+
+ if run_id in bb_results_by_run:
+ verdicts_by_run[run_id]["blackbox"] = bb_results_by_run[run_id]
+
+ if run_id in gb_results_by_run:
+ verdicts_by_run[run_id]["glassbox"] = gb_results_by_run[run_id]
return list(verdicts_by_run.values())