(examples, sequence_parallel_size, sequence_parallel_mode="ulysses")
| 220 | |
| 221 | # sp for Sequence Parallel |
| 222 | def sp_split(examples, sequence_parallel_size, sequence_parallel_mode="ulysses"): |
| 223 | for k, v in examples.items(): |
| 224 | chunks = list() |
| 225 | for row in v: |
| 226 | if k.endswith("attention_mask"): |
| 227 | chunks.extend([row] * sequence_parallel_size) |
| 228 | elif row is None: |
| 229 | chunks.extend([None] * sequence_parallel_size) |
| 230 | else: |
| 231 | chunks.extend( |
| 232 | preprocess_sp_dataset(row, sequence_parallel_size, sequence_parallel_mode) |
| 233 | ) |
| 234 | examples[k] = chunks |
| 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": |
nothing calls this directly
no test coverage detected