summaryrefslogtreecommitdiff
path: root/src/judge/judge.py
diff options
context:
space:
mode:
authorCaptainJack2491 <jayrupnakawala@gmail.com>2026-04-15 21:18:57 +0100
committerCaptainJack2491 <jayrupnakawala@gmail.com>2026-04-15 21:18:57 +0100
commit6cc92304379996f68bb2df96a5e4b637b4c4804b (patch)
tree6cdbc8b5484934927cd6ac989e91fef44a60bb4c /src/judge/judge.py
parentd25f8fc2e2b02c04489ca1960b83ff7a70e7317a (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.py80
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: