File size: 3,784 Bytes
88e15cd | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 | from __future__ import annotations
import math
import re
from .config import RewardConfig
from .schemas import TaskRecord, WorkerResult
def _normalize_text(value: str, case_sensitive: bool) -> str:
normalized = " ".join(value.strip().split())
return normalized if case_sensitive else normalized.casefold()
def deterministic_quality(task: TaskRecord, candidate: str) -> float:
"""Grade tasks whose correctness can be checked locally."""
grader = task.grader
reference = task.reference_answer
if grader.type == "llm_judge":
raise ValueError("llm_judge tasks must be scored by the configured judge model")
if reference is None and grader.type != "regex":
raise ValueError(f"Task {task.task_id} needs reference_answer for {grader.type}")
if grader.type == "exact":
return float(
_normalize_text(candidate, grader.case_sensitive)
== _normalize_text(reference or "", grader.case_sensitive)
)
if grader.type == "contains":
return float(
_normalize_text(reference or "", grader.case_sensitive)
in _normalize_text(candidate, grader.case_sensitive)
)
if grader.type == "numeric":
candidate_match = re.search(r"[-+]?(?:\d+(?:\.\d*)?|\.\d+)(?:[eE][-+]?\d+)?", candidate)
reference_match = re.search(
r"[-+]?(?:\d+(?:\.\d*)?|\.\d+)(?:[eE][-+]?\d+)?", reference or ""
)
if not candidate_match or not reference_match:
return 0.0
candidate_value = float(candidate_match.group(0))
reference_value = float(reference_match.group(0))
return float(math.isclose(candidate_value, reference_value, abs_tol=grader.tolerance))
if grader.type == "regex":
pattern = grader.pattern or reference
if not pattern:
raise ValueError(f"Task {task.task_id} needs grader.pattern")
flags = 0 if grader.case_sensitive else re.IGNORECASE
return float(re.search(pattern, candidate, flags=flags) is not None)
raise ValueError(f"Unsupported grader: {grader.type}")
def compute_utility(result: WorkerResult, config: RewardConfig) -> float:
cost = result.cost_usd or 0.0
cost_penalty = min(cost / config.cost_scale_usd, 5.0)
latency_penalty = min((result.latency_ms / 1000.0) / config.latency_scale_seconds, 5.0)
return float(
config.quality_weight * result.quality
- config.cost_weight * cost_penalty
- config.latency_weight * latency_penalty
)
def soft_targets(rewards, temperature: float):
"""Torch-free reference implementation used by tests and data inspection."""
import numpy as np
values = np.asarray(rewards, dtype=np.float64)
if temperature <= 0:
raise ValueError("temperature must be positive")
shifted = (values - values.max(axis=-1, keepdims=True)) / temperature
probabilities = np.exp(shifted)
return probabilities / probabilities.sum(axis=-1, keepdims=True)
def grader_prompt(task: TaskRecord, candidate: str) -> str:
reference = task.reference_answer or "(no reference answer supplied)"
rubric = task.grader.rubric or "Correct, relevant, complete, and follows the requested format."
return (
f"TASK:\n{task.prompt}\n\n"
f"RUBRIC:\n{rubric}\n\n"
f"REFERENCE:\n{reference}\n\n"
f"CANDIDATE:\n{candidate}\n\n"
"SCORE (0 to 1):"
)
def parse_judge_score(text: str) -> float:
match = re.search(r"(?:^|\s)(?:0(?:\.\d+)?|1(?:\.0+)?)(?:\s|$)", text.strip())
if not match:
match = re.search(r"(?:0(?:\.\d+)?|1(?:\.0+)?)", text)
if not match:
raise ValueError(f"Judge did not return a 0..1 score: {text[:200]!r}")
return min(1.0, max(0.0, float(match.group(0).strip())))
|