Files
2026-09-04 14:58:42 +08:00

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