| 3 | |
| 4 | |
| 5 | class DataLoader: |
| 6 | def __init__(self, data="hotpot_qa", seed=2023): |
| 7 | self.data = data |
| 8 | self.seed = seed |
| 9 | |
| 10 | def load(self, sample_size=None, type="train"): |
| 11 | if self.data == "hotpot_qa": |
| 12 | return self.load_hotpot_qa(sample_size=sample_size, type=type) |
| 13 | elif self.data == "fever": |
| 14 | return self.load_fever(sample_size=sample_size, type=type) |
| 15 | elif self.data == "trivia_qa": |
| 16 | return self.load_trivia_qa(sample_size=sample_size, type=type) |
| 17 | elif self.data == "gsm8k": |
| 18 | return self.load_gsm8k(sample_size=sample_size, type=type) |
| 19 | elif self.data == "physics_question": |
| 20 | return self.load_physics_question(sample_size=sample_size) |
| 21 | elif self.data == "disfl_qa": |
| 22 | return self.load_disfl_qa(sample_size=sample_size) |
| 23 | elif self.data == "sports_understanding": |
| 24 | return self.load_sports_understanding(sample_size=sample_size) |
| 25 | elif self.data == "strategy_qa": |
| 26 | return self.load_strategy_qa(sample_size=sample_size) |
| 27 | elif self.data == "sotu_qa": |
| 28 | return self.load_sotu_qa(sample_size=sample_size) |
| 29 | else: |
| 30 | raise ValueError("Data not supported.") |
| 31 | |
| 32 | def load_hotpot_qa(self, cache_dir="data/hotpot_qa", sample_size=100, type="test"): |
| 33 | assert type in ["train", "validation", "test"] |
| 34 | data = datasets.load_dataset('hotpot_qa', 'fullwiki', cache_dir=cache_dir) |
| 35 | df = data[type].to_pandas() |
| 36 | sampled_df = df.sample(sample_size, random_state=self.seed)[["question", "answer"]].reset_index(drop=True) |
| 37 | return sampled_df |
| 38 | |
| 39 | def load_fever(self, cache_dir="data/fever", sample_size=100, type="test"): |
| 40 | assert type in ["train", "validation", "test"] |
| 41 | data = datasets.load_dataset('copenlu/fever_gold_evidence', cache_dir=cache_dir) |
| 42 | df = data[type].to_pandas() |
| 43 | sampled_df = df.sample(sample_size, random_state=self.seed)[["claim", "label"]].reset_index(drop=True) |
| 44 | return sampled_df |
| 45 | |
| 46 | def load_trivia_qa(self, cache_dir="data/trivia_qa", sample_size=100, type="test"): |
| 47 | assert type in ["train", "validation", "test"] |
| 48 | data = datasets.load_dataset('trivia_qa', 'rc.nocontext', cache_dir=cache_dir) |
| 49 | df = data[type].to_pandas() |
| 50 | sampled_df = df.sample(sample_size, random_state=self.seed)[["question", "answer"]].reset_index(drop=True) |
| 51 | return sampled_df |
| 52 | |
| 53 | def load_gsm8k(self, cache_dir="data/gsm8k", sample_size=100, type="test"): |
| 54 | assert type in ["train", "validation", "test"] |
| 55 | data = datasets.load_dataset('gsm8k', name="main", cache_dir=cache_dir) |
| 56 | df = data[type].to_pandas() |
| 57 | sampled_df = df.sample(sample_size, random_state=self.seed)[["question", "answer"]].reset_index(drop=True) |
| 58 | return sampled_df |
| 59 | |
| 60 | def load_physics_question(self, cache_dir="data/bigbench/physics_question.csv", sample_size=None): |
| 61 | df = pd.read_csv(cache_dir) |
| 62 | if sample_size is not None: |