mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-11 03:40:03 +00:00
feat(evaluation): add check_job_count interface and report search calls in benchmarks
- Extract _counter_key from BaseJob._record_call for reusable counter lookup - Add reme.utils.evaluation_interface.check_job_count read-only helper - Track and report average search calls per query in beam and longmemeval benchmarks
This commit is contained in:
parent
f10a1c9935
commit
4e06d59aca
5 changed files with 109 additions and 8 deletions
|
|
@ -252,13 +252,17 @@ async def answer_question_agentic(app, question: str) -> tuple[str, dict]:
|
|||
|
||||
Returns (answer, metadata)
|
||||
"""
|
||||
from reme.utils.evaluation_interface import check_job_count
|
||||
|
||||
search_calls_before = check_job_count("search", app.context)
|
||||
query_resp = await app.run_job(
|
||||
"agentic_answer",
|
||||
query=question,
|
||||
)
|
||||
answer = (query_resp.answer or "").strip()
|
||||
search_calls = check_job_count("search", app.context) - search_calls_before
|
||||
|
||||
return answer, {"mode": "agentic"}
|
||||
return answer, {"mode": "agentic", "search_calls": search_calls}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -448,6 +452,9 @@ async def evaluate_case(eval_config: dict, case_id: str, eval_only: bool = False
|
|||
if not agentic_answer:
|
||||
agentic_answer = "(no answer generated)"
|
||||
logger.info(f"[Case {case_id}] Agentic answer: {agentic_answer[:200]}...")
|
||||
logger.info(
|
||||
f"[Case {case_id}] Agentic search calls: {agentic_meta.get('search_calls', 0)}",
|
||||
)
|
||||
|
||||
# Judge agentic answer
|
||||
logger.info(f"[Case {case_id}] Judging agentic ({q_type})...")
|
||||
|
|
@ -678,6 +685,7 @@ def main( # pylint: disable=too-many-statements
|
|||
type_binary_scores: dict[str, list[float]] = {}
|
||||
all_scores: list[float] = []
|
||||
all_binary_scores: list[float] = []
|
||||
all_search_calls: list[int] = []
|
||||
|
||||
for case_result in results:
|
||||
if "error" in case_result:
|
||||
|
|
@ -700,6 +708,7 @@ def main( # pylint: disable=too-many-statements
|
|||
type_binary_scores[qtype].append(binary_score)
|
||||
all_scores.append(score)
|
||||
all_binary_scores.append(binary_score)
|
||||
all_search_calls.append(q.get("agentic_metadata", {}).get("search_calls", 0))
|
||||
|
||||
print("\n ── AGENTIC ──")
|
||||
if all_scores:
|
||||
|
|
@ -713,6 +722,8 @@ def main( # pylint: disable=too-many-statements
|
|||
binary_overall = sum(all_binary_scores) / len(all_binary_scores) if all_binary_scores else 0
|
||||
print(f" {'-'*38}")
|
||||
print(f" {'OVERALL':<40s}: {overall:.3f} binary={binary_overall:.3f} ({len(all_scores)} Qs)")
|
||||
avg_search_calls = sum(all_search_calls) / len(all_search_calls)
|
||||
print(f" Average search calls/query: {avg_search_calls:.2f}")
|
||||
else:
|
||||
print(" (no results)")
|
||||
|
||||
|
|
|
|||
|
|
@ -259,6 +259,7 @@ async def evaluate_item(item: dict, eval_config: dict, item_index: int, eval_onl
|
|||
"""
|
||||
from reme import Application
|
||||
from reme.config import resolve_app_config
|
||||
from reme.utils.evaluation_interface import check_job_count
|
||||
|
||||
reme_cfg = eval_config["reme"]
|
||||
dream_trigger_hour = reme_cfg.get("dream_trigger_hour", 23)
|
||||
|
|
@ -432,16 +433,19 @@ async def evaluate_item(item: dict, eval_config: dict, item_index: int, eval_onl
|
|||
f"[Item {item_index}] Asking (agentic): {question[:80]}... query_time={query_time}",
|
||||
)
|
||||
|
||||
search_calls_before = check_job_count("search", app.context)
|
||||
query_resp = await app.run_job(
|
||||
"agentic_answer",
|
||||
query=question,
|
||||
query_time=query_time,
|
||||
)
|
||||
agentic_search_calls = check_job_count("search", app.context) - search_calls_before
|
||||
agentic_response = (query_resp.answer or "").strip()
|
||||
if not agentic_response:
|
||||
agentic_response = "(no answer generated)"
|
||||
|
||||
logger.info(f"[Item {item_index}] Agentic response: {agentic_response[:200]}...")
|
||||
logger.info(f"[Item {item_index}] Agentic search calls: {agentic_search_calls}")
|
||||
|
||||
# ── Phase 5: Judge agentic response (via answer_judge_step) ──────────
|
||||
logger.info(f"[Item {item_index}] Judging agentic (binary, type={item['question_type']})...")
|
||||
|
|
@ -464,6 +468,7 @@ async def evaluate_item(item: dict, eval_config: dict, item_index: int, eval_onl
|
|||
"ground_truth": item["answer"],
|
||||
"agentic_response": agentic_response,
|
||||
"agentic_judgment": agentic_judgment,
|
||||
"agentic_search_calls": agentic_search_calls,
|
||||
"sessions_ingested": len(sorted_sessions),
|
||||
"dreams_triggered": len(dream_dates_triggered),
|
||||
}
|
||||
|
|
@ -711,6 +716,8 @@ def _print_summary(results: list[dict], start_time: float) -> None:
|
|||
# Agentic stats
|
||||
print("\n ── Agentic (ReAct) ──")
|
||||
print(f" Overall accuracy: {agentic_correct}/{total} ({100*agentic_correct/total:.1f}%)")
|
||||
avg_search_calls = sum(r.get("agentic_search_calls", 0) for r in results) / total if total else 0
|
||||
print(f" Average search calls/query: {avg_search_calls:.2f}")
|
||||
print(" Per-type accuracy:")
|
||||
for qtype, stats in sorted(agentic_type_stats.items()):
|
||||
acc = 100 * stats["correct"] / stats["total"] if stats["total"] else 0
|
||||
|
|
|
|||
|
|
@ -58,21 +58,23 @@ class BaseJob(BaseComponent):
|
|||
# dict(params) copies kwargs so steps cannot mutate the shared spec.
|
||||
return [step_cls(**dict(params)) for step_cls, params in self.step_specs]
|
||||
|
||||
def _record_call(self, entry_class: type["BaseJob"]) -> None:
|
||||
"""Increment this job's application-lifetime call counter.
|
||||
def _counter_key(self, entry_class: type["BaseJob"]) -> list[str]:
|
||||
"""Build this job's application-lifetime call counter key.
|
||||
|
||||
The counter path groups calls by the built-in job implementation that
|
||||
accepted the call, then preserves any project-specific subclasses
|
||||
between that implementation and the concrete job class.
|
||||
"""
|
||||
metadata = getattr(self.app_context, "metadata", None)
|
||||
if not isinstance(metadata, dict):
|
||||
return
|
||||
|
||||
mro = type(self).mro()
|
||||
entry_index = mro.index(entry_class)
|
||||
subclass_path = [cls.__name__ for cls in mro[:entry_index]]
|
||||
global_counter_inc(metadata, ["__job_counter", entry_class.__name__, *subclass_path, self.name])
|
||||
return ["__job_counter", entry_class.__name__, *subclass_path, self.name]
|
||||
|
||||
def _record_call(self, entry_class: type["BaseJob"]) -> None:
|
||||
"""Increment this job's application-lifetime call counter."""
|
||||
metadata = getattr(self.app_context, "metadata", None)
|
||||
if isinstance(metadata, dict):
|
||||
global_counter_inc(metadata, self._counter_key(entry_class))
|
||||
|
||||
async def __call__(self, **kwargs) -> Response:
|
||||
"""Run all steps in order, capturing any failure into the response."""
|
||||
|
|
|
|||
30
reme/utils/evaluation_interface.py
Normal file
30
reme/utils/evaluation_interface.py
Normal file
|
|
@ -0,0 +1,30 @@
|
|||
"""Read-only evaluation helpers for application job execution statistics."""
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from ..components.job import BackgroundJob, BaseJob, CronJob, StreamJob
|
||||
from .counter import global_counter_get
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..components.application_context import ApplicationContext
|
||||
|
||||
|
||||
_JOB_ENTRY_CLASSES = {CronJob, StreamJob, BackgroundJob, BaseJob}
|
||||
|
||||
|
||||
def check_job_count(job_name: str, app_context: "ApplicationContext") -> int:
|
||||
"""Return the application-lifetime execution count for a registered job.
|
||||
|
||||
``app_context`` scopes the lookup because ReMe does not maintain a global
|
||||
current Application instance. Unknown job names use the same ``KeyError``
|
||||
contract as :meth:`Application.run_job`.
|
||||
"""
|
||||
job = app_context.jobs.get(job_name)
|
||||
if job is None:
|
||||
raise KeyError(f"Job '{job_name}' not found")
|
||||
|
||||
entry_class = next((cls for cls in type(job).mro() if cls in _JOB_ENTRY_CLASSES), None)
|
||||
if entry_class is None:
|
||||
raise TypeError(f"Job '{job_name}' does not inherit from a supported job implementation")
|
||||
# pylint: disable-next=protected-access
|
||||
return global_counter_get(app_context.metadata, job._counter_key(entry_class))
|
||||
51
tests/unit/test_evaluation_interface.py
Normal file
51
tests/unit/test_evaluation_interface.py
Normal file
|
|
@ -0,0 +1,51 @@
|
|||
"""Tests for read-only job execution count evaluation helpers."""
|
||||
|
||||
import asyncio
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
from reme.components.job import BaseJob, StreamJob
|
||||
from reme.utils.evaluation_interface import check_job_count
|
||||
|
||||
|
||||
def test_check_job_count_reads_registered_base_job_count():
|
||||
"""check_job_count returns the number of completed BaseJob invocations."""
|
||||
|
||||
async def run():
|
||||
app_context = SimpleNamespace(metadata={}, jobs={})
|
||||
job = BaseJob(name="search", app_context=app_context)
|
||||
app_context.jobs[job.name] = job
|
||||
|
||||
await job()
|
||||
await job()
|
||||
|
||||
assert check_job_count("search", app_context) == 2
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
|
||||
def test_check_job_count_resolves_custom_job_inheritance_path():
|
||||
"""check_job_count resolves counters recorded under subclassed job paths."""
|
||||
|
||||
async def run():
|
||||
class ProjectStreamJob(StreamJob):
|
||||
"""Project-specific StreamJob subclass used to exercise MRO lookup."""
|
||||
|
||||
app_context = SimpleNamespace(metadata={}, jobs={})
|
||||
job = ProjectStreamJob(name="chat", app_context=app_context)
|
||||
app_context.jobs[job.name] = job
|
||||
|
||||
await job(stream_queue=asyncio.Queue())
|
||||
|
||||
assert check_job_count("chat", app_context) == 1
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
|
||||
def test_check_job_count_rejects_unknown_job_name():
|
||||
"""Unknown job names raise KeyError, matching Application.run_job."""
|
||||
app_context = SimpleNamespace(metadata={}, jobs={})
|
||||
|
||||
with pytest.raises(KeyError, match="Job 'missing' not found"):
|
||||
check_job_count("missing", app_context)
|
||||
Loading…
Add table
Reference in a new issue