Initial commit
This commit is contained in:
@@ -0,0 +1,191 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
from typing import Any, Protocol
|
||||
|
||||
import requests
|
||||
|
||||
from .identity import composite_id
|
||||
|
||||
|
||||
DEFAULT_MAX_LENGTH = 8192
|
||||
DEFAULT_RM_API_URL = "http://127.0.0.1:28080"
|
||||
DEFAULT_TIMEOUT = 300.0
|
||||
DEFAULT_CONCURRENCY = 8
|
||||
DEFAULT_BATCH_SIZE = 32
|
||||
RETRY_ATTEMPTS = 3
|
||||
|
||||
|
||||
class AgentRMBackend(Protocol):
|
||||
def score(self, requests: list[dict[str, Any]]) -> list[dict[str, Any]]: ...
|
||||
|
||||
|
||||
class HttpAgentRMBackend:
|
||||
"""AgentRM backend backed by the remote ``/score_batch`` API."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_url: str | None = None,
|
||||
*,
|
||||
max_length: int = DEFAULT_MAX_LENGTH,
|
||||
timeout: float = DEFAULT_TIMEOUT,
|
||||
concurrency: int = DEFAULT_CONCURRENCY,
|
||||
batch_size: int = DEFAULT_BATCH_SIZE,
|
||||
) -> None:
|
||||
self.api_url = (
|
||||
api_url or os.environ.get("RM_API_URL", DEFAULT_RM_API_URL)
|
||||
).rstrip("/")
|
||||
self.max_length = max_length
|
||||
self.timeout = timeout
|
||||
self.concurrency = concurrency
|
||||
self.batch_size = batch_size
|
||||
if not self.api_url:
|
||||
raise ValueError("AgentRM API URL cannot be empty")
|
||||
if max_length <= 0 or timeout <= 0 or concurrency <= 0 or batch_size <= 0:
|
||||
raise ValueError(
|
||||
"AgentRM max_length, timeout, concurrency, and batch_size must be positive"
|
||||
)
|
||||
|
||||
def _post_batch(self, batch: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
payload = {
|
||||
"states": [request["state"] for request in batch],
|
||||
"max_length": self.max_length,
|
||||
}
|
||||
last_error: BaseException | None = None
|
||||
for attempt in range(RETRY_ATTEMPTS):
|
||||
try:
|
||||
response = requests.post(
|
||||
f"{self.api_url}/score_batch",
|
||||
json=payload,
|
||||
timeout=self.timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
body = response.json()
|
||||
scores = body.get("scores") if isinstance(body, dict) else None
|
||||
if not isinstance(scores, list) or len(scores) != len(batch):
|
||||
count = len(scores) if isinstance(scores, list) else "invalid"
|
||||
raise ValueError(
|
||||
f"AgentRM returned {count} scores for {len(batch)} states"
|
||||
)
|
||||
if any(not isinstance(score, dict) for score in scores):
|
||||
raise ValueError("AgentRM returned a non-object score item")
|
||||
if any(
|
||||
"score" not in score or "n_tokens" not in score
|
||||
for score in scores
|
||||
):
|
||||
raise ValueError("AgentRM returned a score item with missing fields")
|
||||
return [
|
||||
{
|
||||
"task_name": request["task_name"],
|
||||
"compile_type": request["compile_type"],
|
||||
"test_name": request["test_name"],
|
||||
**score,
|
||||
}
|
||||
for request, score in zip(batch, scores)
|
||||
]
|
||||
except (requests.RequestException, ValueError) as exc:
|
||||
last_error = exc
|
||||
if attempt + 1 < RETRY_ATTEMPTS:
|
||||
time.sleep(2**attempt)
|
||||
assert last_error is not None
|
||||
raise RuntimeError(
|
||||
f"AgentRM request failed after {RETRY_ATTEMPTS} attempts: {last_error}"
|
||||
)
|
||||
|
||||
def score(self, requests_to_score: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
if not requests_to_score:
|
||||
return []
|
||||
batches = [
|
||||
requests_to_score[index : index + self.batch_size]
|
||||
for index in range(0, len(requests_to_score), self.batch_size)
|
||||
]
|
||||
ordered: list[list[dict[str, Any]] | None] = [None] * len(batches)
|
||||
with ThreadPoolExecutor(max_workers=self.concurrency) as pool:
|
||||
futures = {
|
||||
pool.submit(self._post_batch, batch): index
|
||||
for index, batch in enumerate(batches)
|
||||
}
|
||||
for future in as_completed(futures):
|
||||
ordered[futures[future]] = future.result()
|
||||
return [row for batch in ordered if batch is not None for row in batch]
|
||||
|
||||
|
||||
class AgentRM:
|
||||
"""Score AgentRM requests through a validated, replaceable backend."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
backend: AgentRMBackend | None = None,
|
||||
*,
|
||||
api_url: str | None = None,
|
||||
max_length: int = DEFAULT_MAX_LENGTH,
|
||||
timeout: float = DEFAULT_TIMEOUT,
|
||||
concurrency: int = DEFAULT_CONCURRENCY,
|
||||
batch_size: int = DEFAULT_BATCH_SIZE,
|
||||
) -> None:
|
||||
self.backend = (
|
||||
backend
|
||||
if backend is not None
|
||||
else HttpAgentRMBackend(
|
||||
api_url,
|
||||
max_length=max_length,
|
||||
timeout=timeout,
|
||||
concurrency=concurrency,
|
||||
batch_size=batch_size,
|
||||
)
|
||||
)
|
||||
|
||||
def score_requests(
|
||||
self, requests_to_score: list[dict[str, Any]]
|
||||
) -> list[dict[str, Any]]:
|
||||
request_keys = [composite_id(request) for request in requests_to_score]
|
||||
if len(set(request_keys)) != len(request_keys):
|
||||
raise ValueError("AgentRM requests contain duplicate identities")
|
||||
|
||||
responses = self.backend.score(requests_to_score)
|
||||
response_by_key: dict[tuple[str, str, str], dict[str, Any]] = {}
|
||||
for response in responses:
|
||||
key = composite_id(response)
|
||||
if key in response_by_key:
|
||||
raise ValueError(f"AgentRM returned duplicate score identity: {key}")
|
||||
response_by_key[key] = response
|
||||
|
||||
requested = set(request_keys)
|
||||
unexpected = set(response_by_key) - requested
|
||||
missing = requested - set(response_by_key)
|
||||
if unexpected:
|
||||
raise ValueError(
|
||||
f"AgentRM returned unexpected score identities: {sorted(unexpected)}"
|
||||
)
|
||||
if missing:
|
||||
raise ValueError(f"AgentRM returned incomplete scores: {sorted(missing)}")
|
||||
|
||||
result = []
|
||||
for request, key in zip(requests_to_score, request_keys):
|
||||
response = response_by_key[key]
|
||||
try:
|
||||
score = float(response["score"])
|
||||
except (KeyError, TypeError, ValueError) as exc:
|
||||
raise ValueError(f"AgentRM returned an invalid score for {key}") from exc
|
||||
if not math.isfinite(score):
|
||||
raise ValueError(f"AgentRM returned a non-finite score for {key}")
|
||||
try:
|
||||
n_tokens = int(response["n_tokens"])
|
||||
except (KeyError, TypeError, ValueError) as exc:
|
||||
raise ValueError(
|
||||
f"AgentRM returned an invalid n_tokens for {key}"
|
||||
) from exc
|
||||
if n_tokens < 0:
|
||||
raise ValueError(f"AgentRM returned a negative n_tokens for {key}")
|
||||
row = {
|
||||
"task_name": str(request["task_name"]),
|
||||
"compile_type": str(request["compile_type"]),
|
||||
"test_name": str(request["test_name"]),
|
||||
"score": score,
|
||||
"n_tokens": n_tokens,
|
||||
}
|
||||
result.append(row)
|
||||
return result
|
||||
Reference in New Issue
Block a user