Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 16 additions & 16 deletions benchmarks/longmemeval_bench.py
Original file line number Diff line number Diff line change
Expand Up @@ -2917,7 +2917,7 @@ def _load_or_create_split(split_file: str, data: list, dev_size: int = 50, seed:

split_path = Path(split_file)
if split_path.exists():
with open(split_path) as f:
with open(split_path, encoding="utf-8") as f:
return json.load(f)

# Create new split
Expand All @@ -2927,9 +2927,9 @@ def _load_or_create_split(split_file: str, data: list, dev_size: int = 50, seed:
dev_ids = all_ids[:dev_size]
held_out_ids = all_ids[dev_size:]
split = {"dev": dev_ids, "held_out": held_out_ids, "seed": seed, "dev_size": dev_size}
with open(split_path, "w") as f:
with open(split_path, "w", encoding="utf-8") as f:
json.dump(split, f, indent=2)
print(f" Created new split: {len(dev_ids)} dev / {len(held_out_ids)} held-out → {split_path}")
print(f" Created new split: {len(dev_ids)} dev / {len(held_out_ids)} held-out -> {split_path}")
return split


Expand Down Expand Up @@ -2957,7 +2957,7 @@ def run_benchmark(
split_subset: "dev" (50 questions for tuning) or "held_out" (450 for final evaluation).
None = run all questions.
"""
with open(data_file) as f:
with open(data_file, encoding="utf-8") as f:
data = json.load(f)

# Apply train/test split filter before limit/skip
Expand All @@ -2966,7 +2966,7 @@ def run_benchmark(
subset_ids = set(split[split_subset])
before = len(data)
data = [entry for entry in data if entry["question_id"] in subset_ids]
print(f" Split filter ({split_subset}): {before} → {len(data)} questions")
print(f" Split filter ({split_subset}): {before} -> {len(data)} questions")

if limit > 0:
data = data[:limit]
Expand Down Expand Up @@ -2998,7 +2998,7 @@ def run_benchmark(
cache_path = Path(diary_cache_file)
if cache_path.exists():
try:
with open(cache_path) as f:
with open(cache_path, encoding="utf-8") as f:
diary_cache = json.load(f)
print(
f" Diary cache: loaded {len(diary_cache)} sessions from {cache_path.name}"
Expand Down Expand Up @@ -3031,22 +3031,22 @@ def run_benchmark(
# Save progress in case of interruption
if cache_path:
try:
with open(cache_path, "w") as f:
with open(cache_path, "w", encoding="utf-8") as f:
json.dump(diary_cache, f)
except Exception:
pass
print(f" Diary ingest complete: {done} sessions processed")
# Final cache save
if cache_path:
try:
with open(cache_path, "w") as f:
with open(cache_path, "w", encoding="utf-8") as f:
json.dump(diary_cache, f)
print(f" Diary cache saved → {cache_path.name}")
print(f" Diary cache saved -> {cache_path.name}")
except Exception:
pass

print(f"\n{'=' * 60}")
print(" MemPal × LongMemEval Benchmark")
print(" MemPal x LongMemEval Benchmark")
print(f"{'=' * 60}")
print(f" Data: {Path(data_file).name}")
print(f" Questions: {len(data)}")
Expand All @@ -3055,7 +3055,7 @@ def run_benchmark(
rerank_label = f" + LLM re-rank ({model_short})" if llm_rerank_enabled else ""
diary_label = f" [diary ingest: {model_short}]" if mode == "diary" else ""
print(f" Mode: {mode}{diary_label}{rerank_label}")
print(f"{'─' * 60}\n")
print(f"{'-' * 60}\n")

# Collect metrics
ks = [1, 3, 5, 10, 30, 50]
Expand Down Expand Up @@ -3216,7 +3216,7 @@ def run_benchmark(

# Print results
print(f"\n{'=' * 60}")
print(f" RESULTS — MemPal ({mode} mode, {granularity} granularity)")
print(f" RESULTS - MemPal ({mode} mode, {granularity} granularity)")
print(f"{'=' * 60}")
print(f" Time: {elapsed:.1f}s ({elapsed / len(data):.2f}s per question)\n")

Expand Down Expand Up @@ -3245,15 +3245,15 @@ def run_benchmark(
if mode == "diary" and diary_cache and diary_cache_file:
try:
real_cache = {k: v for k, v in diary_cache.items() if v is not None}
with open(diary_cache_file, "w") as f:
with open(diary_cache_file, "w", encoding="utf-8") as f:
json.dump(real_cache, f)
print(f" Diary cache saved: {len(real_cache)} sessions → {diary_cache_file}")
print(f" Diary cache saved: {len(real_cache)} sessions -> {diary_cache_file}")
except Exception as e:
print(f" Warning: could not save diary cache: {e}")

# Save results
if out_file:
with open(out_file, "w") as f:
with open(out_file, "w", encoding="utf-8") as f:
for entry in results_log:
f.write(json.dumps(entry) + "\n")
print(f" Results saved to: {out_file}")
Expand Down Expand Up @@ -3393,7 +3393,7 @@ def run_benchmark(
if args.create_split:
if not args.split_file:
args.split_file = "benchmarks/lme_split_50_450.json"
with open(args.data_file) as f:
with open(args.data_file, encoding="utf-8") as f:
_all_data = json.load(f)
_load_or_create_split(args.split_file, _all_data)
sys.exit(0)
Expand Down
89 changes: 89 additions & 0 deletions tests/benchmarks/test_longmemeval_bench.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,89 @@
import importlib.util
import json
import sys
import types
from pathlib import Path


def _load_longmemeval_bench(monkeypatch):
module_name = "tests._longmemeval_bench_under_test"
module_path = Path(__file__).resolve().parents[2] / "benchmarks" / "longmemeval_bench.py"

class _DummyEphemeralClient:
def delete_collection(self, name):
return None

def create_collection(self, name, embedding_function=None):
return object()

fake_chromadb = types.ModuleType("chromadb")
fake_chromadb.EphemeralClient = lambda: _DummyEphemeralClient()
monkeypatch.setitem(sys.modules, "chromadb", fake_chromadb)

spec = importlib.util.spec_from_file_location(module_name, module_path)
module = importlib.util.module_from_spec(spec)
assert spec.loader is not None
spec.loader.exec_module(module)
return module


def _force_cp1252_default(monkeypatch, target_path):
original_open = open
target_path = Path(target_path)

def fake_open(file, mode="r", *args, **kwargs):
if Path(file) == target_path and "b" not in mode and "encoding" not in kwargs:
kwargs["encoding"] = "cp1252"
return original_open(file, mode, *args, **kwargs)

monkeypatch.setattr("builtins.open", fake_open)


def test_load_or_create_split_reads_utf8_when_windows_default_is_cp1252(tmp_path, monkeypatch):
bench = _load_longmemeval_bench(monkeypatch)
split_path = tmp_path / "lme_split.json"
expected = {"dev": ["frage-ä"], "held_out": ["sess-😀"], "seed": 42, "dev_size": 1}
split_path.write_text(json.dumps(expected, ensure_ascii=False), encoding="utf-8")

_force_cp1252_default(monkeypatch, split_path)

split = bench._load_or_create_split(str(split_path), data=[])

assert split == expected


def test_run_benchmark_reads_utf8_questions_when_windows_default_is_cp1252(
tmp_path, monkeypatch
):
bench = _load_longmemeval_bench(monkeypatch)
question = "Wann hat Zoë das Café 😀 erwähnt?"
answer = "Beim Frühstück im Café 😀."
data = [
{
"question_id": "q1",
"question_type": "single-session-user",
"question": question,
"answer": answer,
"answer_session_ids": ["sess_1"],
"haystack_sessions": [[{"role": "user", "content": "stub memory"}]],
"haystack_session_ids": ["sess_1"],
"haystack_dates": ["2024-01-01"],
}
]
data_path = tmp_path / "longmemeval.json"
data_path.write_text(json.dumps(data, ensure_ascii=False), encoding="utf-8")

captured = {}

def fake_build(entry, granularity="session", n_results=50):
captured["question"] = entry["question"]
captured["answer"] = entry["answer"]
return [0], ["stub memory"], ["sess_1"], ["2024-01-01"]

monkeypatch.setattr(bench, "build_palace_and_retrieve", fake_build)
_force_cp1252_default(monkeypatch, data_path)

bench.run_benchmark(str(data_path), limit=1, out_file=None)

assert captured["question"] == question
assert captured["answer"] == answer