diff options
Diffstat (limited to 'src/judge/batch_providers.py')
| -rw-r--r-- | src/judge/batch_providers.py | 114 |
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", ""), + ) |
