192 lines
7.2 KiB
Python
192 lines
7.2 KiB
Python
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
|