| 35 | return self.tokenizer.decode(t) |
| 36 | |
| 37 | class BaseDataset(Dataset): |
| 38 | def __init__(self, tokenizer=None, max_len=2048, test=False, category="", dedup=False, seed=None): |
| 39 | super().__init__() |
| 40 | self.data = None |
| 41 | self.inputs = None |
| 42 | |
| 43 | if tokenizer is not None: |
| 44 | self.tokenizer = Tokenizer(tokenizer) |
| 45 | if seed is not None: |
| 46 | random.seed(seed) |
| 47 | |
| 48 | self.test = test |
| 49 | self.max_len = max_len |
| 50 | self.category = category |
| 51 | self.dedup = dedup |
| 52 | |
| 53 | def __len__(self): |
| 54 | return len(self.data) |
| 55 | |
| 56 | def get_inputs(self): |
| 57 | inputs = [] |
| 58 | for i in tqdm(range(len(self.data))): |
| 59 | inputs.append(self.pre(i)) |
| 60 | self.inputs = inputs |
| 61 | |
| 62 | def get_all(self): |
| 63 | temp = [] |
| 64 | for i in range(len(self.data)): |
| 65 | temp.append(self.get_history(self.data.iloc[i])) |
| 66 | return temp |
| 67 | |
| 68 | def get_inputs_list(self): |
| 69 | return self.inputs |
| 70 | |
| 71 | def __getitem__(self, idx): |
| 72 | return self.inputs[idx] |
| 73 | |
| 74 | def pre(self, idx): |
| 75 | raise NotImplementedError(None) |
| 76 | |
| 77 | def get_history(self, row): |
| 78 | raise {} |
| 79 | |
| 80 | def generate_prompt(self, data_point): |
| 81 | return f"""### User Input: |
| 82 | {data_point["input"]} |
| 83 | |
| 84 | ### Response:\n{data_point["output"]}""" |
| 85 | |
| 86 | |
| 87 | class CSVBaseDataset(BaseDataset): |
nothing calls this directly
no outgoing calls
no test coverage detected