Spaces:
Running
Running
Download solar_eval/core/runner.py from dev-strender/proofread-demo: direct link, hf CLI and curl.
- Browser
- Download file 19.5 kB
-
https://huggingface.co/spaces/dev-strender/proofread-demo/resolve/483134ace86f21c444532ef540c77505aadb29b9/solar_eval/core/runner.py
- Command line
-
hf download hf://spaces/dev-strender/proofread-demo@483134ace86f21c444532ef540c77505aadb29b9/solar_eval/core/runner.py
-
curl -L -o runner.py https://huggingface.co/spaces/dev-strender/proofread-demo/resolve/483134ace86f21c444532ef540c77505aadb29b9/solar_eval/core/runner.py
19.5 kB
| import asyncio | |
| import logging | |
| import sys | |
| from datetime import datetime, timezone | |
| from typing import Any, Callable | |
| # Ensure runner logs are visible in CLI | |
| logging.basicConfig(stream=sys.stderr, level=logging.INFO, format="%(message)s") | |
| from solar_eval.core.dataset_loader import DatasetLoader | |
| from solar_eval.evaluators.registry import create_evaluator | |
| from solar_eval.models.enums import RunStatus | |
| from solar_eval.models.sample import EvalSample | |
| from solar_eval.pipelines.registry import create_pipeline | |
| from solar_eval.providers.base import BaseProvider | |
| from solar_eval.stores.base import ResultStore | |
| logger = logging.getLogger(__name__) | |
| OnProgress = Callable[[dict[str, Any]], Any] | |
| def build_eval_sample_from_result(result: dict[str, Any], reference: Any) -> EvalSample: | |
| """์ ์ฅ๋ ๊ฒฐ๊ณผ dict(๋ ๊ฑฐ์ 8+trace ํค)์์ ์ฑ์ ์ฉ `EvalSample` ์ ๋์ด๋ฆฐ๋ค. | |
| `insert_run_result` ๊ฐ ์ฐ๋ ํค ์ด๋ฆ(`golden`/`trace`)์ `EvalSample` ํ๋ ์ด๋ฆ | |
| (`reference`/`artifacts`)๊ณผ ๋ค๋ฅด๋ค -- eval-store ํ์ ํธํ์ ์ํด ์ ์ฅ ์คํค๋ง๋ | |
| ๊ทธ๋๋ก ๋๊ณ (๋ง์ด๊ทธ๋ ์ด์ ๊ณํ ยง9-1) ์ฑ์ ์ง์ ์๋ง ์ฌ๊ธฐ์ ๋๋๋ฆฐ๋ค. `execute_run` | |
| ๋ด๋ถ ๋ฃจํ์ `runs eval` CLI(์ฌ์ฑ์ , ์ ์ถ๋ก ์์ด ๋์คํฌ์ `results.jsonl` ์ | |
| ๊ทธ๋๋ก ์) ์์ชฝ์์ ๊ณต์ ํ๋ค. | |
| Args: | |
| result: `insert_run_result` ์ ๋๊ธด ๊ฒ๊ณผ ๊ฐ์ ํํ(๋๋ `JsonlStore. | |
| load_existing_results()` ๋ก ๋์คํฌ์์ ๋ค์ ์ฝ์ ๊ฒ). `input`/`output`/ | |
| `trace` ํค๋ฅผ ์ฝ๋๋ค. | |
| reference: ์ ๋ต(golden) ๊ฐ. ํธ์ถ์๊ฐ ์ง์ ๋๊ธด๋ค -- ๋ ํธ์ถ์์ golden | |
| ์ถ์ถ ๊ฒฝ๋ก๊ฐ ๋ค๋ฅด๊ธฐ ๋๋ฌธ์ด๋ค(์๋ ์ฐธ๊ณ ). ์ด ํจ์๋ ์ด๋ ์ชฝ์ธ์ง ๋ชจ๋ฅธ ์ฑ | |
| ๊ฐ๋ง ๋ฐ์ ๊ทธ๋๋ก `sample.reference` ์ ์ฑ์ด๋ค. | |
| - `execute_run` ๋ด๋ถ ๋ฃจํ: `_golden_raw`(๋์คํฌ์ ์ ๋จ๋ ์คํ ์ค | |
| ์์ ํค, resume ์์๋ ๋งค ์ํ๋ง๋ค ์๋ก ๊ณ์ฐ๋จ)๋ฅผ ๋๊ธด๋ค. | |
| - `runs eval` CLI: ๋์คํฌ์ ์ ์ฅ๋ `golden` ํค๋ฅผ ๊ทธ๋๋ก ๋๊ธด๋ค | |
| (`_golden_raw` ๋ ์ ์ด์ ์ ์ฅ๋์ง ์์ CLI ์ชฝ์ ์๋ค). | |
| Returns: | |
| `input`/`output`/`reference`/`artifacts` ๊ฐ ์ฑ์์ง `EvalSample`. | |
| `contexts` ๋ ์์ง ์๋ฌด๋ ์ ์จ์ ๊ธฐ๋ณธ๊ฐ(`None`) ๊ทธ๋๋ก๋ค. | |
| """ | |
| trace = result.get("trace") or {} | |
| return EvalSample( | |
| input=result.get("input"), | |
| output=result.get("output"), | |
| reference=reference, | |
| artifacts=trace, | |
| ) | |
| class BatchRunner: | |
| """Runs inference + evaluation batches with progress tracking.""" | |
| def __init__( | |
| self, | |
| store: ResultStore, | |
| inference_provider: BaseProvider, | |
| judge_provider: BaseProvider | None = None, | |
| dataset_loader: DatasetLoader | None = None, | |
| ) -> None: | |
| self.store = store | |
| self.inference_provider = inference_provider | |
| self.judge_provider = judge_provider | |
| self.dataset_loader = dataset_loader or DatasetLoader() | |
| async def execute_run( | |
| self, | |
| run_id: str, | |
| project_config: dict[str, Any], | |
| task_config: dict[str, Any], | |
| prompt: dict[str, Any], | |
| on_progress: OnProgress | None = None, | |
| max_workers: int = 5, | |
| limit: int | None = None, | |
| completed_results: list[dict[str, Any]] | None = None, | |
| ) -> None: | |
| """Execute a full inference + evaluation run. | |
| Args: | |
| completed_results: Previously completed results for resume. | |
| Samples with matching sample_idx will be skipped. | |
| ์ํ ๋จ์ ์คํจ(์ถ๋ก /ํ๊ฐ)๋ ์ผํค๊ณ ์ํ(PARTIAL/eval_failed_count)๋ก ๊ธฐ๋กํ์ง๋ง, | |
| run ์ ํต์งธ๋ก ๋ชป ๋๊ฒ ๋ง๋๋ ์์ธ(๋ฐ์ดํฐ์ ๋ก๋ ์คํจ, ํ์ดํ๋ผ์ธ/ํ๊ฐ๊ธฐ ์์ฑ | |
| ์คํจ, ์คํ ์ด ์ฐ๊ธฐ ์คํจ ๋ฑ)๋ ์ํ๋ฅผ FAILED ๋ก ๋จ๊ธด ๋ค ๊ทธ๋๋ก ์ฌ์ ํํ๋ค -- | |
| "๋ฌด์จ ์ผ์ด ์์ด๋ ์์ธ ์์ด ๋ฐํ"์ด ์๋๋ค. ํธ์ถ์๋ ์ด ํจ์๊ฐ raise ํ ์ | |
| ์๋ค๊ณ ๊ฐ์ ํด์ผ ํ๋ค. | |
| """ | |
| try: | |
| # Build set of already-completed sample indices for resume | |
| completed_by_idx: dict[int, dict[str, Any]] = {} | |
| if completed_results: | |
| for r in completed_results: | |
| completed_by_idx[r["sample_idx"]] = r | |
| # Update status to running | |
| await self.store.update_run( | |
| run_id, | |
| { | |
| "status": RunStatus.RUNNING, | |
| "started_at": datetime.now(timezone.utc), | |
| }, | |
| ) | |
| # Load dataset | |
| dataset_config = project_config.get("dataset", {}) | |
| source = dataset_config.get("source", "huggingface") | |
| repo_name = dataset_config.get("repo", "") | |
| dataset_path = task_config.get("dataset_path", "") | |
| data = self.dataset_loader.load_jsonl(repo_name, dataset_path, source=source) | |
| if limit and limit < len(data): | |
| data = data[:limit] | |
| remaining = len(data) - len(completed_by_idx) | |
| await self.store.update_run(run_id, {"total_samples": len(data)}) | |
| if on_progress: | |
| await on_progress({"type": "dataset_loaded", "total": len(data)}) | |
| if completed_by_idx: | |
| logger.info(f"Resuming: {len(completed_by_idx)} done, {remaining} remaining") | |
| # Create pipeline | |
| # project_dir: v24 ์ tool_calling_judge ๊ฐ pmi_lookup ์๋๊ฒฝ๋ก๋ฅผ ํ ๋๋ง | |
| # ์ด๋ค (ยง5-F). dataset_loader.base_dir ์ด projects ๋ฃจํธ์ด๋ฏ๋ก project ์ด๋ฆ์ | |
| # ๋ถ์ด๋ฉด project_dir ์ด ๋๋ค -- CLI(`_start_local`)๊ฐ ์ด๋ฏธ ์ฐ๋ ๊ฒ๊ณผ ๊ฐ์ ๊ด๋ก. | |
| project_name = project_config.get("name") | |
| project_dir = ( | |
| self.dataset_loader.base_dir / project_name | |
| if self.dataset_loader.base_dir and project_name | |
| else None | |
| ) | |
| # config_dir: ๋ ํฌ ๊ด๋ฆฌ ํ๋ก์ ํธ๋ฉด project_loader ๊ฐ ์ฑ์ ๋ config ์ ๋ณธ | |
| # ๊ฒฝ๋ก(03-evaluation). ์นํ ์ฌ์ ๊ฐ์ ์ง์ ์์ฐ์ด ์ฌ๊ธฐ์ ์จ๋ค. | |
| pipeline = create_pipeline( | |
| pipeline_type=task_config.get("pipeline", "single_step"), | |
| input_fields=task_config.get("input_fields", []), | |
| pipeline_config=task_config.get("pipeline_config"), | |
| prompts=prompt.get("step_prompts", {}), | |
| dataset_loader=self.dataset_loader, | |
| project_dir=project_dir, | |
| config_dir=project_config.get("config_dir"), | |
| ) | |
| # Run inference with concurrency control | |
| semaphore = asyncio.Semaphore(max_workers) | |
| completed = len(completed_by_idx) | |
| async def process_sample(idx: int, sample: dict) -> dict[str, Any]: | |
| nonlocal completed | |
| # Skip already-completed samples (resume) | |
| if idx in completed_by_idx: | |
| existing = completed_by_idx[idx] | |
| golden_field = task_config.get("golden_field") | |
| golden = sample.get(golden_field, "") if golden_field else "" | |
| return {**existing, "_golden_raw": golden} | |
| async with semaphore: | |
| input_data = { | |
| field: sample.get(field, "") | |
| for field in task_config.get("input_fields", []) | |
| } | |
| # Get golden reference | |
| golden_field = task_config.get("golden_field") | |
| golden_fields = task_config.get("golden_fields") | |
| if golden_field: | |
| golden = sample.get(golden_field, "") | |
| elif golden_fields: | |
| golden = {k: sample.get(v, "") for k, v in golden_fields.items()} | |
| else: | |
| golden = "" | |
| eval_sample = EvalSample(input=input_data, reference=golden) | |
| try: | |
| eval_sample = await pipeline.run( | |
| sample=eval_sample, | |
| prompts=prompt.get("system_prompt", ""), | |
| provider=self.inference_provider, | |
| model=prompt.get("model", "solar-pro2"), | |
| temperature=prompt.get("temperature", 0.0), | |
| max_tokens=prompt.get("max_tokens", 8000), | |
| reasoning_effort=prompt.get("reasoning_effort"), | |
| messages=prompt.get("messages"), | |
| ) | |
| except Exception as e: | |
| ts = datetime.now().strftime("%H:%M:%S") | |
| logger.warning(f"[{ts}] Sample {idx} failed: {type(e).__name__}: {e}") | |
| raise | |
| # ์ ์ฅ ํํ๋ ์ด์ ๊ณผ ๊ฐ์ 8ํค dict (eval-store ์ resultRowSchema ๊ฐ | |
| # ์ฝ๋ ์ด๋ฆ๋ค๊ณผ ํ์ ํธํ) + trace ๋ฅผ ์ถ๊ฐํ๋ค. trace ๋ artifacts | |
| # ์ ์ฒด๋ฅผ ๊ทธ๋๋ก ๋ฃ๋๋ค -- step_outputs ๋ฟ ์๋๋ผ v24 ๊ฐ ์ฑ์ฐ๋ | |
| # judge_decisions/judge_tool_calls/self_consistency_runs/corrections | |
| # ๋ ์ฌ๊ธฐ ์ ๋ฃ์ผ๋ฉด ์ ์ฅ ์ง์ ์ ํต์งธ๋ก ๋ฒ๋ ค์ง๋ค (์ค์ judge LLM ํธ์ถยท | |
| # self-consistency ๋ฐ๋ณต ํธ์ถ ๋น์ฉ์ด ๋๊ฐ ์ฐ์ถ๋ฌผ์ด๋ค). ํน์ ํค๋ง | |
| # ํ๋์ฝ๋ฉํด ์ฎ๊ธฐ๋ฉด ๋ค์์ ํ์ดํ๋ผ์ธ์ด ์ artifacts ํค๋ฅผ ์ถ๊ฐํ | |
| # ๋๋ง๋ค ์ฌ๊ธฐ๋ฅผ ๋ ๊ณ ์ณ์ผ ํ๋ฏ๋ก ํต์งธ๋ก ๋๊ธด๋ค -- | |
| # resultRowSchema.trace ๋ .passthrough() ๋ผ ์ฌ๋ถ ํค๋ฅผ ๊ทธ๋๋ก ๋ฐ๋๋ค | |
| # (source-schemas.ts:223-228). | |
| run_result = { | |
| "run_id": run_id, | |
| "sample_idx": idx, | |
| "input": input_data, | |
| "output": eval_sample.output, | |
| "golden": golden, | |
| "input_tokens": eval_sample.artifacts.get("usage", {}).get( | |
| "prompt_tokens", 0 | |
| ), | |
| "output_tokens": eval_sample.artifacts.get("usage", {}).get( | |
| "completion_tokens", 0 | |
| ), | |
| "inference_time_ms": eval_sample.artifacts.get("inference_time_ms", 0), | |
| "trace": { | |
| "step_outputs": {}, | |
| **eval_sample.artifacts, | |
| }, | |
| } | |
| await self.store.insert_run_result(run_result) | |
| completed += 1 | |
| ts = datetime.now().strftime("%H:%M:%S") | |
| logger.info(f"[{ts}] Sample {idx} completed ({completed}/{len(data)})") | |
| await self.store.update_run(run_id, {"completed_samples": completed}) | |
| if on_progress: | |
| await on_progress( | |
| { | |
| "type": "inference_progress", | |
| "completed": completed, | |
| "total": len(data), | |
| "sample_idx": idx, | |
| } | |
| ) | |
| return {**run_result, "_golden_raw": golden} | |
| tasks = [process_sample(i, sample) for i, sample in enumerate(data)] | |
| results = await asyncio.gather(*tasks, return_exceptions=True) | |
| # Filter out exceptions | |
| valid_results = [r for r in results if isinstance(r, dict)] | |
| errors = [(i, r) for i, r in enumerate(results) if isinstance(r, Exception)] | |
| if errors: | |
| logger.warning(f"Run {run_id}: {len(errors)}/{len(results)} samples failed") | |
| for sample_idx, err in errors: | |
| logger.warning(f" Sample {sample_idx} error: {type(err).__name__}: {err}") | |
| if on_progress: | |
| await on_progress({"type": "inference_complete", "total": len(valid_results)}) | |
| # Run evaluation | |
| await self.store.update_run(run_id, {"status": RunStatus.EVALUATING}) | |
| if on_progress: | |
| await on_progress({"type": "evaluation_start", "total": len(valid_results)}) | |
| evaluator = create_evaluator(task_config.get("evaluator", {"type": "llm_judge"})) | |
| # eval_results: ์ฑ๊ณตํ evaluate() ๋ฐํ๊ฐ + "sample_idx"(์ ๋ณธ, ยง0.5). | |
| # eval_failures: ์ฑ์ ์์ฒด๊ฐ ์ ๋ ์ํ -- aggregate() ์๋ ์ ๋ ์ ๋๊ธด๋ค | |
| # (lcs_diff.aggregate() ์ฒ๋ผ r["details"]["tp"] ๋ฅผ ์ง์ ์ฝ๋ ๊ตฌํ์ด ์คํจ | |
| # ํญ๋ชฉ์ ๋ง๋๋ฉด KeyError ๋ก ์ฃฝ๋๋ค, ยง0.1). | |
| eval_results: list[dict[str, Any]] = [] | |
| eval_failures: list[dict[str, Any]] = [] | |
| for i, result in enumerate(valid_results): | |
| # F6: enumerate ์์น๊ฐ ์๋๋ผ result["sample_idx"] ๊ฐ ์ ๋ณธ์ด๋ค -- | |
| # ์ค๊ฐ ์ํ์ด ์ถ๋ก ์์ ์คํจํ๋ฉด valid_results ์ ๋ฆฌ์คํธ ์์น์ ์๋ณธ | |
| # sample_idx ๊ฐ ์ด๊ธ๋๋ค. | |
| sample_idx = result["sample_idx"] | |
| try: | |
| eval_sample = build_eval_sample_from_result( | |
| result, reference=result["_golden_raw"] | |
| ) | |
| evaluator.validate_required_fields(eval_sample) | |
| eval_result = await evaluator.evaluate( | |
| sample=eval_sample, | |
| provider=self.judge_provider, | |
| ) | |
| eval_results.append({**eval_result, "sample_idx": sample_idx}) | |
| if on_progress and i % 10 == 0: | |
| await on_progress( | |
| { | |
| "type": "evaluation_progress", | |
| "completed": i + 1, | |
| "total": len(valid_results), | |
| } | |
| ) | |
| except Exception as e: | |
| logger.warning(f"Evaluation failed for sample {sample_idx}: {e}") | |
| eval_failures.append({"sample_idx": sample_idx, "error": str(e)}) | |
| # Aggregate and save evaluation -- eval_results ์๋ ์คํจ ํญ๋ชฉ์ด ์ ์์ฌ | |
| # ์์ผ๋ฏ๋ก ๊ธฐ์กด aggregate() ๊ตฌํ์ด ๊ทธ๋๋ก ๋์ํ๋ค. | |
| aggregated = evaluator.aggregate(eval_results) | |
| # ์ฑ์ ์ ์ฑ๊ณตํ ์ํ์ด ํ๋๋ ์์ผ๋ฉด ์ ์ ์๋ฆฌ๋ฅผ ๋น์ด๋ค -- aggregate() ๋ ๋น | |
| # ์ ๋ ฅ์ 0.0 ์ ๋๋ ค์ฃผ๋๋ฐ, ๊ทธ๊ฑด "0์ ์ ๋ฐ์๋ค"๋ ์ธก์ ๊ฐ์ด๋ผ "์ธก์ ์์ฒด๊ฐ | |
| # ์์๋ค"์ ๊ตฌ๋ถ๋์ง ์๋๋ค. judge ๊ฐ ํต์งธ๋ก ์ฃฝ์ run ์ด evalhub ์ฐจํธ์์ | |
| # ํ์ง ๊ธ๋ฝ์ผ๋ก ๋ณด์ด๋ฉด F1 ์ ๋ฐ๋ง ๊ณ ์น ์ ์ด๋ค. | |
| overall_score = aggregated.get("overall_score", 0.0) if eval_results else None | |
| eval_id = await self.store.create_evaluation( | |
| { | |
| "run_id": run_id, | |
| "scores": aggregated.get("scores", {}), | |
| "overall_score": overall_score, | |
| "eval_model": "gpt-4o", | |
| "eval_success_count": len(eval_results), | |
| "eval_failed_count": len(eval_failures), | |
| "failed_sample_indices": [f["sample_idx"] for f in eval_failures], | |
| } | |
| ) | |
| # Save per-sample eval details -- ์ฑ๊ณต/์คํจ ๋ ๋ค ํ ํ์ฉ ๋จ๊ธด๋ค. | |
| # results.jsonl ๊ณผ eval_details.jsonl ์ด ํญ์ ๊ฐ์ sample_idx ์งํฉ์ | |
| # ๊ฐ๋ฆฌํค๊ฒ ํด์ ๋ถ๋ถ์ ์ผ๋ก๋ง ์ฑ์ ๋ run ์์ "์ด ์ํ์ ์ ์ ๋ณด์ด์ง"๋ฅผ | |
| # ์์ค๋ค. | |
| eval_detail_docs = [] | |
| for er in eval_results: | |
| eval_detail_docs.append( | |
| { | |
| "evaluation_id": eval_id, | |
| "sample_idx": er["sample_idx"], | |
| "category_scores": er.get("category_scores", {}), | |
| "error_counts": { | |
| k: v.get("error_count", 0) | |
| for k, v in er.get("details", {}).items() | |
| if isinstance(v, dict) | |
| }, | |
| "severity": er.get("severity", ""), | |
| "score": er.get("score", 0.0), | |
| } | |
| ) | |
| for f in eval_failures: | |
| eval_detail_docs.append( | |
| { | |
| "evaluation_id": eval_id, | |
| "sample_idx": f["sample_idx"], | |
| "category_scores": {}, | |
| "error_counts": {}, | |
| "severity": None, | |
| # 0.0 ์ด ์๋๋ผ None -- ์ฑ์ ์คํจ๋ฅผ ์ต์ ์ ์์ ๊ตฌ๋ถํ๋ค (F1). | |
| "score": None, | |
| "error": f["error"], | |
| } | |
| ) | |
| await self.store.insert_eval_details(eval_detail_docs) | |
| # Mark run as completed/partial -- ์ถ๋ก ์ด ์ผ๋ถ ์ํ์์ ์คํจํ์ผ๋ฉด | |
| # PARTIAL, ํ๊ฐ ์คํจ๋ ์ด ์ํ์ ์ํฅ์ ์ฃผ์ง ์๋๋ค(ํ๊ฐ ์ฑ๊ณต/์คํจ๋ | |
| # ์ eval_success_count/eval_failed_count ๋ก๋ง ํํํ๋ค, ยง0.4). errors | |
| # ์ ์ธ๋ฑ์ค๋ asyncio.gather ๊ฐ tasks ์์๋ฅผ ๋ณด์กดํ๋ฏ๋ก ์ด๋ฏธ ์ง์ง | |
| # sample_idx ๋ค. | |
| final_status = RunStatus.PARTIAL if errors else RunStatus.COMPLETED | |
| await self.store.update_run( | |
| run_id, | |
| { | |
| "status": final_status, | |
| "completed_at": datetime.now(timezone.utc), | |
| "failed_samples": len(errors), | |
| "failed_sample_indices": [i for i, _ in errors], | |
| }, | |
| ) | |
| if on_progress: | |
| await on_progress( | |
| { | |
| "type": "done", | |
| "overall_score": aggregated.get("overall_score", 0.0), | |
| "scores": aggregated.get("scores", {}), | |
| } | |
| ) | |
| except Exception as e: | |
| logger.exception(f"Run {run_id} failed") | |
| await self.store.update_run( | |
| run_id, | |
| { | |
| "status": RunStatus.FAILED, | |
| "completed_at": datetime.now(timezone.utc), | |
| }, | |
| ) | |
| if on_progress: | |
| await on_progress({"type": "error", "error": str(e)}) | |
| # F4: ์ํ๋ฅผ FAILED ๋ก ๋จ๊ธฐ๊ณ ์งํ ์ํฉ๊น์ง ์๋ฆฐ ๋ค **๊ทธ๋๋ก ์ฌ์ ํํ๋ค**. | |
| # ์ฌ๊ธฐ์ ์ผํค๋ฉด ํธ์ถ์(CLI)๋ "Results saved to ..." ๋ฅผ ์ฐ๊ณ evalhub ์ ์ฌ๊น์ง | |
| # ์๋ํ ๋ค์ exit 0 ์ ๋ธ๋ค -- ๋ฐ์ดํฐ์ ์กฐ์ฐจ ๋ชป ์ฝ์ run ์ด ์ฑ๊ณต์ผ๋ก ๋ณด์ธ๋ค. | |
| raise | |