Merge pull request #228 from ggozad/feat/multimodal-evaluation

Add evaluation dataset for multi-modal q&a
This commit is contained in:
Yiorgis Gozadinos 2026-01-22 17:47:17 +02:00 committed by GitHub
commit 3bddf87454
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 307 additions and 18 deletions

View file

@ -123,3 +123,17 @@ QA accuracy is evaluated over 2000 "hard" questions from the validation dataset.
| Embedding Model | QA Model | Accuracy |
|----------------------|--------------------------|----------|
| `qwen3-embedding:4b` | `gpt-oss:20b` - thinking | 0.86 |
## OpenRAG Bench (ORB)
[OpenRAG Bench](https://huggingface.co/datasets/vectara/open_ragbench) contains ArXiv research papers with multimodal question-answering pairs. Queries include both text-based and image-based questions, testing retrieval over visual content like figures, charts, and diagrams. We use MAP for retrieval evaluation since each query maps to one relevant document.
**Multimodal processing**: Picture descriptions are generated using a Vision Language Model (VLM) during document conversion, making embedded images searchable via text queries. See [Picture Description configuration](configuration/processing.md#picture-description-vlm).
*Results from v0.26.8*
### Retrieval (MAP)
| Embedding Model | MAP | VLM |
|----------------------|--------|----------------------|
| `qwen3-embedding:4b` | 0.9626 | Ollama / ministral-3 |

View file

@ -83,7 +83,7 @@ async def populate_db(
with Progress() as progress:
task = progress.add_task("[green]Populating database...", total=len(corpus))
async with HaikuRAG(db, config=config) as rag:
async with HaikuRAG(db, config=config, create=True) as rag:
docs_since_vacuum = 0
for doc in corpus:
doc_mapping = cast(Mapping[str, Any], doc)
@ -92,7 +92,13 @@ async def populate_db(
progress.advance(task)
continue
existing = await rag.get_document_by_uri(payload.uri)
# Use the actual URI that will be stored in the database
if payload.source_path is not None:
lookup_uri = payload.source_path.absolute().as_uri()
else:
lookup_uri = payload.uri
existing = await rag.get_document_by_uri(lookup_uri)
if existing is not None:
assert existing.id
chunks = await rag.chunk_repository.get_by_document_id(existing.id)
@ -101,13 +107,21 @@ async def populate_db(
continue
await rag.document_repository.delete(existing.id)
await rag.create_document(
content=payload.content,
uri=payload.uri,
title=payload.title,
metadata=payload.metadata,
format=payload.format,
)
if payload.source_path is not None:
await rag.create_document_from_source(
source=payload.source_path,
title=payload.title,
metadata=payload.metadata,
)
else:
assert payload.content is not None
await rag.create_document(
content=payload.content,
uri=payload.uri,
title=payload.title,
metadata=payload.metadata,
format=payload.format,
)
docs_since_vacuum += 1
progress.advance(task)
@ -126,6 +140,7 @@ async def run_retrieval_benchmark(
limit: int | None = None,
name: str | None = None,
db_path: Path | None = None,
multimodal_only: bool = False,
) -> dict[str, float] | None:
if spec.retrieval_loader is None or spec.retrieval_mapper is None:
console.print("Skipping retrieval benchmark; no retrieval config.")
@ -145,9 +160,18 @@ async def run_retrieval_benchmark(
progress.advance(task)
continue
# Filter for multimodal queries if requested
if multimodal_only:
if sample.source_type is None or "image" not in sample.source_type:
progress.advance(task)
continue
case = Case(
inputs=sample.question,
metadata={"relevant_uris": sample.expected_uris},
metadata={
"relevant_uris": sample.expected_uris,
"source_type": sample.source_type,
},
)
cases.append(case)
progress.advance(task)
@ -174,16 +198,22 @@ async def run_retrieval_benchmark(
chunks = await rag.search(query=question, limit=5)
seen = set()
uris = []
identifiers = []
for result in chunks:
if result.document_id is None:
continue
doc = await rag.get_document_by_id(result.document_id)
if doc and doc.uri and doc.uri not in seen:
uris.append(doc.uri)
seen.add(doc.uri)
if doc is None:
continue
# Use arxiv_id from metadata if present, otherwise use URI
doc_id = doc.metadata.get("arxiv_id") if doc.metadata else None
if doc_id is None:
doc_id = doc.uri
if doc_id and doc_id not in seen:
identifiers.append(doc_id)
seen.add(doc_id)
return uris
return identifiers
eval_name = name if name is not None else f"{spec.key}_retrieval_evaluation"
@ -326,6 +356,7 @@ async def evaluate_dataset(
name: str | None,
db_path: Path | None,
vacuum_interval: int = 100,
multimodal_only: bool = False,
) -> None:
if not skip_db:
console.print(f"Using dataset: {spec.key}", style="bold magenta")
@ -336,7 +367,12 @@ async def evaluate_dataset(
if not skip_retrieval:
console.print("Running retrieval benchmarks...", style="bold blue")
await run_retrieval_benchmark(
spec, config, limit=limit, name=name, db_path=db_path
spec,
config,
limit=limit,
name=name,
db_path=db_path,
multimodal_only=multimodal_only,
)
if not skip_qa:
@ -368,6 +404,11 @@ def run(
vacuum_interval: int = typer.Option(
100, "--vacuum-interval", help="Vacuum every N documents during DB population."
),
multimodal_only: bool = typer.Option(
False,
"--multimodal-only",
help="Only evaluate queries requiring image understanding.",
),
) -> None:
spec = DATASETS.get(dataset.lower())
if spec is None:
@ -405,6 +446,7 @@ def run(
name=name,
db_path=db,
vacuum_interval=vacuum_interval,
multimodal_only=multimodal_only,
)
)

View file

@ -11,10 +11,11 @@ from pydantic_evals.evaluators import Evaluator
@dataclass
class DocumentPayload:
uri: str
content: str
content: str | None = None
title: str | None = None
metadata: dict[str, Any] | None = None
format: str = "md"
source_path: Path | None = None
@dataclass
@ -22,6 +23,7 @@ class RetrievalSample:
question: str
expected_uris: tuple[str, ...]
skip: bool = False
source_type: str | None = None
DocumentLoader = Callable[[], Dataset]

View file

@ -1,11 +1,13 @@
from evaluations.config import DatasetSpec
from .hotpotqa import HOTPOTQA_SPEC
from .open_rag_bench import OPEN_RAG_BENCH_SPEC
from .repliqa import REPLIQ_SPEC
from .wix import WIX_SPEC
DATASETS: dict[str, DatasetSpec] = {
spec.key: spec for spec in (REPLIQ_SPEC, WIX_SPEC, HOTPOTQA_SPEC)
spec.key: spec
for spec in (REPLIQ_SPEC, WIX_SPEC, HOTPOTQA_SPEC, OPEN_RAG_BENCH_SPEC)
}
__all__ = ["DATASETS"]

View file

@ -0,0 +1,229 @@
import json
import logging
from collections.abc import Mapping
from pathlib import Path
from typing import Any
import httpx
from datasets import Dataset
from huggingface_hub import hf_hub_download
from pydantic_evals import Case
from evaluations.config import DatasetSpec, DocumentPayload, RetrievalSample
from evaluations.evaluators import MAPEvaluator
logger = logging.getLogger(__name__)
REPO_ID = "vectara/open_ragbench"
PDF_SUBDIR = "pdf/arxiv"
def get_cache_dir() -> Path:
cache_dir = Path.home() / ".cache" / "haiku.rag" / "evaluations" / "arxiv_pdfs"
cache_dir.mkdir(parents=True, exist_ok=True)
return cache_dir
def download_metadata_file(filename: str) -> Path:
return Path(
hf_hub_download(
repo_id=REPO_ID,
filename=f"{PDF_SUBDIR}/{filename}",
repo_type="dataset",
)
)
def load_pdf_urls() -> dict[str, str]:
path = download_metadata_file("pdf_urls.json")
with open(path) as f:
return json.load(f)
def load_queries() -> dict[str, dict[str, str]]:
path = download_metadata_file("queries.json")
with open(path) as f:
return json.load(f)
def load_qrels() -> dict[str, dict[str, Any]]:
path = download_metadata_file("qrels.json")
with open(path) as f:
return json.load(f)
def load_answers() -> dict[str, str]:
path = download_metadata_file("answers.json")
with open(path) as f:
return json.load(f)
def download_pdf(paper_id: str, url: str, cache_dir: Path) -> Path | None:
pdf_path = cache_dir / f"{paper_id}.pdf"
if pdf_path.exists():
return pdf_path
try:
with httpx.Client(timeout=60.0, follow_redirects=True) as client:
response = client.get(url)
response.raise_for_status()
pdf_path.write_bytes(response.content)
return pdf_path
except Exception as e:
logger.warning(f"Failed to download PDF {paper_id}: {e}")
return None
def download_all_pdfs(pdf_urls: dict[str, str]) -> dict[str, Path]:
cache_dir = get_cache_dir()
downloaded = {}
for paper_id, url in pdf_urls.items():
pdf_path = download_pdf(paper_id, url, cache_dir)
if pdf_path is not None:
downloaded[paper_id] = pdf_path
logger.info(f"Downloaded {len(downloaded)}/{len(pdf_urls)} PDFs")
return downloaded
_pdf_urls: dict[str, str] | None = None
_queries: dict[str, dict[str, str]] | None = None
_qrels: dict[str, dict[str, Any]] | None = None
_answers: dict[str, str] | None = None
def ensure_metadata_loaded() -> None:
global _pdf_urls, _queries, _qrels, _answers
if _pdf_urls is None:
_pdf_urls = load_pdf_urls()
if _queries is None:
_queries = load_queries()
if _qrels is None:
_qrels = load_qrels()
if _answers is None:
_answers = load_answers()
def load_orb_corpus() -> Dataset:
ensure_metadata_loaded()
assert _pdf_urls is not None
# Return paper IDs and URLs - PDFs are downloaded lazily during mapping
records = [
{"paper_id": paper_id, "pdf_url": url} for paper_id, url in _pdf_urls.items()
]
return Dataset.from_list(records)
def map_orb_document(doc: Mapping[str, Any]) -> DocumentPayload | None:
paper_id = doc["paper_id"]
pdf_url = doc["pdf_url"]
# Download PDF lazily
cache_dir = get_cache_dir()
pdf_path = download_pdf(paper_id, pdf_url, cache_dir)
if pdf_path is None:
return None
return DocumentPayload(
uri=paper_id,
source_path=pdf_path,
title=paper_id,
metadata={"arxiv_id": paper_id},
)
def load_orb_qa() -> Dataset:
ensure_metadata_loaded()
assert _queries is not None
assert _answers is not None
records = []
for query_id, query_data in _queries.items():
answer = _answers.get(query_id, "")
records.append(
{
"query_id": query_id,
"query": query_data["query"],
"type": query_data["type"],
"source": query_data["source"],
"answer": answer,
}
)
return Dataset.from_list(records)
def build_orb_case(
index: int, doc: Mapping[str, Any]
) -> Case[str, str, dict[str, str]]:
metadata = {
"case_index": str(index),
"query_id": doc["query_id"],
"query_type": doc["type"],
"query_source": doc["source"],
}
return Case(
name=f"{index}_{doc['query_id'][:8]}",
inputs=doc["query"],
expected_output=doc["answer"],
metadata=metadata,
)
def load_orb_retrieval() -> Dataset:
ensure_metadata_loaded()
assert _pdf_urls is not None
assert _queries is not None
assert _qrels is not None
records = []
for query_id, query_data in _queries.items():
qrel = _qrels.get(query_id)
if qrel is None:
continue
doc_id = qrel.get("doc_id")
if doc_id is None or doc_id not in _pdf_urls:
continue
records.append(
{
"query_id": query_id,
"query": query_data["query"],
"type": query_data["type"],
"source": query_data["source"],
"doc_id": doc_id,
}
)
return Dataset.from_list(records)
def map_orb_retrieval(doc: Mapping[str, Any]) -> RetrievalSample | None:
return RetrievalSample(
question=doc["query"],
expected_uris=(doc["doc_id"],),
source_type=doc.get("source"),
)
def is_multimodal_query(source: str) -> bool:
return "image" in source
OPEN_RAG_BENCH_SPEC = DatasetSpec(
key="orb",
db_filename="open_rag_bench.lancedb",
document_loader=load_orb_corpus,
document_mapper=map_orb_document,
qa_loader=load_orb_qa,
qa_case_builder=build_orb_case,
retrieval_loader=load_orb_retrieval,
retrieval_mapper=map_orb_retrieval,
retrieval_evaluator=MAPEvaluator(),
)