diff options
| author | CaptainJack2491 <jayrupnakawala@gmail.com> | 2026-04-10 14:50:40 +0100 |
|---|---|---|
| committer | CaptainJack2491 <jayrupnakawala@gmail.com> | 2026-04-10 14:50:40 +0100 |
| commit | 9155f664a2db0abe841d0d7b9b059ffbe2290c18 (patch) | |
| tree | 8f59187b2b7677727eb19d990b7de377d38e6bad | |
| parent | a72e0d6e93e6ea8d44b4ca86d51903bfbf4e21d5 (diff) | |
refactor(judge): simplify to single judge per prong, remove multi-model aggregation
- Remove multi-model aggregation support from Judge class
- Single blackbox judge: claude-sonnet-4-20250514 (anthropic)
- Single glassbox judge: gpt-4.1 (openai)
- Remove aggregate_results method and all aggregation-related code
- Simplify batch collection to single result per run
- Update config_dryrun.yaml with new judge configuration
| -rw-r--r-- | config_dryrun.yaml | 19 | ||||
| -rw-r--r-- | src/judge/judge.py | 369 |
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()) |
