MCPcopy Create free account
hub / github.com/eric-mitchell/direct-preference-optimization / get_dataset

Function get_dataset

preference_datasets.py:163–177  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

161
162
163def 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
180def get_collate_fn(tokenizer) -> Callable[[List[Dict]], Dict[str, Union[List, torch.Tensor]]]:

Callers 1

get_batch_iteratorFunction · 0.85

Calls 3

get_shpFunction · 0.85
get_hhFunction · 0.85
get_seFunction · 0.85

Tested by

no test coverage detected