30 lines
980 B
Python
30 lines
980 B
Python
"""TriviaQA (official source: mandarjoshi/trivia_qa, rc.nocontext config)."""
|
|
|
|
from ..sample import Sample
|
|
from ..registry import register_dataset
|
|
from ..spec import DatasetSpec
|
|
|
|
|
|
@register_dataset(
|
|
DatasetSpec(
|
|
name='trivia_qa',
|
|
source='mandarjoshi/trivia_qa', # official: https://huggingface.co/datasets/mandarjoshi/trivia_qa
|
|
subset='rc.nocontext',
|
|
split='validation',
|
|
task_type='qa',
|
|
tags=['knowledge', 'openqa'],
|
|
description='TriviaQA open-domain QA without context (official).',
|
|
)
|
|
)
|
|
def trivia_qa():
|
|
def to_sample(record: dict) -> Sample:
|
|
answer = record['answer'] # {'value': ..., 'aliases': [...], ...}
|
|
targets = [answer['value']] + list(answer.get('aliases') or [])
|
|
return Sample(
|
|
input=record['question'],
|
|
target=targets, # multi-target: any alias counts
|
|
metadata={'question_id': record.get('question_id')},
|
|
)
|
|
|
|
return to_sample
|