import json
import re
from typing import Any
from solar_eval.evaluators.base import BaseEvaluator
from solar_eval.models.sample import EvalSample
from solar_eval.providers.base import BaseProvider
class UnknownRuleNameError(ValueError):
"""`RuleBasedEvaluator` 에 등록되지 않은 rule 이름이 설정됐을 때.
예전엔 오타 난 rule 이 1.0(만점)으로 조용히 통과했다 -- 검증이 꺼진 것이 오히려
점수를 올리는 역설. `pipelines/registry.py` 가 모르는 스텝 이름을 생성 시점에
즉시 죽이는 것과 같은 이유로, rule 도 evaluator **생성 시점**에 죽는다
(샘플마다 반복해서 확인할 필요 없이 한 번에 끝난다).
"""
def __init__(self, rule_names: list[str], available: list[str]) -> None:
super().__init__(
f"Unknown rule name(s): {rule_names!r} (registered rules: {', '.join(available)})"
)
self.rule_names = rule_names
class RuleBasedEvaluator(BaseEvaluator):
"""Rule-based evaluator with configurable rules.
각 규칙(`_check_*`)이 실제로 쓰는 필드는 `output`(항상) 과 `reference`(규칙에
따라 다름, `golden` 이 dict 가 아니면 조용히 기본값으로 새는 규칙도 있다) 뿐이다
-- `input` 은 어느 규칙도 읽지 않는다. `required_fields` 는 클래스 속성이라
(구성된 `rules` 리스트와 무관하게) 규칙 종류와 무관하게 항상 필요한 최소
집합만 선언한다.
"""
required_fields = frozenset({"output"})
def __init__(self, rules: list[str]) -> None:
self.rules = rules
self._rule_funcs = {
"title_length_30": self._check_title_length_30,
"title_length_check": self._check_title_length_30,
"subtitle_length_40": self._check_subtitle_length_40,
"format_compliance": self._check_format_compliance,
"token_count_accuracy": self._check_token_count,
"position_accuracy": self._check_position_accuracy,
}
unknown = [r for r in rules if r not in self._rule_funcs]
if unknown:
raise UnknownRuleNameError(unknown, available=sorted(self._rule_funcs))
async def evaluate(
self,
sample: EvalSample,
provider: BaseProvider | None = None,
judge_model: str = "gpt-4o",
) -> dict[str, Any]:
rule_results = {}
for rule_name in self.rules:
# __init__ 이 이미 전 rule 이름을 검증했으므로 KeyError 가 날 수 없다.
func = self._rule_funcs[rule_name]
rule_results[rule_name] = func(sample.input, sample.output, sample.reference)
score = sum(rule_results.values()) / len(rule_results) if rule_results else 0.0
return {"score": score, "rule_results": rule_results, "details": rule_results}
def aggregate(self, results: list[dict[str, Any]]) -> dict[str, Any]:
if not results:
return {"overall_score": 0.0, "scores": {}}
rule_totals: dict[str, list[float]] = {}
for r in results:
for rule, score in r.get("rule_results", {}).items():
rule_totals.setdefault(rule, []).append(score)
avg_rules = {k: sum(v) / len(v) for k, v in rule_totals.items()}
overall = sum(r["score"] for r in results) / len(results)
return {"overall_score": overall, "scores": avg_rules, "num_samples": len(results)}
# --- Rule implementations ---
def _check_title_length_30(self, input_data: dict, output: str, golden: Any) -> float:
try:
parsed = json.loads(output)
title = parsed.get("eng_title", output)
except (json.JSONDecodeError, AttributeError):
title = output
return 1.0 if len(title) <= 80 else 0.0
def _check_subtitle_length_40(self, input_data: dict, output: str, golden: Any) -> float:
try:
parsed = json.loads(output)
subtitle = parsed.get("eng_subtitle", "")
except (json.JSONDecodeError, AttributeError):
subtitle = ""
return 1.0 if len(subtitle) <= 100 else 0.0
def _check_format_compliance(self, input_data: dict, output: str, golden: Any) -> float:
try:
parsed = json.loads(output)
return 1.0 if "eng_title" in parsed else 0.0
except (json.JSONDecodeError, AttributeError):
return 0.0
def _check_token_count(self, input_data: dict, output: str, golden: Any) -> float:
expected = golden.get("num_tokens", 0) if isinstance(golden, dict) else 0
fig_count = len(re.findall(r"", output))
return 1.0 if fig_count == expected else 0.0
def _check_position_accuracy(self, input_data: dict, output: str, golden: Any) -> float:
expected_positions = golden.get("positions", []) if isinstance(golden, dict) else []
paragraphs = output.split("\n\n")
actual_positions = []
for i, para in enumerate(paragraphs):
if "" in para:
actual_positions.append(i)
if not expected_positions:
return 1.0 if not actual_positions else 0.0
matches = sum(1 for a, e in zip(actual_positions, expected_positions) if a == e)
return matches / max(len(expected_positions), len(actual_positions))