| 54 | return pytorch_worker_seed(group=group) + seed * 23 |
| 55 | |
| 56 | class ConfiguredResampledShards(ResampledShards): |
| 57 | def __init__(self, urls, seed, nshards=sys.maxsize, deterministic=True): |
| 58 | from sat.helpers import print_rank0 |
| 59 | try: |
| 60 | from megatron.core.parallel_state import get_data_parallel_group |
| 61 | group = get_data_parallel_group() |
| 62 | print_rank0("Using megatron data parallel group.") |
| 63 | except: |
| 64 | from sat.mpu import get_data_parallel_group |
| 65 | try: |
| 66 | group = get_data_parallel_group() |
| 67 | print_rank0("Using sat data parallel group.") |
| 68 | except AssertionError: |
| 69 | group = None |
| 70 | print_rank0("No data parallel group is specified!") |
| 71 | worker_seed_sat_this = partial(worker_seed_sat, group=group, seed=seed) |
| 72 | super().__init__(urls, nshards, worker_seed_sat_this, deterministic) |
| 73 | |
| 74 | class SimpleDistributedWebDataset(DataPipeline): |
| 75 | def __init__(self, path, process_fn, seed, *, shuffle_buffer=1000): |