From eba9279af2eb9d6ebb6dca3dd32792b42d6707fa Mon Sep 17 00:00:00 2001 From: 3un01a <3un01a.labs@gmail.com> Date: Tue, 24 Feb 2026 06:55:59 +0900 Subject: [PATCH] Oolong Benchmark (#323) * (feat) Add Oolong Benchmarks * (fix) Address issues to fix basedpyright and coderabbit comments * (fix) Address basedpyrwright additional warnings * (fix) Address additional coderabbit issues * (fix) Replace huggingface data loading to local filesystem-based * (fix) Address coderabbit issues regarding data paths * fix: Align with test harness conventions * fix: Code Review Comments * fix: stream data rather than load all at once --------- Co-authored-by: 3un01a <3un01a@plasticlabs.ai> Co-authored-by: Vineeth Voruganti <13438633+VVoruganti@users.noreply.github.com> --- pyproject.toml | 1 + scripts/run_alembic_tests.py | 2 + tests/bench/.gitignore | 1 + tests/bench/README.md | 58 +++ tests/bench/oolong.py | 603 ++++++++++++++++++++++++++++ tests/bench/oolong_common.py | 750 +++++++++++++++++++++++++++++++++++ tests/bench/runner_common.py | 37 +- uv.lock | 2 + 8 files changed, 1447 insertions(+), 7 deletions(-) create mode 100644 tests/bench/oolong.py create mode 100644 tests/bench/oolong_common.py diff --git a/pyproject.toml b/pyproject.toml index 02879c06..e523d964 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -34,6 +34,7 @@ dependencies = [ "json-repair>=0.49.0", "turbopuffer>=1.8.1", "lancedb>=0.25.3", + "pyarrow>=19.0.0", "redis>=7.0.0,<8.0.0", "cashews[redis]==7.4.4", "scikit-learn>=1.6.0", diff --git a/scripts/run_alembic_tests.py b/scripts/run_alembic_tests.py index 8a2439e6..abc3656a 100755 --- a/scripts/run_alembic_tests.py +++ b/scripts/run_alembic_tests.py @@ -108,6 +108,7 @@ def main(): "run", "pytest", "tests/alembic/test_pipeline.py", + "-n0", ] elif revision_ids: # Build a -k expression to filter tests by revision ID @@ -125,6 +126,7 @@ def main(): "run", "pytest", "tests/alembic/test_pipeline.py", + "-n0", "-k", k_expression, ] diff --git a/tests/bench/.gitignore b/tests/bench/.gitignore index 1aa3c022..6b1b9f97 100644 --- a/tests/bench/.gitignore +++ b/tests/bench/.gitignore @@ -4,3 +4,4 @@ perf_metrics beam_data obexeval_data locomo_data +oolongeval_data diff --git a/tests/bench/README.md b/tests/bench/README.md index 48a4b649..02feebbb 100644 --- a/tests/bench/README.md +++ b/tests/bench/README.md @@ -6,6 +6,64 @@ This directory contains benchmarking tools for evaluating Honcho's long-term mem - **LongMemEval**: Tests memory retention across multi-session conversations - **BEAM**: Beyond a Million Tokens - comprehensive long-term memory evaluation across 10 memory abilities +- **LoCoMo**: Long conversation memory benchmark across multi-hop and temporal questions +- **OOLONG**: Long-context aggregation benchmark with `synth` and `real` variants + +## Benchmark Workflow + +Use a harness-first workflow for all benchmark runs: + +1. Start Honcho locally with the benchmark harness: + +```bash +python tests/bench/harness.py +``` + +2. Run one of the benchmark runners in another terminal: + +```bash +# LongMemEval +python -m tests.bench.longmem --test-file tests/bench/longmemeval_data/longmemeval_oracle.json + +# LoCoMo +python -m tests.bench.locomo --data-file tests/bench/locomo_data/locomo10.json + +# BEAM +python -m tests.bench.beam --context-length 100K +``` + +3. For OOLONG, point `--data-dir` at your local dataset clone: + +```bash +# OOLONG-synth +python -m tests.bench.oolong --variant synth --data-dir /path/to/oolong-synth + +# OOLONG-real +python -m tests.bench.oolong --variant real --data-dir /path/to/oolong-real + +# OOLONG-synth with label-augmented context (upstream optional mode) +python -m tests.bench.oolong --variant synth --data-dir /path/to/oolong-synth --labels +``` + +Notes for OOLONG runs: + +- By default, synth uses `context_window_text` (upstream baseline behavior). +- Use `--labels` to switch synth ingestion to `context_window_text_with_labels`. +- Default `--min-context-len` is `1024` and filtering uses strict `>` matching upstream. + +Expected local dataset layout: + +```text +oolong-synth/ + data/ + test-*.parquet + validation-*.parquet + +oolong-real/ + dnd/ + test.jsonl + validation.jsonl +``` ## Development Harness diff --git a/tests/bench/oolong.py b/tests/bench/oolong.py new file mode 100644 index 00000000..d7edc48c --- /dev/null +++ b/tests/bench/oolong.py @@ -0,0 +1,603 @@ +""" +Honcho OOLONG benchmark runner. + +Evaluates long-context reasoning and aggregation on: +- OOLONG-synth: synthetic ICL aggregation tasks +- OOLONG-real: D&D transcript aggregation tasks +""" + +import argparse +import sys +import time +from datetime import datetime +from pathlib import Path +from typing import Any, TypedDict, cast + +from dotenv import load_dotenv +from honcho.api_types import MessageCreateParams +from honcho.session import Session, SessionPeerConfig + +from src.config import settings + +from .oolong_common import ( + calculate_context_length, + calculate_task_statistics, + calculate_timing_statistics, + format_duration, + load_oolong_real_dataset, + load_oolong_synth_dataset, + parse_real_answer, + parse_real_context_messages, + parse_synth_answer, + parse_synth_context_messages, + score_real_response, + score_synth_response, + write_json_summary, +) +from .runner_common import ( + BaseRunner, + ItemContext, + RunnerConfig, + add_common_arguments, + validate_common_arguments, +) + +load_dotenv() + +CONTEXT_SIZE_MAP: dict[str, int] = { + "1K": 1024, + "2K": 2 * 1024, + "4K": 4 * 1024, + "8K": 8 * 1024, + "16K": 16 * 1024, + "32K": 32 * 1024, + "64K": 64 * 1024, + "128K": 128 * 1024, + "256K": 256 * 1024, + "512K": 512 * 1024, + "1M": 1024 * 1024, + "2M": 2 * 1024 * 1024, + "4M": 4 * 1024 * 1024, +} + + +def parse_context_size(size_str: str) -> int: + """Parse a context-size string into an exact token count.""" + normalized = size_str.strip().upper() + + if normalized.isdigit(): + value = int(normalized) + if value <= 0: + raise ValueError("Context size must be positive") + return value + + if normalized in CONTEXT_SIZE_MAP: + return CONTEXT_SIZE_MAP[normalized] + + valid_sizes = ", ".join(CONTEXT_SIZE_MAP) + raise ValueError( + f"Invalid context size '{size_str}'. Use one of [{valid_sizes}] or a positive integer token count." + ) + + +class QueryResult(TypedDict): + """Query execution result for one OOLONG example.""" + + question: str + expected_answer: str + actual_response: str + score: float + context_length_tokens: int + + +class TestResult(TypedDict): + """Single OOLONG example result.""" + + question_id: str + context_window_id: str + task_group: str + dataset: str + answer_type: str + passed: bool + score: float + error: str | None + start_time: float + end_time: float + duration_seconds: float + query_executed: QueryResult | None + output_lines: list[str] + + +class OolongRunner(BaseRunner[TestResult]): + """Execute OOLONG benchmark examples through the shared runner framework.""" + + variant: str + data_dir: Path + split: str + merge_sessions: bool + max_examples: int | None + min_context_len: int | None + max_context_len: int | None + context_window_id: str | None + use_labels: bool + + def __init__( + self, + config: RunnerConfig, + variant: str, + data_dir: Path, + split: str, + merge_sessions: bool, + max_examples: int | None = None, + min_context_len: int | None = None, + max_context_len: int | None = None, + context_window_id: str | None = None, + use_labels: bool = False, + ): + self.variant = variant + self.data_dir = data_dir + self.split = split + self.merge_sessions = merge_sessions + self.max_examples = max_examples + self.min_context_len = min_context_len + self.max_context_len = max_context_len + self.context_window_id = context_window_id + self.use_labels = use_labels + super().__init__(config) + + def get_metrics_prefix(self) -> str: + return "oolong" + + def load_items(self) -> list[Any]: + if self.variant == "synth": + dataset = load_oolong_synth_dataset( + split=self.split, + data_dir=self.data_dir, + max_context_len=self.max_context_len, + min_context_len=self.min_context_len, + max_examples=self.max_examples, + context_window_id=self.context_window_id, + ) + else: + dataset = load_oolong_real_dataset( + split=self.split, + data_dir=self.data_dir, + max_context_len=self.max_context_len, + min_context_len=self.min_context_len, + max_examples=self.max_examples, + context_window_id=self.context_window_id, + ) + return [dataset[i] for i in range(len(dataset))] + + def get_workspace_id(self, item: Any) -> str: + return f"oolong_{self.variant}_{item['id']}" + + def get_session_id(self, item: Any, workspace_id: str) -> str: + return f"{workspace_id}_session" + + async def setup_peers(self, ctx: ItemContext, item: Any) -> None: + ctx.peers["user"] = await ctx.honcho_client.aio.peer(id="user") + + async def setup_session(self, ctx: ItemContext, item: Any) -> None: + if not self.merge_sessions: + ctx.session = None + return + + user_peer = ctx.peers["user"] + ctx.session = await ctx.honcho_client.aio.session( + id=ctx.session_id, configuration=self._get_session_configuration() + ) + await ctx.session.aio.add_peers( + [(user_peer, SessionPeerConfig(observe_me=True, observe_others=False))] + ) + + async def _add_messages_to_session( + self, session: Session, user_peer: Any, messages: list[dict[str, Any]] + ) -> None: + honcho_messages: list[MessageCreateParams] = [] + for msg in messages: + honcho_messages.append( + user_peer.message( + content=msg["content"], + metadata=msg.get("metadata"), + ) + ) + + for i in range(0, len(honcho_messages), 100): + batch = honcho_messages[i : i + 100] + await session.aio.add_messages(batch) + + async def ingest_messages(self, ctx: ItemContext, item: Any) -> int: + context_text = item["context_window_text"] + if self.variant == "synth": + if self.use_labels: + context_text = item.get("context_window_text_with_labels", context_text) + messages = parse_synth_context_messages(context_text) + else: + messages = parse_real_context_messages(context_text) + + user_peer = ctx.peers["user"] + + if self.merge_sessions: + if ctx.session is None: + raise ValueError("Merged mode requires a configured session") + await self._add_messages_to_session(ctx.session, user_peer, messages) + return len(messages) + + chunk_size = 200 + session_ids: list[str] = [] + for idx, start in enumerate(range(0, len(messages), chunk_size)): + chunk = messages[start : start + chunk_size] + session_id = f"{ctx.workspace_id}_session_{idx + 1}" + session = await ctx.honcho_client.aio.session( + id=session_id, configuration=self._get_session_configuration() + ) + await session.aio.add_peers( + [(user_peer, SessionPeerConfig(observe_me=True, observe_others=False))] + ) + await self._add_messages_to_session(session, user_peer, chunk) + session_ids.append(session_id) + + ctx.peers["_session_ids"] = session_ids + return len(messages) + + def get_dream_observers(self, item: Any) -> list[str]: + return ["user"] + + def get_dream_session_ids(self, ctx: ItemContext, _item: Any) -> list[str]: + if self.merge_sessions: + return [ctx.session_id] + + session_ids = ctx.peers.get("_session_ids") + if not isinstance(session_ids, list) or not session_ids: + raise ValueError( + "Non-merged OOLONG mode requires at least one chunk session ID for dreams" + ) + + session_ids_typed = cast(list[object], session_ids) + cleaned_session_ids: list[str] = [] + for maybe_session_id in session_ids_typed: + if isinstance(maybe_session_id, str) and maybe_session_id: + cleaned_session_ids.append(maybe_session_id) + if not cleaned_session_ids: + raise ValueError( + "Non-merged OOLONG mode has no valid chunk session IDs for dreams" + ) + return cleaned_session_ids + + async def execute_questions(self, ctx: ItemContext, item: Any) -> TestResult: + start_time = time.time() + question_id = item["id"] + context_window_id = item["context_window_id"] + question = item["question"] + answer_str = item["answer"] + + if self.variant == "synth": + task_group = item.get("task_group", "unknown") + dataset_name = item.get("dataset", "oolong-synth") + answer_type = item.get("answer_type", "unknown") + gold_answer = parse_synth_answer(answer_str) + else: + task_group = item.get("question_type", "unknown") + dataset_name = "oolong-real" + answer_type = "varied" + gold_answer = parse_real_answer(answer_str) + + context_text = item["context_window_text"] + if self.variant == "synth" and self.use_labels: + context_text = item.get("context_window_text_with_labels", context_text) + context_length = calculate_context_length(context_text) + + result: TestResult = { + "question_id": question_id, + "context_window_id": context_window_id, + "task_group": task_group, + "dataset": dataset_name, + "answer_type": answer_type, + "passed": False, + "score": 0.0, + "error": None, + "start_time": start_time, + "end_time": 0.0, + "duration_seconds": 0.0, + "query_executed": None, + "output_lines": [], + } + + user_peer = ctx.peers["user"] + try: + chat_kwargs: dict[str, Any] = {} + if self.config.reasoning_level: + chat_kwargs["reasoning_level"] = self.config.reasoning_level + if self.merge_sessions and ctx.session is not None: + chat_kwargs["session"] = ctx.session + + response = await user_peer.aio.chat(question, **chat_kwargs) + actual_response = response if isinstance(response, str) else "" + + if self.variant == "synth": + score = score_synth_response(gold_answer, actual_response, answer_type) + else: + score = score_real_response(gold_answer, actual_response) + + result["query_executed"] = QueryResult( + question=question, + expected_answer=str(gold_answer), + actual_response=actual_response, + score=score, + context_length_tokens=context_length, + ) + result["score"] = score + result["passed"] = score >= 0.99 + result["output_lines"] = [ + f"Question: {question}", + f"Expected: {gold_answer}", + f"Score: {score:.3f}", + ] + except Exception as e: + result["error"] = str(e) + result["query_executed"] = QueryResult( + question=question, + expected_answer=str(gold_answer), + actual_response=f"ERROR: {e}", + score=0.0, + context_length_tokens=context_length, + ) + + result["end_time"] = time.time() + result["duration_seconds"] = result["end_time"] - result["start_time"] + return result + + def print_summary(self, results: list[TestResult], total_duration: float) -> None: + total_examples = len(results) + perfect_scores = sum(1 for r in results if r["score"] >= 0.99) + average_score = ( + sum(result["score"] for result in results) / total_examples + if total_examples + else 0.0 + ) + + print(f"\n{'=' * 80}") + print(f"OOLONG-{self.variant.upper()} BENCHMARK SUMMARY") + print(f"{'=' * 80}") + print(f"Total examples: {total_examples}") + print(f"Average score: {average_score:.3f}") + perfect_rate = ( + (perfect_scores / total_examples) * 100 if total_examples else 0.0 + ) + print(f"Perfect scores (>=0.99): {perfect_scores} ({perfect_rate:.1f}%)") + print(f"Total test time: {format_duration(total_duration)}") + + task_stats = calculate_task_statistics(results) + if task_stats: + print("\nTask group statistics:") + for task_name, stats in sorted(task_stats.items()): + print( + f" {task_name}: avg={stats['average_score']:.3f}, perfect={stats['perfect_score_rate']:.1f}% ({stats['total']})" + ) + print(f"{'=' * 80}") + + def generate_output(self, results: list[TestResult], total_duration: float) -> None: + total_examples = len(results) + perfect_scores = sum(1 for r in results if r["score"] >= 0.99) + average_score = ( + sum(result["score"] for result in results) / total_examples + if total_examples + else 0.0 + ) + task_stats = calculate_task_statistics(results) + timing_stats = calculate_timing_statistics(results, total_duration) + + summary: dict[str, Any] = { + "metadata": { + "benchmark": "oolong", + "variant": self.variant, + "split": self.split, + "data_dir": str(self.data_dir), + "execution_timestamp": datetime.now().isoformat(), + "runner_version": "2.0.0", + "base_api_port": self.config.base_api_port, + "pool_size": self.config.pool_size, + "timeout_seconds": self.config.timeout_seconds, + "merge_sessions": self.merge_sessions, + "labels": self.use_labels, + "reasoning_level": self.config.reasoning_level, + "deriver_settings": settings.DERIVER.model_dump(), + "dialectic_settings": settings.DIALECTIC.model_dump(), + "dream_settings": settings.DREAM.model_dump(), + }, + "summary_statistics": { + "total_examples": total_examples, + "perfect_scores": perfect_scores, + "perfect_score_rate": perfect_scores / total_examples + if total_examples + else 0.0, + "average_score": average_score, + "statistics_by_task_group": task_stats, + }, + "timing": timing_stats, + "detailed_results": [ + { + "question_id": result["question_id"], + "context_window_id": result["context_window_id"], + "task_group": result["task_group"], + "dataset": result["dataset"], + "answer_type": result["answer_type"], + "score": result["score"], + "passed": result["passed"], + "error": result["error"], + "duration_seconds": result["duration_seconds"], + "query_executed": result["query_executed"], + } + for result in results + ], + } + + if self.config.json_output: + output_file = self.config.json_output + else: + output_file = Path( + f"tests/bench/eval_results/oolong_{self.variant}_{datetime.now().strftime('%Y%m%d_%H%M%S')}.json" + ) + + write_json_summary(summary, output_file) + + +def main() -> int: + parser = argparse.ArgumentParser( + description="Run OOLONG benchmark tests against a Honcho instance", + formatter_class=argparse.RawDescriptionHelpFormatter, + epilog=""" +Examples: + %(prog)s --variant synth --data-dir /datasets/oolong-synth + %(prog)s --variant real --data-dir /datasets/oolong-real --split validation + %(prog)s --variant synth --data-dir /datasets/oolong-synth --context-size 16K + %(prog)s --variant synth --data-dir /datasets/oolong-synth --no-merge-sessions + """, + ) + + parser.add_argument( + "--variant", + type=str, + default="synth", + choices=["synth", "real"], + help="Which OOLONG variant to run (default: synth)", + ) + parser.add_argument( + "--data-dir", + type=Path, + required=True, + help="Path to the dataset directory", + ) + parser.add_argument( + "--split", + type=str, + default="test", + choices=["test", "validation"], + help="Dataset split to use (default: test)", + ) + parser.add_argument( + "--max-examples", + type=int, + default=None, + help="Maximum number of examples to run (default: all)", + ) + parser.add_argument( + "--context-size", + type=str, + default=None, + help=( + "Context-size bucket cap, e.g. 8K, 16K, 1M, or exact token count like 16384. " + "Sets --max-context-len; --min-context-len is only kept when explicitly passed." + ), + ) + parser.add_argument( + "--max-context-len", + type=int, + default=None, + help="Maximum context length in tokens", + ) + parser.add_argument( + "--labels", + action="store_true", + default=False, + help="Use context_window_text_with_labels for synth examples", + ) + parser.add_argument( + "--min-context-len", + type=int, + default=1024, + help="Minimum context length in tokens (default: 1024, upstream behavior)", + ) + parser.add_argument( + "--context-window-id", + type=str, + default=None, + help="Run only examples with this context_window_id", + ) + parser.add_argument( + "--no-merge-sessions", + action="store_false", + dest="merge_sessions", + default=True, + help="Store context across multiple sessions instead of a merged session", + ) + + add_common_arguments(parser) + args = parser.parse_args() + min_context_len_explicit = "--min-context-len" in sys.argv + + error = validate_common_arguments(args) + if error: + print(error) + return 1 + + if args.use_get_context: + print("Error: --use-get-context is not supported by the OOLONG runner") + return 1 + + if not args.data_dir.exists(): + print(f"Error: data directory does not exist: {args.data_dir}") + return 1 + + if args.max_examples is not None and args.max_examples <= 0: + print(f"Error: max examples must be positive, got {args.max_examples}") + return 1 + + if args.context_size: + try: + exact_size = parse_context_size(args.context_size) + except ValueError as e: + print(f"Error: {e}") + return 1 + + # OOLONG-style behavior: context-size is a bucket cap. Preserve + # an explicit lower bound only when the user provides one. + args.max_context_len = exact_size + if not min_context_len_explicit: + args.min_context_len = None + print( + f"Using context-size cap: <= {exact_size} tokens" + + ( + f" (min: {args.min_context_len})" + if args.min_context_len is not None + else "" + ) + ) + + if args.min_context_len is not None and args.min_context_len < 0: + print(f"Error: min context len must be >= 0, got {args.min_context_len}") + return 1 + + if args.max_context_len is not None and args.max_context_len <= 0: + print(f"Error: max context len must be positive, got {args.max_context_len}") + return 1 + + if ( + args.max_context_len is not None + and args.min_context_len is not None + and args.max_context_len < args.min_context_len + ): + print( + f"Error: max context len must be >= min context len ({args.max_context_len} < {args.min_context_len})" + ) + return 1 + + config = RunnerConfig.from_args(args, default_timeout=600) + runner = OolongRunner( + config=config, + variant=args.variant, + data_dir=args.data_dir, + split=args.split, + merge_sessions=args.merge_sessions, + max_examples=args.max_examples, + min_context_len=args.min_context_len, + max_context_len=args.max_context_len, + context_window_id=args.context_window_id, + use_labels=args.labels, + ) + return runner.run_and_summarize() + + +if __name__ == "__main__": + exit(main()) diff --git a/tests/bench/oolong_common.py b/tests/bench/oolong_common.py new file mode 100644 index 00000000..31ffe2a6 --- /dev/null +++ b/tests/bench/oolong_common.py @@ -0,0 +1,750 @@ +""" +Common utilities for OOLONG benchmark test runners. + +Based on the OOLONG paper: +- OOLONG-synth: Synthetic ICL-based aggregation tasks +- OOLONG-real: Real D&D transcript aggregation tasks +""" + +import ast +import json +import logging +import re +from collections.abc import Callable, Sequence +from datetime import date, datetime +from pathlib import Path +from typing import Any, cast + +import dateutil.parser +import pyarrow.parquet as pq +import tiktoken +from typing_extensions import TypedDict + +logger = logging.getLogger(__name__) + + +class SimpleDataset: + """Simple dataset class that mimics HuggingFace dataset API.""" + + data: list[dict[str, Any]] + column_names: list[str] + + def __init__(self, data: list[dict[str, Any]]): + """Initialize dataset with list of examples. + + Args: + data: List of example dictionaries + """ + self.data = data + self.column_names = list(data[0].keys()) if data else [] + + def __len__(self) -> int: + """Return number of examples.""" + return len(self.data) + + def __getitem__(self, idx: int) -> dict[str, Any]: + """Get example by index.""" + return self.data[idx] + + def filter(self, function: Callable[[dict[str, Any]], bool]) -> "SimpleDataset": + """Filter dataset using a function. + + Args: + function: Filter function + + Returns: + Filtered dataset + """ + filtered_data = [item for item in self.data if function(item)] + return SimpleDataset(filtered_data) + + def select(self, indices: Sequence[int]) -> "SimpleDataset": + """Select examples by indices. + + Args: + indices: List of indices to select + + Returns: + Dataset with selected examples + """ + selected_data = [self.data[i] for i in indices] + return SimpleDataset(selected_data) + + +class BaseQueryResult(TypedDict): + """Base type definition for query execution results.""" + + question: str + expected_answer: str + actual_response: str + score: float + context_length_tokens: int + + +class BaseTestResult(TypedDict): + """Base type definition for test execution results.""" + + question_id: str + task_group: str + dataset: str + passed: bool + score: float + error: str | None + start_time: float + end_time: float + duration_seconds: float + output_lines: list[str] + + +def format_duration(total_seconds: float) -> str: + """Format a duration in seconds into a human-readable string. + + If the duration is at least one minute, this returns a string in the + form "XmYYs" with zero-padded seconds. Otherwise, it returns the + duration in seconds with two decimal places, e.g., "12.34s". + + Args: + total_seconds: The duration in seconds. + + Returns: + A formatted duration string. + """ + minutes = int(total_seconds // 60) + if minutes > 0: + seconds_rounded = int(round(total_seconds - minutes * 60)) + if seconds_rounded == 60: + minutes += 1 + seconds_rounded = 0 + return f"{minutes}m{seconds_rounded:02d}s" + return f"{total_seconds:.2f}s" + + +def calculate_context_length(text: str) -> int: + """Calculate token count for context text using tiktoken. + + Args: + text: Context text to count tokens for + + Returns: + Number of tokens + """ + try: + tokenizer = tiktoken.get_encoding("o200k_base") + return len( + tokenizer.encode( + text, + disallowed_special=(tokenizer.special_tokens_set - {"<|endoftext|>"}), + ) + ) + except Exception: + # Fallback to character-based estimate + return len(text) // 4 + + +def load_oolong_synth_dataset( + split: str = "test", + data_dir: str | Path | None = None, + max_context_len: int | None = None, + min_context_len: int | None = None, + max_examples: int | None = None, + context_window_id: str | None = None, +) -> SimpleDataset: + """Load the OOLONG-synth dataset from filesystem. + + Args: + split: Dataset split to load (default: "test") + data_dir: Path to the oolong-synth dataset directory (must contain a 'data' subdirectory) + max_context_len: Maximum context length in tokens + min_context_len: Minimum context length in tokens (strict >, upstream behavior) + max_examples: Maximum number of examples to return + context_window_id: Specific context window ID to filter to + + Returns: + SimpleDataset object + + Raises: + ValueError: If data_dir is not provided + FileNotFoundError: If no parquet files found for the split + """ + if data_dir is None: + raise ValueError( + "data_dir parameter is required. Please provide the path to the oolong-synth dataset." + ) + + dataset_path = Path(data_dir) / "data" + if not dataset_path.exists(): + raise FileNotFoundError( + f"Expected synth dataset directory at {dataset_path} (layout: data/*.parquet)" + ) + + # Find all parquet files for the given split + parquet_files = sorted(dataset_path.glob(f"{split}-*.parquet")) + + if not parquet_files: + raise FileNotFoundError(f"No {split} parquet files found in {dataset_path}") + + # Stream rows from parquet shards in batches so we can filter and stop early + # without materializing the full split in memory first. + all_data: list[dict[str, Any]] = [] + for parquet_file in parquet_files: + parquet = pq.ParquetFile(parquet_file) + for batch in parquet.iter_batches(): # pyright: ignore[reportUnknownVariableType] + rows = cast( + list[dict[str, Any]], + batch.to_pylist(), # pyright: ignore[reportUnknownMemberType] + ) + for row in rows: + if ( + context_window_id is not None + and str(row.get("context_window_id")) != context_window_id + ): + continue + + context_len = row.get("context_len") + if not isinstance(context_len, int): + context_len = calculate_context_length( + str(row.get("context_window_text", "")) + ) + + if max_context_len is not None and context_len > max_context_len: + continue + + # Keep strict greater-than for consistency with existing behavior. + if min_context_len is not None and context_len <= min_context_len: + continue + + all_data.append(row) + if ( + max_examples is not None + and max_examples > 0 + and len(all_data) >= max_examples + ): + return SimpleDataset(all_data) + + return SimpleDataset(all_data) + + +def load_oolong_real_dataset( + split: str = "test", + data_dir: str | Path | None = None, + max_context_len: int | None = None, + min_context_len: int | None = None, + max_examples: int | None = None, + context_window_id: str | None = None, +) -> SimpleDataset: + """Load the OOLONG-real dataset from filesystem. + + Args: + split: Dataset split to load (default: "test") + data_dir: Path to the oolong-real dataset directory (must contain a 'dnd' subdirectory) + max_context_len: Maximum context length in tokens + min_context_len: Minimum context length in tokens (strict >, upstream behavior) + max_examples: Maximum number of examples to return + context_window_id: Specific context window ID to filter to + + Returns: + SimpleDataset object + + Raises: + ValueError: If data_dir is not provided + FileNotFoundError: If JSONL file not found + """ + if data_dir is None: + raise ValueError( + "data_dir parameter is required. Please provide the path to the oolong-real dataset." + ) + + dataset_path = Path(data_dir) / "dnd" + if not dataset_path.exists(): + raise FileNotFoundError( + f"Expected real dataset directory at {dataset_path} (layout: dnd/*.jsonl)" + ) + jsonl_file = dataset_path / f"{split}.jsonl" + + if not jsonl_file.exists(): + raise FileNotFoundError(f"JSONL file not found: {jsonl_file}") + + # Load JSONL file + data: list[dict[str, Any]] = [] + with open(jsonl_file) as f: + for line in f: + if line.strip(): + row_raw = json.loads(line) + if not isinstance(row_raw, dict): + continue + row = cast(dict[str, Any], row_raw) + + if ( + context_window_id is not None + and str(row.get("context_window_id")) != context_window_id + ): + continue + + context_len = row.get("context_len") + if not isinstance(context_len, int): + context_len = calculate_context_length( + str(row.get("context_window_text", "")) + ) + + if max_context_len is not None and context_len > max_context_len: + continue + + # Keep strict greater-than for consistency with existing behavior. + if min_context_len is not None and context_len <= min_context_len: + continue + + data.append(row) + if ( + max_examples is not None + and max_examples > 0 + and len(data) >= max_examples + ): + return SimpleDataset(data) + + return SimpleDataset(data) + + +def parse_synth_context_messages(context_text: str) -> list[dict[str, Any]]: + """Parse OOLONG-synth context text into individual messages. + + Context format: + Date: YYYY-MM-DD || User: user_XYZ || Instance: [label] + + Args: + context_text: Raw context window text + + Returns: + List of message dictionaries with content and metadata + """ + messages: list[dict[str, Any]] = [] + + # Split by lines and parse each entry + lines = context_text.strip().split("\n") + for line in lines: + if not line.strip(): + continue + + # Parse: Date: ... || User: ... || Instance: ... || Label: ... + try: + parts = line.split(" || ") + if len(parts) < 3: + logger.warning(f"Skipping malformed line: {line}") + continue + + date_part = parts[0].replace("Date: ", "").strip() + user_part = parts[1].replace("User: ", "").strip() + instance_part = parts[2].replace("Instance: ", "").strip() + + # Check if there's a 4th part with label + label = None + if len(parts) >= 4: + label_part = parts[3].replace("Label: ", "").strip() + label = label_part if label_part else None + + # Include label in content so deriver can observe it. + content = f"{instance_part} [Label: {label}]" if label else instance_part + + msg: dict[str, Any] = { + "content": content, + "metadata": { + "date": date_part, + "user_id": user_part, + "label": label, + }, + } + messages.append(msg) + except Exception as e: + logger.warning(f"Error parsing line: {line}. Error: {e}") + continue + + return messages + + +def parse_real_context_messages(context_text: str) -> list[dict[str, Any]]: + """Parse OOLONG-real D&D transcript into individual messages. + + Context format: + Speaker: dialogue text + [multiple lines] + + Args: + context_text: Raw D&D transcript text + + Returns: + List of message dictionaries with content and metadata + """ + messages: list[dict[str, Any]] = [] + + # Split by speaker turns (format: "SPEAKER: text") + lines = context_text.strip().split("\n") + current_speaker = None + current_content = [] + + for line in lines: + if not line.strip(): + continue + + # Check if this is a new speaker turn (must start with speaker label) + # Speaker labels are typically single words or use underscores/hyphens (no spaces) + speaker_match = re.match(r"^\s*([A-Za-z0-9_-]+):", line) + if speaker_match: + # Save previous message if exists + if current_speaker and current_content: + # Include speaker in content for deriver visibility + content_text = " ".join(current_content) + content = f"[Speaker: {current_speaker}] {content_text}" + prev_msg: dict[str, Any] = { + "content": content, + "metadata": { + "speaker": current_speaker, + }, + } + messages.append(prev_msg) + + # Parse new speaker from regex match + current_speaker = speaker_match.group(1).strip() + # Extract content after the colon + content_after_colon = line[speaker_match.end() :].strip() + current_content = [content_after_colon] if content_after_colon else [] + else: + # Continuation of current speaker's dialogue + current_content.append(line.strip()) + + # Save last message + if current_speaker and current_content: + content_text = " ".join(current_content) + content = f"[Speaker: {current_speaker}] {content_text}" + last_msg: dict[str, Any] = { + "content": content, + "metadata": { + "speaker": current_speaker, + }, + } + messages.append(last_msg) + + return messages + + +def parse_synth_answer(answer_str: str) -> Any: + """Parse OOLONG-synth answer string. + + Answers can be: + - Strings (labels, comparisons, dates) + - Numbers (counts) + - Dates (datetime objects) + + Args: + answer_str: Raw answer string from dataset + + Returns: + Parsed answer value + """ + # Handle datetime answers + if "datetime" in answer_str: + try: + # Format: [datetime.date(2023, 5, 15)] + match = re.search(r"datetime\.date\((\d+),\s*(\d+),\s*(\d+)\)", answer_str) + if match: + year, month, day = map(int, match.groups()) + return datetime(year, month, day).date() + except Exception as e: + logger.warning(f"Error parsing datetime answer: {answer_str}. Error: {e}") + return answer_str + + # Try literal eval for lists/primitives + try: + parsed = ast.literal_eval(answer_str) + # If it's a list with one element, return that element + if isinstance(parsed, list) and len(parsed) == 1: # pyright: ignore[reportUnknownArgumentType] + return parsed[0] # pyright: ignore[reportUnknownVariableType] + return parsed # pyright: ignore[reportUnknownVariableType] + except (ValueError, SyntaxError): + # Return as-is if can't parse + return answer_str + + +def parse_real_answer(answer_str: str) -> int | str | list[str]: + """Parse OOLONG-real answer string. + + Answers can be: + - Integers (counts) + - Strings (spell names, roll types) + - Lists (comma-separated spells) + + Args: + answer_str: Raw answer string from dataset + + Returns: + Parsed answer value + """ + # Try to convert to int first + try: + return int(answer_str) + except ValueError: + pass + + # Check if it contains commas (list case) + if "," in answer_str: + return [item.strip() for item in answer_str.split(",") if item.strip()] + + # Otherwise return as string + return answer_str + + +def score_synth_response( + gold_answer: Any, model_answer_str: str, answer_type: str +) -> float: + """Score a response for OOLONG-synth following the paper's scoring rubric. + + Scoring: + - Exact match: 1.0 + - Numeric: 0.75^|y-ŷ| (partial credit) + - Other: 0.0 + + Args: + gold_answer: Expected answer + model_answer_str: Model's response string + answer_type: Type of answer (ANSWER_TYPE.NUMERIC, ANSWER_TYPE.DATE, etc.) + + Returns: + Score between 0.0 and 1.0 + """ + # Try to extract answer from response + # Look for common patterns at start of line or end of response + model_answer = model_answer_str.strip() + + # Try to extract from common formats (use MULTILINE to match at line start, and $ to find at line end) + # Look for patterns at the start of a line (final answer format) + for pattern in [ + r"^[Aa]nswer:\s*(.+)$", # "Answer: X" on its own line + r"^[Ll]abel:\s*(.+)$", # "Label: X" on its own line + r"\n[Ll]abel:\s*(.+)$", # "Label: X" after a newline (final answer) + r"[Aa]nswer:\s*(.+)$", # "Answer: X" at end of response + r"[Uu]ser:\s*(.+)$", # "User: X" at end + r"[Dd]ate:\s*(.+)$", # "Date: X" at end + ]: + match = re.search(pattern, model_answer, re.MULTILINE) + if match: + model_answer = match.group(1).strip() + break + + # Remove formatting like ** or [] + model_answer = re.sub(r"[\*\[\]]", "", model_answer) + + # If still long, try last significant token + if len(model_answer) > 50: + model_answer = model_answer.split()[-1] + + # Exact string match + if str(model_answer).lower() == str(gold_answer).lower(): + return 1.0 + + # Check for comparison answers + if ( + "more common" in model_answer.lower() + and "more common" in str(gold_answer).lower() + ): + return 1.0 + if ( + "less common" in model_answer.lower() + and "less common" in str(gold_answer).lower() + ): + return 1.0 + if ( + "same frequency" in model_answer.lower() + and "same frequency" in str(gold_answer).lower() + ): + return 1.0 + + # Numeric partial credit + if "NUMERIC" in answer_type.upper(): + try: + model_num = float(re.sub(r"[^\d.-]", "", model_answer)) + gold_num = float(gold_answer) + return 0.75 ** abs(gold_num - model_num) + except (ValueError, TypeError): + return 0.0 + + # Date matching + if "DATE" in answer_type.upper(): + try: + model_date = dateutil.parser.parse(model_answer) + if isinstance(gold_answer, datetime): + return 1.0 if model_date.date() == gold_answer.date() else 0.0 + elif isinstance(gold_answer, date): + return 1.0 if model_date.date() == gold_answer else 0.0 + return 0.0 + except (ValueError, TypeError): + return 0.0 + + return 0.0 + + +def score_real_response( + gold_answer: int | str | list[str], model_answer_str: str +) -> float: + """Score a response for OOLONG-real following the paper's scoring rubric. + + Scoring: + - Integer: 0.75^|y-ŷ| (partial credit) + - String: exact match (1.0 or 0.0) + - List: set overlap / |gold| (Jaccard-style) + + Args: + gold_answer: Expected answer + model_answer_str: Model's response string + + Returns: + Score between 0.0 and 1.0 + """ + # Extract answer from \boxed{} format if present + match = re.search(r"\\boxed\{\\text\{([^}]*)\}\}", model_answer_str) or re.search( + r"\\boxed[\{]+([^}]*)[\}]+", model_answer_str + ) + + if match: + model_answer_str = match.group(1) + + # Parse model answer + try: + model_answer: int | str | list[str] = int(model_answer_str) + except ValueError: + if "," in model_answer_str: + model_answer = [ + item.strip() for item in model_answer_str.split(",") if item.strip() + ] + else: + model_answer = model_answer_str.strip() + + # Score based on type + if isinstance(gold_answer, int) and isinstance(model_answer, int): + return 0.75 ** abs(gold_answer - model_answer) + elif isinstance(gold_answer, str) and isinstance(model_answer, str): + return 1.0 if gold_answer.lower() == model_answer.lower() else 0.0 + elif isinstance(gold_answer, list) and isinstance(model_answer, list): + overlap = set(gold_answer) & set(model_answer) + return len(overlap) / len(gold_answer) if gold_answer else 0.0 + else: + return 0.0 + + +def filter_dataset( + dataset: SimpleDataset, + max_context_len: int | None = None, + min_context_len: int | None = None, + max_examples: int | None = None, + context_window_id: str | None = None, +) -> SimpleDataset: + """Filter dataset by context length and example count. + + Args: + dataset: HuggingFace dataset + max_context_len: Maximum context length in tokens + min_context_len: Minimum context length in tokens + max_examples: Maximum number of examples to return + context_window_id: Specific context window ID to filter to + + Returns: + Filtered dataset + """ + if context_window_id is not None: + dataset = dataset.filter(lambda x: x["context_window_id"] == context_window_id) + + if max_context_len is not None: + dataset = dataset.filter( + lambda x: x.get( + "context_len", + calculate_context_length(str(x.get("context_window_text", ""))), + ) + <= max_context_len + ) + + if min_context_len is not None: + dataset = dataset.filter( + lambda x: x.get( + "context_len", + calculate_context_length(str(x.get("context_window_text", ""))), + ) + > min_context_len + ) + + if max_examples is not None and max_examples > 0: + # Get a slice of the dataset + dataset = dataset.select(range(min(max_examples, len(dataset)))) + + return dataset + + +def calculate_task_statistics( + results: Sequence[Any], +) -> dict[str, dict[str, int | float]]: + """Calculate score statistics grouped by task type. + + Args: + results: List of test results + + Returns: + Dictionary mapping task type to statistics + """ + task_stats: dict[str, dict[str, int | float]] = {} + + for result in results: + task_key = result.get("task_group", "unknown") + if task_key not in task_stats: + task_stats[task_key] = {"total": 0, "total_score": 0.0, "perfect_scores": 0} + + task_stats[task_key]["total"] += 1 + score = result.get("score", 0.0) + task_stats[task_key]["total_score"] += score + if score >= 0.99: # Consider >= 0.99 as perfect to account for floating point + task_stats[task_key]["perfect_scores"] += 1 + + # Calculate averages + for task_key in task_stats: + stats = task_stats[task_key] + total = stats["total"] + stats["average_score"] = stats["total_score"] / total if total > 0 else 0.0 + stats["perfect_score_rate"] = ( + (stats["perfect_scores"] / total) * 100 if total > 0 else 0.0 + ) + + return task_stats + + +def calculate_timing_statistics( + results: Sequence[Any], total_elapsed_seconds: float +) -> dict[str, Any]: + """Calculate timing statistics from test results. + + Args: + results: List of test results + total_elapsed_seconds: Total elapsed time for all tests + + Returns: + Dictionary of timing statistics + """ + durations = [r["duration_seconds"] for r in results] + return { + "total_duration_seconds": total_elapsed_seconds, + "individual_test_durations": { + "min_seconds": min(durations) if durations else 0, + "max_seconds": max(durations) if durations else 0, + "mean_seconds": sum(durations) / len(durations) if durations else 0, + "median_seconds": sorted(durations)[len(durations) // 2] + if durations + else 0, + }, + } + + +def write_json_summary(summary: dict[str, Any], output_file: Path) -> None: + """Write a JSON summary to a file. + + Args: + summary: Summary dictionary to write + output_file: Path to output file + """ + output_file.parent.mkdir(parents=True, exist_ok=True) + with open(output_file, "w") as f: + json.dump(summary, f, indent=2, default=str) + print(f"\nJSON summary written to: {output_file}") diff --git a/tests/bench/runner_common.py b/tests/bench/runner_common.py index 31855479..bd1826a2 100644 --- a/tests/bench/runner_common.py +++ b/tests/bench/runner_common.py @@ -373,6 +373,15 @@ class BaseRunner(ABC, Generic[ResultT]): """Return list of peer IDs to trigger dreams for.""" ... + def get_dream_session_ids(self, ctx: ItemContext, _item: Any) -> list[str]: + """Return session IDs to use for dream scheduling. + + Subclasses can override this when ingestion stores messages across + multiple sessions and each session should be included in dream + scheduling. + """ + return [ctx.session_id] + @abstractmethod async def execute_questions(self, ctx: ItemContext, item: Any) -> ResultT: """ @@ -518,15 +527,29 @@ class BaseRunner(ABC, Generic[ResultT]): ) # Trigger dreams - print(f"[{workspace_id}] Deriver queue empty. Triggering dreams...") - for observer in self.get_dream_observers(item): - success = await self._trigger_dream( - ctx.honcho_client, workspace_id, observer, session_id + dream_observers = self.get_dream_observers(item) + dream_session_ids = self.get_dream_session_ids(ctx, item) + if not dream_session_ids: + raise ValueError( + f"No dream session IDs available for {workspace_id}. " + + "Dream scheduling requires at least one session id." ) - if not success: - print( - f"[{workspace_id}] Warning: Dream for {observer} did not complete" + + print( + f"[{workspace_id}] Deriver queue empty. Triggering dreams for " + + f"{len(dream_observers)} observer(s) across " + + f"{len(dream_session_ids)} session(s)..." + ) + for observer in dream_observers: + for dream_session_id in dream_session_ids: + success = await self._trigger_dream( + ctx.honcho_client, workspace_id, observer, dream_session_id ) + if not success: + print( + f"[{workspace_id}] Warning: Dream for {observer} in " + + f"session {dream_session_id} did not complete" + ) # Execute questions print(f"[{workspace_id}] Executing questions...") diff --git a/uv.lock b/uv.lock index bde5710d..89bcbc29 100644 --- a/uv.lock +++ b/uv.lock @@ -1303,6 +1303,7 @@ dependencies = [ { name = "pgvector" }, { name = "prometheus-client" }, { name = "psycopg", extra = ["binary"] }, + { name = "pyarrow" }, { name = "pydantic" }, { name = "pydantic-settings" }, { name = "pyjwt" }, @@ -1359,6 +1360,7 @@ requires-dist = [ { name = "pgvector", specifier = ">=0.2.5" }, { name = "prometheus-client", specifier = ">=0.21.0" }, { name = "psycopg", extras = ["binary"], specifier = ">=3.1.19" }, + { name = "pyarrow", specifier = ">=19.0.0" }, { name = "pydantic", specifier = ">=2.11.7" }, { name = "pydantic-settings", specifier = ">=2.10.1" }, { name = "pyjwt", specifier = ">=2.10.0" },