Spaces:
Sleeping
Sleeping
Download solar_eval/evaluators/rule_based.py from dev-strender/proofread-demo: direct link, hf CLI and curl.
- Browser
- Download file 5.38 kB
-
https://huggingface.co/spaces/dev-strender/proofread-demo/resolve/main/solar_eval/evaluators/rule_based.py
- Command line
-
hf download hf://spaces/dev-strender/proofread-demo/solar_eval/evaluators/rule_based.py
-
curl -L -o rule_based.py https://huggingface.co/spaces/dev-strender/proofread-demo/resolve/main/solar_eval/evaluators/rule_based.py
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)) | |