summaryrefslogtreecommitdiff
path: root/src/judge/batch_providers.py
diff options
context:
space:
mode:
Diffstat (limited to 'src/judge/batch_providers.py')
-rw-r--r--src/judge/batch_providers.py114
1 files changed, 114 insertions, 0 deletions
diff --git a/src/judge/batch_providers.py b/src/judge/batch_providers.py
index 1a6ca27..9213766 100644
--- a/src/judge/batch_providers.py
+++ b/src/judge/batch_providers.py
@@ -12,6 +12,11 @@ from typing import Any, Dict, Iterator, List, Optional
import anthropic
try:
+ import openai
+except ImportError:
+ openai = None
+
+try:
from xai_sdk import Client as XAIClient
except ImportError:
XAIClient = None
@@ -35,6 +40,9 @@ class BatchResult:
custom_id: str
text: str
error: Optional[str] = None
+ usage: Optional[Dict[str, int]] = None
+ latency_ms: Optional[int] = None
+ model: Optional[str] = None
class BatchProvider(ABC):
@@ -218,3 +226,109 @@ class XAIBatchProvider(BatchProvider):
if page.pagination_token is None:
break
pagination_token = page.pagination_token
+
+
+class OpenAIBatchProvider(BatchProvider):
+ def __init__(self, api_key: Optional[str] = None):
+ if openai is None:
+ raise ImportError("openai not installed. Run: uv add openai")
+ key = api_key or os.environ.get("OPENAI_API_KEY")
+ if not key:
+ raise ValueError("OPENAI_API_KEY not set")
+ self.client = openai.OpenAI(api_key=key)
+
+ def build_request(
+ self,
+ custom_id: str,
+ prompt: str,
+ model: str,
+ temperature: float,
+ max_tokens: int = 4096,
+ ) -> BatchRequest:
+ return BatchRequest(
+ custom_id=custom_id,
+ params={
+ "model": model,
+ "max_tokens": max_tokens,
+ "temperature": temperature,
+ "messages": [{"role": "user", "content": prompt}],
+ },
+ )
+
+ def submit_batch(self, requests: List[BatchRequest]) -> str:
+ openai_requests = [
+ {
+ "custom_id": r.custom_id,
+ "method": "POST",
+ "url": "/v1/chat/completions",
+ "body": {
+ "model": r.params["model"],
+ "max_tokens": r.params["max_tokens"],
+ "temperature": r.params["temperature"],
+ "messages": r.params["messages"],
+ },
+ }
+ for r in requests
+ ]
+ response = self.client.batches.create(
+ input_file_id=self._upload_requests(openai_requests),
+ endpoint="/v1/chat/completions",
+ completion_window="24h",
+ )
+ return response.id
+
+ def _upload_requests(self, requests: List[Dict]) -> str:
+ import json
+
+ content = "\n".join(json.dumps(req) for req in requests)
+ import io
+
+ file_obj = io.BytesIO(content.encode("utf-8"))
+ upload = self.client.files.create(file=file_obj, purpose="batch")
+ return upload.id
+
+ def poll_batch(self, batch_id: str, poll_interval: int = 30) -> None:
+ while True:
+ batch = self.client.batches.retrieve(batch_id)
+ status = batch.status
+ counts = batch.request_counts
+ print(
+ f" Batch {batch_id}: {status} "
+ f"(completed={counts.completed}, "
+ f"failed={counts.failed}, "
+ f"total={counts.total})"
+ )
+ if status in ["completed", "failed", "expired"]:
+ return
+ time.sleep(poll_interval)
+
+ def collect_results(self, batch_id: str) -> Iterator[BatchResult]:
+ batch = self.client.batches.retrieve(batch_id)
+ if not batch.output_file_id:
+ return
+
+ response = self.client.files.content(batch.output_file_id)
+ import json
+
+ for line in response.text.strip().split("\n"):
+ if not line:
+ continue
+ data = json.loads(line)
+ custom_id = data.get("custom_id", "")
+ if data.get("error"):
+ yield BatchResult(
+ custom_id=custom_id,
+ text="",
+ error=data["error"].get("message", "Unknown error"),
+ )
+ else:
+ body = data["response"]["body"]
+ msg = body["choices"][0]["message"]
+ text = msg.get("content", "")
+ usage = body.get("usage", {})
+ yield BatchResult(
+ custom_id=custom_id,
+ text=text,
+ usage=usage,
+ model=body.get("model", ""),
+ )