Siddh12334's picture
feat: training space with manual start UI
204fa23 verified
Raw
History Blame Contribute Delete
2.5 kB
import re
from openenv.core.rubrics import Rubric
def _normalize(text: str) -> str:
text = text.lower()
text = re.sub(r"[^\w\s]", "", text)
text = re.sub(r"\s+", " ", text).strip()
return text
def compute_reward(
submitted_answer: str,
ground_truth_answer: str,
flagged_ids: list[int],
corrupt_ids: list[int],
confidence: float,
budget_used: int,
max_budget: int,
) -> tuple[float, dict]:
correct = _normalize(submitted_answer) == _normalize(ground_truth_answer)
answer_score = 0.4 if correct else 0.0
true_positives = [i for i in flagged_ids if i in corrupt_ids]
recall = len(true_positives) / len(corrupt_ids) if corrupt_ids else 0.0
recall_score = 0.3 * recall
false_positives = [i for i in flagged_ids if i not in corrupt_ids]
precision_score = max(0.0, 0.2 - 0.1 * len(false_positives))
confidence = confidence or 0.0
calibration_score = (0.1 * confidence) if correct else (-0.2 * confidence)
efficiency_score = 0.05 * (1 - budget_used / max_budget)
total = answer_score + recall_score + precision_score + calibration_score + efficiency_score
breakdown = {
"answer_correctness": round(answer_score, 4),
"flag_recall": round(recall_score, 4),
"false_positive_penalty": round(precision_score, 4),
"confidence_calibration": round(calibration_score, 4),
"efficiency": round(efficiency_score, 4),
"total": round(total, 4),
}
return round(total, 4), breakdown
class ContextCorruptionRubric(Rubric):
"""Scores a completed episode using compute_reward().
Requires a state_fn closure to access ground-truth env state that is
intentionally hidden from the agent's observation.
"""
def __init__(self, state_fn):
super().__init__()
self._state_fn = state_fn
self.last_breakdown: dict = {}
def forward(self, action, observation) -> float:
if not observation.done:
return 0.0
s = self._state_fn()
reward, breakdown = compute_reward(
submitted_answer=getattr(action, "answer", None) or "",
ground_truth_answer=s["ground_truth"],
flagged_ids=s["flagged_ids"],
corrupt_ids=s["corrupt_ids"],
confidence=getattr(action, "confidence", None) or 0.0,
budget_used=s["budget_used"],
max_budget=s["max_budget"],
)
self.last_breakdown = breakdown
return reward