(stage, tokenizer, cutoff_len=None, sequence_parallel_size=1, sequence_parallel_mode="ulysses")
| 235 | return examples |
| 236 | |
| 237 | def get_sequence_parallel_preprocess(stage, tokenizer, cutoff_len=None, sequence_parallel_size=1, sequence_parallel_mode="ulysses"): |
| 238 | if stage == "pad": |
| 239 | assert cutoff_len is not None |
| 240 | preprocess_func = partial(pad_sequence, cutoff_len=cutoff_len, tokenizer=tokenizer) |
| 241 | elif stage == "split": |
| 242 | preprocess_func = partial(sp_split, sequence_parallel_size=sequence_parallel_size, sequence_parallel_mode=sequence_parallel_mode) |
| 243 | else: |
| 244 | raise NotImplementedError(f"Unexpected stage in sequence_parallel_preprocess: {stage}") |
| 245 | |
| 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", |
no outgoing calls
no test coverage detected