From 9be4c921cdf9773f677beac742a995830095ae63 Mon Sep 17 00:00:00 2001 From: theblackcat102 Date: Wed, 1 Feb 2023 22:33:37 +0000 Subject: [PATCH] [feature] Add OA translated QA --- model/supervised_finetuning/custom_datasets/__init__.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/model/supervised_finetuning/custom_datasets/__init__.py b/model/supervised_finetuning/custom_datasets/__init__.py index ee061f04..f664ceaf 100644 --- a/model/supervised_finetuning/custom_datasets/__init__.py +++ b/model/supervised_finetuning/custom_datasets/__init__.py @@ -2,7 +2,7 @@ High level functions for model training """ from custom_datasets.prompt_dialogue import InstructionTuning, PromptGeneratedDataset -from custom_datasets.qa_datasets import SODA, JokeExplaination, QADataset, SODADialogue, WebGPT +from custom_datasets.qa_datasets import SODA, JokeExplaination, QADataset, SODADialogue, TranslatedQA, WebGPT from custom_datasets.summarization import SummarizationDataset from custom_datasets.toxic_conversation import ProsocialDialogue, ProsocialDialogueExplaination from custom_datasets.translation import WMT2019, DiveMT, TEDTalk @@ -92,6 +92,9 @@ def get_one_dataset(conf, dataset_name): elif dataset_name == "instruct_tuning": dataset = InstructionTuning(conf.cache_dir) train, eval = train_val_dataset(dataset, val_split=0.2) + elif dataset_name == "translate_qa": + dataset = TranslatedQA(conf.cache_dir) + train, eval = train_val_dataset(dataset, val_split=0.01) else: raise ValueError(f"Unknown dataset {dataset_name}")