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:
酱牛肉 2026-07-30 12:09:57 +08:00
parent f10a1c9935
commit 4e06d59aca
5 changed files with 109 additions and 8 deletions

View file

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

View file

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

View file

@ -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."""

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

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