diff options
| author | CaptainJack2491 <jayrupnakawala@gmail.com> | 2026-04-15 21:18:57 +0100 |
|---|---|---|
| committer | CaptainJack2491 <jayrupnakawala@gmail.com> | 2026-04-15 21:18:57 +0100 |
| commit | 6cc92304379996f68bb2df96a5e4b637b4c4804b (patch) | |
| tree | 6cdbc8b5484934927cd6ac989e91fef44a60bb4c /src/judge/judge.py | |
| parent | d25f8fc2e2b02c04489ca1960b83ff7a70e7317a (diff) | |
Study 1: 270 runs complete + judge validation pipeline
- config_study1.yaml: 3 models × 3 oversight × bare framing × n=30
- Judge validation: 54-run subset, gold=Sonnet 4.6, proxy=Grok 4.1 Fast (BB κ=0.702) + GPT-4.1 (GB κ=0.878)
- Fix OpenAI batch provider: BytesIO, method/url fields, response.body parsing
- Scripts: extract_subset.py, judge_validation.py
- Dissertation chapters updated (intro, methodology, results, conclusion)
Diffstat (limited to 'src/judge/judge.py')
| -rw-r--r-- | src/judge/judge.py | 80 |
1 files changed, 68 insertions, 12 deletions
diff --git a/src/judge/judge.py b/src/judge/judge.py index 681d880..71f37b3 100644 --- a/src/judge/judge.py +++ b/src/judge/judge.py @@ -18,9 +18,18 @@ from .judges import BlackboxChecker, GlassboxChecker, RegexChecker import anthropic try: + from dotenv import load_dotenv + + load_dotenv() +except ImportError: + pass + +try: import openai + from openai import BadRequestError except ImportError: openai = None + BadRequestError = None try: from xai_sdk import Client as XAIClient @@ -162,6 +171,61 @@ class Judge: return self.sync_clients[provider_name] return None + def _sync_completion( + self, + client: Any, + provider: str, + model: str, + prompt: str, + temperature: float = 0, + max_tokens: int = 4096, + ) -> str: + """Call sync completion API for any provider. + + Handles differences between providers: + - anthropic: client.messages.create() + - openai: client.chat.completions.create() + - xai: client.chat.create() + chat.append() + chat.sample() + """ + if provider == "xai": + from xai_sdk.chat import system as xai_system, user as xai_user + + chat = client.chat.create( + model=model, + max_tokens=max_tokens, + temperature=temperature, + ) + chat.append(xai_system("You are a helpful assistant.")) + chat.append(xai_user(prompt)) + result = chat.sample() + return result.content + elif provider == "openai": + try: + response = client.chat.completions.create( + model=model, + max_completion_tokens=max_tokens, + temperature=temperature, + messages=[{"role": "user", "content": prompt}], + ) + except BadRequestError as e: + if "does not support 0" in str(e): + response = client.chat.completions.create( + model=model, + max_completion_tokens=max_tokens, + messages=[{"role": "user", "content": prompt}], + ) + else: + raise + return response.choices[0].message.content + else: + response = client.messages.create( + model=model, + max_tokens=max_tokens, + temperature=temperature, + messages=[{"role": "user", "content": prompt}], + ) + return response.content[0].text + def _get_regex_checker(self, scenario_path: str) -> RegexChecker: """Get or create a RegexChecker for a scenario.""" if scenario_path not in self.regex_checker_cache: @@ -291,13 +355,9 @@ class Judge: 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 = self._sync_completion( + sync_client, provider, model_id, bb_prompt, temperature ) - 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: @@ -328,13 +388,9 @@ class Judge: 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 = self._sync_completion( + sync_client, provider, model_id, gb_prompt, temperature ) - 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: |
