Source code for eval_framework.benchmarks.drop

"""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()]