(dataset, num_works, tokenizer=None, cutoff_len=10000,
sequence_parallel_size=1, sequence_parallel_mode="ulysses",
cache_dataset_overwrite=False)
| 246 | return preprocess_func |
| 247 | |
| 248 | def _get_sequence_parallel_dataset(dataset, num_works, tokenizer=None, cutoff_len=10000, |
| 249 | sequence_parallel_size=1, sequence_parallel_mode="ulysses", |
| 250 | cache_dataset_overwrite=False) -> Optional[Union["Dataset", "IterableDataset"]]: |
| 251 | kwargs = dict( |
| 252 | num_proc=num_works, |
| 253 | load_from_cache_file=not cache_dataset_overwrite, |
| 254 | desc="Running padding split on dataset", |
| 255 | ) |
| 256 | pad_sequence_func = get_sequence_parallel_preprocess( |
| 257 | stage="pad", |
| 258 | tokenizer=tokenizer, |
| 259 | cutoff_len=cutoff_len |
| 260 | ) |
| 261 | padded_dataset = dataset.map( |
| 262 | pad_sequence_func, batched=True, batch_size=num_works, **kwargs |
| 263 | ) |
| 264 | kwargs = dict( |
| 265 | num_proc=num_works, |
| 266 | load_from_cache_file=not cache_dataset_overwrite, |
| 267 | desc="Running sequence parallel split on dataset", |
| 268 | ) |
| 269 | sp_dataset_func = get_sequence_parallel_preprocess( |
| 270 | stage="split", |
| 271 | tokenizer=tokenizer, |
| 272 | sequence_parallel_size=sequence_parallel_size, |
| 273 | sequence_parallel_mode=sequence_parallel_mode, |
| 274 | ) |
| 275 | sp_dataset = padded_dataset.map( |
| 276 | sp_dataset_func, batched=True, batch_size=num_works, **kwargs |
| 277 | ) |
| 278 | return sp_dataset |
| 279 | |
| 280 | def packing_dataset(dataset, tokenizer, cutoff_len, num_worker, cache_dataset_overwrite): |
| 281 | preprocess_func = partial(preprocess_packed_supervised_dataset, tokenizer=tokenizer, cutoff_len=cutoff_len) |
no test coverage detected