mirror of
https://github.com/wassname/Open-Assistant.git
synced 2026-09-12 12:03:06 +08:00
[feature] added GSM8k and code refactoring
This commit is contained in:
@@ -4,8 +4,8 @@ from custom_datasets.summarization import SummarizationDataset
|
||||
from sklearn.model_selection import train_test_split
|
||||
from torch.utils.data import Subset
|
||||
|
||||
QA_DATASETS = ["squad_v2", "adversarial_qa", "trivia_qa_context", "trivia_qa_nocontext"]
|
||||
SUMMARIZATION_DATASETS = ["xsum", "cnn_dailymail", "samsum", "multi_news", "scitldr", "billsum", "reddit"]
|
||||
QA_DATASETS = ["squad_v2", "adversarial_qa", "trivia_qa_context", "trivia_qa_nocontext", "gsm8k"]
|
||||
SUMMARIZATION_DATASETS = ["xsum", "cnn_dailymail", "samsum", "multi_news", "scitldr", "billsum"]
|
||||
|
||||
|
||||
def train_val_dataset(dataset, val_split=0.2):
|
||||
@@ -20,11 +20,13 @@ def get_one_dataset(conf, dataset_name):
|
||||
|
||||
if dataset_name in QA_DATASETS:
|
||||
train = QADataset(dataset_name, conf.cache_dir, "train")
|
||||
eval = QADataset(dataset_name, conf.cache_dir, "validation")
|
||||
val_name = "validation" if dataset_name not in ["gsm8k"] else "test"
|
||||
eval = QADataset(dataset_name, conf.cache_dir, val_name)
|
||||
|
||||
elif dataset_name in SUMMARIZATION_DATASETS:
|
||||
train = SummarizationDataset(dataset_name, conf.cache_dir, "train")
|
||||
eval = SummarizationDataset(dataset_name, conf.cache_dir, "validation" if dataset_name != "billsum" else "test")
|
||||
val_name = "validation" if dataset_name not in ["billsum"] else "test"
|
||||
eval = SummarizationDataset(dataset_name, conf.cache_dir, val_name)
|
||||
elif dataset_name == "webgpt":
|
||||
dataset = WebGPT()
|
||||
train, eval = train_val_dataset(dataset, val_split=0.2)
|
||||
|
||||
@@ -39,6 +39,10 @@ def index_adversarial_qa(example):
|
||||
return example["title"] + ". " + example["context"] + " " + example["question"], example["answers"]["text"][0]
|
||||
|
||||
|
||||
def index_gsm8k(example):
|
||||
return example["question"], example["answer"]
|
||||
|
||||
|
||||
class QADataset(Dataset):
|
||||
def __init__(self, dataset, cache_dir, split):
|
||||
if dataset == "squad_v2":
|
||||
@@ -46,13 +50,19 @@ class QADataset(Dataset):
|
||||
self.dataset = load_dataset("squad_v2", cache_dir=cache_dir, split=split)
|
||||
elif dataset == "trivia_qa_nocontext":
|
||||
self.index_fn = index_trivia_qa_nocontext
|
||||
self.dataset = load_dataset("trivia_qa", "rc.nocontext", split=split)
|
||||
self.dataset = load_dataset("trivia_qa", "rc.nocontext", split=split, cache_dir=cache_dir)
|
||||
elif dataset == "trivia_qa_context":
|
||||
self.index_fn = index_trivia_qa_context
|
||||
self.dataset = load_dataset("trivia_qa", "rc", split=split)
|
||||
self.dataset = load_dataset("trivia_qa", "rc", split=split, cache_dir=cache_dir)
|
||||
elif dataset == "adversarial_qa":
|
||||
self.index_fn = index_adversarial_qa
|
||||
self.dataset = load_dataset("adversarial_qa", "adversarialQA", split=split)
|
||||
self.dataset = load_dataset("adversarial_qa", "adversarialQA", split=split, cache_dir=cache_dir)
|
||||
elif dataset == "gsm8k":
|
||||
self.index_fn = index_gsm8k
|
||||
self.dataset = load_dataset("gsm8k", "main", split=split, cache_dir=cache_dir)
|
||||
elif dataset == "adversarial_qa":
|
||||
self.index_fn = index_adversarial_qa
|
||||
self.dataset = load_dataset("adversarial_qa", "adversarialQA", split=split, cache_dir=cache_dir)
|
||||
else:
|
||||
raise ValueError("Unknown dataset : " + dataset)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user