WikiStar / Model /GPT.py
omerehrlich's picture
Per-instance OpenAI client + per-call temp JSONL; remove os.environ mutation
a98a573
Raw
History Blame Contribute Delete
4.69 kB
import json
import os
import tempfile
import time
from typing import Callable, List, Optional
import openai
def _summarize_batch(batch_status) -> dict:
"""Pull the user-relevant bits out of an OpenAI Batch object."""
counts = getattr(batch_status, "request_counts", None)
usage = getattr(batch_status, "usage", None) or {}
completed = getattr(counts, "completed", 0) if counts else 0
failed = getattr(counts, "failed", 0) if counts else 0
total_requests = getattr(counts, "total", 0) if counts else 0
return {
"batch_id": getattr(batch_status, "id", None),
"status": getattr(batch_status, "status", None),
"created_at": getattr(batch_status, "created_at", None),
"in_progress_at": getattr(batch_status, "in_progress_at", None),
"completed_at": getattr(batch_status, "completed_at", None),
"completed": completed,
"failed": failed,
"total_requests": total_requests,
"usage": dict(usage) if isinstance(usage, dict) else {},
}
class GPT:
def __init__(self, model_name="o3-mini", api_key=None):
self.model_name = model_name
self.api_key = api_key or os.getenv("OPENAI_API_KEY")
if not self.api_key:
raise ValueError("API key not provided and OPENAI_API_KEY environment variable not set.")
# Per-instance client so concurrent users with different keys never share auth.
self.client = openai.OpenAI(api_key=self.api_key)
def predict(
self,
prompts: List[str],
status_callback: Optional[Callable[[dict], None]] = None,
) -> List[str]:
# Use a per-call temp file so concurrent requests can't clobber each other's input.
with tempfile.NamedTemporaryFile(
mode="w", suffix=".jsonl", delete=False, encoding="utf-8"
) as tmp:
jsonl_path = tmp.name
try:
self._generate_batch_jsonl(prompts, jsonl_path)
# Upload the JSONL file for batch processing
with open(jsonl_path, "rb") as f:
file_resp = self.client.files.create(file=f, purpose="batch")
file_id = file_resp.id
finally:
try:
os.remove(jsonl_path)
except OSError:
pass
# Create a batch job
batch = self.client.batches.create(
input_file_id=file_id,
endpoint="/v1/chat/completions",
completion_window="24h"
)
batch_id = batch.id
# Wait for completion
while True:
batch_status = self.client.batches.retrieve(batch_id)
print(batch_status)
if status_callback is not None:
try:
status_callback(_summarize_batch(batch_status))
except Exception as cb_err:
print(f"status_callback error (ignored): {cb_err}")
if batch_status.status == "completed":
break
elif batch_status.status in {"failed", "expired", "cancelled"}:
raise RuntimeError(f"Batch job failed with status: {batch_status.status}")
time.sleep(10)
# Handle errors if any
if batch_status.error_file_id:
error_stream = self.client.files.content(batch_status.error_file_id)
error_log = error_stream.read().decode("utf-8")
raise RuntimeError(f"Batch job failed. Error log:\n{error_log}")
# Retrieve output
output_file_id = batch_status.output_file_id
file_stream = self.client.files.content(output_file_id)
content = file_stream.read().decode("utf-8")
# Parse responses
responses_by_id = {}
for line in content.splitlines():
obj = json.loads(line)
idx = int(obj["custom_id"])
text = obj["response"]["body"]["choices"][0]["message"]["content"]
responses_by_id[idx] = text
return [responses_by_id.get(i, "No response") for i in range(len(prompts))]
def _generate_batch_jsonl(self, prompts: List[str], output_path: str):
with open(output_path, "w", encoding="utf-8") as f:
for idx, prompt in enumerate(prompts):
entry = {
"custom_id": str(idx),
"method": "POST",
"url": "/v1/chat/completions",
"body": {
"model": self.model_name,
"messages": [
{"role": "user", "content": prompt}
]
}
}
f.write(json.dumps(entry) + "\n")