Add evaluation dataset for multi-modal q&a
This commit is contained in:
parent
32c64d3d11
commit
e6b239bd6c
4 changed files with 286 additions and 17 deletions
|
|
@ -83,7 +83,7 @@ async def populate_db(
|
||||||
|
|
||||||
with Progress() as progress:
|
with Progress() as progress:
|
||||||
task = progress.add_task("[green]Populating database...", total=len(corpus))
|
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
|
docs_since_vacuum = 0
|
||||||
for doc in corpus:
|
for doc in corpus:
|
||||||
doc_mapping = cast(Mapping[str, Any], doc)
|
doc_mapping = cast(Mapping[str, Any], doc)
|
||||||
|
|
@ -101,13 +101,21 @@ async def populate_db(
|
||||||
continue
|
continue
|
||||||
await rag.document_repository.delete(existing.id)
|
await rag.document_repository.delete(existing.id)
|
||||||
|
|
||||||
await rag.create_document(
|
if payload.source_path is not None:
|
||||||
content=payload.content,
|
await rag.create_document_from_source(
|
||||||
uri=payload.uri,
|
source=payload.source_path,
|
||||||
title=payload.title,
|
title=payload.title,
|
||||||
metadata=payload.metadata,
|
metadata=payload.metadata,
|
||||||
format=payload.format,
|
)
|
||||||
)
|
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
|
docs_since_vacuum += 1
|
||||||
progress.advance(task)
|
progress.advance(task)
|
||||||
|
|
||||||
|
|
@ -126,6 +134,7 @@ async def run_retrieval_benchmark(
|
||||||
limit: int | None = None,
|
limit: int | None = None,
|
||||||
name: str | None = None,
|
name: str | None = None,
|
||||||
db_path: Path | None = None,
|
db_path: Path | None = None,
|
||||||
|
multimodal_only: bool = False,
|
||||||
) -> dict[str, float] | None:
|
) -> dict[str, float] | None:
|
||||||
if spec.retrieval_loader is None or spec.retrieval_mapper is None:
|
if spec.retrieval_loader is None or spec.retrieval_mapper is None:
|
||||||
console.print("Skipping retrieval benchmark; no retrieval config.")
|
console.print("Skipping retrieval benchmark; no retrieval config.")
|
||||||
|
|
@ -145,9 +154,18 @@ async def run_retrieval_benchmark(
|
||||||
progress.advance(task)
|
progress.advance(task)
|
||||||
continue
|
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(
|
case = Case(
|
||||||
inputs=sample.question,
|
inputs=sample.question,
|
||||||
metadata={"relevant_uris": sample.expected_uris},
|
metadata={
|
||||||
|
"relevant_uris": sample.expected_uris,
|
||||||
|
"source_type": sample.source_type,
|
||||||
|
},
|
||||||
)
|
)
|
||||||
cases.append(case)
|
cases.append(case)
|
||||||
progress.advance(task)
|
progress.advance(task)
|
||||||
|
|
@ -174,16 +192,22 @@ async def run_retrieval_benchmark(
|
||||||
chunks = await rag.search(query=question, limit=5)
|
chunks = await rag.search(query=question, limit=5)
|
||||||
|
|
||||||
seen = set()
|
seen = set()
|
||||||
uris = []
|
identifiers = []
|
||||||
for result in chunks:
|
for result in chunks:
|
||||||
if result.document_id is None:
|
if result.document_id is None:
|
||||||
continue
|
continue
|
||||||
doc = await rag.get_document_by_id(result.document_id)
|
doc = await rag.get_document_by_id(result.document_id)
|
||||||
if doc and doc.uri and doc.uri not in seen:
|
if doc is None:
|
||||||
uris.append(doc.uri)
|
continue
|
||||||
seen.add(doc.uri)
|
# 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"
|
eval_name = name if name is not None else f"{spec.key}_retrieval_evaluation"
|
||||||
|
|
||||||
|
|
@ -326,6 +350,7 @@ async def evaluate_dataset(
|
||||||
name: str | None,
|
name: str | None,
|
||||||
db_path: Path | None,
|
db_path: Path | None,
|
||||||
vacuum_interval: int = 100,
|
vacuum_interval: int = 100,
|
||||||
|
multimodal_only: bool = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
if not skip_db:
|
if not skip_db:
|
||||||
console.print(f"Using dataset: {spec.key}", style="bold magenta")
|
console.print(f"Using dataset: {spec.key}", style="bold magenta")
|
||||||
|
|
@ -336,7 +361,12 @@ async def evaluate_dataset(
|
||||||
if not skip_retrieval:
|
if not skip_retrieval:
|
||||||
console.print("Running retrieval benchmarks...", style="bold blue")
|
console.print("Running retrieval benchmarks...", style="bold blue")
|
||||||
await run_retrieval_benchmark(
|
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:
|
if not skip_qa:
|
||||||
|
|
@ -368,6 +398,11 @@ def run(
|
||||||
vacuum_interval: int = typer.Option(
|
vacuum_interval: int = typer.Option(
|
||||||
100, "--vacuum-interval", help="Vacuum every N documents during DB population."
|
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:
|
) -> None:
|
||||||
spec = DATASETS.get(dataset.lower())
|
spec = DATASETS.get(dataset.lower())
|
||||||
if spec is None:
|
if spec is None:
|
||||||
|
|
@ -405,6 +440,7 @@ def run(
|
||||||
name=name,
|
name=name,
|
||||||
db_path=db,
|
db_path=db,
|
||||||
vacuum_interval=vacuum_interval,
|
vacuum_interval=vacuum_interval,
|
||||||
|
multimodal_only=multimodal_only,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -11,10 +11,11 @@ from pydantic_evals.evaluators import Evaluator
|
||||||
@dataclass
|
@dataclass
|
||||||
class DocumentPayload:
|
class DocumentPayload:
|
||||||
uri: str
|
uri: str
|
||||||
content: str
|
content: str | None = None
|
||||||
title: str | None = None
|
title: str | None = None
|
||||||
metadata: dict[str, Any] | None = None
|
metadata: dict[str, Any] | None = None
|
||||||
format: str = "md"
|
format: str = "md"
|
||||||
|
source_path: Path | None = None
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
|
|
@ -22,6 +23,7 @@ class RetrievalSample:
|
||||||
question: str
|
question: str
|
||||||
expected_uris: tuple[str, ...]
|
expected_uris: tuple[str, ...]
|
||||||
skip: bool = False
|
skip: bool = False
|
||||||
|
source_type: str | None = None
|
||||||
|
|
||||||
|
|
||||||
DocumentLoader = Callable[[], Dataset]
|
DocumentLoader = Callable[[], Dataset]
|
||||||
|
|
|
||||||
|
|
@ -1,11 +1,13 @@
|
||||||
from evaluations.config import DatasetSpec
|
from evaluations.config import DatasetSpec
|
||||||
|
|
||||||
from .hotpotqa import HOTPOTQA_SPEC
|
from .hotpotqa import HOTPOTQA_SPEC
|
||||||
|
from .open_rag_bench import OPEN_RAG_BENCH_SPEC
|
||||||
from .repliqa import REPLIQ_SPEC
|
from .repliqa import REPLIQ_SPEC
|
||||||
from .wix import WIX_SPEC
|
from .wix import WIX_SPEC
|
||||||
|
|
||||||
DATASETS: dict[str, DatasetSpec] = {
|
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"]
|
__all__ = ["DATASETS"]
|
||||||
|
|
|
||||||
229
evaluations/evaluations/datasets/open_rag_bench.py
Normal file
229
evaluations/evaluations/datasets/open_rag_bench.py
Normal 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, dict[str, Any]]:
|
||||||
|
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, dict[str, Any]] | 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_data = _answers.get(query_id, {})
|
||||||
|
records.append(
|
||||||
|
{
|
||||||
|
"query_id": query_id,
|
||||||
|
"query": query_data["query"],
|
||||||
|
"type": query_data["type"],
|
||||||
|
"source": query_data["source"],
|
||||||
|
"answer": answer_data.get("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(),
|
||||||
|
)
|
||||||
Loading…
Reference in a new issue