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))