Load the given dataset by name. Supported by default are 'shp', 'hh', and 'se'.
(name: str, split: str, silent: bool = False, cache_dir: str = None)
| 161 | |
| 162 | |
| 163 | def get_dataset(name: str, split: str, silent: bool = False, cache_dir: str = None): |
| 164 | """Load the given dataset by name. Supported by default are 'shp', 'hh', and 'se'.""" |
| 165 | if name == 'shp': |
| 166 | data = get_shp(split, silent=silent, cache_dir=cache_dir) |
| 167 | elif name == 'hh': |
| 168 | data = get_hh(split, silent=silent, cache_dir=cache_dir) |
| 169 | elif name == 'se': |
| 170 | data = get_se(split, silent=silent, cache_dir=cache_dir) |
| 171 | else: |
| 172 | raise ValueError(f"Unknown dataset '{name}'") |
| 173 | |
| 174 | assert set(list(data.values())[0].keys()) == {'responses', 'pairs', 'sft_target'}, \ |
| 175 | f"Unexpected keys in dataset: {list(list(data.values())[0].keys())}" |
| 176 | |
| 177 | return data |
| 178 | |
| 179 | |
| 180 | def get_collate_fn(tokenizer) -> Callable[[List[Dict]], Dict[str, Union[List, torch.Tensor]]]: |
no test coverage detected