"""DROP (Discrete Reasoning Over Paragraphs): https://huggingface.co/datasets/EleutherAI/drop
A passage and a question over it. ``DropCompletion_OLMES`` generates a free-form answer scored by DROP F1 /
exact match (EleutherAI/drop); ``DropMC_OLMES`` scores labelled candidate answers by loglikelihood
(allenai/drop-gen2mc).
"""
from typing import Any, final, override
from eval_framework.answer import ExtractFromCompletion
from eval_framework.choices import ChoiceFields, ChoiceReader
from eval_framework.composed import ComposedBenchmark
from eval_framework.contract import Benchmark
from eval_framework.eval_kind import Generative
from eval_framework.fewshot import FewShot, FewshotExample, FewShotSplit, FunctionRenderer
from eval_framework.metrics.completion.drop_completion import DropF1ExactMatch, DropMetricContext
from eval_framework.tasks.base import Language
from eval_framework.tasks.dataset_loading import DatasetPolicy
from eval_framework.tasks.dataset_revisions import pinned_by_framework
from eval_framework.tasks.task_style import MCStyle
from eval_framework.tasks.utils import get_n_letters
DROP_COMPLETION_DATASET_PATH = "EleutherAI/drop"
DROP_CHOICE_DATASET_PATH = "allenai/drop-gen2mc"
# Generation stops at the next passage/question or a blank line (the boundary between few-shot blocks).
_COMPLETION_STOP_SEQUENCES = ["Passage:", "Question:", "\n\n"]
# Prepended once above the DropCompletion_OLMES prompt (the OLMES reading-comprehension preamble).
_OLMES_PREAMBLE = (
"The following are reading comprehension questions, where the answer to each question is either a "
"segment of text from the corresponding passage, a number, or a date (containing any of the date, "
"month, and/or year components). Some questions may require you to pull together information pieces "
"from the passage and reason over them."
)
def _flatten_validated_answers(validated_answers: dict[str, Any]) -> list[dict[str, Any]]:
"""Flatten validated_answers from a dict of lists to a list of dicts."""
num_list = validated_answers.get("number") or []
date_list = validated_answers.get("date") or []
spans_list = validated_answers.get("spans") or []
n = max(len(num_list), len(date_list), len(spans_list))
return [
{
"number": num_list[i] if i < len(num_list) else "",
"date": date_list[i] if i < len(date_list) else {"day": "", "month": "", "year": ""},
"spans": spans_list[i] if i < len(spans_list) else [],
}
for i in range(n)
]
def _parse_answer(answer: dict[str, Any]) -> tuple[str, ...]:
"""Return a hashable tuple for one answer (number, spans, or date string)."""
if answer.get("number") not in (None, ""):
return (str(answer["number"]),)
spans = answer.get("spans") or []
if spans:
return tuple(spans)
date = answer.get("date") or {}
day = date.get("day") or ""
month = date.get("month") or ""
year = date.get("year") or ""
return (" ".join([day, month, year]).strip(),)
def _get_answers(doc: dict[str, Any]) -> list[tuple[str, ...]]:
"""Deduplicated list of valid answer tuples (main answer + validated_answers)."""
answer = doc.get("answer") or {}
validated = doc.get("validated_answers") or {}
candidates = [answer] + _flatten_validated_answers(validated)
seen: set[tuple[str, ...]] = set()
out = []
for cand in candidates:
if not cand:
continue
parsed = _parse_answer(cand)
if parsed in seen or (len(parsed) == 1 and parsed[0] == ""):
continue
seen.add(parsed)
out.append(parsed)
return out
def _tuple_to_display(tup: tuple[str, ...]) -> str:
"""The single answer string shown to (and generated by) the model."""
return ", ".join(str(x) for x in tup) if tup else ""
def _completion_prompt(item: dict[str, Any]) -> str:
passage = (item.get("passage") or "").strip()
question = item.get("question", "")
return f"Passage: {passage}\nQuestion: {question}\n"
def _completion_ground_truth(item: dict[str, Any]) -> str:
return f" {_tuple_to_display(_get_answers(item)[0])}"
def _completion_context(item: dict[str, Any]) -> DropMetricContext:
# DROP F1 scores the generation against every valid gold answer (each a tuple of spans).
return DropMetricContext(answer_tuples=[list(a) for a in _get_answers(item)])
def _completion_demo(demo: dict[str, Any]) -> FewshotExample:
return FewshotExample(prompt=_completion_prompt(demo), answer=f"Answer:{_completion_ground_truth(demo)}")
[docs]
def drop_completion_olmes(dataset: DatasetPolicy | None = None) -> Benchmark:
"""OLMES: a reading-comprehension preamble, few-shot from the train split, and a 100-token answer budget."""
if dataset is None:
# Only rows whose answer parses to at least one gold tuple are scorable.
dataset = pinned_by_framework(DROP_COMPLETION_DATASET_PATH).subset(
lambda row: bool(_get_answers(row)), "questions with a parseable gold answer"
)
kind = Generative(
build_prompt=_completion_prompt,
cue="Answer:", # the model continues after the cue
ground_truth=_completion_ground_truth,
metrics=[DropF1ExactMatch],
context=_completion_context,
initial_prompt=_OLMES_PREAMBLE,
)
# F1 scores the whole generation; nothing is extracted
answer = ExtractFromCompletion(lambda completion_text: completion_text, _COMPLETION_STOP_SEQUENCES, max_tokens=100)
return ComposedBenchmark.compose(
id="DropCompletion_OLMES",
kind=kind,
answer=answer,
sample_split="validation",
fewshot=FewShot(FewShotSplit("train"), FunctionRenderer(_completion_demo)),
dataset_policy=dataset,
language=Language.ENG,
)
@final
class _DropChoiceReader(ChoiceReader):
"""Reads the gen2mc passage/question/choices; the correct index is the position of ``answerKey``."""
@override
def read(self, item: dict[str, Any]) -> ChoiceFields:
passage = (item.get("passage_original") or "").strip()
question = item.get("question_original", "")
choices = item.get("choices", {})
texts = choices.get("text", [])
labels = choices.get("label") or get_n_letters(len(texts))
return ChoiceFields(
raw_question=f"Passage: {passage}\nQuestion: {question}",
choices=texts,
correct_index=labels.index(item["answerKey"]),
)
[docs]
def drop_mc_olmes(dataset: DatasetPolicy | None = None) -> Benchmark:
"""OLMES lays out the options with a leading space (" A. ...")."""
dataset_policy = dataset if dataset is not None else pinned_by_framework(DROP_CHOICE_DATASET_PATH)
return ComposedBenchmark.choice(
id="DropMC_OLMES",
reader=_DropChoiceReader(),
styler=MCStyle(question_prefix="", space_prefixed_labels=True),
sample_split="validation",
fewshot_split="validation",
dataset_policy=dataset_policy,
language=Language.ENG,
)
DROP_BENCHMARKS: list[Benchmark] = [drop_completion_olmes(), drop_mc_olmes()]