(self, sample_size=None, type="train")
| 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"] |
no test coverage detected