dev-strender's picture
Replace v24-era demo with v34 pipeline demo (engine-vendored bundle)
9c84f9d verified
Raw History Blame Contribute Delete
5.38 kB
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"<fig></fig>", 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 "<fig></fig>" 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))